feat(market-data): 优化同步并增加完整性检查

This commit is contained in:
yuxuanhui
2026-08-11 11:20:13 +08:00
parent 3ce186f977
commit 7ce11543af
42 changed files with 4837 additions and 197 deletions
@@ -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: ...
@@ -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)