feat(market-data): 优化同步并增加完整性检查
This commit is contained in:
@@ -5,6 +5,7 @@ from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
@@ -18,7 +19,7 @@ class Settings(BaseSettings):
|
||||
tushare_token: str = ""
|
||||
market_data_csv_root: Path = Path("./data/market-data")
|
||||
market_data_coverage_threshold: Decimal = Decimal("0.99")
|
||||
market_data_max_workers: int = 4
|
||||
market_data_max_workers: int = Field(default=8, ge=1)
|
||||
market_data_request_interval_seconds: float = 0.2
|
||||
market_data_max_retries: int = 3
|
||||
market_data_retry_backoff_seconds: float = 1.0
|
||||
|
||||
@@ -4,11 +4,17 @@ from fastapi import APIRouter
|
||||
|
||||
from zhixing_server.interfaces.http.system import operational_router, system_router
|
||||
from zhixing_server.modules.market_data.presentation.home import home_router
|
||||
from zhixing_server.modules.market_data.presentation.integrity import integrity_router
|
||||
from zhixing_server.modules.selection.presentation.http import selection_router
|
||||
|
||||
api_v1_router = APIRouter(prefix="/api/v1")
|
||||
api_v1_router.include_router(system_router, prefix="/system", tags=["system"])
|
||||
api_v1_router.include_router(home_router, prefix="/home", tags=["home"])
|
||||
api_v1_router.include_router(
|
||||
integrity_router,
|
||||
prefix="/market-data/integrity-checks",
|
||||
tags=["market-data"],
|
||||
)
|
||||
api_v1_router.include_router(selection_router, prefix="/selection", tags=["selection"])
|
||||
|
||||
__all__ = ["api_v1_router", "operational_router"]
|
||||
|
||||
@@ -1,5 +1,12 @@
|
||||
"""Market data synchronization use cases."""
|
||||
|
||||
from .sync import SyncBatchSummary, SyncMarketData, SyncMarketDataCommand
|
||||
from .integrity import RunMarketIntegrityCheck
|
||||
from .sync import SyncBatchSummary, SyncItemOutcome, SyncMarketData, SyncMarketDataCommand
|
||||
|
||||
__all__ = ["SyncBatchSummary", "SyncMarketData", "SyncMarketDataCommand"]
|
||||
__all__ = [
|
||||
"RunMarketIntegrityCheck",
|
||||
"SyncBatchSummary",
|
||||
"SyncItemOutcome",
|
||||
"SyncMarketData",
|
||||
"SyncMarketDataCommand",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,671 @@
|
||||
"""Read-only PostgreSQL/CSV market-data integrity check use case."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable, Iterable, Iterator, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date
|
||||
from typing import TypeVar, cast
|
||||
|
||||
from ..domain.integrity import (
|
||||
IntegrityCheckNoData,
|
||||
IntegrityCheckNotFound,
|
||||
IntegrityCheckPage,
|
||||
IntegrityCheckQuery,
|
||||
IntegrityCheckRun,
|
||||
IntegrityCheckStore,
|
||||
IntegrityIssue,
|
||||
IntegritySnapshotReader,
|
||||
IntegritySnapshotStore,
|
||||
)
|
||||
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_ISSUE_BATCH_SIZE = 100
|
||||
_DEFAULT_STALE_AFTER_SECONDS = 3_600
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def _empty_issues() -> list[IntegrityIssue]:
|
||||
"""Provide an explicitly typed empty issue buffer for Pyright."""
|
||||
|
||||
return []
|
||||
|
||||
|
||||
def _empty_issue_keys() -> set[str]:
|
||||
"""Provide an explicitly typed empty issue-key set for Pyright."""
|
||||
|
||||
return set()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _IssueCollector:
|
||||
"""Bound issue buffering and count progress without retaining the report."""
|
||||
|
||||
check_id: str
|
||||
store: IntegrityCheckStore
|
||||
issue_count: int = 0
|
||||
pending: list[IntegrityIssue] = field(default_factory=_empty_issues)
|
||||
issue_keys: set[str] = field(default_factory=_empty_issue_keys)
|
||||
|
||||
def add(
|
||||
self,
|
||||
item_kind: str,
|
||||
item_key: str,
|
||||
issue_type: str,
|
||||
message: str,
|
||||
) -> None:
|
||||
"""Add one issue and flush at a bounded batch size."""
|
||||
|
||||
issue = IntegrityIssue.build(
|
||||
self.check_id,
|
||||
item_kind,
|
||||
item_key,
|
||||
issue_type,
|
||||
message,
|
||||
)
|
||||
if issue.issue_key in self.issue_keys:
|
||||
return
|
||||
self.issue_keys.add(issue.issue_key)
|
||||
self.pending.append(issue)
|
||||
self.issue_count += 1
|
||||
if len(self.pending) >= _ISSUE_BATCH_SIZE:
|
||||
self.flush()
|
||||
|
||||
def flush(self) -> None:
|
||||
"""Persist buffered issues and release their row objects."""
|
||||
|
||||
if not self.pending:
|
||||
return
|
||||
self.store.record_issues(tuple(self.pending))
|
||||
self.pending.clear()
|
||||
|
||||
|
||||
class RunMarketIntegrityCheck:
|
||||
"""Prepare, execute, and query a persisted integrity report.
|
||||
|
||||
The application depends only on read ports and an independent check store.
|
||||
It never constructs a Tushare adapter and never receives a market-data
|
||||
write port, which makes the no-fact-mutation boundary explicit.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reader: IntegritySnapshotReader,
|
||||
snapshots: IntegritySnapshotStore,
|
||||
store: IntegrityCheckStore,
|
||||
*,
|
||||
lock_key: int = 7_380_521,
|
||||
stale_after_seconds: int = _DEFAULT_STALE_AFTER_SECONDS,
|
||||
) -> None:
|
||||
if stale_after_seconds < 1:
|
||||
raise ValueError("stale_after_seconds must be positive")
|
||||
self.reader = reader
|
||||
self.snapshots = snapshots
|
||||
self.store = store
|
||||
self.lock_key = lock_key
|
||||
self.stale_after_seconds = stale_after_seconds
|
||||
|
||||
def prepare(self) -> IntegrityCheckRun:
|
||||
"""Atomically claim a check against the newest completed window."""
|
||||
|
||||
self.store.recover_stale_running(self.lock_key, self.stale_after_seconds)
|
||||
window = self.reader.latest_successful_window()
|
||||
if window is None:
|
||||
raise IntegrityCheckNoData("no completed market-data synchronization is available")
|
||||
stocks = tuple(self.reader.active_stocks())
|
||||
target_count = self._target_count(window, stocks)
|
||||
return self.store.create_running(window, target_count)
|
||||
|
||||
def execute(self, check_id: str) -> None:
|
||||
"""Run a claimed check and always converge unexpected errors to failed."""
|
||||
|
||||
try:
|
||||
page = self.store.get(check_id, IntegrityCheckQuery(page=1, page_size=1))
|
||||
except Exception as exc: # noqa: BLE001 - background boundary must converge state
|
||||
self._finish_failed(check_id, 0, "check_error", _safe_error(exc))
|
||||
return
|
||||
if page is None:
|
||||
raise IntegrityCheckNotFound(f"integrity check not found: {check_id}")
|
||||
if page.run.status != "running":
|
||||
return
|
||||
collector = _IssueCollector(check_id, self.store, issue_count=page.run.issue_count)
|
||||
try:
|
||||
with self.reader.advisory_lock(self.lock_key) as acquired:
|
||||
if not acquired:
|
||||
self._finish_failed(
|
||||
check_id,
|
||||
collector.issue_count,
|
||||
"lock_unavailable",
|
||||
"market-data synchronization or another integrity check is running",
|
||||
)
|
||||
return
|
||||
try:
|
||||
self._execute_locked(page.run, collector)
|
||||
except Exception as exc: # noqa: BLE001 - worker boundary must converge state
|
||||
logger.exception("market_integrity_check_failed check_id=%s", check_id)
|
||||
self._finish_failed(
|
||||
check_id,
|
||||
collector.issue_count,
|
||||
"check_error",
|
||||
_safe_error(exc),
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - lock/storage boundary must be observable
|
||||
logger.exception("market_integrity_check_worker_failed check_id=%s", check_id)
|
||||
self._finish_failed(check_id, collector.issue_count, "check_error", _safe_error(exc))
|
||||
|
||||
def get(
|
||||
self,
|
||||
check_id: str,
|
||||
*,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
query: IntegrityCheckQuery | None = None,
|
||||
) -> IntegrityCheckPage | None:
|
||||
"""Read one check and a bounded issue page."""
|
||||
|
||||
resolved = query or IntegrityCheckQuery(page=page, page_size=page_size)
|
||||
return self.store.get(check_id, resolved)
|
||||
|
||||
def get_latest(
|
||||
self,
|
||||
*,
|
||||
page: int = 1,
|
||||
page_size: int = 10,
|
||||
query: IntegrityCheckQuery | None = None,
|
||||
) -> IntegrityCheckPage | None:
|
||||
"""Read the newest check or ``None`` when no check has been run."""
|
||||
|
||||
resolved = query or IntegrityCheckQuery(page=page, page_size=page_size)
|
||||
return self.store.get_latest(resolved)
|
||||
|
||||
def _target_count(self, window: SyncWindow, stocks: Sequence[Stock]) -> int:
|
||||
"""Count comparison groups using keys only, never fact rows."""
|
||||
|
||||
bar_codes = set(_string_keys(self.reader, "list_bar_codes", window))
|
||||
bar_codes.update(_string_keys(self.snapshots, "list_bar_codes", window))
|
||||
bar_codes.update(stock.ts_code for stock in stocks)
|
||||
basic_dates = set(_date_keys(self.reader, "list_daily_basic_dates", window))
|
||||
basic_dates.update(_date_keys(self.snapshots, "list_daily_basic_dates", window))
|
||||
return max(1, 1 + len(bar_codes) + len(basic_dates))
|
||||
|
||||
def _execute_locked(self, run: IntegrityCheckRun, collector: _IssueCollector) -> None:
|
||||
"""Compare each group and persist progress after it completes."""
|
||||
|
||||
if run.window is None:
|
||||
raise ValueError("integrity check has no comparison window")
|
||||
window = run.window
|
||||
self._compare_stocks(run.id, collector)
|
||||
checked_count = 1
|
||||
collector.flush()
|
||||
self.store.update_progress(run.id, checked_count, collector.issue_count)
|
||||
|
||||
checked_count = self._compare_bars(
|
||||
run.id,
|
||||
window,
|
||||
collector,
|
||||
checked_count,
|
||||
)
|
||||
checked_count = self._compare_daily_basic(
|
||||
run.id,
|
||||
window,
|
||||
collector,
|
||||
checked_count,
|
||||
)
|
||||
collector.flush()
|
||||
self.store.update_progress(run.id, checked_count, collector.issue_count)
|
||||
self.store.finish(
|
||||
run.id,
|
||||
"issues_found" if collector.issue_count else "passed",
|
||||
issue_count=collector.issue_count,
|
||||
)
|
||||
|
||||
def _compare_stocks(self, check_id: str, collector: _IssueCollector) -> None:
|
||||
"""Compare current stock master rows as one small bounded group."""
|
||||
|
||||
database_rows = tuple(self.reader.active_stocks())
|
||||
try:
|
||||
csv_rows = self.snapshots.read_stocks()
|
||||
except Exception as exc: # noqa: BLE001 - one malformed file must not stop other groups
|
||||
collector.add("stock", "current", _issue_type(exc), "stock-master CSV cannot be parsed")
|
||||
return
|
||||
if csv_rows is None:
|
||||
if database_rows:
|
||||
collector.add(
|
||||
"stock",
|
||||
"current",
|
||||
"missing_csv",
|
||||
"current stock-master CSV is missing",
|
||||
)
|
||||
return
|
||||
database_by_code = {row.ts_code: row for row in database_rows}
|
||||
csv_by_code = {row.ts_code: row for row in csv_rows}
|
||||
for code in sorted(database_by_code.keys() - csv_by_code.keys()):
|
||||
collector.add("stock", code, "missing_csv", "active stock is missing from CSV")
|
||||
for code in sorted(csv_by_code.keys() - database_by_code.keys()):
|
||||
collector.add("stock", code, "extra_csv", "CSV stock is not active in PostgreSQL")
|
||||
for code in sorted(database_by_code.keys() & csv_by_code.keys()):
|
||||
if database_by_code[code] != csv_by_code[code]:
|
||||
collector.add("stock", code, "content_mismatch", "stock-master fields differ")
|
||||
|
||||
def _compare_bars(
|
||||
self,
|
||||
check_id: str,
|
||||
window: SyncWindow,
|
||||
collector: _IssueCollector,
|
||||
checked_count: int,
|
||||
) -> int:
|
||||
"""Merge PostgreSQL and per-stock CSV bars one stock at a time."""
|
||||
|
||||
database_codes = set(_string_keys(self.reader, "list_bar_codes", window))
|
||||
csv_codes = set(_string_keys(self.snapshots, "list_bar_codes", window))
|
||||
expected_codes = (
|
||||
database_codes | csv_codes | {stock.ts_code for stock in self.reader.active_stocks()}
|
||||
)
|
||||
for code in _string_keys(self.reader, "list_invalid_bar_codes", window):
|
||||
collector.add("bar", code, "content_mismatch", "database bar source_adj is not qfq")
|
||||
database_groups = _grouped(self.reader.iter_bars(window), lambda row: row.ts_code)
|
||||
parse_failed_codes: set[str] = set()
|
||||
current = next(database_groups, None)
|
||||
for code in sorted(expected_codes):
|
||||
while current is not None and current[0] < code:
|
||||
self._compare_bar_group(current[0], current[1], None, window, collector)
|
||||
checked_count += 1
|
||||
current = next(database_groups, None)
|
||||
self._progress(collector, checked_count)
|
||||
database_rows: tuple[Bar, ...] | None = None
|
||||
if current is not None and current[0] == code:
|
||||
database_rows = current[1]
|
||||
current = next(database_groups, None)
|
||||
csv_rows = (
|
||||
self._read_bars(code, collector, parse_failed_codes) if code in csv_codes else None
|
||||
)
|
||||
self._compare_bar_group(
|
||||
code,
|
||||
database_rows,
|
||||
csv_rows,
|
||||
window,
|
||||
collector,
|
||||
csv_parse_failed=code in parse_failed_codes,
|
||||
)
|
||||
checked_count += 1
|
||||
self._progress(collector, checked_count)
|
||||
while current is not None:
|
||||
self._compare_bar_group(current[0], current[1], None, window, collector)
|
||||
checked_count += 1
|
||||
current = next(database_groups, None)
|
||||
self._progress(collector, checked_count)
|
||||
return checked_count
|
||||
|
||||
def _compare_daily_basic(
|
||||
self,
|
||||
check_id: str,
|
||||
window: SyncWindow,
|
||||
collector: _IssueCollector,
|
||||
checked_count: int,
|
||||
) -> int:
|
||||
"""Merge PostgreSQL and per-date CSV daily-basic groups."""
|
||||
|
||||
database_dates = set(_date_keys(self.reader, "list_daily_basic_dates", window))
|
||||
csv_dates = set(_date_keys(self.snapshots, "list_daily_basic_dates", window))
|
||||
expected_dates = database_dates | csv_dates
|
||||
database_groups = _grouped(self.reader.iter_daily_basic(window), lambda row: row.trade_date)
|
||||
parse_failed_dates: set[date] = set()
|
||||
invalid_path_method = getattr(
|
||||
self.snapshots, "list_invalid_daily_basic_snapshot_files", None
|
||||
)
|
||||
if callable(invalid_path_method):
|
||||
try:
|
||||
invalid_paths = cast(Iterable[object], invalid_path_method())
|
||||
except Exception: # noqa: BLE001 - malformed path listing must not abort other groups
|
||||
invalid_paths = ()
|
||||
for value in invalid_paths:
|
||||
if not isinstance(value, tuple):
|
||||
continue
|
||||
parts = cast(tuple[object, ...], value)
|
||||
if len(parts) != 2:
|
||||
continue
|
||||
item_key, issue_type = parts
|
||||
collector.add(
|
||||
"daily_basic",
|
||||
str(item_key),
|
||||
str(issue_type) or "parse_error",
|
||||
"daily-basic snapshot path cannot be parsed",
|
||||
)
|
||||
current = next(database_groups, None)
|
||||
for trade_date in sorted(expected_dates):
|
||||
while current is not None and current[0] < trade_date:
|
||||
self._compare_basic_group(current[0], current[1], None, window, collector)
|
||||
checked_count += 1
|
||||
current = next(database_groups, None)
|
||||
self._progress(collector, checked_count)
|
||||
database_rows: tuple[DailyBasic, ...] | None = None
|
||||
if current is not None and current[0] == trade_date:
|
||||
database_rows = current[1]
|
||||
current = next(database_groups, None)
|
||||
csv_rows = (
|
||||
self._read_daily_basic(trade_date, collector, parse_failed_dates)
|
||||
if trade_date in csv_dates
|
||||
else None
|
||||
)
|
||||
self._compare_basic_group(
|
||||
trade_date,
|
||||
database_rows,
|
||||
csv_rows,
|
||||
window,
|
||||
collector,
|
||||
csv_parse_failed=trade_date in parse_failed_dates,
|
||||
)
|
||||
checked_count += 1
|
||||
self._progress(collector, checked_count)
|
||||
while current is not None:
|
||||
self._compare_basic_group(current[0], current[1], None, window, collector)
|
||||
checked_count += 1
|
||||
current = next(database_groups, None)
|
||||
self._progress(collector, checked_count)
|
||||
return checked_count
|
||||
|
||||
def _read_bars(
|
||||
self,
|
||||
code: str,
|
||||
collector: _IssueCollector,
|
||||
parse_failed_codes: set[str],
|
||||
) -> tuple[Bar, ...] | None:
|
||||
"""Read one formal bar file and isolate its parse failure."""
|
||||
|
||||
try:
|
||||
return self.snapshots.read_bars(code)
|
||||
except Exception as exc: # noqa: BLE001 - continue with other stocks
|
||||
parse_failed_codes.add(code)
|
||||
collector.add("bar", code, _issue_type(exc), "bar CSV cannot be parsed")
|
||||
return None
|
||||
|
||||
def _read_daily_basic(
|
||||
self,
|
||||
trade_date: date,
|
||||
collector: _IssueCollector,
|
||||
parse_failed_dates: set[date],
|
||||
) -> tuple[DailyBasic, ...] | None:
|
||||
"""Read one formal daily-basic file and isolate its parse failure."""
|
||||
|
||||
try:
|
||||
return self.snapshots.read_daily_basic(trade_date)
|
||||
except Exception as exc: # noqa: BLE001 - continue with other dates
|
||||
parse_failed_dates.add(trade_date)
|
||||
collector.add(
|
||||
"daily_basic",
|
||||
trade_date.isoformat(),
|
||||
_issue_type(exc),
|
||||
"daily-basic CSV cannot be parsed",
|
||||
)
|
||||
return None
|
||||
|
||||
def _compare_bar_group(
|
||||
self,
|
||||
code: str,
|
||||
database_rows: tuple[Bar, ...] | None,
|
||||
csv_rows: tuple[Bar, ...] | None,
|
||||
window: SyncWindow,
|
||||
collector: _IssueCollector,
|
||||
*,
|
||||
csv_parse_failed: bool = False,
|
||||
) -> None:
|
||||
"""Compare one stock's rows without building a market-wide map."""
|
||||
|
||||
database_inside = self._keep_bar_rows_in_window(
|
||||
"database", code, database_rows, window, collector
|
||||
)
|
||||
if csv_parse_failed:
|
||||
return
|
||||
csv_inside = self._keep_bar_rows_in_window("csv", code, csv_rows, window, collector)
|
||||
if database_inside is None and csv_inside is None:
|
||||
return
|
||||
if database_inside is None:
|
||||
if csv_inside:
|
||||
self._compare_rows(
|
||||
"bar",
|
||||
code,
|
||||
(),
|
||||
csv_inside,
|
||||
lambda row: row.trade_date,
|
||||
window,
|
||||
collector,
|
||||
row_date=lambda row: row.trade_date,
|
||||
)
|
||||
return
|
||||
if csv_inside is None:
|
||||
if database_inside:
|
||||
self._compare_rows(
|
||||
"bar",
|
||||
code,
|
||||
database_inside,
|
||||
(),
|
||||
lambda row: row.trade_date,
|
||||
window,
|
||||
collector,
|
||||
row_date=lambda row: row.trade_date,
|
||||
)
|
||||
return
|
||||
self._compare_rows(
|
||||
"bar",
|
||||
code,
|
||||
database_inside,
|
||||
csv_inside,
|
||||
lambda row: row.trade_date,
|
||||
window,
|
||||
collector,
|
||||
row_date=lambda row: row.trade_date,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _keep_bar_rows_in_window(
|
||||
source: str,
|
||||
code: str,
|
||||
rows: tuple[Bar, ...] | None,
|
||||
window: SyncWindow,
|
||||
collector: _IssueCollector,
|
||||
) -> tuple[Bar, ...] | None:
|
||||
"""Report and discard out-of-window CSV/DB rows before merging."""
|
||||
|
||||
if rows is None:
|
||||
return None
|
||||
inside: list[Bar] = []
|
||||
for row in rows:
|
||||
if window.contains(row.trade_date):
|
||||
inside.append(row)
|
||||
else:
|
||||
collector.add(
|
||||
"bar",
|
||||
f"{code}:{row.trade_date.isoformat()}",
|
||||
"window_out_of_bounds",
|
||||
f"{source} bar row is outside the check window",
|
||||
)
|
||||
return tuple(inside)
|
||||
|
||||
def _compare_basic_group(
|
||||
self,
|
||||
trade_date: date,
|
||||
database_rows: tuple[DailyBasic, ...] | None,
|
||||
csv_rows: tuple[DailyBasic, ...] | None,
|
||||
window: SyncWindow,
|
||||
collector: _IssueCollector,
|
||||
*,
|
||||
csv_parse_failed: bool = False,
|
||||
) -> None:
|
||||
"""Compare one trading-day's daily-basic rows."""
|
||||
|
||||
group_key = trade_date.isoformat()
|
||||
if not window.contains(trade_date):
|
||||
collector.add(
|
||||
"daily_basic",
|
||||
group_key,
|
||||
"window_out_of_bounds",
|
||||
"daily-basic date is outside the check window",
|
||||
)
|
||||
return
|
||||
if csv_parse_failed:
|
||||
return
|
||||
if database_rows is None and csv_rows is None:
|
||||
return
|
||||
if database_rows is None:
|
||||
collector.add(
|
||||
"daily_basic",
|
||||
group_key,
|
||||
"extra_csv",
|
||||
"daily-basic date exists only in CSV",
|
||||
)
|
||||
return
|
||||
if csv_rows is None:
|
||||
collector.add(
|
||||
"daily_basic", group_key, "missing_csv", "daily-basic date is missing in CSV"
|
||||
)
|
||||
return
|
||||
self._compare_rows(
|
||||
"daily_basic",
|
||||
group_key,
|
||||
database_rows,
|
||||
csv_rows,
|
||||
lambda row: row.ts_code,
|
||||
window,
|
||||
collector,
|
||||
row_date=lambda row: row.trade_date,
|
||||
)
|
||||
|
||||
def _compare_rows(
|
||||
self,
|
||||
item_kind: str,
|
||||
group_key: str,
|
||||
database_rows: Sequence[T],
|
||||
csv_rows: Sequence[T],
|
||||
key: Callable[[T], object],
|
||||
window: SyncWindow,
|
||||
collector: _IssueCollector,
|
||||
row_date: Callable[[T], date] | None = None,
|
||||
) -> None:
|
||||
"""Merge two bounded groups and report duplicate/content differences."""
|
||||
|
||||
database_by_key, database_duplicates = _index_rows(database_rows, key)
|
||||
csv_by_key, csv_duplicates = _index_rows(csv_rows, key)
|
||||
for row_key in sorted(database_duplicates | csv_duplicates, key=str):
|
||||
collector.add(
|
||||
item_kind,
|
||||
f"{group_key}:{row_key}",
|
||||
"duplicate_key",
|
||||
"comparison group contains duplicate business keys",
|
||||
)
|
||||
for row_key in sorted(database_by_key.keys() | csv_by_key.keys(), key=str):
|
||||
database_row = database_by_key.get(row_key)
|
||||
csv_row = csv_by_key.get(row_key)
|
||||
identity = f"{group_key}:{row_key}"
|
||||
if row_date is not None:
|
||||
for candidate in (database_row, csv_row):
|
||||
if candidate is not None and not window.contains(row_date(candidate)):
|
||||
collector.add(
|
||||
item_kind,
|
||||
identity,
|
||||
"window_out_of_bounds",
|
||||
"row is outside the check window",
|
||||
)
|
||||
if database_row is None:
|
||||
collector.add(item_kind, identity, "extra_csv", "row exists only in CSV")
|
||||
elif csv_row is None:
|
||||
collector.add(item_kind, identity, "missing_csv", "row is missing in CSV")
|
||||
elif database_row != csv_row:
|
||||
collector.add(
|
||||
item_kind, identity, "content_mismatch", "database and CSV fields differ"
|
||||
)
|
||||
|
||||
def _progress(self, collector: _IssueCollector, checked_count: int) -> None:
|
||||
"""Persist issues and a heartbeat after one comparison group."""
|
||||
|
||||
collector.flush()
|
||||
self.store.update_progress(collector.check_id, checked_count, collector.issue_count)
|
||||
|
||||
def _finish_failed(
|
||||
self,
|
||||
check_id: str,
|
||||
issue_count: int,
|
||||
error_type: str,
|
||||
message: str,
|
||||
) -> None:
|
||||
"""Best-effort terminal write for background worker failures."""
|
||||
|
||||
try:
|
||||
self.store.finish(
|
||||
check_id,
|
||||
"failed",
|
||||
issue_count=issue_count,
|
||||
error_type=error_type,
|
||||
error_message=message,
|
||||
)
|
||||
except Exception: # noqa: BLE001 - nothing safer can be persisted here
|
||||
logger.exception("market_integrity_check_failure_persist_failed check_id=%s", check_id)
|
||||
|
||||
|
||||
def _grouped[T, K](rows: Iterable[T], key: Callable[[T], K]) -> Iterator[tuple[K, tuple[T, ...]]]:
|
||||
"""Consume an ordered iterator one comparison group at a time."""
|
||||
|
||||
iterator = iter(rows)
|
||||
pending = next(iterator, None)
|
||||
while pending is not None:
|
||||
group_key = key(pending)
|
||||
group: list[T] = [pending]
|
||||
pending = next(iterator, None)
|
||||
while pending is not None and key(pending) == group_key:
|
||||
group.append(pending)
|
||||
pending = next(iterator, None)
|
||||
yield group_key, tuple(group)
|
||||
|
||||
|
||||
def _index_rows[T](
|
||||
rows: Sequence[T], key: Callable[[T], object]
|
||||
) -> tuple[dict[object, T], set[object]]:
|
||||
"""Index one bounded group and retain duplicate identities only."""
|
||||
|
||||
indexed: dict[object, T] = {}
|
||||
duplicates: set[object] = set()
|
||||
for row in rows:
|
||||
row_key = key(row)
|
||||
if row_key in indexed:
|
||||
duplicates.add(row_key)
|
||||
else:
|
||||
indexed[row_key] = row
|
||||
return indexed, duplicates
|
||||
|
||||
|
||||
def _string_keys(adapter: object, method_name: str, window: SyncWindow) -> tuple[str, ...]:
|
||||
"""Call an optional string-key method without hiding storage failures."""
|
||||
|
||||
method = getattr(adapter, method_name, None)
|
||||
if not callable(method):
|
||||
return ()
|
||||
values = cast(Iterable[object], method(window))
|
||||
return tuple(str(value) for value in values)
|
||||
|
||||
|
||||
def _date_keys(adapter: object, method_name: str, window: SyncWindow) -> tuple[date, ...]:
|
||||
"""Call an optional date-key method without hiding storage failures."""
|
||||
|
||||
method = getattr(adapter, method_name, None)
|
||||
if not callable(method):
|
||||
return ()
|
||||
values = cast(Iterable[object], method(window))
|
||||
return tuple(value for value in values if isinstance(value, date))
|
||||
|
||||
|
||||
def _issue_type(error: BaseException) -> str:
|
||||
"""Use adapter-provided stable issue categories when available."""
|
||||
|
||||
value = getattr(error, "issue_type", "parse_error")
|
||||
return str(value) if value else "parse_error"
|
||||
|
||||
|
||||
def _safe_error(error: Exception) -> str:
|
||||
"""Keep worker failures bounded and free of tracebacks or secrets."""
|
||||
|
||||
return " ".join(str(error).split())[:500] or error.__class__.__name__
|
||||
|
||||
|
||||
__all__ = ["RunMarketIntegrityCheck"]
|
||||
@@ -4,11 +4,12 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Iterable, Sequence
|
||||
from collections.abc import Callable, Iterable, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
from typing import Literal, cast
|
||||
|
||||
from ..domain.fingerprint import SnapshotChange, compare_snapshots
|
||||
from ..domain.models import Bar, Stock, SyncWindow
|
||||
@@ -40,6 +41,24 @@ class SyncFailure:
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SyncItemOutcome:
|
||||
"""Immutable result returned by one synchronization item.
|
||||
|
||||
A worker owns all side effects for its stock, then hands only this value
|
||||
back to the coordinating thread. Keeping audit and aggregate mutation out
|
||||
of the worker makes completion order irrelevant and prevents one failed
|
||||
future from corrupting the batch counters.
|
||||
"""
|
||||
|
||||
item_kind: str
|
||||
item_key: str
|
||||
status: Literal["success", "failed"]
|
||||
result: WriteResult = field(default_factory=WriteResult)
|
||||
fingerprint: str | None = None
|
||||
failure: SyncFailure | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SyncBatchSummary:
|
||||
"""Stable output contract for CLI, cron, and later strategy callers."""
|
||||
@@ -63,7 +82,9 @@ class SyncBatchSummary:
|
||||
|
||||
if self.status == "failed":
|
||||
return 1
|
||||
return 0 if self.strategy_eligible else 2
|
||||
if self.status != "success" or not self.strategy_eligible:
|
||||
return 2
|
||||
return 0
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
"""Serialize the summary without credentials or raw vendor responses."""
|
||||
@@ -107,36 +128,55 @@ class SyncMarketData:
|
||||
coverage_threshold: Decimal = Decimal("0.99"),
|
||||
lock_key: int = 7_380_521,
|
||||
today: date | None = None,
|
||||
max_workers: int = 8,
|
||||
) -> None:
|
||||
if not Decimal("0") <= coverage_threshold <= Decimal("1"):
|
||||
raise ValueError("coverage_threshold must be between 0 and 1")
|
||||
if max_workers < 1:
|
||||
raise ValueError("max_workers must be at least 1")
|
||||
self.source = source
|
||||
self.snapshots = snapshots
|
||||
self.repository = repository
|
||||
self.coverage_threshold = coverage_threshold
|
||||
self.lock_key = lock_key
|
||||
self.today = today or date.today()
|
||||
self.max_workers = max_workers
|
||||
|
||||
def execute(self, command: SyncMarketDataCommand | None = None) -> SyncBatchSummary:
|
||||
"""Run one synchronization and retain successful items on partial failure."""
|
||||
|
||||
command = command or SyncMarketDataCommand()
|
||||
with self.repository.advisory_lock(self.lock_key) as acquired:
|
||||
if not acquired:
|
||||
return SyncBatchSummary(
|
||||
batch_id=None,
|
||||
target_trade_date=None,
|
||||
window=None,
|
||||
status="failed",
|
||||
target_count=0,
|
||||
valid_count=0,
|
||||
coverage=Decimal("0"),
|
||||
strategy_eligible=False,
|
||||
failures=(
|
||||
SyncFailure("batch", "lock", "sync_locked", "another sync is running"),
|
||||
),
|
||||
)
|
||||
return self._execute_locked(command)
|
||||
try:
|
||||
with self.repository.advisory_lock(self.lock_key) as acquired:
|
||||
if not acquired:
|
||||
return SyncBatchSummary(
|
||||
batch_id=None,
|
||||
target_trade_date=None,
|
||||
window=None,
|
||||
status="failed",
|
||||
target_count=0,
|
||||
valid_count=0,
|
||||
coverage=Decimal("0"),
|
||||
strategy_eligible=False,
|
||||
failures=(
|
||||
SyncFailure("batch", "lock", "sync_locked", "another sync is running"),
|
||||
),
|
||||
)
|
||||
return self._execute_locked(command)
|
||||
except Exception as exc:
|
||||
# Connection/pool failures while acquiring the advisory lock must
|
||||
# still produce the CLI's infrastructure-failure exit code.
|
||||
return SyncBatchSummary(
|
||||
batch_id=None,
|
||||
target_trade_date=command.target_trade_date,
|
||||
window=None,
|
||||
status="failed",
|
||||
target_count=0,
|
||||
valid_count=0,
|
||||
coverage=Decimal("0"),
|
||||
strategy_eligible=False,
|
||||
failures=(self._failure("batch", "lock", exc),),
|
||||
)
|
||||
|
||||
def _execute_locked(self, command: SyncMarketDataCommand) -> SyncBatchSummary:
|
||||
started_at = time.monotonic()
|
||||
@@ -145,7 +185,20 @@ class SyncMarketData:
|
||||
command.mode,
|
||||
command.target_trade_date or "auto",
|
||||
)
|
||||
target_trade_date = self._resolve_target(command.target_trade_date)
|
||||
try:
|
||||
target_trade_date = self._resolve_target(command.target_trade_date)
|
||||
except Exception as exc:
|
||||
return SyncBatchSummary(
|
||||
None,
|
||||
command.target_trade_date,
|
||||
None,
|
||||
"failed",
|
||||
0,
|
||||
0,
|
||||
Decimal("0"),
|
||||
False,
|
||||
failures=(self._failure("batch", "target", exc),),
|
||||
)
|
||||
window = SyncWindow.from_target(target_trade_date)
|
||||
logger.info(
|
||||
"market_data_sync_target target_trade_date=%s window_start=%s window_end=%s",
|
||||
@@ -153,7 +206,20 @@ class SyncMarketData:
|
||||
window.start,
|
||||
window.end,
|
||||
)
|
||||
all_stocks = filter_current_hs_a_stocks(self.source.fetch_stocks())
|
||||
try:
|
||||
all_stocks = filter_current_hs_a_stocks(self.source.fetch_stocks())
|
||||
except Exception as exc:
|
||||
return SyncBatchSummary(
|
||||
None,
|
||||
target_trade_date,
|
||||
window,
|
||||
"failed",
|
||||
0,
|
||||
0,
|
||||
Decimal("0"),
|
||||
False,
|
||||
failures=(self._failure("stock", "universe", exc),),
|
||||
)
|
||||
if not all_stocks:
|
||||
return SyncBatchSummary(
|
||||
None,
|
||||
@@ -168,13 +234,44 @@ class SyncMarketData:
|
||||
SyncFailure("stock", "universe", "empty_universe", "no eligible stocks"),
|
||||
),
|
||||
)
|
||||
batch_id = self.repository.create_batch(
|
||||
target_trade_date,
|
||||
window,
|
||||
command.mode,
|
||||
command.parent_batch_id,
|
||||
len(all_stocks),
|
||||
)
|
||||
try:
|
||||
retry_items: set[tuple[str, str]] = (
|
||||
self._retry_items(command.parent_batch_id) if command.mode == "retry" else set()
|
||||
)
|
||||
dates = self._dates_to_process(window, target_trade_date, command.mode, retry_items)
|
||||
except Exception as exc:
|
||||
return SyncBatchSummary(
|
||||
None,
|
||||
target_trade_date,
|
||||
window,
|
||||
"failed",
|
||||
len(all_stocks),
|
||||
0,
|
||||
Decimal("0"),
|
||||
False,
|
||||
failures=(self._failure("batch", "prepare", exc),),
|
||||
)
|
||||
try:
|
||||
batch_id = self.repository.create_batch(
|
||||
target_trade_date,
|
||||
window,
|
||||
command.mode,
|
||||
command.parent_batch_id,
|
||||
len(all_stocks),
|
||||
)
|
||||
except Exception as exc:
|
||||
failure = self._failure("batch", "create", exc)
|
||||
return SyncBatchSummary(
|
||||
None,
|
||||
target_trade_date,
|
||||
window,
|
||||
"failed",
|
||||
len(all_stocks),
|
||||
0,
|
||||
Decimal("0"),
|
||||
False,
|
||||
failures=(failure,),
|
||||
)
|
||||
logger.info(
|
||||
"market_data_sync_batch_created batch_id=%s target_count=%d",
|
||||
batch_id,
|
||||
@@ -182,12 +279,19 @@ class SyncMarketData:
|
||||
)
|
||||
failures: list[SyncFailure] = []
|
||||
totals = [0, 0, 0]
|
||||
pending_outcomes: list[SyncItemOutcome] = []
|
||||
audit_failed = False
|
||||
stock_codes = {stock.ts_code for stock in all_stocks}
|
||||
retry_items: set[tuple[str, str]] = (
|
||||
self._retry_items(command.parent_batch_id) if command.mode == "retry" else set()
|
||||
)
|
||||
|
||||
self._process_stock_master(batch_id, all_stocks, failures, totals)
|
||||
stock_outcome = self._process_stock_master(all_stocks)
|
||||
audit_failed = self._consume_outcome(
|
||||
batch_id,
|
||||
stock_outcome,
|
||||
pending_outcomes,
|
||||
failures,
|
||||
totals,
|
||||
audit_failed,
|
||||
)
|
||||
self._log_progress(
|
||||
batch_id=batch_id,
|
||||
stage="stock_master",
|
||||
@@ -199,7 +303,6 @@ class SyncMarketData:
|
||||
started_at=started_at,
|
||||
force=bool(failures),
|
||||
)
|
||||
dates = self._dates_to_process(window, target_trade_date, command.mode, retry_items)
|
||||
logger.info(
|
||||
"market_data_sync_stage_started batch_id=%s stage=daily_basic total=%d",
|
||||
batch_id,
|
||||
@@ -207,7 +310,19 @@ class SyncMarketData:
|
||||
)
|
||||
for current, trade_date in enumerate(dates, start=1):
|
||||
failure_count = len(failures)
|
||||
self._process_daily_basic(batch_id, trade_date, stock_codes, window, failures, totals)
|
||||
daily_basic_outcome = self._process_daily_basic(
|
||||
trade_date,
|
||||
stock_codes,
|
||||
window,
|
||||
)
|
||||
audit_failed = self._consume_outcome(
|
||||
batch_id,
|
||||
daily_basic_outcome,
|
||||
pending_outcomes,
|
||||
failures,
|
||||
totals,
|
||||
audit_failed,
|
||||
)
|
||||
self._log_progress(
|
||||
batch_id=batch_id,
|
||||
stage="daily_basic",
|
||||
@@ -234,19 +349,55 @@ class SyncMarketData:
|
||||
batch_id,
|
||||
len(stocks_to_process),
|
||||
)
|
||||
for current, stock in enumerate(stocks_to_process, start=1):
|
||||
failure_count = len(failures)
|
||||
self._process_bar(batch_id, stock, window, failures, totals)
|
||||
self._log_progress(
|
||||
batch_id=batch_id,
|
||||
stage="bar",
|
||||
current=current,
|
||||
total=len(stocks_to_process),
|
||||
item_key=stock.ts_code,
|
||||
totals=totals,
|
||||
failures=failures,
|
||||
started_at=started_at,
|
||||
force=len(failures) > failure_count,
|
||||
try:
|
||||
with ThreadPoolExecutor(
|
||||
max_workers=self.max_workers,
|
||||
thread_name_prefix="market-data-bar",
|
||||
) as executor:
|
||||
futures = {
|
||||
executor.submit(self._process_bar, stock, window): stock.ts_code
|
||||
for stock in stocks_to_process
|
||||
}
|
||||
for current, future in enumerate(as_completed(futures), start=1):
|
||||
item_key = futures[future]
|
||||
failure_count = len(failures)
|
||||
try:
|
||||
bar_outcome = future.result()
|
||||
except Exception as exc: # pragma: no cover - defensive future boundary
|
||||
bar_outcome = SyncItemOutcome(
|
||||
"bar",
|
||||
item_key,
|
||||
"failed",
|
||||
failure=self._failure("bar", item_key, exc),
|
||||
)
|
||||
audit_failed = self._consume_outcome(
|
||||
batch_id,
|
||||
bar_outcome,
|
||||
pending_outcomes,
|
||||
failures,
|
||||
totals,
|
||||
audit_failed,
|
||||
)
|
||||
self._log_progress(
|
||||
batch_id=batch_id,
|
||||
stage="bar",
|
||||
current=current,
|
||||
total=len(stocks_to_process),
|
||||
item_key=bar_outcome.item_key,
|
||||
totals=totals,
|
||||
failures=failures,
|
||||
started_at=started_at,
|
||||
force=len(failures) > failure_count,
|
||||
)
|
||||
except Exception as exc:
|
||||
# Executor construction/submission is a batch-level failure. Any
|
||||
# facts committed by already completed workers remain committed.
|
||||
failure = self._failure("batch", "executor", exc)
|
||||
failures.append(failure)
|
||||
logger.exception(
|
||||
"market_data_sync_executor_failed batch_id=%s error_type=%s",
|
||||
batch_id,
|
||||
failure.error_type,
|
||||
)
|
||||
logger.info(
|
||||
"market_data_sync_stage_completed batch_id=%s stage=bar total=%d",
|
||||
@@ -254,32 +405,62 @@ class SyncMarketData:
|
||||
len(stocks_to_process),
|
||||
)
|
||||
|
||||
if pending_outcomes:
|
||||
audit_failure = self._record_outcomes(batch_id, pending_outcomes)
|
||||
pending_outcomes.clear()
|
||||
if audit_failure is not None and not audit_failed:
|
||||
failures.append(audit_failure)
|
||||
audit_failed = True
|
||||
|
||||
if not failures:
|
||||
try:
|
||||
self.repository.purge_before(window)
|
||||
self.snapshots.clean_daily_basic_before(window.start)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
except Exception as exc:
|
||||
failure = self._failure("batch", "retention", exc)
|
||||
failures.append(failure)
|
||||
self.repository.record_item(
|
||||
audit_failure = self._record_outcomes(
|
||||
batch_id,
|
||||
"batch",
|
||||
"retention",
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
[
|
||||
SyncItemOutcome(
|
||||
"batch",
|
||||
"retention",
|
||||
"failed",
|
||||
failure=failure,
|
||||
)
|
||||
],
|
||||
)
|
||||
valid_count = sum(
|
||||
1
|
||||
for stock in all_stocks
|
||||
if self.repository.has_bar(stock.ts_code, target_trade_date)
|
||||
and self.repository.has_daily_basic(stock.ts_code, target_trade_date)
|
||||
)
|
||||
if audit_failure is not None and not audit_failed:
|
||||
failures.append(audit_failure)
|
||||
audit_failed = True
|
||||
try:
|
||||
count_valid_stocks = getattr(self.repository, "count_valid_stocks", None)
|
||||
if callable(count_valid_stocks):
|
||||
count_fn = cast(Callable[[date], int], count_valid_stocks)
|
||||
valid_count = int(count_fn(target_trade_date))
|
||||
else:
|
||||
# Compatibility fallback for older in-memory adapters. The
|
||||
# PostgreSQL adapter always takes the set-based path above.
|
||||
valid_count = sum(
|
||||
1
|
||||
for stock in all_stocks
|
||||
if self.repository.has_bar(stock.ts_code, target_trade_date)
|
||||
and self.repository.has_daily_basic(stock.ts_code, target_trade_date)
|
||||
)
|
||||
except Exception as exc:
|
||||
failure = self._failure("batch", "coverage", exc)
|
||||
failures.append(failure)
|
||||
valid_count = 0
|
||||
coverage = Decimal(valid_count) / Decimal(len(all_stocks))
|
||||
status = "success" if not failures else "partial_success" if valid_count else "failed"
|
||||
eligible = coverage >= self.coverage_threshold
|
||||
self.repository.record_batch(batch_id, status, valid_count, coverage, eligible)
|
||||
try:
|
||||
self.repository.record_batch(batch_id, status, valid_count, coverage, eligible)
|
||||
except Exception as exc:
|
||||
failure = self._failure("batch", "record", exc)
|
||||
failures.append(failure)
|
||||
status = "failed"
|
||||
eligible = False
|
||||
logger.info(
|
||||
"market_data_sync_finished batch_id=%s status=%s target_count=%d valid_count=%d "
|
||||
"coverage=%s failures=%d inserted=%d updated=%d unchanged=%d elapsed_seconds=%.1f",
|
||||
@@ -294,6 +475,12 @@ class SyncMarketData:
|
||||
totals[2],
|
||||
time.monotonic() - started_at,
|
||||
)
|
||||
ordered_failures = tuple(
|
||||
sorted(
|
||||
failures,
|
||||
key=lambda failure: (failure.item_kind, failure.item_key, failure.error_type),
|
||||
)
|
||||
)
|
||||
return SyncBatchSummary(
|
||||
batch_id,
|
||||
target_trade_date,
|
||||
@@ -306,7 +493,7 @@ class SyncMarketData:
|
||||
totals[0],
|
||||
totals[1],
|
||||
totals[2],
|
||||
tuple(failures),
|
||||
ordered_failures,
|
||||
)
|
||||
|
||||
def _resolve_target(self, requested: date | None) -> date:
|
||||
@@ -346,57 +533,37 @@ class SyncMarketData:
|
||||
|
||||
def _process_stock_master(
|
||||
self,
|
||||
batch_id: str,
|
||||
stocks: Sequence[Stock],
|
||||
failures: list[SyncFailure],
|
||||
totals: list[int],
|
||||
) -> None:
|
||||
) -> SyncItemOutcome:
|
||||
staged = None
|
||||
try:
|
||||
staged = self.snapshots.stage_stocks(stocks)
|
||||
result = self.repository.upsert_stocks(stocks)
|
||||
self.snapshots.publish(staged)
|
||||
totals[0] += result.inserted
|
||||
totals[1] += result.updated
|
||||
totals[2] += result.unchanged
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
return SyncItemOutcome(
|
||||
"stock",
|
||||
"current",
|
||||
"success",
|
||||
result,
|
||||
staged.fingerprint,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
except Exception as exc:
|
||||
if staged is not None:
|
||||
self.snapshots.discard(staged)
|
||||
failure = self._failure("stock", "current", exc)
|
||||
failures.append(failure)
|
||||
logger.warning(
|
||||
"market_data_sync_item_failed batch_id=%s stage=stock_master item=current "
|
||||
"error_type=%s",
|
||||
batch_id,
|
||||
failure.error_type,
|
||||
)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
return SyncItemOutcome(
|
||||
"stock",
|
||||
"current",
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
failure=failure,
|
||||
)
|
||||
|
||||
def _process_daily_basic(
|
||||
self,
|
||||
batch_id: str,
|
||||
trade_date: date,
|
||||
stock_codes: set[str],
|
||||
window: SyncWindow,
|
||||
failures: list[SyncFailure],
|
||||
totals: list[int],
|
||||
) -> None:
|
||||
) -> SyncItemOutcome:
|
||||
staged = None
|
||||
key = trade_date.isoformat()
|
||||
try:
|
||||
@@ -410,50 +577,35 @@ class SyncMarketData:
|
||||
staged = self.snapshots.stage_daily_basic(trade_date, rows)
|
||||
result = self.repository.upsert_daily_basic(rows, window)
|
||||
self.snapshots.publish(staged)
|
||||
self._add_counts(totals, result)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
return SyncItemOutcome(
|
||||
"daily_basic",
|
||||
key,
|
||||
"success",
|
||||
result,
|
||||
staged.fingerprint,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
except Exception as exc:
|
||||
if staged is not None:
|
||||
self.snapshots.discard(staged)
|
||||
failure = self._failure("daily_basic", key, exc)
|
||||
failures.append(failure)
|
||||
logger.warning(
|
||||
"market_data_sync_item_failed batch_id=%s stage=daily_basic item=%s error_type=%s",
|
||||
batch_id,
|
||||
key,
|
||||
failure.error_type,
|
||||
)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
return SyncItemOutcome(
|
||||
"daily_basic",
|
||||
key,
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
failure=failure,
|
||||
)
|
||||
|
||||
def _process_bar(
|
||||
self,
|
||||
batch_id: str,
|
||||
stock: Stock,
|
||||
window: SyncWindow,
|
||||
failures: list[SyncFailure],
|
||||
totals: list[int],
|
||||
) -> None:
|
||||
) -> SyncItemOutcome:
|
||||
staged = None
|
||||
try:
|
||||
rows = tuple(self.source.fetch_bars(stock.ts_code, window))
|
||||
staged = self.snapshots.stage_bars(stock.ts_code, rows)
|
||||
old_rows = self.snapshots.read_bars(stock.ts_code)
|
||||
comparison = compare_snapshots(old_rows, rows)
|
||||
staged = self.snapshots.stage_bars(stock.ts_code, rows)
|
||||
if comparison.change is SnapshotChange.UNCHANGED:
|
||||
result = WriteResult(unchanged=len(rows))
|
||||
else:
|
||||
@@ -468,34 +620,22 @@ class SyncMarketData:
|
||||
in {SnapshotChange.INITIAL, SnapshotChange.CHANGED},
|
||||
)
|
||||
self.snapshots.publish(staged)
|
||||
self._add_counts(totals, result)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
return SyncItemOutcome(
|
||||
"bar",
|
||||
stock.ts_code,
|
||||
"success",
|
||||
result,
|
||||
staged.fingerprint,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
except Exception as exc:
|
||||
if staged is not None:
|
||||
self.snapshots.discard(staged)
|
||||
failure = self._failure("bar", stock.ts_code, exc)
|
||||
failures.append(failure)
|
||||
logger.warning(
|
||||
"market_data_sync_item_failed batch_id=%s stage=bar item=%s error_type=%s",
|
||||
batch_id,
|
||||
stock.ts_code,
|
||||
failure.error_type,
|
||||
)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
return SyncItemOutcome(
|
||||
"bar",
|
||||
stock.ts_code,
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
failure=failure,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -504,6 +644,68 @@ class SyncMarketData:
|
||||
totals[1] += result.updated
|
||||
totals[2] += result.unchanged
|
||||
|
||||
def _consume_outcome(
|
||||
self,
|
||||
batch_id: str,
|
||||
outcome: SyncItemOutcome,
|
||||
pending_outcomes: list[SyncItemOutcome],
|
||||
failures: list[SyncFailure],
|
||||
totals: list[int],
|
||||
audit_failed: bool,
|
||||
) -> bool:
|
||||
"""Aggregate one completed item and flush audit rows in bounded batches."""
|
||||
|
||||
if outcome.failure is not None:
|
||||
failures.append(outcome.failure)
|
||||
logger.warning(
|
||||
"market_data_sync_item_failed batch_id=%s stage=%s item=%s error_type=%s",
|
||||
batch_id,
|
||||
outcome.item_kind,
|
||||
outcome.item_key,
|
||||
outcome.failure.error_type,
|
||||
)
|
||||
else:
|
||||
self._add_counts(totals, outcome.result)
|
||||
pending_outcomes.append(outcome)
|
||||
if len(pending_outcomes) < _PROGRESS_LOG_INTERVAL:
|
||||
return audit_failed
|
||||
audit_failure = self._record_outcomes(batch_id, pending_outcomes)
|
||||
pending_outcomes.clear()
|
||||
if audit_failure is not None and not audit_failed:
|
||||
failures.append(audit_failure)
|
||||
return True
|
||||
return audit_failed
|
||||
|
||||
def _record_outcomes(
|
||||
self,
|
||||
batch_id: str,
|
||||
outcomes: Sequence[SyncItemOutcome],
|
||||
) -> SyncFailure | None:
|
||||
"""Persist a batch of audit outcomes, with a legacy adapter fallback."""
|
||||
|
||||
if not outcomes:
|
||||
return None
|
||||
try:
|
||||
record_items = getattr(self.repository, "record_items", None)
|
||||
if callable(record_items):
|
||||
record_items(batch_id, tuple(outcomes))
|
||||
else:
|
||||
for outcome in outcomes:
|
||||
failure = outcome.failure
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
outcome.item_kind,
|
||||
outcome.item_key,
|
||||
outcome.status,
|
||||
outcome.result,
|
||||
outcome.fingerprint,
|
||||
failure.error_type if failure else None,
|
||||
failure.message if failure else None,
|
||||
)
|
||||
except Exception as exc:
|
||||
return self._failure("batch", "audit", exc)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _log_progress(
|
||||
*,
|
||||
|
||||
@@ -1,12 +1,32 @@
|
||||
"""Pure market data domain types and ports."""
|
||||
|
||||
from .fingerprint import SnapshotChange, SnapshotComparison, compare_snapshots, snapshot_fingerprint
|
||||
from .integrity import (
|
||||
IntegrityCheckInProgress,
|
||||
IntegrityCheckNoData,
|
||||
IntegrityCheckNotFound,
|
||||
IntegrityCheckPage,
|
||||
IntegrityCheckQuery,
|
||||
IntegrityCheckRun,
|
||||
IntegrityCheckStoreError,
|
||||
IntegrityIssue,
|
||||
IntegrityStatus,
|
||||
)
|
||||
from .models import Bar, DailyBasic, Stock, SyncWindow
|
||||
from .rules import filter_current_hs_a_stocks, is_current_hs_a_stock
|
||||
|
||||
__all__ = [
|
||||
"Bar",
|
||||
"DailyBasic",
|
||||
"IntegrityCheckInProgress",
|
||||
"IntegrityCheckNoData",
|
||||
"IntegrityCheckNotFound",
|
||||
"IntegrityCheckPage",
|
||||
"IntegrityCheckQuery",
|
||||
"IntegrityCheckRun",
|
||||
"IntegrityCheckStoreError",
|
||||
"IntegrityIssue",
|
||||
"IntegrityStatus",
|
||||
"SnapshotChange",
|
||||
"SnapshotComparison",
|
||||
"Stock",
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
"""Domain contracts for read-only market-data integrity checks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Sequence
|
||||
from contextlib import AbstractContextManager
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime
|
||||
from hashlib import sha256
|
||||
from typing import Literal, Protocol
|
||||
|
||||
from .models import Bar, DailyBasic, Stock, SyncWindow
|
||||
|
||||
IntegrityStatus = Literal["running", "passed", "issues_found", "failed"]
|
||||
|
||||
|
||||
class IntegrityCheckInProgress(RuntimeError):
|
||||
"""Raised when an integrity check is already running."""
|
||||
|
||||
|
||||
class IntegrityCheckNoData(RuntimeError):
|
||||
"""Raised when no completed market-data window is available to check."""
|
||||
|
||||
|
||||
class IntegrityCheckNotFound(RuntimeError):
|
||||
"""Raised when a requested check id does not exist."""
|
||||
|
||||
|
||||
class IntegrityCheckStoreError(RuntimeError):
|
||||
"""Raised when check metadata or issue records cannot be read or written."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class IntegrityIssue:
|
||||
"""One stable, safe integrity discrepancy.
|
||||
|
||||
``issue_key`` is derived from the object and issue category instead of a
|
||||
provider error string. This keeps pagination and retries deterministic
|
||||
while allowing the human-readable message to evolve independently.
|
||||
"""
|
||||
|
||||
check_id: str
|
||||
issue_key: str
|
||||
item_kind: str
|
||||
item_key: str
|
||||
issue_type: str
|
||||
message: str
|
||||
created_at: datetime | None = None
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
check_id: str,
|
||||
item_kind: str,
|
||||
item_key: str,
|
||||
issue_type: str,
|
||||
message: str,
|
||||
) -> IntegrityIssue:
|
||||
"""Build a deterministic issue key from stable comparison identity."""
|
||||
|
||||
identity = "\x00".join((item_kind, item_key, issue_type)).encode("utf-8")
|
||||
issue_key = sha256(identity).hexdigest()
|
||||
return cls(
|
||||
check_id=check_id,
|
||||
issue_key=issue_key,
|
||||
item_kind=item_kind,
|
||||
item_key=item_key,
|
||||
issue_type=issue_type,
|
||||
message=" ".join(message.split())[:500],
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class IntegrityCheckRun:
|
||||
"""Persisted progress and terminal state for one integrity check."""
|
||||
|
||||
id: str
|
||||
status: IntegrityStatus
|
||||
window: SyncWindow | None
|
||||
target_count: int
|
||||
checked_count: int = 0
|
||||
issue_count: int = 0
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
created_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
|
||||
@property
|
||||
def window_start(self) -> date | None:
|
||||
"""Return the inclusive window start for response adapters."""
|
||||
|
||||
return self.window.start if self.window is not None else None
|
||||
|
||||
@property
|
||||
def window_end(self) -> date | None:
|
||||
"""Return the inclusive window end for response adapters."""
|
||||
|
||||
return self.window.end if self.window is not None else None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class IntegrityCheckQuery:
|
||||
"""Validated issue pagination passed from HTTP or another caller."""
|
||||
|
||||
page: int = 1
|
||||
page_size: int = 10
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.page < 1:
|
||||
raise ValueError("page must be at least 1")
|
||||
if not 1 <= self.page_size <= 100:
|
||||
raise ValueError("page_size must be between 1 and 100")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class IntegrityCheckPage:
|
||||
"""One check run plus a bounded page of persisted issues."""
|
||||
|
||||
run: IntegrityCheckRun
|
||||
page: int
|
||||
page_size: int
|
||||
issues_total: int
|
||||
issues: tuple[IntegrityIssue, ...] = ()
|
||||
|
||||
@property
|
||||
def check(self) -> IntegrityCheckRun:
|
||||
"""Alias used by callers that name the aggregate ``check``."""
|
||||
|
||||
return self.run
|
||||
|
||||
|
||||
class IntegritySnapshotReader(Protocol):
|
||||
"""Read-only PostgreSQL and CSV comparison input port."""
|
||||
|
||||
def latest_successful_window(self) -> SyncWindow | None: ...
|
||||
|
||||
def active_stocks(self) -> Sequence[Stock]: ...
|
||||
|
||||
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]: ...
|
||||
|
||||
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]: ...
|
||||
|
||||
def list_bar_codes(self, window: SyncWindow) -> Sequence[str]: ...
|
||||
|
||||
def list_daily_basic_dates(self, window: SyncWindow) -> Sequence[date]: ...
|
||||
|
||||
def advisory_lock(self, key: int) -> AbstractContextManager[bool]: ...
|
||||
|
||||
|
||||
class IntegritySnapshotStore(Protocol):
|
||||
"""Read-only formal CSV snapshot port."""
|
||||
|
||||
def read_stocks(self) -> tuple[Stock, ...] | None: ...
|
||||
|
||||
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None: ...
|
||||
|
||||
def read_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...] | None: ...
|
||||
|
||||
def list_bar_codes(self, window: SyncWindow) -> Sequence[str]: ...
|
||||
|
||||
def list_daily_basic_dates(self, window: SyncWindow) -> Sequence[date]: ...
|
||||
|
||||
|
||||
class IntegrityCheckStore(Protocol):
|
||||
"""Persistence port that writes only integrity metadata and issues."""
|
||||
|
||||
def recover_stale_running(self, lock_key: int, stale_after_seconds: int) -> int: ...
|
||||
|
||||
def create_running(self, window: SyncWindow, target_count: int) -> IntegrityCheckRun: ...
|
||||
|
||||
def update_progress(self, check_id: str, checked_count: int, issue_count: int) -> None: ...
|
||||
|
||||
def record_issues(self, issues: Iterable[IntegrityIssue]) -> None: ...
|
||||
|
||||
def finish(
|
||||
self,
|
||||
check_id: str,
|
||||
status: IntegrityStatus,
|
||||
*,
|
||||
issue_count: int,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> None: ...
|
||||
|
||||
def get(self, check_id: str, query: IntegrityCheckQuery) -> IntegrityCheckPage | None: ...
|
||||
|
||||
def get_latest(self, query: IntegrityCheckQuery) -> IntegrityCheckPage | None: ...
|
||||
@@ -46,6 +46,8 @@ def decimal_text(value: Decimal | None) -> str:
|
||||
|
||||
if value is None:
|
||||
return ""
|
||||
if not value.is_finite():
|
||||
raise ValueError("numeric value must be finite")
|
||||
text = format(value, "f")
|
||||
if "." in text:
|
||||
text = text.rstrip("0").rstrip(".")
|
||||
|
||||
@@ -124,3 +124,17 @@ class MarketDataRepository(Protocol):
|
||||
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]: ...
|
||||
|
||||
def advisory_lock(self, key: int) -> AbstractContextManager[bool]: ...
|
||||
|
||||
|
||||
class MarketDataBatchRepository(Protocol):
|
||||
"""Optional batch APIs used by the optimized PostgreSQL adapter.
|
||||
|
||||
Keeping these methods in a refinement protocol preserves compatibility
|
||||
with small in-memory repositories used by legacy application tests. The
|
||||
production PostgreSQL repository implements both this protocol and
|
||||
``MarketDataRepository``.
|
||||
"""
|
||||
|
||||
def count_valid_stocks(self, trade_date: date) -> int: ...
|
||||
|
||||
def record_items(self, batch_id: str, outcomes: Iterable[object]) -> None: ...
|
||||
|
||||
+172
-6
@@ -55,12 +55,50 @@ STOCK_COLUMNS = ("ts_code", "name", "market", "exchange", "list_status", "list_d
|
||||
_DAILY_DECIMAL_FIELDS = DAILY_BASIC_COLUMNS[2:]
|
||||
|
||||
|
||||
class SnapshotReadError(ValueError):
|
||||
"""A formal CSV could not be parsed without exposing file internals."""
|
||||
|
||||
issue_type = "parse_error"
|
||||
|
||||
|
||||
class SnapshotDuplicateKeyError(SnapshotReadError):
|
||||
"""A formal CSV contains more than one row for a business key."""
|
||||
|
||||
issue_type = "duplicate_key"
|
||||
|
||||
|
||||
class SnapshotWindowError(SnapshotReadError):
|
||||
"""A CSV path or row does not match its declared snapshot window."""
|
||||
|
||||
issue_type = "window_out_of_bounds"
|
||||
|
||||
|
||||
def _code_path(ts_code: str) -> str:
|
||||
if not _SAFE_CODE.fullmatch(ts_code):
|
||||
raise ValueError(f"unsupported stock code: {ts_code!r}")
|
||||
return ts_code
|
||||
|
||||
|
||||
def _require_header(fieldnames: Iterable[str] | None, expected: tuple[str, ...]) -> None:
|
||||
"""Reject schema drift before a row can be mistaken for valid data."""
|
||||
|
||||
if tuple(fieldnames or ()) != expected:
|
||||
raise SnapshotReadError("snapshot header does not match the market-data contract")
|
||||
|
||||
|
||||
def _date_from_snapshot_path(path: Path) -> date:
|
||||
"""Parse ``daily-basic/YYYY/YYYYMMDD.csv`` path components."""
|
||||
|
||||
if path.parent.parent.name == "daily-basic" and len(path.parent.name) == 4:
|
||||
text = path.stem
|
||||
if len(text) == 8 and text.isdigit() and text[:4] == path.parent.name:
|
||||
try:
|
||||
return date.fromisoformat(f"{text[:4]}-{text[4:6]}-{text[6:]}")
|
||||
except ValueError as exc:
|
||||
raise SnapshotWindowError("daily-basic path contains an invalid date") from exc
|
||||
raise SnapshotWindowError("daily-basic path does not match the formal layout")
|
||||
|
||||
|
||||
def _bar_row(row: Bar) -> dict[str, str]:
|
||||
return {
|
||||
"ts_code": row.ts_code,
|
||||
@@ -125,19 +163,139 @@ class CsvSnapshotStore:
|
||||
|
||||
return self.root / "stock-basic" / "current.csv"
|
||||
|
||||
def list_bar_snapshot_files(self) -> tuple[Path, ...]:
|
||||
"""List formal bar files without opening or creating any file."""
|
||||
|
||||
root = self.root / "bars"
|
||||
if not root.exists():
|
||||
return ()
|
||||
return tuple(sorted(root.glob("*.csv"), key=lambda path: path.name))
|
||||
|
||||
def list_daily_basic_snapshot_files(self) -> tuple[Path, ...]:
|
||||
"""List formal daily-basic files without loading their rows."""
|
||||
|
||||
root = self.root / "daily-basic"
|
||||
if not root.exists():
|
||||
return ()
|
||||
return tuple(sorted(root.glob("*/*.csv"), key=lambda path: str(path)))
|
||||
|
||||
def list_bar_codes(self, window: object | None = None) -> tuple[str, ...]:
|
||||
"""Return stock codes represented by formal bar paths.
|
||||
|
||||
The optional window is accepted for the integrity reader protocol; the
|
||||
path itself has no date, so row-level window validation happens while
|
||||
the file is read.
|
||||
"""
|
||||
|
||||
del window
|
||||
return tuple(path.stem for path in self.list_bar_snapshot_files())
|
||||
|
||||
def list_daily_basic_dates(self, window: object | None = None) -> tuple[date, ...]:
|
||||
"""Return parseable dates represented by formal daily-basic paths."""
|
||||
|
||||
del window
|
||||
result: list[date] = []
|
||||
for path in self.list_daily_basic_snapshot_files():
|
||||
try:
|
||||
result.append(_date_from_snapshot_path(path))
|
||||
except ValueError:
|
||||
continue
|
||||
return tuple(sorted(set(result)))
|
||||
|
||||
def list_invalid_daily_basic_snapshot_files(self) -> tuple[tuple[str, str], ...]:
|
||||
"""Return malformed daily-basic paths without opening their contents.
|
||||
|
||||
A path that cannot identify a valid ``YYYY/YYYYMMDD`` date is itself
|
||||
an integrity discrepancy. The normal date listing intentionally
|
||||
skips such paths so callers can still inspect every valid date; the
|
||||
integrity application consumes this explicit error listing separately.
|
||||
"""
|
||||
|
||||
errors: list[tuple[str, str]] = []
|
||||
for path in self.list_daily_basic_snapshot_files():
|
||||
try:
|
||||
_date_from_snapshot_path(path)
|
||||
except SnapshotReadError as exc:
|
||||
relative_path = path.relative_to(self.root).as_posix()
|
||||
errors.append((relative_path, getattr(exc, "issue_type", "parse_error")))
|
||||
return tuple(errors)
|
||||
|
||||
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None:
|
||||
"""Read the current formal snapshot, returning ``None`` if absent."""
|
||||
|
||||
path = self.bars_path(ts_code)
|
||||
if not path.exists():
|
||||
return None
|
||||
with path.open(newline="", encoding="utf-8") as handle:
|
||||
reader = csv.DictReader(handle)
|
||||
if tuple(reader.fieldnames or ()) != BAR_COLUMNS:
|
||||
raise ValueError(f"unexpected bar CSV header: {path}")
|
||||
rows = tuple(Bar.from_mapping(row) for row in reader)
|
||||
try:
|
||||
with path.open(newline="", encoding="utf-8") as handle:
|
||||
reader = csv.DictReader(handle)
|
||||
_require_header(reader.fieldnames, BAR_COLUMNS)
|
||||
rows = tuple(Bar.from_mapping(row) for row in reader)
|
||||
except SnapshotReadError:
|
||||
raise
|
||||
except (OSError, csv.Error, TypeError, ValueError) as exc:
|
||||
raise SnapshotReadError(f"bar snapshot cannot be parsed for {ts_code}") from exc
|
||||
if not rows:
|
||||
raise SnapshotReadError(f"bar snapshot is empty for {ts_code}")
|
||||
keys = [(row.ts_code, row.trade_date) for row in rows]
|
||||
if any(row.ts_code != ts_code for row in rows):
|
||||
raise SnapshotReadError(f"bar snapshot contains an unexpected stock code: {ts_code}")
|
||||
if len(keys) != len(set(keys)):
|
||||
raise SnapshotDuplicateKeyError(f"bar snapshot contains duplicate keys: {ts_code}")
|
||||
return tuple(sorted(rows, key=lambda row: row.trade_date))
|
||||
|
||||
def read_stocks(self) -> tuple[Stock, ...] | None:
|
||||
"""Read the formal current stock-master file without publishing."""
|
||||
|
||||
path = self.stock_basic_path
|
||||
if not path.exists():
|
||||
return None
|
||||
try:
|
||||
with path.open(newline="", encoding="utf-8") as handle:
|
||||
reader = csv.DictReader(handle)
|
||||
_require_header(reader.fieldnames, STOCK_COLUMNS)
|
||||
rows = tuple(Stock.from_mapping(row) for row in reader)
|
||||
except SnapshotReadError:
|
||||
raise
|
||||
except (OSError, csv.Error, TypeError, ValueError) as exc:
|
||||
raise SnapshotReadError("stock-master snapshot cannot be parsed") from exc
|
||||
if not rows:
|
||||
raise SnapshotReadError("stock-master snapshot is empty")
|
||||
keys = [row.ts_code for row in rows]
|
||||
if len(keys) != len(set(keys)):
|
||||
raise SnapshotDuplicateKeyError("stock-master snapshot contains duplicate codes")
|
||||
return tuple(sorted(rows, key=lambda row: row.ts_code))
|
||||
|
||||
def read_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...] | None:
|
||||
"""Read one formal daily-basic date snapshot without side effects."""
|
||||
|
||||
path = self.daily_basic_path(trade_date)
|
||||
if not path.exists():
|
||||
return None
|
||||
try:
|
||||
with path.open(newline="", encoding="utf-8") as handle:
|
||||
reader = csv.DictReader(handle)
|
||||
_require_header(reader.fieldnames, DAILY_BASIC_COLUMNS)
|
||||
rows = tuple(DailyBasic.from_mapping(row) for row in reader)
|
||||
except SnapshotReadError:
|
||||
raise
|
||||
except (OSError, csv.Error, TypeError, ValueError) as exc:
|
||||
raise SnapshotReadError(
|
||||
f"daily-basic snapshot cannot be parsed for {trade_date}"
|
||||
) from exc
|
||||
if not rows:
|
||||
raise SnapshotReadError(f"daily-basic snapshot is empty for {trade_date}")
|
||||
keys = [(row.ts_code, row.trade_date) for row in rows]
|
||||
if any(row.trade_date != trade_date for row in rows):
|
||||
raise SnapshotWindowError(
|
||||
f"daily-basic snapshot contains an unexpected date: {trade_date}"
|
||||
)
|
||||
if len(keys) != len(set(keys)):
|
||||
raise SnapshotDuplicateKeyError(
|
||||
f"daily-basic snapshot contains duplicate keys: {trade_date}"
|
||||
)
|
||||
return tuple(sorted(rows, key=lambda row: row.ts_code))
|
||||
|
||||
def stage_bars(self, ts_code: str, rows: Iterable[Bar]) -> StagedSnapshot:
|
||||
"""Validate and write a temporary bar snapshot."""
|
||||
|
||||
@@ -146,7 +304,15 @@ class CsvSnapshotStore:
|
||||
raise ValueError("bar snapshot must be non-empty and belong to one stock")
|
||||
final_path = self.bars_path(ts_code)
|
||||
temp_path = self._write_csv(final_path, BAR_COLUMNS, (_bar_row(row) for row in normalized))
|
||||
return StagedSnapshot(temp_path, final_path, snapshot_fingerprint(normalized))
|
||||
try:
|
||||
fingerprint = snapshot_fingerprint(normalized)
|
||||
except BaseException:
|
||||
# The CSV has already been fully staged, so validation failures
|
||||
# (for example an invalid OHLC relation) must not leave a hidden
|
||||
# temporary file behind for a later run to mistake for a snapshot.
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
return StagedSnapshot(temp_path, final_path, fingerprint)
|
||||
|
||||
def stage_daily_basic(self, trade_date: date, rows: Iterable[DailyBasic]) -> StagedSnapshot:
|
||||
"""Validate and stage one daily-basic date snapshot."""
|
||||
|
||||
@@ -0,0 +1,422 @@
|
||||
"""PostgreSQL and CSV seams used by the read-only integrity application."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator, Iterable
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from typing import Any, Protocol, cast
|
||||
from uuid import uuid4
|
||||
|
||||
import psycopg
|
||||
|
||||
from ..domain.integrity import (
|
||||
IntegrityCheckInProgress,
|
||||
IntegrityCheckPage,
|
||||
IntegrityCheckQuery,
|
||||
IntegrityCheckRun,
|
||||
IntegrityCheckStoreError,
|
||||
IntegrityIssue,
|
||||
IntegrityStatus,
|
||||
)
|
||||
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||
from ..domain.ports import MarketDataRepositoryError
|
||||
|
||||
|
||||
class _PostgresRepository(Protocol):
|
||||
"""Small structural seam shared with the existing pooled repository."""
|
||||
|
||||
def connection(self) -> Any: ...
|
||||
|
||||
def open(self) -> None: ...
|
||||
|
||||
def close(self) -> None: ...
|
||||
|
||||
def advisory_lock(self, key: int) -> Any: ...
|
||||
|
||||
def latest_successful_window(self) -> SyncWindow | None: ...
|
||||
|
||||
def active_stocks(self) -> tuple[Stock, ...]: ...
|
||||
|
||||
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]: ...
|
||||
|
||||
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]: ...
|
||||
|
||||
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]: ...
|
||||
|
||||
def list_invalid_bar_codes(self, window: SyncWindow) -> tuple[str, ...]: ...
|
||||
|
||||
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]: ...
|
||||
|
||||
|
||||
class PostgresIntegritySnapshotReader:
|
||||
"""Expose the existing repository's read-only integrity operations."""
|
||||
|
||||
def __init__(self, repository: _PostgresRepository | str, *, max_connections: int = 10) -> None:
|
||||
if isinstance(repository, str):
|
||||
# Keep the import lazy: ``postgres.py`` re-exports this adapter for
|
||||
# compatibility, so importing it at module load time would cycle.
|
||||
from .postgres import PostgresMarketDataRepository
|
||||
|
||||
self.repository: _PostgresRepository = PostgresMarketDataRepository(
|
||||
repository,
|
||||
max_connections=max_connections,
|
||||
)
|
||||
else:
|
||||
self.repository = repository
|
||||
|
||||
def open(self) -> None:
|
||||
"""Open the underlying pool when this reader owns a URL-backed one."""
|
||||
|
||||
self.repository.open()
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the underlying pool when this reader owns a URL-backed one."""
|
||||
|
||||
self.repository.close()
|
||||
|
||||
def latest_successful_window(self) -> SyncWindow | None:
|
||||
"""Return the newest completed synchronization window."""
|
||||
|
||||
return self.repository.latest_successful_window()
|
||||
|
||||
def active_stocks(self) -> tuple[Stock, ...]:
|
||||
"""Return the active stock master in deterministic order."""
|
||||
|
||||
return self.repository.active_stocks()
|
||||
|
||||
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]:
|
||||
"""Return a streaming bar iterator."""
|
||||
|
||||
return self.repository.iter_bars(window)
|
||||
|
||||
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]:
|
||||
"""Return a streaming daily-basic iterator."""
|
||||
|
||||
return self.repository.iter_daily_basic(window)
|
||||
|
||||
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||
"""Return bar group keys without reading their rows."""
|
||||
|
||||
return self.repository.list_bar_codes(window)
|
||||
|
||||
def list_invalid_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||
"""Return bar groups whose stored adjustment source is not qfq."""
|
||||
|
||||
return self.repository.list_invalid_bar_codes(window)
|
||||
|
||||
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
|
||||
"""Return daily-basic group keys without reading their rows."""
|
||||
|
||||
return self.repository.list_daily_basic_dates(window)
|
||||
|
||||
def advisory_lock(self, key: int) -> Any:
|
||||
"""Use exactly the same connection-bound advisory lock as sync."""
|
||||
|
||||
return self.repository.advisory_lock(key)
|
||||
|
||||
|
||||
class PostgresIntegrityCheckStore:
|
||||
"""Persist integrity progress and issues through the shared pool.
|
||||
|
||||
The store deliberately has no methods that write market facts. Its only
|
||||
mutation surface is the two ``market_integrity_*`` tables.
|
||||
"""
|
||||
|
||||
def __init__(self, repository: _PostgresRepository | str, *, max_connections: int = 10) -> None:
|
||||
if isinstance(repository, str):
|
||||
from .postgres import PostgresMarketDataRepository
|
||||
|
||||
self.repository: _PostgresRepository = PostgresMarketDataRepository(
|
||||
repository,
|
||||
max_connections=max_connections,
|
||||
)
|
||||
else:
|
||||
self.repository = repository
|
||||
|
||||
def open(self) -> None:
|
||||
"""Open the underlying pool."""
|
||||
|
||||
self.repository.open()
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the underlying pool."""
|
||||
|
||||
self.repository.close()
|
||||
|
||||
def recover_stale_running(self, lock_key: int, stale_after_seconds: int) -> int:
|
||||
"""Mark old workers failed only while the shared lock is available."""
|
||||
|
||||
if stale_after_seconds < 1:
|
||||
raise ValueError("stale_after_seconds must be positive")
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
acquired = False
|
||||
count = 0
|
||||
try:
|
||||
with connection.transaction():
|
||||
acquired_row = connection.execute(
|
||||
"SELECT pg_try_advisory_lock(%s)",
|
||||
(lock_key,),
|
||||
).fetchone()
|
||||
acquired = bool(acquired_row[0]) if acquired_row is not None else False
|
||||
if acquired:
|
||||
result = connection.execute(
|
||||
"""
|
||||
UPDATE market_integrity_check
|
||||
SET status = 'failed',
|
||||
error_type = 'stale_worker',
|
||||
error_message = 'integrity check worker became stale',
|
||||
updated_at = now(),
|
||||
finished_at = now()
|
||||
WHERE status = 'running'
|
||||
AND updated_at < now() - (%s * interval '1 second')
|
||||
""",
|
||||
(stale_after_seconds,),
|
||||
)
|
||||
count = int(result.rowcount)
|
||||
finally:
|
||||
if acquired:
|
||||
connection.execute("SELECT pg_advisory_unlock(%s)", (lock_key,))
|
||||
except MarketDataRepositoryError as exc:
|
||||
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||
return count
|
||||
|
||||
def create_running(self, window: SyncWindow, target_count: int) -> IntegrityCheckRun:
|
||||
"""Atomically claim the single running-check slot."""
|
||||
|
||||
check_id = str(uuid4())
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
row = connection.execute(
|
||||
"""
|
||||
INSERT INTO market_integrity_check
|
||||
(id, status, window_start, window_end, target_count)
|
||||
VALUES (%s, 'running', %s, %s, %s)
|
||||
RETURNING id, status, window_start, window_end, target_count,
|
||||
checked_count, issue_count, error_type, error_message,
|
||||
created_at, finished_at
|
||||
""",
|
||||
(check_id, window.start, window.end, target_count),
|
||||
).fetchone()
|
||||
except MarketDataRepositoryError as exc:
|
||||
if isinstance(exc.__cause__, psycopg.errors.UniqueViolation):
|
||||
raise IntegrityCheckInProgress("another integrity check is running") from exc
|
||||
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||
if row is None:
|
||||
raise IntegrityCheckStoreError("integrity check claim returned no row")
|
||||
return _check_from_row(row)
|
||||
|
||||
def update_progress(self, check_id: str, checked_count: int, issue_count: int) -> None:
|
||||
"""Persist heartbeat and counters after each comparison group."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
UPDATE market_integrity_check
|
||||
SET checked_count = %s, issue_count = %s, updated_at = now()
|
||||
WHERE id = %s AND status = 'running'
|
||||
""",
|
||||
(checked_count, issue_count, check_id),
|
||||
)
|
||||
except MarketDataRepositoryError as exc:
|
||||
raise IntegrityCheckStoreError("integrity check progress cannot be saved") from exc
|
||||
|
||||
def record_issues(self, issues: Iterable[IntegrityIssue]) -> None:
|
||||
"""Insert a bounded issue batch idempotently."""
|
||||
|
||||
records = tuple(issues)
|
||||
if not records:
|
||||
return
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.executemany(
|
||||
"""
|
||||
INSERT INTO market_integrity_issue
|
||||
(check_id, issue_key, item_kind, item_key, issue_type, message)
|
||||
VALUES (%s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (check_id, issue_key) DO UPDATE SET
|
||||
message = EXCLUDED.message
|
||||
""",
|
||||
[
|
||||
(
|
||||
issue.check_id,
|
||||
issue.issue_key,
|
||||
issue.item_kind,
|
||||
issue.item_key,
|
||||
issue.issue_type,
|
||||
_safe_message(issue.message),
|
||||
)
|
||||
for issue in records
|
||||
],
|
||||
)
|
||||
except MarketDataRepositoryError as exc:
|
||||
raise IntegrityCheckStoreError("integrity issues cannot be saved") from exc
|
||||
|
||||
def finish(
|
||||
self,
|
||||
check_id: str,
|
||||
status: IntegrityStatus,
|
||||
*,
|
||||
issue_count: int,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
"""Converge one running check to a terminal state."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
UPDATE market_integrity_check
|
||||
SET status = %s, issue_count = %s, error_type = %s,
|
||||
error_message = %s, updated_at = now(), finished_at = now()
|
||||
WHERE id = %s
|
||||
""",
|
||||
(
|
||||
status,
|
||||
issue_count,
|
||||
error_type,
|
||||
_safe_message(error_message),
|
||||
check_id,
|
||||
),
|
||||
)
|
||||
except MarketDataRepositoryError as exc:
|
||||
raise IntegrityCheckStoreError("integrity check result cannot be saved") from exc
|
||||
|
||||
def get(self, check_id: str, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
|
||||
"""Read one check and a stable issue page."""
|
||||
|
||||
return self._read_page("WHERE id = %s", (check_id,), query)
|
||||
|
||||
def get_latest(self, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
|
||||
"""Read the newest check by creation time."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT id
|
||||
FROM market_integrity_check
|
||||
ORDER BY created_at DESC, id DESC
|
||||
LIMIT 1
|
||||
"""
|
||||
).fetchone()
|
||||
except MarketDataRepositoryError as exc:
|
||||
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||
if row is None:
|
||||
return None
|
||||
return self.get(str(row[0]), query)
|
||||
|
||||
def _read_page(
|
||||
self,
|
||||
predicate: str,
|
||||
parameters: tuple[object, ...],
|
||||
query: IntegrityCheckQuery,
|
||||
) -> IntegrityCheckPage | None:
|
||||
offset = (query.page - 1) * query.page_size
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
row = connection.execute(
|
||||
f"""
|
||||
SELECT id, status, window_start, window_end, target_count,
|
||||
checked_count, issue_count, error_type, error_message,
|
||||
created_at, finished_at
|
||||
FROM market_integrity_check
|
||||
{predicate}
|
||||
""",
|
||||
parameters,
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
total_row = connection.execute(
|
||||
"SELECT count(*) FROM market_integrity_issue WHERE check_id = %s",
|
||||
(row[0],),
|
||||
).fetchone()
|
||||
issue_rows = connection.execute(
|
||||
"""
|
||||
SELECT check_id, issue_key, item_kind, item_key, issue_type,
|
||||
message, created_at
|
||||
FROM market_integrity_issue
|
||||
WHERE check_id = %s
|
||||
ORDER BY issue_key
|
||||
LIMIT %s OFFSET %s
|
||||
""",
|
||||
(row[0], query.page_size, offset),
|
||||
).fetchall()
|
||||
except MarketDataRepositoryError as exc:
|
||||
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||
return IntegrityCheckPage(
|
||||
run=_check_from_row(row),
|
||||
page=query.page,
|
||||
page_size=query.page_size,
|
||||
issues_total=int(total_row[0]) if total_row is not None else 0,
|
||||
issues=tuple(_issue_from_row(issue_row) for issue_row in issue_rows),
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
def _connection(self) -> Generator[Any, None, None]:
|
||||
"""Borrow a pooled connection and translate driver errors safely."""
|
||||
|
||||
try:
|
||||
self.repository.open()
|
||||
with self.repository.connection() as connection:
|
||||
yield connection
|
||||
except psycopg.Error as exc:
|
||||
raise MarketDataRepositoryError("market data database operation failed") from exc
|
||||
|
||||
|
||||
def _check_from_row(row: tuple[Any, ...]) -> IntegrityCheckRun:
|
||||
"""Translate one persistence row into the domain run model."""
|
||||
|
||||
window = None
|
||||
if row[2] is not None and row[3] is not None:
|
||||
window = SyncWindow(start=row[2], end=row[3])
|
||||
status = cast(IntegrityStatus, str(row[1]))
|
||||
return IntegrityCheckRun(
|
||||
id=str(row[0]),
|
||||
status=status,
|
||||
window=window,
|
||||
target_count=int(row[4]),
|
||||
checked_count=int(row[5]),
|
||||
issue_count=int(row[6]),
|
||||
error_type=str(row[7]) if row[7] is not None else None,
|
||||
error_message=str(row[8]) if row[8] is not None else None,
|
||||
created_at=row[9],
|
||||
finished_at=row[10],
|
||||
)
|
||||
|
||||
|
||||
def _issue_from_row(row: tuple[Any, ...]) -> IntegrityIssue:
|
||||
"""Translate one persistence row into a domain issue."""
|
||||
|
||||
return IntegrityIssue(
|
||||
check_id=str(row[0]),
|
||||
issue_key=str(row[1]),
|
||||
item_kind=str(row[2]),
|
||||
item_key=str(row[3]),
|
||||
issue_type=str(row[4]),
|
||||
message=str(row[5]),
|
||||
created_at=row[6],
|
||||
)
|
||||
|
||||
|
||||
def _safe_message(message: str | None) -> str | None:
|
||||
"""Persist bounded, whitespace-normalized messages only."""
|
||||
|
||||
if message is None:
|
||||
return None
|
||||
return " ".join(message.split())[:500]
|
||||
|
||||
|
||||
# This name reads naturally at the application seam and keeps an intuitive
|
||||
# compatibility alias for callers that call every persistence adapter a repo.
|
||||
PostgresIntegrityCheckRepository = PostgresIntegrityCheckStore
|
||||
|
||||
|
||||
__all__ = [
|
||||
"PostgresIntegrityCheckRepository",
|
||||
"PostgresIntegrityCheckStore",
|
||||
"PostgresIntegritySnapshotReader",
|
||||
]
|
||||
@@ -2,14 +2,15 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator, Iterable
|
||||
import threading
|
||||
from collections.abc import Generator, Iterable, Iterator
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
import psycopg
|
||||
from psycopg_pool import ConnectionPool
|
||||
|
||||
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||
from ..domain.overview import (
|
||||
@@ -25,12 +26,69 @@ from ..domain.ports import MarketDataRepositoryError, WriteResult
|
||||
class PostgresMarketDataRepository:
|
||||
"""Persist market data without an ORM identity map.
|
||||
|
||||
Each public write opens a short transaction. The caller publishes the
|
||||
matching CSV only after this method returns successfully.
|
||||
Each public write borrows an independent connection from a bounded pool
|
||||
for a short transaction. The caller publishes the matching CSV only
|
||||
after this method returns successfully.
|
||||
"""
|
||||
|
||||
def __init__(self, database_url: str) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
database_url: str,
|
||||
*,
|
||||
max_connections: int = 10,
|
||||
pool: ConnectionPool[Any] | None = None,
|
||||
) -> None:
|
||||
if max_connections < 1:
|
||||
raise ValueError("max_connections must be at least 1")
|
||||
self.database_url = database_url
|
||||
self.max_connections = max_connections
|
||||
self.pool = (
|
||||
pool
|
||||
if pool is not None
|
||||
else ConnectionPool(
|
||||
conninfo=database_url,
|
||||
min_size=1,
|
||||
max_size=max_connections,
|
||||
open=False,
|
||||
)
|
||||
)
|
||||
self._owns_pool = pool is None
|
||||
self._pool_open = False
|
||||
self._pool_state_lock = threading.Lock()
|
||||
|
||||
def open(self) -> None:
|
||||
"""Open the pool and wait until its minimum connections are ready."""
|
||||
|
||||
with self._pool_state_lock:
|
||||
if self._pool_open:
|
||||
return
|
||||
if bool(getattr(self.pool, "_opened", False)):
|
||||
self._pool_open = True
|
||||
return
|
||||
self.pool.open(wait=True)
|
||||
self._pool_open = True
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close this repository's pool after all borrowed connections return."""
|
||||
|
||||
with self._pool_state_lock:
|
||||
if self._owns_pool and (self._pool_open or bool(getattr(self.pool, "_opened", False))):
|
||||
self.pool.close()
|
||||
self._pool_open = False
|
||||
|
||||
def __enter__(self) -> PostgresMarketDataRepository:
|
||||
self.open()
|
||||
return self
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[Any, None, None]:
|
||||
"""Expose a safe pooled connection seam to sibling adapters."""
|
||||
|
||||
with self._connection() as connection:
|
||||
yield connection
|
||||
|
||||
def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None:
|
||||
self.close()
|
||||
|
||||
def read_overview(self) -> MarketDataOverview:
|
||||
"""Aggregate the active stock pool and latest synchronization facts.
|
||||
@@ -356,6 +414,171 @@ class PostgresMarketDataRepository:
|
||||
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool:
|
||||
return self._exists("market_daily_basic", ts_code, trade_date)
|
||||
|
||||
def count_valid_stocks(self, trade_date: date) -> int:
|
||||
"""Count complete active-stock facts with one set-based query."""
|
||||
|
||||
with self._connection() as connection:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT count(*)
|
||||
FROM market_stock AS stock
|
||||
WHERE stock.is_active = true
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM market_daily_bar AS bar
|
||||
WHERE bar.ts_code = stock.ts_code
|
||||
AND bar.trade_date = %s
|
||||
)
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM market_daily_basic AS basic
|
||||
WHERE basic.ts_code = stock.ts_code
|
||||
AND basic.trade_date = %s
|
||||
)
|
||||
""",
|
||||
(trade_date, trade_date),
|
||||
).fetchone()
|
||||
return int(row[0]) if row is not None else 0
|
||||
|
||||
def latest_successful_window(self) -> SyncWindow | None:
|
||||
"""Return the newest completed synchronization window.
|
||||
|
||||
Integrity checks compare persisted facts to a completed window only;
|
||||
an in-progress or failed batch must never redefine the comparison
|
||||
boundary.
|
||||
"""
|
||||
|
||||
with self._connection() as connection:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT window_start, target_trade_date
|
||||
FROM market_sync_batch
|
||||
WHERE status IN ('success', 'partial_success')
|
||||
ORDER BY target_trade_date DESC, finished_at DESC NULLS LAST,
|
||||
created_at DESC, id DESC
|
||||
LIMIT 1
|
||||
"""
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
return SyncWindow(start=row[0], end=row[1])
|
||||
|
||||
def active_stocks(self) -> tuple[Stock, ...]:
|
||||
"""Read the current active stock master in stable code order."""
|
||||
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
"""
|
||||
SELECT ts_code, name, market, exchange, list_status, list_date
|
||||
FROM market_stock
|
||||
WHERE is_active = true
|
||||
ORDER BY ts_code
|
||||
"""
|
||||
).fetchall()
|
||||
return tuple(
|
||||
Stock(
|
||||
ts_code=str(row[0]),
|
||||
name=str(row[1]),
|
||||
market=str(row[2]),
|
||||
exchange=str(row[3]),
|
||||
list_status=str(row[4]),
|
||||
list_date=row[5],
|
||||
)
|
||||
for row in rows
|
||||
)
|
||||
|
||||
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||
"""List bar groups in a window without materializing their rows."""
|
||||
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
"""
|
||||
SELECT DISTINCT ts_code
|
||||
FROM market_daily_bar
|
||||
WHERE trade_date BETWEEN %s AND %s
|
||||
ORDER BY ts_code
|
||||
""",
|
||||
(window.start, window.end),
|
||||
).fetchall()
|
||||
return tuple(str(row[0]) for row in rows)
|
||||
|
||||
def list_invalid_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||
"""List qfq groups whose persisted adjustment source is not qfq."""
|
||||
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
"""
|
||||
SELECT DISTINCT ts_code
|
||||
FROM market_daily_bar
|
||||
WHERE trade_date BETWEEN %s AND %s
|
||||
AND source_adj <> 'qfq'
|
||||
ORDER BY ts_code
|
||||
""",
|
||||
(window.start, window.end),
|
||||
).fetchall()
|
||||
return tuple(str(row[0]) for row in rows)
|
||||
|
||||
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
|
||||
"""List daily-basic groups in a window without loading their rows."""
|
||||
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
"""
|
||||
SELECT DISTINCT trade_date
|
||||
FROM market_daily_basic
|
||||
WHERE trade_date BETWEEN %s AND %s
|
||||
ORDER BY trade_date
|
||||
""",
|
||||
(window.start, window.end),
|
||||
).fetchall()
|
||||
return tuple(row[0] for row in rows)
|
||||
|
||||
def iter_bars(self, window: SyncWindow) -> Iterator[Bar]:
|
||||
"""Stream qfq bars ordered by stock and trade date.
|
||||
|
||||
A named PostgreSQL cursor keeps the roughly seven-million-row fact
|
||||
table out of process memory. The connection and transaction remain
|
||||
open only for the lifetime of this iterator.
|
||||
"""
|
||||
|
||||
query = """
|
||||
SELECT ts_code, trade_date, open, high, low, close, pre_close,
|
||||
change, pct_chg, vol, amount
|
||||
FROM market_daily_bar
|
||||
WHERE trade_date BETWEEN %s AND %s
|
||||
ORDER BY ts_code, trade_date
|
||||
"""
|
||||
with (
|
||||
self._connection() as connection,
|
||||
connection.transaction(),
|
||||
connection.cursor(name=f"integrity-bars-{uuid4().hex}") as cursor,
|
||||
):
|
||||
cursor.execute(query, (window.start, window.end))
|
||||
while rows := cursor.fetchmany(2_000):
|
||||
for row in rows:
|
||||
yield _bar_from_database_row(row)
|
||||
|
||||
def iter_daily_basic(self, window: SyncWindow) -> Iterator[DailyBasic]:
|
||||
"""Stream daily-basic rows ordered by trade date and stock code."""
|
||||
|
||||
query = """
|
||||
SELECT ts_code, trade_date, close, turnover_rate, turnover_rate_f,
|
||||
volume_ratio, pe, pe_ttm, pb, ps, ps_ttm, dv_ratio, dv_ttm,
|
||||
total_share, float_share, free_share, total_mv, circ_mv
|
||||
FROM market_daily_basic
|
||||
WHERE trade_date BETWEEN %s AND %s
|
||||
ORDER BY trade_date, ts_code
|
||||
"""
|
||||
with (
|
||||
self._connection() as connection,
|
||||
connection.transaction(),
|
||||
connection.cursor(name=f"integrity-basic-{uuid4().hex}") as cursor,
|
||||
):
|
||||
cursor.execute(query, (window.start, window.end))
|
||||
while rows := cursor.fetchmany(2_000):
|
||||
for row in rows:
|
||||
yield _daily_basic_from_database_row(row)
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
target_trade_date: date,
|
||||
@@ -441,6 +664,50 @@ class PostgresMarketDataRepository:
|
||||
),
|
||||
)
|
||||
|
||||
def record_items(self, batch_id: str, outcomes: Iterable[object]) -> None:
|
||||
"""Batch-upsert synchronization audit outcomes in one transaction."""
|
||||
|
||||
records = tuple(outcomes)
|
||||
if not records:
|
||||
return
|
||||
parameters: list[tuple[object, ...]] = []
|
||||
for outcome in records:
|
||||
result = getattr(outcome, "result", WriteResult())
|
||||
failure = getattr(outcome, "failure", None)
|
||||
item_kind_attribute = "item_kind"
|
||||
item_key_attribute = "item_key"
|
||||
status_attribute = "status"
|
||||
parameters.append(
|
||||
(
|
||||
batch_id,
|
||||
str(getattr(outcome, item_kind_attribute)),
|
||||
str(getattr(outcome, item_key_attribute)),
|
||||
str(getattr(outcome, status_attribute)),
|
||||
int(getattr(result, "inserted", 0)),
|
||||
int(getattr(result, "updated", 0)),
|
||||
int(getattr(result, "unchanged", 0)),
|
||||
getattr(outcome, "fingerprint", None),
|
||||
getattr(failure, "error_type", None),
|
||||
self._safe_error(getattr(failure, "message", None)),
|
||||
)
|
||||
)
|
||||
query = """
|
||||
INSERT INTO market_sync_item
|
||||
(batch_id, item_kind, item_key, status, inserted_count,
|
||||
updated_count, unchanged_count, fingerprint, error_type, error_message)
|
||||
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (batch_id, item_kind, item_key) DO UPDATE SET
|
||||
status = EXCLUDED.status,
|
||||
inserted_count = EXCLUDED.inserted_count,
|
||||
updated_count = EXCLUDED.updated_count,
|
||||
unchanged_count = EXCLUDED.unchanged_count,
|
||||
fingerprint = EXCLUDED.fingerprint,
|
||||
error_type = EXCLUDED.error_type,
|
||||
error_message = EXCLUDED.error_message
|
||||
"""
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.cursor().executemany(query, parameters)
|
||||
|
||||
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]:
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
@@ -490,12 +757,13 @@ class PostgresMarketDataRepository:
|
||||
|
||||
@contextmanager
|
||||
def _connection(self) -> Generator[Any, None, None]:
|
||||
"""Translate driver failures into a safe application-level error."""
|
||||
"""Borrow one thread-safe pool connection and hide driver failures."""
|
||||
|
||||
try:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
self.open()
|
||||
with self.pool.connection() as connection:
|
||||
yield connection
|
||||
except psycopg.Error as exc:
|
||||
except Exception as exc: # noqa: BLE001 - repository boundary redacts driver/pool errors
|
||||
raise MarketDataRepositoryError("market data database operation failed") from exc
|
||||
|
||||
@staticmethod
|
||||
@@ -610,11 +878,11 @@ class PostgresMarketDataRepository:
|
||||
market_daily_bar.open, market_daily_bar.high, market_daily_bar.low,
|
||||
market_daily_bar.close, market_daily_bar.pre_close,
|
||||
market_daily_bar.change, market_daily_bar.pct_chg,
|
||||
market_daily_bar.vol, market_daily_bar.amount
|
||||
market_daily_bar.vol, market_daily_bar.amount, market_daily_bar.source_adj
|
||||
) IS DISTINCT FROM (
|
||||
EXCLUDED.open, EXCLUDED.high, EXCLUDED.low, EXCLUDED.close,
|
||||
EXCLUDED.pre_close, EXCLUDED.change, EXCLUDED.pct_chg,
|
||||
EXCLUDED.vol, EXCLUDED.amount
|
||||
EXCLUDED.vol, EXCLUDED.amount, 'qfq'
|
||||
)
|
||||
RETURNING (xmax = 0) AS inserted
|
||||
"""
|
||||
@@ -673,3 +941,77 @@ class PostgresMarketDataRepository:
|
||||
if message is None:
|
||||
return None
|
||||
return " ".join(message.split())[:500]
|
||||
|
||||
|
||||
def _database_decimal(value: object) -> Decimal | None:
|
||||
"""Normalize a PostgreSQL numeric value for the domain model."""
|
||||
|
||||
if value is None:
|
||||
return None
|
||||
result = Decimal(str(value))
|
||||
if not result.is_finite():
|
||||
raise ValueError("database numeric value must be finite")
|
||||
return result
|
||||
|
||||
|
||||
def _bar_from_database_row(row: tuple[Any, ...]) -> Bar:
|
||||
"""Translate one ordered PostgreSQL bar row into a domain value."""
|
||||
|
||||
return Bar(
|
||||
ts_code=str(row[0]),
|
||||
trade_date=row[1],
|
||||
open=_database_decimal(row[2]),
|
||||
high=_database_decimal(row[3]),
|
||||
low=_database_decimal(row[4]),
|
||||
close=_database_decimal(row[5]),
|
||||
pre_close=_database_decimal(row[6]),
|
||||
change=_database_decimal(row[7]),
|
||||
pct_chg=_database_decimal(row[8]),
|
||||
vol=_database_decimal(row[9]),
|
||||
amount=_database_decimal(row[10]),
|
||||
)
|
||||
|
||||
|
||||
def _daily_basic_from_database_row(row: tuple[Any, ...]) -> DailyBasic:
|
||||
"""Translate one ordered PostgreSQL daily-basic row into a domain value."""
|
||||
|
||||
fields = (
|
||||
"close",
|
||||
"turnover_rate",
|
||||
"turnover_rate_f",
|
||||
"volume_ratio",
|
||||
"pe",
|
||||
"pe_ttm",
|
||||
"pb",
|
||||
"ps",
|
||||
"ps_ttm",
|
||||
"dv_ratio",
|
||||
"dv_ttm",
|
||||
"total_share",
|
||||
"float_share",
|
||||
"free_share",
|
||||
"total_mv",
|
||||
"circ_mv",
|
||||
)
|
||||
values = {field: _database_decimal(row[index + 2]) for index, field in enumerate(fields)}
|
||||
return DailyBasic(
|
||||
ts_code=str(row[0]),
|
||||
trade_date=row[1],
|
||||
**values,
|
||||
)
|
||||
|
||||
|
||||
# Re-export the read-only integrity adapters from the established PostgreSQL
|
||||
# module so existing infrastructure import paths remain discoverable.
|
||||
from .integrity import ( # noqa: E402
|
||||
PostgresIntegrityCheckRepository,
|
||||
PostgresIntegrityCheckStore,
|
||||
PostgresIntegritySnapshotReader,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"PostgresIntegrityCheckRepository",
|
||||
"PostgresIntegrityCheckStore",
|
||||
"PostgresIntegritySnapshotReader",
|
||||
"PostgresMarketDataRepository",
|
||||
]
|
||||
|
||||
@@ -16,6 +16,7 @@ from sqlalchemy import (
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
func,
|
||||
text,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
@@ -173,6 +174,41 @@ selection_signal = Table(
|
||||
PrimaryKeyConstraint("run_id", "ts_code", "category"),
|
||||
)
|
||||
|
||||
market_integrity_check = Table(
|
||||
"market_integrity_check",
|
||||
metadata,
|
||||
Column("id", String(36), primary_key=True),
|
||||
Column("status", String(24), nullable=False),
|
||||
Column("window_start", Date),
|
||||
Column("window_end", Date),
|
||||
Column("target_count", Integer, nullable=False, server_default="0"),
|
||||
Column("checked_count", Integer, nullable=False, server_default="0"),
|
||||
Column("issue_count", Integer, nullable=False, server_default="0"),
|
||||
Column("error_type", String(64)),
|
||||
Column("error_message", Text),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("finished_at", DateTime(timezone=True)),
|
||||
)
|
||||
|
||||
market_integrity_issue = Table(
|
||||
"market_integrity_issue",
|
||||
metadata,
|
||||
Column(
|
||||
"check_id",
|
||||
String(36),
|
||||
ForeignKey("market_integrity_check.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
Column("issue_key", String(64), nullable=False),
|
||||
Column("item_kind", String(24), nullable=False),
|
||||
Column("item_key", String(128), nullable=False),
|
||||
Column("issue_type", String(64), nullable=False),
|
||||
Column("message", Text, nullable=False),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
PrimaryKeyConstraint("check_id", "issue_key"),
|
||||
)
|
||||
|
||||
# Keep the declarative metadata aligned with the indexes created by the
|
||||
# Alembic revisions. Alembic uses this object for both offline inspection
|
||||
# and future autogeneration, so omitting these indexes would make the schema
|
||||
@@ -193,3 +229,15 @@ Index(
|
||||
selection_signal.c.target_trade_date,
|
||||
selection_signal.c.ts_code,
|
||||
)
|
||||
Index(
|
||||
"ix_market_integrity_check_status_created_at",
|
||||
market_integrity_check.c.status,
|
||||
market_integrity_check.c.created_at,
|
||||
)
|
||||
Index(
|
||||
"uq_market_integrity_check_running",
|
||||
market_integrity_check.c.status,
|
||||
unique=True,
|
||||
postgresql_where=text("status = 'running'"),
|
||||
)
|
||||
Index("ix_market_integrity_issue_check_id", market_integrity_issue.c.check_id)
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
"""Tushare source adapter and current-universe filtering."""
|
||||
"""Tushare source adapter, request coordination, and current-universe filtering."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from datetime import date
|
||||
from typing import cast
|
||||
|
||||
@@ -14,18 +15,226 @@ from ..domain.rules import filter_current_hs_a_stocks
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_RATE_LIMIT_COOLDOWNS = (60.0, 120.0, 180.0)
|
||||
_RATE_LIMIT_MESSAGES = (
|
||||
"访问频繁",
|
||||
"请稍后",
|
||||
"超过频率",
|
||||
"频率限制",
|
||||
"too many requests",
|
||||
"rate limit",
|
||||
"rate_limit",
|
||||
"http 429",
|
||||
"status code: 429",
|
||||
"429",
|
||||
"http 403",
|
||||
"status code: 403",
|
||||
"403",
|
||||
)
|
||||
|
||||
|
||||
class TushareSourceError(RuntimeError):
|
||||
"""A vendor request failed after the configured retry budget."""
|
||||
|
||||
|
||||
class TushareAdapter:
|
||||
"""Translate Tushare SDK responses into domain records.
|
||||
class RequestCoordinator:
|
||||
"""Coordinate retry and shared rate-limit cooling for one token client.
|
||||
|
||||
The SDK is kept behind this adapter so ordinary domain/application tests
|
||||
can inject a tiny fake client and never need a network token.
|
||||
Normal requests are deliberately not serialized. Only a provider rate
|
||||
limit creates a shared cooldown, so independent worker calls can proceed
|
||||
concurrently during ordinary traffic. ``clock`` and ``wait_fn`` are
|
||||
injectable to make long cooldown behavior deterministic in unit tests.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
random_fn: Callable[[], float] = random.random,
|
||||
clock: Callable[[], float] = time.monotonic,
|
||||
wait_fn: Callable[[float], None] = time.sleep,
|
||||
sleep_fn: Callable[[float], None] | None = None,
|
||||
) -> None:
|
||||
cooldowns = tuple(float(value) for value in cooldown_seconds)
|
||||
if not cooldowns or any(value < 0 for value in cooldowns):
|
||||
raise ValueError("cooldown_seconds must contain non-negative values")
|
||||
self.max_retries = max(0, max_retries)
|
||||
self.backoff_seconds = max(0.0, backoff_seconds)
|
||||
self.cooldown_seconds = cooldowns
|
||||
self.random_fn = random_fn
|
||||
self.clock = clock
|
||||
self.wait_fn = wait_fn
|
||||
self.sleep_fn = sleep_fn or wait_fn
|
||||
self._condition = threading.Condition()
|
||||
self._cooldown_until = 0.0
|
||||
self._rate_limit_count = 0
|
||||
|
||||
@property
|
||||
def cooldown_until(self) -> float:
|
||||
"""Return the current monotonic cooldown deadline."""
|
||||
|
||||
with self._condition:
|
||||
return self._cooldown_until
|
||||
|
||||
def call(self, method_name: str, request: Callable[[], object]) -> object:
|
||||
"""Execute one provider request with bounded, shared retry behavior."""
|
||||
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
self._wait_for_cooldown(method_name)
|
||||
try:
|
||||
result = request()
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
if self.is_rate_limited(exc):
|
||||
cooldown = self._set_rate_limit_cooldown()
|
||||
logger.warning(
|
||||
"tushare_rate_limit method=%s attempt=%d max_attempts=%d "
|
||||
"cooldown_seconds=%.1f",
|
||||
method_name,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
cooldown,
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
continue
|
||||
break
|
||||
if not self._is_retryable(exc):
|
||||
raise
|
||||
if attempt == self.max_retries:
|
||||
break
|
||||
delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn())
|
||||
logger.warning(
|
||||
"tushare_request_retry method=%s attempt=%d max_attempts=%d "
|
||||
"backoff_seconds=%.1f",
|
||||
method_name,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
delay,
|
||||
)
|
||||
self.sleep_fn(delay)
|
||||
else:
|
||||
self._clear_rate_limit_after_success()
|
||||
return result
|
||||
logger.error(
|
||||
"tushare_request_failed method=%s attempts=%d",
|
||||
method_name,
|
||||
self.max_retries + 1,
|
||||
)
|
||||
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
|
||||
|
||||
def request(self, method_name: str, operation: Callable[[], object]) -> object:
|
||||
"""Alias for ``call`` for adapters that model requests as a port."""
|
||||
|
||||
return self.call(method_name, operation)
|
||||
|
||||
def _wait_for_cooldown(self, method_name: str) -> None:
|
||||
while True:
|
||||
with self._condition:
|
||||
delay = self._cooldown_until - self.clock()
|
||||
if delay <= 0:
|
||||
return
|
||||
logger.info(
|
||||
"tushare_rate_limit_wait method=%s wait_seconds=%.1f",
|
||||
method_name,
|
||||
delay,
|
||||
)
|
||||
# A single injected wait hook makes fake-clock tests independent
|
||||
# from wall time. After waiting, re-check because another worker
|
||||
# may have extended the shared deadline.
|
||||
self.wait_fn(delay)
|
||||
|
||||
def _set_rate_limit_cooldown(self) -> float:
|
||||
with self._condition:
|
||||
self._rate_limit_count += 1
|
||||
index = min(self._rate_limit_count - 1, len(self.cooldown_seconds) - 1)
|
||||
duration = self.cooldown_seconds[index]
|
||||
self._cooldown_until = max(self._cooldown_until, self.clock() + duration)
|
||||
self._condition.notify_all()
|
||||
return duration
|
||||
|
||||
def _clear_rate_limit_after_success(self) -> None:
|
||||
with self._condition:
|
||||
# A request that was already in flight when another worker hit a
|
||||
# limit may succeed during the shared cooldown. Do not erase the
|
||||
# escalation history until the cooldown has actually elapsed.
|
||||
if self.clock() >= self._cooldown_until:
|
||||
self._rate_limit_count = 0
|
||||
|
||||
@staticmethod
|
||||
def is_rate_limited(error: BaseException) -> bool:
|
||||
"""Classify stable provider rate-limit signals without logging details."""
|
||||
|
||||
for attribute in ("status_code", "status", "code"):
|
||||
value = getattr(error, attribute, None)
|
||||
if str(value).strip() in {"403", "429"}:
|
||||
return True
|
||||
message = str(error).casefold()
|
||||
return any(marker.casefold() in message for marker in _RATE_LIMIT_MESSAGES)
|
||||
|
||||
@staticmethod
|
||||
def _is_retryable(error: BaseException) -> bool:
|
||||
return isinstance(error, (OSError, RuntimeError, TimeoutError))
|
||||
|
||||
|
||||
# The longer name is useful to callers that want to make the infrastructure
|
||||
# boundary explicit, while the short name remains convenient in unit tests.
|
||||
TushareRequestCoordinator = RequestCoordinator
|
||||
|
||||
|
||||
class CoordinatedTushareClient:
|
||||
"""Proxy that routes qfq's real ``daily`` and ``adj_factor`` calls."""
|
||||
|
||||
def __init__(self, client: object, coordinator: RequestCoordinator) -> None:
|
||||
self._client = client
|
||||
self._coordinator = coordinator
|
||||
|
||||
def __getattr__(self, name: str) -> object:
|
||||
method = getattr(self._client, name)
|
||||
if name not in {"daily", "adj_factor"} or not callable(method):
|
||||
return method
|
||||
|
||||
def coordinated(*args: object, **kwargs: object) -> object:
|
||||
return self._coordinator.call(
|
||||
name,
|
||||
lambda: method(*args, **kwargs),
|
||||
)
|
||||
|
||||
return coordinated
|
||||
|
||||
|
||||
def _bind_coordinated_methods(client: object, coordinator: RequestCoordinator) -> object:
|
||||
"""Bind wrappers on a mutable SDK client while preserving ``api is client``.
|
||||
|
||||
Tushare's ``pro_bar`` accepts the token client through its ``api``
|
||||
parameter. Binding the two methods in place keeps that identity contract
|
||||
(and avoids a proxy-visible behavior change); objects that do not allow
|
||||
attributes fall back to the proxy.
|
||||
"""
|
||||
|
||||
def build_wrapper(name: str, method: Callable[..., object]) -> Callable[..., object]:
|
||||
def coordinated(*args: object, **kwargs: object) -> object:
|
||||
return coordinator.call(name, lambda: method(*args, **kwargs))
|
||||
|
||||
return coordinated
|
||||
|
||||
for name in ("daily", "adj_factor"):
|
||||
method = getattr(client, name, None)
|
||||
if not callable(method):
|
||||
continue
|
||||
|
||||
try:
|
||||
setattr(client, name, build_wrapper(name, method))
|
||||
except (AttributeError, TypeError):
|
||||
return CoordinatedTushareClient(client, coordinator)
|
||||
return client
|
||||
|
||||
|
||||
class TushareAdapter:
|
||||
"""Translate Tushare SDK responses into domain records."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client: object,
|
||||
@@ -36,14 +245,32 @@ class TushareAdapter:
|
||||
request_interval_seconds: float = 0.2,
|
||||
random_fn: Callable[[], float] = random.random,
|
||||
sleep_fn: Callable[[float], None] = time.sleep,
|
||||
request_coordinator: RequestCoordinator | None = None,
|
||||
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
coordinated_client: object | None = None,
|
||||
pro_bar_coordinated: bool = False,
|
||||
) -> None:
|
||||
self.client = client
|
||||
self.pro_bar = pro_bar
|
||||
self.max_retries = max(0, max_retries)
|
||||
self.backoff_seconds = max(0.0, backoff_seconds)
|
||||
self.request_interval_seconds = max(0.0, request_interval_seconds)
|
||||
self.random_fn = random_fn
|
||||
self.sleep_fn = sleep_fn
|
||||
self.request_coordinator = request_coordinator or RequestCoordinator(
|
||||
max_retries=self.max_retries,
|
||||
backoff_seconds=self.backoff_seconds,
|
||||
cooldown_seconds=cooldown_seconds,
|
||||
random_fn=random_fn,
|
||||
wait_fn=sleep_fn,
|
||||
sleep_fn=sleep_fn,
|
||||
)
|
||||
self.coordinated_client = (
|
||||
coordinated_client
|
||||
if coordinated_client is not None
|
||||
else _bind_coordinated_methods(client, self.request_coordinator)
|
||||
)
|
||||
self.pro_bar = pro_bar
|
||||
self.pro_bar_coordinated = pro_bar_coordinated
|
||||
|
||||
@classmethod
|
||||
def from_token(
|
||||
@@ -53,6 +280,7 @@ class TushareAdapter:
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
request_interval_seconds: float = 0.2,
|
||||
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
) -> TushareAdapter:
|
||||
"""Create a production adapter from a token without exposing it."""
|
||||
|
||||
@@ -61,13 +289,22 @@ class TushareAdapter:
|
||||
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
||||
|
||||
client = cast(object, ts.pro_api(token))
|
||||
coordinator = RequestCoordinator(
|
||||
max_retries=max_retries,
|
||||
backoff_seconds=backoff_seconds,
|
||||
cooldown_seconds=cooldown_seconds,
|
||||
)
|
||||
coordinated_client = _bind_coordinated_methods(client, coordinator)
|
||||
pro_bar = cast(
|
||||
Callable[..., object],
|
||||
ts.pro_bar, # pyright: ignore[reportUnknownMemberType]
|
||||
)
|
||||
|
||||
def pro_bar_with_client(**kwargs: object) -> object:
|
||||
return pro_bar(api=client, **kwargs)
|
||||
# qfq's SDK helper otherwise performs hidden retries outside the
|
||||
# shared coordinator and can amplify a provider rate limit.
|
||||
kwargs["retry_count"] = 1
|
||||
return pro_bar(api=coordinated_client, **kwargs)
|
||||
|
||||
return cls(
|
||||
client,
|
||||
@@ -75,6 +312,9 @@ class TushareAdapter:
|
||||
max_retries=max_retries,
|
||||
backoff_seconds=backoff_seconds,
|
||||
request_interval_seconds=request_interval_seconds,
|
||||
request_coordinator=coordinator,
|
||||
coordinated_client=coordinated_client,
|
||||
pro_bar_coordinated=True,
|
||||
)
|
||||
|
||||
def fetch_stocks(self) -> tuple[Stock, ...]:
|
||||
@@ -137,43 +377,29 @@ class TushareAdapter:
|
||||
return tuple(sorted(metrics, key=lambda row: row.ts_code))
|
||||
|
||||
def _records(self, method_name: str, **kwargs: object) -> tuple[Mapping[str, object], ...]:
|
||||
"""Call an SDK method with bounded retry and normalize its tabular output."""
|
||||
"""Call the SDK through the coordinator and normalize tabular output."""
|
||||
|
||||
def request() -> object:
|
||||
if method_name == "pro_bar":
|
||||
if self.pro_bar is not None:
|
||||
return self.pro_bar(**kwargs)
|
||||
method = getattr(self.client, method_name, None)
|
||||
method = getattr(self.coordinated_client, method_name, None)
|
||||
else:
|
||||
method = getattr(self.client, method_name, None)
|
||||
method = getattr(self.coordinated_client, method_name, None)
|
||||
if not callable(method):
|
||||
raise TypeError(f"Tushare client has no callable {method_name}")
|
||||
return method(**kwargs)
|
||||
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
result = request()
|
||||
self.sleep_fn(self.request_interval_seconds)
|
||||
return self._as_records(result)
|
||||
except (OSError, RuntimeError, TimeoutError) as exc:
|
||||
last_error = exc
|
||||
if attempt == self.max_retries:
|
||||
break
|
||||
delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn())
|
||||
logger.warning(
|
||||
"tushare_request_retry method=%s attempt=%d max_attempts=%d",
|
||||
method_name,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
)
|
||||
self.sleep_fn(delay)
|
||||
logger.error(
|
||||
"tushare_request_failed method=%s attempts=%d",
|
||||
method_name,
|
||||
self.max_retries + 1,
|
||||
# ``pro_bar`` itself is a qfq composition helper; its nested daily and
|
||||
# adj_factor methods are bound to the coordinator. Wrapping the helper
|
||||
# as a second retry layer would hide the useful error classification.
|
||||
result = (
|
||||
request()
|
||||
if method_name == "pro_bar" and self.pro_bar_coordinated
|
||||
else self.request_coordinator.call(method_name, request)
|
||||
)
|
||||
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
|
||||
self.sleep_fn(self.request_interval_seconds)
|
||||
return self._as_records(result)
|
||||
|
||||
@staticmethod
|
||||
def _as_records(result: object) -> tuple[Mapping[str, object], ...]:
|
||||
|
||||
@@ -70,14 +70,24 @@ def main(argv: Sequence[str] | None = None) -> int:
|
||||
backoff_seconds=settings.market_data_retry_backoff_seconds,
|
||||
request_interval_seconds=settings.market_data_request_interval_seconds,
|
||||
)
|
||||
use_case = SyncMarketData(
|
||||
source,
|
||||
CsvSnapshotStore(settings.market_data_csv_root),
|
||||
PostgresMarketDataRepository(settings.database_url),
|
||||
coverage_threshold=settings.market_data_coverage_threshold,
|
||||
lock_key=settings.market_data_advisory_lock_key,
|
||||
repository = PostgresMarketDataRepository(
|
||||
settings.database_url,
|
||||
# Keep one control/advisory connection and one main-thread connection
|
||||
# available in addition to the worker connections.
|
||||
max_connections=settings.market_data_max_workers + 2,
|
||||
)
|
||||
summary = use_case.execute(command)
|
||||
try:
|
||||
use_case = SyncMarketData(
|
||||
source,
|
||||
CsvSnapshotStore(settings.market_data_csv_root),
|
||||
repository,
|
||||
coverage_threshold=settings.market_data_coverage_threshold,
|
||||
lock_key=settings.market_data_advisory_lock_key,
|
||||
max_workers=settings.market_data_max_workers,
|
||||
)
|
||||
summary = use_case.execute(command)
|
||||
finally:
|
||||
repository.close()
|
||||
print(json.dumps(summary.as_dict(), ensure_ascii=False, sort_keys=True))
|
||||
return summary.exit_code
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""HTTP presentation for the Home market-data overview."""
|
||||
|
||||
import atexit
|
||||
import threading
|
||||
from datetime import date, datetime
|
||||
from typing import Annotated, Literal
|
||||
from zoneinfo import ZoneInfo
|
||||
@@ -24,6 +26,8 @@ from zhixing_server.modules.market_data.infrastructure.postgres import (
|
||||
|
||||
home_router = APIRouter()
|
||||
_SHANGHAI = ZoneInfo("Asia/Shanghai")
|
||||
_REPOSITORY_CACHE_LOCK = threading.Lock()
|
||||
_REPOSITORY_CACHE: dict[tuple[str, int], PostgresMarketDataRepository] = {}
|
||||
|
||||
OverviewStatusValue = Literal["updating", "success", "partial_success", "failed", "no_data"]
|
||||
UpdateStatusValue = Literal["updating", "success", "partial_success", "failed"]
|
||||
@@ -138,9 +142,37 @@ def _to_shanghai(value: datetime | None) -> datetime | None:
|
||||
def get_market_data_overview_reader(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> MarketDataOverviewReader:
|
||||
"""Build the PostgreSQL read adapter for one request."""
|
||||
"""Return the process-cached PostgreSQL read adapter for this database."""
|
||||
|
||||
return PostgresMarketDataRepository(settings.database_url)
|
||||
return get_market_data_repository(settings)
|
||||
|
||||
|
||||
def get_market_data_repository(settings: Settings) -> PostgresMarketDataRepository:
|
||||
"""Return the shared process-cached PostgreSQL pool for market data."""
|
||||
|
||||
key = (settings.database_url, settings.market_data_max_workers + 2)
|
||||
with _REPOSITORY_CACHE_LOCK:
|
||||
repository = _REPOSITORY_CACHE.get(key)
|
||||
if repository is None:
|
||||
repository = PostgresMarketDataRepository(
|
||||
settings.database_url,
|
||||
max_connections=key[1],
|
||||
)
|
||||
_REPOSITORY_CACHE[key] = repository
|
||||
return repository
|
||||
|
||||
|
||||
def _close_cached_repositories() -> None:
|
||||
"""Release read pools when the HTTP process exits."""
|
||||
|
||||
with _REPOSITORY_CACHE_LOCK:
|
||||
repositories = tuple(_REPOSITORY_CACHE.values())
|
||||
_REPOSITORY_CACHE.clear()
|
||||
for repository in repositories:
|
||||
repository.close()
|
||||
|
||||
|
||||
atexit.register(_close_cached_repositories)
|
||||
|
||||
|
||||
@home_router.get("/overview", response_model=HomeOverviewResponse)
|
||||
|
||||
Reference in New Issue
Block a user