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
+6
View File
@@ -17,4 +17,10 @@ ZHIXING_POSTGRES_PASSWORD=zhixing
# ZHIXING_DATABASE_URL=postgresql://zhixing-system:<url-encoded-password>@postgresql:5432/zhixing-system?sslmode=disable
ZHIXING_TUSHARE_TOKEN=
ZHIXING_MARKET_DATA_CSV_ROOT=/app/data/market-data
ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD=0.99
ZHIXING_MARKET_DATA_MAX_WORKERS=8
ZHIXING_MARKET_DATA_MAX_RETRIES=3
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS=1.0
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS=0.2
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY=7380521
API_UPSTREAM=http://server:8000
+12
View File
@@ -30,6 +30,12 @@ services:
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:-postgresql://zhixing:zhixing@postgres:5432/zhixing}
ZHIXING_LOG_LEVEL: ${ZHIXING_LOG_LEVEL:-INFO}
ZHIXING_MARKET_DATA_CSV_ROOT: /app/data/market-data
ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD: ${ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD:-0.99}
ZHIXING_MARKET_DATA_MAX_WORKERS: ${ZHIXING_MARKET_DATA_MAX_WORKERS:-8}
ZHIXING_MARKET_DATA_MAX_RETRIES: ${ZHIXING_MARKET_DATA_MAX_RETRIES:-3}
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS: ${ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS:-1.0}
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS: ${ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS:-0.2}
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY: ${ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY:-7380521}
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
init: true
ports:
@@ -92,6 +98,12 @@ services:
environment:
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:-postgresql://zhixing:zhixing@postgres:5432/zhixing}
ZHIXING_MARKET_DATA_CSV_ROOT: /app/data/market-data
ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD: ${ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD:-0.99}
ZHIXING_MARKET_DATA_MAX_WORKERS: ${ZHIXING_MARKET_DATA_MAX_WORKERS:-8}
ZHIXING_MARKET_DATA_MAX_RETRIES: ${ZHIXING_MARKET_DATA_MAX_RETRIES:-3}
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS: ${ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS:-1.0}
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS: ${ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS:-0.2}
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY: ${ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY:-7380521}
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
volumes:
- ./zhixing-server:/app
+12
View File
@@ -12,6 +12,12 @@ services:
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:?Set ZHIXING_DATABASE_URL to the 1Panel PostgreSQL URL}
ZHIXING_LOG_LEVEL: ${ZHIXING_LOG_LEVEL:-INFO}
ZHIXING_MARKET_DATA_CSV_ROOT: /app/data/market-data
ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD: ${ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD:-0.99}
ZHIXING_MARKET_DATA_MAX_WORKERS: ${ZHIXING_MARKET_DATA_MAX_WORKERS:-8}
ZHIXING_MARKET_DATA_MAX_RETRIES: ${ZHIXING_MARKET_DATA_MAX_RETRIES:-3}
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS: ${ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS:-1.0}
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS: ${ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS:-0.2}
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY: ${ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY:-7380521}
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
init: true
expose:
@@ -91,6 +97,12 @@ services:
TZ: Asia/Shanghai
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:?Set ZHIXING_DATABASE_URL to the 1Panel PostgreSQL URL}
ZHIXING_MARKET_DATA_CSV_ROOT: /app/data/market-data
ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD: ${ZHIXING_MARKET_DATA_COVERAGE_THRESHOLD:-0.99}
ZHIXING_MARKET_DATA_MAX_WORKERS: ${ZHIXING_MARKET_DATA_MAX_WORKERS:-8}
ZHIXING_MARKET_DATA_MAX_RETRIES: ${ZHIXING_MARKET_DATA_MAX_RETRIES:-3}
ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS: ${ZHIXING_MARKET_DATA_RETRY_BACKOFF_SECONDS:-1.0}
ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS: ${ZHIXING_MARKET_DATA_REQUEST_INTERVAL_SECONDS:-0.2}
ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY: ${ZHIXING_MARKET_DATA_ADVISORY_LOCK_KEY:-7380521}
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
volumes:
- market-data:/app/data/market-data
+6
View File
@@ -62,6 +62,12 @@ docker compose -f docker-compose.prod.yml --profile jobs run --rm market-sync --
`success` 返回 0;覆盖率低于配置阈值或出现部分失败返回 2;没有可用成功结果或基础设施失败返回 1。CLI 输出不包含 Tushare token 或数据库密码。
行情阶段默认使用 8 路固定 worker;可通过 `ZHIXING_MARKET_DATA_MAX_WORKERS` 调低或调高,取值必须
至少为 1。每个 worker 从有上限的 PostgreSQL 连接池借用独立连接,主线程批量写入同步审计并通过
一次集合查询计算覆盖率。所有 worker 共用 Tushare 频控协调器;普通请求保持并发,命中 403、429 或
“访问频繁”等提示时共享 60/120/180 秒冷却窗口。若供应商频控持续发生,先把 worker 降到 1,再通过
`--retry-batch-id` 只恢复失败对象。
## Cron 与重试
宿主机 cron 只负责启动临时容器,不写入容器内部的 crontab。下面的示例每天工作日 18:00 触发;交易日历、唯一约束和 PostgreSQL advisory lock 使周末、节假日、重复触发和重叠触发保持安全:
@@ -0,0 +1,90 @@
"""Create independent market-data integrity check and issue tables."""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "0003_market_integrity_checks"
down_revision: str | None = "0002_selection_results"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Create only check metadata; market facts remain untouched."""
op.create_table(
"market_integrity_check",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("status", sa.String(24), nullable=False),
sa.Column("window_start", sa.Date()),
sa.Column("window_end", sa.Date()),
sa.Column("target_count", sa.Integer(), nullable=False, server_default="0"),
sa.Column("checked_count", sa.Integer(), nullable=False, server_default="0"),
sa.Column("issue_count", sa.Integer(), nullable=False, server_default="0"),
sa.Column("error_type", sa.String(64)),
sa.Column("error_message", sa.Text()),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.Column("finished_at", sa.DateTime(timezone=True)),
)
op.create_index(
"ix_market_integrity_check_status_created_at",
"market_integrity_check",
["status", "created_at"],
)
# A partial unique index makes the running claim atomic under concurrent
# POST requests; stale recovery first moves an abandoned row to failed.
op.create_index(
"uq_market_integrity_check_running",
"market_integrity_check",
["status"],
unique=True,
postgresql_where=sa.text("status = 'running'"),
)
op.create_table(
"market_integrity_issue",
sa.Column(
"check_id",
sa.String(36),
sa.ForeignKey("market_integrity_check.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("issue_key", sa.String(64), nullable=False),
sa.Column("item_kind", sa.String(24), nullable=False),
sa.Column("item_key", sa.String(128), nullable=False),
sa.Column("issue_type", sa.String(64), nullable=False),
sa.Column("message", sa.Text(), nullable=False),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.PrimaryKeyConstraint("check_id", "issue_key"),
)
op.create_index("ix_market_integrity_issue_check_id", "market_integrity_issue", ["check_id"])
def downgrade() -> None:
"""Drop only the integrity report tables in dependency-safe order."""
op.drop_index("ix_market_integrity_issue_check_id", table_name="market_integrity_issue")
op.drop_table("market_integrity_issue")
op.drop_index("uq_market_integrity_check_running", table_name="market_integrity_check")
op.drop_index(
"ix_market_integrity_check_status_created_at",
table_name="market_integrity_check",
)
op.drop_table("market_integrity_check")
+1 -1
View File
@@ -9,7 +9,7 @@ dependencies = [
"fastapi>=0.141.1",
"numpy>=2.4.0",
"pandas>=2.3.3",
"psycopg[binary]>=3.3.2",
"psycopg[binary,pool]>=3.3.2",
"pydantic-settings>=2.14.2",
"sqlalchemy>=2.0.46",
"tushare>=1.4.24",
@@ -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,20 +128,25 @@ 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()
try:
with self.repository.advisory_lock(self.lock_key) as acquired:
if not acquired:
return SyncBatchSummary(
@@ -137,6 +163,20 @@ class SyncMarketData:
),
)
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",
)
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,
)
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,6 +234,24 @@ class SyncMarketData:
SyncFailure("stock", "universe", "empty_universe", "no eligible 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,
@@ -175,6 +259,19 @@ class SyncMarketData:
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,52 +349,118 @@ class SyncMarketData:
batch_id,
len(stocks_to_process),
)
for current, stock in enumerate(stocks_to_process, start=1):
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)
self._process_bar(batch_id, stock, window, failures, totals)
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=stock.ts_code,
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",
batch_id,
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,
[
SyncItemOutcome(
"batch",
"retention",
"failed",
WriteResult(),
error_type=failure.error_type,
error_message=failure.message,
failure=failure,
)
],
)
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
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
try:
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}")
_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()
# ``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)
)
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,
)
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
@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,
)
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,
)
try:
use_case = SyncMarketData(
source,
CsvSnapshotStore(settings.market_data_csv_root),
PostgresMarketDataRepository(settings.database_url),
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)
@@ -33,6 +33,8 @@ def test_postgres_migration_creates_market_data_contract(
"market_daily_basic",
"market_sync_batch",
"market_sync_item",
"market_integrity_check",
"market_integrity_issue",
"selection_run",
"selection_run_item",
"selection_signal",
@@ -0,0 +1,141 @@
from __future__ import annotations
import os
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from pathlib import Path
from typing import Any
import psycopg
import pytest
from alembic import command
from alembic.config import Config
from zhixing_server.modules.market_data.application.sync import SyncFailure, SyncItemOutcome
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, SyncWindow
from zhixing_server.modules.market_data.domain.ports import MarketDataRepositoryError, WriteResult
from zhixing_server.modules.market_data.infrastructure.postgres import (
PostgresMarketDataRepository,
)
TEST_CODES = ("991901.SZ", "991902.SZ")
TARGET = date(2024, 1, 2)
WINDOW = SyncWindow(start=TARGET, end=TARGET)
@contextmanager
def _database(database_url: str) -> Any:
with psycopg.connect(database_url) as connection, connection.transaction():
yield connection
def _cleanup(database_url: str) -> None:
with _database(database_url) as connection:
connection.execute(
"DELETE FROM market_sync_item WHERE batch_id IN "
"(SELECT id FROM market_sync_batch WHERE id LIKE '9919%')"
)
connection.execute("DELETE FROM market_sync_batch WHERE id LIKE '9919%'")
connection.execute(
"DELETE FROM market_daily_basic WHERE ts_code = ANY(%s)",
(list(TEST_CODES),),
)
connection.execute(
"DELETE FROM market_daily_bar WHERE ts_code = ANY(%s)",
(list(TEST_CODES),),
)
connection.execute(
"DELETE FROM market_stock WHERE ts_code = ANY(%s)",
(list(TEST_CODES),),
)
def _prepare_database(database_url: str) -> None:
server_root = Path(__file__).parents[2]
config = Config(str(server_root / "alembic.ini"))
config.set_main_option(
"sqlalchemy.url",
database_url.replace("%", "%%").replace("postgresql://", "postgresql+psycopg://"),
)
command.upgrade(config, "head")
_cleanup(database_url)
with _database(database_url) as connection:
connection.executemany(
"""
INSERT INTO market_stock
(ts_code, name, market, exchange, list_status, is_active)
VALUES (%s, %s, '主板', 'SZSE', 'L', true)
""",
[(code, f"测试{code}") for code in TEST_CODES],
)
@pytest.mark.integration
def test_pool_upsert_rollback_batch_audit_and_set_coverage() -> None:
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
if not database_url:
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
_prepare_database(database_url)
repository = PostgresMarketDataRepository(database_url, max_connections=4)
batch_id: str | None = None
try:
bars = {code: Bar(code, TARGET, close=Decimal("10")) for code in TEST_CODES}
with repository:
def upsert(code: str) -> WriteResult:
return repository.upsert_bars((bars[code],), WINDOW, full_snapshot=True)
with ThreadPoolExecutor(max_workers=2) as executor:
results = tuple(executor.map(upsert, TEST_CODES))
assert [result.inserted for result in results] == [1, 1]
with pytest.raises(MarketDataRepositoryError):
repository.upsert_bars(
(
Bar(TEST_CODES[0], TARGET, close=Decimal("11")),
Bar(TEST_CODES[0], TARGET, close=Decimal("12")),
),
WINDOW,
full_snapshot=True,
)
unchanged = repository.upsert_daily_basic(
(DailyBasic(TEST_CODES[0], TARGET, close=Decimal("10")),),
WINDOW,
)
assert unchanged.inserted == 1
assert repository.count_valid_stocks(TARGET) == 1
batch_id = "9919-pool-test"
outcomes = (
SyncItemOutcome("bar", TEST_CODES[0], "success", WriteResult(inserted=1)),
SyncItemOutcome(
"bar",
TEST_CODES[1],
"failed",
failure=SyncFailure("bar", TEST_CODES[1], "source_error", "safe failure"),
),
)
# Use the real schema row directly because the test repository's
# create_batch API generates a UUID for normal production calls.
with _database(database_url) as connection:
connection.execute(
"""
INSERT INTO market_sync_batch
(id, target_trade_date, window_start, mode, status, target_count)
VALUES (%s, %s, %s, 'daily', 'running', 2)
""",
(batch_id, TARGET, TARGET),
)
repository.record_items(batch_id, outcomes)
with _database(database_url) as connection:
count = connection.execute(
"SELECT count(*) FROM market_sync_item WHERE batch_id = %s",
(batch_id,),
).fetchone()[0]
assert count == 2
finally:
repository.close()
_cleanup(database_url)
@@ -5,7 +5,11 @@ from pathlib import Path
import pytest
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import CsvSnapshotStore
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import (
BAR_COLUMNS,
CsvSnapshotStore,
SnapshotReadError,
)
def make_bar(trade_date: date, close: str = "10") -> Bar:
@@ -44,3 +48,24 @@ def test_daily_basic_snapshot_rejects_duplicate_codes(tmp_path: Path) -> None:
with pytest.raises(ValueError, match="duplicate"):
store.stage_daily_basic(target, (make_basic(target), make_basic(target)))
def test_readers_reject_empty_formal_snapshots(tmp_path: Path) -> None:
store = CsvSnapshotStore(tmp_path)
path = store.bars_path("000001.SZ")
path.parent.mkdir(parents=True)
path.write_text(",".join(BAR_COLUMNS) + "\n", encoding="utf-8")
with pytest.raises(SnapshotReadError, match="empty"):
store.read_bars("000001.SZ")
def test_invalid_daily_basic_path_is_reported_without_opening_content(tmp_path: Path) -> None:
store = CsvSnapshotStore(tmp_path)
invalid_path = tmp_path / "daily-basic" / "2024" / "20241301.csv"
invalid_path.parent.mkdir(parents=True)
invalid_path.write_text("not parsed", encoding="utf-8")
assert store.list_invalid_daily_basic_snapshot_files() == (
("daily-basic/2024/20241301.csv", "window_out_of_bounds"),
)
@@ -13,6 +13,7 @@ from zhixing_server.modules.market_data.domain.models import (
DailyBasic,
Stock,
SyncWindow,
decimal_text,
)
from zhixing_server.modules.market_data.domain.rules import filter_current_hs_a_stocks
@@ -65,6 +66,12 @@ def test_daily_basic_maps_nan_to_none_but_rejects_infinite_values() -> None:
)
def test_decimal_text_rejects_non_finite_domain_values() -> None:
for value in (Decimal("NaN"), Decimal("Infinity"), Decimal("-Infinity")):
with pytest.raises(ValueError, match="must be finite"):
decimal_text(value)
def test_universe_keeps_current_non_st_hs_a_stocks() -> None:
stocks = (
Stock("000001.SZ", "平安银行", exchange="SZSE", list_status="L"),
@@ -0,0 +1,345 @@
from collections.abc import Generator, Iterable
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from fastapi.testclient import TestClient
from zhixing_server.bootstrap.app import create_app
from zhixing_server.modules.market_data.application.integrity import RunMarketIntegrityCheck
from zhixing_server.modules.market_data.domain.integrity import (
IntegrityCheckInProgress,
IntegrityCheckPage,
IntegrityCheckQuery,
IntegrityCheckRun,
IntegrityCheckStoreError,
IntegrityIssue,
)
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, Stock, SyncWindow
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import SnapshotReadError
from zhixing_server.modules.market_data.presentation.integrity import get_market_integrity_service
WINDOW = SyncWindow(date(2024, 1, 2), date(2024, 1, 3))
STOCK = Stock("000001.SZ", "平安银行", exchange="SZSE")
BARS = (
Bar(STOCK.ts_code, date(2024, 1, 2), close=Decimal("10")),
Bar(STOCK.ts_code, date(2024, 1, 3), close=Decimal("11")),
)
BASICS = (
DailyBasic(STOCK.ts_code, date(2024, 1, 2), close=Decimal("10")),
DailyBasic(STOCK.ts_code, date(2024, 1, 3), close=Decimal("11")),
)
class FakeReader:
def __init__(self, *, acquired: bool = True) -> None:
self.acquired = acquired
self.stocks = (STOCK,)
self.bars: tuple[Bar, ...] = BARS
self.basics: tuple[DailyBasic, ...] = BASICS
self.bar_rows_read = 0
self.basic_rows_read = 0
def latest_successful_window(self) -> SyncWindow:
return WINDOW
def active_stocks(self) -> tuple[Stock, ...]:
return self.stocks
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
return (STOCK.ts_code,)
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
return (date(2024, 1, 2), date(2024, 1, 3))
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]:
for row in self.bars:
self.bar_rows_read += 1
yield row
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]:
for row in self.basics:
self.basic_rows_read += 1
yield row
@contextmanager
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
yield self.acquired
class BrokenKeyReader(FakeReader):
"""Reader whose storage key query fails before a check can be claimed."""
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
raise RuntimeError("database unavailable")
class FakeSnapshots:
def __init__(self) -> None:
self.stocks: tuple[Stock, ...] | None = (STOCK,)
self.bars: tuple[Bar, ...] | None = BARS
self.basics: dict[date, tuple[DailyBasic, ...]] = {
date(2024, 1, 2): (BASICS[0],),
date(2024, 1, 3): (BASICS[1],),
}
self.parse_bar = False
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
return (STOCK.ts_code,) if self.bars is not None else ()
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
return tuple(self.basics)
def read_stocks(self) -> tuple[Stock, ...] | None:
return self.stocks
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None:
if self.parse_bar:
raise SnapshotReadError("bad bar csv")
return self.bars
def read_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...] | None:
return self.basics.get(trade_date)
class FakeStore:
def __init__(self) -> None:
self.run = IntegrityCheckRun("check-1", "running", WINDOW, 4)
self.issues: list[IntegrityIssue] = []
self.progress: tuple[int, int] = (0, 0)
def recover_stale_running(self, lock_key: int, stale_after_seconds: int) -> int:
return 0
def create_running(self, window: SyncWindow, target_count: int) -> IntegrityCheckRun:
self.run = IntegrityCheckRun("check-1", "running", window, target_count)
return self.run
def update_progress(self, check_id: str, checked_count: int, issue_count: int) -> None:
self.progress = (checked_count, issue_count)
self.run = IntegrityCheckRun(
self.run.id,
self.run.status,
self.run.window,
self.run.target_count,
checked_count,
issue_count,
)
def record_issues(self, issues: Iterable[IntegrityIssue]) -> None:
self.issues.extend(issues)
def finish(
self,
check_id: str,
status: str,
*,
issue_count: int,
error_type: str | None = None,
error_message: str | None = None,
) -> None:
self.run = IntegrityCheckRun(
self.run.id,
status, # type: ignore[arg-type]
self.run.window,
self.run.target_count,
self.progress[0],
issue_count,
error_type,
error_message,
)
def get(self, check_id: str, query: IntegrityCheckQuery) -> IntegrityCheckPage:
return IntegrityCheckPage(
self.run,
query.page,
query.page_size,
len(self.issues),
tuple(self.issues[(query.page - 1) * query.page_size : query.page * query.page_size]),
)
def get_latest(self, query: IntegrityCheckQuery) -> IntegrityCheckPage:
return self.get(self.run.id, query)
def test_integrity_check_passes_without_writing_facts() -> None:
reader = FakeReader()
snapshots = FakeSnapshots()
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
run = service.prepare()
service.execute(run.id)
assert store.run.status == "passed"
assert store.issues == []
assert reader.bar_rows_read == len(BARS)
assert reader.basic_rows_read == len(BASICS)
def test_prepare_does_not_hide_storage_key_query_failure() -> None:
service = RunMarketIntegrityCheck(BrokenKeyReader(), FakeSnapshots(), FakeStore())
try:
service.prepare()
except RuntimeError as exc:
assert str(exc) == "database unavailable"
else:
raise AssertionError("storage failure must not claim a running check")
def test_integrity_check_reports_stable_content_and_missing_types() -> None:
reader = FakeReader()
snapshots = FakeSnapshots()
snapshots.bars = (Bar(STOCK.ts_code, date(2024, 1, 2), close=Decimal("99")),)
snapshots.basics.pop(date(2024, 1, 3))
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
service.execute(service.prepare().id)
assert store.run.status == "issues_found"
assert {issue.issue_type for issue in store.issues} >= {
"content_mismatch",
"missing_csv",
}
keys = [issue.issue_key for issue in store.issues]
assert keys == list(dict.fromkeys(keys))
def test_csv_parse_error_does_not_stop_other_groups() -> None:
reader = FakeReader()
snapshots = FakeSnapshots()
snapshots.parse_bar = True
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
service.execute(service.prepare().id)
assert store.run.status == "issues_found"
assert any(issue.issue_type == "parse_error" for issue in store.issues)
assert not any(
issue.item_kind == "bar" and issue.issue_type == "missing_csv" for issue in store.issues
)
assert reader.basic_rows_read == len(BASICS)
def test_window_only_bar_group_does_not_become_extra_csv() -> None:
reader = FakeReader()
reader.bars = ()
snapshots = FakeSnapshots()
snapshots.bars = (Bar(STOCK.ts_code, date(2025, 1, 2), close=Decimal("10")),)
store = FakeStore()
service = RunMarketIntegrityCheck(reader, snapshots, store)
service.execute(service.prepare().id)
bar_issues = [issue for issue in store.issues if issue.item_kind == "bar"]
assert {issue.issue_type for issue in bar_issues} == {"window_out_of_bounds"}
def test_integrity_lock_conflict_converges_to_failed() -> None:
reader = FakeReader(acquired=False)
store = FakeStore()
service = RunMarketIntegrityCheck(reader, FakeSnapshots(), store)
service.execute(service.prepare().id)
assert store.run.status == "failed"
assert store.run.error_type == "lock_unavailable"
class FakeHttpService:
def __init__(self, result: IntegrityCheckPage | None = None) -> None:
self.result = result
self.executed = False
self.mode = "ok"
def prepare(self) -> IntegrityCheckRun:
if self.mode == "in_progress":
raise IntegrityCheckInProgress("already running")
if self.mode == "storage_error":
raise IntegrityCheckStoreError("storage unavailable")
return IntegrityCheckRun("http-check", "running", WINDOW, 1)
def execute(self, check_id: str) -> None:
self.executed = True
def get(self, check_id: str, *, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
if self.mode == "storage_error":
raise IntegrityCheckStoreError("storage unavailable")
return self.result
def get_latest(self, *, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
if self.mode == "storage_error":
raise IntegrityCheckStoreError("storage unavailable")
return self.result
def test_integrity_http_returns_202_and_paginates_report() -> None:
issue = IntegrityIssue.build("http-check", "bar", "000001.SZ:2024-01-02", "parse_error", "bad")
result = IntegrityCheckPage(
IntegrityCheckRun("http-check", "issues_found", WINDOW, 1, 1, 1),
page=2,
page_size=1,
issues_total=1,
issues=(issue,),
)
service = FakeHttpService(result)
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: service
client = TestClient(app)
accepted = client.post("/api/v1/market-data/integrity-checks")
report = client.get(
"/api/v1/market-data/integrity-checks/http-check",
params={"page": 2, "page_size": 1},
)
assert accepted.status_code == 202
assert accepted.json()["check_id"] == "http-check"
assert service.executed is True
assert report.status_code == 200
assert report.json()["issues"][0]["issue_type"] == "parse_error"
assert report.json()["page"] == 2
def test_integrity_http_latest_without_report_is_no_data() -> None:
service = FakeHttpService(None)
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: service
response = TestClient(app).get("/api/v1/market-data/integrity-checks/latest")
assert response.status_code == 200
assert response.json()["status"] == "no_data"
assert response.json()["issues"] == []
def test_integrity_http_maps_running_conflict_to_409() -> None:
service = FakeHttpService()
service.mode = "in_progress"
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: service
response = TestClient(app).post("/api/v1/market-data/integrity-checks")
assert response.status_code == 409
assert response.json()["detail"]["code"] == "integrity_check_in_progress"
def test_integrity_http_maps_missing_report_and_storage_to_404_and_503() -> None:
missing_service = FakeHttpService(None)
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: missing_service
missing = TestClient(app).get("/api/v1/market-data/integrity-checks/unknown")
assert missing.status_code == 404
assert missing.json()["detail"]["code"] == "integrity_check_not_found"
storage_service = FakeHttpService(None)
storage_service.mode = "storage_error"
app = create_app()
app.dependency_overrides[get_market_integrity_service] = lambda: storage_service
unavailable = TestClient(app).get("/api/v1/market-data/integrity-checks/latest")
assert unavailable.status_code == 503
assert unavailable.json()["detail"]["code"] == "integrity_storage_unavailable"
@@ -0,0 +1,241 @@
from __future__ import annotations
import threading
import time
from collections.abc import Generator, Iterable, Sequence
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from pathlib import Path
from zhixing_server.modules.market_data.application.sync import (
SyncBatchSummary,
SyncItemOutcome,
SyncMarketData,
SyncMarketDataCommand,
)
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, Stock, SyncWindow
from zhixing_server.modules.market_data.domain.ports import WriteResult
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import CsvSnapshotStore
class ConcurrentSource:
def __init__(
self, target: date, delays: dict[str, float], failed_code: str | None = None
) -> None:
self.target = target
self.stocks = tuple(
Stock(f"{index:06d}.SZ", f"测试{index}", exchange="SZSE", list_status="L")
for index in range(1, 6)
)
self.delays = delays
self.failed_code = failed_code
self._lock = threading.Lock()
self.active = 0
self.max_active = 0
def fetch_stocks(self) -> Sequence[Stock]:
return self.stocks
def fetch_open_dates(self, start: date, end: date) -> Sequence[date]:
return (self.target,) if start <= self.target <= end else ()
def fetch_daily_basic(self, trade_date: date) -> Sequence[DailyBasic]:
return tuple(
DailyBasic(stock.ts_code, trade_date, close=Decimal("10")) for stock in self.stocks
)
def fetch_bars(self, ts_code: str, window: SyncWindow) -> Sequence[Bar]:
with self._lock:
self.active += 1
self.max_active = max(self.max_active, self.active)
try:
time.sleep(self.delays.get(ts_code, 0))
if ts_code == self.failed_code:
raise RuntimeError("simulated source failure")
return (Bar(ts_code, self.target, close=Decimal("10")),)
finally:
with self._lock:
self.active -= 1
class ConcurrentRepository:
def __init__(self) -> None:
self.bars: dict[tuple[str, date], Bar] = {}
self.daily_basic: dict[tuple[str, date], DailyBasic] = {}
self.audit_batches: list[tuple[SyncItemOutcome, ...]] = []
self.batch_summary: tuple[str, int, Decimal, bool] | None = None
self._lock = threading.Lock()
@contextmanager
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
yield True
def upsert_stocks(self, rows: Iterable[Stock]) -> WriteResult:
return WriteResult(inserted=len(tuple(rows)))
def upsert_bars(
self,
rows: Iterable[Bar],
window: SyncWindow,
*,
full_snapshot: bool,
) -> WriteResult:
records = tuple(rows)
inserted = 0
for row in records:
with self._lock:
if (row.ts_code, row.trade_date) not in self.bars:
inserted += 1
self.bars[(row.ts_code, row.trade_date)] = row
return WriteResult(inserted=inserted)
def upsert_daily_basic(self, rows: Iterable[DailyBasic], window: SyncWindow) -> WriteResult:
records = tuple(rows)
with self._lock:
for row in records:
self.daily_basic[(row.ts_code, row.trade_date)] = row
return WriteResult(inserted=len(records))
def purge_before(self, window: SyncWindow) -> None:
return None
def create_batch(
self,
target_trade_date: date,
window: SyncWindow,
mode: str,
parent_batch_id: str | None,
target_count: int,
) -> str:
return "batch-test"
def record_batch(
self,
batch_id: str,
status: str,
valid_count: int,
coverage: Decimal,
strategy_eligible: bool,
) -> None:
self.batch_summary = (status, valid_count, coverage, strategy_eligible)
def record_items(self, batch_id: str, outcomes: Iterable[object]) -> None:
self.audit_batches.append(tuple(outcomes)) # type: ignore[arg-type]
def record_item(
self,
batch_id: str,
item_kind: str,
item_key: str,
status: str,
result: WriteResult,
fingerprint: str | None = None,
error_type: str | None = None,
error_message: str | None = None,
) -> None:
return None
def count_valid_stocks(self, trade_date: date) -> int:
return sum(
1
for ts_code, current_date in self.bars
if current_date == trade_date and (ts_code, trade_date) in self.daily_basic
)
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]:
return ()
def has_bar(self, ts_code: str, trade_date: date) -> bool:
raise AssertionError("coverage must use count_valid_stocks")
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool:
raise AssertionError("coverage must use count_valid_stocks")
def _run_sync(
tmp_path: Path, delays: dict[str, float], failed_code: str | None = None
) -> tuple[SyncBatchSummary, ConcurrentSource, ConcurrentRepository]:
target = date(2024, 1, 2)
source = ConcurrentSource(target, delays, failed_code)
repository = ConcurrentRepository()
summary = SyncMarketData(
source,
CsvSnapshotStore(tmp_path),
repository,
today=target,
max_workers=2,
).execute(SyncMarketDataCommand(target_trade_date=target))
return summary, source, repository
def test_bar_workers_are_bounded_and_audit_is_batched(tmp_path: Path) -> None:
summary, source, repository = _run_sync(
tmp_path,
{"000001.SZ": 0.03, "000002.SZ": 0.01, "000003.SZ": 0.02, "000004.SZ": 0},
)
assert summary.status == "success"
assert source.max_active <= 2
assert summary.valid_count == 5
assert sum(len(batch) for batch in repository.audit_batches) == 7
def test_one_bar_failure_does_not_publish_a_csv_or_reduce_other_facts(tmp_path: Path) -> None:
failed_code = "000003.SZ"
summary, _, repository = _run_sync(
tmp_path,
{"000001.SZ": 0.02, "000002.SZ": 0.01, failed_code: 0},
failed_code,
)
assert summary.status == "partial_success"
assert summary.valid_count == 4
assert [failure.item_key for failure in summary.failures if failure.item_kind == "bar"] == [
failed_code
]
assert not (tmp_path / "bars" / f"{failed_code}.csv").exists()
assert not list((tmp_path / "bars").glob(f".{failed_code}.csv.*.tmp"))
assert len(repository.bars) == 4
def test_completion_order_does_not_change_aggregate_counts(tmp_path: Path) -> None:
first, _, _ = _run_sync(
tmp_path / "first",
{"000001.SZ": 0.03, "000002.SZ": 0, "000003.SZ": 0.02},
)
second, _, _ = _run_sync(
tmp_path / "second",
{"000001.SZ": 0, "000002.SZ": 0.03, "000003.SZ": 0.01},
)
assert (
first.status,
first.valid_count,
first.inserted_count,
first.updated_count,
first.unchanged_count,
first.failures,
) == (
second.status,
second.valid_count,
second.inserted_count,
second.updated_count,
second.unchanged_count,
second.failures,
)
def test_partial_success_always_uses_incomplete_exit_code() -> None:
summary = SyncBatchSummary(
batch_id="batch-test",
target_trade_date=date(2024, 1, 2),
window=SyncWindow(start=date(2018, 1, 2), end=date(2024, 1, 2)),
status="partial_success",
target_count=100,
valid_count=99,
coverage=Decimal("0.99"),
strategy_eligible=True,
)
assert summary.exit_code == 2
@@ -4,7 +4,10 @@ import pytest
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
from zhixing_server.modules.market_data.domain.models import SyncWindow
from zhixing_server.modules.market_data.infrastructure.tushare import TushareAdapter
from zhixing_server.modules.market_data.infrastructure.tushare import (
RequestCoordinator,
TushareAdapter,
)
def test_from_token_reuses_api_client_for_pro_bar(
@@ -44,3 +47,83 @@ def test_from_token_reuses_api_client_for_pro_bar(
assert calls
assert calls[0]["api"] is created_client
assert calls[0]["adj"] == "qfq"
assert calls[0]["retry_count"] == 1
def test_rate_limit_cooldown_is_shared_by_following_requests() -> None:
current = [0.0]
waits: list[float] = []
calls: list[float] = []
def clock() -> float:
return current[0]
def wait(seconds: float) -> None:
waits.append(seconds)
current[0] += seconds
coordinator = RequestCoordinator(
max_retries=0,
cooldown_seconds=(60, 120, 180),
clock=clock,
wait_fn=wait,
sleep_fn=wait,
)
def rate_limited() -> object:
calls.append(current[0])
raise RuntimeError("HTTP 429: too many requests")
with pytest.raises(RuntimeError):
coordinator.call("daily", rate_limited)
assert calls == [0]
assert coordinator.cooldown_until == 60
result = coordinator.call("adj_factor", lambda: calls.append(current[0]) or "ok")
assert result == "ok"
assert calls == [0, 60]
assert waits == [60]
def test_pro_bar_qfq_calls_are_bound_to_the_shared_coordinator(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class Client:
def __init__(self) -> None:
self.calls: list[str] = []
def daily(self, **kwargs: object) -> object:
self.calls.append("daily")
return object()
def adj_factor(self, **kwargs: object) -> object:
self.calls.append("adj_factor")
return object()
client = Client()
pro_bar_calls: list[dict[str, object]] = []
def fake_pro_api(token: str) -> object:
return client
def fake_pro_bar(**kwargs: object) -> list[dict[str, object]]:
pro_bar_calls.append(kwargs)
api = kwargs["api"]
assert api is client
api.daily(ts_code="000001.SZ") # type: ignore[union-attr]
api.adj_factor(ts_code="000001.SZ") # type: ignore[union-attr]
return [{"ts_code": "000001.SZ", "trade_date": "20240102", "close": "10"}]
monkeypatch.setattr(ts, "pro_api", fake_pro_api)
monkeypatch.setattr(ts, "pro_bar", fake_pro_bar)
adapter = TushareAdapter.from_token("test-token", request_interval_seconds=0)
adapter.fetch_bars(
"000001.SZ",
SyncWindow(start=date(2024, 1, 2), end=date(2024, 1, 2)),
)
assert client.calls == ["daily", "adj_factor"]
assert pro_bar_calls[0]["retry_count"] == 1
+16
View File
@@ -0,0 +1,16 @@
import pytest
from pydantic import ValidationError
from zhixing_server.bootstrap.config import Settings
def test_market_data_workers_default_to_eight(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("ZHIXING_MARKET_DATA_MAX_WORKERS", raising=False)
settings = Settings()
assert settings.market_data_max_workers == 8
def test_market_data_workers_must_be_positive() -> None:
with pytest.raises(ValidationError):
Settings(market_data_max_workers=0)
+17 -2
View File
@@ -412,6 +412,9 @@ wheels = [
binary = [
{ name = "psycopg-binary", marker = "implementation_name != 'pypy'" },
]
pool = [
{ name = "psycopg-pool" },
]
[[package]]
name = "psycopg-binary"
@@ -431,6 +434,18 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/25/8f/81dcbc2e8454b74d14881275ea45f00791052dac531a9fa8be1730d1685b/psycopg_binary-3.3.4-cp312-cp312-win_amd64.whl", hash = "sha256:494ca54901be8cf9eb7e02c25b731f2317c378efa44f43e8f9bd0e1184ae7be4", size = 3560782, upload-time = "2026-05-01T23:29:11.967Z" },
]
[[package]]
name = "psycopg-pool"
version = "3.3.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/90/82/7a23d26039827ecd4ebe93905651029ddd307c5182ad59296dfb6f67b528/psycopg_pool-3.3.1.tar.gz", hash = "sha256:b10b10b7a175d5cc1592147dc5b7eec8a9e0834eb3ed2c4a92c858e2f51eb63c", size = 31661, upload-time = "2026-05-01T23:31:59.809Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/37/ed/89c2c620af0e1660354cd8aabf9f5b21f911597ce22acb37c805d6c86bc8/psycopg_pool-3.3.1-py3-none-any.whl", hash = "sha256:2af5b432941c4c9ad5c87b3fa410aec910ec8f7c122855897983a06c45f2e4b5", size = 40023, upload-time = "2026-05-01T23:31:53.136Z" },
]
[[package]]
name = "pydantic"
version = "2.13.4"
@@ -879,7 +894,7 @@ dependencies = [
{ name = "fastapi" },
{ name = "numpy" },
{ name = "pandas" },
{ name = "psycopg", extra = ["binary"] },
{ name = "psycopg", extra = ["binary", "pool"] },
{ name = "pydantic-settings" },
{ name = "sqlalchemy" },
{ name = "tushare" },
@@ -902,7 +917,7 @@ requires-dist = [
{ name = "fastapi", specifier = ">=0.141.1" },
{ name = "numpy", specifier = ">=2.4.0" },
{ name = "pandas", specifier = ">=2.3.3" },
{ name = "psycopg", extras = ["binary"], specifier = ">=3.3.2" },
{ name = "psycopg", extras = ["binary", "pool"], specifier = ">=3.3.2" },
{ name = "pydantic-settings", specifier = ">=2.14.2" },
{ name = "sqlalchemy", specifier = ">=2.0.46" },
{ name = "tushare", specifier = ">=1.4.24" },
+2 -1
View File
@@ -24,6 +24,7 @@ import {
function useActiveRoutePath() {
const matchRoute = useMatchRoute()
if (matchRoute({ to: "/selection", fuzzy: true })) return "/selection"
if (matchRoute({ to: "/sync", fuzzy: true })) return "/sync"
if (matchRoute({ to: "/components", fuzzy: true })) return "/components"
if (matchRoute({ to: "/", fuzzy: false })) return "/"
return "/"
@@ -244,7 +245,7 @@ function MoreSheet({
<DialogHeader>
<DialogTitle>更多功能</DialogTitle>
<DialogDescription>
行情数据与同步任务尚未上线,暂无更多可用入口。
行情数据尚未上线,暂无更多可用入口。
</DialogDescription>
</DialogHeader>
</DialogContent>
@@ -0,0 +1,28 @@
import { describe, expect, it } from "vitest"
import {
mobilePrimaryNavigation,
primaryNavigation,
routePresentation,
} from "./navigation"
describe("sync navigation", () => {
it("exposes the sync route in desktop and mobile navigation", () => {
const sync = primaryNavigation.find((item) => item.id === "sync")
expect(sync).toMatchObject({
availability: "available",
label: "同步任务",
to: "/sync",
})
expect(mobilePrimaryNavigation).toContain(sync)
})
it("provides the active route presentation for sync", () => {
expect(routePresentation["/sync"]).toEqual({
breadcrumb: "研究工作台",
id: "sync",
title: "市场数据完整性检查",
})
})
})
+7 -2
View File
@@ -35,9 +35,9 @@ export const primaryNavigation: readonly NavigationItem[] = [
{
id: "sync",
label: "同步任务",
to: null,
to: "/sync",
icon: ClipboardCheck,
availability: "unavailable",
availability: "available",
},
{
id: "selection",
@@ -79,6 +79,11 @@ export const routePresentation: Record<string, RoutePresentation> = {
breadcrumb: "研究工作台",
title: "知行 B1 执行结果",
},
"/sync": {
id: "sync",
breadcrumb: "研究工作台",
title: "市场数据完整性检查",
},
"/components": {
id: "components",
breadcrumb: "研究工作台",
@@ -0,0 +1,59 @@
import { beforeEach, describe, expect, it, vi } from "vitest"
const requestJson = vi.hoisted(() => vi.fn())
vi.mock("@/shared/api/request-json", () => ({ requestJson }))
import {
getIntegrityCheck,
getLatestIntegrityCheck,
triggerIntegrityCheck,
} from "./sync.api"
describe("sync API adapters", () => {
beforeEach(() => {
requestJson.mockReset()
})
it("triggers one read-only integrity check through the POST endpoint", async () => {
const signal = new AbortController().signal
await triggerIntegrityCheck(signal)
expect(requestJson).toHaveBeenCalledWith(
"/api/v1/market-data/integrity-checks",
{
method: "POST",
signal,
},
)
})
it("uses the latest report pagination contract", async () => {
await getLatestIntegrityCheck({ page: 2, pageSize: 20 })
const [input, init] = requestJson.mock.calls[0] as [
string,
{ signal?: AbortSignal },
]
const params = new URL(input, "http://localhost").searchParams
expect(input).toContain("/api/v1/market-data/integrity-checks/latest?")
expect(params.get("page")).toBe("2")
expect(params.get("page_size")).toBe("20")
expect(init).toEqual({ signal: undefined })
})
it("encodes a check id while keeping the paged check contract", async () => {
await getIntegrityCheck("check/with spaces", { page: 3, pageSize: 50 })
const [input] = requestJson.mock.calls[0] as [string]
const params = new URL(input, "http://localhost").searchParams
expect(input).toContain(
"/api/v1/market-data/integrity-checks/check%2Fwith%20spaces?",
)
expect(params.get("page")).toBe("3")
expect(params.get("page_size")).toBe("50")
})
})
@@ -0,0 +1,50 @@
import { requestJson } from "@/shared/api/request-json"
import {
defaultMarketIntegrityQuery,
type IntegrityCheckAccepted,
type MarketIntegrityCheck,
type MarketIntegrityQuery,
} from "./sync.types"
const integrityChecksPath = "/api/v1/market-data/integrity-checks"
export function triggerIntegrityCheck(signal?: AbortSignal) {
return requestJson<IntegrityCheckAccepted>(integrityChecksPath, {
method: "POST",
signal,
})
}
export function getLatestIntegrityCheck(
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
signal?: AbortSignal,
) {
return requestJson<MarketIntegrityCheck>(
`${integrityChecksPath}/latest?${buildQuery(query)}`,
{ signal },
)
}
export function getIntegrityCheck(
checkId: string,
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
signal?: AbortSignal,
) {
return requestJson<MarketIntegrityCheck>(
`${integrityChecksPath}/${encodeURIComponent(checkId)}?${buildQuery(query)}`,
{ signal },
)
}
function buildQuery(query: MarketIntegrityQuery) {
const params = new URLSearchParams({
page: String(query.page),
page_size: String(query.pageSize),
})
return params.toString()
}
export const getMarketIntegrityCheck = getIntegrityCheck
export const getLatestMarketIntegrityCheck = getLatestIntegrityCheck
export const triggerMarketIntegrityCheck = triggerIntegrityCheck
@@ -0,0 +1,111 @@
import { beforeEach, describe, expect, it, vi } from "vitest"
import { ApiError } from "@/shared/api/request-json"
const queryHooks = vi.hoisted(() => ({
useMutation: vi.fn(),
useQuery: vi.fn(),
useQueryClient: vi.fn(),
}))
const api = vi.hoisted(() => ({
getIntegrityCheck: vi.fn(),
getLatestIntegrityCheck: vi.fn(),
triggerIntegrityCheck: vi.fn(),
}))
vi.mock("@tanstack/react-query", () => queryHooks)
vi.mock("./sync.api", () => api)
import {
marketIntegrityCheckQueryKey,
marketIntegrityLatestQueryKey,
useIntegrityCheck,
useLatestIntegrityCheck,
useTriggerIntegrityCheck,
} from "./sync.query"
interface QueryOptionsForTest {
queryKey: readonly unknown[]
enabled?: boolean
refetchInterval?: (query: {
state: { data?: { status?: string } }
}) => number | false
}
interface MutationOptionsForTest {
onError?: (error: unknown) => void
onSuccess?: () => void
}
describe("sync query hooks", () => {
const queryClient = { invalidateQueries: vi.fn() }
beforeEach(() => {
vi.clearAllMocks()
queryHooks.useQuery.mockImplementation((options) => options)
queryHooks.useQueryClient.mockReturnValue(queryClient)
queryHooks.useMutation.mockImplementation((options) => options)
})
it("keeps latest and check keys under marketIntegrity", () => {
expect(marketIntegrityLatestQueryKey({ page: 2, pageSize: 20 })).toEqual([
"marketIntegrity",
"latest",
2,
20,
])
expect(
marketIntegrityCheckQueryKey("check-1", { page: 3, pageSize: 50 }),
).toEqual(["marketIntegrity", "check", "check-1", 3, 50])
})
it("polls only while a check response is running", () => {
useLatestIntegrityCheck({ page: 1, pageSize: 10 })
useIntegrityCheck("check-1", { page: 1, pageSize: 10 })
const latest = queryHooks.useQuery.mock.calls[0]?.[0] as QueryOptionsForTest
const check = queryHooks.useQuery.mock.calls[1]?.[0] as QueryOptionsForTest
expect(latest.queryKey).toEqual(["marketIntegrity", "latest", 1, 10])
expect(check.queryKey).toEqual([
"marketIntegrity",
"check",
"check-1",
1,
10,
])
expect(check.enabled).toBe(true)
if (!check.refetchInterval) throw new Error("polling callback is missing")
expect(
check.refetchInterval({ state: { data: { status: "running" } } }),
).toBe(1500)
expect(
check.refetchInterval({ state: { data: { status: "passed" } } }),
).toBe(false)
})
it("disables a check query without an active id", () => {
useIntegrityCheck(null)
const check = queryHooks.useQuery.mock.calls[0]?.[0] as QueryOptionsForTest
expect(check.enabled).toBe(false)
expect(check.queryKey).toEqual(["marketIntegrity", "check", "none", 1, 10])
})
it("invalidates latest reports after success and on a 409 conflict", () => {
useTriggerIntegrityCheck()
const mutation = queryHooks.useMutation.mock
.calls[0]?.[0] as MutationOptionsForTest
mutation.onSuccess?.()
mutation.onError?.(new ApiError(409, "already running"))
expect(queryClient.invalidateQueries).toHaveBeenCalledTimes(2)
expect(queryClient.invalidateQueries).toHaveBeenNthCalledWith(1, {
queryKey: ["marketIntegrity", "latest"],
})
expect(queryClient.invalidateQueries).toHaveBeenNthCalledWith(2, {
queryKey: ["marketIntegrity", "latest"],
})
})
})
@@ -0,0 +1,93 @@
import {
useMutation,
useQuery,
useQueryClient,
type QueryClient,
} from "@tanstack/react-query"
import { ApiError } from "@/shared/api/request-json"
import {
getIntegrityCheck,
getLatestIntegrityCheck,
triggerIntegrityCheck,
} from "./sync.api"
import {
defaultMarketIntegrityQuery,
type IntegrityCheckAccepted,
type MarketIntegrityCheck,
type MarketIntegrityQuery,
} from "./sync.types"
export const marketIntegrityQueryKey = ["marketIntegrity"] as const
export const marketIntegrityLatestQueryPrefix = [
"marketIntegrity",
"latest",
] as const
export const marketIntegrityLatestQueryKey = (
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
) => [...marketIntegrityLatestQueryPrefix, query.page, query.pageSize] as const
export const marketIntegrityCheckQueryKey = (
checkId: string,
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
) => ["marketIntegrity", "check", checkId, query.page, query.pageSize] as const
const MARKET_INTEGRITY_POLL_INTERVAL_MS = 1500
export function useLatestIntegrityCheck(
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
) {
return useQuery({
queryFn: ({ signal }) => getLatestIntegrityCheck(query, signal),
queryKey: marketIntegrityLatestQueryKey(query),
})
}
export function useIntegrityCheck(
checkId: string | null,
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
) {
return useQuery({
enabled: Boolean(checkId),
queryFn: ({ signal }) => getIntegrityCheck(checkId ?? "", query, signal),
queryKey: marketIntegrityCheckQueryKey(checkId ?? "none", query),
refetchInterval: (currentQuery) =>
currentQuery.state.data?.status === "running"
? MARKET_INTEGRITY_POLL_INTERVAL_MS
: false,
})
}
export function useTriggerIntegrityCheck() {
const queryClient = useQueryClient()
return useMutation({
mutationFn: ({ signal }: { signal?: AbortSignal } = {}) =>
triggerIntegrityCheck(signal),
onError: (error) => {
if (error instanceof ApiError && error.status === 409) {
void invalidateLatestIntegrityChecks(queryClient)
}
},
onSuccess: () => {
void invalidateLatestIntegrityChecks(queryClient)
},
})
}
export function invalidateLatestIntegrityChecks(queryClient: QueryClient) {
return queryClient.invalidateQueries({
queryKey: marketIntegrityLatestQueryPrefix,
})
}
export const useLatestMarketIntegrityCheck = useLatestIntegrityCheck
export const useMarketIntegrityCheck = useIntegrityCheck
export const useTriggerMarketIntegrityCheck = useTriggerIntegrityCheck
export type MarketIntegrityMutation = ReturnType<
typeof useTriggerIntegrityCheck
>
export type MarketIntegrityQueryResult = ReturnType<typeof useIntegrityCheck>
export type MarketIntegrityAccepted = IntegrityCheckAccepted
export type MarketIntegrityData = MarketIntegrityCheck
@@ -0,0 +1,45 @@
export type MarketIntegrityStatus =
"no_data" | "running" | "passed" | "issues_found" | "failed"
export interface MarketIntegrityQuery {
page: number
pageSize: number
}
export interface IntegrityCheckAccepted {
check_id: string
status: "running"
window_start: string
window_end: string
}
export interface IntegrityIssue {
issue_key: string
item_kind: string
item_key: string
issue_type: string
message: string
}
export interface MarketIntegrityCheck {
check_id: string | null
status: MarketIntegrityStatus
window_start: string | null
window_end: string | null
target_count: number
checked_count: number
issue_count: number
error_type: string | null
error_message: string | null
created_at: string | null
finished_at: string | null
page: number
page_size: number
issues_total: number
issues: IntegrityIssue[]
}
export const defaultMarketIntegrityQuery: MarketIntegrityQuery = {
page: 1,
pageSize: 10,
}
@@ -0,0 +1,214 @@
import { fireEvent, render, screen, waitFor } from "@testing-library/react"
import { beforeEach, describe, expect, it, vi } from "vitest"
import { ApiError } from "@/shared/api/request-json"
import { SyncPage } from "./sync-page"
import type {
IntegrityIssue,
MarketIntegrityCheck,
MarketIntegrityStatus,
} from "../api/sync.types"
const syncHooks = vi.hoisted(() => ({
useIntegrityCheck: vi.fn(),
useLatestIntegrityCheck: vi.fn(),
useTriggerIntegrityCheck: vi.fn(),
}))
vi.mock("../api/sync.query", () => syncHooks)
const issue: IntegrityIssue = {
issue_key: "issue-1",
item_kind: "bar",
item_key: "000001.SZ:2026-08-08",
issue_type: "missing_csv_row",
message: "仅报告 PostgreSQL 与 CSV 的差异,不会自动修复。",
}
function buildReport(
status: MarketIntegrityStatus,
overrides: Partial<MarketIntegrityCheck> = {},
): MarketIntegrityCheck {
const hasRun = status !== "no_data"
return {
check_id: hasRun ? "check-1" : null,
status,
window_start: hasRun ? "2026-08-01" : null,
window_end: hasRun ? "2026-08-08" : null,
target_count: 10,
checked_count: 10,
issue_count: status === "issues_found" ? 1 : 0,
error_type: null,
error_message: null,
created_at: "2026-08-08T09:00:00+08:00",
finished_at: "2026-08-08T09:02:00+08:00",
page: 1,
page_size: 10,
issues_total: status === "issues_found" ? 1 : 0,
issues: status === "issues_found" ? [issue] : [],
...overrides,
}
}
function configureHooks(
latestData: MarketIntegrityCheck | undefined,
options: {
checkData?: MarketIntegrityCheck
latestError?: unknown
triggerError?: unknown
} = {},
) {
const refetchLatest = vi.fn()
const refetchCheck = vi.fn()
const mutate = vi.fn()
const reset = vi.fn()
syncHooks.useLatestIntegrityCheck.mockReturnValue({
data: latestData,
error: options.latestError,
isError: Boolean(options.latestError),
isPending: false,
refetch: refetchLatest,
})
syncHooks.useIntegrityCheck.mockReturnValue({
data: options.checkData,
error: undefined,
isError: false,
isPending: false,
refetch: refetchCheck,
})
syncHooks.useTriggerIntegrityCheck.mockReturnValue({
data: undefined,
error: options.triggerError,
isError: Boolean(options.triggerError),
isPending: false,
mutate,
reset,
})
return { mutate, refetchCheck, refetchLatest, reset }
}
describe("SyncPage", () => {
beforeEach(() => {
vi.clearAllMocks()
configureHooks(buildReport("no_data"))
})
it("explains the no-data state and starts a check once", () => {
const { mutate } = configureHooks(buildReport("no_data"))
render(<SyncPage />)
expect(screen.getByText("暂无检查记录")).toBeInTheDocument()
expect(screen.getAllByText(/只报告,不自动修复/).length).toBeGreaterThan(0)
fireEvent.click(screen.getByRole("button", { name: "开始完整性检查" }))
expect(mutate).toHaveBeenCalledOnce()
})
it("shows running progress and disables duplicate triggers", () => {
configureHooks(buildReport("running", { checked_count: 4 }))
render(<SyncPage />)
expect(screen.getByText("完整性检查进行中")).toBeInTheDocument()
expect(screen.getByText("4 / 10")).toBeInTheDocument()
expect(screen.getByRole("button", { name: "检查进行中" })).toBeDisabled()
})
it("recovers a running check id from the latest report after refresh", async () => {
configureHooks(buildReport("running"), {
checkData: buildReport("running", { checked_count: 3 }),
})
render(<SyncPage />)
await waitFor(() => {
expect(syncHooks.useIntegrityCheck).toHaveBeenCalledWith("check-1", {
page: 1,
pageSize: 10,
})
})
})
it("refreshes latest once after a running check reaches a terminal state", async () => {
const { refetchLatest } = configureHooks(buildReport("running"), {
checkData: buildReport("passed"),
})
render(<SyncPage />)
await waitFor(() => {
expect(refetchLatest).toHaveBeenCalledOnce()
})
})
it("renders a passed report without an issue table", () => {
configureHooks(buildReport("passed"))
render(<SyncPage />)
expect(screen.getByText("检查通过")).toBeInTheDocument()
expect(screen.getAllByText(/未发现 PostgreSQL\/CSV 不一致/)).toHaveLength(2)
expect(screen.queryByRole("table")).not.toBeInTheDocument()
})
it("renders paginated issue details and the safe-report notice", () => {
configureHooks(buildReport("issues_found"))
render(<SyncPage />)
expect(
screen.getByRole("table", { name: "完整性检查问题报告" }),
).toBeInTheDocument()
expect(screen.getByText("bar")).toBeInTheDocument()
expect(screen.getByText("000001.SZ:2026-08-08")).toBeInTheDocument()
expect(screen.getByText("missing_csv_row")).toBeInTheDocument()
expect(screen.getByText(/只展示对象和差异说明/)).toBeInTheDocument()
expect(
screen.getByRole("contentinfo", { name: "表格分页" }),
).toBeInTheDocument()
})
it("shows a failed report and offers a retry entry", () => {
const { mutate } = configureHooks(
buildReport("failed", {
error_type: "check_error",
error_message: "存储暂时不可用",
}),
)
render(<SyncPage />)
expect(screen.getByText("检查未完成")).toBeInTheDocument()
expect(screen.getByText(/不能据此判断市场数据已经损坏/)).toBeInTheDocument()
fireEvent.click(screen.getByRole("button", { name: "重新发起检查" }))
expect(mutate).toHaveBeenCalledOnce()
})
it("makes 503 and network failures visible with a retry button", () => {
const { refetchLatest } = configureHooks(undefined, {
latestError: new ApiError(503, "unavailable"),
})
render(<SyncPage />)
expect(screen.getByText("完整性检查结果暂时不可用")).toBeInTheDocument()
fireEvent.click(screen.getByRole("button", { name: "重试" }))
expect(refetchLatest).toHaveBeenCalledOnce()
})
it("explains a 409 and lets the page take over the running report", () => {
configureHooks(buildReport("running"), {
triggerError: new ApiError(409, "conflict"),
})
render(<SyncPage />)
expect(screen.getByText(/已刷新并接管该检查/)).toBeInTheDocument()
expect(screen.getByText("完整性检查进行中")).toBeInTheDocument()
})
})
@@ -0,0 +1,665 @@
import {
AlertTriangle,
CheckCircle2,
ClipboardCheck,
Database,
RefreshCw,
ShieldCheck,
} from "lucide-react"
import { useEffect, useRef, useState } from "react"
import { PageLayout } from "@/app/layout/page-layout"
import { ApiError } from "@/shared/api/request-json"
import { Badge } from "@/shared/ui/badge"
import { Button } from "@/shared/ui/button"
import {
Card,
CardContent,
CardDescription,
CardHeader,
CardTitle,
} from "@/shared/ui/card"
import { Pagination } from "@/shared/ui/pagination"
import { Progress, ProgressLabel, ProgressValue } from "@/shared/ui/progress"
import {
useIntegrityCheck,
useLatestIntegrityCheck,
useTriggerIntegrityCheck,
} from "../api/sync.query"
import {
defaultMarketIntegrityQuery,
type IntegrityCheckAccepted,
type MarketIntegrityCheck,
} from "../api/sync.types"
export function SyncPage() {
const [triggeredCheckId, setTriggeredCheckId] = useState<string | null>(null)
const [page, setPage] = useState(defaultMarketIntegrityQuery.page)
const [pageSize, setPageSize] = useState(defaultMarketIntegrityQuery.pageSize)
const query = { page, pageSize }
const latest = useLatestIntegrityCheck(query)
const latestRunningCheckId =
latest.data?.status === "running" ? latest.data.check_id : null
const activeCheckId = triggeredCheckId ?? latestRunningCheckId
const check = useIntegrityCheck(activeCheckId, query)
const trigger = useTriggerIntegrityCheck()
const { refetch: refetchLatest } = latest
const terminalRefreshId = useRef<string | null>(null)
useEffect(() => {
if (!activeCheckId || !check.data || check.data.status === "running") {
if (!activeCheckId) terminalRefreshId.current = null
return
}
if (terminalRefreshId.current === activeCheckId) return
terminalRefreshId.current = activeCheckId
void refetchLatest()
}, [activeCheckId, check.data, refetchLatest])
const acceptedReport = trigger.data
? acceptedCheckAsReport(trigger.data)
: undefined
const latestMatchesActive =
!activeCheckId || latest.data?.check_id === activeCheckId
const report =
check.data ??
(activeCheckId && latestMatchesActive ? latest.data : undefined) ??
(trigger.data ? acceptedReport : latest.data)
function startCheck() {
setPage(1)
trigger.reset()
trigger.mutate(
{},
{
onError: (error) => {
if (getErrorStatus(error) === 409) setTriggeredCheckId(null)
},
onSuccess: (accepted) => setTriggeredCheckId(accepted.check_id),
},
)
}
function retryActiveCheck() {
void check.refetch()
}
function retryLatest() {
void latest.refetch()
}
const triggerStatus = getErrorStatus(trigger.error)
const triggerError =
trigger.isError && triggerStatus !== 409 ? trigger.error : null
const queryPending = activeCheckId ? check.isPending : latest.isPending
const queryError = activeCheckId ? check.error : latest.error
const queryHasData = Boolean(report)
return (
<PageLayout>
<div className="mx-auto w-full max-w-6xl space-y-4">
<header className="space-y-1">
<p className="text-xs font-medium tracking-[0.08em] text-muted-foreground uppercase">
只读数据校验
</p>
<h1 className="text-2xl font-semibold tracking-tight">
市场数据完整性检查
</h1>
<p className="max-w-3xl text-sm text-muted-foreground">
检查 PostgreSQL 与 CSV 快照中的市场数据是否一致,帮助定位数据问题。
</p>
</header>
<ReadOnlyNotice />
{triggerStatus === 409 ? <ConflictNotice /> : null}
{triggerError ? (
<RequestError
description={getTriggerErrorDescription(triggerError)}
onRetry={startCheck}
title="完整性检查暂时无法发起"
/>
) : null}
{!queryHasData && queryPending ? <LoadingState /> : null}
{!queryPending && queryError ? (
<RequestError
description={getQueryErrorDescription(queryError)}
onRetry={activeCheckId ? retryActiveCheck : retryLatest}
title="完整性检查结果暂时不可用"
/>
) : null}
{report?.status === "no_data" ? (
<NoDataState isPending={trigger.isPending} onStart={startCheck} />
) : null}
{report?.status === "running" ? <RunningState report={report} /> : null}
{report?.status === "passed" ? (
<PassedState
isPending={trigger.isPending}
onStart={startCheck}
report={report}
/>
) : null}
{report?.status === "issues_found" ? (
<IssuesFoundState
onPageChange={setPage}
onPageSizeChange={(nextPageSize) => {
setPage(1)
setPageSize(nextPageSize)
}}
report={report}
selectedPage={page}
selectedPageSize={pageSize}
isPending={trigger.isPending}
onStart={startCheck}
/>
) : null}
{report?.status === "failed" ? (
<FailedState
isPending={trigger.isPending}
onStart={startCheck}
report={report}
/>
) : null}
</div>
</PageLayout>
)
}
function ReadOnlyNotice() {
return (
<Card className="border-primary/20 bg-primary/5">
<CardContent className="flex gap-3 p-4">
<ShieldCheck
className="mt-0.5 size-5 shrink-0 text-primary"
aria-hidden="true"
/>
<div className="space-y-1">
<p className="font-medium">只读检查,安全报告</p>
<p className="text-sm text-muted-foreground">
本检查只比较 PostgreSQL 与 CSV,不访问
Tushare,不修改市场事实。发现差异时只报告,不自动修复。
</p>
</div>
</CardContent>
</Card>
)
}
function ConflictNotice() {
return (
<div
role="status"
className="flex items-start gap-2 rounded-md border border-warning/40 bg-warning/10 p-3 text-sm text-foreground"
>
<RefreshCw
className="mt-0.5 size-4 shrink-0 text-warning"
aria-hidden="true"
/>
<p>已有完整性检查正在运行,页面已刷新并接管该检查的进度。</p>
</div>
)
}
function LoadingState() {
return (
<Card aria-label="正在加载完整性检查" role="status">
<CardHeader className="gap-3">
<div className="h-6 w-52 animate-pulse rounded-md bg-muted" />
<div className="h-4 w-80 max-w-full animate-pulse rounded-md bg-muted" />
</CardHeader>
<CardContent>
<div className="h-20 animate-pulse rounded-md bg-muted" />
</CardContent>
</Card>
)
}
function NoDataState({
isPending,
onStart,
}: {
isPending: boolean
onStart: () => void
}) {
return (
<Card>
<CardHeader>
<CardTitle className="flex items-center gap-2 text-xl">
<Database className="size-5 text-primary" aria-hidden="true" />
暂无检查记录
</CardTitle>
<CardDescription>
当前没有可查看的完整性检查。发起后,系统会比较最近完成同步窗口中的
PostgreSQL 与 CSV 数据。
</CardDescription>
</CardHeader>
<CardContent className="space-y-3">
<Button disabled={isPending} onClick={onStart}>
<ClipboardCheck aria-hidden="true" />
{isPending ? "正在发起检查…" : "开始完整性检查"}
</Button>
<p className="text-xs text-muted-foreground">
检查只读数据并生成报告,不会访问 Tushare,也不会自动修复。
</p>
</CardContent>
</Card>
)
}
function RunningState({ report }: { report: MarketIntegrityCheck }) {
const percentage = progressPercentage(
report.checked_count,
report.target_count,
)
return (
<Card aria-live="polite">
<CardHeader className="gap-3">
<div className="flex flex-wrap items-start justify-between gap-3">
<div>
<CardTitle className="flex items-center gap-2 text-xl">
<ClipboardCheck
className="size-5 text-primary"
aria-hidden="true"
/>
完整性检查进行中
</CardTitle>
<CardDescription>
检查完成后会自动停止刷新;运行期间不会修改 PostgreSQL 或 CSV
数据。
</CardDescription>
</div>
<Badge variant="outline">运行中</Badge>
</div>
</CardHeader>
<CardContent className="space-y-5">
<Progress aria-label="完整性检查进度" max={100} value={percentage}>
<ProgressLabel>检查进度</ProgressLabel>
<ProgressValue>
{() => `${report.checked_count}/${report.target_count}`}
</ProgressValue>
</Progress>
<div className="grid gap-3 sm:grid-cols-3">
<Metric
label="已检查对象"
value={`${report.checked_count} / ${report.target_count}`}
/>
<Metric label="当前问题数" value={String(report.issue_count)} />
<Metric label="检查窗口" value={formatWindow(report)} />
</div>
<div className="flex flex-wrap items-center justify-between gap-3 border-t border-border/60 pt-4">
<Timestamp label="开始时间" value={report.created_at} />
<Button disabled variant="outline">
检查进行中
</Button>
</div>
</CardContent>
</Card>
)
}
function PassedState({
isPending,
onStart,
report,
}: {
isPending: boolean
onStart: () => void
report: MarketIntegrityCheck
}) {
return (
<Card>
<CardHeader className="gap-3">
<div className="flex flex-wrap items-start justify-between gap-3">
<div>
<CardTitle className="flex items-center gap-2 text-xl">
<CheckCircle2
className="size-5 text-success"
aria-hidden="true"
/>
检查通过
</CardTitle>
<CardDescription>
未发现 PostgreSQL/CSV 不一致,本次检查只报告结果,没有修改数据。
</CardDescription>
</div>
<Badge className="border-transparent bg-success text-success-foreground">
已通过
</Badge>
</div>
</CardHeader>
<CardContent className="space-y-4">
<div className="rounded-md border border-success/30 bg-success/10 p-4 text-sm">
未发现 PostgreSQL/CSV 不一致。
</div>
<div className="grid gap-3 sm:grid-cols-3">
<Metric
label="已检查对象"
value={`${report.checked_count} / ${report.target_count}`}
/>
<Metric label="问题数" value={String(report.issue_count)} />
<Metric label="检查窗口" value={formatWindow(report)} />
</div>
<div className="flex flex-wrap items-end justify-between gap-3 border-t border-border/60 pt-4">
<div className="flex flex-wrap gap-x-5 gap-y-2">
<Timestamp label="开始时间" value={report.created_at} />
<Timestamp label="完成时间" value={report.finished_at} />
</div>
<Button disabled={isPending} onClick={onStart} variant="outline">
<RefreshCw aria-hidden="true" />
{isPending ? "正在发起…" : "再次检查"}
</Button>
</div>
</CardContent>
</Card>
)
}
function IssuesFoundState({
isPending,
onPageChange,
onPageSizeChange,
onStart,
report,
selectedPage,
selectedPageSize,
}: {
isPending: boolean
onPageChange: (page: number) => void
onPageSizeChange: (pageSize: number) => void
onStart: () => void
report: MarketIntegrityCheck
selectedPage: number
selectedPageSize: number
}) {
return (
<Card>
<CardHeader className="gap-3">
<div className="flex flex-wrap items-start justify-between gap-3">
<div>
<CardTitle className="flex items-center gap-2 text-xl">
<AlertTriangle
className="size-5 text-warning"
aria-hidden="true"
/>
发现数据差异
</CardTitle>
<CardDescription>
检查完成并发现 {report.issue_count}{" "}
个问题。以下内容仅供定位,系统不会自动修复。
</CardDescription>
</div>
<Badge className="border-transparent bg-warning text-warning-foreground">
{report.issue_count} 个问题
</Badge>
</div>
</CardHeader>
<CardContent className="space-y-4">
<div className="grid gap-3 sm:grid-cols-3">
<Metric
label="已检查对象"
value={`${report.checked_count} / ${report.target_count}`}
/>
<Metric label="问题数" value={String(report.issue_count)} />
<Metric label="检查窗口" value={formatWindow(report)} />
</div>
<div className="rounded-md border border-warning/40 bg-warning/10 p-4 text-sm">
安全说明:问题报告只展示对象和差异说明,不会执行修复或重新同步。
</div>
<IssueTable report={report} />
<Pagination
onPageChange={onPageChange}
onPageSizeChange={onPageSizeChange}
page={selectedPage}
pageSize={selectedPageSize}
pageSizeOptions={[10, 20, 50]}
total={report.issues_total}
/>
<div className="flex flex-wrap items-end justify-between gap-3 border-t border-border/60 pt-4">
<div className="flex flex-wrap gap-x-5 gap-y-2">
<Timestamp label="开始时间" value={report.created_at} />
<Timestamp label="完成时间" value={report.finished_at} />
</div>
<Button disabled={isPending} onClick={onStart} variant="outline">
<RefreshCw aria-hidden="true" />
{isPending ? "正在发起…" : "再次检查"}
</Button>
</div>
</CardContent>
</Card>
)
}
function IssueTable({ report }: { report: MarketIntegrityCheck }) {
if (report.issues.length === 0) {
return (
<div className="rounded-md border p-4 text-sm text-muted-foreground">
当前页没有问题记录。
</div>
)
}
return (
<div className="overflow-x-auto rounded-md border">
<table className="w-full min-w-[680px] text-left text-sm">
<caption className="sr-only">完整性检查问题报告</caption>
<thead className="bg-muted/50 text-xs text-muted-foreground">
<tr>
<th className="px-3 py-2.5 font-medium" scope="col">
对象类型
</th>
<th className="px-3 py-2.5 font-medium" scope="col">
对象 key
</th>
<th className="px-3 py-2.5 font-medium" scope="col">
问题类型
</th>
<th className="px-3 py-2.5 font-medium" scope="col">
安全说明
</th>
</tr>
</thead>
<tbody className="divide-y divide-border/60">
{report.issues.map((issue) => (
<tr key={issue.issue_key} className="align-top">
<td className="px-3 py-3 font-medium">{issue.item_kind}</td>
<td className="break-all px-3 py-3 font-mono text-xs">
{issue.item_key}
</td>
<td className="px-3 py-3">{issue.issue_type}</td>
<td className="max-w-md px-3 py-3 text-muted-foreground">
{issue.message}
</td>
</tr>
))}
</tbody>
</table>
</div>
)
}
function FailedState({
isPending,
onStart,
report,
}: {
isPending: boolean
onStart: () => void
report: MarketIntegrityCheck
}) {
return (
<Card>
<CardHeader className="gap-3">
<div className="flex flex-wrap items-start justify-between gap-3">
<div>
<CardTitle className="flex items-center gap-2 text-xl">
<AlertTriangle
className="size-5 text-destructive"
aria-hidden="true"
/>
检查未完成
</CardTitle>
<CardDescription>
本次检查遇到安全错误,不能据此判断市场数据已经损坏。
</CardDescription>
</div>
<Badge variant="destructive">失败</Badge>
</div>
</CardHeader>
<CardContent className="space-y-4">
<div
role="alert"
className="rounded-md border border-destructive/30 bg-destructive/10 p-4 text-sm"
>
<p className="font-medium">未判断市场数据状态</p>
<p className="mt-1 text-muted-foreground">
{report.error_type
? `错误类型:${report.error_type}`
: "检查服务暂时无法完成比较。"}
</p>
{report.error_message ? (
<p className="mt-1 text-muted-foreground">{report.error_message}</p>
) : null}
</div>
<div className="grid gap-3 sm:grid-cols-3">
<Metric
label="已检查对象"
value={`${report.checked_count} / ${report.target_count}`}
/>
<Metric label="问题数" value={String(report.issue_count)} />
<Metric label="检查窗口" value={formatWindow(report)} />
</div>
<div className="flex flex-wrap items-end justify-between gap-3 border-t border-border/60 pt-4">
<div className="flex flex-wrap gap-x-5 gap-y-2">
<Timestamp label="开始时间" value={report.created_at} />
<Timestamp label="结束时间" value={report.finished_at} />
</div>
<Button disabled={isPending} onClick={onStart}>
<RefreshCw aria-hidden="true" />
{isPending ? "正在重新发起…" : "重新发起检查"}
</Button>
</div>
</CardContent>
</Card>
)
}
function RequestError({
description,
onRetry,
title,
}: {
description: string
onRetry: () => void
title: string
}) {
return (
<Card role="alert">
<CardHeader>
<CardTitle className="flex items-center gap-2 text-xl">
<AlertTriangle
className="size-5 text-destructive"
aria-hidden="true"
/>
{title}
</CardTitle>
<CardDescription>{description}</CardDescription>
</CardHeader>
<CardContent>
<Button onClick={onRetry} variant="outline">
<RefreshCw aria-hidden="true" />
重试
</Button>
</CardContent>
</Card>
)
}
function Metric({ label, value }: { label: string; value: string }) {
return (
<div className="rounded-md border bg-muted/30 p-3">
<p className="text-xs text-muted-foreground">{label}</p>
<p className="mt-1 break-words font-medium tabular-nums">{value}</p>
</div>
)
}
function Timestamp({ label, value }: { label: string; value: string | null }) {
return (
<p className="text-xs text-muted-foreground">
{label}:
<time className="text-foreground" dateTime={value ?? undefined}>
{value ?? "—"}
</time>
</p>
)
}
function formatWindow(report: MarketIntegrityCheck) {
if (!report.window_start || !report.window_end) return "—"
return `${report.window_start} 至 ${report.window_end}`
}
function progressPercentage(checked: number, target: number) {
if (target <= 0) return 0
return Math.min(100, Math.max(0, (checked / target) * 100))
}
function acceptedCheckAsReport(
accepted: IntegrityCheckAccepted,
): MarketIntegrityCheck {
return {
check_id: accepted.check_id,
status: "running",
window_start: accepted.window_start,
window_end: accepted.window_end,
target_count: 0,
checked_count: 0,
issue_count: 0,
error_type: null,
error_message: null,
created_at: null,
finished_at: null,
page: 1,
page_size: defaultMarketIntegrityQuery.pageSize,
issues_total: 0,
issues: [],
}
}
function getErrorStatus(error: unknown) {
if (error instanceof ApiError) return error.status
if (typeof error === "object" && error !== null && "status" in error) {
const status = error.status
return typeof status === "number" ? status : null
}
return null
}
function getQueryErrorDescription(error: unknown) {
const status = getErrorStatus(error)
if (status === 503) {
return "检查结果存储暂时不可用,未将这次请求解释为市场数据损坏。"
}
return "检查结果暂时无法获取,可能是网络或服务异常。请稍后重试。"
}
function getTriggerErrorDescription(error: unknown) {
const status = getErrorStatus(error)
if (status === 422) {
return "当前没有可供比较的已完成市场数据窗口,请先完成一次市场数据同步。"
}
if (status === 503) {
return "检查服务暂时不可用,未发起新的检查。请稍后重试。"
}
return "检查请求可能没有到达服务,请确认网络后重试。"
}
+13 -1
View File
@@ -8,6 +8,7 @@ import {
type SelectionCategoryFilter,
} from "@/features/selection/api/selection.types"
import { SelectionResultsPage } from "@/features/selection/pages/selection-results-page"
import { SyncPage } from "@/features/sync/pages/sync-page"
const rootRoute = createRootRoute({
component: () => <Outlet />,
@@ -48,6 +49,12 @@ const selectionRoute = createRoute({
component: SelectionResultsPage,
})
const syncRoute = createRoute({
getParentRoute: () => workspaceRoute,
path: "/sync",
component: SyncPage,
})
const componentsRoute = createRoute({
getParentRoute: () => workspaceRoute,
path: "/components",
@@ -55,5 +62,10 @@ const componentsRoute = createRoute({
})
export const routeTree = rootRoute.addChildren([
workspaceRoute.addChildren([indexRoute, selectionRoute, componentsRoute]),
workspaceRoute.addChildren([
indexRoute,
selectionRoute,
syncRoute,
componentsRoute,
]),
])