Merge branch 'develop' into codex/point

This commit is contained in:
yuxuanhui
2026-08-31 16:14:35 +08:00
86 changed files with 13384 additions and 217 deletions
@@ -39,6 +39,13 @@ def test_postgres_migration_creates_market_data_contract(
"selection_run",
"selection_run_item",
"selection_signal",
"sector_radar_source_snapshot",
"sector_radar_membership",
"sector_radar_stock_fact",
"sector_radar_publication",
"sector_radar_ranking",
"sector_radar_daily_aggregate",
"sector_radar_publication_source",
} <= tables
item_columns = {column["name"] for column in inspector.get_columns("selection_run_item")}
assert {
@@ -0,0 +1,175 @@
import os
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
from pathlib import Path
import psycopg
import pytest
from alembic import command
from alembic.config import Config
from zhixing_server.bootstrap.config import sqlalchemy_database_url
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
PublicationStatus,
RadarPublication,
SectorType,
StockFactStatus,
)
from zhixing_server.modules.sector_radar.domain.persistence import (
MembershipRecord,
StockFactRecord,
)
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
PostgresSectorRadarRepository,
)
TARGET_DATE = date(2099, 1, 4)
STARTED_AT = datetime(2099, 1, 4, 17, 30, tzinfo=UTC)
def prepare_database(database_url: str) -> None:
server_root = Path(__file__).parents[2]
config = Config(str(server_root / "alembic.ini"))
sqlalchemy_url = sqlalchemy_database_url(database_url)
config.set_main_option("sqlalchemy.url", sqlalchemy_url.replace("%", "%%"))
command.upgrade(config, "head")
@pytest.mark.integration
def test_postgres_sector_radar_revisions_and_last_good() -> None:
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
if not database_url:
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
prepare_database(database_url)
snapshot = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20990104"},
rows=(
{
"trade_date": "20990104",
"ts_code": "BKTEST.DC",
"con_code": "000001.SZ",
"name": "测试股票",
},
),
target_trade_date=TARGET_DATE,
observed_at=STARTED_AT,
)
fact_revision = "b" * 64
publication_ids = ("test-sector-radar-success", "test-sector-radar-failed")
with psycopg.connect(database_url) as connection, connection.transaction():
connection.execute(
"DELETE FROM sector_radar_publication WHERE id = ANY(%s)",
(list(publication_ids),),
)
connection.execute(
"DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s",
(fact_revision,),
)
connection.execute(
"DELETE FROM sector_radar_source_snapshot WHERE id = %s",
(snapshot.snapshot_id,),
)
repository = PostgresSectorRadarRepository(database_url, max_connections=2)
try:
assert repository.save_source_snapshots((snapshot,)).inserted == 1
assert (
repository.save_source_snapshots(
(replace(snapshot, observed_at=STARTED_AT + timedelta(minutes=1)),)
).unchanged
== 1
)
assert (
repository.save_memberships(
(
MembershipRecord(
source_snapshot_id=snapshot.snapshot_id,
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BKTEST.DC",
sector_name="测试概念",
stock_code="000001.SZ",
stock_name="测试股票",
status=MembershipStatus.AVAILABLE,
),
)
).inserted
== 1
)
assert (
repository.save_stock_facts(
(
StockFactRecord(
fact_revision=fact_revision,
source_snapshot_ids=(snapshot.snapshot_id,),
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("1000"),
net_amount_yuan=Decimal("100"),
),
)
).inserted
== 1
)
running = RadarPublication(
publication_id=publication_ids[0],
target_trade_date=TARGET_DATE,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version=snapshot.content_sha256,
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal(0),
started_at=STARTED_AT,
)
repository.create_publication(running)
repository.finish_publication(
replace(
running,
status=PublicationStatus.SUCCESS,
input_hash="a" * 64,
coverage=Decimal(1),
finished_at=STARTED_AT + timedelta(minutes=5),
)
)
failed = replace(
running,
publication_id=publication_ids[1],
started_at=STARTED_AT + timedelta(minutes=6),
)
repository.create_publication(failed)
repository.finish_publication(
replace(
failed,
status=PublicationStatus.FAILED,
coverage=Decimal("0.8"),
finished_at=STARTED_AT + timedelta(minutes=7),
error_summary="safe_error",
)
)
last_good = repository.get_last_good_publication(TARGET_DATE)
assert last_good is not None
assert last_good.publication_id == publication_ids[0]
finally:
repository.close()
with psycopg.connect(database_url) as connection, connection.transaction():
connection.execute(
"DELETE FROM sector_radar_publication WHERE id = ANY(%s)",
(list(publication_ids),),
)
connection.execute(
"DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s",
(fact_revision,),
)
connection.execute(
"DELETE FROM sector_radar_source_snapshot WHERE id = %s",
(snapshot.snapshot_id,),
)
@@ -0,0 +1,230 @@
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
import pytest
from fastapi.testclient import TestClient
from pydantic import ValidationError
from zhixing_server.bootstrap.app import create_app
from zhixing_server.modules.sector_radar.application.read import (
RadarDateIndex,
RadarMetricDefinition,
RadarQuery,
RadarView,
RankingPage,
)
from zhixing_server.modules.sector_radar.domain.metrics import AmountNetStrategy
from zhixing_server.modules.sector_radar.domain.models import (
MetricKind,
MetricObservation,
MetricQuality,
MetricUnit,
PublicationStatus,
RadarPublication,
RankChange,
RankedMetric,
RankSide,
SectorType,
)
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
SectorRadarRepositoryError,
)
from zhixing_server.modules.sector_radar.presentation.http import (
RadarRankingRowResponse,
get_sector_radar_reader,
)
TARGET_DATE = date(2026, 8, 28)
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
def _publication(
publication_id: str,
status: PublicationStatus,
*,
trade_date: date = TARGET_DATE,
) -> RadarPublication:
return RadarPublication(
publication_id=publication_id,
target_trade_date=trade_date,
status=status,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-v1",
metric_versions=(AmountNetStrategy.metric_version,),
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
coverage=Decimal(1) if status is PublicationStatus.SUCCESS else Decimal("0.8"),
started_at=NOW,
finished_at=None if status is PublicationStatus.RUNNING else NOW + timedelta(minutes=5),
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
)
def _ranking() -> RankedMetric:
return RankedMetric(
observation=MetricObservation(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="机器人",
metric_kind=MetricKind.AMOUNT,
metric_version=AmountNetStrategy.metric_version,
implementation_kind="independent",
unit=MetricUnit.CNY_100M,
value=Decimal("12.5"),
quality=MetricQuality.AVAILABLE,
member_count=20,
valid_sample_count=19,
membership_coverage=Decimal(1),
moneyflow_coverage=Decimal("0.95"),
),
rank_position=1,
rank_percentile=Decimal(100),
rank_changes=(RankChange(days=5, value=3),),
)
class FakeReader:
def __init__(self, *, no_data: bool = False, fail: bool = False) -> None:
self.fail = fail
self.last_query: RadarQuery | None = None
success = _publication("publication-success", PublicationStatus.SUCCESS)
current = _publication(
"publication-partial",
PublicationStatus.PARTIAL,
trade_date=TARGET_DATE + timedelta(days=1),
)
self.date_index = RadarDateIndex(
available_dates=() if no_data else (TARGET_DATE,),
current_attempt=None if no_data else current,
last_good=None if no_data else success,
)
query = RadarQuery()
self.page = RankingPage(
status="no_data" if no_data else "success",
query=query,
publication=None if no_data else success,
definition=RadarMetricDefinition(
metric_kind=MetricKind.AMOUNT,
metric_version=AmountNetStrategy.metric_version,
label="主力净流入(知行独立实现)",
unit=MetricUnit.CNY_100M,
),
rows=() if no_data else (_ranking(),),
total=0 if no_data else 1,
)
def list_dates(self) -> RadarDateIndex:
if self.fail:
raise SectorRadarRepositoryError("private database detail")
return self.date_index
def query(self, query: RadarQuery) -> RankingPage:
if self.fail:
raise SectorRadarRepositoryError("private database detail")
self.last_query = query
return replace(self.page, query=query)
def _client(reader: FakeReader) -> TestClient:
application = create_app()
application.dependency_overrides[get_sector_radar_reader] = lambda: reader
return TestClient(application)
def test_dates_exposes_partial_attempt_without_replacing_last_good() -> None:
response = _client(FakeReader()).get("/api/v1/sector-radar/dates")
assert response.status_code == 200
payload = response.json()
assert payload["status"] == "success"
assert payload["available_dates"] == ["2026-08-28"]
assert payload["current_attempt"]["status"] == "partial"
assert payload["last_good"]["status"] == "success"
assert payload["last_good"]["coverage"] == "1"
def test_rankings_maps_filters_and_independent_metric_contract() -> None:
reader = FakeReader()
response = _client(reader).get(
"/api/v1/sector-radar/rankings",
params={
"trade_date": "2026-08-28",
"sector_type": "concept",
"view": "rank_change",
"rank_change_metric": "amount",
"rank_change_days": 5,
"side": "top",
"search": " 机器人 ",
"page": 2,
"page_size": 10,
},
)
assert response.status_code == 200
assert reader.last_query == RadarQuery(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
view=RadarView.RANK_CHANGE,
rank_change_metric=MetricKind.AMOUNT,
rank_change_days=5,
side=RankSide.TOP,
search="机器人",
page=2,
page_size=10,
)
payload = response.json()
assert payload["definition"]["metric_version"] == "zhixing_amount_net_bn_v1"
assert payload["definition"]["implementation_kind"] == "independent"
assert "知行独立实现" in payload["definition"]["disclaimer"]
assert payload["rows"][0]["rank_change"] == 3
assert payload["rows"][0]["unit"] == "CNY_100M"
def test_no_data_is_a_stable_200_response() -> None:
client = _client(FakeReader(no_data=True))
dates = client.get("/api/v1/sector-radar/dates")
rankings = client.get("/api/v1/sector-radar/rankings")
assert dates.status_code == 200
assert dates.json()["status"] == "no_data"
assert rankings.status_code == 200
assert rankings.json()["status"] == "no_data"
assert rankings.json()["rows"] == []
def test_http_contract_rejects_zero_rank_percentile() -> None:
payload = _client(FakeReader()).get("/api/v1/sector-radar/rankings").json()["rows"][0]
payload["rank_percentile"] = "0"
with pytest.raises(ValidationError):
RadarRankingRowResponse.model_validate(payload)
def test_invalid_query_values_return_422() -> None:
client = _client(FakeReader())
for params in (
{"rank_change_days": 0},
{"rank_change_days": 6},
{"page": 0},
{"page_size": 101},
{"sector_type": "region"},
{"view": "unknown"},
{"side": "unknown"},
):
assert client.get("/api/v1/sector-radar/rankings", params=params).status_code == 422
def test_repository_error_maps_to_redacted_503() -> None:
response = _client(FakeReader(fail=True)).get("/api/v1/sector-radar/rankings")
assert response.status_code == 503
assert response.json() == {
"detail": {
"code": "sector_radar_storage_unavailable",
"message": "sector radar storage is unavailable",
}
}
assert "private database detail" not in response.text
@@ -1,3 +1,4 @@
import threading
from datetime import date
import pytest
@@ -87,6 +88,109 @@ def test_rate_limit_cooldown_is_shared_by_following_requests() -> None:
assert waits == [60]
def test_request_start_interval_allows_overlapping_provider_calls() -> None:
current = [0.0]
state_lock = threading.Lock()
first_started = threading.Event()
release_first = threading.Event()
waits: list[float] = []
starts: list[tuple[str, float]] = []
errors: list[BaseException] = []
def clock() -> float:
with state_lock:
return current[0]
def wait(seconds: float) -> None:
with state_lock:
waits.append(seconds)
current[0] += seconds
coordinator = RequestCoordinator(
max_retries=0,
request_interval_seconds=0.2,
clock=clock,
wait_fn=wait,
sleep_fn=wait,
)
def first_request() -> object:
starts.append(("first", clock()))
first_started.set()
if not release_first.wait(timeout=2):
raise AssertionError("first provider call was not released")
return "first"
def run_first() -> None:
try:
coordinator.call("first", first_request)
except BaseException as exc: # pragma: no cover - surfaced by the assertion below
errors.append(exc)
first_thread = threading.Thread(target=run_first)
first_thread.start()
assert first_started.wait(timeout=2)
second = coordinator.call(
"second",
lambda: starts.append(("second", clock())) or "second",
)
assert second == "second"
assert first_thread.is_alive()
release_first.set()
first_thread.join(timeout=2)
assert not first_thread.is_alive()
assert errors == []
assert starts == [("first", 0.0), ("second", 0.2)]
assert waits == [0.2]
def test_request_start_interval_is_disabled_by_default() -> None:
waits: list[float] = []
starts: list[str] = []
coordinator = RequestCoordinator(
max_retries=0,
clock=lambda: 0.0,
wait_fn=waits.append,
)
coordinator.call("first", lambda: starts.append("first"))
coordinator.call("second", lambda: starts.append("second"))
assert starts == ["first", "second"]
assert waits == []
def test_request_start_interval_applies_to_retry_attempts() -> None:
current = [0.0]
waits: list[float] = []
starts: list[float] = []
def wait(seconds: float) -> None:
waits.append(seconds)
current[0] += seconds
coordinator = RequestCoordinator(
max_retries=1,
backoff_seconds=0,
request_interval_seconds=0.2,
clock=lambda: current[0],
wait_fn=wait,
sleep_fn=wait,
)
def request() -> object:
starts.append(current[0])
if len(starts) == 1:
raise RuntimeError("transient provider failure")
return "ok"
assert coordinator.call("daily", request) == "ok"
assert starts == [0.0, 0.2]
assert waits == [0.0, 0.2]
def test_pro_bar_qfq_calls_are_bound_to_the_shared_coordinator(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -0,0 +1,641 @@
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)
@@ -0,0 +1,129 @@
import json
from datetime import date
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.application.build import (
BuildDateOutcome,
BuildOutcomeStatus,
BuildSummary,
)
from zhixing_server.modules.sector_radar.presentation import cli
from zhixing_server.modules.sector_radar.presentation.cli import build_parser
def test_sector_radar_cli_parses_single_range_and_retry_modes() -> None:
parser = build_parser()
single = parser.parse_args(["--trade-date", "2026-08-28"])
date_range = parser.parse_args(["--start-date", "2026-08-18", "--end-date", "2026-08-28"])
retry = parser.parse_args(["--retry-publication-id", "publication-a"])
assert single.trade_date == date(2026, 8, 28)
assert date_range.start_date == date(2026, 8, 18)
assert date_range.end_date == date(2026, 8, 28)
assert retry.retry_publication_id == "publication-a"
class FakeSettings:
log_level = "INFO"
tushare_token = "secret-token"
database_url = "postgresql://unused"
sector_radar_max_retries = 3
sector_radar_retry_backoff_seconds = 1.0
sector_radar_request_interval_seconds = 0.2
sector_radar_advisory_lock_key = 7_380_522
sector_radar_coverage_threshold = Decimal("0.99")
class FakeRepository:
closed = False
def __init__(self, database_url: str, *, advisory_lock_key: int) -> None:
assert database_url == "postgresql://unused"
assert advisory_lock_key == 7_380_522
def close(self) -> None:
self.closed = True
@pytest.mark.parametrize(
("outcome_status", "coverage", "expected_code"),
(("success", Decimal(1), 0), ("partial", Decimal("0.8"), 2), ("failed", Decimal(0), 1)),
)
def test_cli_main_returns_summary_exit_code_and_json(
monkeypatch: pytest.MonkeyPatch,
capsys: pytest.CaptureFixture[str],
outcome_status: BuildOutcomeStatus,
coverage: Decimal,
expected_code: int,
) -> None:
summary = BuildSummary(
(
BuildDateOutcome(
date(2026, 8, 28),
outcome_status,
"publication-a",
coverage,
2,
6,
),
)
)
class FakeSourceFactory:
@staticmethod
def from_token(token: str, **kwargs: object) -> object:
assert token == "secret-token"
assert kwargs["max_retries"] == 3
assert kwargs["backoff_seconds"] == 1.0
assert kwargs["request_interval_seconds"] == 0.2
return object()
class FakeBuild:
def __init__(self, source: object, repository: object, **kwargs: object) -> None:
assert source is not None
assert repository is not None
assert kwargs
def execute(self, command: object) -> BuildSummary:
assert command is not None
return summary
monkeypatch.setattr(cli, "get_settings", FakeSettings)
monkeypatch.setattr(cli, "TushareSectorRadarAdapter", FakeSourceFactory)
monkeypatch.setattr(cli, "PostgresSectorRadarRepository", FakeRepository)
monkeypatch.setattr(cli, "BuildSectorRadar", FakeBuild)
exit_code = cli.main(["--trade-date", "2026-08-28"])
output = json.loads(capsys.readouterr().out)
assert exit_code == expected_code
assert output["status"] == summary.status
assert output["exit_code"] == expected_code
assert "secret-token" not in str(output)
def test_cli_initialization_failure_is_redacted(
monkeypatch: pytest.MonkeyPatch,
capsys: pytest.CaptureFixture[str],
) -> None:
class FailingSourceFactory:
@staticmethod
def from_token(token: str, **kwargs: object) -> object:
del token, kwargs
raise RuntimeError("private provider detail secret-token")
monkeypatch.setattr(cli, "get_settings", FakeSettings)
monkeypatch.setattr(cli, "TushareSectorRadarAdapter", FailingSourceFactory)
exit_code = cli.main(["--trade-date", "2026-08-28"])
captured = capsys.readouterr()
output = json.loads(captured.out)
assert exit_code == 1
assert output["status"] == "failed"
assert output["error_type"] == "RuntimeError"
assert "private provider detail" not in captured.out
assert "secret-token" not in captured.out
@@ -0,0 +1,146 @@
from datetime import UTC, date, datetime
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.domain.facts import aggregate_sector_snapshot
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
PublicationStatus,
RadarPublication,
SectorMembershipSnapshot,
SectorType,
StockDailyFact,
StockFactStatus,
)
TARGET_DATE = date(2026, 8, 28)
def test_point_in_time_aggregation_distinguishes_suspension_missing_and_zero() -> None:
snapshot = SectorMembershipSnapshot(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
member_codes=("000001.SZ", "000002.SZ", "000003.SZ", "000004.SZ"),
status=MembershipStatus.AVAILABLE,
source_version="dc-member-20260828-a",
)
facts = (
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("1000"),
net_amount_yuan=Decimal("100"),
),
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000002.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("2000"),
net_amount_yuan=Decimal("0"),
),
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000003.SZ",
status=StockFactStatus.SUSPENDED,
),
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000004.SZ",
status=StockFactStatus.MISSING_MONEYFLOW,
),
)
aggregate = aggregate_sector_snapshot(snapshot, facts)
assert aggregate.member_count == 4
assert aggregate.valid_sample_count == 2
assert aggregate.net_amount_yuan == Decimal("100")
assert aggregate.turnover_yuan == Decimal("3000")
assert aggregate.membership_coverage == Decimal("1")
assert aggregate.moneyflow_coverage == Decimal("2") / Decimal("3")
def test_unknown_membership_never_falls_back_to_available_stock_facts() -> None:
snapshot = SectorMembershipSnapshot(
trade_date=TARGET_DATE,
sector_type=SectorType.INDUSTRY,
sector_code="BK1001.DC",
sector_name="示例行业",
member_codes=(),
status=MembershipStatus.UNKNOWN,
source_version="dc-member-missing",
)
fact = StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("1000"),
net_amount_yuan=Decimal("100"),
)
aggregate = aggregate_sector_snapshot(snapshot, (fact,))
assert aggregate.member_count == 0
assert aggregate.net_amount_yuan is None
assert aggregate.turnover_yuan is None
assert aggregate.membership_coverage == Decimal("0")
def test_stock_fact_rejects_non_finite_values_and_invalid_status_payloads() -> None:
with pytest.raises(ValueError, match="finite"):
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("Infinity"),
net_amount_yuan=Decimal("1"),
)
with pytest.raises(ValueError, match="must not expose amounts"):
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.SUSPENDED,
turnover_yuan=Decimal("0"),
)
def test_publication_requires_terminal_completion_and_replay_identity() -> None:
started_at = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
publication = RadarPublication(
publication_id="radar-20260828-a",
target_trade_date=TARGET_DATE,
status=PublicationStatus.SUCCESS,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-20260828-a",
metric_versions=(
"zhixing_amount_net_bn_v1",
"zhixing_ratio_turnover_v1",
"zhixing_swing_equal_3_10_v1",
),
input_hash="a" * 64,
coverage=Decimal("0.995"),
started_at=started_at,
finished_at=datetime(2026, 8, 28, 17, 35, tzinfo=UTC),
)
assert publication.status is PublicationStatus.SUCCESS
with pytest.raises(ValueError, match="finished_at"):
RadarPublication(
publication_id="radar-20260828-running",
target_trade_date=TARGET_DATE,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-20260828-a",
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal("0"),
started_at=started_at,
finished_at=started_at,
)
@@ -0,0 +1,112 @@
from datetime import date
from decimal import Decimal
from zhixing_server.modules.sector_radar.domain.metrics import (
AmountNetStrategy,
RatioTurnoverStrategy,
SwingEqualThreeToTenStrategy,
)
from zhixing_server.modules.sector_radar.domain.models import (
MetricQuality,
SectorDailyAggregate,
SectorType,
)
TARGET_DATE = date(2026, 8, 28)
def make_aggregate(
*,
net_amount_yuan: Decimal | None = Decimal("125000000"),
turnover_yuan: Decimal | None = Decimal("5000000000"),
) -> SectorDailyAggregate:
return SectorDailyAggregate(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
member_count=10,
valid_sample_count=10,
net_amount_yuan=net_amount_yuan,
turnover_yuan=turnover_yuan,
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
def test_amount_and_ratio_strategies_expose_independent_versioned_values() -> None:
aggregate = make_aggregate()
amount = AmountNetStrategy().evaluate((aggregate,), TARGET_DATE)
ratio = RatioTurnoverStrategy().evaluate((aggregate,), TARGET_DATE)
assert amount.value == Decimal("1.25")
assert amount.metric_version == "zhixing_amount_net_bn_v1"
assert amount.implementation_kind == "independent"
assert amount.unit == "CNY_100M"
assert amount.quality is MetricQuality.AVAILABLE
assert ratio.value == Decimal("0.025")
assert ratio.metric_version == "zhixing_ratio_turnover_v1"
assert ratio.implementation_kind == "independent"
assert ratio.unit == "ratio"
def test_missing_moneyflow_is_unavailable_but_zero_remains_a_real_value() -> None:
missing = AmountNetStrategy().evaluate((make_aggregate(net_amount_yuan=None),), TARGET_DATE)
zero = AmountNetStrategy().evaluate(
(make_aggregate(net_amount_yuan=Decimal("0")),), TARGET_DATE
)
assert missing.value is None
assert missing.quality is MetricQuality.UNAVAILABLE
assert zero.value == Decimal("0")
assert zero.quality is MetricQuality.AVAILABLE
def test_swing_strategy_uses_each_days_point_in_time_aggregate() -> None:
history = tuple(
SectorDailyAggregate(
trade_date=date(2026, 8, 18 + offset),
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
member_count=6 + offset,
valid_sample_count=6 + offset,
net_amount_yuan=Decimal(str(offset + 1)),
turnover_yuan=Decimal("100"),
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
for offset in range(10)
)
result = SwingEqualThreeToTenStrategy().evaluate(history, date(2026, 8, 27))
# The worked 3..10-day window ratios average to exactly 0.0725.
assert result.value == Decimal("0.0725")
assert result.metric_version == "zhixing_swing_equal_3_10_v1"
assert result.member_count == 15
def test_swing_strategy_carries_forward_limited_historical_sample_quality() -> None:
history = tuple(
SectorDailyAggregate(
trade_date=date(2026, 8, 18 + offset),
sector_type=SectorType.INDUSTRY,
sector_code="BK1001.DC",
sector_name="示例行业",
member_count=10,
valid_sample_count=4 if offset == 0 else 10,
net_amount_yuan=Decimal("10"),
turnover_yuan=Decimal("100"),
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
for offset in range(10)
)
result = SwingEqualThreeToTenStrategy().evaluate(history, date(2026, 8, 27))
assert result.value == Decimal("0.1")
assert result.quality is MetricQuality.AVAILABLE_LIMITED_SAMPLE
@@ -0,0 +1,215 @@
from datetime import UTC, date, datetime
from decimal import Decimal
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
SectorType,
StockFactStatus,
)
from zhixing_server.modules.sector_radar.domain.normalize import (
normalize_memberships,
normalize_stock_facts,
)
from zhixing_server.modules.sector_radar.domain.source import (
DailyRow,
MoneyflowDcRow,
SectorIndexRow,
SectorMemberRow,
SourceResult,
StockBasicRow,
SuspendRow,
build_source_snapshot,
)
TARGET_DATE = date(2026, 8, 28)
OBSERVED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
def result[T](api_name: str, rows: tuple[T, ...]) -> SourceResult[T]:
snapshot = build_source_snapshot(
api_name=api_name,
params={"trade_date": "20260828"},
rows=(),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
)
return SourceResult((snapshot,), rows)
def basic(
ts_code: str,
*,
list_date: date = date(2020, 1, 1),
market: str | None = "主板",
) -> StockBasicRow:
return StockBasicRow(
ts_code=ts_code,
symbol=ts_code.split(".")[0],
name=ts_code,
market=market,
exchange="SZSE",
list_status="L",
list_date=list_date,
delist_date=None,
)
def test_missing_market_keeps_hs_a_stock_but_code_rules_still_exclude_bse_and_b_shares() -> None:
codes = ("000001.SZ", "920001.BJ", "200001.SZ", "900001.SH")
basics = tuple(basic(code, market=None) for code in codes)
daily_rows = tuple(daily(code, Decimal("1")) for code in codes)
moneyflow_rows = tuple(moneyflow(code, Decimal("1")) for code in codes)
facts = normalize_stock_facts(
target_trade_date=TARGET_DATE,
candidate_codes=codes,
stock_basics=result("stock_basic", basics),
suspensions=result("suspend_d", ()),
daily=result("daily", daily_rows),
moneyflow=result("moneyflow_dc", moneyflow_rows),
)
by_code = {fact.ts_code: fact for fact in facts}
assert by_code["000001.SZ"].status is StockFactStatus.AVAILABLE
assert by_code["920001.BJ"].status is StockFactStatus.LIFECYCLE_INVALID
assert by_code["200001.SZ"].status is StockFactStatus.LIFECYCLE_INVALID
assert by_code["900001.SH"].status is StockFactStatus.LIFECYCLE_INVALID
def daily(ts_code: str, amount: Decimal | None) -> DailyRow:
return DailyRow(
ts_code=ts_code,
trade_date=TARGET_DATE,
close=Decimal("10"),
pre_close=Decimal("10"),
pct_chg=Decimal(0),
volume=Decimal(0),
amount_thousand_yuan=amount,
)
def moneyflow(ts_code: str, amount: Decimal | None) -> MoneyflowDcRow:
return MoneyflowDcRow(
trade_date=TARGET_DATE,
ts_code=ts_code,
name=ts_code,
net_amount_ten_thousand_yuan=amount,
net_amount_rate=Decimal(0),
pct_change=Decimal(0),
close=Decimal("10"),
)
def test_membership_normalization_persists_an_explicit_unknown_sector() -> None:
indices = (
SectorIndexRow(
TARGET_DATE,
SectorType.CONCEPT,
"BK0001.DC",
"机器人",
"一级",
Decimal(1),
None,
),
SectorIndexRow(
TARGET_DATE,
SectorType.CONCEPT,
"BK0002.DC",
"低空经济",
"一级",
Decimal(1),
None,
),
)
member = SectorMemberRow(
TARGET_DATE,
"BK0001.DC",
"000001.SZ",
"平安银行",
)
all_snapshot = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20260828"},
rows=(
{
"trade_date": "20260828",
"ts_code": member.sector_code,
"con_code": member.stock_code,
"name": member.stock_name,
},
),
target_trade_date=TARGET_DATE,
partition_key="all",
observed_at=OBSERVED_AT,
)
empty_partition = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20260828", "ts_code": "BK0002.DC"},
rows=(),
target_trade_date=TARGET_DATE,
partition_key="BK0002.DC",
observed_at=OBSERVED_AT,
)
records = normalize_memberships(
indices,
SourceResult((all_snapshot, empty_partition), (member,)),
)
assert records[0].status is MembershipStatus.AVAILABLE
assert records[0].stock_code == "000001.SZ"
assert records[1].status is MembershipStatus.UNKNOWN
assert records[1].stock_code is None
assert records[1].membership_key == "__membership_unknown__"
def test_stock_fact_normalization_preserves_all_missing_and_zero_states() -> None:
codes = tuple(f"00000{index}.SZ" for index in range(1, 9))
basics = tuple(
basic(code, list_date=date(2027, 1, 1) if code == codes[7] else date(2020, 1, 1))
for code in codes
)
daily_rows = (
daily(codes[0], Decimal("1")),
daily(codes[3], None),
daily(codes[4], Decimal("1")),
daily(codes[5], Decimal("1")),
daily(codes[6], Decimal("0")),
daily(codes[7], Decimal("1")),
)
moneyflow_rows = (
moneyflow(codes[0], Decimal("0")),
moneyflow(codes[3], Decimal("1")),
moneyflow(codes[5], None),
moneyflow(codes[6], Decimal("0")),
moneyflow(codes[7], Decimal("1")),
)
suspensions = (
SuspendRow(
ts_code=codes[1],
trade_date=TARGET_DATE,
suspend_timing=None,
suspend_type="停牌",
),
)
facts = normalize_stock_facts(
target_trade_date=TARGET_DATE,
candidate_codes=codes,
stock_basics=result("stock_basic", basics),
suspensions=result("suspend_d", suspensions),
daily=result("daily", daily_rows),
moneyflow=result("moneyflow_dc", moneyflow_rows),
)
by_code = {fact.ts_code: fact for fact in facts}
assert by_code[codes[0]].status is StockFactStatus.AVAILABLE
assert by_code[codes[0]].turnover_yuan == Decimal("1000")
assert by_code[codes[0]].net_amount_yuan == Decimal("0")
assert by_code[codes[1]].status is StockFactStatus.SUSPENDED
assert by_code[codes[2]].status is StockFactStatus.MISSING_DAILY
assert by_code[codes[3]].status is StockFactStatus.NULL_DAILY_AMOUNT
assert by_code[codes[4]].status is StockFactStatus.MISSING_MONEYFLOW
assert by_code[codes[5]].status is StockFactStatus.NULL_MONEYFLOW
assert by_code[codes[6]].status is StockFactStatus.LOW_LIQUIDITY
assert by_code[codes[7]].status is StockFactStatus.LIFECYCLE_INVALID
@@ -0,0 +1,170 @@
from collections.abc import Generator
from contextlib import contextmanager
from datetime import UTC, date, datetime
from decimal import Decimal
from typing import Any, cast
from psycopg_pool import ConnectionPool
from zhixing_server.modules.sector_radar.domain.models import PublicationStatus
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
PostgresSectorRadarRepository,
)
TARGET_DATE = date(2026, 8, 28)
class FakeResult:
def __init__(
self,
row: tuple[object, ...] | None = None,
rows: tuple[tuple[object, ...], ...] | None = None,
) -> None:
self.row = row
self.rows = rows or (() if row is None else (row,))
def fetchone(self) -> tuple[object, ...] | None:
return self.row
def fetchall(self) -> tuple[tuple[object, ...], ...]:
return self.rows
class FakeConnection:
def __init__(self) -> None:
self.statements: list[tuple[str, tuple[object, ...]]] = []
def execute(
self,
query: str,
parameters: tuple[object, ...] = (),
) -> FakeResult:
self.statements.append((query, parameters))
if "FROM sector_radar_ranking" in query:
return FakeResult(
rows=(
(
TARGET_DATE,
"concept",
"BK0001.DC",
"机器人",
"amount",
"zhixing_amount_net_bn_v1",
"independent",
"CNY_100M",
Decimal("12.5"),
"available",
20,
19,
Decimal(1),
Decimal("0.95"),
1,
Decimal(100),
{"1": 3, "2": None},
),
)
)
if "FROM sector_radar_publication" in query:
return FakeResult(
(
"publication-a",
TARGET_DATE,
"success",
"tushare-pro-v1",
"eastmoney-dc-v1",
["zhixing_amount_net_bn_v1"],
"a" * 64,
Decimal("1"),
datetime(2026, 8, 28, 17, 30, tzinfo=UTC),
datetime(2026, 8, 28, 17, 35, tzinfo=UTC),
None,
)
)
if "pg_try_advisory_lock" in query:
return FakeResult((True,))
return FakeResult((True,))
class FakePool:
def __init__(self, connection: FakeConnection) -> None:
self._connection = connection
self._opened = False
def open(self, *, wait: bool) -> None:
assert wait
self._opened = True
def close(self) -> None:
self._opened = False
@contextmanager
def connection(self) -> Generator[FakeConnection]:
yield self._connection
def make_repository(connection: FakeConnection) -> PostgresSectorRadarRepository:
pool = cast(ConnectionPool[Any], cast(object, FakePool(connection)))
return PostgresSectorRadarRepository("postgresql://unused", pool=pool)
def test_last_good_query_strictly_filters_success_and_date() -> None:
connection = FakeConnection()
publication = make_repository(connection).get_last_good_publication(TARGET_DATE)
assert publication is not None
assert publication.status is PublicationStatus.SUCCESS
query, parameters = connection.statements[0]
assert "status = 'success'" in query
assert "partial" not in query
assert "target_trade_date <= %s" in query
assert parameters == (TARGET_DATE,)
def test_advisory_lock_uses_target_date_and_releases_same_key() -> None:
connection = FakeConnection()
with make_repository(connection).advisory_lock(TARGET_DATE) as acquired:
assert acquired
assert len(connection.statements) == 2
assert "pg_try_advisory_lock" in connection.statements[0][0]
assert "2026-08-28" in str(connection.statements[0][1][0])
assert "pg_advisory_unlock" in connection.statements[1][0]
assert connection.statements[0][1] == connection.statements[1][1]
def test_exact_success_and_latest_attempt_queries_use_distinct_semantics() -> None:
connection = FakeConnection()
repository = make_repository(connection)
exact = repository.get_successful_publication(TARGET_DATE)
latest = repository.get_latest_publication()
assert exact is not None
assert latest is not None
exact_query, exact_parameters = connection.statements[0]
latest_query, latest_parameters = connection.statements[1]
assert "status = 'success' AND target_trade_date = %s" in exact_query
assert exact_parameters == (TARGET_DATE,)
assert "status = 'success'" not in latest_query
assert "started_at DESC" in latest_query
assert latest_parameters == ()
def test_load_rankings_reconstructs_values_and_rank_changes() -> None:
connection = FakeConnection()
rankings = make_repository(connection).load_rankings("publication-a")
assert len(rankings) == 1
ranking = rankings[0]
assert ranking.observation.metric_version == "zhixing_amount_net_bn_v1"
assert ranking.observation.value == Decimal("12.5")
assert ranking.rank_position == 1
assert ranking.rank_change(1) == 3
assert ranking.rank_change(2) is None
query, parameters = connection.statements[0]
assert "WHERE publication_id = %s" in query
assert "rank_position NULLS LAST" in query
assert parameters == ("publication-a",)
@@ -0,0 +1,147 @@
from datetime import date
from decimal import Decimal
from zhixing_server.modules.sector_radar.domain.models import (
MetricKind,
MetricObservation,
MetricQuality,
MetricUnit,
RankSide,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.ranking import (
rank_metric_observations,
select_percentile_side,
select_rank_change_side,
with_rank_changes,
)
TARGET_DATE = date(2026, 8, 28)
def make_observation(
sector_code: str,
sector_type: SectorType,
value: str | None,
*,
trade_date: date = TARGET_DATE,
) -> MetricObservation:
metric_value = Decimal(value) if value is not None else None
return MetricObservation(
trade_date=trade_date,
sector_type=sector_type,
sector_code=sector_code,
sector_name=sector_code,
metric_kind=MetricKind.AMOUNT,
metric_version="zhixing_amount_net_bn_v1",
implementation_kind="independent",
unit=MetricUnit.CNY_100M,
value=metric_value,
quality=(
MetricQuality.AVAILABLE if metric_value is not None else MetricQuality.UNAVAILABLE
),
member_count=10,
valid_sample_count=10 if metric_value is not None else 0,
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
def test_ranking_separates_types_and_uses_code_as_stable_tie_breaker() -> None:
observations = (
make_observation("BK2002.DC", SectorType.INDUSTRY, "20"),
make_observation("BK1002.DC", SectorType.CONCEPT, "30"),
make_observation("BK2001.DC", SectorType.INDUSTRY, "20"),
make_observation("BK1001.DC", SectorType.CONCEPT, "10"),
)
ranked = rank_metric_observations(tuple(reversed(observations)))
by_code = {row.observation.sector_code: row for row in ranked}
assert by_code["BK1002.DC"].rank_position == 1
assert by_code["BK1002.DC"].rank_percentile == Decimal("100")
assert by_code["BK1001.DC"].rank_position == 2
assert by_code["BK1001.DC"].rank_percentile == Decimal("50")
assert by_code["BK2001.DC"].rank_position == 1
assert by_code["BK2002.DC"].rank_position == 2
def test_ranking_handles_empty_and_single_element_pools() -> None:
assert rank_metric_observations(()) == ()
[single] = rank_metric_observations((make_observation("BK0001.DC", SectorType.CONCEPT, "0"),))
assert single.rank_position == 1
assert single.rank_percentile == Decimal("100")
def test_percentile_sides_use_confirmed_inclusive_thresholds() -> None:
ranked = rank_metric_observations(
make_observation(f"BK{position:04d}.DC", SectorType.CONCEPT, str(11 - position))
for position in range(1, 11)
)
top = select_percentile_side(ranked, RankSide.TOP)
bottom = select_percentile_side(ranked, RankSide.BOTTOM)
assert [row.observation.sector_code for row in top] == ["BK0001.DC", "BK0002.DC"]
assert [row.observation.sector_code for row in bottom] == ["BK0010.DC"]
def test_rank_change_is_past_rank_minus_current_and_preserves_missing_history() -> None:
current = rank_metric_observations(
(
make_observation("BK0001.DC", SectorType.CONCEPT, "30"),
make_observation("BK0002.DC", SectorType.CONCEPT, "20"),
)
)
previous = rank_metric_observations(
(
make_observation(
"BK0001.DC",
SectorType.CONCEPT,
"10",
trade_date=date(2026, 8, 27),
),
make_observation(
"BK0002.DC",
SectorType.CONCEPT,
"40",
trade_date=date(2026, 8, 27),
),
)
)
changed = with_rank_changes(current, {1: previous, 5: ()})
by_code = {row.observation.sector_code: row for row in changed}
assert by_code["BK0001.DC"].rank_change(1) == 1
assert by_code["BK0002.DC"].rank_change(1) == -1
assert by_code["BK0001.DC"].rank_change(5) is None
def test_rank_change_sides_take_ceiling_ten_percent_per_pool() -> None:
current = rank_metric_observations(
make_observation(f"BK{position:04d}.DC", SectorType.CONCEPT, str(12 - position))
for position in range(1, 12)
)
previous = rank_metric_observations(
make_observation(
f"BK{position:04d}.DC",
SectorType.CONCEPT,
str(position),
trade_date=date(2026, 8, 27),
)
for position in range(1, 12)
)
changed = with_rank_changes(current, {1: previous})
top = select_rank_change_side(changed, days=1, side=RankSide.TOP)
bottom = select_rank_change_side(changed, days=1, side=RankSide.BOTTOM)
assert [row.observation.sector_code for row in top] == ["BK0001.DC", "BK0002.DC"]
assert [row.observation.sector_code for row in bottom] == [
"BK0011.DC",
"BK0010.DC",
]
@@ -0,0 +1,189 @@
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
from zhixing_server.modules.sector_radar.application.read import (
RadarQuery,
RadarView,
ReadSectorRadar,
)
from zhixing_server.modules.sector_radar.domain.metrics import AmountNetStrategy
from zhixing_server.modules.sector_radar.domain.models import (
MetricKind,
MetricObservation,
MetricQuality,
MetricUnit,
PublicationStatus,
RadarPublication,
RankChange,
RankedMetric,
RankSide,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.persistence import RankingRecord
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
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)
def _running(publication_id: str, trade_date: date) -> RadarPublication:
return RadarPublication(
publication_id=publication_id,
target_trade_date=trade_date,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-v1",
metric_versions=(AmountNetStrategy.metric_version,),
input_hash=None,
coverage=Decimal(0),
started_at=NOW,
)
def _finish(
publication: RadarPublication,
status: PublicationStatus,
) -> RadarPublication:
return replace(
publication,
status=status,
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
coverage=Decimal(1) if status is PublicationStatus.SUCCESS else Decimal("0.8"),
finished_at=publication.started_at + timedelta(minutes=5),
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
)
def _amount_rankings() -> tuple[RankedMetric, ...]:
observations = tuple(
MetricObservation(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code=f"BK{index:04d}.DC",
sector_name=f"概念{index}",
metric_kind=MetricKind.AMOUNT,
metric_version=AmountNetStrategy.metric_version,
implementation_kind="independent",
unit=MetricUnit.CNY_100M,
value=Decimal(11 - index),
quality=MetricQuality.AVAILABLE,
member_count=5,
valid_sample_count=5,
membership_coverage=Decimal(1),
moneyflow_coverage=Decimal(1),
)
for index in range(1, 11)
)
rankings = rank_metric_observations(observations)
return tuple(
replace(
row,
rank_changes=tuple(
RankChange(
days=days,
value=(
None
if row.observation.sector_code == "BK0005.DC" and days == 5
else (row.rank_position or 0) - 5
),
)
for days in range(1, 6)
),
)
for row in rankings
)
def _published_repository() -> InMemorySectorRadarRepository:
repository = InMemorySectorRadarRepository()
publication = _running("publication-success", TARGET_DATE)
repository.create_publication(publication)
repository.finish_publication(_finish(publication, PublicationStatus.SUCCESS))
repository.save_rankings(
RankingRecord(publication.publication_id, ranking) for ranking in _amount_rankings()
)
return repository
def test_no_successful_publication_returns_stable_no_data() -> None:
reader = ReadSectorRadar(InMemorySectorRadarRepository())
dates = reader.list_dates()
rankings = reader.query(RadarQuery())
assert dates.status == "no_data"
assert dates.available_dates == ()
assert rankings.status == "no_data"
assert rankings.publication is None
assert rankings.total == 0
assert rankings.definition.metric_version == AmountNetStrategy.metric_version
def test_explicit_date_never_falls_back_to_an_earlier_last_good() -> None:
reader = ReadSectorRadar(_published_repository())
missing = reader.query(RadarQuery(trade_date=TARGET_DATE + timedelta(days=1)))
assert missing.status == "no_data"
assert missing.publication is None
def test_percentile_side_is_selected_before_search_and_pagination() -> None:
reader = ReadSectorRadar(_published_repository())
top = reader.query(RadarQuery(side=RankSide.TOP, page_size=1))
second_page = reader.query(RadarQuery(side=RankSide.TOP, page=2, page_size=1))
searched = reader.query(RadarQuery(side=RankSide.TOP, search="概念2"))
bottom = reader.query(RadarQuery(side=RankSide.BOTTOM))
assert top.total == 2
assert top.rows[0].observation.sector_code == "BK0001.DC"
assert second_page.rows[0].observation.sector_code == "BK0002.DC"
assert searched.total == 1
assert searched.rows[0].observation.sector_name == "概念2"
assert bottom.total == 1
assert bottom.rows[0].observation.sector_code == "BK0010.DC"
def test_rank_change_uses_selected_metric_days_and_pool_sides() -> None:
reader = ReadSectorRadar(_published_repository())
query = RadarQuery(
view=RadarView.RANK_CHANGE,
rank_change_metric=MetricKind.AMOUNT,
rank_change_days=5,
)
top = reader.query(replace(query, side=RankSide.TOP))
bottom = reader.query(replace(query, side=RankSide.BOTTOM))
all_rows = reader.query(query)
assert top.total == 1
assert top.rows[0].rank_change(5) == 5
assert bottom.total == 1
assert bottom.rows[0].rank_change(5) == -4
assert all_rows.total == 10
assert all_rows.rows[-1].observation.sector_code == "BK0005.DC"
assert all_rows.rows[-1].rank_change(5) is None
def test_latest_partial_attempt_is_visible_but_does_not_replace_last_good() -> None:
repository = _published_repository()
partial = replace(
_running("publication-partial", TARGET_DATE + timedelta(days=1)),
started_at=NOW + timedelta(days=1),
)
repository.create_publication(partial)
repository.finish_publication(_finish(partial, PublicationStatus.PARTIAL))
index = ReadSectorRadar(repository).list_dates()
assert index.status == "success"
assert index.current_attempt is not None
assert index.current_attempt.status is PublicationStatus.PARTIAL
assert index.last_good is not None
assert index.last_good.publication_id == "publication-success"
assert index.available_dates == (TARGET_DATE,)
@@ -0,0 +1,121 @@
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.domain.models import (
PublicationStatus,
RadarPublication,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.persistence import MembershipRecord
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
from zhixing_server.modules.sector_radar.infrastructure.memory import (
InMemorySectorRadarRepository,
)
TARGET_DATE = date(2026, 8, 28)
STARTED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
def make_running(publication_id: str, target_trade_date: date = TARGET_DATE) -> RadarPublication:
return RadarPublication(
publication_id=publication_id,
target_trade_date=target_trade_date,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-v1",
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal(0),
started_at=STARTED_AT,
)
def finish(
publication: RadarPublication,
status: PublicationStatus,
*,
offset_minutes: int = 5,
) -> RadarPublication:
return replace(
publication,
status=status,
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
coverage=Decimal("1") if status is PublicationStatus.SUCCESS else Decimal("0.8"),
finished_at=publication.started_at + timedelta(minutes=offset_minutes),
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
)
def test_source_and_membership_revisions_are_idempotent_but_not_overwritable() -> None:
repository = InMemorySectorRadarRepository()
snapshot = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20260828"},
rows=(
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "平安银行",
},
),
target_trade_date=TARGET_DATE,
observed_at=STARTED_AT,
)
member = MembershipRecord(
source_snapshot_id=snapshot.snapshot_id,
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
stock_code="000001.SZ",
stock_name="平安银行",
)
assert repository.save_source_snapshots((snapshot,)).inserted == 1
assert repository.save_source_snapshots((snapshot,)).unchanged == 1
assert repository.save_memberships((member,)).inserted == 1
assert repository.save_memberships((member,)).unchanged == 1
with pytest.raises(ValueError, match="cannot change content"):
repository.save_memberships((replace(member, stock_name="已改变"),))
def test_partial_and_failed_revisions_never_replace_last_good() -> None:
repository = InMemorySectorRadarRepository()
successful = make_running("success-a")
partial = make_running("partial-b")
failed = make_running("failed-c", TARGET_DATE + timedelta(days=1))
repository.create_publication(successful)
repository.finish_publication(finish(successful, PublicationStatus.SUCCESS))
repository.create_publication(partial)
repository.finish_publication(finish(partial, PublicationStatus.PARTIAL, offset_minutes=6))
repository.create_publication(failed)
repository.finish_publication(finish(failed, PublicationStatus.FAILED, offset_minutes=7))
last_good = repository.get_last_good_publication()
assert last_good is not None
assert last_good.publication_id == "success-a"
assert repository.list_successful_dates() == (TARGET_DATE,)
def test_publication_identity_allows_sequential_same_date_revisions() -> None:
repository = InMemorySectorRadarRepository()
first = make_running("revision-a")
second = make_running("revision-b")
assert repository.create_publication(first).inserted == 1
with pytest.raises(ValueError, match="already has a running"):
repository.create_publication(second)
with pytest.raises(ValueError, match="terminal"):
repository.finish_publication(first)
repository.finish_publication(finish(first, PublicationStatus.FAILED))
assert repository.create_publication(second).inserted == 1
with pytest.raises(ValueError, match="running status"):
repository.finish_publication(finish(first, PublicationStatus.SUCCESS))
@@ -0,0 +1,669 @@
import logging
import threading
from collections.abc import Mapping
from datetime import UTC, date, datetime
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.domain.models import SectorType
from zhixing_server.modules.sector_radar.domain.source import (
CapabilityStatus,
SourceContractError,
build_source_snapshot,
)
from zhixing_server.modules.sector_radar.infrastructure import tushare as source_module
from zhixing_server.modules.sector_radar.infrastructure.tushare import (
TushareSectorRadarAdapter,
)
TARGET_DATE = date(2026, 8, 28)
OBSERVED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
class QueryClient:
def __init__(self, responses: Mapping[tuple[str, str], object]) -> None:
self.responses = dict(responses)
self.calls: list[tuple[str, dict[str, object]]] = []
self._lock = threading.Lock()
def query(self, api_name: str, **kwargs: object) -> object:
partition = str(kwargs.get("ts_code") or kwargs.get("list_status") or "")
with self._lock:
self.calls.append((api_name, kwargs))
response = self.responses.get((api_name, partition), ())
if isinstance(response, BaseException):
raise response
return response
def make_adapter(client: object) -> TushareSectorRadarAdapter:
return TushareSectorRadarAdapter(
client,
max_retries=0,
request_interval_seconds=0,
sleep_fn=lambda _: None,
now_fn=lambda: OBSERVED_AT,
)
def moneyflow_record(
ts_code: str,
*,
trade_date: str = "20260828",
) -> dict[str, object]:
return {
"trade_date": trade_date,
"ts_code": ts_code,
"name": ts_code,
"net_amount": "1",
"net_amount_rate": "0.1",
"pct_change": "1",
"close": "10",
}
def test_daily_and_moneyflow_keep_source_units_and_distinguish_missing_from_zero() -> None:
client = QueryClient(
{
(
"daily",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"close": "10",
"pre_close": "9.5",
"pct_chg": "1.5",
"vol": "100",
"amount": "12.5",
},
{
"ts_code": "000002.SZ",
"trade_date": "20260828",
"close": "20",
"pre_close": "20",
"pct_chg": "0",
"vol": "0",
"amount": float("nan"),
},
),
(
"moneyflow_dc",
"",
): (
{
"trade_date": "20260828",
"ts_code": "000001.SZ",
"name": "平安银行",
"net_amount": "2.5",
"net_amount_rate": "0.2",
"pct_change": "1.5",
"close": "10",
},
{
"trade_date": "20260828",
"ts_code": "000002.SZ",
"name": "示例股票",
"net_amount": "0",
"net_amount_rate": "0",
"pct_change": "0",
"close": "20",
},
),
}
)
adapter = make_adapter(client)
daily = adapter.fetch_daily(TARGET_DATE)
moneyflow = adapter.fetch_moneyflow_dc(TARGET_DATE, ("000001.SZ", "000002.SZ"))
assert daily.rows[0].amount_thousand_yuan == Decimal("12.5")
assert daily.rows[0].turnover_yuan == Decimal("12500.0")
assert daily.rows[1].amount_thousand_yuan is None
assert moneyflow.rows[0].net_amount_ten_thousand_yuan == Decimal("2.5")
assert moneyflow.rows[0].net_amount_yuan == Decimal("25000.0")
assert moneyflow.rows[1].net_amount_yuan == Decimal("0")
assert client.calls[0][1]["fields"] == ",".join(source_module.FIELDS["daily"])
def test_moneyflow_accepts_a_full_initial_snapshot_at_the_provider_limit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", 2)
client = QueryClient(
{
("moneyflow_dc", ""): (
moneyflow_record("000001.SZ"),
moneyflow_record("000002.SZ"),
)
}
)
result = make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
("000001.SZ", "000002.SZ"),
)
assert result.snapshots[0].limit_reached is True
assert [row.ts_code for row in result.rows] == ["000001.SZ", "000002.SZ"]
assert len(client.calls) == 1
@pytest.mark.parametrize(
("initial_rows", "message"),
(
((moneyflow_record("000001.SZ", trade_date="20260827"),), "trade_date"),
(
(moneyflow_record("000001.SZ"), moneyflow_record("000001.SZ")),
"duplicate business keys",
),
),
)
def test_moneyflow_initial_contract_errors_fail_closed(
initial_rows: tuple[dict[str, object], ...],
message: str,
) -> None:
client = QueryClient({("moneyflow_dc", ""): initial_rows})
with pytest.raises(SourceContractError, match=message):
make_adapter(client).fetch_moneyflow_dc(TARGET_DATE, ())
def test_moneyflow_refills_only_missing_codes_in_stable_snapshot_order(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", 3)
third_finished = threading.Event()
completion_order: list[str] = []
completion_lock = threading.Lock()
class ReverseCompletionClient(QueryClient):
def query(self, api_name: str, **kwargs: object) -> object:
response = super().query(api_name, **kwargs)
ts_code = str(kwargs.get("ts_code") or "")
if ts_code == "000004.SZ":
if not third_finished.wait(timeout=2):
raise AssertionError("second moneyflow worker did not start")
elif ts_code == "000005.SZ":
third_finished.set()
if ts_code:
with completion_lock:
completion_order.append(ts_code)
return response
client = ReverseCompletionClient(
{
("moneyflow_dc", ""): tuple(
moneyflow_record(f"00000{index}.SZ") for index in range(1, 4)
),
("moneyflow_dc", "000004.SZ"): (moneyflow_record("000004.SZ"),),
("moneyflow_dc", "000005.SZ"): (moneyflow_record("000005.SZ"),),
}
)
result = make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
tuple(f"00000{index}.SZ" for index in range(1, 6)),
)
assert completion_order == ["000005.SZ", "000004.SZ"]
assert [snapshot.partition_key for snapshot in result.snapshots] == [
"all",
"000004.SZ",
"000005.SZ",
]
assert [row.ts_code for row in result.rows] == [
"000001.SZ",
"000002.SZ",
"000003.SZ",
"000004.SZ",
"000005.SZ",
]
assert len(client.calls) == 3
def test_moneyflow_empty_and_exhausted_refills_remain_real_gaps(
caplog: pytest.LogCaptureFixture,
) -> None:
client = QueryClient(
{
("moneyflow_dc", ""): (moneyflow_record("000001.SZ"),),
("moneyflow_dc", "000002.SZ"): (),
("moneyflow_dc", "000003.SZ"): RuntimeError("private provider payload"),
}
)
caplog.set_level(
logging.WARNING,
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
)
result = make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
("000001.SZ", "000002.SZ", "000003.SZ"),
)
assert [row.ts_code for row in result.rows] == ["000001.SZ"]
assert [snapshot.partition_key for snapshot in result.snapshots] == ["all"]
messages = "\n".join(record.getMessage() for record in caplog.records)
assert "partition_empty partition_key=000002.SZ" in messages
assert "partition_failed partition_key=000003.SZ" in messages
assert "private provider payload" not in messages
@pytest.mark.parametrize(
("partition_rows", "row_limit", "message"),
(
((moneyflow_record("000002.SZ", trade_date="20260827"),), 6_000, "trade_date"),
((moneyflow_record("000099.SZ"),), 6_000, "different ts_code"),
(
(moneyflow_record("000002.SZ"), moneyflow_record("000002.SZ")),
6_000,
"duplicate business keys",
),
(
(moneyflow_record("000002.SZ"), moneyflow_record("000002.SZ")),
2,
"provider row limit",
),
),
)
def test_moneyflow_partition_contract_errors_fail_closed(
monkeypatch: pytest.MonkeyPatch,
partition_rows: tuple[dict[str, object], ...],
row_limit: int,
message: str,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", row_limit)
client = QueryClient(
{
("moneyflow_dc", ""): (moneyflow_record("000001.SZ"),),
("moneyflow_dc", "000002.SZ"): partition_rows,
}
)
with pytest.raises(SourceContractError, match=message):
make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
("000001.SZ", "000002.SZ"),
)
def test_non_finite_source_values_are_rejected() -> None:
client = QueryClient(
{
(
"daily",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"close": "Infinity",
"pre_close": "9.5",
"pct_chg": "1.5",
"vol": "100",
"amount": "12.5",
},
)
}
)
with pytest.raises(SourceContractError, match="finite"):
make_adapter(client).fetch_daily(TARGET_DATE)
def test_contract_failure_log_identifies_member_partition_without_payload(
caplog: pytest.LogCaptureFixture,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "dc_member", 2)
client = QueryClient(
{
(
"dc_member",
"",
): (
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "private-payload-marker",
},
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000002.SZ",
"name": "private-payload-marker",
},
),
(
"dc_member",
"BK0001.DC",
): (
{
"trade_date": "20260828",
"ts_code": "BK9999.DC",
"con_code": "000001.SZ",
"name": "private-payload-marker",
},
),
}
)
caplog.set_level(
logging.ERROR,
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
)
with pytest.raises(SourceContractError, match="different sector"):
make_adapter(client).fetch_sector_members(TARGET_DATE, ("BK0001.DC",))
messages = "\n".join(record.getMessage() for record in caplog.records)
assert "sector_radar_source_contract_failed" in messages
assert "api_name=dc_member" in messages
assert "partition_key=BK0001.DC" in messages
assert "validation=dc_member partition returned a different sector" in messages
assert "private-payload-marker" not in messages
assert len(caplog.records) == 1
def test_merged_member_contract_failure_has_one_interface_level_log(
caplog: pytest.LogCaptureFixture,
) -> None:
duplicate = {
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "private-payload-marker",
}
client = QueryClient({("dc_member", ""): (duplicate, duplicate)})
caplog.set_level(
logging.ERROR,
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
)
with pytest.raises(SourceContractError, match="duplicate business keys"):
make_adapter(client).fetch_sector_members(TARGET_DATE, ("BK0001.DC",))
messages = "\n".join(record.getMessage() for record in caplog.records)
assert "api_name=dc_member" in messages
assert "partition_key=merged" in messages
assert "validation=dc_member returned duplicate business keys" in messages
assert "private-payload-marker" not in messages
assert len(caplog.records) == 1
def test_dc_member_reloads_by_sector_when_the_all_market_call_hits_limit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "dc_member", 2)
client = QueryClient(
{
(
"dc_member",
"",
): (
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "A",
},
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000002.SZ",
"name": "B",
},
),
(
"dc_member",
"BK0001.DC",
): (
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "A",
},
),
(
"dc_member",
"BK0002.DC",
): (
{
"trade_date": "20260828",
"ts_code": "BK0002.DC",
"con_code": "600000.SH",
"name": "C",
},
),
}
)
result = make_adapter(client).fetch_sector_members(
TARGET_DATE,
("BK0001.DC", "BK0002.DC"),
)
assert [row.stock_code for row in result.rows] == ["000001.SZ", "600000.SH"]
assert [snapshot.partition_key for snapshot in result.snapshots] == [
"all",
"BK0001.DC",
"BK0002.DC",
]
def test_dc_member_preserves_an_explicit_empty_partition() -> None:
client = QueryClient(
{
(
"dc_member",
"",
): (
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "A",
},
),
("dc_member", "BK0002.DC"): (),
}
)
result = make_adapter(client).fetch_sector_members(
TARGET_DATE,
("BK0001.DC", "BK0002.DC"),
)
assert [row.sector_code for row in result.rows] == ["BK0001.DC"]
assert [snapshot.partition_key for snapshot in result.snapshots] == [
"all",
"BK0002.DC",
]
assert result.snapshots[1].row_count == 0
def test_stock_basic_requests_only_current_listings() -> None:
client = QueryClient(
{
(
"stock_basic",
"L",
): (
{
"ts_code": "000001.SZ",
"symbol": "000001",
"name": "L",
"market": "主板",
"exchange": "SZSE",
"list_status": "L",
"list_date": "20200101",
"delist_date": None,
},
)
}
)
result = make_adapter(client).fetch_stock_basics()
assert {row.list_status for row in result.rows} == {"L"}
assert [snapshot.partition_key for snapshot in result.snapshots] == ["L"]
assert [call[1]["list_status"] for call in client.calls] == ["L"]
def test_stock_basic_rejects_a_non_listed_row_from_the_l_partition() -> None:
client = QueryClient(
{
("stock_basic", "L"): (
{
"ts_code": "000001.SZ",
"symbol": "000001",
"name": "unexpected",
"market": "主板",
"exchange": "SZSE",
"list_status": "D",
"list_date": "20200101",
"delist_date": "20260828",
},
)
}
)
with pytest.raises(SourceContractError, match="unexpected list_status"):
make_adapter(client).fetch_stock_basics()
def test_suspend_timing_may_be_missing_while_suspend_type_remains_required() -> None:
client = QueryClient(
{
(
"suspend_d",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"suspend_timing": None,
"suspend_type": "S",
},
)
}
)
result = make_adapter(client).fetch_suspensions(TARGET_DATE)
assert result.rows[0].suspend_timing is None
assert result.rows[0].suspend_type == "S"
missing_type = QueryClient(
{
(
"suspend_d",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"suspend_timing": None,
"suspend_type": None,
},
)
}
)
with pytest.raises(SourceContractError, match="suspend_type must be a non-empty string"):
make_adapter(missing_type).fetch_suspensions(TARGET_DATE)
def test_source_snapshot_hash_is_order_stable_and_excludes_token_params() -> None:
first = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828", "token": "secret"},
rows=({"ts_code": "2"}, {"ts_code": "1"}),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
)
second = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "1"}, {"ts_code": "2"}),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
)
assert first.snapshot_id == second.snapshot_id
assert "secret" not in repr(first)
def test_source_snapshot_identity_includes_schema_and_limit_metadata() -> None:
first = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "000001.SZ"},),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
returned_fields=("ts_code",),
row_limit=1,
)
changed_schema = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "000001.SZ"},),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
returned_fields=("name", "ts_code"),
row_limit=1,
)
changed_limit = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "000001.SZ"},),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
returned_fields=("ts_code",),
row_limit=2,
)
assert first.content_sha256 != changed_schema.content_sha256
assert first.snapshot_id != changed_schema.snapshot_id
assert first.content_sha256 != changed_limit.content_sha256
assert first.snapshot_id != changed_limit.snapshot_id
def test_capability_probe_classifies_errors_without_exposing_provider_text() -> None:
client = QueryClient({("daily", ""): RuntimeError("权限不足 private-detail")})
probe = make_adapter(client).probe(TARGET_DATE)
by_name = {result.api_name: result for result in probe.interfaces}
assert by_name["daily"].status is CapabilityStatus.FORBIDDEN
assert "private-detail" not in repr(probe)
assert len(probe.interfaces) == 7
def test_sector_index_uses_independent_concept_and_industry_params() -> None:
client = QueryClient(
{
(
"dc_index",
"",
): (
{
"ts_code": "BK0001.DC",
"trade_date": "20260828",
"name": "示例",
"idx_type": "概念板块",
"level": "一级",
"pct_change": "1",
"leading_code": "000001.SZ",
},
)
}
)
result = make_adapter(client).fetch_sector_indices(TARGET_DATE, SectorType.CONCEPT)
assert result.rows[0].sector_type is SectorType.CONCEPT
assert client.calls[0][1]["idx_type"] == "概念板块"