Files

429 lines
17 KiB
Python

"""Regression coverage for publication-scoped, bounded radar read paths."""
import os
from collections.abc import Generator, Sequence
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
from pathlib import Path
from unittest.mock import patch
from uuid import uuid4
import psycopg
import pytest
from alembic import command
from alembic.config import Config
from fastapi.testclient import TestClient
from psycopg.conninfo import make_conninfo
from psycopg.sql import SQL, Identifier
from zhixing_server.bootstrap.app import create_app
from zhixing_server.bootstrap.config import Settings, sqlalchemy_database_url
from zhixing_server.modules.sector_radar.application.details import ReadRadarDetails, metric_at
from zhixing_server.modules.sector_radar.application.read import RadarQuery, 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,
RankSide,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.persistence import (
PublicationSourceGroup,
PublicationSourceRecord,
RankingRecord,
SectorRadarRepository,
)
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
from zhixing_server.modules.sector_radar.domain.source import SourceScalar, build_source_snapshot
from zhixing_server.modules.sector_radar.infrastructure.memory import InMemorySectorRadarRepository
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
PostgresSectorRadarRepository,
)
from zhixing_server.modules.sector_radar.presentation.http import get_sector_radar_reader
TARGET = date(2026, 9, 4)
NOW = datetime(2026, 9, 4, 18, tzinfo=UTC)
CODE = "BK1147.DC"
def seed_detail(repository: SectorRadarRepository) -> None:
"""Seed raw-only legacy publications with missing facts and unrelated rows."""
for offset in (-1, 0):
day = TARGET + timedelta(days=offset)
publication = RadarPublication(
publication_id=f"detail-{offset}",
target_trade_date=day,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="test-v1",
metric_versions=(AmountNetStrategy.metric_version,),
input_hash=None,
coverage=Decimal(0),
started_at=NOW,
)
repository.create_publication(publication)
observations = tuple(
MetricObservation(
trade_date=day,
sector_type=SectorType.CONCEPT,
sector_code=code,
sector_name=code,
metric_kind=MetricKind.AMOUNT,
metric_version=AmountNetStrategy.metric_version,
implementation_kind="independent",
unit=MetricUnit.CNY_100M,
value=Decimal(20 - index),
quality=MetricQuality.AVAILABLE,
member_count=2,
valid_sample_count=2,
membership_coverage=Decimal(1),
moneyflow_coverage=Decimal(1),
)
for index, code in enumerate((CODE, "OTHER") if offset == 0 else ("OTHER",))
)
repository.save_rankings(
RankingRecord(publication.publication_id, row)
for row in rank_metric_observations(observations)
)
repository.finish_publication(
replace(
publication,
status=PublicationStatus.SUCCESS,
input_hash="a" * 64,
coverage=Decimal(1),
finished_at=NOW + timedelta(seconds=1),
)
)
groups: dict[PublicationSourceGroup, tuple[dict[str, SourceScalar], ...]] = {
PublicationSourceGroup.CALENDAR: tuple(
{"exchange": "SSE", "cal_date": (TARGET + timedelta(days=i)).isoformat(), "is_open": 1}
for i in (-2, -1, 0, 1)
),
PublicationSourceGroup.CONCEPT_INDICES: tuple(
{
"ts_code": code,
"name": code,
"trade_date": TARGET.isoformat(),
"pct_change": "1.25",
"leading_code": "000003.SZ",
}
for code in (CODE, "OTHER")
),
PublicationSourceGroup.STOCK_BASICS: tuple(
{
"ts_code": code,
"symbol": code[:6],
"name": code,
"exchange": "SZSE",
"list_status": status,
"list_date": "20200101",
}
for code, status in (
("000001.SZ", "L"),
("000002.SZ", "L"),
("000003.SZ", "L"),
("000004.SZ", "D"),
)
),
PublicationSourceGroup.MEMBERS: tuple(
{"ts_code": sector, "con_code": code, "name": code, "trade_date": "20260904"}
for sector, code in (
(CODE, "000001.SZ"),
(CODE, "000002.SZ"),
(CODE, "000004.SZ"),
("OTHER", "000002.SZ"),
("OTHER", "000003.SZ"),
)
),
PublicationSourceGroup.DAILY: (
{"ts_code": "000001.SZ", "trade_date": "20260904", "pct_chg": "2.50"},
{"ts_code": "000003.SZ", "trade_date": "20260904", "pct_chg": "5"},
{"ts_code": "000001.SZ", "trade_date": "20260903", "pct_chg": "99"},
),
PublicationSourceGroup.MONEYFLOW_DC: (
{
"ts_code": "000002.SZ",
"trade_date": "2026-09-04",
"net_amount": "1.2345",
"name": "000002.SZ",
},
),
PublicationSourceGroup.MONEYFLOW: (
{
"ts_code": "000001.SZ",
"trade_date": "20260904",
"net_mf_amount": "11",
"buy_lg_amount": "10",
"sell_lg_amount": "2",
"buy_elg_amount": "4",
"sell_elg_amount": "1",
},
),
PublicationSourceGroup.SUSPENSIONS: ({"unused": "must not be read"},),
}
for group, rows in groups.items():
snapshot = build_source_snapshot(
api_name=group.value,
params={},
rows=rows,
target_trade_date=TARGET,
observed_at=NOW,
)
repository.save_source_snapshots((snapshot,))
repository.save_publication_sources(
(PublicationSourceRecord("detail-0", group, 0, snapshot),)
)
class TrackingRepository(InMemorySectorRadarRepository):
def __init__(self) -> None:
super().__init__()
self.reads: list[
tuple[tuple[PublicationSourceGroup, ...], date | None, tuple[str, ...] | None]
] = []
def load_publication_sources(self, publication_id: str) -> Sequence[PublicationSourceRecord]:
raise AssertionError("HTTP read must not load complete audit snapshots")
def load_publication_rows(
self,
publication_id: str,
source_groups: Sequence[PublicationSourceGroup],
*,
trade_date: date | None = None,
ts_codes: Sequence[str] | None = None,
) -> dict[PublicationSourceGroup, tuple[dict[str, SourceScalar], ...]]:
self.reads.append(
(tuple(source_groups), trade_date, None if ts_codes is None else tuple(ts_codes))
)
return super().load_publication_rows(
publication_id,
source_groups,
trade_date=trade_date,
ts_codes=ts_codes,
)
def test_detail_http_reads_only_required_rows_and_preserves_independent_missing_values() -> None:
repository = TrackingRepository()
seed_detail(repository)
app = create_app()
app.dependency_overrides[get_sector_radar_reader] = lambda: ReadSectorRadar(repository)
with TestClient(app) as client:
response = client.get(
f"/api/v1/sector-radar/sectors/concept/{CODE}/detail?trade_date={TARGET}"
)
assert response.status_code == 200
body = response.json()
assert body["members"] == [
{
"ts_code": "000001.SZ",
"name": "000001.SZ",
"pct_change": "2.50",
"net_amount_yuan": None,
"active_buy_net_amount_yuan": "110000",
},
{
"ts_code": "000002.SZ",
"name": "000002.SZ",
"pct_change": None,
"net_amount_yuan": "12345.0000",
"active_buy_net_amount_yuan": None,
},
]
assert body["similar_sectors"][0]["intersection_count"] == 1
assert body["similar_sectors"][0]["union_count"] == 3
assert body["history"]["available_days"] == 1
assert [point["amount"]["pool_size"] for point in body["history"]["points"]] == [0, 1, 2]
assert len(repository.reads) == 3
assert repository.reads[0][0] == (PublicationSourceGroup.CALENDAR,)
assert repository.reads[-1][1:] == (TARGET, ("000001.SZ", "000002.SZ"))
assert all(PublicationSourceGroup.SUSPENSIONS not in read[0] for read in repository.reads)
def test_history_and_ranking_enrichment_do_not_read_members_or_stock_facts() -> None:
repository = TrackingRepository()
seed_detail(repository)
history = ReadRadarDetails(repository).history(TARGET, SectorType.CONCEPT, CODE)
assert history.status == "success"
assert [read[0] for read in repository.reads] == [(PublicationSourceGroup.CALENDAR,)]
repository.reads.clear()
page = ReadSectorRadar(repository).query(
RadarQuery(trade_date=TARGET, page_size=1, side=RankSide.TOP)
)
assert page.extras[CODE].on_list_count == 1
assert page.extras[CODE].pct_change == Decimal("1.25")
assert len(repository.reads) == 2
assert all(
set(read[0])
<= {
PublicationSourceGroup.CALENDAR,
PublicationSourceGroup.CONCEPT_INDICES,
PublicationSourceGroup.INDUSTRY_INDICES,
}
for read in repository.reads
)
def test_filtered_history_matches_full_pool_metrics_even_when_sector_is_absent() -> None:
repository = InMemorySectorRadarRepository()
seed_detail(repository)
records = repository.load_ranked_history(("detail--1", "detail-0"), (CODE,))
assert {record.ranking.observation.sector_code for record in records if record.ranking} == {
CODE
}
assert any(record.ranking is None and record.pool_size == 1 for record in records)
history = ReadRadarDetails(repository).history(TARGET, SectorType.CONCEPT, CODE)
for point in history.points:
rows = repository.load_rankings(point.publication_id) if point.publication_id else ()
assert point.amount == metric_at(rows, SectorType.CONCEPT, CODE, MetricKind.AMOUNT)
def test_history_keeps_selected_publication_when_rebuild_finishes_during_read() -> None:
repository = InMemorySectorRadarRepository()
seed_detail(repository)
selected = repository.get_successful_publication(TARGET)
assert selected is not None
replacement = replace(selected, publication_id="newer-revision")
with patch.object(repository, "load_history_publications", return_value=(replacement,)):
history = ReadRadarDetails(repository).history(TARGET, SectorType.CONCEPT, CODE)
assert history.publication == selected
assert history.points[-1].publication_id == selected.publication_id
assert history.points[-1].amount.pool_size == 2
def test_history_does_not_mix_source_versions_or_metric_versions() -> None:
repository = InMemorySectorRadarRepository()
seed_detail(repository)
prior = repository.publications["detail--1"]
repository.publications[prior.publication_id] = replace(prior, source_version="incompatible")
current = repository.load_rankings("detail-0")[0]
repository.save_rankings(
(
RankingRecord(
"detail-0",
replace(
current,
observation=replace(
current.observation, metric_version="old-amount", value=Decimal("999")
),
),
),
)
)
history = ReadRadarDetails(repository).history(TARGET, SectorType.CONCEPT, CODE)
assert history.points[-2].amount.pool_size == 0
assert history.points[-2].amount.missing
assert history.points[-1].amount.pool_size == 2
assert history.points[-1].amount.metric_value == Decimal(20)
@pytest.fixture
def postgres_read_repository() -> Generator[PostgresSectorRadarRepository]:
"""Use an isolated schema so read regression tests never overwrite other fixtures."""
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
if not database_url:
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
schema = "radar_read_" + uuid4().hex
with psycopg.connect(database_url, autocommit=True) as connection:
connection.execute(SQL("CREATE SCHEMA {}").format(Identifier(schema)))
isolated_url = (
database_url + ("&" if "?" in database_url else "?") + f"options=-csearch_path%3D{schema}"
)
repository = PostgresSectorRadarRepository(
make_conninfo(database_url, options=f"-c search_path={schema}")
)
try:
config = Config(str(Path(__file__).parents[3] / "alembic.ini"))
config.set_main_option(
"sqlalchemy.url", sqlalchemy_database_url(isolated_url).replace("%", "%%")
)
config.config_file_name = None
with patch(
"zhixing_server.bootstrap.config.get_settings",
return_value=Settings(database_url=isolated_url),
):
command.upgrade(config, "head")
yield repository
finally:
repository.close()
with psycopg.connect(database_url, autocommit=True) as connection:
connection.execute(SQL("DROP SCHEMA {} CASCADE").format(Identifier(schema)))
@pytest.mark.integration
def test_postgres_bounded_reads_match_legacy_raw_only_data(
postgres_read_repository: PostgresSectorRadarRepository,
) -> None:
repository = postgres_read_repository
memory = InMemorySectorRadarRepository()
for repo in (repository, memory):
seed_detail(repo)
ranked = repo.load_rankings("detail-0")[0]
repo.save_rankings(
(
RankingRecord(
"detail-0",
replace(
ranked,
rank_position=None,
rank_percentile=None,
observation=replace(
ranked.observation,
sector_code="MISSING",
value=None,
quality=MetricQuality.UNAVAILABLE,
),
),
),
RankingRecord(
"detail-0",
replace(
ranked,
observation=replace(ranked.observation, sector_type=SectorType.INDUSTRY),
),
),
)
)
expected = ReadRadarDetails(memory).detail(TARGET, SectorType.CONCEPT, CODE)
actual = ReadRadarDetails(repository).detail(TARGET, SectorType.CONCEPT, CODE)
assert actual == expected
assert actual.summary["amount"].pool_size == 2
assert repository.load_ranked_history(
("detail--1", "detail-0"), (CODE,)
) == memory.load_ranked_history(("detail--1", "detail-0"), (CODE,))
groups = (PublicationSourceGroup.DAILY, PublicationSourceGroup.MONEYFLOW_DC)
assert repository.load_publication_rows(
"detail-0", groups, trade_date=TARGET, ts_codes=("000002.SZ",)
) == memory.load_publication_rows(
"detail-0", groups, trade_date=TARGET, ts_codes=("000002.SZ",)
)
assert repository.load_publication_rows("detail-0", groups, ts_codes=()) == {}
assert repository.load_publication_rows("detail-0", ()) == {}
assert repository.load_ranked_history((), (CODE,)) == ()
assert repository.load_ranked_history(("detail-0",), ()) == ()
assert repository.load_publication_rows("unknown-publication", groups) == {}
# Later snapshot rows win over earlier ones even for the same stock/date.
replacement = build_source_snapshot(
api_name="daily",
params={"retry": "1"},
rows=({"ts_code": "000001.SZ", "trade_date": "20260904", "pct_chg": "8.5"},),
target_trade_date=TARGET,
observed_at=NOW,
)
repository.save_source_snapshots((replacement,))
repository.save_publication_sources(
(PublicationSourceRecord("detail-0", PublicationSourceGroup.DAILY, 1, replacement),)
)
rows = repository.load_publication_rows(
"detail-0", (PublicationSourceGroup.DAILY,), trade_date=TARGET, ts_codes=("000001.SZ",)
)
assert [row["pct_chg"] for row in rows[PublicationSourceGroup.DAILY]] == ["2.50", "8.5"]