diff --git a/.env.example b/.env.example index 76678bd..6b201f6 100644 --- a/.env.example +++ b/.env.example @@ -17,4 +17,10 @@ ZHIXING_POSTGRES_PASSWORD=zhixing # ZHIXING_DATABASE_URL=postgresql://zhixing-system:@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 diff --git a/docker-compose.dev.yml b/docker-compose.dev.yml index 600dfc7..6386ad0 100644 --- a/docker-compose.dev.yml +++ b/docker-compose.dev.yml @@ -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 diff --git a/docker-compose.prod.yml b/docker-compose.prod.yml index 8b9e01c..b61c934 100644 --- a/docker-compose.prod.yml +++ b/docker-compose.prod.yml @@ -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 diff --git a/docs/market-data-sync.md b/docs/market-data-sync.md index 5f93f31..7cb162b 100644 --- a/docs/market-data-sync.md +++ b/docs/market-data-sync.md @@ -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 使周末、节假日、重复触发和重叠触发保持安全: diff --git a/zhixing-server/migrations/versions/0003_market_integrity_checks.py b/zhixing-server/migrations/versions/0003_market_integrity_checks.py new file mode 100644 index 0000000..2d1f565 --- /dev/null +++ b/zhixing-server/migrations/versions/0003_market_integrity_checks.py @@ -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") diff --git a/zhixing-server/pyproject.toml b/zhixing-server/pyproject.toml index 894ebb2..024b63b 100644 --- a/zhixing-server/pyproject.toml +++ b/zhixing-server/pyproject.toml @@ -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", diff --git a/zhixing-server/src/zhixing_server/bootstrap/config.py b/zhixing-server/src/zhixing_server/bootstrap/config.py index 8889338..0455cc3 100644 --- a/zhixing-server/src/zhixing_server/bootstrap/config.py +++ b/zhixing-server/src/zhixing_server/bootstrap/config.py @@ -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 diff --git a/zhixing-server/src/zhixing_server/interfaces/http/router.py b/zhixing-server/src/zhixing_server/interfaces/http/router.py index ee642ba..0e2971c 100644 --- a/zhixing-server/src/zhixing_server/interfaces/http/router.py +++ b/zhixing-server/src/zhixing_server/interfaces/http/router.py @@ -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"] diff --git a/zhixing-server/src/zhixing_server/modules/market_data/application/__init__.py b/zhixing-server/src/zhixing_server/modules/market_data/application/__init__.py index 924b661..ea4fb38 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/application/__init__.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/application/__init__.py @@ -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", +] diff --git a/zhixing-server/src/zhixing_server/modules/market_data/application/integrity.py b/zhixing-server/src/zhixing_server/modules/market_data/application/integrity.py new file mode 100644 index 0000000..70e03fa --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/market_data/application/integrity.py @@ -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"] diff --git a/zhixing-server/src/zhixing_server/modules/market_data/application/sync.py b/zhixing-server/src/zhixing_server/modules/market_data/application/sync.py index 8d0f1f0..5d0d43b 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/application/sync.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/application/sync.py @@ -4,11 +4,12 @@ from __future__ import annotations import logging import time -from collections.abc import Iterable, Sequence +from collections.abc import Callable, Iterable, Sequence +from concurrent.futures import ThreadPoolExecutor, as_completed from dataclasses import dataclass, field from datetime import date, timedelta from decimal import Decimal -from typing import Literal +from typing import Literal, cast from ..domain.fingerprint import SnapshotChange, compare_snapshots from ..domain.models import Bar, Stock, SyncWindow @@ -40,6 +41,24 @@ class SyncFailure: message: str +@dataclass(frozen=True, slots=True) +class SyncItemOutcome: + """Immutable result returned by one synchronization item. + + A worker owns all side effects for its stock, then hands only this value + back to the coordinating thread. Keeping audit and aggregate mutation out + of the worker makes completion order irrelevant and prevents one failed + future from corrupting the batch counters. + """ + + item_kind: str + item_key: str + status: Literal["success", "failed"] + result: WriteResult = field(default_factory=WriteResult) + fingerprint: str | None = None + failure: SyncFailure | None = None + + @dataclass(frozen=True, slots=True) class SyncBatchSummary: """Stable output contract for CLI, cron, and later strategy callers.""" @@ -63,7 +82,9 @@ class SyncBatchSummary: if self.status == "failed": return 1 - return 0 if self.strategy_eligible else 2 + if self.status != "success" or not self.strategy_eligible: + return 2 + return 0 def as_dict(self) -> dict[str, object]: """Serialize the summary without credentials or raw vendor responses.""" @@ -107,36 +128,55 @@ class SyncMarketData: coverage_threshold: Decimal = Decimal("0.99"), lock_key: int = 7_380_521, today: date | None = None, + max_workers: int = 8, ) -> None: if not Decimal("0") <= coverage_threshold <= Decimal("1"): raise ValueError("coverage_threshold must be between 0 and 1") + if max_workers < 1: + raise ValueError("max_workers must be at least 1") self.source = source self.snapshots = snapshots self.repository = repository self.coverage_threshold = coverage_threshold self.lock_key = lock_key self.today = today or date.today() + self.max_workers = max_workers def execute(self, command: SyncMarketDataCommand | None = None) -> SyncBatchSummary: """Run one synchronization and retain successful items on partial failure.""" command = command or SyncMarketDataCommand() - with self.repository.advisory_lock(self.lock_key) as acquired: - if not acquired: - return SyncBatchSummary( - batch_id=None, - target_trade_date=None, - window=None, - status="failed", - target_count=0, - valid_count=0, - coverage=Decimal("0"), - strategy_eligible=False, - failures=( - SyncFailure("batch", "lock", "sync_locked", "another sync is running"), - ), - ) - return self._execute_locked(command) + try: + with self.repository.advisory_lock(self.lock_key) as acquired: + if not acquired: + return SyncBatchSummary( + batch_id=None, + target_trade_date=None, + window=None, + status="failed", + target_count=0, + valid_count=0, + coverage=Decimal("0"), + strategy_eligible=False, + failures=( + SyncFailure("batch", "lock", "sync_locked", "another sync is running"), + ), + ) + return self._execute_locked(command) + except Exception as exc: + # Connection/pool failures while acquiring the advisory lock must + # still produce the CLI's infrastructure-failure exit code. + return SyncBatchSummary( + batch_id=None, + target_trade_date=command.target_trade_date, + window=None, + status="failed", + target_count=0, + valid_count=0, + coverage=Decimal("0"), + strategy_eligible=False, + failures=(self._failure("batch", "lock", exc),), + ) def _execute_locked(self, command: SyncMarketDataCommand) -> SyncBatchSummary: started_at = time.monotonic() @@ -145,7 +185,20 @@ class SyncMarketData: command.mode, command.target_trade_date or "auto", ) - target_trade_date = self._resolve_target(command.target_trade_date) + try: + target_trade_date = self._resolve_target(command.target_trade_date) + except Exception as exc: + return SyncBatchSummary( + None, + command.target_trade_date, + None, + "failed", + 0, + 0, + Decimal("0"), + False, + failures=(self._failure("batch", "target", exc),), + ) window = SyncWindow.from_target(target_trade_date) logger.info( "market_data_sync_target target_trade_date=%s window_start=%s window_end=%s", @@ -153,7 +206,20 @@ class SyncMarketData: window.start, window.end, ) - all_stocks = filter_current_hs_a_stocks(self.source.fetch_stocks()) + try: + all_stocks = filter_current_hs_a_stocks(self.source.fetch_stocks()) + except Exception as exc: + return SyncBatchSummary( + None, + target_trade_date, + window, + "failed", + 0, + 0, + Decimal("0"), + False, + failures=(self._failure("stock", "universe", exc),), + ) if not all_stocks: return SyncBatchSummary( None, @@ -168,13 +234,44 @@ class SyncMarketData: SyncFailure("stock", "universe", "empty_universe", "no eligible stocks"), ), ) - batch_id = self.repository.create_batch( - target_trade_date, - window, - command.mode, - command.parent_batch_id, - len(all_stocks), - ) + try: + retry_items: set[tuple[str, str]] = ( + self._retry_items(command.parent_batch_id) if command.mode == "retry" else set() + ) + dates = self._dates_to_process(window, target_trade_date, command.mode, retry_items) + except Exception as exc: + return SyncBatchSummary( + None, + target_trade_date, + window, + "failed", + len(all_stocks), + 0, + Decimal("0"), + False, + failures=(self._failure("batch", "prepare", exc),), + ) + try: + batch_id = self.repository.create_batch( + target_trade_date, + window, + command.mode, + command.parent_batch_id, + len(all_stocks), + ) + except Exception as exc: + failure = self._failure("batch", "create", exc) + return SyncBatchSummary( + None, + target_trade_date, + window, + "failed", + len(all_stocks), + 0, + Decimal("0"), + False, + failures=(failure,), + ) logger.info( "market_data_sync_batch_created batch_id=%s target_count=%d", batch_id, @@ -182,12 +279,19 @@ class SyncMarketData: ) failures: list[SyncFailure] = [] totals = [0, 0, 0] + pending_outcomes: list[SyncItemOutcome] = [] + audit_failed = False stock_codes = {stock.ts_code for stock in all_stocks} - retry_items: set[tuple[str, str]] = ( - self._retry_items(command.parent_batch_id) if command.mode == "retry" else set() - ) - self._process_stock_master(batch_id, all_stocks, failures, totals) + stock_outcome = self._process_stock_master(all_stocks) + audit_failed = self._consume_outcome( + batch_id, + stock_outcome, + pending_outcomes, + failures, + totals, + audit_failed, + ) self._log_progress( batch_id=batch_id, stage="stock_master", @@ -199,7 +303,6 @@ class SyncMarketData: started_at=started_at, force=bool(failures), ) - dates = self._dates_to_process(window, target_trade_date, command.mode, retry_items) logger.info( "market_data_sync_stage_started batch_id=%s stage=daily_basic total=%d", batch_id, @@ -207,7 +310,19 @@ class SyncMarketData: ) for current, trade_date in enumerate(dates, start=1): failure_count = len(failures) - self._process_daily_basic(batch_id, trade_date, stock_codes, window, failures, totals) + daily_basic_outcome = self._process_daily_basic( + trade_date, + stock_codes, + window, + ) + audit_failed = self._consume_outcome( + batch_id, + daily_basic_outcome, + pending_outcomes, + failures, + totals, + audit_failed, + ) self._log_progress( batch_id=batch_id, stage="daily_basic", @@ -234,19 +349,55 @@ class SyncMarketData: batch_id, len(stocks_to_process), ) - for current, stock in enumerate(stocks_to_process, start=1): - failure_count = len(failures) - self._process_bar(batch_id, stock, window, failures, totals) - self._log_progress( - batch_id=batch_id, - stage="bar", - current=current, - total=len(stocks_to_process), - item_key=stock.ts_code, - totals=totals, - failures=failures, - started_at=started_at, - force=len(failures) > failure_count, + try: + with ThreadPoolExecutor( + max_workers=self.max_workers, + thread_name_prefix="market-data-bar", + ) as executor: + futures = { + executor.submit(self._process_bar, stock, window): stock.ts_code + for stock in stocks_to_process + } + for current, future in enumerate(as_completed(futures), start=1): + item_key = futures[future] + failure_count = len(failures) + try: + bar_outcome = future.result() + except Exception as exc: # pragma: no cover - defensive future boundary + bar_outcome = SyncItemOutcome( + "bar", + item_key, + "failed", + failure=self._failure("bar", item_key, exc), + ) + audit_failed = self._consume_outcome( + batch_id, + bar_outcome, + pending_outcomes, + failures, + totals, + audit_failed, + ) + self._log_progress( + batch_id=batch_id, + stage="bar", + current=current, + total=len(stocks_to_process), + item_key=bar_outcome.item_key, + totals=totals, + failures=failures, + started_at=started_at, + force=len(failures) > failure_count, + ) + except Exception as exc: + # Executor construction/submission is a batch-level failure. Any + # facts committed by already completed workers remain committed. + failure = self._failure("batch", "executor", exc) + failures.append(failure) + logger.exception( + "market_data_sync_executor_failed batch_id=%s error_type=%s", + batch_id, + failure.error_type, ) logger.info( "market_data_sync_stage_completed batch_id=%s stage=bar total=%d", @@ -254,32 +405,62 @@ class SyncMarketData: len(stocks_to_process), ) + if pending_outcomes: + audit_failure = self._record_outcomes(batch_id, pending_outcomes) + pending_outcomes.clear() + if audit_failure is not None and not audit_failed: + failures.append(audit_failure) + audit_failed = True + if not failures: try: self.repository.purge_before(window) self.snapshots.clean_daily_basic_before(window.start) - except (OSError, RuntimeError, TypeError, ValueError) as exc: + except Exception as exc: failure = self._failure("batch", "retention", exc) failures.append(failure) - self.repository.record_item( + audit_failure = self._record_outcomes( batch_id, - "batch", - "retention", - "failed", - WriteResult(), - error_type=failure.error_type, - error_message=failure.message, + [ + SyncItemOutcome( + "batch", + "retention", + "failed", + failure=failure, + ) + ], ) - valid_count = sum( - 1 - for stock in all_stocks - if self.repository.has_bar(stock.ts_code, target_trade_date) - and self.repository.has_daily_basic(stock.ts_code, target_trade_date) - ) + if audit_failure is not None and not audit_failed: + failures.append(audit_failure) + audit_failed = True + try: + count_valid_stocks = getattr(self.repository, "count_valid_stocks", None) + if callable(count_valid_stocks): + count_fn = cast(Callable[[date], int], count_valid_stocks) + valid_count = int(count_fn(target_trade_date)) + else: + # Compatibility fallback for older in-memory adapters. The + # PostgreSQL adapter always takes the set-based path above. + valid_count = sum( + 1 + for stock in all_stocks + if self.repository.has_bar(stock.ts_code, target_trade_date) + and self.repository.has_daily_basic(stock.ts_code, target_trade_date) + ) + except Exception as exc: + failure = self._failure("batch", "coverage", exc) + failures.append(failure) + valid_count = 0 coverage = Decimal(valid_count) / Decimal(len(all_stocks)) status = "success" if not failures else "partial_success" if valid_count else "failed" eligible = coverage >= self.coverage_threshold - self.repository.record_batch(batch_id, status, valid_count, coverage, eligible) + try: + self.repository.record_batch(batch_id, status, valid_count, coverage, eligible) + except Exception as exc: + failure = self._failure("batch", "record", exc) + failures.append(failure) + status = "failed" + eligible = False logger.info( "market_data_sync_finished batch_id=%s status=%s target_count=%d valid_count=%d " "coverage=%s failures=%d inserted=%d updated=%d unchanged=%d elapsed_seconds=%.1f", @@ -294,6 +475,12 @@ class SyncMarketData: totals[2], time.monotonic() - started_at, ) + ordered_failures = tuple( + sorted( + failures, + key=lambda failure: (failure.item_kind, failure.item_key, failure.error_type), + ) + ) return SyncBatchSummary( batch_id, target_trade_date, @@ -306,7 +493,7 @@ class SyncMarketData: totals[0], totals[1], totals[2], - tuple(failures), + ordered_failures, ) def _resolve_target(self, requested: date | None) -> date: @@ -346,57 +533,37 @@ class SyncMarketData: def _process_stock_master( self, - batch_id: str, stocks: Sequence[Stock], - failures: list[SyncFailure], - totals: list[int], - ) -> None: + ) -> SyncItemOutcome: staged = None try: staged = self.snapshots.stage_stocks(stocks) result = self.repository.upsert_stocks(stocks) self.snapshots.publish(staged) - totals[0] += result.inserted - totals[1] += result.updated - totals[2] += result.unchanged - self.repository.record_item( - batch_id, + return SyncItemOutcome( "stock", "current", "success", result, staged.fingerprint, ) - except (OSError, RuntimeError, TypeError, ValueError) as exc: + except Exception as exc: if staged is not None: self.snapshots.discard(staged) failure = self._failure("stock", "current", exc) - failures.append(failure) - logger.warning( - "market_data_sync_item_failed batch_id=%s stage=stock_master item=current " - "error_type=%s", - batch_id, - failure.error_type, - ) - self.repository.record_item( - batch_id, + return SyncItemOutcome( "stock", "current", "failed", - WriteResult(), - error_type=failure.error_type, - error_message=failure.message, + failure=failure, ) def _process_daily_basic( self, - batch_id: str, trade_date: date, stock_codes: set[str], window: SyncWindow, - failures: list[SyncFailure], - totals: list[int], - ) -> None: + ) -> SyncItemOutcome: staged = None key = trade_date.isoformat() try: @@ -410,50 +577,35 @@ class SyncMarketData: staged = self.snapshots.stage_daily_basic(trade_date, rows) result = self.repository.upsert_daily_basic(rows, window) self.snapshots.publish(staged) - self._add_counts(totals, result) - self.repository.record_item( - batch_id, + return SyncItemOutcome( "daily_basic", key, "success", result, staged.fingerprint, ) - except (OSError, RuntimeError, TypeError, ValueError) as exc: + except Exception as exc: if staged is not None: self.snapshots.discard(staged) failure = self._failure("daily_basic", key, exc) - failures.append(failure) - logger.warning( - "market_data_sync_item_failed batch_id=%s stage=daily_basic item=%s error_type=%s", - batch_id, - key, - failure.error_type, - ) - self.repository.record_item( - batch_id, + return SyncItemOutcome( "daily_basic", key, "failed", - WriteResult(), - error_type=failure.error_type, - error_message=failure.message, + failure=failure, ) def _process_bar( self, - batch_id: str, stock: Stock, window: SyncWindow, - failures: list[SyncFailure], - totals: list[int], - ) -> None: + ) -> SyncItemOutcome: staged = None try: rows = tuple(self.source.fetch_bars(stock.ts_code, window)) + staged = self.snapshots.stage_bars(stock.ts_code, rows) old_rows = self.snapshots.read_bars(stock.ts_code) comparison = compare_snapshots(old_rows, rows) - staged = self.snapshots.stage_bars(stock.ts_code, rows) if comparison.change is SnapshotChange.UNCHANGED: result = WriteResult(unchanged=len(rows)) else: @@ -468,34 +620,22 @@ class SyncMarketData: in {SnapshotChange.INITIAL, SnapshotChange.CHANGED}, ) self.snapshots.publish(staged) - self._add_counts(totals, result) - self.repository.record_item( - batch_id, + return SyncItemOutcome( "bar", stock.ts_code, "success", result, staged.fingerprint, ) - except (OSError, RuntimeError, TypeError, ValueError) as exc: + except Exception as exc: if staged is not None: self.snapshots.discard(staged) failure = self._failure("bar", stock.ts_code, exc) - failures.append(failure) - logger.warning( - "market_data_sync_item_failed batch_id=%s stage=bar item=%s error_type=%s", - batch_id, - stock.ts_code, - failure.error_type, - ) - self.repository.record_item( - batch_id, + return SyncItemOutcome( "bar", stock.ts_code, "failed", - WriteResult(), - error_type=failure.error_type, - error_message=failure.message, + failure=failure, ) @staticmethod @@ -504,6 +644,68 @@ class SyncMarketData: totals[1] += result.updated totals[2] += result.unchanged + def _consume_outcome( + self, + batch_id: str, + outcome: SyncItemOutcome, + pending_outcomes: list[SyncItemOutcome], + failures: list[SyncFailure], + totals: list[int], + audit_failed: bool, + ) -> bool: + """Aggregate one completed item and flush audit rows in bounded batches.""" + + if outcome.failure is not None: + failures.append(outcome.failure) + logger.warning( + "market_data_sync_item_failed batch_id=%s stage=%s item=%s error_type=%s", + batch_id, + outcome.item_kind, + outcome.item_key, + outcome.failure.error_type, + ) + else: + self._add_counts(totals, outcome.result) + pending_outcomes.append(outcome) + if len(pending_outcomes) < _PROGRESS_LOG_INTERVAL: + return audit_failed + audit_failure = self._record_outcomes(batch_id, pending_outcomes) + pending_outcomes.clear() + if audit_failure is not None and not audit_failed: + failures.append(audit_failure) + return True + return audit_failed + + def _record_outcomes( + self, + batch_id: str, + outcomes: Sequence[SyncItemOutcome], + ) -> SyncFailure | None: + """Persist a batch of audit outcomes, with a legacy adapter fallback.""" + + if not outcomes: + return None + try: + record_items = getattr(self.repository, "record_items", None) + if callable(record_items): + record_items(batch_id, tuple(outcomes)) + else: + for outcome in outcomes: + failure = outcome.failure + self.repository.record_item( + batch_id, + outcome.item_kind, + outcome.item_key, + outcome.status, + outcome.result, + outcome.fingerprint, + failure.error_type if failure else None, + failure.message if failure else None, + ) + except Exception as exc: + return self._failure("batch", "audit", exc) + return None + @staticmethod def _log_progress( *, diff --git a/zhixing-server/src/zhixing_server/modules/market_data/domain/__init__.py b/zhixing-server/src/zhixing_server/modules/market_data/domain/__init__.py index b405f83..f8fee77 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/domain/__init__.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/domain/__init__.py @@ -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", diff --git a/zhixing-server/src/zhixing_server/modules/market_data/domain/integrity.py b/zhixing-server/src/zhixing_server/modules/market_data/domain/integrity.py new file mode 100644 index 0000000..5ab64bc --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/market_data/domain/integrity.py @@ -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: ... diff --git a/zhixing-server/src/zhixing_server/modules/market_data/domain/models.py b/zhixing-server/src/zhixing_server/modules/market_data/domain/models.py index adb3be1..236c192 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/domain/models.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/domain/models.py @@ -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(".") diff --git a/zhixing-server/src/zhixing_server/modules/market_data/domain/ports.py b/zhixing-server/src/zhixing_server/modules/market_data/domain/ports.py index 0b2eb28..9fba8ca 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/domain/ports.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/domain/ports.py @@ -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: ... diff --git a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/csv_snapshot.py b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/csv_snapshot.py index 614a1ee..33cdc03 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/csv_snapshot.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/csv_snapshot.py @@ -55,12 +55,50 @@ STOCK_COLUMNS = ("ts_code", "name", "market", "exchange", "list_status", "list_d _DAILY_DECIMAL_FIELDS = DAILY_BASIC_COLUMNS[2:] +class SnapshotReadError(ValueError): + """A formal CSV could not be parsed without exposing file internals.""" + + issue_type = "parse_error" + + +class SnapshotDuplicateKeyError(SnapshotReadError): + """A formal CSV contains more than one row for a business key.""" + + issue_type = "duplicate_key" + + +class SnapshotWindowError(SnapshotReadError): + """A CSV path or row does not match its declared snapshot window.""" + + issue_type = "window_out_of_bounds" + + def _code_path(ts_code: str) -> str: if not _SAFE_CODE.fullmatch(ts_code): raise ValueError(f"unsupported stock code: {ts_code!r}") return ts_code +def _require_header(fieldnames: Iterable[str] | None, expected: tuple[str, ...]) -> None: + """Reject schema drift before a row can be mistaken for valid data.""" + + if tuple(fieldnames or ()) != expected: + raise SnapshotReadError("snapshot header does not match the market-data contract") + + +def _date_from_snapshot_path(path: Path) -> date: + """Parse ``daily-basic/YYYY/YYYYMMDD.csv`` path components.""" + + if path.parent.parent.name == "daily-basic" and len(path.parent.name) == 4: + text = path.stem + if len(text) == 8 and text.isdigit() and text[:4] == path.parent.name: + try: + return date.fromisoformat(f"{text[:4]}-{text[4:6]}-{text[6:]}") + except ValueError as exc: + raise SnapshotWindowError("daily-basic path contains an invalid date") from exc + raise SnapshotWindowError("daily-basic path does not match the formal layout") + + def _bar_row(row: Bar) -> dict[str, str]: return { "ts_code": row.ts_code, @@ -125,19 +163,139 @@ class CsvSnapshotStore: return self.root / "stock-basic" / "current.csv" + def list_bar_snapshot_files(self) -> tuple[Path, ...]: + """List formal bar files without opening or creating any file.""" + + root = self.root / "bars" + if not root.exists(): + return () + return tuple(sorted(root.glob("*.csv"), key=lambda path: path.name)) + + def list_daily_basic_snapshot_files(self) -> tuple[Path, ...]: + """List formal daily-basic files without loading their rows.""" + + root = self.root / "daily-basic" + if not root.exists(): + return () + return tuple(sorted(root.glob("*/*.csv"), key=lambda path: str(path))) + + def list_bar_codes(self, window: object | None = None) -> tuple[str, ...]: + """Return stock codes represented by formal bar paths. + + The optional window is accepted for the integrity reader protocol; the + path itself has no date, so row-level window validation happens while + the file is read. + """ + + del window + return tuple(path.stem for path in self.list_bar_snapshot_files()) + + def list_daily_basic_dates(self, window: object | None = None) -> tuple[date, ...]: + """Return parseable dates represented by formal daily-basic paths.""" + + del window + result: list[date] = [] + for path in self.list_daily_basic_snapshot_files(): + try: + result.append(_date_from_snapshot_path(path)) + except ValueError: + continue + return tuple(sorted(set(result))) + + def list_invalid_daily_basic_snapshot_files(self) -> tuple[tuple[str, str], ...]: + """Return malformed daily-basic paths without opening their contents. + + A path that cannot identify a valid ``YYYY/YYYYMMDD`` date is itself + an integrity discrepancy. The normal date listing intentionally + skips such paths so callers can still inspect every valid date; the + integrity application consumes this explicit error listing separately. + """ + + errors: list[tuple[str, str]] = [] + for path in self.list_daily_basic_snapshot_files(): + try: + _date_from_snapshot_path(path) + except SnapshotReadError as exc: + relative_path = path.relative_to(self.root).as_posix() + errors.append((relative_path, getattr(exc, "issue_type", "parse_error"))) + return tuple(errors) + def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None: """Read the current formal snapshot, returning ``None`` if absent.""" path = self.bars_path(ts_code) if not path.exists(): return None - with path.open(newline="", encoding="utf-8") as handle: - reader = csv.DictReader(handle) - if tuple(reader.fieldnames or ()) != BAR_COLUMNS: - raise ValueError(f"unexpected bar CSV header: {path}") - rows = tuple(Bar.from_mapping(row) for row in reader) + try: + with path.open(newline="", encoding="utf-8") as handle: + reader = csv.DictReader(handle) + _require_header(reader.fieldnames, BAR_COLUMNS) + rows = tuple(Bar.from_mapping(row) for row in reader) + except SnapshotReadError: + raise + except (OSError, csv.Error, TypeError, ValueError) as exc: + raise SnapshotReadError(f"bar snapshot cannot be parsed for {ts_code}") from exc + if not rows: + raise SnapshotReadError(f"bar snapshot is empty for {ts_code}") + keys = [(row.ts_code, row.trade_date) for row in rows] + if any(row.ts_code != ts_code for row in rows): + raise SnapshotReadError(f"bar snapshot contains an unexpected stock code: {ts_code}") + if len(keys) != len(set(keys)): + raise SnapshotDuplicateKeyError(f"bar snapshot contains duplicate keys: {ts_code}") return tuple(sorted(rows, key=lambda row: row.trade_date)) + def read_stocks(self) -> tuple[Stock, ...] | None: + """Read the formal current stock-master file without publishing.""" + + path = self.stock_basic_path + if not path.exists(): + return None + try: + with path.open(newline="", encoding="utf-8") as handle: + reader = csv.DictReader(handle) + _require_header(reader.fieldnames, STOCK_COLUMNS) + rows = tuple(Stock.from_mapping(row) for row in reader) + except SnapshotReadError: + raise + except (OSError, csv.Error, TypeError, ValueError) as exc: + raise SnapshotReadError("stock-master snapshot cannot be parsed") from exc + if not rows: + raise SnapshotReadError("stock-master snapshot is empty") + keys = [row.ts_code for row in rows] + if len(keys) != len(set(keys)): + raise SnapshotDuplicateKeyError("stock-master snapshot contains duplicate codes") + return tuple(sorted(rows, key=lambda row: row.ts_code)) + + def read_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...] | None: + """Read one formal daily-basic date snapshot without side effects.""" + + path = self.daily_basic_path(trade_date) + if not path.exists(): + return None + try: + with path.open(newline="", encoding="utf-8") as handle: + reader = csv.DictReader(handle) + _require_header(reader.fieldnames, DAILY_BASIC_COLUMNS) + rows = tuple(DailyBasic.from_mapping(row) for row in reader) + except SnapshotReadError: + raise + except (OSError, csv.Error, TypeError, ValueError) as exc: + raise SnapshotReadError( + f"daily-basic snapshot cannot be parsed for {trade_date}" + ) from exc + if not rows: + raise SnapshotReadError(f"daily-basic snapshot is empty for {trade_date}") + keys = [(row.ts_code, row.trade_date) for row in rows] + if any(row.trade_date != trade_date for row in rows): + raise SnapshotWindowError( + f"daily-basic snapshot contains an unexpected date: {trade_date}" + ) + if len(keys) != len(set(keys)): + raise SnapshotDuplicateKeyError( + f"daily-basic snapshot contains duplicate keys: {trade_date}" + ) + return tuple(sorted(rows, key=lambda row: row.ts_code)) + def stage_bars(self, ts_code: str, rows: Iterable[Bar]) -> StagedSnapshot: """Validate and write a temporary bar snapshot.""" @@ -146,7 +304,15 @@ class CsvSnapshotStore: raise ValueError("bar snapshot must be non-empty and belong to one stock") final_path = self.bars_path(ts_code) temp_path = self._write_csv(final_path, BAR_COLUMNS, (_bar_row(row) for row in normalized)) - return StagedSnapshot(temp_path, final_path, snapshot_fingerprint(normalized)) + try: + fingerprint = snapshot_fingerprint(normalized) + except BaseException: + # The CSV has already been fully staged, so validation failures + # (for example an invalid OHLC relation) must not leave a hidden + # temporary file behind for a later run to mistake for a snapshot. + temp_path.unlink(missing_ok=True) + raise + return StagedSnapshot(temp_path, final_path, fingerprint) def stage_daily_basic(self, trade_date: date, rows: Iterable[DailyBasic]) -> StagedSnapshot: """Validate and stage one daily-basic date snapshot.""" diff --git a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/integrity.py b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/integrity.py new file mode 100644 index 0000000..7aaf0fc --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/integrity.py @@ -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", +] diff --git a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/postgres.py b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/postgres.py index 492d5c0..9ef3f0a 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/postgres.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/postgres.py @@ -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", +] diff --git a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/schema.py b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/schema.py index b5405f2..c2aab64 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/schema.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/schema.py @@ -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) diff --git a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py index 6a055f3..9ccf80e 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py @@ -1,11 +1,12 @@ -"""Tushare source adapter and current-universe filtering.""" +"""Tushare source adapter, request coordination, and current-universe filtering.""" from __future__ import annotations import logging import random +import threading import time -from collections.abc import Callable, Iterable, Mapping +from collections.abc import Callable, Iterable, Mapping, Sequence from datetime import date from typing import cast @@ -14,18 +15,226 @@ from ..domain.rules import filter_current_hs_a_stocks logger = logging.getLogger(__name__) +DEFAULT_RATE_LIMIT_COOLDOWNS = (60.0, 120.0, 180.0) +_RATE_LIMIT_MESSAGES = ( + "访问频繁", + "请稍后", + "超过频率", + "频率限制", + "too many requests", + "rate limit", + "rate_limit", + "http 429", + "status code: 429", + "429", + "http 403", + "status code: 403", + "403", +) + class TushareSourceError(RuntimeError): """A vendor request failed after the configured retry budget.""" -class TushareAdapter: - """Translate Tushare SDK responses into domain records. +class RequestCoordinator: + """Coordinate retry and shared rate-limit cooling for one token client. - The SDK is kept behind this adapter so ordinary domain/application tests - can inject a tiny fake client and never need a network token. + Normal requests are deliberately not serialized. Only a provider rate + limit creates a shared cooldown, so independent worker calls can proceed + concurrently during ordinary traffic. ``clock`` and ``wait_fn`` are + injectable to make long cooldown behavior deterministic in unit tests. """ + def __init__( + self, + *, + max_retries: int = 3, + backoff_seconds: float = 1.0, + cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS, + random_fn: Callable[[], float] = random.random, + clock: Callable[[], float] = time.monotonic, + wait_fn: Callable[[float], None] = time.sleep, + sleep_fn: Callable[[float], None] | None = None, + ) -> None: + cooldowns = tuple(float(value) for value in cooldown_seconds) + if not cooldowns or any(value < 0 for value in cooldowns): + raise ValueError("cooldown_seconds must contain non-negative values") + self.max_retries = max(0, max_retries) + self.backoff_seconds = max(0.0, backoff_seconds) + self.cooldown_seconds = cooldowns + self.random_fn = random_fn + self.clock = clock + self.wait_fn = wait_fn + self.sleep_fn = sleep_fn or wait_fn + self._condition = threading.Condition() + self._cooldown_until = 0.0 + self._rate_limit_count = 0 + + @property + def cooldown_until(self) -> float: + """Return the current monotonic cooldown deadline.""" + + with self._condition: + return self._cooldown_until + + def call(self, method_name: str, request: Callable[[], object]) -> object: + """Execute one provider request with bounded, shared retry behavior.""" + + last_error: BaseException | None = None + for attempt in range(self.max_retries + 1): + self._wait_for_cooldown(method_name) + try: + result = request() + except Exception as exc: + last_error = exc + if self.is_rate_limited(exc): + cooldown = self._set_rate_limit_cooldown() + logger.warning( + "tushare_rate_limit method=%s attempt=%d max_attempts=%d " + "cooldown_seconds=%.1f", + method_name, + attempt + 1, + self.max_retries + 1, + cooldown, + ) + if attempt < self.max_retries: + continue + break + if not self._is_retryable(exc): + raise + if attempt == self.max_retries: + break + delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn()) + logger.warning( + "tushare_request_retry method=%s attempt=%d max_attempts=%d " + "backoff_seconds=%.1f", + method_name, + attempt + 1, + self.max_retries + 1, + delay, + ) + self.sleep_fn(delay) + else: + self._clear_rate_limit_after_success() + return result + logger.error( + "tushare_request_failed method=%s attempts=%d", + method_name, + self.max_retries + 1, + ) + raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error + + def request(self, method_name: str, operation: Callable[[], object]) -> object: + """Alias for ``call`` for adapters that model requests as a port.""" + + return self.call(method_name, operation) + + def _wait_for_cooldown(self, method_name: str) -> None: + while True: + with self._condition: + delay = self._cooldown_until - self.clock() + if delay <= 0: + return + logger.info( + "tushare_rate_limit_wait method=%s wait_seconds=%.1f", + method_name, + delay, + ) + # A single injected wait hook makes fake-clock tests independent + # from wall time. After waiting, re-check because another worker + # may have extended the shared deadline. + self.wait_fn(delay) + + def _set_rate_limit_cooldown(self) -> float: + with self._condition: + self._rate_limit_count += 1 + index = min(self._rate_limit_count - 1, len(self.cooldown_seconds) - 1) + duration = self.cooldown_seconds[index] + self._cooldown_until = max(self._cooldown_until, self.clock() + duration) + self._condition.notify_all() + return duration + + def _clear_rate_limit_after_success(self) -> None: + with self._condition: + # A request that was already in flight when another worker hit a + # limit may succeed during the shared cooldown. Do not erase the + # escalation history until the cooldown has actually elapsed. + if self.clock() >= self._cooldown_until: + self._rate_limit_count = 0 + + @staticmethod + def is_rate_limited(error: BaseException) -> bool: + """Classify stable provider rate-limit signals without logging details.""" + + for attribute in ("status_code", "status", "code"): + value = getattr(error, attribute, None) + if str(value).strip() in {"403", "429"}: + return True + message = str(error).casefold() + return any(marker.casefold() in message for marker in _RATE_LIMIT_MESSAGES) + + @staticmethod + def _is_retryable(error: BaseException) -> bool: + return isinstance(error, (OSError, RuntimeError, TimeoutError)) + + +# The longer name is useful to callers that want to make the infrastructure +# boundary explicit, while the short name remains convenient in unit tests. +TushareRequestCoordinator = RequestCoordinator + + +class CoordinatedTushareClient: + """Proxy that routes qfq's real ``daily`` and ``adj_factor`` calls.""" + + def __init__(self, client: object, coordinator: RequestCoordinator) -> None: + self._client = client + self._coordinator = coordinator + + def __getattr__(self, name: str) -> object: + method = getattr(self._client, name) + if name not in {"daily", "adj_factor"} or not callable(method): + return method + + def coordinated(*args: object, **kwargs: object) -> object: + return self._coordinator.call( + name, + lambda: method(*args, **kwargs), + ) + + return coordinated + + +def _bind_coordinated_methods(client: object, coordinator: RequestCoordinator) -> object: + """Bind wrappers on a mutable SDK client while preserving ``api is client``. + + Tushare's ``pro_bar`` accepts the token client through its ``api`` + parameter. Binding the two methods in place keeps that identity contract + (and avoids a proxy-visible behavior change); objects that do not allow + attributes fall back to the proxy. + """ + + def build_wrapper(name: str, method: Callable[..., object]) -> Callable[..., object]: + def coordinated(*args: object, **kwargs: object) -> object: + return coordinator.call(name, lambda: method(*args, **kwargs)) + + return coordinated + + for name in ("daily", "adj_factor"): + method = getattr(client, name, None) + if not callable(method): + continue + + try: + setattr(client, name, build_wrapper(name, method)) + except (AttributeError, TypeError): + return CoordinatedTushareClient(client, coordinator) + return client + + +class TushareAdapter: + """Translate Tushare SDK responses into domain records.""" + def __init__( self, client: object, @@ -36,14 +245,32 @@ class TushareAdapter: request_interval_seconds: float = 0.2, random_fn: Callable[[], float] = random.random, sleep_fn: Callable[[float], None] = time.sleep, + request_coordinator: RequestCoordinator | None = None, + cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS, + coordinated_client: object | None = None, + pro_bar_coordinated: bool = False, ) -> None: self.client = client - self.pro_bar = pro_bar self.max_retries = max(0, max_retries) self.backoff_seconds = max(0.0, backoff_seconds) self.request_interval_seconds = max(0.0, request_interval_seconds) self.random_fn = random_fn self.sleep_fn = sleep_fn + self.request_coordinator = request_coordinator or RequestCoordinator( + max_retries=self.max_retries, + backoff_seconds=self.backoff_seconds, + cooldown_seconds=cooldown_seconds, + random_fn=random_fn, + wait_fn=sleep_fn, + sleep_fn=sleep_fn, + ) + self.coordinated_client = ( + coordinated_client + if coordinated_client is not None + else _bind_coordinated_methods(client, self.request_coordinator) + ) + self.pro_bar = pro_bar + self.pro_bar_coordinated = pro_bar_coordinated @classmethod def from_token( @@ -53,6 +280,7 @@ class TushareAdapter: max_retries: int = 3, backoff_seconds: float = 1.0, request_interval_seconds: float = 0.2, + cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS, ) -> TushareAdapter: """Create a production adapter from a token without exposing it.""" @@ -61,13 +289,22 @@ class TushareAdapter: import tushare as ts # pyright: ignore[reportMissingTypeStubs] client = cast(object, ts.pro_api(token)) + coordinator = RequestCoordinator( + max_retries=max_retries, + backoff_seconds=backoff_seconds, + cooldown_seconds=cooldown_seconds, + ) + coordinated_client = _bind_coordinated_methods(client, coordinator) pro_bar = cast( Callable[..., object], ts.pro_bar, # pyright: ignore[reportUnknownMemberType] ) def pro_bar_with_client(**kwargs: object) -> object: - return pro_bar(api=client, **kwargs) + # qfq's SDK helper otherwise performs hidden retries outside the + # shared coordinator and can amplify a provider rate limit. + kwargs["retry_count"] = 1 + return pro_bar(api=coordinated_client, **kwargs) return cls( client, @@ -75,6 +312,9 @@ class TushareAdapter: max_retries=max_retries, backoff_seconds=backoff_seconds, request_interval_seconds=request_interval_seconds, + request_coordinator=coordinator, + coordinated_client=coordinated_client, + pro_bar_coordinated=True, ) def fetch_stocks(self) -> tuple[Stock, ...]: @@ -137,43 +377,29 @@ class TushareAdapter: return tuple(sorted(metrics, key=lambda row: row.ts_code)) def _records(self, method_name: str, **kwargs: object) -> tuple[Mapping[str, object], ...]: - """Call an SDK method with bounded retry and normalize its tabular output.""" + """Call the SDK through the coordinator and normalize tabular output.""" def request() -> object: if method_name == "pro_bar": if self.pro_bar is not None: return self.pro_bar(**kwargs) - method = getattr(self.client, method_name, None) + method = getattr(self.coordinated_client, method_name, None) else: - method = getattr(self.client, method_name, None) + method = getattr(self.coordinated_client, method_name, None) if not callable(method): raise TypeError(f"Tushare client has no callable {method_name}") return method(**kwargs) - last_error: BaseException | None = None - for attempt in range(self.max_retries + 1): - try: - result = request() - self.sleep_fn(self.request_interval_seconds) - return self._as_records(result) - except (OSError, RuntimeError, TimeoutError) as exc: - last_error = exc - if attempt == self.max_retries: - break - delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn()) - logger.warning( - "tushare_request_retry method=%s attempt=%d max_attempts=%d", - method_name, - attempt + 1, - self.max_retries + 1, - ) - self.sleep_fn(delay) - logger.error( - "tushare_request_failed method=%s attempts=%d", - method_name, - self.max_retries + 1, + # ``pro_bar`` itself is a qfq composition helper; its nested daily and + # adj_factor methods are bound to the coordinator. Wrapping the helper + # as a second retry layer would hide the useful error classification. + result = ( + request() + if method_name == "pro_bar" and self.pro_bar_coordinated + else self.request_coordinator.call(method_name, request) ) - raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error + self.sleep_fn(self.request_interval_seconds) + return self._as_records(result) @staticmethod def _as_records(result: object) -> tuple[Mapping[str, object], ...]: diff --git a/zhixing-server/src/zhixing_server/modules/market_data/presentation/cli.py b/zhixing-server/src/zhixing_server/modules/market_data/presentation/cli.py index 4e87ce5..cec5b77 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/presentation/cli.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/presentation/cli.py @@ -70,14 +70,24 @@ def main(argv: Sequence[str] | None = None) -> int: backoff_seconds=settings.market_data_retry_backoff_seconds, request_interval_seconds=settings.market_data_request_interval_seconds, ) - use_case = SyncMarketData( - source, - CsvSnapshotStore(settings.market_data_csv_root), - PostgresMarketDataRepository(settings.database_url), - coverage_threshold=settings.market_data_coverage_threshold, - lock_key=settings.market_data_advisory_lock_key, + repository = PostgresMarketDataRepository( + settings.database_url, + # Keep one control/advisory connection and one main-thread connection + # available in addition to the worker connections. + max_connections=settings.market_data_max_workers + 2, ) - summary = use_case.execute(command) + try: + use_case = SyncMarketData( + source, + CsvSnapshotStore(settings.market_data_csv_root), + repository, + coverage_threshold=settings.market_data_coverage_threshold, + lock_key=settings.market_data_advisory_lock_key, + max_workers=settings.market_data_max_workers, + ) + summary = use_case.execute(command) + finally: + repository.close() print(json.dumps(summary.as_dict(), ensure_ascii=False, sort_keys=True)) return summary.exit_code diff --git a/zhixing-server/src/zhixing_server/modules/market_data/presentation/home.py b/zhixing-server/src/zhixing_server/modules/market_data/presentation/home.py index 871f0f2..80e0ab2 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/presentation/home.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/presentation/home.py @@ -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) diff --git a/zhixing-server/tests/integration/test_market_data_migration.py b/zhixing-server/tests/integration/test_market_data_migration.py index 52e9c11..12b93fa 100644 --- a/zhixing-server/tests/integration/test_market_data_migration.py +++ b/zhixing-server/tests/integration/test_market_data_migration.py @@ -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", diff --git a/zhixing-server/tests/integration/test_market_data_repository_pool.py b/zhixing-server/tests/integration/test_market_data_repository_pool.py new file mode 100644 index 0000000..978d969 --- /dev/null +++ b/zhixing-server/tests/integration/test_market_data_repository_pool.py @@ -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) diff --git a/zhixing-server/tests/unit/market_data/test_csv_snapshot.py b/zhixing-server/tests/unit/market_data/test_csv_snapshot.py index 2b5ab91..66d7f96 100644 --- a/zhixing-server/tests/unit/market_data/test_csv_snapshot.py +++ b/zhixing-server/tests/unit/market_data/test_csv_snapshot.py @@ -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"), + ) diff --git a/zhixing-server/tests/unit/market_data/test_domain.py b/zhixing-server/tests/unit/market_data/test_domain.py index f9e41d1..a9a0742 100644 --- a/zhixing-server/tests/unit/market_data/test_domain.py +++ b/zhixing-server/tests/unit/market_data/test_domain.py @@ -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"), diff --git a/zhixing-server/tests/unit/market_data/test_integrity.py b/zhixing-server/tests/unit/market_data/test_integrity.py new file mode 100644 index 0000000..1b0cff9 --- /dev/null +++ b/zhixing-server/tests/unit/market_data/test_integrity.py @@ -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" diff --git a/zhixing-server/tests/unit/market_data/test_sync_concurrency.py b/zhixing-server/tests/unit/market_data/test_sync_concurrency.py new file mode 100644 index 0000000..cf73131 --- /dev/null +++ b/zhixing-server/tests/unit/market_data/test_sync_concurrency.py @@ -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 diff --git a/zhixing-server/tests/unit/market_data/test_tushare.py b/zhixing-server/tests/unit/market_data/test_tushare.py index 5bff483..04d74fd 100644 --- a/zhixing-server/tests/unit/market_data/test_tushare.py +++ b/zhixing-server/tests/unit/market_data/test_tushare.py @@ -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 diff --git a/zhixing-server/tests/unit/test_config.py b/zhixing-server/tests/unit/test_config.py new file mode 100644 index 0000000..9a02e59 --- /dev/null +++ b/zhixing-server/tests/unit/test_config.py @@ -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) diff --git a/zhixing-server/uv.lock b/zhixing-server/uv.lock index ec0d629..960515a 100644 --- a/zhixing-server/uv.lock +++ b/zhixing-server/uv.lock @@ -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" }, diff --git a/zhixing-web/src/app/layout/app-layout.tsx b/zhixing-web/src/app/layout/app-layout.tsx index a4bf0c8..71f7d68 100644 --- a/zhixing-web/src/app/layout/app-layout.tsx +++ b/zhixing-web/src/app/layout/app-layout.tsx @@ -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({ 更多功能 - 行情数据与同步任务尚未上线,暂无更多可用入口。 + 行情数据尚未上线,暂无更多可用入口。 diff --git a/zhixing-web/src/app/layout/navigation.test.ts b/zhixing-web/src/app/layout/navigation.test.ts new file mode 100644 index 0000000..812e50b --- /dev/null +++ b/zhixing-web/src/app/layout/navigation.test.ts @@ -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: "市场数据完整性检查", + }) + }) +}) diff --git a/zhixing-web/src/app/layout/navigation.ts b/zhixing-web/src/app/layout/navigation.ts index 8528827..065f1cd 100644 --- a/zhixing-web/src/app/layout/navigation.ts +++ b/zhixing-web/src/app/layout/navigation.ts @@ -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 = { breadcrumb: "研究工作台", title: "知行 B1 执行结果", }, + "/sync": { + id: "sync", + breadcrumb: "研究工作台", + title: "市场数据完整性检查", + }, "/components": { id: "components", breadcrumb: "研究工作台", diff --git a/zhixing-web/src/features/sync/api/sync.api.test.ts b/zhixing-web/src/features/sync/api/sync.api.test.ts new file mode 100644 index 0000000..e4e85ed --- /dev/null +++ b/zhixing-web/src/features/sync/api/sync.api.test.ts @@ -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") + }) +}) diff --git a/zhixing-web/src/features/sync/api/sync.api.ts b/zhixing-web/src/features/sync/api/sync.api.ts new file mode 100644 index 0000000..a89b9c1 --- /dev/null +++ b/zhixing-web/src/features/sync/api/sync.api.ts @@ -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(integrityChecksPath, { + method: "POST", + signal, + }) +} + +export function getLatestIntegrityCheck( + query: MarketIntegrityQuery = defaultMarketIntegrityQuery, + signal?: AbortSignal, +) { + return requestJson( + `${integrityChecksPath}/latest?${buildQuery(query)}`, + { signal }, + ) +} + +export function getIntegrityCheck( + checkId: string, + query: MarketIntegrityQuery = defaultMarketIntegrityQuery, + signal?: AbortSignal, +) { + return requestJson( + `${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 diff --git a/zhixing-web/src/features/sync/api/sync.query.test.ts b/zhixing-web/src/features/sync/api/sync.query.test.ts new file mode 100644 index 0000000..0b1baaa --- /dev/null +++ b/zhixing-web/src/features/sync/api/sync.query.test.ts @@ -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"], + }) + }) +}) diff --git a/zhixing-web/src/features/sync/api/sync.query.ts b/zhixing-web/src/features/sync/api/sync.query.ts new file mode 100644 index 0000000..fb538f1 --- /dev/null +++ b/zhixing-web/src/features/sync/api/sync.query.ts @@ -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 +export type MarketIntegrityAccepted = IntegrityCheckAccepted +export type MarketIntegrityData = MarketIntegrityCheck diff --git a/zhixing-web/src/features/sync/api/sync.types.ts b/zhixing-web/src/features/sync/api/sync.types.ts new file mode 100644 index 0000000..e90b5ae --- /dev/null +++ b/zhixing-web/src/features/sync/api/sync.types.ts @@ -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, +} diff --git a/zhixing-web/src/features/sync/pages/sync-page.test.tsx b/zhixing-web/src/features/sync/pages/sync-page.test.tsx new file mode 100644 index 0000000..266434d --- /dev/null +++ b/zhixing-web/src/features/sync/pages/sync-page.test.tsx @@ -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 { + 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() + + 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() + + 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() + + 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() + + await waitFor(() => { + expect(refetchLatest).toHaveBeenCalledOnce() + }) + }) + + it("renders a passed report without an issue table", () => { + configureHooks(buildReport("passed")) + + render() + + 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() + + 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() + + 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() + + 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() + + expect(screen.getByText(/已刷新并接管该检查/)).toBeInTheDocument() + expect(screen.getByText("完整性检查进行中")).toBeInTheDocument() + }) +}) diff --git a/zhixing-web/src/features/sync/pages/sync-page.tsx b/zhixing-web/src/features/sync/pages/sync-page.tsx new file mode 100644 index 0000000..0340fec --- /dev/null +++ b/zhixing-web/src/features/sync/pages/sync-page.tsx @@ -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(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(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 ( + +
+
+

+ 只读数据校验 +

+

+ 市场数据完整性检查 +

+

+ 检查 PostgreSQL 与 CSV 快照中的市场数据是否一致,帮助定位数据问题。 +

+
+ + + + {triggerStatus === 409 ? : null} + + {triggerError ? ( + + ) : null} + + {!queryHasData && queryPending ? : null} + + {!queryPending && queryError ? ( + + ) : null} + + {report?.status === "no_data" ? ( + + ) : null} + + {report?.status === "running" ? : null} + + {report?.status === "passed" ? ( + + ) : null} + + {report?.status === "issues_found" ? ( + { + setPage(1) + setPageSize(nextPageSize) + }} + report={report} + selectedPage={page} + selectedPageSize={pageSize} + isPending={trigger.isPending} + onStart={startCheck} + /> + ) : null} + + {report?.status === "failed" ? ( + + ) : null} +
+
+ ) +} + +function ReadOnlyNotice() { + return ( + + + + + ) +} + +function ConflictNotice() { + return ( +
+
+ ) +} + +function LoadingState() { + return ( + + +
+
+ + +
+ + + ) +} + +function NoDataState({ + isPending, + onStart, +}: { + isPending: boolean + onStart: () => void +}) { + return ( + + + + + + 当前没有可查看的完整性检查。发起后,系统会比较最近完成同步窗口中的 + PostgreSQL 与 CSV 数据。 + + + + +

+ 检查只读数据并生成报告,不会访问 Tushare,也不会自动修复。 +

+
+
+ ) +} + +function RunningState({ report }: { report: MarketIntegrityCheck }) { + const percentage = progressPercentage( + report.checked_count, + report.target_count, + ) + return ( + + +
+
+ + + + 检查完成后会自动停止刷新;运行期间不会修改 PostgreSQL 或 CSV + 数据。 + +
+ 运行中 +
+
+ + + 检查进度 + + {() => `${report.checked_count}/${report.target_count}`} + + +
+ + + +
+
+ + +
+
+
+ ) +} + +function PassedState({ + isPending, + onStart, + report, +}: { + isPending: boolean + onStart: () => void + report: MarketIntegrityCheck +}) { + return ( + + +
+
+ + + + 未发现 PostgreSQL/CSV 不一致,本次检查只报告结果,没有修改数据。 + +
+ + 已通过 + +
+
+ +
+ 未发现 PostgreSQL/CSV 不一致。 +
+
+ + + +
+
+
+ + +
+ +
+
+
+ ) +} + +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 ( + + +
+
+ + + + 检查完成并发现 {report.issue_count}{" "} + 个问题。以下内容仅供定位,系统不会自动修复。 + +
+ + {report.issue_count} 个问题 + +
+
+ +
+ + + +
+
+ 安全说明:问题报告只展示对象和差异说明,不会执行修复或重新同步。 +
+ + +
+
+ + +
+ +
+
+
+ ) +} + +function IssueTable({ report }: { report: MarketIntegrityCheck }) { + if (report.issues.length === 0) { + return ( +
+ 当前页没有问题记录。 +
+ ) + } + + return ( +
+ + + + + + + + + + + + {report.issues.map((issue) => ( + + + + + + + ))} + +
完整性检查问题报告
+ 对象类型 + + 对象 key + + 问题类型 + + 安全说明 +
{issue.item_kind} + {issue.item_key} + {issue.issue_type} + {issue.message} +
+
+ ) +} + +function FailedState({ + isPending, + onStart, + report, +}: { + isPending: boolean + onStart: () => void + report: MarketIntegrityCheck +}) { + return ( + + +
+
+ + + + 本次检查遇到安全错误,不能据此判断市场数据已经损坏。 + +
+ 失败 +
+
+ +
+

未判断市场数据状态

+

+ {report.error_type + ? `错误类型:${report.error_type}` + : "检查服务暂时无法完成比较。"} +

+ {report.error_message ? ( +

{report.error_message}

+ ) : null} +
+
+ + + +
+
+
+ + +
+ +
+
+
+ ) +} + +function RequestError({ + description, + onRetry, + title, +}: { + description: string + onRetry: () => void + title: string +}) { + return ( + + + + + {description} + + + + + + ) +} + +function Metric({ label, value }: { label: string; value: string }) { + return ( +
+

{label}

+

{value}

+
+ ) +} + +function Timestamp({ label, value }: { label: string; value: string | null }) { + return ( +

+ {label}: + +

+ ) +} + +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 "检查请求可能没有到达服务,请确认网络后重试。" +} diff --git a/zhixing-web/src/routes/route-tree.tsx b/zhixing-web/src/routes/route-tree.tsx index 1718c26..fbcdc0b 100644 --- a/zhixing-web/src/routes/route-tree.tsx +++ b/zhixing-web/src/routes/route-tree.tsx @@ -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: () => , @@ -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, + ]), ])