Files

1014 lines
41 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_weighted_v2"
)
assert len(swing) == 2
assert all(
ranking.observation.value == pytest.approx(Decimal(150000) / Decimal(5000100))
for ranking in swing
)
assert all(ranking.observation.weighted_score is not None 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)
def test_detail_history_is_thirty_sessions_past_only_and_latest_revision() -> None:
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
repository = InMemorySectorRadarRepository()
end = TARGET_DATE + timedelta(days=33)
summary = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end)
)
assert summary.status == "success"
target = end - timedelta(days=1)
replacement = BuildSectorRadar(
FakeRadarSource(net_scale=Decimal(2)),
repository,
now_fn=lambda: NOW + timedelta(hours=1),
).execute(BuildSectorRadarCommand(trade_date=target))
reader = ReadRadarDetails(repository)
history = reader.history(target, SectorType.CONCEPT, "BK0001.DC")
assert len(history.points) == 30
assert history.points[0].trade_date == target - timedelta(days=29)
assert history.points[-1].trade_date == target
assert len({point.trade_date for point in history.points}) == 30
assert history.points[-1].publication_id == replacement.outcomes[0].publication_id
assert history.points[-1].amount.metric_value == Decimal("0.003")
assert history.available_days == 30
assert (
reader.history(TARGET_DATE - timedelta(days=1), SectorType.CONCEPT, "BK0001.DC").status
== "no_data"
)
old = repository.get_successful_publication(target - timedelta(days=1))
assert old is not None
repository.publications[old.publication_id] = replace(old, source_version="incompatible-v2")
isolated = reader.history(target, SectorType.CONCEPT, "BK0001.DC")
assert isolated.points[-2].amount.missing
assert isolated.available_days == 29
def test_optional_moneyflow_failure_keeps_main_rankings_and_nullable_details() -> None:
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
from zhixing_server.modules.sector_radar.domain.source import MoneyflowRow
class UnavailableActiveSource(FakeRadarSource):
def fetch_moneyflow(self, trade_date: date) -> SourceResult[MoneyflowRow]:
raise SourceContractError("optional provider unavailable")
repository = InMemorySectorRadarRepository()
summary = BuildSectorRadar(UnavailableActiveSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
assert summary.status == "success"
detail = ReadRadarDetails(repository).detail(TARGET_DATE, SectorType.CONCEPT, "BK0001.DC")
assert detail.pct_change == Decimal(1)
assert len(detail.members) == 5
assert all(member.active_buy_net_amount_yuan is None for member in detail.members)
assert detail.members[0].net_amount_yuan == Decimal(10000)
assert detail.members[0].pct_change == Decimal(0)
assert detail.leaders["active_buy_net_amount_yuan"].top == ()
assert detail.summary["amount"].metric_value == Decimal("0.0015")
def test_detail_jaccard_uses_same_day_current_listed_members_and_independent_values() -> None:
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
from zhixing_server.modules.sector_radar.domain.source import MoneyflowRow
class DetailSource(FakeRadarSource):
def fetch_sector_members(
self, trade_date: date, sector_codes: Sequence[str]
) -> SourceResult[SectorMemberRow]:
rows = tuple(
SectorMemberRow(trade_date, code, f"00000{index}.SZ", f"股票{index}")
for code in sector_codes
for index in ((1, 2, 6) if code == "BK0001.DC" else (2, 3))
)
return self._result("dc_member", trade_date, rows)
def fetch_moneyflow(self, trade_date: date) -> SourceResult[MoneyflowRow]:
raw = (
{
"trade_date": trade_date.isoformat(),
"ts_code": "000001.SZ",
"net_mf_amount": "-2.5",
},
)
snapshot = build_source_snapshot(
api_name="moneyflow",
params={},
rows=raw,
target_trade_date=trade_date,
observed_at=NOW,
)
return SourceResult((snapshot,), tuple(MoneyflowRow.from_mapping(row) for row in raw))
repository = InMemorySectorRadarRepository()
result = BuildSectorRadar(DetailSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
assert result.status == "success"
reader = ReadRadarDetails(repository)
detail = reader.detail(TARGET_DATE, SectorType.CONCEPT, "BK0001.DC")
assert [member.ts_code for member in detail.members] == ["000001.SZ", "000002.SZ"]
assert detail.members[0].active_buy_net_amount_yuan == Decimal(-25000)
assert detail.members[0].net_amount_yuan == Decimal(10000)
assert detail.members[1].active_buy_net_amount_yuan is None
assert len(detail.leaders["active_buy_net_amount_yuan"].top) == 1
assert detail.similar_sectors[0].intersection_count == 1
assert detail.similar_sectors[0].union_count == 3
assert detail.similar_sectors[0].overlap_ratio == Decimal(1) / Decimal(3)
assert detail.similar_sectors[0].sector_type is SectorType.INDUSTRY
# A later build with different membership cannot alter the older detail.
BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE + timedelta(days=1))
)
assert reader.detail(TARGET_DATE, SectorType.CONCEPT, "BK0001.DC") == detail
def test_detail_history_and_ranking_extras_http_use_the_same_publication() -> None:
from fastapi.testclient import TestClient
from zhixing_server.bootstrap.app import create_app
from zhixing_server.modules.sector_radar.application.read import ReadSectorRadar
from zhixing_server.modules.sector_radar.presentation.http import get_sector_radar_reader
repository = InMemorySectorRadarRepository()
result = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
app = create_app()
app.dependency_overrides[get_sector_radar_reader] = lambda: ReadSectorRadar(repository)
with TestClient(app) as client:
params = {"trade_date": TARGET_DATE.isoformat()}
base = "/api/v1/sector-radar/sectors/concept/BK0001.DC"
history = client.get(base + "/history", params=params)
detail = client.get(base + "/detail", params=params)
ranking = client.get(
"/api/v1/sector-radar/rankings", params={**params, "view": "amount", "side": "top"}
)
assert history.status_code == detail.status_code == ranking.status_code == 200
payload = detail.json()
assert payload["publication"]["publication_id"] == result.outcomes[0].publication_id
assert payload["history"] == history.json()
assert payload["pct_change"] == "1"
assert payload["members"][0]["active_buy_net_amount_yuan"] is None
row = ranking.json()["rows"][0]
assert row["pct_change"] == "1"
assert Decimal(row["daily_net_amount_yuan"]) == 150000
assert Decimal(row["daily_ratio"]) == Decimal(150000) / Decimal(5000100)
assert row["on_list_count"] == row["history_available_days"] == 1
assert client.get(base + "/detail").status_code == 422
absent = client.get(base + "/detail", params={"trade_date": "2020-01-01"})
assert absent.json()["status"] == "no_data"
assert absent.json()["members"] == []
@pytest.mark.integration
def test_postgres_detail_migration_and_build_roundtrip(monkeypatch: pytest.MonkeyPatch) -> None:
import os
from pathlib import Path
import psycopg
from alembic import command
from alembic.config import Config
from zhixing_server.bootstrap.config import sqlalchemy_database_url
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
from zhixing_server.modules.sector_radar.application.recompute import RecomputeSectorRadar
from zhixing_server.modules.sector_radar.domain.metrics import (
AmountNetStrategy,
RatioTurnoverStrategy,
SwingEqualThreeToTenStrategy,
)
from zhixing_server.modules.sector_radar.domain.models import MetricKind
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
PostgresSectorRadarRepository,
)
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
if not database_url:
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
config = Config(str(Path(__file__).parents[3] / "alembic.ini"))
config.set_main_option(
"sqlalchemy.url", sqlalchemy_database_url(database_url).replace("%", "%%")
)
# Alembic intentionally reads Settings; bind the explicit test DSN and avoid
# fileConfig disabling unrelated test loggers in the same pytest process.
config.config_file_name = None
with monkeypatch.context() as context:
context.setenv("ZHIXING_DATABASE_URL", database_url)
from zhixing_server.bootstrap.config import get_settings
get_settings.cache_clear()
try:
command.upgrade(config, "head")
finally:
get_settings.cache_clear()
target = date(2098, 12, 1)
repository = PostgresSectorRadarRepository(database_url, max_connections=2)
try:
result = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(trade_date=target)
)
assert result.status in ("success", "unchanged")
detail = ReadRadarDetails(repository).detail(target, SectorType.CONCEPT, "BK0001.DC")
assert detail.history.status == "success"
assert detail.pct_change == Decimal(1)
assert len(detail.members) == 5
assert detail.members[0].active_buy_net_amount_yuan is None
assert detail.summary["amount"].metric_value == Decimal("0.0015")
with psycopg.connect(database_url) as connection:
assert connection.execute("SELECT version_num FROM alembic_version").fetchone() == (
"0010_radar_weighted_score",
)
row = connection.execute(
"SELECT pct_change, leading_code FROM sector_radar_daily_aggregate "
"WHERE publication_id = %s AND sector_type = 'concept'",
(result.outcomes[0].publication_id,),
).fetchone()
assert row == (Decimal(1), "000001.SZ")
facts = connection.execute(
"SELECT pct_change, active_buy_net_amount_yuan "
"FROM sector_radar_stock_fact WHERE trade_date = %s",
(target,),
).fetchall()
assert facts and all(row == (Decimal(0), None) for row in facts)
with pytest.raises(psycopg.errors.CheckViolation), connection.transaction():
connection.execute(
"UPDATE sector_radar_stock_fact "
"SET active_buy_net_amount_yuan = 'NaN'::numeric WHERE trade_date = %s",
(target,),
)
# Exercise every ranking projection against persisted v2 scores, while
# keeping the original source publication and its NULL score readable.
start, end = target + timedelta(days=31), target + timedelta(days=46)
interval = BuildSectorRadarCommand(start_date=start, end_date=end)
source = FakeRadarSource()
legacy = BuildSectorRadar(
source,
repository,
now_fn=lambda: NOW,
strategies=(
AmountNetStrategy(),
RatioTurnoverStrategy(),
SwingEqualThreeToTenStrategy(),
),
).execute(interval)
assert legacy.status in ("success", "unchanged")
old_id = legacy.outcomes[-1].publication_id
assert old_id is not None
old_rows = tuple(repository.load_rankings(old_id))
assert all(row.observation.weighted_score is None for row in old_rows)
source.calls.clear()
rescore = RecomputeSectorRadar(repository, now_fn=lambda: NOW + timedelta(hours=1))
updated = rescore.execute(interval)
assert updated.status in ("success", "unchanged")
updated_id = updated.outcomes[-1].publication_id
assert updated_id is not None and updated_id != old_id
assert source.calls == []
assert tuple(repository.load_rankings(old_id)) == old_rows
current_rows = tuple(repository.load_rankings(updated_id))
scored = next(
row
for row in current_rows
if row.observation.metric_kind is MetricKind.SWING
and row.observation.sector_type is SectorType.CONCEPT
)
expected = ((Decimal(5000000) + 1).log10() * 100).quantize(Decimal("0.000000000001"))
assert scored.observation.weighted_score == expected
assert scored.rank_change(5) == 0
assert dict(repository.load_publication_rankings((updated_id,)))[updated_id] == current_rows
previous = dict(repository.load_previous_rankings(end + timedelta(days=1), limit_dates=1))
assert previous[end] == current_rows
historical = repository.load_ranked_history((updated_id,), ("BK0001.DC",))
assert any(row.ranking == scored for row in historical)
detail = ReadRadarDetails(repository).detail(end, SectorType.CONCEPT, "BK0001.DC")
assert detail.summary["swing"].weighted_score == expected
assert detail.pct_change == 1 and len(detail.members) == 5
assert rescore.execute(interval).status == "unchanged"
with (
psycopg.connect(database_url) as connection,
pytest.raises(psycopg.errors.CheckViolation),
connection.transaction(),
):
connection.execute(
"UPDATE sector_radar_ranking SET weighted_score = 'NaN'::numeric "
"WHERE publication_id = %s",
(updated_id,),
)
finally:
repository.close()
def test_offline_rescore_is_idempotent_preserves_old_publications_and_details() -> None:
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails
from zhixing_server.modules.sector_radar.application.read import (
RadarQuery,
RadarView,
ReadSectorRadar,
)
from zhixing_server.modules.sector_radar.application.recompute import RecomputeSectorRadar
from zhixing_server.modules.sector_radar.domain.metrics import (
AmountNetStrategy,
RatioTurnoverStrategy,
SwingEqualThreeToTenStrategy,
)
source = FakeRadarSource()
repository = InMemorySectorRadarRepository()
end = TARGET_DATE + timedelta(days=15)
command = BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end)
initial = BuildSectorRadar(
source,
repository,
now_fn=lambda: NOW,
strategies=(AmountNetStrategy(), RatioTurnoverStrategy(), SwingEqualThreeToTenStrategy()),
).execute(command)
assert initial.status == "success"
old_publications = repository.publications.copy()
old_rankings = repository.rankings.copy()
snapshot_count = len(repository.source_snapshots)
source.calls.clear()
source.fail_daily = True
service = RecomputeSectorRadar(repository, now_fn=lambda: NOW + timedelta(hours=1))
result = service.execute(command)
assert result.status == "success"
assert all(repository.publications[key] == value for key, value in old_publications.items())
assert all(repository.rankings[key] == value for key, value in old_rankings.items())
assert len(repository.source_snapshots) == snapshot_count
assert source.calls == []
latest = repository.get_last_good_publication(end)
assert latest is not None
assert "zhixing_swing_weighted_v2" in latest.metric_versions
count = len(repository.publications)
assert service.execute(command).status == "unchanged"
assert len(repository.publications) == count
reader = ReadSectorRadar(repository)
page = reader.query(RadarQuery(trade_date=end, view=RadarView.SWING))
assert len(page.rows) == 1
assert page.rows[0].observation.weighted_score is not None
assert page.rows[0].rank_change(5) == 0
detail = ReadRadarDetails(repository).detail(end, SectorType.CONCEPT, "BK0001.DC")
assert len(detail.members) == 5
assert detail.pct_change == 1
assert detail.summary["swing"].weighted_score == page.rows[0].observation.weighted_score
assert detail.summary["swing"].rank_position == page.rows[0].rank_position
def test_offline_rescore_failure_does_not_replace_last_good(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from unittest.mock import Mock
from zhixing_server.modules.sector_radar.application.recompute import RecomputeSectorRadar
repository = InMemorySectorRadarRepository()
command = BuildSectorRadarCommand(trade_date=TARGET_DATE)
BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(command)
before = repository.get_last_good_publication()
monkeypatch.setattr(
repository, "finalize_publication", Mock(side_effect=RuntimeError("injected failure"))
)
result = RecomputeSectorRadar(repository, now_fn=lambda: NOW + timedelta(hours=1)).execute(
command
)
assert result.status == "failed"
assert repository.get_last_good_publication() == before
latest = repository.get_latest_publication()
assert latest is not None and latest.status is PublicationStatus.FAILED