Files

346 lines
12 KiB
Python

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"