Files

346 lines
12 KiB
Python
Raw Permalink Normal View History

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"