feat(market-data): 优化同步并增加完整性检查
This commit is contained in:
@@ -33,6 +33,8 @@ def test_postgres_migration_creates_market_data_contract(
|
||||
"market_daily_basic",
|
||||
"market_sync_batch",
|
||||
"market_sync_item",
|
||||
"market_integrity_check",
|
||||
"market_integrity_issue",
|
||||
"selection_run",
|
||||
"selection_run_item",
|
||||
"selection_signal",
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import psycopg
|
||||
import pytest
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
|
||||
from zhixing_server.modules.market_data.application.sync import SyncFailure, SyncItemOutcome
|
||||
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, SyncWindow
|
||||
from zhixing_server.modules.market_data.domain.ports import MarketDataRepositoryError, WriteResult
|
||||
from zhixing_server.modules.market_data.infrastructure.postgres import (
|
||||
PostgresMarketDataRepository,
|
||||
)
|
||||
|
||||
TEST_CODES = ("991901.SZ", "991902.SZ")
|
||||
TARGET = date(2024, 1, 2)
|
||||
WINDOW = SyncWindow(start=TARGET, end=TARGET)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _database(database_url: str) -> Any:
|
||||
with psycopg.connect(database_url) as connection, connection.transaction():
|
||||
yield connection
|
||||
|
||||
|
||||
def _cleanup(database_url: str) -> None:
|
||||
with _database(database_url) as connection:
|
||||
connection.execute(
|
||||
"DELETE FROM market_sync_item WHERE batch_id IN "
|
||||
"(SELECT id FROM market_sync_batch WHERE id LIKE '9919%')"
|
||||
)
|
||||
connection.execute("DELETE FROM market_sync_batch WHERE id LIKE '9919%'")
|
||||
connection.execute(
|
||||
"DELETE FROM market_daily_basic WHERE ts_code = ANY(%s)",
|
||||
(list(TEST_CODES),),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM market_daily_bar WHERE ts_code = ANY(%s)",
|
||||
(list(TEST_CODES),),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM market_stock WHERE ts_code = ANY(%s)",
|
||||
(list(TEST_CODES),),
|
||||
)
|
||||
|
||||
|
||||
def _prepare_database(database_url: str) -> None:
|
||||
server_root = Path(__file__).parents[2]
|
||||
config = Config(str(server_root / "alembic.ini"))
|
||||
config.set_main_option(
|
||||
"sqlalchemy.url",
|
||||
database_url.replace("%", "%%").replace("postgresql://", "postgresql+psycopg://"),
|
||||
)
|
||||
command.upgrade(config, "head")
|
||||
_cleanup(database_url)
|
||||
with _database(database_url) as connection:
|
||||
connection.executemany(
|
||||
"""
|
||||
INSERT INTO market_stock
|
||||
(ts_code, name, market, exchange, list_status, is_active)
|
||||
VALUES (%s, %s, '主板', 'SZSE', 'L', true)
|
||||
""",
|
||||
[(code, f"测试{code}") for code in TEST_CODES],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_pool_upsert_rollback_batch_audit_and_set_coverage() -> None:
|
||||
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
|
||||
if not database_url:
|
||||
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
|
||||
_prepare_database(database_url)
|
||||
repository = PostgresMarketDataRepository(database_url, max_connections=4)
|
||||
batch_id: str | None = None
|
||||
try:
|
||||
bars = {code: Bar(code, TARGET, close=Decimal("10")) for code in TEST_CODES}
|
||||
with repository:
|
||||
|
||||
def upsert(code: str) -> WriteResult:
|
||||
return repository.upsert_bars((bars[code],), WINDOW, full_snapshot=True)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
results = tuple(executor.map(upsert, TEST_CODES))
|
||||
assert [result.inserted for result in results] == [1, 1]
|
||||
|
||||
with pytest.raises(MarketDataRepositoryError):
|
||||
repository.upsert_bars(
|
||||
(
|
||||
Bar(TEST_CODES[0], TARGET, close=Decimal("11")),
|
||||
Bar(TEST_CODES[0], TARGET, close=Decimal("12")),
|
||||
),
|
||||
WINDOW,
|
||||
full_snapshot=True,
|
||||
)
|
||||
|
||||
unchanged = repository.upsert_daily_basic(
|
||||
(DailyBasic(TEST_CODES[0], TARGET, close=Decimal("10")),),
|
||||
WINDOW,
|
||||
)
|
||||
assert unchanged.inserted == 1
|
||||
assert repository.count_valid_stocks(TARGET) == 1
|
||||
|
||||
batch_id = "9919-pool-test"
|
||||
outcomes = (
|
||||
SyncItemOutcome("bar", TEST_CODES[0], "success", WriteResult(inserted=1)),
|
||||
SyncItemOutcome(
|
||||
"bar",
|
||||
TEST_CODES[1],
|
||||
"failed",
|
||||
failure=SyncFailure("bar", TEST_CODES[1], "source_error", "safe failure"),
|
||||
),
|
||||
)
|
||||
# Use the real schema row directly because the test repository's
|
||||
# create_batch API generates a UUID for normal production calls.
|
||||
with _database(database_url) as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO market_sync_batch
|
||||
(id, target_trade_date, window_start, mode, status, target_count)
|
||||
VALUES (%s, %s, %s, 'daily', 'running', 2)
|
||||
""",
|
||||
(batch_id, TARGET, TARGET),
|
||||
)
|
||||
repository.record_items(batch_id, outcomes)
|
||||
with _database(database_url) as connection:
|
||||
count = connection.execute(
|
||||
"SELECT count(*) FROM market_sync_item WHERE batch_id = %s",
|
||||
(batch_id,),
|
||||
).fetchone()[0]
|
||||
assert count == 2
|
||||
finally:
|
||||
repository.close()
|
||||
_cleanup(database_url)
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from zhixing_server.bootstrap.config import Settings
|
||||
|
||||
|
||||
def test_market_data_workers_default_to_eight(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("ZHIXING_MARKET_DATA_MAX_WORKERS", raising=False)
|
||||
settings = Settings()
|
||||
|
||||
assert settings.market_data_max_workers == 8
|
||||
|
||||
|
||||
def test_market_data_workers_must_be_positive() -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
Settings(market_data_max_workers=0)
|
||||
Reference in New Issue
Block a user