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

868 lines
34 KiB
Python
Raw Normal View History

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)
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("0.03")
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.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() == (
"0009_radar_sector_detail",
)
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,),
)
finally:
repository.close()