Files
zhixing-system/zhixing-server/tests/unit/sector_radar/test_build.py
T

642 lines
23 KiB
Python

import logging
from collections.abc import Sequence
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.application.build import (
BuildSectorRadar,
BuildSectorRadarCommand,
)
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
PublicationStatus,
RadarPublication,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.persistence import PublicationSourceGroup
from zhixing_server.modules.sector_radar.domain.source import (
CapabilityProbeResult,
DailyRow,
MoneyflowDcRow,
SectorIndexRow,
SectorMemberRow,
SourceContractError,
SourceResult,
SourceSnapshot,
StockBasicRow,
SuspendRow,
TradeCalendarRow,
build_source_snapshot,
)
from zhixing_server.modules.sector_radar.infrastructure.memory import (
InMemorySectorRadarRepository,
)
TARGET_DATE = date(2026, 8, 28)
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
class FakeRadarSource:
def __init__(
self,
*,
missing_moneyflow: bool = False,
missing_membership: bool = False,
net_scale: Decimal = Decimal(1),
) -> None:
self.missing_moneyflow = missing_moneyflow
self.missing_membership = missing_membership
self.net_scale = net_scale
self.fail_daily = False
self.calls: list[str] = []
self.moneyflow_candidate_codes: list[tuple[str, ...]] = []
def _result[T](
self, api_name: str, target: date | None, rows: tuple[T, ...]
) -> SourceResult[T]:
snapshot = build_source_snapshot(
api_name=api_name,
params={
"trade_date": target.isoformat() if target is not None else "all",
"fixture_fingerprint": repr(rows),
},
rows=tuple(self._raw_row(row) for row in rows),
target_trade_date=target,
partition_key="all" if api_name == "dc_member" else None,
observed_at=NOW,
)
return SourceResult((snapshot,), rows)
@staticmethod
def _raw_row(row: object) -> dict[str, object]:
if isinstance(row, TradeCalendarRow):
return {
"exchange": row.exchange,
"cal_date": row.cal_date,
"is_open": int(row.is_open),
"pretrade_date": row.pretrade_date,
}
if isinstance(row, SectorIndexRow):
return {
"trade_date": row.trade_date,
"ts_code": row.sector_code,
"name": row.name,
"level": row.level,
"pct_change": row.pct_change,
"leading_code": row.leading_code,
}
if isinstance(row, SectorMemberRow):
return {
"trade_date": row.trade_date,
"ts_code": row.sector_code,
"con_code": row.stock_code,
"name": row.stock_name,
}
if isinstance(row, StockBasicRow):
return {
"ts_code": row.ts_code,
"symbol": row.symbol,
"name": row.name,
"market": row.market,
"exchange": row.exchange,
"list_status": row.list_status,
"list_date": row.list_date,
"delist_date": row.delist_date,
}
if isinstance(row, SuspendRow):
return {
"ts_code": row.ts_code,
"trade_date": row.trade_date,
"suspend_timing": row.suspend_timing,
"suspend_type": row.suspend_type,
}
if isinstance(row, DailyRow):
return {
"ts_code": row.ts_code,
"trade_date": row.trade_date,
"close": row.close,
"pre_close": row.pre_close,
"pct_chg": row.pct_chg,
"vol": row.volume,
"amount": row.amount_thousand_yuan,
}
if isinstance(row, MoneyflowDcRow):
return {
"trade_date": row.trade_date,
"ts_code": row.ts_code,
"name": row.name,
"net_amount": row.net_amount_ten_thousand_yuan,
"net_amount_rate": row.net_amount_rate,
"pct_change": row.pct_change,
"close": row.close,
}
raise TypeError(f"unsupported fake source row: {type(row).__name__}")
def fetch_trade_calendar(self, start: date, end: date) -> SourceResult[TradeCalendarRow]:
self.calls.append("calendar")
rows = tuple(
TradeCalendarRow("SSE", start + timedelta(days=offset), True, None)
for offset in range((end - start).days + 1)
)
return self._result("trade_cal", end, rows)
def fetch_sector_indices(
self, trade_date: date, sector_type: SectorType
) -> SourceResult[SectorIndexRow]:
self.calls.append(f"{sector_type.value}_indices")
prefix = "BK0" if sector_type is SectorType.CONCEPT else "BK1"
row = SectorIndexRow(
trade_date,
sector_type,
f"{prefix}001.DC",
"示例概念" if sector_type is SectorType.CONCEPT else "示例行业",
"一级",
Decimal(1),
"000001.SZ",
)
return self._result(f"dc_index_{sector_type.value}", trade_date, (row,))
def fetch_sector_members(
self, trade_date: date, sector_codes: Sequence[str]
) -> SourceResult[SectorMemberRow]:
self.calls.append("members")
if self.missing_membership:
snapshots: list[SourceSnapshot] = []
member_rows: list[SectorMemberRow] = []
for sector_code in sector_codes:
sector_rows = (
()
if sector_code == sector_codes[-1]
else tuple(
SectorMemberRow(
trade_date,
sector_code,
f"00000{index}.SZ",
f"股票{index}",
)
for index in range(1, 6)
)
)
snapshots.append(
build_source_snapshot(
api_name="dc_member",
params={
"trade_date": trade_date.isoformat(),
"ts_code": sector_code,
},
rows=tuple(self._raw_row(row) for row in sector_rows),
target_trade_date=trade_date,
partition_key=sector_code,
observed_at=NOW,
)
)
member_rows.extend(sector_rows)
return SourceResult(tuple(snapshots), tuple(member_rows))
rows = tuple(
SectorMemberRow(trade_date, sector_code, f"00000{index}.SZ", f"股票{index}")
for sector_code in sector_codes
for index in range(1, 6)
)
return self._result("dc_member", trade_date, rows)
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
self.calls.append("stock_basics")
rows = tuple(
StockBasicRow(
f"00000{index}.SZ",
f"00000{index}",
f"股票{index}",
"主板",
"SZSE",
"L",
date(2020, 1, 1),
None,
)
for index in range(1, 6)
)
return self._result("stock_basic", None, rows)
def fetch_suspensions(self, trade_date: date) -> SourceResult[SuspendRow]:
self.calls.append("suspensions")
return self._result("suspend_d", trade_date, ())
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
self.calls.append("daily")
if self.fail_daily:
raise RuntimeError("private provider detail")
rows = tuple(
DailyRow(
f"00000{index}.SZ",
trade_date,
Decimal(10),
Decimal(10),
Decimal(0),
Decimal(100),
Decimal(1000),
)
for index in range(1, 6)
)
return self._result("daily", trade_date, rows)
def fetch_moneyflow_dc(
self,
trade_date: date,
candidate_codes: Sequence[str],
) -> SourceResult[MoneyflowDcRow]:
self.calls.append("moneyflow_dc")
self.moneyflow_candidate_codes.append(tuple(candidate_codes))
count = 4 if self.missing_moneyflow else 5
rows = tuple(
MoneyflowDcRow(
trade_date,
f"00000{index}.SZ",
f"股票{index}",
Decimal(index) * self.net_scale,
Decimal(0),
Decimal(0),
Decimal(10),
)
for index in range(1, count + 1)
)
return self._result("moneyflow_dc", trade_date, rows)
def probe(self, trade_date: date) -> CapabilityProbeResult:
return CapabilityProbeResult(NOW, ())
def test_successful_build_is_idempotent_and_failed_retry_preserves_last_good() -> None:
source = FakeRadarSource()
repository = InMemorySectorRadarRepository()
use_case = BuildSectorRadar(
source,
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
)
first = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
repeated = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
source.fail_daily = True
failed = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
assert first.status == "success"
assert first.exit_code == 0
assert first.outcomes[0].ranking_count == 6
assert repeated.status == "unchanged"
assert repeated.outcomes[0].publication_id == first.outcomes[0].publication_id
assert failed.status == "failed"
assert "private provider detail" not in str(failed.as_dict())
last_good = repository.get_last_good_publication()
assert last_good is not None
assert last_good.publication_id == first.outcomes[0].publication_id
assert any(item.status is PublicationStatus.FAILED for item in repository.publications.values())
def test_build_passes_stable_current_listing_member_intersection_to_moneyflow() -> None:
class FutureListingSource(FakeRadarSource):
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
result = super().fetch_stock_basics()
rows = result.rows[:-1] + (replace(result.rows[-1], list_date=date(2027, 1, 1)),)
return self._result("stock_basic", None, rows)
source = FutureListingSource()
summary = BuildSectorRadar(
source,
InMemorySectorRadarRepository(),
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
assert summary.status == "success"
assert source.moneyflow_candidate_codes == [
("000001.SZ", "000002.SZ", "000003.SZ", "000004.SZ")
]
def test_source_contract_failure_is_logged_with_safe_build_context(
caplog: pytest.LogCaptureFixture,
) -> None:
class InvalidDailySource(FakeRadarSource):
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
self.calls.append("daily")
raise SourceContractError("daily returned duplicate business keys")
repository = InMemorySectorRadarRepository()
caplog.set_level(
logging.ERROR,
logger="zhixing_server.modules.sector_radar.application.build",
)
failed = BuildSectorRadar(
InvalidDailySource(),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
messages = "\n".join(record.getMessage() for record in caplog.records)
assert failed.status == "failed"
assert failed.outcomes[0].error_message == "input or source contract validation failed"
assert "sector_radar_source_group_contract_failed" in messages
assert "source_group=daily" in messages
assert "publication_id=radar-20260828-running-" in messages
assert "validation=daily returned duplicate business keys" in messages
assert len(caplog.records) == 1
def test_partial_coverage_and_lock_have_distinct_exit_codes() -> None:
repository = InMemorySectorRadarRepository()
partial = BuildSectorRadar(
FakeRadarSource(missing_moneyflow=True),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
publication_count = len(repository.publications)
partial_repeated = BuildSectorRadar(
FakeRadarSource(missing_moneyflow=True),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
repository.lock_available = False
locked = BuildSectorRadar(
FakeRadarSource(),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
assert partial.status == "partial"
assert partial.exit_code == 2
assert partial.outcomes[0].coverage == Decimal("0.8")
assert partial_repeated.status == "partial"
assert len(repository.publications) == publication_count
assert repository.get_last_good_publication() is None
assert {
record.source_group
for record in repository.load_publication_sources(partial.outcomes[0].publication_id or "")
if record.refresh_on_retry
} == {PublicationSourceGroup.MONEYFLOW_DC}
assert locked.status == "failed"
assert locked.exit_code == 1
assert locked.outcomes[0].status == "locked"
def test_unknown_membership_is_persisted_as_partial_and_retried_independently() -> None:
repository = InMemorySectorRadarRepository()
source = FakeRadarSource(missing_membership=True)
use_case = BuildSectorRadar(source, repository, today=TARGET_DATE, now_fn=lambda: NOW)
partial = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
partial_id = partial.outcomes[0].publication_id
assert partial_id is not None
publication = repository.get_publication(partial_id)
assert partial.status == "partial"
assert publication is not None
assert publication.error_summary == "membership_unknown"
assert any(item.status is MembershipStatus.UNKNOWN for item in repository.memberships.values())
assert {
record.source_group
for record in repository.load_publication_sources(partial_id)
if record.refresh_on_retry
} == {PublicationSourceGroup.MEMBERS}
source.missing_membership = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
assert retried.status == "success"
assert source.calls == ["members"]
def test_membership_retry_refreshes_moneyflow_when_replay_misses_new_candidates() -> None:
class ExpandingMembershipSource(FakeRadarSource):
def fetch_sector_members(
self,
trade_date: date,
sector_codes: Sequence[str],
) -> SourceResult[SectorMemberRow]:
self.calls.append("members")
rows: list[SectorMemberRow] = []
snapshots: list[SourceSnapshot] = []
for index, sector_code in enumerate(sector_codes, start=1):
sector_rows = (
()
if self.missing_membership and index == len(sector_codes)
else (
SectorMemberRow(
trade_date,
sector_code,
f"00000{index}.SZ",
f"股票{index}",
),
)
)
snapshots.append(
build_source_snapshot(
api_name="dc_member",
params={
"trade_date": trade_date.isoformat(),
"ts_code": sector_code,
},
rows=tuple(self._raw_row(row) for row in sector_rows),
target_trade_date=trade_date,
partition_key=sector_code,
observed_at=NOW,
)
)
rows.extend(sector_rows)
return SourceResult(tuple(snapshots), tuple(rows))
def fetch_moneyflow_dc(
self,
trade_date: date,
candidate_codes: Sequence[str],
) -> SourceResult[MoneyflowDcRow]:
self.calls.append("moneyflow_dc")
self.moneyflow_candidate_codes.append(tuple(candidate_codes))
rows = tuple(
MoneyflowDcRow(
trade_date,
code,
code,
Decimal(1),
Decimal(0),
Decimal(0),
Decimal(10),
)
for code in candidate_codes
)
return self._result("moneyflow_dc", trade_date, rows)
repository = InMemorySectorRadarRepository()
source = ExpandingMembershipSource(missing_membership=True)
use_case = BuildSectorRadar(source, repository, today=TARGET_DATE, now_fn=lambda: NOW)
partial = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
partial_id = partial.outcomes[0].publication_id
assert partial.status == "partial"
assert partial_id is not None
assert source.moneyflow_candidate_codes == [("000001.SZ",)]
source.missing_membership = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
assert retried.status == "success"
assert source.calls == ["members", "moneyflow_dc"]
assert source.moneyflow_candidate_codes[-1] == ("000001.SZ", "000002.SZ")
def test_range_builds_dates_in_order_and_retry_uses_old_target() -> None:
repository = InMemorySectorRadarRepository()
source = FakeRadarSource(missing_moneyflow=True)
use_case = BuildSectorRadar(source, repository, now_fn=lambda: NOW)
end = TARGET_DATE + timedelta(days=1)
summary = use_case.execute(BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end))
partial_id = summary.outcomes[0].publication_id
assert partial_id is not None
source.missing_moneyflow = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
assert [item.target_trade_date for item in summary.outcomes] == [TARGET_DATE, end]
assert all(item.status == "partial" for item in summary.outcomes)
assert retried.outcomes[0].target_trade_date == TARGET_DATE
assert retried.status == "success"
assert source.calls == ["moneyflow_dc"]
def test_failed_retry_reuses_every_completed_source_group() -> None:
repository = InMemorySectorRadarRepository()
source = FakeRadarSource()
source.fail_daily = True
use_case = BuildSectorRadar(source, repository, now_fn=lambda: NOW)
failed = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
failed_id = failed.outcomes[0].publication_id
assert failed_id is not None
source.fail_daily = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=failed_id))
assert retried.status == "success"
assert source.calls == ["daily", "moneyflow_dc"]
old_groups = {record.source_group for record in repository.load_publication_sources(failed_id)}
assert len(old_groups) == 6
def test_date_lock_recovers_an_orphaned_running_publication() -> None:
repository = InMemorySectorRadarRepository()
stale = RadarPublication(
publication_id="stale-running",
target_trade_date=TARGET_DATE,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="pending",
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal(0),
started_at=NOW - timedelta(hours=1),
)
repository.create_publication(stale)
summary = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
recovered = repository.get_publication("stale-running")
assert summary.status == "success"
assert recovered is not None
assert recovered.status is PublicationStatus.FAILED
assert recovered.error_summary == "recovered_stale_running"
def test_default_target_excludes_today_before_closing_data_is_ready() -> None:
before_close = datetime(2026, 8, 28, 6, 0, tzinfo=UTC)
after_close = datetime(2026, 8, 28, 8, 0, tzinfo=UTC)
before = BuildSectorRadar(
FakeRadarSource(),
InMemorySectorRadarRepository(),
today=TARGET_DATE,
now_fn=lambda: before_close,
).execute()
after = BuildSectorRadar(
FakeRadarSource(),
InMemorySectorRadarRepository(),
today=TARGET_DATE,
now_fn=lambda: after_close,
).execute()
assert before.outcomes[0].target_trade_date == TARGET_DATE - timedelta(days=1)
assert after.outcomes[0].target_trade_date == TARGET_DATE
def test_tenth_trading_day_publishes_swing_and_five_rank_changes() -> None:
repository = InMemorySectorRadarRepository()
end = TARGET_DATE + timedelta(days=9)
summary = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end)
)
publication = repository.get_last_good_publication(end)
assert summary.status == "success"
assert publication is not None
current = tuple(
record.ranking
for record in repository.rankings.values()
if record.publication_id == publication.publication_id
)
swing = tuple(
ranking
for ranking in current
if ranking.observation.metric_version == "zhixing_swing_equal_3_10_v1"
)
assert len(swing) == 2
assert all(ranking.observation.value == Decimal("0.03") for ranking in swing)
assert all(
tuple(change.value for change in ranking.rank_changes) == (None, None, None, None, None)
for ranking in swing
)
amount = tuple(
ranking
for ranking in current
if ranking.observation.metric_version == "zhixing_amount_net_bn_v1"
)
assert all(
tuple(change.value for change in ranking.rank_changes) == (0, 0, 0, 0, 0)
for ranking in amount
)
def test_history_uses_latest_successful_input_revision_for_a_date() -> None:
repository = InMemorySectorRadarRepository()
clock = [NOW]
first_source = FakeRadarSource(net_scale=Decimal(1))
second_source = FakeRadarSource(net_scale=Decimal(2))
first = BuildSectorRadar(first_source, repository, now_fn=lambda: clock[0]).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
clock[0] = NOW + timedelta(minutes=5)
second = BuildSectorRadar(second_source, repository, now_fn=lambda: clock[0]).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
history = repository.load_daily_aggregate_history(
TARGET_DATE + timedelta(days=1), limit_dates=1
)
assert first.status == "success"
assert second.status == "success"
assert len(history) == 2
assert all(item.net_amount_yuan == Decimal(300_000) for item in history)