feat(sector-radar): 接入Tushare事实与版本化存储

This commit is contained in:
yuxuanhui
2026-08-29 17:22:39 +08:00
parent 3789008ea6
commit 284c480a90
22 changed files with 3243 additions and 180 deletions
@@ -50,7 +50,7 @@ def test_point_in_time_aggregation_distinguishes_suspension_missing_and_zero() -
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000004.SZ",
status=StockFactStatus.MISSING,
status=StockFactStatus.MISSING_MONEYFLOW,
),
)
@@ -0,0 +1,116 @@
from datetime import UTC, date, datetime
from decimal import Decimal
from zhixing_server.modules.sector_radar.domain.models import StockFactStatus
from zhixing_server.modules.sector_radar.domain.normalize import normalize_stock_facts
from zhixing_server.modules.sector_radar.domain.source import (
DailyRow,
MoneyflowDcRow,
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)) -> StockBasicRow:
return StockBasicRow(
ts_code=ts_code,
symbol=ts_code.split(".")[0],
name=ts_code,
market="主板",
exchange="SZSE",
list_status="L",
list_date=list_date,
delist_date=None,
)
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_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="09:30",
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,102 @@
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) -> None:
self.row = row
def fetchone(self) -> tuple[object, ...] | None:
return self.row
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_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]
@@ -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,277 @@
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]]] = []
def query(self, api_name: str, **kwargs: object) -> object:
self.calls.append((api_name, kwargs))
partition = str(kwargs.get("ts_code") or kwargs.get("list_status") or "")
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 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)
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_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_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_stock_basic_explicitly_requests_all_lifecycle_statuses() -> None:
responses = {
(
"stock_basic",
status,
): (
{
"ts_code": f"00000{index}.SZ",
"symbol": f"00000{index}",
"name": status,
"market": "主板",
"exchange": "SZSE",
"list_status": status,
"list_date": "20200101",
"delist_date": None,
},
)
for index, status in enumerate(("L", "D", "P", "G", "UN"), start=1)
}
client = QueryClient(responses)
result = make_adapter(client).fetch_stock_basics()
assert {row.list_status for row in result.rows} == {"L", "D", "P", "G", "UN"}
assert [call[1]["list_status"] for call in client.calls] == ["L", "D", "P", "G", "UN"]
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_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"] == "概念板块"