346 lines
12 KiB
Python
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"
|