feat(market-data): 优化同步并增加完整性检查

This commit is contained in:
yuxuanhui
2026-08-11 11:20:13 +08:00
parent 3ce186f977
commit 7ce11543af
42 changed files with 4837 additions and 197 deletions
@@ -5,7 +5,11 @@ from pathlib import Path
import pytest
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import CsvSnapshotStore
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import (
BAR_COLUMNS,
CsvSnapshotStore,
SnapshotReadError,
)
def make_bar(trade_date: date, close: str = "10") -> Bar:
@@ -44,3 +48,24 @@ def test_daily_basic_snapshot_rejects_duplicate_codes(tmp_path: Path) -> None:
with pytest.raises(ValueError, match="duplicate"):
store.stage_daily_basic(target, (make_basic(target), make_basic(target)))
def test_readers_reject_empty_formal_snapshots(tmp_path: Path) -> None:
store = CsvSnapshotStore(tmp_path)
path = store.bars_path("000001.SZ")
path.parent.mkdir(parents=True)
path.write_text(",".join(BAR_COLUMNS) + "\n", encoding="utf-8")
with pytest.raises(SnapshotReadError, match="empty"):
store.read_bars("000001.SZ")
def test_invalid_daily_basic_path_is_reported_without_opening_content(tmp_path: Path) -> None:
store = CsvSnapshotStore(tmp_path)
invalid_path = tmp_path / "daily-basic" / "2024" / "20241301.csv"
invalid_path.parent.mkdir(parents=True)
invalid_path.write_text("not parsed", encoding="utf-8")
assert store.list_invalid_daily_basic_snapshot_files() == (
("daily-basic/2024/20241301.csv", "window_out_of_bounds"),
)
@@ -13,6 +13,7 @@ from zhixing_server.modules.market_data.domain.models import (
DailyBasic,
Stock,
SyncWindow,
decimal_text,
)
from zhixing_server.modules.market_data.domain.rules import filter_current_hs_a_stocks
@@ -65,6 +66,12 @@ def test_daily_basic_maps_nan_to_none_but_rejects_infinite_values() -> None:
)
def test_decimal_text_rejects_non_finite_domain_values() -> None:
for value in (Decimal("NaN"), Decimal("Infinity"), Decimal("-Infinity")):
with pytest.raises(ValueError, match="must be finite"):
decimal_text(value)
def test_universe_keeps_current_non_st_hs_a_stocks() -> None:
stocks = (
Stock("000001.SZ", "平安银行", exchange="SZSE", list_status="L"),
@@ -0,0 +1,345 @@
from collections.abc import Generator, Iterable
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from fastapi.testclient import TestClient
from zhixing_server.bootstrap.app import create_app
from zhixing_server.modules.market_data.application.integrity import RunMarketIntegrityCheck
from zhixing_server.modules.market_data.domain.integrity import (
IntegrityCheckInProgress,
IntegrityCheckPage,
IntegrityCheckQuery,
IntegrityCheckRun,
IntegrityCheckStoreError,
IntegrityIssue,
)
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, Stock, SyncWindow
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import SnapshotReadError
from zhixing_server.modules.market_data.presentation.integrity import get_market_integrity_service
WINDOW = SyncWindow(date(2024, 1, 2), date(2024, 1, 3))
STOCK = Stock("000001.SZ", "平安银行", exchange="SZSE")
BARS = (
Bar(STOCK.ts_code, date(2024, 1, 2), close=Decimal("10")),
Bar(STOCK.ts_code, date(2024, 1, 3), close=Decimal("11")),
)
BASICS = (
DailyBasic(STOCK.ts_code, date(2024, 1, 2), close=Decimal("10")),
DailyBasic(STOCK.ts_code, date(2024, 1, 3), close=Decimal("11")),
)
class FakeReader:
def __init__(self, *, acquired: bool = True) -> None:
self.acquired = acquired
self.stocks = (STOCK,)
self.bars: tuple[Bar, ...] = BARS
self.basics: tuple[DailyBasic, ...] = BASICS
self.bar_rows_read = 0
self.basic_rows_read = 0
def latest_successful_window(self) -> SyncWindow:
return WINDOW
def active_stocks(self) -> tuple[Stock, ...]:
return self.stocks
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
return (STOCK.ts_code,)
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
return (date(2024, 1, 2), date(2024, 1, 3))
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]:
for row in self.bars:
self.bar_rows_read += 1
yield row
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]:
for row in self.basics:
self.basic_rows_read += 1
yield row
@contextmanager
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
yield self.acquired
class BrokenKeyReader(FakeReader):
"""Reader whose storage key query fails before a check can be claimed."""
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
raise RuntimeError("database unavailable")
class FakeSnapshots:
def __init__(self) -> None:
self.stocks: tuple[Stock, ...] | None = (STOCK,)
self.bars: tuple[Bar, ...] | None = BARS
self.basics: dict[date, tuple[DailyBasic, ...]] = {
date(2024, 1, 2): (BASICS[0],),
date(2024, 1, 3): (BASICS[1],),
}
self.parse_bar = False
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
return (STOCK.ts_code,) if self.bars is not None else ()
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
return tuple(self.basics)
def read_stocks(self) -> tuple[Stock, ...] | None:
return self.stocks
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None:
if self.parse_bar:
raise SnapshotReadError("bad bar csv")
return self.bars
def read_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...] | None:
return self.basics.get(trade_date)
class FakeStore:
def __init__(self) -> None:
self.run = IntegrityCheckRun("check-1", "running", WINDOW, 4)
self.issues: list[IntegrityIssue] = []
self.progress: tuple[int, int] = (0, 0)
def recover_stale_running(self, lock_key: int, stale_after_seconds: int) -> int:
return 0
def create_running(self, window: SyncWindow, target_count: int) -> IntegrityCheckRun:
self.run = IntegrityCheckRun("check-1", "running", window, target_count)
return self.run
def update_progress(self, check_id: str, checked_count: int, issue_count: int) -> None:
self.progress = (checked_count, issue_count)
self.run = IntegrityCheckRun(
self.run.id,
self.run.status,
self.run.window,
self.run.target_count,
checked_count,
issue_count,
)
def record_issues(self, issues: Iterable[IntegrityIssue]) -> None:
self.issues.extend(issues)
def finish(
self,
check_id: str,
status: str,
*,
issue_count: int,
error_type: str | None = None,
error_message: str | None = None,
) -> None:
self.run = IntegrityCheckRun(
self.run.id,
status, # type: ignore[arg-type]
self.run.window,
self.run.target_count,
self.progress[0],
issue_count,
error_type,
error_message,
)
def get(self, check_id: str, query: IntegrityCheckQuery) -> IntegrityCheckPage:
return IntegrityCheckPage(
self.run,
query.page,
query.page_size,
len(self.issues),
tuple(self.issues[(query.page - 1) * query.page_size : query.page * query.page_size]),
)
def get_latest(self, query: IntegrityCheckQuery) -> IntegrityCheckPage:
return self.get(self.run.id, query)
def test_integrity_check_passes_without_writing_facts() -> None:
reader = FakeReader()
snapshots = FakeSnapshots()
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
run = service.prepare()
service.execute(run.id)
assert store.run.status == "passed"
assert store.issues == []
assert reader.bar_rows_read == len(BARS)
assert reader.basic_rows_read == len(BASICS)
def test_prepare_does_not_hide_storage_key_query_failure() -> None:
service = RunMarketIntegrityCheck(BrokenKeyReader(), FakeSnapshots(), FakeStore())
try:
service.prepare()
except RuntimeError as exc:
assert str(exc) == "database unavailable"
else:
raise AssertionError("storage failure must not claim a running check")
def test_integrity_check_reports_stable_content_and_missing_types() -> None:
reader = FakeReader()
snapshots = FakeSnapshots()
snapshots.bars = (Bar(STOCK.ts_code, date(2024, 1, 2), close=Decimal("99")),)
snapshots.basics.pop(date(2024, 1, 3))
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
service.execute(service.prepare().id)
assert store.run.status == "issues_found"
assert {issue.issue_type for issue in store.issues} >= {
"content_mismatch",
"missing_csv",
}
keys = [issue.issue_key for issue in store.issues]
assert keys == list(dict.fromkeys(keys))
def test_csv_parse_error_does_not_stop_other_groups() -> None:
reader = FakeReader()
snapshots = FakeSnapshots()
snapshots.parse_bar = True
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
service.execute(service.prepare().id)
assert store.run.status == "issues_found"
assert any(issue.issue_type == "parse_error" for issue in store.issues)
assert not any(
issue.item_kind == "bar" and issue.issue_type == "missing_csv" for issue in store.issues
)
assert reader.basic_rows_read == len(BASICS)
def test_window_only_bar_group_does_not_become_extra_csv() -> None:
reader = FakeReader()
reader.bars = ()
snapshots = FakeSnapshots()
snapshots.bars = (Bar(STOCK.ts_code, date(2025, 1, 2), close=Decimal("10")),)
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
service.execute(service.prepare().id)
bar_issues = [issue for issue in store.issues if issue.item_kind == "bar"]
assert {issue.issue_type for issue in bar_issues} == {"window_out_of_bounds"}
def test_integrity_lock_conflict_converges_to_failed() -> None:
reader = FakeReader(acquired=False)
store = FakeStore()
service = RunMarketIntegrityCheck(reader, FakeSnapshots(), store)
service.execute(service.prepare().id)
assert store.run.status == "failed"
assert store.run.error_type == "lock_unavailable"
class FakeHttpService:
def __init__(self, result: IntegrityCheckPage | None = None) -> None:
self.result = result
self.executed = False
self.mode = "ok"
def prepare(self) -> IntegrityCheckRun:
if self.mode == "in_progress":
raise IntegrityCheckInProgress("already running")
if self.mode == "storage_error":
raise IntegrityCheckStoreError("storage unavailable")
return IntegrityCheckRun("http-check", "running", WINDOW, 1)
def execute(self, check_id: str) -> None:
self.executed = True
def get(self, check_id: str, *, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
if self.mode == "storage_error":
raise IntegrityCheckStoreError("storage unavailable")
return self.result
def get_latest(self, *, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
if self.mode == "storage_error":
raise IntegrityCheckStoreError("storage unavailable")
return self.result
def test_integrity_http_returns_202_and_paginates_report() -> None:
issue = IntegrityIssue.build("http-check", "bar", "000001.SZ:2024-01-02", "parse_error", "bad")
result = IntegrityCheckPage(
IntegrityCheckRun("http-check", "issues_found", WINDOW, 1, 1, 1),
page=2,
page_size=1,
issues_total=1,
issues=(issue,),
)
service = FakeHttpService(result)
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: service
client = TestClient(app)
accepted = client.post("/api/v1/market-data/integrity-checks")
report = client.get(
"/api/v1/market-data/integrity-checks/http-check",
params={"page": 2, "page_size": 1},
)
assert accepted.status_code == 202
assert accepted.json()["check_id"] == "http-check"
assert service.executed is True
assert report.status_code == 200
assert report.json()["issues"][0]["issue_type"] == "parse_error"
assert report.json()["page"] == 2
def test_integrity_http_latest_without_report_is_no_data() -> None:
service = FakeHttpService(None)
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: service
response = TestClient(app).get("/api/v1/market-data/integrity-checks/latest")
assert response.status_code == 200
assert response.json()["status"] == "no_data"
assert response.json()["issues"] == []
def test_integrity_http_maps_running_conflict_to_409() -> None:
service = FakeHttpService()
service.mode = "in_progress"
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: service
response = TestClient(app).post("/api/v1/market-data/integrity-checks")
assert response.status_code == 409
assert response.json()["detail"]["code"] == "integrity_check_in_progress"
def test_integrity_http_maps_missing_report_and_storage_to_404_and_503() -> None:
missing_service = FakeHttpService(None)
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: missing_service
missing = TestClient(app).get("/api/v1/market-data/integrity-checks/unknown")
assert missing.status_code == 404
assert missing.json()["detail"]["code"] == "integrity_check_not_found"
storage_service = FakeHttpService(None)
storage_service.mode = "storage_error"
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: storage_service
unavailable = TestClient(app).get("/api/v1/market-data/integrity-checks/latest")
assert unavailable.status_code == 503
assert unavailable.json()["detail"]["code"] == "integrity_storage_unavailable"
@@ -0,0 +1,241 @@
from __future__ import annotations
import threading
import time
from collections.abc import Generator, Iterable, Sequence
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from pathlib import Path
from zhixing_server.modules.market_data.application.sync import (
SyncBatchSummary,
SyncItemOutcome,
SyncMarketData,
SyncMarketDataCommand,
)
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, Stock, SyncWindow
from zhixing_server.modules.market_data.domain.ports import WriteResult
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import CsvSnapshotStore
class ConcurrentSource:
def __init__(
self, target: date, delays: dict[str, float], failed_code: str | None = None
) -> None:
self.target = target
self.stocks = tuple(
Stock(f"{index:06d}.SZ", f"测试{index}", exchange="SZSE", list_status="L")
for index in range(1, 6)
)
self.delays = delays
self.failed_code = failed_code
self._lock = threading.Lock()
self.active = 0
self.max_active = 0
def fetch_stocks(self) -> Sequence[Stock]:
return self.stocks
def fetch_open_dates(self, start: date, end: date) -> Sequence[date]:
return (self.target,) if start <= self.target <= end else ()
def fetch_daily_basic(self, trade_date: date) -> Sequence[DailyBasic]:
return tuple(
DailyBasic(stock.ts_code, trade_date, close=Decimal("10")) for stock in self.stocks
)
def fetch_bars(self, ts_code: str, window: SyncWindow) -> Sequence[Bar]:
with self._lock:
self.active += 1
self.max_active = max(self.max_active, self.active)
try:
time.sleep(self.delays.get(ts_code, 0))
if ts_code == self.failed_code:
raise RuntimeError("simulated source failure")
return (Bar(ts_code, self.target, close=Decimal("10")),)
finally:
with self._lock:
self.active -= 1
class ConcurrentRepository:
def __init__(self) -> None:
self.bars: dict[tuple[str, date], Bar] = {}
self.daily_basic: dict[tuple[str, date], DailyBasic] = {}
self.audit_batches: list[tuple[SyncItemOutcome, ...]] = []
self.batch_summary: tuple[str, int, Decimal, bool] | None = None
self._lock = threading.Lock()
@contextmanager
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
yield True
def upsert_stocks(self, rows: Iterable[Stock]) -> WriteResult:
return WriteResult(inserted=len(tuple(rows)))
def upsert_bars(
self,
rows: Iterable[Bar],
window: SyncWindow,
*,
full_snapshot: bool,
) -> WriteResult:
records = tuple(rows)
inserted = 0
for row in records:
with self._lock:
if (row.ts_code, row.trade_date) not in self.bars:
inserted += 1
self.bars[(row.ts_code, row.trade_date)] = row
return WriteResult(inserted=inserted)
def upsert_daily_basic(self, rows: Iterable[DailyBasic], window: SyncWindow) -> WriteResult:
records = tuple(rows)
with self._lock:
for row in records:
self.daily_basic[(row.ts_code, row.trade_date)] = row
return WriteResult(inserted=len(records))
def purge_before(self, window: SyncWindow) -> None:
return None
def create_batch(
self,
target_trade_date: date,
window: SyncWindow,
mode: str,
parent_batch_id: str | None,
target_count: int,
) -> str:
return "batch-test"
def record_batch(
self,
batch_id: str,
status: str,
valid_count: int,
coverage: Decimal,
strategy_eligible: bool,
) -> None:
self.batch_summary = (status, valid_count, coverage, strategy_eligible)
def record_items(self, batch_id: str, outcomes: Iterable[object]) -> None:
self.audit_batches.append(tuple(outcomes)) # type: ignore[arg-type]
def record_item(
self,
batch_id: str,
item_kind: str,
item_key: str,
status: str,
result: WriteResult,
fingerprint: str | None = None,
error_type: str | None = None,
error_message: str | None = None,
) -> None:
return None
def count_valid_stocks(self, trade_date: date) -> int:
return sum(
1
for ts_code, current_date in self.bars
if current_date == trade_date and (ts_code, trade_date) in self.daily_basic
)
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]:
return ()
def has_bar(self, ts_code: str, trade_date: date) -> bool:
raise AssertionError("coverage must use count_valid_stocks")
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool:
raise AssertionError("coverage must use count_valid_stocks")
def _run_sync(
tmp_path: Path, delays: dict[str, float], failed_code: str | None = None
) -> tuple[SyncBatchSummary, ConcurrentSource, ConcurrentRepository]:
target = date(2024, 1, 2)
source = ConcurrentSource(target, delays, failed_code)
repository = ConcurrentRepository()
summary = SyncMarketData(
source,
CsvSnapshotStore(tmp_path),
repository,
today=target,
max_workers=2,
).execute(SyncMarketDataCommand(target_trade_date=target))
return summary, source, repository
def test_bar_workers_are_bounded_and_audit_is_batched(tmp_path: Path) -> None:
summary, source, repository = _run_sync(
tmp_path,
{"000001.SZ": 0.03, "000002.SZ": 0.01, "000003.SZ": 0.02, "000004.SZ": 0},
)
assert summary.status == "success"
assert source.max_active <= 2
assert summary.valid_count == 5
assert sum(len(batch) for batch in repository.audit_batches) == 7
def test_one_bar_failure_does_not_publish_a_csv_or_reduce_other_facts(tmp_path: Path) -> None:
failed_code = "000003.SZ"
summary, _, repository = _run_sync(
tmp_path,
{"000001.SZ": 0.02, "000002.SZ": 0.01, failed_code: 0},
failed_code,
)
assert summary.status == "partial_success"
assert summary.valid_count == 4
assert [failure.item_key for failure in summary.failures if failure.item_kind == "bar"] == [
failed_code
]
assert not (tmp_path / "bars" / f"{failed_code}.csv").exists()
assert not list((tmp_path / "bars").glob(f".{failed_code}.csv.*.tmp"))
assert len(repository.bars) == 4
def test_completion_order_does_not_change_aggregate_counts(tmp_path: Path) -> None:
first, _, _ = _run_sync(
tmp_path / "first",
{"000001.SZ": 0.03, "000002.SZ": 0, "000003.SZ": 0.02},
)
second, _, _ = _run_sync(
tmp_path / "second",
{"000001.SZ": 0, "000002.SZ": 0.03, "000003.SZ": 0.01},
)
assert (
first.status,
first.valid_count,
first.inserted_count,
first.updated_count,
first.unchanged_count,
first.failures,
) == (
second.status,
second.valid_count,
second.inserted_count,
second.updated_count,
second.unchanged_count,
second.failures,
)
def test_partial_success_always_uses_incomplete_exit_code() -> None:
summary = SyncBatchSummary(
batch_id="batch-test",
target_trade_date=date(2024, 1, 2),
window=SyncWindow(start=date(2018, 1, 2), end=date(2024, 1, 2)),
status="partial_success",
target_count=100,
valid_count=99,
coverage=Decimal("0.99"),
strategy_eligible=True,
)
assert summary.exit_code == 2
@@ -4,7 +4,10 @@ import pytest
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
from zhixing_server.modules.market_data.domain.models import SyncWindow
from zhixing_server.modules.market_data.infrastructure.tushare import TushareAdapter
from zhixing_server.modules.market_data.infrastructure.tushare import (
RequestCoordinator,
TushareAdapter,
)
def test_from_token_reuses_api_client_for_pro_bar(
@@ -44,3 +47,83 @@ def test_from_token_reuses_api_client_for_pro_bar(
assert calls
assert calls[0]["api"] is created_client
assert calls[0]["adj"] == "qfq"
assert calls[0]["retry_count"] == 1
def test_rate_limit_cooldown_is_shared_by_following_requests() -> None:
current = [0.0]
waits: list[float] = []
calls: list[float] = []
def clock() -> float:
return current[0]
def wait(seconds: float) -> None:
waits.append(seconds)
current[0] += seconds
coordinator = RequestCoordinator(
max_retries=0,
cooldown_seconds=(60, 120, 180),
clock=clock,
wait_fn=wait,
sleep_fn=wait,
)
def rate_limited() -> object:
calls.append(current[0])
raise RuntimeError("HTTP 429: too many requests")
with pytest.raises(RuntimeError):
coordinator.call("daily", rate_limited)
assert calls == [0]
assert coordinator.cooldown_until == 60
result = coordinator.call("adj_factor", lambda: calls.append(current[0]) or "ok")
assert result == "ok"
assert calls == [0, 60]
assert waits == [60]
def test_pro_bar_qfq_calls_are_bound_to_the_shared_coordinator(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class Client:
def __init__(self) -> None:
self.calls: list[str] = []
def daily(self, **kwargs: object) -> object:
self.calls.append("daily")
return object()
def adj_factor(self, **kwargs: object) -> object:
self.calls.append("adj_factor")
return object()
client = Client()
pro_bar_calls: list[dict[str, object]] = []
def fake_pro_api(token: str) -> object:
return client
def fake_pro_bar(**kwargs: object) -> list[dict[str, object]]:
pro_bar_calls.append(kwargs)
api = kwargs["api"]
assert api is client
api.daily(ts_code="000001.SZ") # type: ignore[union-attr]
api.adj_factor(ts_code="000001.SZ") # type: ignore[union-attr]
return [{"ts_code": "000001.SZ", "trade_date": "20240102", "close": "10"}]
monkeypatch.setattr(ts, "pro_api", fake_pro_api)
monkeypatch.setattr(ts, "pro_bar", fake_pro_bar)
adapter = TushareAdapter.from_token("test-token", request_interval_seconds=0)
adapter.fetch_bars(
"000001.SZ",
SyncWindow(start=date(2024, 1, 2), end=date(2024, 1, 2)),
)
assert client.calls == ["daily", "adj_factor"]
assert pro_bar_calls[0]["retry_count"] == 1