feat(market-data): 优化同步并增加完整性检查
This commit is contained in:
@@ -17,4 +17,10 @@ ZHIXING_POSTGRES_PASSWORD=zhixing
|
|||||||
# ZHIXING_DATABASE_URL=postgresql://zhixing-system:<url-encoded-password>@postgresql:5432/zhixing-system?sslmode=disable
|
# ZHIXING_DATABASE_URL=postgresql://zhixing-system:<url-encoded-password>@postgresql:5432/zhixing-system?sslmode=disable
|
||||||
ZHIXING_TUSHARE_TOKEN=
|
ZHIXING_TUSHARE_TOKEN=
|
||||||
ZHIXING_MARKET_DATA_CSV_ROOT=/app/data/market-data
|
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
|
API_UPSTREAM=http://server:8000
|
||||||
|
|||||||
@@ -30,6 +30,12 @@ services:
|
|||||||
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:-postgresql://zhixing:zhixing@postgres:5432/zhixing}
|
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:-postgresql://zhixing:zhixing@postgres:5432/zhixing}
|
||||||
ZHIXING_LOG_LEVEL: ${ZHIXING_LOG_LEVEL:-INFO}
|
ZHIXING_LOG_LEVEL: ${ZHIXING_LOG_LEVEL:-INFO}
|
||||||
ZHIXING_MARKET_DATA_CSV_ROOT: /app/data/market-data
|
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:-}
|
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||||
init: true
|
init: true
|
||||||
ports:
|
ports:
|
||||||
@@ -92,6 +98,12 @@ services:
|
|||||||
environment:
|
environment:
|
||||||
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:-postgresql://zhixing:zhixing@postgres:5432/zhixing}
|
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:-postgresql://zhixing:zhixing@postgres:5432/zhixing}
|
||||||
ZHIXING_MARKET_DATA_CSV_ROOT: /app/data/market-data
|
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:-}
|
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||||
volumes:
|
volumes:
|
||||||
- ./zhixing-server:/app
|
- ./zhixing-server:/app
|
||||||
|
|||||||
@@ -12,6 +12,12 @@ services:
|
|||||||
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:?Set ZHIXING_DATABASE_URL to the 1Panel PostgreSQL URL}
|
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:?Set ZHIXING_DATABASE_URL to the 1Panel PostgreSQL URL}
|
||||||
ZHIXING_LOG_LEVEL: ${ZHIXING_LOG_LEVEL:-INFO}
|
ZHIXING_LOG_LEVEL: ${ZHIXING_LOG_LEVEL:-INFO}
|
||||||
ZHIXING_MARKET_DATA_CSV_ROOT: /app/data/market-data
|
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:-}
|
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||||
init: true
|
init: true
|
||||||
expose:
|
expose:
|
||||||
@@ -91,6 +97,12 @@ services:
|
|||||||
TZ: Asia/Shanghai
|
TZ: Asia/Shanghai
|
||||||
ZHIXING_DATABASE_URL: ${ZHIXING_DATABASE_URL:?Set ZHIXING_DATABASE_URL to the 1Panel PostgreSQL URL}
|
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_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:-}
|
ZHIXING_TUSHARE_TOKEN: ${ZHIXING_TUSHARE_TOKEN:-}
|
||||||
volumes:
|
volumes:
|
||||||
- market-data:/app/data/market-data
|
- market-data:/app/data/market-data
|
||||||
|
|||||||
@@ -62,6 +62,12 @@ docker compose -f docker-compose.prod.yml --profile jobs run --rm market-sync --
|
|||||||
|
|
||||||
`success` 返回 0;覆盖率低于配置阈值或出现部分失败返回 2;没有可用成功结果或基础设施失败返回 1。CLI 输出不包含 Tushare token 或数据库密码。
|
`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 与重试
|
||||||
|
|
||||||
宿主机 cron 只负责启动临时容器,不写入容器内部的 crontab。下面的示例每天工作日 18:00 触发;交易日历、唯一约束和 PostgreSQL advisory lock 使周末、节假日、重复触发和重叠触发保持安全:
|
宿主机 cron 只负责启动临时容器,不写入容器内部的 crontab。下面的示例每天工作日 18:00 触发;交易日历、唯一约束和 PostgreSQL advisory lock 使周末、节假日、重复触发和重叠触发保持安全:
|
||||||
|
|||||||
@@ -0,0 +1,90 @@
|
|||||||
|
"""Create independent market-data integrity check and issue tables."""
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "0003_market_integrity_checks"
|
||||||
|
down_revision: str | None = "0002_selection_results"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
"""Create only check metadata; market facts remain untouched."""
|
||||||
|
|
||||||
|
op.create_table(
|
||||||
|
"market_integrity_check",
|
||||||
|
sa.Column("id", sa.String(36), primary_key=True),
|
||||||
|
sa.Column("status", sa.String(24), nullable=False),
|
||||||
|
sa.Column("window_start", sa.Date()),
|
||||||
|
sa.Column("window_end", sa.Date()),
|
||||||
|
sa.Column("target_count", sa.Integer(), nullable=False, server_default="0"),
|
||||||
|
sa.Column("checked_count", sa.Integer(), nullable=False, server_default="0"),
|
||||||
|
sa.Column("issue_count", sa.Integer(), nullable=False, server_default="0"),
|
||||||
|
sa.Column("error_type", sa.String(64)),
|
||||||
|
sa.Column("error_message", sa.Text()),
|
||||||
|
sa.Column(
|
||||||
|
"created_at",
|
||||||
|
sa.DateTime(timezone=True),
|
||||||
|
nullable=False,
|
||||||
|
server_default=sa.text("now()"),
|
||||||
|
),
|
||||||
|
sa.Column(
|
||||||
|
"updated_at",
|
||||||
|
sa.DateTime(timezone=True),
|
||||||
|
nullable=False,
|
||||||
|
server_default=sa.text("now()"),
|
||||||
|
),
|
||||||
|
sa.Column("finished_at", sa.DateTime(timezone=True)),
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"ix_market_integrity_check_status_created_at",
|
||||||
|
"market_integrity_check",
|
||||||
|
["status", "created_at"],
|
||||||
|
)
|
||||||
|
# A partial unique index makes the running claim atomic under concurrent
|
||||||
|
# POST requests; stale recovery first moves an abandoned row to failed.
|
||||||
|
op.create_index(
|
||||||
|
"uq_market_integrity_check_running",
|
||||||
|
"market_integrity_check",
|
||||||
|
["status"],
|
||||||
|
unique=True,
|
||||||
|
postgresql_where=sa.text("status = 'running'"),
|
||||||
|
)
|
||||||
|
op.create_table(
|
||||||
|
"market_integrity_issue",
|
||||||
|
sa.Column(
|
||||||
|
"check_id",
|
||||||
|
sa.String(36),
|
||||||
|
sa.ForeignKey("market_integrity_check.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.Column("issue_key", sa.String(64), nullable=False),
|
||||||
|
sa.Column("item_kind", sa.String(24), nullable=False),
|
||||||
|
sa.Column("item_key", sa.String(128), nullable=False),
|
||||||
|
sa.Column("issue_type", sa.String(64), nullable=False),
|
||||||
|
sa.Column("message", sa.Text(), nullable=False),
|
||||||
|
sa.Column(
|
||||||
|
"created_at",
|
||||||
|
sa.DateTime(timezone=True),
|
||||||
|
nullable=False,
|
||||||
|
server_default=sa.text("now()"),
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("check_id", "issue_key"),
|
||||||
|
)
|
||||||
|
op.create_index("ix_market_integrity_issue_check_id", "market_integrity_issue", ["check_id"])
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
"""Drop only the integrity report tables in dependency-safe order."""
|
||||||
|
|
||||||
|
op.drop_index("ix_market_integrity_issue_check_id", table_name="market_integrity_issue")
|
||||||
|
op.drop_table("market_integrity_issue")
|
||||||
|
op.drop_index("uq_market_integrity_check_running", table_name="market_integrity_check")
|
||||||
|
op.drop_index(
|
||||||
|
"ix_market_integrity_check_status_created_at",
|
||||||
|
table_name="market_integrity_check",
|
||||||
|
)
|
||||||
|
op.drop_table("market_integrity_check")
|
||||||
@@ -9,7 +9,7 @@ dependencies = [
|
|||||||
"fastapi>=0.141.1",
|
"fastapi>=0.141.1",
|
||||||
"numpy>=2.4.0",
|
"numpy>=2.4.0",
|
||||||
"pandas>=2.3.3",
|
"pandas>=2.3.3",
|
||||||
"psycopg[binary]>=3.3.2",
|
"psycopg[binary,pool]>=3.3.2",
|
||||||
"pydantic-settings>=2.14.2",
|
"pydantic-settings>=2.14.2",
|
||||||
"sqlalchemy>=2.0.46",
|
"sqlalchemy>=2.0.46",
|
||||||
"tushare>=1.4.24",
|
"tushare>=1.4.24",
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from functools import lru_cache
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
|
from pydantic import Field
|
||||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||||
|
|
||||||
|
|
||||||
@@ -18,7 +19,7 @@ class Settings(BaseSettings):
|
|||||||
tushare_token: str = ""
|
tushare_token: str = ""
|
||||||
market_data_csv_root: Path = Path("./data/market-data")
|
market_data_csv_root: Path = Path("./data/market-data")
|
||||||
market_data_coverage_threshold: Decimal = Decimal("0.99")
|
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_request_interval_seconds: float = 0.2
|
||||||
market_data_max_retries: int = 3
|
market_data_max_retries: int = 3
|
||||||
market_data_retry_backoff_seconds: float = 1.0
|
market_data_retry_backoff_seconds: float = 1.0
|
||||||
|
|||||||
@@ -4,11 +4,17 @@ from fastapi import APIRouter
|
|||||||
|
|
||||||
from zhixing_server.interfaces.http.system import operational_router, system_router
|
from zhixing_server.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.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
|
from zhixing_server.modules.selection.presentation.http import selection_router
|
||||||
|
|
||||||
api_v1_router = APIRouter(prefix="/api/v1")
|
api_v1_router = APIRouter(prefix="/api/v1")
|
||||||
api_v1_router.include_router(system_router, prefix="/system", tags=["system"])
|
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(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"])
|
api_v1_router.include_router(selection_router, prefix="/selection", tags=["selection"])
|
||||||
|
|
||||||
__all__ = ["api_v1_router", "operational_router"]
|
__all__ = ["api_v1_router", "operational_router"]
|
||||||
|
|||||||
@@ -1,5 +1,12 @@
|
|||||||
"""Market data synchronization use cases."""
|
"""Market data synchronization use cases."""
|
||||||
|
|
||||||
from .sync import SyncBatchSummary, SyncMarketData, SyncMarketDataCommand
|
from .integrity import RunMarketIntegrityCheck
|
||||||
|
from .sync import SyncBatchSummary, SyncItemOutcome, SyncMarketData, SyncMarketDataCommand
|
||||||
|
|
||||||
__all__ = ["SyncBatchSummary", "SyncMarketData", "SyncMarketDataCommand"]
|
__all__ = [
|
||||||
|
"RunMarketIntegrityCheck",
|
||||||
|
"SyncBatchSummary",
|
||||||
|
"SyncItemOutcome",
|
||||||
|
"SyncMarketData",
|
||||||
|
"SyncMarketDataCommand",
|
||||||
|
]
|
||||||
|
|||||||
@@ -0,0 +1,671 @@
|
|||||||
|
"""Read-only PostgreSQL/CSV market-data integrity check use case."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from collections.abc import Callable, Iterable, Iterator, Sequence
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from datetime import date
|
||||||
|
from typing import TypeVar, cast
|
||||||
|
|
||||||
|
from ..domain.integrity import (
|
||||||
|
IntegrityCheckNoData,
|
||||||
|
IntegrityCheckNotFound,
|
||||||
|
IntegrityCheckPage,
|
||||||
|
IntegrityCheckQuery,
|
||||||
|
IntegrityCheckRun,
|
||||||
|
IntegrityCheckStore,
|
||||||
|
IntegrityIssue,
|
||||||
|
IntegritySnapshotReader,
|
||||||
|
IntegritySnapshotStore,
|
||||||
|
)
|
||||||
|
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
_ISSUE_BATCH_SIZE = 100
|
||||||
|
_DEFAULT_STALE_AFTER_SECONDS = 3_600
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
def _empty_issues() -> list[IntegrityIssue]:
|
||||||
|
"""Provide an explicitly typed empty issue buffer for Pyright."""
|
||||||
|
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def _empty_issue_keys() -> set[str]:
|
||||||
|
"""Provide an explicitly typed empty issue-key set for Pyright."""
|
||||||
|
|
||||||
|
return set()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class _IssueCollector:
|
||||||
|
"""Bound issue buffering and count progress without retaining the report."""
|
||||||
|
|
||||||
|
check_id: str
|
||||||
|
store: IntegrityCheckStore
|
||||||
|
issue_count: int = 0
|
||||||
|
pending: list[IntegrityIssue] = field(default_factory=_empty_issues)
|
||||||
|
issue_keys: set[str] = field(default_factory=_empty_issue_keys)
|
||||||
|
|
||||||
|
def add(
|
||||||
|
self,
|
||||||
|
item_kind: str,
|
||||||
|
item_key: str,
|
||||||
|
issue_type: str,
|
||||||
|
message: str,
|
||||||
|
) -> None:
|
||||||
|
"""Add one issue and flush at a bounded batch size."""
|
||||||
|
|
||||||
|
issue = IntegrityIssue.build(
|
||||||
|
self.check_id,
|
||||||
|
item_kind,
|
||||||
|
item_key,
|
||||||
|
issue_type,
|
||||||
|
message,
|
||||||
|
)
|
||||||
|
if issue.issue_key in self.issue_keys:
|
||||||
|
return
|
||||||
|
self.issue_keys.add(issue.issue_key)
|
||||||
|
self.pending.append(issue)
|
||||||
|
self.issue_count += 1
|
||||||
|
if len(self.pending) >= _ISSUE_BATCH_SIZE:
|
||||||
|
self.flush()
|
||||||
|
|
||||||
|
def flush(self) -> None:
|
||||||
|
"""Persist buffered issues and release their row objects."""
|
||||||
|
|
||||||
|
if not self.pending:
|
||||||
|
return
|
||||||
|
self.store.record_issues(tuple(self.pending))
|
||||||
|
self.pending.clear()
|
||||||
|
|
||||||
|
|
||||||
|
class RunMarketIntegrityCheck:
|
||||||
|
"""Prepare, execute, and query a persisted integrity report.
|
||||||
|
|
||||||
|
The application depends only on read ports and an independent check store.
|
||||||
|
It never constructs a Tushare adapter and never receives a market-data
|
||||||
|
write port, which makes the no-fact-mutation boundary explicit.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
reader: IntegritySnapshotReader,
|
||||||
|
snapshots: IntegritySnapshotStore,
|
||||||
|
store: IntegrityCheckStore,
|
||||||
|
*,
|
||||||
|
lock_key: int = 7_380_521,
|
||||||
|
stale_after_seconds: int = _DEFAULT_STALE_AFTER_SECONDS,
|
||||||
|
) -> None:
|
||||||
|
if stale_after_seconds < 1:
|
||||||
|
raise ValueError("stale_after_seconds must be positive")
|
||||||
|
self.reader = reader
|
||||||
|
self.snapshots = snapshots
|
||||||
|
self.store = store
|
||||||
|
self.lock_key = lock_key
|
||||||
|
self.stale_after_seconds = stale_after_seconds
|
||||||
|
|
||||||
|
def prepare(self) -> IntegrityCheckRun:
|
||||||
|
"""Atomically claim a check against the newest completed window."""
|
||||||
|
|
||||||
|
self.store.recover_stale_running(self.lock_key, self.stale_after_seconds)
|
||||||
|
window = self.reader.latest_successful_window()
|
||||||
|
if window is None:
|
||||||
|
raise IntegrityCheckNoData("no completed market-data synchronization is available")
|
||||||
|
stocks = tuple(self.reader.active_stocks())
|
||||||
|
target_count = self._target_count(window, stocks)
|
||||||
|
return self.store.create_running(window, target_count)
|
||||||
|
|
||||||
|
def execute(self, check_id: str) -> None:
|
||||||
|
"""Run a claimed check and always converge unexpected errors to failed."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
page = self.store.get(check_id, IntegrityCheckQuery(page=1, page_size=1))
|
||||||
|
except Exception as exc: # noqa: BLE001 - background boundary must converge state
|
||||||
|
self._finish_failed(check_id, 0, "check_error", _safe_error(exc))
|
||||||
|
return
|
||||||
|
if page is None:
|
||||||
|
raise IntegrityCheckNotFound(f"integrity check not found: {check_id}")
|
||||||
|
if page.run.status != "running":
|
||||||
|
return
|
||||||
|
collector = _IssueCollector(check_id, self.store, issue_count=page.run.issue_count)
|
||||||
|
try:
|
||||||
|
with self.reader.advisory_lock(self.lock_key) as acquired:
|
||||||
|
if not acquired:
|
||||||
|
self._finish_failed(
|
||||||
|
check_id,
|
||||||
|
collector.issue_count,
|
||||||
|
"lock_unavailable",
|
||||||
|
"market-data synchronization or another integrity check is running",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
self._execute_locked(page.run, collector)
|
||||||
|
except Exception as exc: # noqa: BLE001 - worker boundary must converge state
|
||||||
|
logger.exception("market_integrity_check_failed check_id=%s", check_id)
|
||||||
|
self._finish_failed(
|
||||||
|
check_id,
|
||||||
|
collector.issue_count,
|
||||||
|
"check_error",
|
||||||
|
_safe_error(exc),
|
||||||
|
)
|
||||||
|
except Exception as exc: # noqa: BLE001 - lock/storage boundary must be observable
|
||||||
|
logger.exception("market_integrity_check_worker_failed check_id=%s", check_id)
|
||||||
|
self._finish_failed(check_id, collector.issue_count, "check_error", _safe_error(exc))
|
||||||
|
|
||||||
|
def get(
|
||||||
|
self,
|
||||||
|
check_id: str,
|
||||||
|
*,
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 10,
|
||||||
|
query: IntegrityCheckQuery | None = None,
|
||||||
|
) -> IntegrityCheckPage | None:
|
||||||
|
"""Read one check and a bounded issue page."""
|
||||||
|
|
||||||
|
resolved = query or IntegrityCheckQuery(page=page, page_size=page_size)
|
||||||
|
return self.store.get(check_id, resolved)
|
||||||
|
|
||||||
|
def get_latest(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 10,
|
||||||
|
query: IntegrityCheckQuery | None = None,
|
||||||
|
) -> IntegrityCheckPage | None:
|
||||||
|
"""Read the newest check or ``None`` when no check has been run."""
|
||||||
|
|
||||||
|
resolved = query or IntegrityCheckQuery(page=page, page_size=page_size)
|
||||||
|
return self.store.get_latest(resolved)
|
||||||
|
|
||||||
|
def _target_count(self, window: SyncWindow, stocks: Sequence[Stock]) -> int:
|
||||||
|
"""Count comparison groups using keys only, never fact rows."""
|
||||||
|
|
||||||
|
bar_codes = set(_string_keys(self.reader, "list_bar_codes", window))
|
||||||
|
bar_codes.update(_string_keys(self.snapshots, "list_bar_codes", window))
|
||||||
|
bar_codes.update(stock.ts_code for stock in stocks)
|
||||||
|
basic_dates = set(_date_keys(self.reader, "list_daily_basic_dates", window))
|
||||||
|
basic_dates.update(_date_keys(self.snapshots, "list_daily_basic_dates", window))
|
||||||
|
return max(1, 1 + len(bar_codes) + len(basic_dates))
|
||||||
|
|
||||||
|
def _execute_locked(self, run: IntegrityCheckRun, collector: _IssueCollector) -> None:
|
||||||
|
"""Compare each group and persist progress after it completes."""
|
||||||
|
|
||||||
|
if run.window is None:
|
||||||
|
raise ValueError("integrity check has no comparison window")
|
||||||
|
window = run.window
|
||||||
|
self._compare_stocks(run.id, collector)
|
||||||
|
checked_count = 1
|
||||||
|
collector.flush()
|
||||||
|
self.store.update_progress(run.id, checked_count, collector.issue_count)
|
||||||
|
|
||||||
|
checked_count = self._compare_bars(
|
||||||
|
run.id,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
checked_count,
|
||||||
|
)
|
||||||
|
checked_count = self._compare_daily_basic(
|
||||||
|
run.id,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
checked_count,
|
||||||
|
)
|
||||||
|
collector.flush()
|
||||||
|
self.store.update_progress(run.id, checked_count, collector.issue_count)
|
||||||
|
self.store.finish(
|
||||||
|
run.id,
|
||||||
|
"issues_found" if collector.issue_count else "passed",
|
||||||
|
issue_count=collector.issue_count,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _compare_stocks(self, check_id: str, collector: _IssueCollector) -> None:
|
||||||
|
"""Compare current stock master rows as one small bounded group."""
|
||||||
|
|
||||||
|
database_rows = tuple(self.reader.active_stocks())
|
||||||
|
try:
|
||||||
|
csv_rows = self.snapshots.read_stocks()
|
||||||
|
except Exception as exc: # noqa: BLE001 - one malformed file must not stop other groups
|
||||||
|
collector.add("stock", "current", _issue_type(exc), "stock-master CSV cannot be parsed")
|
||||||
|
return
|
||||||
|
if csv_rows is None:
|
||||||
|
if database_rows:
|
||||||
|
collector.add(
|
||||||
|
"stock",
|
||||||
|
"current",
|
||||||
|
"missing_csv",
|
||||||
|
"current stock-master CSV is missing",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
database_by_code = {row.ts_code: row for row in database_rows}
|
||||||
|
csv_by_code = {row.ts_code: row for row in csv_rows}
|
||||||
|
for code in sorted(database_by_code.keys() - csv_by_code.keys()):
|
||||||
|
collector.add("stock", code, "missing_csv", "active stock is missing from CSV")
|
||||||
|
for code in sorted(csv_by_code.keys() - database_by_code.keys()):
|
||||||
|
collector.add("stock", code, "extra_csv", "CSV stock is not active in PostgreSQL")
|
||||||
|
for code in sorted(database_by_code.keys() & csv_by_code.keys()):
|
||||||
|
if database_by_code[code] != csv_by_code[code]:
|
||||||
|
collector.add("stock", code, "content_mismatch", "stock-master fields differ")
|
||||||
|
|
||||||
|
def _compare_bars(
|
||||||
|
self,
|
||||||
|
check_id: str,
|
||||||
|
window: SyncWindow,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
checked_count: int,
|
||||||
|
) -> int:
|
||||||
|
"""Merge PostgreSQL and per-stock CSV bars one stock at a time."""
|
||||||
|
|
||||||
|
database_codes = set(_string_keys(self.reader, "list_bar_codes", window))
|
||||||
|
csv_codes = set(_string_keys(self.snapshots, "list_bar_codes", window))
|
||||||
|
expected_codes = (
|
||||||
|
database_codes | csv_codes | {stock.ts_code for stock in self.reader.active_stocks()}
|
||||||
|
)
|
||||||
|
for code in _string_keys(self.reader, "list_invalid_bar_codes", window):
|
||||||
|
collector.add("bar", code, "content_mismatch", "database bar source_adj is not qfq")
|
||||||
|
database_groups = _grouped(self.reader.iter_bars(window), lambda row: row.ts_code)
|
||||||
|
parse_failed_codes: set[str] = set()
|
||||||
|
current = next(database_groups, None)
|
||||||
|
for code in sorted(expected_codes):
|
||||||
|
while current is not None and current[0] < code:
|
||||||
|
self._compare_bar_group(current[0], current[1], None, window, collector)
|
||||||
|
checked_count += 1
|
||||||
|
current = next(database_groups, None)
|
||||||
|
self._progress(collector, checked_count)
|
||||||
|
database_rows: tuple[Bar, ...] | None = None
|
||||||
|
if current is not None and current[0] == code:
|
||||||
|
database_rows = current[1]
|
||||||
|
current = next(database_groups, None)
|
||||||
|
csv_rows = (
|
||||||
|
self._read_bars(code, collector, parse_failed_codes) if code in csv_codes else None
|
||||||
|
)
|
||||||
|
self._compare_bar_group(
|
||||||
|
code,
|
||||||
|
database_rows,
|
||||||
|
csv_rows,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
csv_parse_failed=code in parse_failed_codes,
|
||||||
|
)
|
||||||
|
checked_count += 1
|
||||||
|
self._progress(collector, checked_count)
|
||||||
|
while current is not None:
|
||||||
|
self._compare_bar_group(current[0], current[1], None, window, collector)
|
||||||
|
checked_count += 1
|
||||||
|
current = next(database_groups, None)
|
||||||
|
self._progress(collector, checked_count)
|
||||||
|
return checked_count
|
||||||
|
|
||||||
|
def _compare_daily_basic(
|
||||||
|
self,
|
||||||
|
check_id: str,
|
||||||
|
window: SyncWindow,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
checked_count: int,
|
||||||
|
) -> int:
|
||||||
|
"""Merge PostgreSQL and per-date CSV daily-basic groups."""
|
||||||
|
|
||||||
|
database_dates = set(_date_keys(self.reader, "list_daily_basic_dates", window))
|
||||||
|
csv_dates = set(_date_keys(self.snapshots, "list_daily_basic_dates", window))
|
||||||
|
expected_dates = database_dates | csv_dates
|
||||||
|
database_groups = _grouped(self.reader.iter_daily_basic(window), lambda row: row.trade_date)
|
||||||
|
parse_failed_dates: set[date] = set()
|
||||||
|
invalid_path_method = getattr(
|
||||||
|
self.snapshots, "list_invalid_daily_basic_snapshot_files", None
|
||||||
|
)
|
||||||
|
if callable(invalid_path_method):
|
||||||
|
try:
|
||||||
|
invalid_paths = cast(Iterable[object], invalid_path_method())
|
||||||
|
except Exception: # noqa: BLE001 - malformed path listing must not abort other groups
|
||||||
|
invalid_paths = ()
|
||||||
|
for value in invalid_paths:
|
||||||
|
if not isinstance(value, tuple):
|
||||||
|
continue
|
||||||
|
parts = cast(tuple[object, ...], value)
|
||||||
|
if len(parts) != 2:
|
||||||
|
continue
|
||||||
|
item_key, issue_type = parts
|
||||||
|
collector.add(
|
||||||
|
"daily_basic",
|
||||||
|
str(item_key),
|
||||||
|
str(issue_type) or "parse_error",
|
||||||
|
"daily-basic snapshot path cannot be parsed",
|
||||||
|
)
|
||||||
|
current = next(database_groups, None)
|
||||||
|
for trade_date in sorted(expected_dates):
|
||||||
|
while current is not None and current[0] < trade_date:
|
||||||
|
self._compare_basic_group(current[0], current[1], None, window, collector)
|
||||||
|
checked_count += 1
|
||||||
|
current = next(database_groups, None)
|
||||||
|
self._progress(collector, checked_count)
|
||||||
|
database_rows: tuple[DailyBasic, ...] | None = None
|
||||||
|
if current is not None and current[0] == trade_date:
|
||||||
|
database_rows = current[1]
|
||||||
|
current = next(database_groups, None)
|
||||||
|
csv_rows = (
|
||||||
|
self._read_daily_basic(trade_date, collector, parse_failed_dates)
|
||||||
|
if trade_date in csv_dates
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self._compare_basic_group(
|
||||||
|
trade_date,
|
||||||
|
database_rows,
|
||||||
|
csv_rows,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
csv_parse_failed=trade_date in parse_failed_dates,
|
||||||
|
)
|
||||||
|
checked_count += 1
|
||||||
|
self._progress(collector, checked_count)
|
||||||
|
while current is not None:
|
||||||
|
self._compare_basic_group(current[0], current[1], None, window, collector)
|
||||||
|
checked_count += 1
|
||||||
|
current = next(database_groups, None)
|
||||||
|
self._progress(collector, checked_count)
|
||||||
|
return checked_count
|
||||||
|
|
||||||
|
def _read_bars(
|
||||||
|
self,
|
||||||
|
code: str,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
parse_failed_codes: set[str],
|
||||||
|
) -> tuple[Bar, ...] | None:
|
||||||
|
"""Read one formal bar file and isolate its parse failure."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
return self.snapshots.read_bars(code)
|
||||||
|
except Exception as exc: # noqa: BLE001 - continue with other stocks
|
||||||
|
parse_failed_codes.add(code)
|
||||||
|
collector.add("bar", code, _issue_type(exc), "bar CSV cannot be parsed")
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _read_daily_basic(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
parse_failed_dates: set[date],
|
||||||
|
) -> tuple[DailyBasic, ...] | None:
|
||||||
|
"""Read one formal daily-basic file and isolate its parse failure."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
return self.snapshots.read_daily_basic(trade_date)
|
||||||
|
except Exception as exc: # noqa: BLE001 - continue with other dates
|
||||||
|
parse_failed_dates.add(trade_date)
|
||||||
|
collector.add(
|
||||||
|
"daily_basic",
|
||||||
|
trade_date.isoformat(),
|
||||||
|
_issue_type(exc),
|
||||||
|
"daily-basic CSV cannot be parsed",
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _compare_bar_group(
|
||||||
|
self,
|
||||||
|
code: str,
|
||||||
|
database_rows: tuple[Bar, ...] | None,
|
||||||
|
csv_rows: tuple[Bar, ...] | None,
|
||||||
|
window: SyncWindow,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
*,
|
||||||
|
csv_parse_failed: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""Compare one stock's rows without building a market-wide map."""
|
||||||
|
|
||||||
|
database_inside = self._keep_bar_rows_in_window(
|
||||||
|
"database", code, database_rows, window, collector
|
||||||
|
)
|
||||||
|
if csv_parse_failed:
|
||||||
|
return
|
||||||
|
csv_inside = self._keep_bar_rows_in_window("csv", code, csv_rows, window, collector)
|
||||||
|
if database_inside is None and csv_inside is None:
|
||||||
|
return
|
||||||
|
if database_inside is None:
|
||||||
|
if csv_inside:
|
||||||
|
self._compare_rows(
|
||||||
|
"bar",
|
||||||
|
code,
|
||||||
|
(),
|
||||||
|
csv_inside,
|
||||||
|
lambda row: row.trade_date,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
row_date=lambda row: row.trade_date,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if csv_inside is None:
|
||||||
|
if database_inside:
|
||||||
|
self._compare_rows(
|
||||||
|
"bar",
|
||||||
|
code,
|
||||||
|
database_inside,
|
||||||
|
(),
|
||||||
|
lambda row: row.trade_date,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
row_date=lambda row: row.trade_date,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
self._compare_rows(
|
||||||
|
"bar",
|
||||||
|
code,
|
||||||
|
database_inside,
|
||||||
|
csv_inside,
|
||||||
|
lambda row: row.trade_date,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
row_date=lambda row: row.trade_date,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _keep_bar_rows_in_window(
|
||||||
|
source: str,
|
||||||
|
code: str,
|
||||||
|
rows: tuple[Bar, ...] | None,
|
||||||
|
window: SyncWindow,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
) -> tuple[Bar, ...] | None:
|
||||||
|
"""Report and discard out-of-window CSV/DB rows before merging."""
|
||||||
|
|
||||||
|
if rows is None:
|
||||||
|
return None
|
||||||
|
inside: list[Bar] = []
|
||||||
|
for row in rows:
|
||||||
|
if window.contains(row.trade_date):
|
||||||
|
inside.append(row)
|
||||||
|
else:
|
||||||
|
collector.add(
|
||||||
|
"bar",
|
||||||
|
f"{code}:{row.trade_date.isoformat()}",
|
||||||
|
"window_out_of_bounds",
|
||||||
|
f"{source} bar row is outside the check window",
|
||||||
|
)
|
||||||
|
return tuple(inside)
|
||||||
|
|
||||||
|
def _compare_basic_group(
|
||||||
|
self,
|
||||||
|
trade_date: date,
|
||||||
|
database_rows: tuple[DailyBasic, ...] | None,
|
||||||
|
csv_rows: tuple[DailyBasic, ...] | None,
|
||||||
|
window: SyncWindow,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
*,
|
||||||
|
csv_parse_failed: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""Compare one trading-day's daily-basic rows."""
|
||||||
|
|
||||||
|
group_key = trade_date.isoformat()
|
||||||
|
if not window.contains(trade_date):
|
||||||
|
collector.add(
|
||||||
|
"daily_basic",
|
||||||
|
group_key,
|
||||||
|
"window_out_of_bounds",
|
||||||
|
"daily-basic date is outside the check window",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if csv_parse_failed:
|
||||||
|
return
|
||||||
|
if database_rows is None and csv_rows is None:
|
||||||
|
return
|
||||||
|
if database_rows is None:
|
||||||
|
collector.add(
|
||||||
|
"daily_basic",
|
||||||
|
group_key,
|
||||||
|
"extra_csv",
|
||||||
|
"daily-basic date exists only in CSV",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if csv_rows is None:
|
||||||
|
collector.add(
|
||||||
|
"daily_basic", group_key, "missing_csv", "daily-basic date is missing in CSV"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
self._compare_rows(
|
||||||
|
"daily_basic",
|
||||||
|
group_key,
|
||||||
|
database_rows,
|
||||||
|
csv_rows,
|
||||||
|
lambda row: row.ts_code,
|
||||||
|
window,
|
||||||
|
collector,
|
||||||
|
row_date=lambda row: row.trade_date,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _compare_rows(
|
||||||
|
self,
|
||||||
|
item_kind: str,
|
||||||
|
group_key: str,
|
||||||
|
database_rows: Sequence[T],
|
||||||
|
csv_rows: Sequence[T],
|
||||||
|
key: Callable[[T], object],
|
||||||
|
window: SyncWindow,
|
||||||
|
collector: _IssueCollector,
|
||||||
|
row_date: Callable[[T], date] | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Merge two bounded groups and report duplicate/content differences."""
|
||||||
|
|
||||||
|
database_by_key, database_duplicates = _index_rows(database_rows, key)
|
||||||
|
csv_by_key, csv_duplicates = _index_rows(csv_rows, key)
|
||||||
|
for row_key in sorted(database_duplicates | csv_duplicates, key=str):
|
||||||
|
collector.add(
|
||||||
|
item_kind,
|
||||||
|
f"{group_key}:{row_key}",
|
||||||
|
"duplicate_key",
|
||||||
|
"comparison group contains duplicate business keys",
|
||||||
|
)
|
||||||
|
for row_key in sorted(database_by_key.keys() | csv_by_key.keys(), key=str):
|
||||||
|
database_row = database_by_key.get(row_key)
|
||||||
|
csv_row = csv_by_key.get(row_key)
|
||||||
|
identity = f"{group_key}:{row_key}"
|
||||||
|
if row_date is not None:
|
||||||
|
for candidate in (database_row, csv_row):
|
||||||
|
if candidate is not None and not window.contains(row_date(candidate)):
|
||||||
|
collector.add(
|
||||||
|
item_kind,
|
||||||
|
identity,
|
||||||
|
"window_out_of_bounds",
|
||||||
|
"row is outside the check window",
|
||||||
|
)
|
||||||
|
if database_row is None:
|
||||||
|
collector.add(item_kind, identity, "extra_csv", "row exists only in CSV")
|
||||||
|
elif csv_row is None:
|
||||||
|
collector.add(item_kind, identity, "missing_csv", "row is missing in CSV")
|
||||||
|
elif database_row != csv_row:
|
||||||
|
collector.add(
|
||||||
|
item_kind, identity, "content_mismatch", "database and CSV fields differ"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _progress(self, collector: _IssueCollector, checked_count: int) -> None:
|
||||||
|
"""Persist issues and a heartbeat after one comparison group."""
|
||||||
|
|
||||||
|
collector.flush()
|
||||||
|
self.store.update_progress(collector.check_id, checked_count, collector.issue_count)
|
||||||
|
|
||||||
|
def _finish_failed(
|
||||||
|
self,
|
||||||
|
check_id: str,
|
||||||
|
issue_count: int,
|
||||||
|
error_type: str,
|
||||||
|
message: str,
|
||||||
|
) -> None:
|
||||||
|
"""Best-effort terminal write for background worker failures."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
self.store.finish(
|
||||||
|
check_id,
|
||||||
|
"failed",
|
||||||
|
issue_count=issue_count,
|
||||||
|
error_type=error_type,
|
||||||
|
error_message=message,
|
||||||
|
)
|
||||||
|
except Exception: # noqa: BLE001 - nothing safer can be persisted here
|
||||||
|
logger.exception("market_integrity_check_failure_persist_failed check_id=%s", check_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _grouped[T, K](rows: Iterable[T], key: Callable[[T], K]) -> Iterator[tuple[K, tuple[T, ...]]]:
|
||||||
|
"""Consume an ordered iterator one comparison group at a time."""
|
||||||
|
|
||||||
|
iterator = iter(rows)
|
||||||
|
pending = next(iterator, None)
|
||||||
|
while pending is not None:
|
||||||
|
group_key = key(pending)
|
||||||
|
group: list[T] = [pending]
|
||||||
|
pending = next(iterator, None)
|
||||||
|
while pending is not None and key(pending) == group_key:
|
||||||
|
group.append(pending)
|
||||||
|
pending = next(iterator, None)
|
||||||
|
yield group_key, tuple(group)
|
||||||
|
|
||||||
|
|
||||||
|
def _index_rows[T](
|
||||||
|
rows: Sequence[T], key: Callable[[T], object]
|
||||||
|
) -> tuple[dict[object, T], set[object]]:
|
||||||
|
"""Index one bounded group and retain duplicate identities only."""
|
||||||
|
|
||||||
|
indexed: dict[object, T] = {}
|
||||||
|
duplicates: set[object] = set()
|
||||||
|
for row in rows:
|
||||||
|
row_key = key(row)
|
||||||
|
if row_key in indexed:
|
||||||
|
duplicates.add(row_key)
|
||||||
|
else:
|
||||||
|
indexed[row_key] = row
|
||||||
|
return indexed, duplicates
|
||||||
|
|
||||||
|
|
||||||
|
def _string_keys(adapter: object, method_name: str, window: SyncWindow) -> tuple[str, ...]:
|
||||||
|
"""Call an optional string-key method without hiding storage failures."""
|
||||||
|
|
||||||
|
method = getattr(adapter, method_name, None)
|
||||||
|
if not callable(method):
|
||||||
|
return ()
|
||||||
|
values = cast(Iterable[object], method(window))
|
||||||
|
return tuple(str(value) for value in values)
|
||||||
|
|
||||||
|
|
||||||
|
def _date_keys(adapter: object, method_name: str, window: SyncWindow) -> tuple[date, ...]:
|
||||||
|
"""Call an optional date-key method without hiding storage failures."""
|
||||||
|
|
||||||
|
method = getattr(adapter, method_name, None)
|
||||||
|
if not callable(method):
|
||||||
|
return ()
|
||||||
|
values = cast(Iterable[object], method(window))
|
||||||
|
return tuple(value for value in values if isinstance(value, date))
|
||||||
|
|
||||||
|
|
||||||
|
def _issue_type(error: BaseException) -> str:
|
||||||
|
"""Use adapter-provided stable issue categories when available."""
|
||||||
|
|
||||||
|
value = getattr(error, "issue_type", "parse_error")
|
||||||
|
return str(value) if value else "parse_error"
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_error(error: Exception) -> str:
|
||||||
|
"""Keep worker failures bounded and free of tracebacks or secrets."""
|
||||||
|
|
||||||
|
return " ".join(str(error).split())[:500] or error.__class__.__name__
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["RunMarketIntegrityCheck"]
|
||||||
@@ -4,11 +4,12 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import time
|
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 dataclasses import dataclass, field
|
||||||
from datetime import date, timedelta
|
from datetime import date, timedelta
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from typing import Literal
|
from typing import Literal, cast
|
||||||
|
|
||||||
from ..domain.fingerprint import SnapshotChange, compare_snapshots
|
from ..domain.fingerprint import SnapshotChange, compare_snapshots
|
||||||
from ..domain.models import Bar, Stock, SyncWindow
|
from ..domain.models import Bar, Stock, SyncWindow
|
||||||
@@ -40,6 +41,24 @@ class SyncFailure:
|
|||||||
message: str
|
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)
|
@dataclass(frozen=True, slots=True)
|
||||||
class SyncBatchSummary:
|
class SyncBatchSummary:
|
||||||
"""Stable output contract for CLI, cron, and later strategy callers."""
|
"""Stable output contract for CLI, cron, and later strategy callers."""
|
||||||
@@ -63,7 +82,9 @@ class SyncBatchSummary:
|
|||||||
|
|
||||||
if self.status == "failed":
|
if self.status == "failed":
|
||||||
return 1
|
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]:
|
def as_dict(self) -> dict[str, object]:
|
||||||
"""Serialize the summary without credentials or raw vendor responses."""
|
"""Serialize the summary without credentials or raw vendor responses."""
|
||||||
@@ -107,36 +128,55 @@ class SyncMarketData:
|
|||||||
coverage_threshold: Decimal = Decimal("0.99"),
|
coverage_threshold: Decimal = Decimal("0.99"),
|
||||||
lock_key: int = 7_380_521,
|
lock_key: int = 7_380_521,
|
||||||
today: date | None = None,
|
today: date | None = None,
|
||||||
|
max_workers: int = 8,
|
||||||
) -> None:
|
) -> None:
|
||||||
if not Decimal("0") <= coverage_threshold <= Decimal("1"):
|
if not Decimal("0") <= coverage_threshold <= Decimal("1"):
|
||||||
raise ValueError("coverage_threshold must be between 0 and 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.source = source
|
||||||
self.snapshots = snapshots
|
self.snapshots = snapshots
|
||||||
self.repository = repository
|
self.repository = repository
|
||||||
self.coverage_threshold = coverage_threshold
|
self.coverage_threshold = coverage_threshold
|
||||||
self.lock_key = lock_key
|
self.lock_key = lock_key
|
||||||
self.today = today or date.today()
|
self.today = today or date.today()
|
||||||
|
self.max_workers = max_workers
|
||||||
|
|
||||||
def execute(self, command: SyncMarketDataCommand | None = None) -> SyncBatchSummary:
|
def execute(self, command: SyncMarketDataCommand | None = None) -> SyncBatchSummary:
|
||||||
"""Run one synchronization and retain successful items on partial failure."""
|
"""Run one synchronization and retain successful items on partial failure."""
|
||||||
|
|
||||||
command = command or SyncMarketDataCommand()
|
command = command or SyncMarketDataCommand()
|
||||||
with self.repository.advisory_lock(self.lock_key) as acquired:
|
try:
|
||||||
if not acquired:
|
with self.repository.advisory_lock(self.lock_key) as acquired:
|
||||||
return SyncBatchSummary(
|
if not acquired:
|
||||||
batch_id=None,
|
return SyncBatchSummary(
|
||||||
target_trade_date=None,
|
batch_id=None,
|
||||||
window=None,
|
target_trade_date=None,
|
||||||
status="failed",
|
window=None,
|
||||||
target_count=0,
|
status="failed",
|
||||||
valid_count=0,
|
target_count=0,
|
||||||
coverage=Decimal("0"),
|
valid_count=0,
|
||||||
strategy_eligible=False,
|
coverage=Decimal("0"),
|
||||||
failures=(
|
strategy_eligible=False,
|
||||||
SyncFailure("batch", "lock", "sync_locked", "another sync is running"),
|
failures=(
|
||||||
),
|
SyncFailure("batch", "lock", "sync_locked", "another sync is running"),
|
||||||
)
|
),
|
||||||
return self._execute_locked(command)
|
)
|
||||||
|
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:
|
def _execute_locked(self, command: SyncMarketDataCommand) -> SyncBatchSummary:
|
||||||
started_at = time.monotonic()
|
started_at = time.monotonic()
|
||||||
@@ -145,7 +185,20 @@ class SyncMarketData:
|
|||||||
command.mode,
|
command.mode,
|
||||||
command.target_trade_date or "auto",
|
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)
|
window = SyncWindow.from_target(target_trade_date)
|
||||||
logger.info(
|
logger.info(
|
||||||
"market_data_sync_target target_trade_date=%s window_start=%s window_end=%s",
|
"market_data_sync_target target_trade_date=%s window_start=%s window_end=%s",
|
||||||
@@ -153,7 +206,20 @@ class SyncMarketData:
|
|||||||
window.start,
|
window.start,
|
||||||
window.end,
|
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:
|
if not all_stocks:
|
||||||
return SyncBatchSummary(
|
return SyncBatchSummary(
|
||||||
None,
|
None,
|
||||||
@@ -168,13 +234,44 @@ class SyncMarketData:
|
|||||||
SyncFailure("stock", "universe", "empty_universe", "no eligible stocks"),
|
SyncFailure("stock", "universe", "empty_universe", "no eligible stocks"),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
batch_id = self.repository.create_batch(
|
try:
|
||||||
target_trade_date,
|
retry_items: set[tuple[str, str]] = (
|
||||||
window,
|
self._retry_items(command.parent_batch_id) if command.mode == "retry" else set()
|
||||||
command.mode,
|
)
|
||||||
command.parent_batch_id,
|
dates = self._dates_to_process(window, target_trade_date, command.mode, retry_items)
|
||||||
len(all_stocks),
|
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(
|
logger.info(
|
||||||
"market_data_sync_batch_created batch_id=%s target_count=%d",
|
"market_data_sync_batch_created batch_id=%s target_count=%d",
|
||||||
batch_id,
|
batch_id,
|
||||||
@@ -182,12 +279,19 @@ class SyncMarketData:
|
|||||||
)
|
)
|
||||||
failures: list[SyncFailure] = []
|
failures: list[SyncFailure] = []
|
||||||
totals = [0, 0, 0]
|
totals = [0, 0, 0]
|
||||||
|
pending_outcomes: list[SyncItemOutcome] = []
|
||||||
|
audit_failed = False
|
||||||
stock_codes = {stock.ts_code for stock in all_stocks}
|
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(
|
self._log_progress(
|
||||||
batch_id=batch_id,
|
batch_id=batch_id,
|
||||||
stage="stock_master",
|
stage="stock_master",
|
||||||
@@ -199,7 +303,6 @@ class SyncMarketData:
|
|||||||
started_at=started_at,
|
started_at=started_at,
|
||||||
force=bool(failures),
|
force=bool(failures),
|
||||||
)
|
)
|
||||||
dates = self._dates_to_process(window, target_trade_date, command.mode, retry_items)
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"market_data_sync_stage_started batch_id=%s stage=daily_basic total=%d",
|
"market_data_sync_stage_started batch_id=%s stage=daily_basic total=%d",
|
||||||
batch_id,
|
batch_id,
|
||||||
@@ -207,7 +310,19 @@ class SyncMarketData:
|
|||||||
)
|
)
|
||||||
for current, trade_date in enumerate(dates, start=1):
|
for current, trade_date in enumerate(dates, start=1):
|
||||||
failure_count = len(failures)
|
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(
|
self._log_progress(
|
||||||
batch_id=batch_id,
|
batch_id=batch_id,
|
||||||
stage="daily_basic",
|
stage="daily_basic",
|
||||||
@@ -234,19 +349,55 @@ class SyncMarketData:
|
|||||||
batch_id,
|
batch_id,
|
||||||
len(stocks_to_process),
|
len(stocks_to_process),
|
||||||
)
|
)
|
||||||
for current, stock in enumerate(stocks_to_process, start=1):
|
try:
|
||||||
failure_count = len(failures)
|
with ThreadPoolExecutor(
|
||||||
self._process_bar(batch_id, stock, window, failures, totals)
|
max_workers=self.max_workers,
|
||||||
self._log_progress(
|
thread_name_prefix="market-data-bar",
|
||||||
batch_id=batch_id,
|
) as executor:
|
||||||
stage="bar",
|
futures = {
|
||||||
current=current,
|
executor.submit(self._process_bar, stock, window): stock.ts_code
|
||||||
total=len(stocks_to_process),
|
for stock in stocks_to_process
|
||||||
item_key=stock.ts_code,
|
}
|
||||||
totals=totals,
|
for current, future in enumerate(as_completed(futures), start=1):
|
||||||
failures=failures,
|
item_key = futures[future]
|
||||||
started_at=started_at,
|
failure_count = len(failures)
|
||||||
force=len(failures) > failure_count,
|
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(
|
logger.info(
|
||||||
"market_data_sync_stage_completed batch_id=%s stage=bar total=%d",
|
"market_data_sync_stage_completed batch_id=%s stage=bar total=%d",
|
||||||
@@ -254,32 +405,62 @@ class SyncMarketData:
|
|||||||
len(stocks_to_process),
|
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:
|
if not failures:
|
||||||
try:
|
try:
|
||||||
self.repository.purge_before(window)
|
self.repository.purge_before(window)
|
||||||
self.snapshots.clean_daily_basic_before(window.start)
|
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)
|
failure = self._failure("batch", "retention", exc)
|
||||||
failures.append(failure)
|
failures.append(failure)
|
||||||
self.repository.record_item(
|
audit_failure = self._record_outcomes(
|
||||||
batch_id,
|
batch_id,
|
||||||
"batch",
|
[
|
||||||
"retention",
|
SyncItemOutcome(
|
||||||
"failed",
|
"batch",
|
||||||
WriteResult(),
|
"retention",
|
||||||
error_type=failure.error_type,
|
"failed",
|
||||||
error_message=failure.message,
|
failure=failure,
|
||||||
|
)
|
||||||
|
],
|
||||||
)
|
)
|
||||||
valid_count = sum(
|
if audit_failure is not None and not audit_failed:
|
||||||
1
|
failures.append(audit_failure)
|
||||||
for stock in all_stocks
|
audit_failed = True
|
||||||
if self.repository.has_bar(stock.ts_code, target_trade_date)
|
try:
|
||||||
and self.repository.has_daily_basic(stock.ts_code, target_trade_date)
|
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))
|
coverage = Decimal(valid_count) / Decimal(len(all_stocks))
|
||||||
status = "success" if not failures else "partial_success" if valid_count else "failed"
|
status = "success" if not failures else "partial_success" if valid_count else "failed"
|
||||||
eligible = coverage >= self.coverage_threshold
|
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(
|
logger.info(
|
||||||
"market_data_sync_finished batch_id=%s status=%s target_count=%d valid_count=%d "
|
"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",
|
"coverage=%s failures=%d inserted=%d updated=%d unchanged=%d elapsed_seconds=%.1f",
|
||||||
@@ -294,6 +475,12 @@ class SyncMarketData:
|
|||||||
totals[2],
|
totals[2],
|
||||||
time.monotonic() - started_at,
|
time.monotonic() - started_at,
|
||||||
)
|
)
|
||||||
|
ordered_failures = tuple(
|
||||||
|
sorted(
|
||||||
|
failures,
|
||||||
|
key=lambda failure: (failure.item_kind, failure.item_key, failure.error_type),
|
||||||
|
)
|
||||||
|
)
|
||||||
return SyncBatchSummary(
|
return SyncBatchSummary(
|
||||||
batch_id,
|
batch_id,
|
||||||
target_trade_date,
|
target_trade_date,
|
||||||
@@ -306,7 +493,7 @@ class SyncMarketData:
|
|||||||
totals[0],
|
totals[0],
|
||||||
totals[1],
|
totals[1],
|
||||||
totals[2],
|
totals[2],
|
||||||
tuple(failures),
|
ordered_failures,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _resolve_target(self, requested: date | None) -> date:
|
def _resolve_target(self, requested: date | None) -> date:
|
||||||
@@ -346,57 +533,37 @@ class SyncMarketData:
|
|||||||
|
|
||||||
def _process_stock_master(
|
def _process_stock_master(
|
||||||
self,
|
self,
|
||||||
batch_id: str,
|
|
||||||
stocks: Sequence[Stock],
|
stocks: Sequence[Stock],
|
||||||
failures: list[SyncFailure],
|
) -> SyncItemOutcome:
|
||||||
totals: list[int],
|
|
||||||
) -> None:
|
|
||||||
staged = None
|
staged = None
|
||||||
try:
|
try:
|
||||||
staged = self.snapshots.stage_stocks(stocks)
|
staged = self.snapshots.stage_stocks(stocks)
|
||||||
result = self.repository.upsert_stocks(stocks)
|
result = self.repository.upsert_stocks(stocks)
|
||||||
self.snapshots.publish(staged)
|
self.snapshots.publish(staged)
|
||||||
totals[0] += result.inserted
|
return SyncItemOutcome(
|
||||||
totals[1] += result.updated
|
|
||||||
totals[2] += result.unchanged
|
|
||||||
self.repository.record_item(
|
|
||||||
batch_id,
|
|
||||||
"stock",
|
"stock",
|
||||||
"current",
|
"current",
|
||||||
"success",
|
"success",
|
||||||
result,
|
result,
|
||||||
staged.fingerprint,
|
staged.fingerprint,
|
||||||
)
|
)
|
||||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
except Exception as exc:
|
||||||
if staged is not None:
|
if staged is not None:
|
||||||
self.snapshots.discard(staged)
|
self.snapshots.discard(staged)
|
||||||
failure = self._failure("stock", "current", exc)
|
failure = self._failure("stock", "current", exc)
|
||||||
failures.append(failure)
|
return SyncItemOutcome(
|
||||||
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,
|
|
||||||
"stock",
|
"stock",
|
||||||
"current",
|
"current",
|
||||||
"failed",
|
"failed",
|
||||||
WriteResult(),
|
failure=failure,
|
||||||
error_type=failure.error_type,
|
|
||||||
error_message=failure.message,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _process_daily_basic(
|
def _process_daily_basic(
|
||||||
self,
|
self,
|
||||||
batch_id: str,
|
|
||||||
trade_date: date,
|
trade_date: date,
|
||||||
stock_codes: set[str],
|
stock_codes: set[str],
|
||||||
window: SyncWindow,
|
window: SyncWindow,
|
||||||
failures: list[SyncFailure],
|
) -> SyncItemOutcome:
|
||||||
totals: list[int],
|
|
||||||
) -> None:
|
|
||||||
staged = None
|
staged = None
|
||||||
key = trade_date.isoformat()
|
key = trade_date.isoformat()
|
||||||
try:
|
try:
|
||||||
@@ -410,50 +577,35 @@ class SyncMarketData:
|
|||||||
staged = self.snapshots.stage_daily_basic(trade_date, rows)
|
staged = self.snapshots.stage_daily_basic(trade_date, rows)
|
||||||
result = self.repository.upsert_daily_basic(rows, window)
|
result = self.repository.upsert_daily_basic(rows, window)
|
||||||
self.snapshots.publish(staged)
|
self.snapshots.publish(staged)
|
||||||
self._add_counts(totals, result)
|
return SyncItemOutcome(
|
||||||
self.repository.record_item(
|
|
||||||
batch_id,
|
|
||||||
"daily_basic",
|
"daily_basic",
|
||||||
key,
|
key,
|
||||||
"success",
|
"success",
|
||||||
result,
|
result,
|
||||||
staged.fingerprint,
|
staged.fingerprint,
|
||||||
)
|
)
|
||||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
except Exception as exc:
|
||||||
if staged is not None:
|
if staged is not None:
|
||||||
self.snapshots.discard(staged)
|
self.snapshots.discard(staged)
|
||||||
failure = self._failure("daily_basic", key, exc)
|
failure = self._failure("daily_basic", key, exc)
|
||||||
failures.append(failure)
|
return SyncItemOutcome(
|
||||||
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,
|
|
||||||
"daily_basic",
|
"daily_basic",
|
||||||
key,
|
key,
|
||||||
"failed",
|
"failed",
|
||||||
WriteResult(),
|
failure=failure,
|
||||||
error_type=failure.error_type,
|
|
||||||
error_message=failure.message,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _process_bar(
|
def _process_bar(
|
||||||
self,
|
self,
|
||||||
batch_id: str,
|
|
||||||
stock: Stock,
|
stock: Stock,
|
||||||
window: SyncWindow,
|
window: SyncWindow,
|
||||||
failures: list[SyncFailure],
|
) -> SyncItemOutcome:
|
||||||
totals: list[int],
|
|
||||||
) -> None:
|
|
||||||
staged = None
|
staged = None
|
||||||
try:
|
try:
|
||||||
rows = tuple(self.source.fetch_bars(stock.ts_code, window))
|
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)
|
old_rows = self.snapshots.read_bars(stock.ts_code)
|
||||||
comparison = compare_snapshots(old_rows, rows)
|
comparison = compare_snapshots(old_rows, rows)
|
||||||
staged = self.snapshots.stage_bars(stock.ts_code, rows)
|
|
||||||
if comparison.change is SnapshotChange.UNCHANGED:
|
if comparison.change is SnapshotChange.UNCHANGED:
|
||||||
result = WriteResult(unchanged=len(rows))
|
result = WriteResult(unchanged=len(rows))
|
||||||
else:
|
else:
|
||||||
@@ -468,34 +620,22 @@ class SyncMarketData:
|
|||||||
in {SnapshotChange.INITIAL, SnapshotChange.CHANGED},
|
in {SnapshotChange.INITIAL, SnapshotChange.CHANGED},
|
||||||
)
|
)
|
||||||
self.snapshots.publish(staged)
|
self.snapshots.publish(staged)
|
||||||
self._add_counts(totals, result)
|
return SyncItemOutcome(
|
||||||
self.repository.record_item(
|
|
||||||
batch_id,
|
|
||||||
"bar",
|
"bar",
|
||||||
stock.ts_code,
|
stock.ts_code,
|
||||||
"success",
|
"success",
|
||||||
result,
|
result,
|
||||||
staged.fingerprint,
|
staged.fingerprint,
|
||||||
)
|
)
|
||||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
except Exception as exc:
|
||||||
if staged is not None:
|
if staged is not None:
|
||||||
self.snapshots.discard(staged)
|
self.snapshots.discard(staged)
|
||||||
failure = self._failure("bar", stock.ts_code, exc)
|
failure = self._failure("bar", stock.ts_code, exc)
|
||||||
failures.append(failure)
|
return SyncItemOutcome(
|
||||||
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,
|
|
||||||
"bar",
|
"bar",
|
||||||
stock.ts_code,
|
stock.ts_code,
|
||||||
"failed",
|
"failed",
|
||||||
WriteResult(),
|
failure=failure,
|
||||||
error_type=failure.error_type,
|
|
||||||
error_message=failure.message,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -504,6 +644,68 @@ class SyncMarketData:
|
|||||||
totals[1] += result.updated
|
totals[1] += result.updated
|
||||||
totals[2] += result.unchanged
|
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
|
@staticmethod
|
||||||
def _log_progress(
|
def _log_progress(
|
||||||
*,
|
*,
|
||||||
|
|||||||
@@ -1,12 +1,32 @@
|
|||||||
"""Pure market data domain types and ports."""
|
"""Pure market data domain types and ports."""
|
||||||
|
|
||||||
from .fingerprint import SnapshotChange, SnapshotComparison, compare_snapshots, snapshot_fingerprint
|
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 .models import Bar, DailyBasic, Stock, SyncWindow
|
||||||
from .rules import filter_current_hs_a_stocks, is_current_hs_a_stock
|
from .rules import filter_current_hs_a_stocks, is_current_hs_a_stock
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Bar",
|
"Bar",
|
||||||
"DailyBasic",
|
"DailyBasic",
|
||||||
|
"IntegrityCheckInProgress",
|
||||||
|
"IntegrityCheckNoData",
|
||||||
|
"IntegrityCheckNotFound",
|
||||||
|
"IntegrityCheckPage",
|
||||||
|
"IntegrityCheckQuery",
|
||||||
|
"IntegrityCheckRun",
|
||||||
|
"IntegrityCheckStoreError",
|
||||||
|
"IntegrityIssue",
|
||||||
|
"IntegrityStatus",
|
||||||
"SnapshotChange",
|
"SnapshotChange",
|
||||||
"SnapshotComparison",
|
"SnapshotComparison",
|
||||||
"Stock",
|
"Stock",
|
||||||
|
|||||||
@@ -0,0 +1,187 @@
|
|||||||
|
"""Domain contracts for read-only market-data integrity checks."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Iterable, Sequence
|
||||||
|
from contextlib import AbstractContextManager
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from datetime import date, datetime
|
||||||
|
from hashlib import sha256
|
||||||
|
from typing import Literal, Protocol
|
||||||
|
|
||||||
|
from .models import Bar, DailyBasic, Stock, SyncWindow
|
||||||
|
|
||||||
|
IntegrityStatus = Literal["running", "passed", "issues_found", "failed"]
|
||||||
|
|
||||||
|
|
||||||
|
class IntegrityCheckInProgress(RuntimeError):
|
||||||
|
"""Raised when an integrity check is already running."""
|
||||||
|
|
||||||
|
|
||||||
|
class IntegrityCheckNoData(RuntimeError):
|
||||||
|
"""Raised when no completed market-data window is available to check."""
|
||||||
|
|
||||||
|
|
||||||
|
class IntegrityCheckNotFound(RuntimeError):
|
||||||
|
"""Raised when a requested check id does not exist."""
|
||||||
|
|
||||||
|
|
||||||
|
class IntegrityCheckStoreError(RuntimeError):
|
||||||
|
"""Raised when check metadata or issue records cannot be read or written."""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class IntegrityIssue:
|
||||||
|
"""One stable, safe integrity discrepancy.
|
||||||
|
|
||||||
|
``issue_key`` is derived from the object and issue category instead of a
|
||||||
|
provider error string. This keeps pagination and retries deterministic
|
||||||
|
while allowing the human-readable message to evolve independently.
|
||||||
|
"""
|
||||||
|
|
||||||
|
check_id: str
|
||||||
|
issue_key: str
|
||||||
|
item_kind: str
|
||||||
|
item_key: str
|
||||||
|
issue_type: str
|
||||||
|
message: str
|
||||||
|
created_at: datetime | None = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def build(
|
||||||
|
cls,
|
||||||
|
check_id: str,
|
||||||
|
item_kind: str,
|
||||||
|
item_key: str,
|
||||||
|
issue_type: str,
|
||||||
|
message: str,
|
||||||
|
) -> IntegrityIssue:
|
||||||
|
"""Build a deterministic issue key from stable comparison identity."""
|
||||||
|
|
||||||
|
identity = "\x00".join((item_kind, item_key, issue_type)).encode("utf-8")
|
||||||
|
issue_key = sha256(identity).hexdigest()
|
||||||
|
return cls(
|
||||||
|
check_id=check_id,
|
||||||
|
issue_key=issue_key,
|
||||||
|
item_kind=item_kind,
|
||||||
|
item_key=item_key,
|
||||||
|
issue_type=issue_type,
|
||||||
|
message=" ".join(message.split())[:500],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class IntegrityCheckRun:
|
||||||
|
"""Persisted progress and terminal state for one integrity check."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
status: IntegrityStatus
|
||||||
|
window: SyncWindow | None
|
||||||
|
target_count: int
|
||||||
|
checked_count: int = 0
|
||||||
|
issue_count: int = 0
|
||||||
|
error_type: str | None = None
|
||||||
|
error_message: str | None = None
|
||||||
|
created_at: datetime | None = None
|
||||||
|
finished_at: datetime | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def window_start(self) -> date | None:
|
||||||
|
"""Return the inclusive window start for response adapters."""
|
||||||
|
|
||||||
|
return self.window.start if self.window is not None else None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def window_end(self) -> date | None:
|
||||||
|
"""Return the inclusive window end for response adapters."""
|
||||||
|
|
||||||
|
return self.window.end if self.window is not None else None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class IntegrityCheckQuery:
|
||||||
|
"""Validated issue pagination passed from HTTP or another caller."""
|
||||||
|
|
||||||
|
page: int = 1
|
||||||
|
page_size: int = 10
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
if self.page < 1:
|
||||||
|
raise ValueError("page must be at least 1")
|
||||||
|
if not 1 <= self.page_size <= 100:
|
||||||
|
raise ValueError("page_size must be between 1 and 100")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class IntegrityCheckPage:
|
||||||
|
"""One check run plus a bounded page of persisted issues."""
|
||||||
|
|
||||||
|
run: IntegrityCheckRun
|
||||||
|
page: int
|
||||||
|
page_size: int
|
||||||
|
issues_total: int
|
||||||
|
issues: tuple[IntegrityIssue, ...] = ()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def check(self) -> IntegrityCheckRun:
|
||||||
|
"""Alias used by callers that name the aggregate ``check``."""
|
||||||
|
|
||||||
|
return self.run
|
||||||
|
|
||||||
|
|
||||||
|
class IntegritySnapshotReader(Protocol):
|
||||||
|
"""Read-only PostgreSQL and CSV comparison input port."""
|
||||||
|
|
||||||
|
def latest_successful_window(self) -> SyncWindow | None: ...
|
||||||
|
|
||||||
|
def active_stocks(self) -> Sequence[Stock]: ...
|
||||||
|
|
||||||
|
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]: ...
|
||||||
|
|
||||||
|
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]: ...
|
||||||
|
|
||||||
|
def list_bar_codes(self, window: SyncWindow) -> Sequence[str]: ...
|
||||||
|
|
||||||
|
def list_daily_basic_dates(self, window: SyncWindow) -> Sequence[date]: ...
|
||||||
|
|
||||||
|
def advisory_lock(self, key: int) -> AbstractContextManager[bool]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class IntegritySnapshotStore(Protocol):
|
||||||
|
"""Read-only formal CSV snapshot port."""
|
||||||
|
|
||||||
|
def read_stocks(self) -> tuple[Stock, ...] | None: ...
|
||||||
|
|
||||||
|
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None: ...
|
||||||
|
|
||||||
|
def read_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...] | None: ...
|
||||||
|
|
||||||
|
def list_bar_codes(self, window: SyncWindow) -> Sequence[str]: ...
|
||||||
|
|
||||||
|
def list_daily_basic_dates(self, window: SyncWindow) -> Sequence[date]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class IntegrityCheckStore(Protocol):
|
||||||
|
"""Persistence port that writes only integrity metadata and issues."""
|
||||||
|
|
||||||
|
def recover_stale_running(self, lock_key: int, stale_after_seconds: int) -> int: ...
|
||||||
|
|
||||||
|
def create_running(self, window: SyncWindow, target_count: int) -> IntegrityCheckRun: ...
|
||||||
|
|
||||||
|
def update_progress(self, check_id: str, checked_count: int, issue_count: int) -> None: ...
|
||||||
|
|
||||||
|
def record_issues(self, issues: Iterable[IntegrityIssue]) -> None: ...
|
||||||
|
|
||||||
|
def finish(
|
||||||
|
self,
|
||||||
|
check_id: str,
|
||||||
|
status: IntegrityStatus,
|
||||||
|
*,
|
||||||
|
issue_count: int,
|
||||||
|
error_type: str | None = None,
|
||||||
|
error_message: str | None = None,
|
||||||
|
) -> None: ...
|
||||||
|
|
||||||
|
def get(self, check_id: str, query: IntegrityCheckQuery) -> IntegrityCheckPage | None: ...
|
||||||
|
|
||||||
|
def get_latest(self, query: IntegrityCheckQuery) -> IntegrityCheckPage | None: ...
|
||||||
@@ -46,6 +46,8 @@ def decimal_text(value: Decimal | None) -> str:
|
|||||||
|
|
||||||
if value is None:
|
if value is None:
|
||||||
return ""
|
return ""
|
||||||
|
if not value.is_finite():
|
||||||
|
raise ValueError("numeric value must be finite")
|
||||||
text = format(value, "f")
|
text = format(value, "f")
|
||||||
if "." in text:
|
if "." in text:
|
||||||
text = text.rstrip("0").rstrip(".")
|
text = text.rstrip("0").rstrip(".")
|
||||||
|
|||||||
@@ -124,3 +124,17 @@ class MarketDataRepository(Protocol):
|
|||||||
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]: ...
|
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]: ...
|
||||||
|
|
||||||
def advisory_lock(self, key: int) -> AbstractContextManager[bool]: ...
|
def advisory_lock(self, key: int) -> AbstractContextManager[bool]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class MarketDataBatchRepository(Protocol):
|
||||||
|
"""Optional batch APIs used by the optimized PostgreSQL adapter.
|
||||||
|
|
||||||
|
Keeping these methods in a refinement protocol preserves compatibility
|
||||||
|
with small in-memory repositories used by legacy application tests. The
|
||||||
|
production PostgreSQL repository implements both this protocol and
|
||||||
|
``MarketDataRepository``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def count_valid_stocks(self, trade_date: date) -> int: ...
|
||||||
|
|
||||||
|
def record_items(self, batch_id: str, outcomes: Iterable[object]) -> None: ...
|
||||||
|
|||||||
+172
-6
@@ -55,12 +55,50 @@ STOCK_COLUMNS = ("ts_code", "name", "market", "exchange", "list_status", "list_d
|
|||||||
_DAILY_DECIMAL_FIELDS = DAILY_BASIC_COLUMNS[2:]
|
_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:
|
def _code_path(ts_code: str) -> str:
|
||||||
if not _SAFE_CODE.fullmatch(ts_code):
|
if not _SAFE_CODE.fullmatch(ts_code):
|
||||||
raise ValueError(f"unsupported stock code: {ts_code!r}")
|
raise ValueError(f"unsupported stock code: {ts_code!r}")
|
||||||
return ts_code
|
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]:
|
def _bar_row(row: Bar) -> dict[str, str]:
|
||||||
return {
|
return {
|
||||||
"ts_code": row.ts_code,
|
"ts_code": row.ts_code,
|
||||||
@@ -125,19 +163,139 @@ class CsvSnapshotStore:
|
|||||||
|
|
||||||
return self.root / "stock-basic" / "current.csv"
|
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:
|
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None:
|
||||||
"""Read the current formal snapshot, returning ``None`` if absent."""
|
"""Read the current formal snapshot, returning ``None`` if absent."""
|
||||||
|
|
||||||
path = self.bars_path(ts_code)
|
path = self.bars_path(ts_code)
|
||||||
if not path.exists():
|
if not path.exists():
|
||||||
return None
|
return None
|
||||||
with path.open(newline="", encoding="utf-8") as handle:
|
try:
|
||||||
reader = csv.DictReader(handle)
|
with path.open(newline="", encoding="utf-8") as handle:
|
||||||
if tuple(reader.fieldnames or ()) != BAR_COLUMNS:
|
reader = csv.DictReader(handle)
|
||||||
raise ValueError(f"unexpected bar CSV header: {path}")
|
_require_header(reader.fieldnames, BAR_COLUMNS)
|
||||||
rows = tuple(Bar.from_mapping(row) for row in reader)
|
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))
|
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:
|
def stage_bars(self, ts_code: str, rows: Iterable[Bar]) -> StagedSnapshot:
|
||||||
"""Validate and write a temporary bar snapshot."""
|
"""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")
|
raise ValueError("bar snapshot must be non-empty and belong to one stock")
|
||||||
final_path = self.bars_path(ts_code)
|
final_path = self.bars_path(ts_code)
|
||||||
temp_path = self._write_csv(final_path, BAR_COLUMNS, (_bar_row(row) for row in normalized))
|
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:
|
def stage_daily_basic(self, trade_date: date, rows: Iterable[DailyBasic]) -> StagedSnapshot:
|
||||||
"""Validate and stage one daily-basic date snapshot."""
|
"""Validate and stage one daily-basic date snapshot."""
|
||||||
|
|||||||
@@ -0,0 +1,422 @@
|
|||||||
|
"""PostgreSQL and CSV seams used by the read-only integrity application."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Generator, Iterable
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from datetime import date
|
||||||
|
from typing import Any, Protocol, cast
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
|
||||||
|
from ..domain.integrity import (
|
||||||
|
IntegrityCheckInProgress,
|
||||||
|
IntegrityCheckPage,
|
||||||
|
IntegrityCheckQuery,
|
||||||
|
IntegrityCheckRun,
|
||||||
|
IntegrityCheckStoreError,
|
||||||
|
IntegrityIssue,
|
||||||
|
IntegrityStatus,
|
||||||
|
)
|
||||||
|
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||||
|
from ..domain.ports import MarketDataRepositoryError
|
||||||
|
|
||||||
|
|
||||||
|
class _PostgresRepository(Protocol):
|
||||||
|
"""Small structural seam shared with the existing pooled repository."""
|
||||||
|
|
||||||
|
def connection(self) -> Any: ...
|
||||||
|
|
||||||
|
def open(self) -> None: ...
|
||||||
|
|
||||||
|
def close(self) -> None: ...
|
||||||
|
|
||||||
|
def advisory_lock(self, key: int) -> Any: ...
|
||||||
|
|
||||||
|
def latest_successful_window(self) -> SyncWindow | None: ...
|
||||||
|
|
||||||
|
def active_stocks(self) -> tuple[Stock, ...]: ...
|
||||||
|
|
||||||
|
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]: ...
|
||||||
|
|
||||||
|
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]: ...
|
||||||
|
|
||||||
|
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]: ...
|
||||||
|
|
||||||
|
def list_invalid_bar_codes(self, window: SyncWindow) -> tuple[str, ...]: ...
|
||||||
|
|
||||||
|
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class PostgresIntegritySnapshotReader:
|
||||||
|
"""Expose the existing repository's read-only integrity operations."""
|
||||||
|
|
||||||
|
def __init__(self, repository: _PostgresRepository | str, *, max_connections: int = 10) -> None:
|
||||||
|
if isinstance(repository, str):
|
||||||
|
# Keep the import lazy: ``postgres.py`` re-exports this adapter for
|
||||||
|
# compatibility, so importing it at module load time would cycle.
|
||||||
|
from .postgres import PostgresMarketDataRepository
|
||||||
|
|
||||||
|
self.repository: _PostgresRepository = PostgresMarketDataRepository(
|
||||||
|
repository,
|
||||||
|
max_connections=max_connections,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.repository = repository
|
||||||
|
|
||||||
|
def open(self) -> None:
|
||||||
|
"""Open the underlying pool when this reader owns a URL-backed one."""
|
||||||
|
|
||||||
|
self.repository.open()
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
"""Close the underlying pool when this reader owns a URL-backed one."""
|
||||||
|
|
||||||
|
self.repository.close()
|
||||||
|
|
||||||
|
def latest_successful_window(self) -> SyncWindow | None:
|
||||||
|
"""Return the newest completed synchronization window."""
|
||||||
|
|
||||||
|
return self.repository.latest_successful_window()
|
||||||
|
|
||||||
|
def active_stocks(self) -> tuple[Stock, ...]:
|
||||||
|
"""Return the active stock master in deterministic order."""
|
||||||
|
|
||||||
|
return self.repository.active_stocks()
|
||||||
|
|
||||||
|
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]:
|
||||||
|
"""Return a streaming bar iterator."""
|
||||||
|
|
||||||
|
return self.repository.iter_bars(window)
|
||||||
|
|
||||||
|
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]:
|
||||||
|
"""Return a streaming daily-basic iterator."""
|
||||||
|
|
||||||
|
return self.repository.iter_daily_basic(window)
|
||||||
|
|
||||||
|
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||||
|
"""Return bar group keys without reading their rows."""
|
||||||
|
|
||||||
|
return self.repository.list_bar_codes(window)
|
||||||
|
|
||||||
|
def list_invalid_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||||
|
"""Return bar groups whose stored adjustment source is not qfq."""
|
||||||
|
|
||||||
|
return self.repository.list_invalid_bar_codes(window)
|
||||||
|
|
||||||
|
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
|
||||||
|
"""Return daily-basic group keys without reading their rows."""
|
||||||
|
|
||||||
|
return self.repository.list_daily_basic_dates(window)
|
||||||
|
|
||||||
|
def advisory_lock(self, key: int) -> Any:
|
||||||
|
"""Use exactly the same connection-bound advisory lock as sync."""
|
||||||
|
|
||||||
|
return self.repository.advisory_lock(key)
|
||||||
|
|
||||||
|
|
||||||
|
class PostgresIntegrityCheckStore:
|
||||||
|
"""Persist integrity progress and issues through the shared pool.
|
||||||
|
|
||||||
|
The store deliberately has no methods that write market facts. Its only
|
||||||
|
mutation surface is the two ``market_integrity_*`` tables.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, repository: _PostgresRepository | str, *, max_connections: int = 10) -> None:
|
||||||
|
if isinstance(repository, str):
|
||||||
|
from .postgres import PostgresMarketDataRepository
|
||||||
|
|
||||||
|
self.repository: _PostgresRepository = PostgresMarketDataRepository(
|
||||||
|
repository,
|
||||||
|
max_connections=max_connections,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.repository = repository
|
||||||
|
|
||||||
|
def open(self) -> None:
|
||||||
|
"""Open the underlying pool."""
|
||||||
|
|
||||||
|
self.repository.open()
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
"""Close the underlying pool."""
|
||||||
|
|
||||||
|
self.repository.close()
|
||||||
|
|
||||||
|
def recover_stale_running(self, lock_key: int, stale_after_seconds: int) -> int:
|
||||||
|
"""Mark old workers failed only while the shared lock is available."""
|
||||||
|
|
||||||
|
if stale_after_seconds < 1:
|
||||||
|
raise ValueError("stale_after_seconds must be positive")
|
||||||
|
try:
|
||||||
|
with self._connection() as connection:
|
||||||
|
acquired = False
|
||||||
|
count = 0
|
||||||
|
try:
|
||||||
|
with connection.transaction():
|
||||||
|
acquired_row = connection.execute(
|
||||||
|
"SELECT pg_try_advisory_lock(%s)",
|
||||||
|
(lock_key,),
|
||||||
|
).fetchone()
|
||||||
|
acquired = bool(acquired_row[0]) if acquired_row is not None else False
|
||||||
|
if acquired:
|
||||||
|
result = connection.execute(
|
||||||
|
"""
|
||||||
|
UPDATE market_integrity_check
|
||||||
|
SET status = 'failed',
|
||||||
|
error_type = 'stale_worker',
|
||||||
|
error_message = 'integrity check worker became stale',
|
||||||
|
updated_at = now(),
|
||||||
|
finished_at = now()
|
||||||
|
WHERE status = 'running'
|
||||||
|
AND updated_at < now() - (%s * interval '1 second')
|
||||||
|
""",
|
||||||
|
(stale_after_seconds,),
|
||||||
|
)
|
||||||
|
count = int(result.rowcount)
|
||||||
|
finally:
|
||||||
|
if acquired:
|
||||||
|
connection.execute("SELECT pg_advisory_unlock(%s)", (lock_key,))
|
||||||
|
except MarketDataRepositoryError as exc:
|
||||||
|
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||||
|
return count
|
||||||
|
|
||||||
|
def create_running(self, window: SyncWindow, target_count: int) -> IntegrityCheckRun:
|
||||||
|
"""Atomically claim the single running-check slot."""
|
||||||
|
|
||||||
|
check_id = str(uuid4())
|
||||||
|
try:
|
||||||
|
with self._connection() as connection, connection.transaction():
|
||||||
|
row = connection.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO market_integrity_check
|
||||||
|
(id, status, window_start, window_end, target_count)
|
||||||
|
VALUES (%s, 'running', %s, %s, %s)
|
||||||
|
RETURNING id, status, window_start, window_end, target_count,
|
||||||
|
checked_count, issue_count, error_type, error_message,
|
||||||
|
created_at, finished_at
|
||||||
|
""",
|
||||||
|
(check_id, window.start, window.end, target_count),
|
||||||
|
).fetchone()
|
||||||
|
except MarketDataRepositoryError as exc:
|
||||||
|
if isinstance(exc.__cause__, psycopg.errors.UniqueViolation):
|
||||||
|
raise IntegrityCheckInProgress("another integrity check is running") from exc
|
||||||
|
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||||
|
if row is None:
|
||||||
|
raise IntegrityCheckStoreError("integrity check claim returned no row")
|
||||||
|
return _check_from_row(row)
|
||||||
|
|
||||||
|
def update_progress(self, check_id: str, checked_count: int, issue_count: int) -> None:
|
||||||
|
"""Persist heartbeat and counters after each comparison group."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
with self._connection() as connection, connection.transaction():
|
||||||
|
connection.execute(
|
||||||
|
"""
|
||||||
|
UPDATE market_integrity_check
|
||||||
|
SET checked_count = %s, issue_count = %s, updated_at = now()
|
||||||
|
WHERE id = %s AND status = 'running'
|
||||||
|
""",
|
||||||
|
(checked_count, issue_count, check_id),
|
||||||
|
)
|
||||||
|
except MarketDataRepositoryError as exc:
|
||||||
|
raise IntegrityCheckStoreError("integrity check progress cannot be saved") from exc
|
||||||
|
|
||||||
|
def record_issues(self, issues: Iterable[IntegrityIssue]) -> None:
|
||||||
|
"""Insert a bounded issue batch idempotently."""
|
||||||
|
|
||||||
|
records = tuple(issues)
|
||||||
|
if not records:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
with self._connection() as connection, connection.transaction():
|
||||||
|
connection.executemany(
|
||||||
|
"""
|
||||||
|
INSERT INTO market_integrity_issue
|
||||||
|
(check_id, issue_key, item_kind, item_key, issue_type, message)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s)
|
||||||
|
ON CONFLICT (check_id, issue_key) DO UPDATE SET
|
||||||
|
message = EXCLUDED.message
|
||||||
|
""",
|
||||||
|
[
|
||||||
|
(
|
||||||
|
issue.check_id,
|
||||||
|
issue.issue_key,
|
||||||
|
issue.item_kind,
|
||||||
|
issue.item_key,
|
||||||
|
issue.issue_type,
|
||||||
|
_safe_message(issue.message),
|
||||||
|
)
|
||||||
|
for issue in records
|
||||||
|
],
|
||||||
|
)
|
||||||
|
except MarketDataRepositoryError as exc:
|
||||||
|
raise IntegrityCheckStoreError("integrity issues cannot be saved") from exc
|
||||||
|
|
||||||
|
def finish(
|
||||||
|
self,
|
||||||
|
check_id: str,
|
||||||
|
status: IntegrityStatus,
|
||||||
|
*,
|
||||||
|
issue_count: int,
|
||||||
|
error_type: str | None = None,
|
||||||
|
error_message: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Converge one running check to a terminal state."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
with self._connection() as connection, connection.transaction():
|
||||||
|
connection.execute(
|
||||||
|
"""
|
||||||
|
UPDATE market_integrity_check
|
||||||
|
SET status = %s, issue_count = %s, error_type = %s,
|
||||||
|
error_message = %s, updated_at = now(), finished_at = now()
|
||||||
|
WHERE id = %s
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
status,
|
||||||
|
issue_count,
|
||||||
|
error_type,
|
||||||
|
_safe_message(error_message),
|
||||||
|
check_id,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
except MarketDataRepositoryError as exc:
|
||||||
|
raise IntegrityCheckStoreError("integrity check result cannot be saved") from exc
|
||||||
|
|
||||||
|
def get(self, check_id: str, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
|
||||||
|
"""Read one check and a stable issue page."""
|
||||||
|
|
||||||
|
return self._read_page("WHERE id = %s", (check_id,), query)
|
||||||
|
|
||||||
|
def get_latest(self, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
|
||||||
|
"""Read the newest check by creation time."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
with self._connection() as connection:
|
||||||
|
row = connection.execute(
|
||||||
|
"""
|
||||||
|
SELECT id
|
||||||
|
FROM market_integrity_check
|
||||||
|
ORDER BY created_at DESC, id DESC
|
||||||
|
LIMIT 1
|
||||||
|
"""
|
||||||
|
).fetchone()
|
||||||
|
except MarketDataRepositoryError as exc:
|
||||||
|
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||||
|
if row is None:
|
||||||
|
return None
|
||||||
|
return self.get(str(row[0]), query)
|
||||||
|
|
||||||
|
def _read_page(
|
||||||
|
self,
|
||||||
|
predicate: str,
|
||||||
|
parameters: tuple[object, ...],
|
||||||
|
query: IntegrityCheckQuery,
|
||||||
|
) -> IntegrityCheckPage | None:
|
||||||
|
offset = (query.page - 1) * query.page_size
|
||||||
|
try:
|
||||||
|
with self._connection() as connection:
|
||||||
|
row = connection.execute(
|
||||||
|
f"""
|
||||||
|
SELECT id, status, window_start, window_end, target_count,
|
||||||
|
checked_count, issue_count, error_type, error_message,
|
||||||
|
created_at, finished_at
|
||||||
|
FROM market_integrity_check
|
||||||
|
{predicate}
|
||||||
|
""",
|
||||||
|
parameters,
|
||||||
|
).fetchone()
|
||||||
|
if row is None:
|
||||||
|
return None
|
||||||
|
total_row = connection.execute(
|
||||||
|
"SELECT count(*) FROM market_integrity_issue WHERE check_id = %s",
|
||||||
|
(row[0],),
|
||||||
|
).fetchone()
|
||||||
|
issue_rows = connection.execute(
|
||||||
|
"""
|
||||||
|
SELECT check_id, issue_key, item_kind, item_key, issue_type,
|
||||||
|
message, created_at
|
||||||
|
FROM market_integrity_issue
|
||||||
|
WHERE check_id = %s
|
||||||
|
ORDER BY issue_key
|
||||||
|
LIMIT %s OFFSET %s
|
||||||
|
""",
|
||||||
|
(row[0], query.page_size, offset),
|
||||||
|
).fetchall()
|
||||||
|
except MarketDataRepositoryError as exc:
|
||||||
|
raise IntegrityCheckStoreError("integrity check storage is unavailable") from exc
|
||||||
|
return IntegrityCheckPage(
|
||||||
|
run=_check_from_row(row),
|
||||||
|
page=query.page,
|
||||||
|
page_size=query.page_size,
|
||||||
|
issues_total=int(total_row[0]) if total_row is not None else 0,
|
||||||
|
issues=tuple(_issue_from_row(issue_row) for issue_row in issue_rows),
|
||||||
|
)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _connection(self) -> Generator[Any, None, None]:
|
||||||
|
"""Borrow a pooled connection and translate driver errors safely."""
|
||||||
|
|
||||||
|
try:
|
||||||
|
self.repository.open()
|
||||||
|
with self.repository.connection() as connection:
|
||||||
|
yield connection
|
||||||
|
except psycopg.Error as exc:
|
||||||
|
raise MarketDataRepositoryError("market data database operation failed") from exc
|
||||||
|
|
||||||
|
|
||||||
|
def _check_from_row(row: tuple[Any, ...]) -> IntegrityCheckRun:
|
||||||
|
"""Translate one persistence row into the domain run model."""
|
||||||
|
|
||||||
|
window = None
|
||||||
|
if row[2] is not None and row[3] is not None:
|
||||||
|
window = SyncWindow(start=row[2], end=row[3])
|
||||||
|
status = cast(IntegrityStatus, str(row[1]))
|
||||||
|
return IntegrityCheckRun(
|
||||||
|
id=str(row[0]),
|
||||||
|
status=status,
|
||||||
|
window=window,
|
||||||
|
target_count=int(row[4]),
|
||||||
|
checked_count=int(row[5]),
|
||||||
|
issue_count=int(row[6]),
|
||||||
|
error_type=str(row[7]) if row[7] is not None else None,
|
||||||
|
error_message=str(row[8]) if row[8] is not None else None,
|
||||||
|
created_at=row[9],
|
||||||
|
finished_at=row[10],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _issue_from_row(row: tuple[Any, ...]) -> IntegrityIssue:
|
||||||
|
"""Translate one persistence row into a domain issue."""
|
||||||
|
|
||||||
|
return IntegrityIssue(
|
||||||
|
check_id=str(row[0]),
|
||||||
|
issue_key=str(row[1]),
|
||||||
|
item_kind=str(row[2]),
|
||||||
|
item_key=str(row[3]),
|
||||||
|
issue_type=str(row[4]),
|
||||||
|
message=str(row[5]),
|
||||||
|
created_at=row[6],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_message(message: str | None) -> str | None:
|
||||||
|
"""Persist bounded, whitespace-normalized messages only."""
|
||||||
|
|
||||||
|
if message is None:
|
||||||
|
return None
|
||||||
|
return " ".join(message.split())[:500]
|
||||||
|
|
||||||
|
|
||||||
|
# This name reads naturally at the application seam and keeps an intuitive
|
||||||
|
# compatibility alias for callers that call every persistence adapter a repo.
|
||||||
|
PostgresIntegrityCheckRepository = PostgresIntegrityCheckStore
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"PostgresIntegrityCheckRepository",
|
||||||
|
"PostgresIntegrityCheckStore",
|
||||||
|
"PostgresIntegritySnapshotReader",
|
||||||
|
]
|
||||||
@@ -2,14 +2,15 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from collections.abc import Generator, Iterable
|
import threading
|
||||||
|
from collections.abc import Generator, Iterable, Iterator
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
import psycopg
|
from psycopg_pool import ConnectionPool
|
||||||
|
|
||||||
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||||
from ..domain.overview import (
|
from ..domain.overview import (
|
||||||
@@ -25,12 +26,69 @@ from ..domain.ports import MarketDataRepositoryError, WriteResult
|
|||||||
class PostgresMarketDataRepository:
|
class PostgresMarketDataRepository:
|
||||||
"""Persist market data without an ORM identity map.
|
"""Persist market data without an ORM identity map.
|
||||||
|
|
||||||
Each public write opens a short transaction. The caller publishes the
|
Each public write borrows an independent connection from a bounded pool
|
||||||
matching CSV only after this method returns successfully.
|
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.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:
|
def read_overview(self) -> MarketDataOverview:
|
||||||
"""Aggregate the active stock pool and latest synchronization facts.
|
"""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:
|
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool:
|
||||||
return self._exists("market_daily_basic", ts_code, trade_date)
|
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(
|
def create_batch(
|
||||||
self,
|
self,
|
||||||
target_trade_date: date,
|
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], ...]:
|
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]:
|
||||||
with self._connection() as connection:
|
with self._connection() as connection:
|
||||||
rows = connection.execute(
|
rows = connection.execute(
|
||||||
@@ -490,12 +757,13 @@ class PostgresMarketDataRepository:
|
|||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def _connection(self) -> Generator[Any, None, None]:
|
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:
|
try:
|
||||||
with psycopg.connect(self.database_url) as connection:
|
self.open()
|
||||||
|
with self.pool.connection() as connection:
|
||||||
yield 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
|
raise MarketDataRepositoryError("market data database operation failed") from exc
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -610,11 +878,11 @@ class PostgresMarketDataRepository:
|
|||||||
market_daily_bar.open, market_daily_bar.high, market_daily_bar.low,
|
market_daily_bar.open, market_daily_bar.high, market_daily_bar.low,
|
||||||
market_daily_bar.close, market_daily_bar.pre_close,
|
market_daily_bar.close, market_daily_bar.pre_close,
|
||||||
market_daily_bar.change, market_daily_bar.pct_chg,
|
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 (
|
) IS DISTINCT FROM (
|
||||||
EXCLUDED.open, EXCLUDED.high, EXCLUDED.low, EXCLUDED.close,
|
EXCLUDED.open, EXCLUDED.high, EXCLUDED.low, EXCLUDED.close,
|
||||||
EXCLUDED.pre_close, EXCLUDED.change, EXCLUDED.pct_chg,
|
EXCLUDED.pre_close, EXCLUDED.change, EXCLUDED.pct_chg,
|
||||||
EXCLUDED.vol, EXCLUDED.amount
|
EXCLUDED.vol, EXCLUDED.amount, 'qfq'
|
||||||
)
|
)
|
||||||
RETURNING (xmax = 0) AS inserted
|
RETURNING (xmax = 0) AS inserted
|
||||||
"""
|
"""
|
||||||
@@ -673,3 +941,77 @@ class PostgresMarketDataRepository:
|
|||||||
if message is None:
|
if message is None:
|
||||||
return None
|
return None
|
||||||
return " ".join(message.split())[:500]
|
return " ".join(message.split())[:500]
|
||||||
|
|
||||||
|
|
||||||
|
def _database_decimal(value: object) -> Decimal | None:
|
||||||
|
"""Normalize a PostgreSQL numeric value for the domain model."""
|
||||||
|
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
result = Decimal(str(value))
|
||||||
|
if not result.is_finite():
|
||||||
|
raise ValueError("database numeric value must be finite")
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _bar_from_database_row(row: tuple[Any, ...]) -> Bar:
|
||||||
|
"""Translate one ordered PostgreSQL bar row into a domain value."""
|
||||||
|
|
||||||
|
return Bar(
|
||||||
|
ts_code=str(row[0]),
|
||||||
|
trade_date=row[1],
|
||||||
|
open=_database_decimal(row[2]),
|
||||||
|
high=_database_decimal(row[3]),
|
||||||
|
low=_database_decimal(row[4]),
|
||||||
|
close=_database_decimal(row[5]),
|
||||||
|
pre_close=_database_decimal(row[6]),
|
||||||
|
change=_database_decimal(row[7]),
|
||||||
|
pct_chg=_database_decimal(row[8]),
|
||||||
|
vol=_database_decimal(row[9]),
|
||||||
|
amount=_database_decimal(row[10]),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _daily_basic_from_database_row(row: tuple[Any, ...]) -> DailyBasic:
|
||||||
|
"""Translate one ordered PostgreSQL daily-basic row into a domain value."""
|
||||||
|
|
||||||
|
fields = (
|
||||||
|
"close",
|
||||||
|
"turnover_rate",
|
||||||
|
"turnover_rate_f",
|
||||||
|
"volume_ratio",
|
||||||
|
"pe",
|
||||||
|
"pe_ttm",
|
||||||
|
"pb",
|
||||||
|
"ps",
|
||||||
|
"ps_ttm",
|
||||||
|
"dv_ratio",
|
||||||
|
"dv_ttm",
|
||||||
|
"total_share",
|
||||||
|
"float_share",
|
||||||
|
"free_share",
|
||||||
|
"total_mv",
|
||||||
|
"circ_mv",
|
||||||
|
)
|
||||||
|
values = {field: _database_decimal(row[index + 2]) for index, field in enumerate(fields)}
|
||||||
|
return DailyBasic(
|
||||||
|
ts_code=str(row[0]),
|
||||||
|
trade_date=row[1],
|
||||||
|
**values,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Re-export the read-only integrity adapters from the established PostgreSQL
|
||||||
|
# module so existing infrastructure import paths remain discoverable.
|
||||||
|
from .integrity import ( # noqa: E402
|
||||||
|
PostgresIntegrityCheckRepository,
|
||||||
|
PostgresIntegrityCheckStore,
|
||||||
|
PostgresIntegritySnapshotReader,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"PostgresIntegrityCheckRepository",
|
||||||
|
"PostgresIntegrityCheckStore",
|
||||||
|
"PostgresIntegritySnapshotReader",
|
||||||
|
"PostgresMarketDataRepository",
|
||||||
|
]
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from sqlalchemy import (
|
|||||||
Text,
|
Text,
|
||||||
UniqueConstraint,
|
UniqueConstraint,
|
||||||
func,
|
func,
|
||||||
|
text,
|
||||||
)
|
)
|
||||||
from sqlalchemy.dialects.postgresql import JSONB
|
from sqlalchemy.dialects.postgresql import JSONB
|
||||||
|
|
||||||
@@ -173,6 +174,41 @@ selection_signal = Table(
|
|||||||
PrimaryKeyConstraint("run_id", "ts_code", "category"),
|
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
|
# Keep the declarative metadata aligned with the indexes created by the
|
||||||
# Alembic revisions. Alembic uses this object for both offline inspection
|
# Alembic revisions. Alembic uses this object for both offline inspection
|
||||||
# and future autogeneration, so omitting these indexes would make the schema
|
# 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.target_trade_date,
|
||||||
selection_signal.c.ts_code,
|
selection_signal.c.ts_code,
|
||||||
)
|
)
|
||||||
|
Index(
|
||||||
|
"ix_market_integrity_check_status_created_at",
|
||||||
|
market_integrity_check.c.status,
|
||||||
|
market_integrity_check.c.created_at,
|
||||||
|
)
|
||||||
|
Index(
|
||||||
|
"uq_market_integrity_check_running",
|
||||||
|
market_integrity_check.c.status,
|
||||||
|
unique=True,
|
||||||
|
postgresql_where=text("status = 'running'"),
|
||||||
|
)
|
||||||
|
Index("ix_market_integrity_issue_check_id", market_integrity_issue.c.check_id)
|
||||||
|
|||||||
@@ -1,11 +1,12 @@
|
|||||||
"""Tushare source adapter and current-universe filtering."""
|
"""Tushare source adapter, request coordination, and current-universe filtering."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import random
|
import random
|
||||||
|
import threading
|
||||||
import time
|
import time
|
||||||
from collections.abc import Callable, Iterable, Mapping
|
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||||
from datetime import date
|
from datetime import date
|
||||||
from typing import cast
|
from typing import cast
|
||||||
|
|
||||||
@@ -14,18 +15,226 @@ from ..domain.rules import filter_current_hs_a_stocks
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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):
|
class TushareSourceError(RuntimeError):
|
||||||
"""A vendor request failed after the configured retry budget."""
|
"""A vendor request failed after the configured retry budget."""
|
||||||
|
|
||||||
|
|
||||||
class TushareAdapter:
|
class RequestCoordinator:
|
||||||
"""Translate Tushare SDK responses into domain records.
|
"""Coordinate retry and shared rate-limit cooling for one token client.
|
||||||
|
|
||||||
The SDK is kept behind this adapter so ordinary domain/application tests
|
Normal requests are deliberately not serialized. Only a provider rate
|
||||||
can inject a tiny fake client and never need a network token.
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
client: object,
|
client: object,
|
||||||
@@ -36,14 +245,32 @@ class TushareAdapter:
|
|||||||
request_interval_seconds: float = 0.2,
|
request_interval_seconds: float = 0.2,
|
||||||
random_fn: Callable[[], float] = random.random,
|
random_fn: Callable[[], float] = random.random,
|
||||||
sleep_fn: Callable[[float], None] = time.sleep,
|
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:
|
) -> None:
|
||||||
self.client = client
|
self.client = client
|
||||||
self.pro_bar = pro_bar
|
|
||||||
self.max_retries = max(0, max_retries)
|
self.max_retries = max(0, max_retries)
|
||||||
self.backoff_seconds = max(0.0, backoff_seconds)
|
self.backoff_seconds = max(0.0, backoff_seconds)
|
||||||
self.request_interval_seconds = max(0.0, request_interval_seconds)
|
self.request_interval_seconds = max(0.0, request_interval_seconds)
|
||||||
self.random_fn = random_fn
|
self.random_fn = random_fn
|
||||||
self.sleep_fn = sleep_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
|
@classmethod
|
||||||
def from_token(
|
def from_token(
|
||||||
@@ -53,6 +280,7 @@ class TushareAdapter:
|
|||||||
max_retries: int = 3,
|
max_retries: int = 3,
|
||||||
backoff_seconds: float = 1.0,
|
backoff_seconds: float = 1.0,
|
||||||
request_interval_seconds: float = 0.2,
|
request_interval_seconds: float = 0.2,
|
||||||
|
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||||
) -> TushareAdapter:
|
) -> TushareAdapter:
|
||||||
"""Create a production adapter from a token without exposing it."""
|
"""Create a production adapter from a token without exposing it."""
|
||||||
|
|
||||||
@@ -61,13 +289,22 @@ class TushareAdapter:
|
|||||||
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
||||||
|
|
||||||
client = cast(object, ts.pro_api(token))
|
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(
|
pro_bar = cast(
|
||||||
Callable[..., object],
|
Callable[..., object],
|
||||||
ts.pro_bar, # pyright: ignore[reportUnknownMemberType]
|
ts.pro_bar, # pyright: ignore[reportUnknownMemberType]
|
||||||
)
|
)
|
||||||
|
|
||||||
def pro_bar_with_client(**kwargs: object) -> object:
|
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(
|
return cls(
|
||||||
client,
|
client,
|
||||||
@@ -75,6 +312,9 @@ class TushareAdapter:
|
|||||||
max_retries=max_retries,
|
max_retries=max_retries,
|
||||||
backoff_seconds=backoff_seconds,
|
backoff_seconds=backoff_seconds,
|
||||||
request_interval_seconds=request_interval_seconds,
|
request_interval_seconds=request_interval_seconds,
|
||||||
|
request_coordinator=coordinator,
|
||||||
|
coordinated_client=coordinated_client,
|
||||||
|
pro_bar_coordinated=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
def fetch_stocks(self) -> tuple[Stock, ...]:
|
def fetch_stocks(self) -> tuple[Stock, ...]:
|
||||||
@@ -137,43 +377,29 @@ class TushareAdapter:
|
|||||||
return tuple(sorted(metrics, key=lambda row: row.ts_code))
|
return tuple(sorted(metrics, key=lambda row: row.ts_code))
|
||||||
|
|
||||||
def _records(self, method_name: str, **kwargs: object) -> tuple[Mapping[str, object], ...]:
|
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:
|
def request() -> object:
|
||||||
if method_name == "pro_bar":
|
if method_name == "pro_bar":
|
||||||
if self.pro_bar is not None:
|
if self.pro_bar is not None:
|
||||||
return self.pro_bar(**kwargs)
|
return self.pro_bar(**kwargs)
|
||||||
method = getattr(self.client, method_name, None)
|
method = getattr(self.coordinated_client, method_name, None)
|
||||||
else:
|
else:
|
||||||
method = getattr(self.client, method_name, None)
|
method = getattr(self.coordinated_client, method_name, None)
|
||||||
if not callable(method):
|
if not callable(method):
|
||||||
raise TypeError(f"Tushare client has no callable {method_name}")
|
raise TypeError(f"Tushare client has no callable {method_name}")
|
||||||
return method(**kwargs)
|
return method(**kwargs)
|
||||||
|
|
||||||
last_error: BaseException | None = None
|
# ``pro_bar`` itself is a qfq composition helper; its nested daily and
|
||||||
for attempt in range(self.max_retries + 1):
|
# adj_factor methods are bound to the coordinator. Wrapping the helper
|
||||||
try:
|
# as a second retry layer would hide the useful error classification.
|
||||||
result = request()
|
result = (
|
||||||
self.sleep_fn(self.request_interval_seconds)
|
request()
|
||||||
return self._as_records(result)
|
if method_name == "pro_bar" and self.pro_bar_coordinated
|
||||||
except (OSError, RuntimeError, TimeoutError) as exc:
|
else self.request_coordinator.call(method_name, request)
|
||||||
last_error = exc
|
|
||||||
if attempt == self.max_retries:
|
|
||||||
break
|
|
||||||
delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn())
|
|
||||||
logger.warning(
|
|
||||||
"tushare_request_retry method=%s attempt=%d max_attempts=%d",
|
|
||||||
method_name,
|
|
||||||
attempt + 1,
|
|
||||||
self.max_retries + 1,
|
|
||||||
)
|
|
||||||
self.sleep_fn(delay)
|
|
||||||
logger.error(
|
|
||||||
"tushare_request_failed method=%s attempts=%d",
|
|
||||||
method_name,
|
|
||||||
self.max_retries + 1,
|
|
||||||
)
|
)
|
||||||
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
|
self.sleep_fn(self.request_interval_seconds)
|
||||||
|
return self._as_records(result)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _as_records(result: object) -> tuple[Mapping[str, object], ...]:
|
def _as_records(result: object) -> tuple[Mapping[str, object], ...]:
|
||||||
|
|||||||
@@ -70,14 +70,24 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||||||
backoff_seconds=settings.market_data_retry_backoff_seconds,
|
backoff_seconds=settings.market_data_retry_backoff_seconds,
|
||||||
request_interval_seconds=settings.market_data_request_interval_seconds,
|
request_interval_seconds=settings.market_data_request_interval_seconds,
|
||||||
)
|
)
|
||||||
use_case = SyncMarketData(
|
repository = PostgresMarketDataRepository(
|
||||||
source,
|
settings.database_url,
|
||||||
CsvSnapshotStore(settings.market_data_csv_root),
|
# Keep one control/advisory connection and one main-thread connection
|
||||||
PostgresMarketDataRepository(settings.database_url),
|
# available in addition to the worker connections.
|
||||||
coverage_threshold=settings.market_data_coverage_threshold,
|
max_connections=settings.market_data_max_workers + 2,
|
||||||
lock_key=settings.market_data_advisory_lock_key,
|
|
||||||
)
|
)
|
||||||
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))
|
print(json.dumps(summary.as_dict(), ensure_ascii=False, sort_keys=True))
|
||||||
return summary.exit_code
|
return summary.exit_code
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
"""HTTP presentation for the Home market-data overview."""
|
"""HTTP presentation for the Home market-data overview."""
|
||||||
|
|
||||||
|
import atexit
|
||||||
|
import threading
|
||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
from typing import Annotated, Literal
|
from typing import Annotated, Literal
|
||||||
from zoneinfo import ZoneInfo
|
from zoneinfo import ZoneInfo
|
||||||
@@ -24,6 +26,8 @@ from zhixing_server.modules.market_data.infrastructure.postgres import (
|
|||||||
|
|
||||||
home_router = APIRouter()
|
home_router = APIRouter()
|
||||||
_SHANGHAI = ZoneInfo("Asia/Shanghai")
|
_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"]
|
OverviewStatusValue = Literal["updating", "success", "partial_success", "failed", "no_data"]
|
||||||
UpdateStatusValue = Literal["updating", "success", "partial_success", "failed"]
|
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(
|
def get_market_data_overview_reader(
|
||||||
settings: Annotated[Settings, Depends(get_settings)],
|
settings: Annotated[Settings, Depends(get_settings)],
|
||||||
) -> MarketDataOverviewReader:
|
) -> 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)
|
@home_router.get("/overview", response_model=HomeOverviewResponse)
|
||||||
|
|||||||
@@ -33,6 +33,8 @@ def test_postgres_migration_creates_market_data_contract(
|
|||||||
"market_daily_basic",
|
"market_daily_basic",
|
||||||
"market_sync_batch",
|
"market_sync_batch",
|
||||||
"market_sync_item",
|
"market_sync_item",
|
||||||
|
"market_integrity_check",
|
||||||
|
"market_integrity_issue",
|
||||||
"selection_run",
|
"selection_run",
|
||||||
"selection_run_item",
|
"selection_run_item",
|
||||||
"selection_signal",
|
"selection_signal",
|
||||||
|
|||||||
@@ -0,0 +1,141 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
import pytest
|
||||||
|
from alembic import command
|
||||||
|
from alembic.config import Config
|
||||||
|
|
||||||
|
from zhixing_server.modules.market_data.application.sync import SyncFailure, SyncItemOutcome
|
||||||
|
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, SyncWindow
|
||||||
|
from zhixing_server.modules.market_data.domain.ports import MarketDataRepositoryError, WriteResult
|
||||||
|
from zhixing_server.modules.market_data.infrastructure.postgres import (
|
||||||
|
PostgresMarketDataRepository,
|
||||||
|
)
|
||||||
|
|
||||||
|
TEST_CODES = ("991901.SZ", "991902.SZ")
|
||||||
|
TARGET = date(2024, 1, 2)
|
||||||
|
WINDOW = SyncWindow(start=TARGET, end=TARGET)
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _database(database_url: str) -> Any:
|
||||||
|
with psycopg.connect(database_url) as connection, connection.transaction():
|
||||||
|
yield connection
|
||||||
|
|
||||||
|
|
||||||
|
def _cleanup(database_url: str) -> None:
|
||||||
|
with _database(database_url) as connection:
|
||||||
|
connection.execute(
|
||||||
|
"DELETE FROM market_sync_item WHERE batch_id IN "
|
||||||
|
"(SELECT id FROM market_sync_batch WHERE id LIKE '9919%')"
|
||||||
|
)
|
||||||
|
connection.execute("DELETE FROM market_sync_batch WHERE id LIKE '9919%'")
|
||||||
|
connection.execute(
|
||||||
|
"DELETE FROM market_daily_basic WHERE ts_code = ANY(%s)",
|
||||||
|
(list(TEST_CODES),),
|
||||||
|
)
|
||||||
|
connection.execute(
|
||||||
|
"DELETE FROM market_daily_bar WHERE ts_code = ANY(%s)",
|
||||||
|
(list(TEST_CODES),),
|
||||||
|
)
|
||||||
|
connection.execute(
|
||||||
|
"DELETE FROM market_stock WHERE ts_code = ANY(%s)",
|
||||||
|
(list(TEST_CODES),),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _prepare_database(database_url: str) -> None:
|
||||||
|
server_root = Path(__file__).parents[2]
|
||||||
|
config = Config(str(server_root / "alembic.ini"))
|
||||||
|
config.set_main_option(
|
||||||
|
"sqlalchemy.url",
|
||||||
|
database_url.replace("%", "%%").replace("postgresql://", "postgresql+psycopg://"),
|
||||||
|
)
|
||||||
|
command.upgrade(config, "head")
|
||||||
|
_cleanup(database_url)
|
||||||
|
with _database(database_url) as connection:
|
||||||
|
connection.executemany(
|
||||||
|
"""
|
||||||
|
INSERT INTO market_stock
|
||||||
|
(ts_code, name, market, exchange, list_status, is_active)
|
||||||
|
VALUES (%s, %s, '主板', 'SZSE', 'L', true)
|
||||||
|
""",
|
||||||
|
[(code, f"测试{code}") for code in TEST_CODES],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_pool_upsert_rollback_batch_audit_and_set_coverage() -> None:
|
||||||
|
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
|
||||||
|
if not database_url:
|
||||||
|
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
|
||||||
|
_prepare_database(database_url)
|
||||||
|
repository = PostgresMarketDataRepository(database_url, max_connections=4)
|
||||||
|
batch_id: str | None = None
|
||||||
|
try:
|
||||||
|
bars = {code: Bar(code, TARGET, close=Decimal("10")) for code in TEST_CODES}
|
||||||
|
with repository:
|
||||||
|
|
||||||
|
def upsert(code: str) -> WriteResult:
|
||||||
|
return repository.upsert_bars((bars[code],), WINDOW, full_snapshot=True)
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||||
|
results = tuple(executor.map(upsert, TEST_CODES))
|
||||||
|
assert [result.inserted for result in results] == [1, 1]
|
||||||
|
|
||||||
|
with pytest.raises(MarketDataRepositoryError):
|
||||||
|
repository.upsert_bars(
|
||||||
|
(
|
||||||
|
Bar(TEST_CODES[0], TARGET, close=Decimal("11")),
|
||||||
|
Bar(TEST_CODES[0], TARGET, close=Decimal("12")),
|
||||||
|
),
|
||||||
|
WINDOW,
|
||||||
|
full_snapshot=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
unchanged = repository.upsert_daily_basic(
|
||||||
|
(DailyBasic(TEST_CODES[0], TARGET, close=Decimal("10")),),
|
||||||
|
WINDOW,
|
||||||
|
)
|
||||||
|
assert unchanged.inserted == 1
|
||||||
|
assert repository.count_valid_stocks(TARGET) == 1
|
||||||
|
|
||||||
|
batch_id = "9919-pool-test"
|
||||||
|
outcomes = (
|
||||||
|
SyncItemOutcome("bar", TEST_CODES[0], "success", WriteResult(inserted=1)),
|
||||||
|
SyncItemOutcome(
|
||||||
|
"bar",
|
||||||
|
TEST_CODES[1],
|
||||||
|
"failed",
|
||||||
|
failure=SyncFailure("bar", TEST_CODES[1], "source_error", "safe failure"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
# Use the real schema row directly because the test repository's
|
||||||
|
# create_batch API generates a UUID for normal production calls.
|
||||||
|
with _database(database_url) as connection:
|
||||||
|
connection.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO market_sync_batch
|
||||||
|
(id, target_trade_date, window_start, mode, status, target_count)
|
||||||
|
VALUES (%s, %s, %s, 'daily', 'running', 2)
|
||||||
|
""",
|
||||||
|
(batch_id, TARGET, TARGET),
|
||||||
|
)
|
||||||
|
repository.record_items(batch_id, outcomes)
|
||||||
|
with _database(database_url) as connection:
|
||||||
|
count = connection.execute(
|
||||||
|
"SELECT count(*) FROM market_sync_item WHERE batch_id = %s",
|
||||||
|
(batch_id,),
|
||||||
|
).fetchone()[0]
|
||||||
|
assert count == 2
|
||||||
|
finally:
|
||||||
|
repository.close()
|
||||||
|
_cleanup(database_url)
|
||||||
@@ -5,7 +5,11 @@ from pathlib import Path
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic
|
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:
|
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"):
|
with pytest.raises(ValueError, match="duplicate"):
|
||||||
store.stage_daily_basic(target, (make_basic(target), make_basic(target)))
|
store.stage_daily_basic(target, (make_basic(target), make_basic(target)))
|
||||||
|
|
||||||
|
|
||||||
|
def test_readers_reject_empty_formal_snapshots(tmp_path: Path) -> None:
|
||||||
|
store = CsvSnapshotStore(tmp_path)
|
||||||
|
path = store.bars_path("000001.SZ")
|
||||||
|
path.parent.mkdir(parents=True)
|
||||||
|
path.write_text(",".join(BAR_COLUMNS) + "\n", encoding="utf-8")
|
||||||
|
|
||||||
|
with pytest.raises(SnapshotReadError, match="empty"):
|
||||||
|
store.read_bars("000001.SZ")
|
||||||
|
|
||||||
|
|
||||||
|
def test_invalid_daily_basic_path_is_reported_without_opening_content(tmp_path: Path) -> None:
|
||||||
|
store = CsvSnapshotStore(tmp_path)
|
||||||
|
invalid_path = tmp_path / "daily-basic" / "2024" / "20241301.csv"
|
||||||
|
invalid_path.parent.mkdir(parents=True)
|
||||||
|
invalid_path.write_text("not parsed", encoding="utf-8")
|
||||||
|
|
||||||
|
assert store.list_invalid_daily_basic_snapshot_files() == (
|
||||||
|
("daily-basic/2024/20241301.csv", "window_out_of_bounds"),
|
||||||
|
)
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from zhixing_server.modules.market_data.domain.models import (
|
|||||||
DailyBasic,
|
DailyBasic,
|
||||||
Stock,
|
Stock,
|
||||||
SyncWindow,
|
SyncWindow,
|
||||||
|
decimal_text,
|
||||||
)
|
)
|
||||||
from zhixing_server.modules.market_data.domain.rules import filter_current_hs_a_stocks
|
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:
|
def test_universe_keeps_current_non_st_hs_a_stocks() -> None:
|
||||||
stocks = (
|
stocks = (
|
||||||
Stock("000001.SZ", "平安银行", exchange="SZSE", list_status="L"),
|
Stock("000001.SZ", "平安银行", exchange="SZSE", list_status="L"),
|
||||||
|
|||||||
@@ -0,0 +1,345 @@
|
|||||||
|
from collections.abc import Generator, Iterable
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
from zhixing_server.bootstrap.app import create_app
|
||||||
|
from zhixing_server.modules.market_data.application.integrity import RunMarketIntegrityCheck
|
||||||
|
from zhixing_server.modules.market_data.domain.integrity import (
|
||||||
|
IntegrityCheckInProgress,
|
||||||
|
IntegrityCheckPage,
|
||||||
|
IntegrityCheckQuery,
|
||||||
|
IntegrityCheckRun,
|
||||||
|
IntegrityCheckStoreError,
|
||||||
|
IntegrityIssue,
|
||||||
|
)
|
||||||
|
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||||
|
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import SnapshotReadError
|
||||||
|
from zhixing_server.modules.market_data.presentation.integrity import get_market_integrity_service
|
||||||
|
|
||||||
|
WINDOW = SyncWindow(date(2024, 1, 2), date(2024, 1, 3))
|
||||||
|
STOCK = Stock("000001.SZ", "平安银行", exchange="SZSE")
|
||||||
|
BARS = (
|
||||||
|
Bar(STOCK.ts_code, date(2024, 1, 2), close=Decimal("10")),
|
||||||
|
Bar(STOCK.ts_code, date(2024, 1, 3), close=Decimal("11")),
|
||||||
|
)
|
||||||
|
BASICS = (
|
||||||
|
DailyBasic(STOCK.ts_code, date(2024, 1, 2), close=Decimal("10")),
|
||||||
|
DailyBasic(STOCK.ts_code, date(2024, 1, 3), close=Decimal("11")),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class FakeReader:
|
||||||
|
def __init__(self, *, acquired: bool = True) -> None:
|
||||||
|
self.acquired = acquired
|
||||||
|
self.stocks = (STOCK,)
|
||||||
|
self.bars: tuple[Bar, ...] = BARS
|
||||||
|
self.basics: tuple[DailyBasic, ...] = BASICS
|
||||||
|
self.bar_rows_read = 0
|
||||||
|
self.basic_rows_read = 0
|
||||||
|
|
||||||
|
def latest_successful_window(self) -> SyncWindow:
|
||||||
|
return WINDOW
|
||||||
|
|
||||||
|
def active_stocks(self) -> tuple[Stock, ...]:
|
||||||
|
return self.stocks
|
||||||
|
|
||||||
|
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||||
|
return (STOCK.ts_code,)
|
||||||
|
|
||||||
|
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
|
||||||
|
return (date(2024, 1, 2), date(2024, 1, 3))
|
||||||
|
|
||||||
|
def iter_bars(self, window: SyncWindow) -> Iterable[Bar]:
|
||||||
|
for row in self.bars:
|
||||||
|
self.bar_rows_read += 1
|
||||||
|
yield row
|
||||||
|
|
||||||
|
def iter_daily_basic(self, window: SyncWindow) -> Iterable[DailyBasic]:
|
||||||
|
for row in self.basics:
|
||||||
|
self.basic_rows_read += 1
|
||||||
|
yield row
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
|
||||||
|
yield self.acquired
|
||||||
|
|
||||||
|
|
||||||
|
class BrokenKeyReader(FakeReader):
|
||||||
|
"""Reader whose storage key query fails before a check can be claimed."""
|
||||||
|
|
||||||
|
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||||
|
raise RuntimeError("database unavailable")
|
||||||
|
|
||||||
|
|
||||||
|
class FakeSnapshots:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.stocks: tuple[Stock, ...] | None = (STOCK,)
|
||||||
|
self.bars: tuple[Bar, ...] | None = BARS
|
||||||
|
self.basics: dict[date, tuple[DailyBasic, ...]] = {
|
||||||
|
date(2024, 1, 2): (BASICS[0],),
|
||||||
|
date(2024, 1, 3): (BASICS[1],),
|
||||||
|
}
|
||||||
|
self.parse_bar = False
|
||||||
|
|
||||||
|
def list_bar_codes(self, window: SyncWindow) -> tuple[str, ...]:
|
||||||
|
return (STOCK.ts_code,) if self.bars is not None else ()
|
||||||
|
|
||||||
|
def list_daily_basic_dates(self, window: SyncWindow) -> tuple[date, ...]:
|
||||||
|
return tuple(self.basics)
|
||||||
|
|
||||||
|
def read_stocks(self) -> tuple[Stock, ...] | None:
|
||||||
|
return self.stocks
|
||||||
|
|
||||||
|
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None:
|
||||||
|
if self.parse_bar:
|
||||||
|
raise SnapshotReadError("bad bar csv")
|
||||||
|
return self.bars
|
||||||
|
|
||||||
|
def read_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...] | None:
|
||||||
|
return self.basics.get(trade_date)
|
||||||
|
|
||||||
|
|
||||||
|
class FakeStore:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.run = IntegrityCheckRun("check-1", "running", WINDOW, 4)
|
||||||
|
self.issues: list[IntegrityIssue] = []
|
||||||
|
self.progress: tuple[int, int] = (0, 0)
|
||||||
|
|
||||||
|
def recover_stale_running(self, lock_key: int, stale_after_seconds: int) -> int:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def create_running(self, window: SyncWindow, target_count: int) -> IntegrityCheckRun:
|
||||||
|
self.run = IntegrityCheckRun("check-1", "running", window, target_count)
|
||||||
|
return self.run
|
||||||
|
|
||||||
|
def update_progress(self, check_id: str, checked_count: int, issue_count: int) -> None:
|
||||||
|
self.progress = (checked_count, issue_count)
|
||||||
|
self.run = IntegrityCheckRun(
|
||||||
|
self.run.id,
|
||||||
|
self.run.status,
|
||||||
|
self.run.window,
|
||||||
|
self.run.target_count,
|
||||||
|
checked_count,
|
||||||
|
issue_count,
|
||||||
|
)
|
||||||
|
|
||||||
|
def record_issues(self, issues: Iterable[IntegrityIssue]) -> None:
|
||||||
|
self.issues.extend(issues)
|
||||||
|
|
||||||
|
def finish(
|
||||||
|
self,
|
||||||
|
check_id: str,
|
||||||
|
status: str,
|
||||||
|
*,
|
||||||
|
issue_count: int,
|
||||||
|
error_type: str | None = None,
|
||||||
|
error_message: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.run = IntegrityCheckRun(
|
||||||
|
self.run.id,
|
||||||
|
status, # type: ignore[arg-type]
|
||||||
|
self.run.window,
|
||||||
|
self.run.target_count,
|
||||||
|
self.progress[0],
|
||||||
|
issue_count,
|
||||||
|
error_type,
|
||||||
|
error_message,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get(self, check_id: str, query: IntegrityCheckQuery) -> IntegrityCheckPage:
|
||||||
|
return IntegrityCheckPage(
|
||||||
|
self.run,
|
||||||
|
query.page,
|
||||||
|
query.page_size,
|
||||||
|
len(self.issues),
|
||||||
|
tuple(self.issues[(query.page - 1) * query.page_size : query.page * query.page_size]),
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_latest(self, query: IntegrityCheckQuery) -> IntegrityCheckPage:
|
||||||
|
return self.get(self.run.id, query)
|
||||||
|
|
||||||
|
|
||||||
|
def test_integrity_check_passes_without_writing_facts() -> None:
|
||||||
|
reader = FakeReader()
|
||||||
|
snapshots = FakeSnapshots()
|
||||||
|
store = FakeStore()
|
||||||
|
service = RunMarketIntegrityCheck(reader, snapshots, store)
|
||||||
|
|
||||||
|
run = service.prepare()
|
||||||
|
service.execute(run.id)
|
||||||
|
|
||||||
|
assert store.run.status == "passed"
|
||||||
|
assert store.issues == []
|
||||||
|
assert reader.bar_rows_read == len(BARS)
|
||||||
|
assert reader.basic_rows_read == len(BASICS)
|
||||||
|
|
||||||
|
|
||||||
|
def test_prepare_does_not_hide_storage_key_query_failure() -> None:
|
||||||
|
service = RunMarketIntegrityCheck(BrokenKeyReader(), FakeSnapshots(), FakeStore())
|
||||||
|
|
||||||
|
try:
|
||||||
|
service.prepare()
|
||||||
|
except RuntimeError as exc:
|
||||||
|
assert str(exc) == "database unavailable"
|
||||||
|
else:
|
||||||
|
raise AssertionError("storage failure must not claim a running check")
|
||||||
|
|
||||||
|
|
||||||
|
def test_integrity_check_reports_stable_content_and_missing_types() -> None:
|
||||||
|
reader = FakeReader()
|
||||||
|
snapshots = FakeSnapshots()
|
||||||
|
snapshots.bars = (Bar(STOCK.ts_code, date(2024, 1, 2), close=Decimal("99")),)
|
||||||
|
snapshots.basics.pop(date(2024, 1, 3))
|
||||||
|
store = FakeStore()
|
||||||
|
service = RunMarketIntegrityCheck(reader, snapshots, store)
|
||||||
|
|
||||||
|
service.execute(service.prepare().id)
|
||||||
|
|
||||||
|
assert store.run.status == "issues_found"
|
||||||
|
assert {issue.issue_type for issue in store.issues} >= {
|
||||||
|
"content_mismatch",
|
||||||
|
"missing_csv",
|
||||||
|
}
|
||||||
|
keys = [issue.issue_key for issue in store.issues]
|
||||||
|
assert keys == list(dict.fromkeys(keys))
|
||||||
|
|
||||||
|
|
||||||
|
def test_csv_parse_error_does_not_stop_other_groups() -> None:
|
||||||
|
reader = FakeReader()
|
||||||
|
snapshots = FakeSnapshots()
|
||||||
|
snapshots.parse_bar = True
|
||||||
|
store = FakeStore()
|
||||||
|
service = RunMarketIntegrityCheck(reader, snapshots, store)
|
||||||
|
|
||||||
|
service.execute(service.prepare().id)
|
||||||
|
|
||||||
|
assert store.run.status == "issues_found"
|
||||||
|
assert any(issue.issue_type == "parse_error" for issue in store.issues)
|
||||||
|
assert not any(
|
||||||
|
issue.item_kind == "bar" and issue.issue_type == "missing_csv" for issue in store.issues
|
||||||
|
)
|
||||||
|
assert reader.basic_rows_read == len(BASICS)
|
||||||
|
|
||||||
|
|
||||||
|
def test_window_only_bar_group_does_not_become_extra_csv() -> None:
|
||||||
|
reader = FakeReader()
|
||||||
|
reader.bars = ()
|
||||||
|
snapshots = FakeSnapshots()
|
||||||
|
snapshots.bars = (Bar(STOCK.ts_code, date(2025, 1, 2), close=Decimal("10")),)
|
||||||
|
store = FakeStore()
|
||||||
|
service = RunMarketIntegrityCheck(reader, snapshots, store)
|
||||||
|
|
||||||
|
service.execute(service.prepare().id)
|
||||||
|
|
||||||
|
bar_issues = [issue for issue in store.issues if issue.item_kind == "bar"]
|
||||||
|
assert {issue.issue_type for issue in bar_issues} == {"window_out_of_bounds"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_integrity_lock_conflict_converges_to_failed() -> None:
|
||||||
|
reader = FakeReader(acquired=False)
|
||||||
|
store = FakeStore()
|
||||||
|
service = RunMarketIntegrityCheck(reader, FakeSnapshots(), store)
|
||||||
|
|
||||||
|
service.execute(service.prepare().id)
|
||||||
|
|
||||||
|
assert store.run.status == "failed"
|
||||||
|
assert store.run.error_type == "lock_unavailable"
|
||||||
|
|
||||||
|
|
||||||
|
class FakeHttpService:
|
||||||
|
def __init__(self, result: IntegrityCheckPage | None = None) -> None:
|
||||||
|
self.result = result
|
||||||
|
self.executed = False
|
||||||
|
self.mode = "ok"
|
||||||
|
|
||||||
|
def prepare(self) -> IntegrityCheckRun:
|
||||||
|
if self.mode == "in_progress":
|
||||||
|
raise IntegrityCheckInProgress("already running")
|
||||||
|
if self.mode == "storage_error":
|
||||||
|
raise IntegrityCheckStoreError("storage unavailable")
|
||||||
|
return IntegrityCheckRun("http-check", "running", WINDOW, 1)
|
||||||
|
|
||||||
|
def execute(self, check_id: str) -> None:
|
||||||
|
self.executed = True
|
||||||
|
|
||||||
|
def get(self, check_id: str, *, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
|
||||||
|
if self.mode == "storage_error":
|
||||||
|
raise IntegrityCheckStoreError("storage unavailable")
|
||||||
|
return self.result
|
||||||
|
|
||||||
|
def get_latest(self, *, query: IntegrityCheckQuery) -> IntegrityCheckPage | None:
|
||||||
|
if self.mode == "storage_error":
|
||||||
|
raise IntegrityCheckStoreError("storage unavailable")
|
||||||
|
return self.result
|
||||||
|
|
||||||
|
|
||||||
|
def test_integrity_http_returns_202_and_paginates_report() -> None:
|
||||||
|
issue = IntegrityIssue.build("http-check", "bar", "000001.SZ:2024-01-02", "parse_error", "bad")
|
||||||
|
result = IntegrityCheckPage(
|
||||||
|
IntegrityCheckRun("http-check", "issues_found", WINDOW, 1, 1, 1),
|
||||||
|
page=2,
|
||||||
|
page_size=1,
|
||||||
|
issues_total=1,
|
||||||
|
issues=(issue,),
|
||||||
|
)
|
||||||
|
service = FakeHttpService(result)
|
||||||
|
app = create_app()
|
||||||
|
app.dependency_overrides[get_market_integrity_service] = lambda: service
|
||||||
|
client = TestClient(app)
|
||||||
|
|
||||||
|
accepted = client.post("/api/v1/market-data/integrity-checks")
|
||||||
|
report = client.get(
|
||||||
|
"/api/v1/market-data/integrity-checks/http-check",
|
||||||
|
params={"page": 2, "page_size": 1},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert accepted.status_code == 202
|
||||||
|
assert accepted.json()["check_id"] == "http-check"
|
||||||
|
assert service.executed is True
|
||||||
|
assert report.status_code == 200
|
||||||
|
assert report.json()["issues"][0]["issue_type"] == "parse_error"
|
||||||
|
assert report.json()["page"] == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_integrity_http_latest_without_report_is_no_data() -> None:
|
||||||
|
service = FakeHttpService(None)
|
||||||
|
app = create_app()
|
||||||
|
app.dependency_overrides[get_market_integrity_service] = lambda: service
|
||||||
|
|
||||||
|
response = TestClient(app).get("/api/v1/market-data/integrity-checks/latest")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json()["status"] == "no_data"
|
||||||
|
assert response.json()["issues"] == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_integrity_http_maps_running_conflict_to_409() -> None:
|
||||||
|
service = FakeHttpService()
|
||||||
|
service.mode = "in_progress"
|
||||||
|
app = create_app()
|
||||||
|
app.dependency_overrides[get_market_integrity_service] = lambda: service
|
||||||
|
|
||||||
|
response = TestClient(app).post("/api/v1/market-data/integrity-checks")
|
||||||
|
|
||||||
|
assert response.status_code == 409
|
||||||
|
assert response.json()["detail"]["code"] == "integrity_check_in_progress"
|
||||||
|
|
||||||
|
|
||||||
|
def test_integrity_http_maps_missing_report_and_storage_to_404_and_503() -> None:
|
||||||
|
missing_service = FakeHttpService(None)
|
||||||
|
app = create_app()
|
||||||
|
app.dependency_overrides[get_market_integrity_service] = lambda: missing_service
|
||||||
|
missing = TestClient(app).get("/api/v1/market-data/integrity-checks/unknown")
|
||||||
|
assert missing.status_code == 404
|
||||||
|
assert missing.json()["detail"]["code"] == "integrity_check_not_found"
|
||||||
|
|
||||||
|
storage_service = FakeHttpService(None)
|
||||||
|
storage_service.mode = "storage_error"
|
||||||
|
app = create_app()
|
||||||
|
app.dependency_overrides[get_market_integrity_service] = lambda: storage_service
|
||||||
|
unavailable = TestClient(app).get("/api/v1/market-data/integrity-checks/latest")
|
||||||
|
assert unavailable.status_code == 503
|
||||||
|
assert unavailable.json()["detail"]["code"] == "integrity_storage_unavailable"
|
||||||
@@ -0,0 +1,241 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
from collections.abc import Generator, Iterable, Sequence
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from zhixing_server.modules.market_data.application.sync import (
|
||||||
|
SyncBatchSummary,
|
||||||
|
SyncItemOutcome,
|
||||||
|
SyncMarketData,
|
||||||
|
SyncMarketDataCommand,
|
||||||
|
)
|
||||||
|
from zhixing_server.modules.market_data.domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||||
|
from zhixing_server.modules.market_data.domain.ports import WriteResult
|
||||||
|
from zhixing_server.modules.market_data.infrastructure.csv_snapshot import CsvSnapshotStore
|
||||||
|
|
||||||
|
|
||||||
|
class ConcurrentSource:
|
||||||
|
def __init__(
|
||||||
|
self, target: date, delays: dict[str, float], failed_code: str | None = None
|
||||||
|
) -> None:
|
||||||
|
self.target = target
|
||||||
|
self.stocks = tuple(
|
||||||
|
Stock(f"{index:06d}.SZ", f"测试{index}", exchange="SZSE", list_status="L")
|
||||||
|
for index in range(1, 6)
|
||||||
|
)
|
||||||
|
self.delays = delays
|
||||||
|
self.failed_code = failed_code
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self.active = 0
|
||||||
|
self.max_active = 0
|
||||||
|
|
||||||
|
def fetch_stocks(self) -> Sequence[Stock]:
|
||||||
|
return self.stocks
|
||||||
|
|
||||||
|
def fetch_open_dates(self, start: date, end: date) -> Sequence[date]:
|
||||||
|
return (self.target,) if start <= self.target <= end else ()
|
||||||
|
|
||||||
|
def fetch_daily_basic(self, trade_date: date) -> Sequence[DailyBasic]:
|
||||||
|
return tuple(
|
||||||
|
DailyBasic(stock.ts_code, trade_date, close=Decimal("10")) for stock in self.stocks
|
||||||
|
)
|
||||||
|
|
||||||
|
def fetch_bars(self, ts_code: str, window: SyncWindow) -> Sequence[Bar]:
|
||||||
|
with self._lock:
|
||||||
|
self.active += 1
|
||||||
|
self.max_active = max(self.max_active, self.active)
|
||||||
|
try:
|
||||||
|
time.sleep(self.delays.get(ts_code, 0))
|
||||||
|
if ts_code == self.failed_code:
|
||||||
|
raise RuntimeError("simulated source failure")
|
||||||
|
return (Bar(ts_code, self.target, close=Decimal("10")),)
|
||||||
|
finally:
|
||||||
|
with self._lock:
|
||||||
|
self.active -= 1
|
||||||
|
|
||||||
|
|
||||||
|
class ConcurrentRepository:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.bars: dict[tuple[str, date], Bar] = {}
|
||||||
|
self.daily_basic: dict[tuple[str, date], DailyBasic] = {}
|
||||||
|
self.audit_batches: list[tuple[SyncItemOutcome, ...]] = []
|
||||||
|
self.batch_summary: tuple[str, int, Decimal, bool] | None = None
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
|
||||||
|
yield True
|
||||||
|
|
||||||
|
def upsert_stocks(self, rows: Iterable[Stock]) -> WriteResult:
|
||||||
|
return WriteResult(inserted=len(tuple(rows)))
|
||||||
|
|
||||||
|
def upsert_bars(
|
||||||
|
self,
|
||||||
|
rows: Iterable[Bar],
|
||||||
|
window: SyncWindow,
|
||||||
|
*,
|
||||||
|
full_snapshot: bool,
|
||||||
|
) -> WriteResult:
|
||||||
|
records = tuple(rows)
|
||||||
|
inserted = 0
|
||||||
|
for row in records:
|
||||||
|
with self._lock:
|
||||||
|
if (row.ts_code, row.trade_date) not in self.bars:
|
||||||
|
inserted += 1
|
||||||
|
self.bars[(row.ts_code, row.trade_date)] = row
|
||||||
|
return WriteResult(inserted=inserted)
|
||||||
|
|
||||||
|
def upsert_daily_basic(self, rows: Iterable[DailyBasic], window: SyncWindow) -> WriteResult:
|
||||||
|
records = tuple(rows)
|
||||||
|
with self._lock:
|
||||||
|
for row in records:
|
||||||
|
self.daily_basic[(row.ts_code, row.trade_date)] = row
|
||||||
|
return WriteResult(inserted=len(records))
|
||||||
|
|
||||||
|
def purge_before(self, window: SyncWindow) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def create_batch(
|
||||||
|
self,
|
||||||
|
target_trade_date: date,
|
||||||
|
window: SyncWindow,
|
||||||
|
mode: str,
|
||||||
|
parent_batch_id: str | None,
|
||||||
|
target_count: int,
|
||||||
|
) -> str:
|
||||||
|
return "batch-test"
|
||||||
|
|
||||||
|
def record_batch(
|
||||||
|
self,
|
||||||
|
batch_id: str,
|
||||||
|
status: str,
|
||||||
|
valid_count: int,
|
||||||
|
coverage: Decimal,
|
||||||
|
strategy_eligible: bool,
|
||||||
|
) -> None:
|
||||||
|
self.batch_summary = (status, valid_count, coverage, strategy_eligible)
|
||||||
|
|
||||||
|
def record_items(self, batch_id: str, outcomes: Iterable[object]) -> None:
|
||||||
|
self.audit_batches.append(tuple(outcomes)) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
def record_item(
|
||||||
|
self,
|
||||||
|
batch_id: str,
|
||||||
|
item_kind: str,
|
||||||
|
item_key: str,
|
||||||
|
status: str,
|
||||||
|
result: WriteResult,
|
||||||
|
fingerprint: str | None = None,
|
||||||
|
error_type: str | None = None,
|
||||||
|
error_message: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def count_valid_stocks(self, trade_date: date) -> int:
|
||||||
|
return sum(
|
||||||
|
1
|
||||||
|
for ts_code, current_date in self.bars
|
||||||
|
if current_date == trade_date and (ts_code, trade_date) in self.daily_basic
|
||||||
|
)
|
||||||
|
|
||||||
|
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]:
|
||||||
|
return ()
|
||||||
|
|
||||||
|
def has_bar(self, ts_code: str, trade_date: date) -> bool:
|
||||||
|
raise AssertionError("coverage must use count_valid_stocks")
|
||||||
|
|
||||||
|
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool:
|
||||||
|
raise AssertionError("coverage must use count_valid_stocks")
|
||||||
|
|
||||||
|
|
||||||
|
def _run_sync(
|
||||||
|
tmp_path: Path, delays: dict[str, float], failed_code: str | None = None
|
||||||
|
) -> tuple[SyncBatchSummary, ConcurrentSource, ConcurrentRepository]:
|
||||||
|
target = date(2024, 1, 2)
|
||||||
|
source = ConcurrentSource(target, delays, failed_code)
|
||||||
|
repository = ConcurrentRepository()
|
||||||
|
summary = SyncMarketData(
|
||||||
|
source,
|
||||||
|
CsvSnapshotStore(tmp_path),
|
||||||
|
repository,
|
||||||
|
today=target,
|
||||||
|
max_workers=2,
|
||||||
|
).execute(SyncMarketDataCommand(target_trade_date=target))
|
||||||
|
return summary, source, repository
|
||||||
|
|
||||||
|
|
||||||
|
def test_bar_workers_are_bounded_and_audit_is_batched(tmp_path: Path) -> None:
|
||||||
|
summary, source, repository = _run_sync(
|
||||||
|
tmp_path,
|
||||||
|
{"000001.SZ": 0.03, "000002.SZ": 0.01, "000003.SZ": 0.02, "000004.SZ": 0},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert summary.status == "success"
|
||||||
|
assert source.max_active <= 2
|
||||||
|
assert summary.valid_count == 5
|
||||||
|
assert sum(len(batch) for batch in repository.audit_batches) == 7
|
||||||
|
|
||||||
|
|
||||||
|
def test_one_bar_failure_does_not_publish_a_csv_or_reduce_other_facts(tmp_path: Path) -> None:
|
||||||
|
failed_code = "000003.SZ"
|
||||||
|
summary, _, repository = _run_sync(
|
||||||
|
tmp_path,
|
||||||
|
{"000001.SZ": 0.02, "000002.SZ": 0.01, failed_code: 0},
|
||||||
|
failed_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert summary.status == "partial_success"
|
||||||
|
assert summary.valid_count == 4
|
||||||
|
assert [failure.item_key for failure in summary.failures if failure.item_kind == "bar"] == [
|
||||||
|
failed_code
|
||||||
|
]
|
||||||
|
assert not (tmp_path / "bars" / f"{failed_code}.csv").exists()
|
||||||
|
assert not list((tmp_path / "bars").glob(f".{failed_code}.csv.*.tmp"))
|
||||||
|
assert len(repository.bars) == 4
|
||||||
|
|
||||||
|
|
||||||
|
def test_completion_order_does_not_change_aggregate_counts(tmp_path: Path) -> None:
|
||||||
|
first, _, _ = _run_sync(
|
||||||
|
tmp_path / "first",
|
||||||
|
{"000001.SZ": 0.03, "000002.SZ": 0, "000003.SZ": 0.02},
|
||||||
|
)
|
||||||
|
second, _, _ = _run_sync(
|
||||||
|
tmp_path / "second",
|
||||||
|
{"000001.SZ": 0, "000002.SZ": 0.03, "000003.SZ": 0.01},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert (
|
||||||
|
first.status,
|
||||||
|
first.valid_count,
|
||||||
|
first.inserted_count,
|
||||||
|
first.updated_count,
|
||||||
|
first.unchanged_count,
|
||||||
|
first.failures,
|
||||||
|
) == (
|
||||||
|
second.status,
|
||||||
|
second.valid_count,
|
||||||
|
second.inserted_count,
|
||||||
|
second.updated_count,
|
||||||
|
second.unchanged_count,
|
||||||
|
second.failures,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_partial_success_always_uses_incomplete_exit_code() -> None:
|
||||||
|
summary = SyncBatchSummary(
|
||||||
|
batch_id="batch-test",
|
||||||
|
target_trade_date=date(2024, 1, 2),
|
||||||
|
window=SyncWindow(start=date(2018, 1, 2), end=date(2024, 1, 2)),
|
||||||
|
status="partial_success",
|
||||||
|
target_count=100,
|
||||||
|
valid_count=99,
|
||||||
|
coverage=Decimal("0.99"),
|
||||||
|
strategy_eligible=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert summary.exit_code == 2
|
||||||
@@ -4,7 +4,10 @@ import pytest
|
|||||||
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
||||||
|
|
||||||
from zhixing_server.modules.market_data.domain.models import SyncWindow
|
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(
|
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
|
||||||
assert calls[0]["api"] is created_client
|
assert calls[0]["api"] is created_client
|
||||||
assert calls[0]["adj"] == "qfq"
|
assert calls[0]["adj"] == "qfq"
|
||||||
|
assert calls[0]["retry_count"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_rate_limit_cooldown_is_shared_by_following_requests() -> None:
|
||||||
|
current = [0.0]
|
||||||
|
waits: list[float] = []
|
||||||
|
calls: list[float] = []
|
||||||
|
|
||||||
|
def clock() -> float:
|
||||||
|
return current[0]
|
||||||
|
|
||||||
|
def wait(seconds: float) -> None:
|
||||||
|
waits.append(seconds)
|
||||||
|
current[0] += seconds
|
||||||
|
|
||||||
|
coordinator = RequestCoordinator(
|
||||||
|
max_retries=0,
|
||||||
|
cooldown_seconds=(60, 120, 180),
|
||||||
|
clock=clock,
|
||||||
|
wait_fn=wait,
|
||||||
|
sleep_fn=wait,
|
||||||
|
)
|
||||||
|
|
||||||
|
def rate_limited() -> object:
|
||||||
|
calls.append(current[0])
|
||||||
|
raise RuntimeError("HTTP 429: too many requests")
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError):
|
||||||
|
coordinator.call("daily", rate_limited)
|
||||||
|
|
||||||
|
assert calls == [0]
|
||||||
|
assert coordinator.cooldown_until == 60
|
||||||
|
|
||||||
|
result = coordinator.call("adj_factor", lambda: calls.append(current[0]) or "ok")
|
||||||
|
|
||||||
|
assert result == "ok"
|
||||||
|
assert calls == [0, 60]
|
||||||
|
assert waits == [60]
|
||||||
|
|
||||||
|
|
||||||
|
def test_pro_bar_qfq_calls_are_bound_to_the_shared_coordinator(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
class Client:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.calls: list[str] = []
|
||||||
|
|
||||||
|
def daily(self, **kwargs: object) -> object:
|
||||||
|
self.calls.append("daily")
|
||||||
|
return object()
|
||||||
|
|
||||||
|
def adj_factor(self, **kwargs: object) -> object:
|
||||||
|
self.calls.append("adj_factor")
|
||||||
|
return object()
|
||||||
|
|
||||||
|
client = Client()
|
||||||
|
pro_bar_calls: list[dict[str, object]] = []
|
||||||
|
|
||||||
|
def fake_pro_api(token: str) -> object:
|
||||||
|
return client
|
||||||
|
|
||||||
|
def fake_pro_bar(**kwargs: object) -> list[dict[str, object]]:
|
||||||
|
pro_bar_calls.append(kwargs)
|
||||||
|
api = kwargs["api"]
|
||||||
|
assert api is client
|
||||||
|
api.daily(ts_code="000001.SZ") # type: ignore[union-attr]
|
||||||
|
api.adj_factor(ts_code="000001.SZ") # type: ignore[union-attr]
|
||||||
|
return [{"ts_code": "000001.SZ", "trade_date": "20240102", "close": "10"}]
|
||||||
|
|
||||||
|
monkeypatch.setattr(ts, "pro_api", fake_pro_api)
|
||||||
|
monkeypatch.setattr(ts, "pro_bar", fake_pro_bar)
|
||||||
|
|
||||||
|
adapter = TushareAdapter.from_token("test-token", request_interval_seconds=0)
|
||||||
|
adapter.fetch_bars(
|
||||||
|
"000001.SZ",
|
||||||
|
SyncWindow(start=date(2024, 1, 2), end=date(2024, 1, 2)),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert client.calls == ["daily", "adj_factor"]
|
||||||
|
assert pro_bar_calls[0]["retry_count"] == 1
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
import pytest
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
from zhixing_server.bootstrap.config import Settings
|
||||||
|
|
||||||
|
|
||||||
|
def test_market_data_workers_default_to_eight(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
monkeypatch.delenv("ZHIXING_MARKET_DATA_MAX_WORKERS", raising=False)
|
||||||
|
settings = Settings()
|
||||||
|
|
||||||
|
assert settings.market_data_max_workers == 8
|
||||||
|
|
||||||
|
|
||||||
|
def test_market_data_workers_must_be_positive() -> None:
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
Settings(market_data_max_workers=0)
|
||||||
Generated
+17
-2
@@ -412,6 +412,9 @@ wheels = [
|
|||||||
binary = [
|
binary = [
|
||||||
{ name = "psycopg-binary", marker = "implementation_name != 'pypy'" },
|
{ name = "psycopg-binary", marker = "implementation_name != 'pypy'" },
|
||||||
]
|
]
|
||||||
|
pool = [
|
||||||
|
{ name = "psycopg-pool" },
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "psycopg-binary"
|
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" },
|
{ 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]]
|
[[package]]
|
||||||
name = "pydantic"
|
name = "pydantic"
|
||||||
version = "2.13.4"
|
version = "2.13.4"
|
||||||
@@ -879,7 +894,7 @@ dependencies = [
|
|||||||
{ name = "fastapi" },
|
{ name = "fastapi" },
|
||||||
{ name = "numpy" },
|
{ name = "numpy" },
|
||||||
{ name = "pandas" },
|
{ name = "pandas" },
|
||||||
{ name = "psycopg", extra = ["binary"] },
|
{ name = "psycopg", extra = ["binary", "pool"] },
|
||||||
{ name = "pydantic-settings" },
|
{ name = "pydantic-settings" },
|
||||||
{ name = "sqlalchemy" },
|
{ name = "sqlalchemy" },
|
||||||
{ name = "tushare" },
|
{ name = "tushare" },
|
||||||
@@ -902,7 +917,7 @@ requires-dist = [
|
|||||||
{ name = "fastapi", specifier = ">=0.141.1" },
|
{ name = "fastapi", specifier = ">=0.141.1" },
|
||||||
{ name = "numpy", specifier = ">=2.4.0" },
|
{ name = "numpy", specifier = ">=2.4.0" },
|
||||||
{ name = "pandas", specifier = ">=2.3.3" },
|
{ 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 = "pydantic-settings", specifier = ">=2.14.2" },
|
||||||
{ name = "sqlalchemy", specifier = ">=2.0.46" },
|
{ name = "sqlalchemy", specifier = ">=2.0.46" },
|
||||||
{ name = "tushare", specifier = ">=1.4.24" },
|
{ name = "tushare", specifier = ">=1.4.24" },
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ import {
|
|||||||
function useActiveRoutePath() {
|
function useActiveRoutePath() {
|
||||||
const matchRoute = useMatchRoute()
|
const matchRoute = useMatchRoute()
|
||||||
if (matchRoute({ to: "/selection", fuzzy: true })) return "/selection"
|
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: "/components", fuzzy: true })) return "/components"
|
||||||
if (matchRoute({ to: "/", fuzzy: false })) return "/"
|
if (matchRoute({ to: "/", fuzzy: false })) return "/"
|
||||||
return "/"
|
return "/"
|
||||||
@@ -244,7 +245,7 @@ function MoreSheet({
|
|||||||
<DialogHeader>
|
<DialogHeader>
|
||||||
<DialogTitle>更多功能</DialogTitle>
|
<DialogTitle>更多功能</DialogTitle>
|
||||||
<DialogDescription>
|
<DialogDescription>
|
||||||
行情数据与同步任务尚未上线,暂无更多可用入口。
|
行情数据尚未上线,暂无更多可用入口。
|
||||||
</DialogDescription>
|
</DialogDescription>
|
||||||
</DialogHeader>
|
</DialogHeader>
|
||||||
</DialogContent>
|
</DialogContent>
|
||||||
|
|||||||
@@ -0,0 +1,28 @@
|
|||||||
|
import { describe, expect, it } from "vitest"
|
||||||
|
|
||||||
|
import {
|
||||||
|
mobilePrimaryNavigation,
|
||||||
|
primaryNavigation,
|
||||||
|
routePresentation,
|
||||||
|
} from "./navigation"
|
||||||
|
|
||||||
|
describe("sync navigation", () => {
|
||||||
|
it("exposes the sync route in desktop and mobile navigation", () => {
|
||||||
|
const sync = primaryNavigation.find((item) => item.id === "sync")
|
||||||
|
|
||||||
|
expect(sync).toMatchObject({
|
||||||
|
availability: "available",
|
||||||
|
label: "同步任务",
|
||||||
|
to: "/sync",
|
||||||
|
})
|
||||||
|
expect(mobilePrimaryNavigation).toContain(sync)
|
||||||
|
})
|
||||||
|
|
||||||
|
it("provides the active route presentation for sync", () => {
|
||||||
|
expect(routePresentation["/sync"]).toEqual({
|
||||||
|
breadcrumb: "研究工作台",
|
||||||
|
id: "sync",
|
||||||
|
title: "市场数据完整性检查",
|
||||||
|
})
|
||||||
|
})
|
||||||
|
})
|
||||||
@@ -35,9 +35,9 @@ export const primaryNavigation: readonly NavigationItem[] = [
|
|||||||
{
|
{
|
||||||
id: "sync",
|
id: "sync",
|
||||||
label: "同步任务",
|
label: "同步任务",
|
||||||
to: null,
|
to: "/sync",
|
||||||
icon: ClipboardCheck,
|
icon: ClipboardCheck,
|
||||||
availability: "unavailable",
|
availability: "available",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
id: "selection",
|
id: "selection",
|
||||||
@@ -79,6 +79,11 @@ export const routePresentation: Record<string, RoutePresentation> = {
|
|||||||
breadcrumb: "研究工作台",
|
breadcrumb: "研究工作台",
|
||||||
title: "知行 B1 执行结果",
|
title: "知行 B1 执行结果",
|
||||||
},
|
},
|
||||||
|
"/sync": {
|
||||||
|
id: "sync",
|
||||||
|
breadcrumb: "研究工作台",
|
||||||
|
title: "市场数据完整性检查",
|
||||||
|
},
|
||||||
"/components": {
|
"/components": {
|
||||||
id: "components",
|
id: "components",
|
||||||
breadcrumb: "研究工作台",
|
breadcrumb: "研究工作台",
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||||
|
|
||||||
|
const requestJson = vi.hoisted(() => vi.fn())
|
||||||
|
|
||||||
|
vi.mock("@/shared/api/request-json", () => ({ requestJson }))
|
||||||
|
|
||||||
|
import {
|
||||||
|
getIntegrityCheck,
|
||||||
|
getLatestIntegrityCheck,
|
||||||
|
triggerIntegrityCheck,
|
||||||
|
} from "./sync.api"
|
||||||
|
|
||||||
|
describe("sync API adapters", () => {
|
||||||
|
beforeEach(() => {
|
||||||
|
requestJson.mockReset()
|
||||||
|
})
|
||||||
|
|
||||||
|
it("triggers one read-only integrity check through the POST endpoint", async () => {
|
||||||
|
const signal = new AbortController().signal
|
||||||
|
|
||||||
|
await triggerIntegrityCheck(signal)
|
||||||
|
|
||||||
|
expect(requestJson).toHaveBeenCalledWith(
|
||||||
|
"/api/v1/market-data/integrity-checks",
|
||||||
|
{
|
||||||
|
method: "POST",
|
||||||
|
signal,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
})
|
||||||
|
|
||||||
|
it("uses the latest report pagination contract", async () => {
|
||||||
|
await getLatestIntegrityCheck({ page: 2, pageSize: 20 })
|
||||||
|
|
||||||
|
const [input, init] = requestJson.mock.calls[0] as [
|
||||||
|
string,
|
||||||
|
{ signal?: AbortSignal },
|
||||||
|
]
|
||||||
|
const params = new URL(input, "http://localhost").searchParams
|
||||||
|
|
||||||
|
expect(input).toContain("/api/v1/market-data/integrity-checks/latest?")
|
||||||
|
expect(params.get("page")).toBe("2")
|
||||||
|
expect(params.get("page_size")).toBe("20")
|
||||||
|
expect(init).toEqual({ signal: undefined })
|
||||||
|
})
|
||||||
|
|
||||||
|
it("encodes a check id while keeping the paged check contract", async () => {
|
||||||
|
await getIntegrityCheck("check/with spaces", { page: 3, pageSize: 50 })
|
||||||
|
|
||||||
|
const [input] = requestJson.mock.calls[0] as [string]
|
||||||
|
const params = new URL(input, "http://localhost").searchParams
|
||||||
|
|
||||||
|
expect(input).toContain(
|
||||||
|
"/api/v1/market-data/integrity-checks/check%2Fwith%20spaces?",
|
||||||
|
)
|
||||||
|
expect(params.get("page")).toBe("3")
|
||||||
|
expect(params.get("page_size")).toBe("50")
|
||||||
|
})
|
||||||
|
})
|
||||||
@@ -0,0 +1,50 @@
|
|||||||
|
import { requestJson } from "@/shared/api/request-json"
|
||||||
|
|
||||||
|
import {
|
||||||
|
defaultMarketIntegrityQuery,
|
||||||
|
type IntegrityCheckAccepted,
|
||||||
|
type MarketIntegrityCheck,
|
||||||
|
type MarketIntegrityQuery,
|
||||||
|
} from "./sync.types"
|
||||||
|
|
||||||
|
const integrityChecksPath = "/api/v1/market-data/integrity-checks"
|
||||||
|
|
||||||
|
export function triggerIntegrityCheck(signal?: AbortSignal) {
|
||||||
|
return requestJson<IntegrityCheckAccepted>(integrityChecksPath, {
|
||||||
|
method: "POST",
|
||||||
|
signal,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
export function getLatestIntegrityCheck(
|
||||||
|
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
|
||||||
|
signal?: AbortSignal,
|
||||||
|
) {
|
||||||
|
return requestJson<MarketIntegrityCheck>(
|
||||||
|
`${integrityChecksPath}/latest?${buildQuery(query)}`,
|
||||||
|
{ signal },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
export function getIntegrityCheck(
|
||||||
|
checkId: string,
|
||||||
|
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
|
||||||
|
signal?: AbortSignal,
|
||||||
|
) {
|
||||||
|
return requestJson<MarketIntegrityCheck>(
|
||||||
|
`${integrityChecksPath}/${encodeURIComponent(checkId)}?${buildQuery(query)}`,
|
||||||
|
{ signal },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function buildQuery(query: MarketIntegrityQuery) {
|
||||||
|
const params = new URLSearchParams({
|
||||||
|
page: String(query.page),
|
||||||
|
page_size: String(query.pageSize),
|
||||||
|
})
|
||||||
|
return params.toString()
|
||||||
|
}
|
||||||
|
|
||||||
|
export const getMarketIntegrityCheck = getIntegrityCheck
|
||||||
|
export const getLatestMarketIntegrityCheck = getLatestIntegrityCheck
|
||||||
|
export const triggerMarketIntegrityCheck = triggerIntegrityCheck
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||||
|
|
||||||
|
import { ApiError } from "@/shared/api/request-json"
|
||||||
|
|
||||||
|
const queryHooks = vi.hoisted(() => ({
|
||||||
|
useMutation: vi.fn(),
|
||||||
|
useQuery: vi.fn(),
|
||||||
|
useQueryClient: vi.fn(),
|
||||||
|
}))
|
||||||
|
|
||||||
|
const api = vi.hoisted(() => ({
|
||||||
|
getIntegrityCheck: vi.fn(),
|
||||||
|
getLatestIntegrityCheck: vi.fn(),
|
||||||
|
triggerIntegrityCheck: vi.fn(),
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock("@tanstack/react-query", () => queryHooks)
|
||||||
|
vi.mock("./sync.api", () => api)
|
||||||
|
|
||||||
|
import {
|
||||||
|
marketIntegrityCheckQueryKey,
|
||||||
|
marketIntegrityLatestQueryKey,
|
||||||
|
useIntegrityCheck,
|
||||||
|
useLatestIntegrityCheck,
|
||||||
|
useTriggerIntegrityCheck,
|
||||||
|
} from "./sync.query"
|
||||||
|
|
||||||
|
interface QueryOptionsForTest {
|
||||||
|
queryKey: readonly unknown[]
|
||||||
|
enabled?: boolean
|
||||||
|
refetchInterval?: (query: {
|
||||||
|
state: { data?: { status?: string } }
|
||||||
|
}) => number | false
|
||||||
|
}
|
||||||
|
|
||||||
|
interface MutationOptionsForTest {
|
||||||
|
onError?: (error: unknown) => void
|
||||||
|
onSuccess?: () => void
|
||||||
|
}
|
||||||
|
|
||||||
|
describe("sync query hooks", () => {
|
||||||
|
const queryClient = { invalidateQueries: vi.fn() }
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.clearAllMocks()
|
||||||
|
queryHooks.useQuery.mockImplementation((options) => options)
|
||||||
|
queryHooks.useQueryClient.mockReturnValue(queryClient)
|
||||||
|
queryHooks.useMutation.mockImplementation((options) => options)
|
||||||
|
})
|
||||||
|
|
||||||
|
it("keeps latest and check keys under marketIntegrity", () => {
|
||||||
|
expect(marketIntegrityLatestQueryKey({ page: 2, pageSize: 20 })).toEqual([
|
||||||
|
"marketIntegrity",
|
||||||
|
"latest",
|
||||||
|
2,
|
||||||
|
20,
|
||||||
|
])
|
||||||
|
expect(
|
||||||
|
marketIntegrityCheckQueryKey("check-1", { page: 3, pageSize: 50 }),
|
||||||
|
).toEqual(["marketIntegrity", "check", "check-1", 3, 50])
|
||||||
|
})
|
||||||
|
|
||||||
|
it("polls only while a check response is running", () => {
|
||||||
|
useLatestIntegrityCheck({ page: 1, pageSize: 10 })
|
||||||
|
useIntegrityCheck("check-1", { page: 1, pageSize: 10 })
|
||||||
|
const latest = queryHooks.useQuery.mock.calls[0]?.[0] as QueryOptionsForTest
|
||||||
|
const check = queryHooks.useQuery.mock.calls[1]?.[0] as QueryOptionsForTest
|
||||||
|
|
||||||
|
expect(latest.queryKey).toEqual(["marketIntegrity", "latest", 1, 10])
|
||||||
|
expect(check.queryKey).toEqual([
|
||||||
|
"marketIntegrity",
|
||||||
|
"check",
|
||||||
|
"check-1",
|
||||||
|
1,
|
||||||
|
10,
|
||||||
|
])
|
||||||
|
expect(check.enabled).toBe(true)
|
||||||
|
if (!check.refetchInterval) throw new Error("polling callback is missing")
|
||||||
|
expect(
|
||||||
|
check.refetchInterval({ state: { data: { status: "running" } } }),
|
||||||
|
).toBe(1500)
|
||||||
|
expect(
|
||||||
|
check.refetchInterval({ state: { data: { status: "passed" } } }),
|
||||||
|
).toBe(false)
|
||||||
|
})
|
||||||
|
|
||||||
|
it("disables a check query without an active id", () => {
|
||||||
|
useIntegrityCheck(null)
|
||||||
|
const check = queryHooks.useQuery.mock.calls[0]?.[0] as QueryOptionsForTest
|
||||||
|
|
||||||
|
expect(check.enabled).toBe(false)
|
||||||
|
expect(check.queryKey).toEqual(["marketIntegrity", "check", "none", 1, 10])
|
||||||
|
})
|
||||||
|
|
||||||
|
it("invalidates latest reports after success and on a 409 conflict", () => {
|
||||||
|
useTriggerIntegrityCheck()
|
||||||
|
const mutation = queryHooks.useMutation.mock
|
||||||
|
.calls[0]?.[0] as MutationOptionsForTest
|
||||||
|
|
||||||
|
mutation.onSuccess?.()
|
||||||
|
mutation.onError?.(new ApiError(409, "already running"))
|
||||||
|
|
||||||
|
expect(queryClient.invalidateQueries).toHaveBeenCalledTimes(2)
|
||||||
|
expect(queryClient.invalidateQueries).toHaveBeenNthCalledWith(1, {
|
||||||
|
queryKey: ["marketIntegrity", "latest"],
|
||||||
|
})
|
||||||
|
expect(queryClient.invalidateQueries).toHaveBeenNthCalledWith(2, {
|
||||||
|
queryKey: ["marketIntegrity", "latest"],
|
||||||
|
})
|
||||||
|
})
|
||||||
|
})
|
||||||
@@ -0,0 +1,93 @@
|
|||||||
|
import {
|
||||||
|
useMutation,
|
||||||
|
useQuery,
|
||||||
|
useQueryClient,
|
||||||
|
type QueryClient,
|
||||||
|
} from "@tanstack/react-query"
|
||||||
|
|
||||||
|
import { ApiError } from "@/shared/api/request-json"
|
||||||
|
|
||||||
|
import {
|
||||||
|
getIntegrityCheck,
|
||||||
|
getLatestIntegrityCheck,
|
||||||
|
triggerIntegrityCheck,
|
||||||
|
} from "./sync.api"
|
||||||
|
import {
|
||||||
|
defaultMarketIntegrityQuery,
|
||||||
|
type IntegrityCheckAccepted,
|
||||||
|
type MarketIntegrityCheck,
|
||||||
|
type MarketIntegrityQuery,
|
||||||
|
} from "./sync.types"
|
||||||
|
|
||||||
|
export const marketIntegrityQueryKey = ["marketIntegrity"] as const
|
||||||
|
export const marketIntegrityLatestQueryPrefix = [
|
||||||
|
"marketIntegrity",
|
||||||
|
"latest",
|
||||||
|
] as const
|
||||||
|
export const marketIntegrityLatestQueryKey = (
|
||||||
|
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
|
||||||
|
) => [...marketIntegrityLatestQueryPrefix, query.page, query.pageSize] as const
|
||||||
|
export const marketIntegrityCheckQueryKey = (
|
||||||
|
checkId: string,
|
||||||
|
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
|
||||||
|
) => ["marketIntegrity", "check", checkId, query.page, query.pageSize] as const
|
||||||
|
|
||||||
|
const MARKET_INTEGRITY_POLL_INTERVAL_MS = 1500
|
||||||
|
|
||||||
|
export function useLatestIntegrityCheck(
|
||||||
|
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
|
||||||
|
) {
|
||||||
|
return useQuery({
|
||||||
|
queryFn: ({ signal }) => getLatestIntegrityCheck(query, signal),
|
||||||
|
queryKey: marketIntegrityLatestQueryKey(query),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
export function useIntegrityCheck(
|
||||||
|
checkId: string | null,
|
||||||
|
query: MarketIntegrityQuery = defaultMarketIntegrityQuery,
|
||||||
|
) {
|
||||||
|
return useQuery({
|
||||||
|
enabled: Boolean(checkId),
|
||||||
|
queryFn: ({ signal }) => getIntegrityCheck(checkId ?? "", query, signal),
|
||||||
|
queryKey: marketIntegrityCheckQueryKey(checkId ?? "none", query),
|
||||||
|
refetchInterval: (currentQuery) =>
|
||||||
|
currentQuery.state.data?.status === "running"
|
||||||
|
? MARKET_INTEGRITY_POLL_INTERVAL_MS
|
||||||
|
: false,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
export function useTriggerIntegrityCheck() {
|
||||||
|
const queryClient = useQueryClient()
|
||||||
|
|
||||||
|
return useMutation({
|
||||||
|
mutationFn: ({ signal }: { signal?: AbortSignal } = {}) =>
|
||||||
|
triggerIntegrityCheck(signal),
|
||||||
|
onError: (error) => {
|
||||||
|
if (error instanceof ApiError && error.status === 409) {
|
||||||
|
void invalidateLatestIntegrityChecks(queryClient)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
onSuccess: () => {
|
||||||
|
void invalidateLatestIntegrityChecks(queryClient)
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
export function invalidateLatestIntegrityChecks(queryClient: QueryClient) {
|
||||||
|
return queryClient.invalidateQueries({
|
||||||
|
queryKey: marketIntegrityLatestQueryPrefix,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
export const useLatestMarketIntegrityCheck = useLatestIntegrityCheck
|
||||||
|
export const useMarketIntegrityCheck = useIntegrityCheck
|
||||||
|
export const useTriggerMarketIntegrityCheck = useTriggerIntegrityCheck
|
||||||
|
|
||||||
|
export type MarketIntegrityMutation = ReturnType<
|
||||||
|
typeof useTriggerIntegrityCheck
|
||||||
|
>
|
||||||
|
export type MarketIntegrityQueryResult = ReturnType<typeof useIntegrityCheck>
|
||||||
|
export type MarketIntegrityAccepted = IntegrityCheckAccepted
|
||||||
|
export type MarketIntegrityData = MarketIntegrityCheck
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
export type MarketIntegrityStatus =
|
||||||
|
"no_data" | "running" | "passed" | "issues_found" | "failed"
|
||||||
|
|
||||||
|
export interface MarketIntegrityQuery {
|
||||||
|
page: number
|
||||||
|
pageSize: number
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface IntegrityCheckAccepted {
|
||||||
|
check_id: string
|
||||||
|
status: "running"
|
||||||
|
window_start: string
|
||||||
|
window_end: string
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface IntegrityIssue {
|
||||||
|
issue_key: string
|
||||||
|
item_kind: string
|
||||||
|
item_key: string
|
||||||
|
issue_type: string
|
||||||
|
message: string
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface MarketIntegrityCheck {
|
||||||
|
check_id: string | null
|
||||||
|
status: MarketIntegrityStatus
|
||||||
|
window_start: string | null
|
||||||
|
window_end: string | null
|
||||||
|
target_count: number
|
||||||
|
checked_count: number
|
||||||
|
issue_count: number
|
||||||
|
error_type: string | null
|
||||||
|
error_message: string | null
|
||||||
|
created_at: string | null
|
||||||
|
finished_at: string | null
|
||||||
|
page: number
|
||||||
|
page_size: number
|
||||||
|
issues_total: number
|
||||||
|
issues: IntegrityIssue[]
|
||||||
|
}
|
||||||
|
|
||||||
|
export const defaultMarketIntegrityQuery: MarketIntegrityQuery = {
|
||||||
|
page: 1,
|
||||||
|
pageSize: 10,
|
||||||
|
}
|
||||||
@@ -0,0 +1,214 @@
|
|||||||
|
import { fireEvent, render, screen, waitFor } from "@testing-library/react"
|
||||||
|
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||||
|
|
||||||
|
import { ApiError } from "@/shared/api/request-json"
|
||||||
|
|
||||||
|
import { SyncPage } from "./sync-page"
|
||||||
|
import type {
|
||||||
|
IntegrityIssue,
|
||||||
|
MarketIntegrityCheck,
|
||||||
|
MarketIntegrityStatus,
|
||||||
|
} from "../api/sync.types"
|
||||||
|
|
||||||
|
const syncHooks = vi.hoisted(() => ({
|
||||||
|
useIntegrityCheck: vi.fn(),
|
||||||
|
useLatestIntegrityCheck: vi.fn(),
|
||||||
|
useTriggerIntegrityCheck: vi.fn(),
|
||||||
|
}))
|
||||||
|
|
||||||
|
vi.mock("../api/sync.query", () => syncHooks)
|
||||||
|
|
||||||
|
const issue: IntegrityIssue = {
|
||||||
|
issue_key: "issue-1",
|
||||||
|
item_kind: "bar",
|
||||||
|
item_key: "000001.SZ:2026-08-08",
|
||||||
|
issue_type: "missing_csv_row",
|
||||||
|
message: "仅报告 PostgreSQL 与 CSV 的差异,不会自动修复。",
|
||||||
|
}
|
||||||
|
|
||||||
|
function buildReport(
|
||||||
|
status: MarketIntegrityStatus,
|
||||||
|
overrides: Partial<MarketIntegrityCheck> = {},
|
||||||
|
): MarketIntegrityCheck {
|
||||||
|
const hasRun = status !== "no_data"
|
||||||
|
return {
|
||||||
|
check_id: hasRun ? "check-1" : null,
|
||||||
|
status,
|
||||||
|
window_start: hasRun ? "2026-08-01" : null,
|
||||||
|
window_end: hasRun ? "2026-08-08" : null,
|
||||||
|
target_count: 10,
|
||||||
|
checked_count: 10,
|
||||||
|
issue_count: status === "issues_found" ? 1 : 0,
|
||||||
|
error_type: null,
|
||||||
|
error_message: null,
|
||||||
|
created_at: "2026-08-08T09:00:00+08:00",
|
||||||
|
finished_at: "2026-08-08T09:02:00+08:00",
|
||||||
|
page: 1,
|
||||||
|
page_size: 10,
|
||||||
|
issues_total: status === "issues_found" ? 1 : 0,
|
||||||
|
issues: status === "issues_found" ? [issue] : [],
|
||||||
|
...overrides,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function configureHooks(
|
||||||
|
latestData: MarketIntegrityCheck | undefined,
|
||||||
|
options: {
|
||||||
|
checkData?: MarketIntegrityCheck
|
||||||
|
latestError?: unknown
|
||||||
|
triggerError?: unknown
|
||||||
|
} = {},
|
||||||
|
) {
|
||||||
|
const refetchLatest = vi.fn()
|
||||||
|
const refetchCheck = vi.fn()
|
||||||
|
const mutate = vi.fn()
|
||||||
|
const reset = vi.fn()
|
||||||
|
|
||||||
|
syncHooks.useLatestIntegrityCheck.mockReturnValue({
|
||||||
|
data: latestData,
|
||||||
|
error: options.latestError,
|
||||||
|
isError: Boolean(options.latestError),
|
||||||
|
isPending: false,
|
||||||
|
refetch: refetchLatest,
|
||||||
|
})
|
||||||
|
syncHooks.useIntegrityCheck.mockReturnValue({
|
||||||
|
data: options.checkData,
|
||||||
|
error: undefined,
|
||||||
|
isError: false,
|
||||||
|
isPending: false,
|
||||||
|
refetch: refetchCheck,
|
||||||
|
})
|
||||||
|
syncHooks.useTriggerIntegrityCheck.mockReturnValue({
|
||||||
|
data: undefined,
|
||||||
|
error: options.triggerError,
|
||||||
|
isError: Boolean(options.triggerError),
|
||||||
|
isPending: false,
|
||||||
|
mutate,
|
||||||
|
reset,
|
||||||
|
})
|
||||||
|
|
||||||
|
return { mutate, refetchCheck, refetchLatest, reset }
|
||||||
|
}
|
||||||
|
|
||||||
|
describe("SyncPage", () => {
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.clearAllMocks()
|
||||||
|
configureHooks(buildReport("no_data"))
|
||||||
|
})
|
||||||
|
|
||||||
|
it("explains the no-data state and starts a check once", () => {
|
||||||
|
const { mutate } = configureHooks(buildReport("no_data"))
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
expect(screen.getByText("暂无检查记录")).toBeInTheDocument()
|
||||||
|
expect(screen.getAllByText(/只报告,不自动修复/).length).toBeGreaterThan(0)
|
||||||
|
|
||||||
|
fireEvent.click(screen.getByRole("button", { name: "开始完整性检查" }))
|
||||||
|
|
||||||
|
expect(mutate).toHaveBeenCalledOnce()
|
||||||
|
})
|
||||||
|
|
||||||
|
it("shows running progress and disables duplicate triggers", () => {
|
||||||
|
configureHooks(buildReport("running", { checked_count: 4 }))
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
expect(screen.getByText("完整性检查进行中")).toBeInTheDocument()
|
||||||
|
expect(screen.getByText("4 / 10")).toBeInTheDocument()
|
||||||
|
expect(screen.getByRole("button", { name: "检查进行中" })).toBeDisabled()
|
||||||
|
})
|
||||||
|
|
||||||
|
it("recovers a running check id from the latest report after refresh", async () => {
|
||||||
|
configureHooks(buildReport("running"), {
|
||||||
|
checkData: buildReport("running", { checked_count: 3 }),
|
||||||
|
})
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
await waitFor(() => {
|
||||||
|
expect(syncHooks.useIntegrityCheck).toHaveBeenCalledWith("check-1", {
|
||||||
|
page: 1,
|
||||||
|
pageSize: 10,
|
||||||
|
})
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
it("refreshes latest once after a running check reaches a terminal state", async () => {
|
||||||
|
const { refetchLatest } = configureHooks(buildReport("running"), {
|
||||||
|
checkData: buildReport("passed"),
|
||||||
|
})
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
await waitFor(() => {
|
||||||
|
expect(refetchLatest).toHaveBeenCalledOnce()
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
it("renders a passed report without an issue table", () => {
|
||||||
|
configureHooks(buildReport("passed"))
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
expect(screen.getByText("检查通过")).toBeInTheDocument()
|
||||||
|
expect(screen.getAllByText(/未发现 PostgreSQL\/CSV 不一致/)).toHaveLength(2)
|
||||||
|
expect(screen.queryByRole("table")).not.toBeInTheDocument()
|
||||||
|
})
|
||||||
|
|
||||||
|
it("renders paginated issue details and the safe-report notice", () => {
|
||||||
|
configureHooks(buildReport("issues_found"))
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
expect(
|
||||||
|
screen.getByRole("table", { name: "完整性检查问题报告" }),
|
||||||
|
).toBeInTheDocument()
|
||||||
|
expect(screen.getByText("bar")).toBeInTheDocument()
|
||||||
|
expect(screen.getByText("000001.SZ:2026-08-08")).toBeInTheDocument()
|
||||||
|
expect(screen.getByText("missing_csv_row")).toBeInTheDocument()
|
||||||
|
expect(screen.getByText(/只展示对象和差异说明/)).toBeInTheDocument()
|
||||||
|
expect(
|
||||||
|
screen.getByRole("contentinfo", { name: "表格分页" }),
|
||||||
|
).toBeInTheDocument()
|
||||||
|
})
|
||||||
|
|
||||||
|
it("shows a failed report and offers a retry entry", () => {
|
||||||
|
const { mutate } = configureHooks(
|
||||||
|
buildReport("failed", {
|
||||||
|
error_type: "check_error",
|
||||||
|
error_message: "存储暂时不可用",
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
expect(screen.getByText("检查未完成")).toBeInTheDocument()
|
||||||
|
expect(screen.getByText(/不能据此判断市场数据已经损坏/)).toBeInTheDocument()
|
||||||
|
fireEvent.click(screen.getByRole("button", { name: "重新发起检查" }))
|
||||||
|
expect(mutate).toHaveBeenCalledOnce()
|
||||||
|
})
|
||||||
|
|
||||||
|
it("makes 503 and network failures visible with a retry button", () => {
|
||||||
|
const { refetchLatest } = configureHooks(undefined, {
|
||||||
|
latestError: new ApiError(503, "unavailable"),
|
||||||
|
})
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
expect(screen.getByText("完整性检查结果暂时不可用")).toBeInTheDocument()
|
||||||
|
fireEvent.click(screen.getByRole("button", { name: "重试" }))
|
||||||
|
expect(refetchLatest).toHaveBeenCalledOnce()
|
||||||
|
})
|
||||||
|
|
||||||
|
it("explains a 409 and lets the page take over the running report", () => {
|
||||||
|
configureHooks(buildReport("running"), {
|
||||||
|
triggerError: new ApiError(409, "conflict"),
|
||||||
|
})
|
||||||
|
|
||||||
|
render(<SyncPage />)
|
||||||
|
|
||||||
|
expect(screen.getByText(/已刷新并接管该检查/)).toBeInTheDocument()
|
||||||
|
expect(screen.getByText("完整性检查进行中")).toBeInTheDocument()
|
||||||
|
})
|
||||||
|
})
|
||||||
@@ -0,0 +1,665 @@
|
|||||||
|
import {
|
||||||
|
AlertTriangle,
|
||||||
|
CheckCircle2,
|
||||||
|
ClipboardCheck,
|
||||||
|
Database,
|
||||||
|
RefreshCw,
|
||||||
|
ShieldCheck,
|
||||||
|
} from "lucide-react"
|
||||||
|
import { useEffect, useRef, useState } from "react"
|
||||||
|
|
||||||
|
import { PageLayout } from "@/app/layout/page-layout"
|
||||||
|
import { ApiError } from "@/shared/api/request-json"
|
||||||
|
import { Badge } from "@/shared/ui/badge"
|
||||||
|
import { Button } from "@/shared/ui/button"
|
||||||
|
import {
|
||||||
|
Card,
|
||||||
|
CardContent,
|
||||||
|
CardDescription,
|
||||||
|
CardHeader,
|
||||||
|
CardTitle,
|
||||||
|
} from "@/shared/ui/card"
|
||||||
|
import { Pagination } from "@/shared/ui/pagination"
|
||||||
|
import { Progress, ProgressLabel, ProgressValue } from "@/shared/ui/progress"
|
||||||
|
|
||||||
|
import {
|
||||||
|
useIntegrityCheck,
|
||||||
|
useLatestIntegrityCheck,
|
||||||
|
useTriggerIntegrityCheck,
|
||||||
|
} from "../api/sync.query"
|
||||||
|
import {
|
||||||
|
defaultMarketIntegrityQuery,
|
||||||
|
type IntegrityCheckAccepted,
|
||||||
|
type MarketIntegrityCheck,
|
||||||
|
} from "../api/sync.types"
|
||||||
|
|
||||||
|
export function SyncPage() {
|
||||||
|
const [triggeredCheckId, setTriggeredCheckId] = useState<string | null>(null)
|
||||||
|
const [page, setPage] = useState(defaultMarketIntegrityQuery.page)
|
||||||
|
const [pageSize, setPageSize] = useState(defaultMarketIntegrityQuery.pageSize)
|
||||||
|
const query = { page, pageSize }
|
||||||
|
const latest = useLatestIntegrityCheck(query)
|
||||||
|
const latestRunningCheckId =
|
||||||
|
latest.data?.status === "running" ? latest.data.check_id : null
|
||||||
|
const activeCheckId = triggeredCheckId ?? latestRunningCheckId
|
||||||
|
const check = useIntegrityCheck(activeCheckId, query)
|
||||||
|
const trigger = useTriggerIntegrityCheck()
|
||||||
|
const { refetch: refetchLatest } = latest
|
||||||
|
const terminalRefreshId = useRef<string | null>(null)
|
||||||
|
|
||||||
|
useEffect(() => {
|
||||||
|
if (!activeCheckId || !check.data || check.data.status === "running") {
|
||||||
|
if (!activeCheckId) terminalRefreshId.current = null
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if (terminalRefreshId.current === activeCheckId) return
|
||||||
|
|
||||||
|
terminalRefreshId.current = activeCheckId
|
||||||
|
void refetchLatest()
|
||||||
|
}, [activeCheckId, check.data, refetchLatest])
|
||||||
|
|
||||||
|
const acceptedReport = trigger.data
|
||||||
|
? acceptedCheckAsReport(trigger.data)
|
||||||
|
: undefined
|
||||||
|
const latestMatchesActive =
|
||||||
|
!activeCheckId || latest.data?.check_id === activeCheckId
|
||||||
|
const report =
|
||||||
|
check.data ??
|
||||||
|
(activeCheckId && latestMatchesActive ? latest.data : undefined) ??
|
||||||
|
(trigger.data ? acceptedReport : latest.data)
|
||||||
|
|
||||||
|
function startCheck() {
|
||||||
|
setPage(1)
|
||||||
|
trigger.reset()
|
||||||
|
trigger.mutate(
|
||||||
|
{},
|
||||||
|
{
|
||||||
|
onError: (error) => {
|
||||||
|
if (getErrorStatus(error) === 409) setTriggeredCheckId(null)
|
||||||
|
},
|
||||||
|
onSuccess: (accepted) => setTriggeredCheckId(accepted.check_id),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function retryActiveCheck() {
|
||||||
|
void check.refetch()
|
||||||
|
}
|
||||||
|
|
||||||
|
function retryLatest() {
|
||||||
|
void latest.refetch()
|
||||||
|
}
|
||||||
|
|
||||||
|
const triggerStatus = getErrorStatus(trigger.error)
|
||||||
|
const triggerError =
|
||||||
|
trigger.isError && triggerStatus !== 409 ? trigger.error : null
|
||||||
|
const queryPending = activeCheckId ? check.isPending : latest.isPending
|
||||||
|
const queryError = activeCheckId ? check.error : latest.error
|
||||||
|
const queryHasData = Boolean(report)
|
||||||
|
|
||||||
|
return (
|
||||||
|
<PageLayout>
|
||||||
|
<div className="mx-auto w-full max-w-6xl space-y-4">
|
||||||
|
<header className="space-y-1">
|
||||||
|
<p className="text-xs font-medium tracking-[0.08em] text-muted-foreground uppercase">
|
||||||
|
只读数据校验
|
||||||
|
</p>
|
||||||
|
<h1 className="text-2xl font-semibold tracking-tight">
|
||||||
|
市场数据完整性检查
|
||||||
|
</h1>
|
||||||
|
<p className="max-w-3xl text-sm text-muted-foreground">
|
||||||
|
检查 PostgreSQL 与 CSV 快照中的市场数据是否一致,帮助定位数据问题。
|
||||||
|
</p>
|
||||||
|
</header>
|
||||||
|
|
||||||
|
<ReadOnlyNotice />
|
||||||
|
|
||||||
|
{triggerStatus === 409 ? <ConflictNotice /> : null}
|
||||||
|
|
||||||
|
{triggerError ? (
|
||||||
|
<RequestError
|
||||||
|
description={getTriggerErrorDescription(triggerError)}
|
||||||
|
onRetry={startCheck}
|
||||||
|
title="完整性检查暂时无法发起"
|
||||||
|
/>
|
||||||
|
) : null}
|
||||||
|
|
||||||
|
{!queryHasData && queryPending ? <LoadingState /> : null}
|
||||||
|
|
||||||
|
{!queryPending && queryError ? (
|
||||||
|
<RequestError
|
||||||
|
description={getQueryErrorDescription(queryError)}
|
||||||
|
onRetry={activeCheckId ? retryActiveCheck : retryLatest}
|
||||||
|
title="完整性检查结果暂时不可用"
|
||||||
|
/>
|
||||||
|
) : null}
|
||||||
|
|
||||||
|
{report?.status === "no_data" ? (
|
||||||
|
<NoDataState isPending={trigger.isPending} onStart={startCheck} />
|
||||||
|
) : null}
|
||||||
|
|
||||||
|
{report?.status === "running" ? <RunningState report={report} /> : null}
|
||||||
|
|
||||||
|
{report?.status === "passed" ? (
|
||||||
|
<PassedState
|
||||||
|
isPending={trigger.isPending}
|
||||||
|
onStart={startCheck}
|
||||||
|
report={report}
|
||||||
|
/>
|
||||||
|
) : null}
|
||||||
|
|
||||||
|
{report?.status === "issues_found" ? (
|
||||||
|
<IssuesFoundState
|
||||||
|
onPageChange={setPage}
|
||||||
|
onPageSizeChange={(nextPageSize) => {
|
||||||
|
setPage(1)
|
||||||
|
setPageSize(nextPageSize)
|
||||||
|
}}
|
||||||
|
report={report}
|
||||||
|
selectedPage={page}
|
||||||
|
selectedPageSize={pageSize}
|
||||||
|
isPending={trigger.isPending}
|
||||||
|
onStart={startCheck}
|
||||||
|
/>
|
||||||
|
) : null}
|
||||||
|
|
||||||
|
{report?.status === "failed" ? (
|
||||||
|
<FailedState
|
||||||
|
isPending={trigger.isPending}
|
||||||
|
onStart={startCheck}
|
||||||
|
report={report}
|
||||||
|
/>
|
||||||
|
) : null}
|
||||||
|
</div>
|
||||||
|
</PageLayout>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function ReadOnlyNotice() {
|
||||||
|
return (
|
||||||
|
<Card className="border-primary/20 bg-primary/5">
|
||||||
|
<CardContent className="flex gap-3 p-4">
|
||||||
|
<ShieldCheck
|
||||||
|
className="mt-0.5 size-5 shrink-0 text-primary"
|
||||||
|
aria-hidden="true"
|
||||||
|
/>
|
||||||
|
<div className="space-y-1">
|
||||||
|
<p className="font-medium">只读检查,安全报告</p>
|
||||||
|
<p className="text-sm text-muted-foreground">
|
||||||
|
本检查只比较 PostgreSQL 与 CSV,不访问
|
||||||
|
Tushare,不修改市场事实。发现差异时只报告,不自动修复。
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function ConflictNotice() {
|
||||||
|
return (
|
||||||
|
<div
|
||||||
|
role="status"
|
||||||
|
className="flex items-start gap-2 rounded-md border border-warning/40 bg-warning/10 p-3 text-sm text-foreground"
|
||||||
|
>
|
||||||
|
<RefreshCw
|
||||||
|
className="mt-0.5 size-4 shrink-0 text-warning"
|
||||||
|
aria-hidden="true"
|
||||||
|
/>
|
||||||
|
<p>已有完整性检查正在运行,页面已刷新并接管该检查的进度。</p>
|
||||||
|
</div>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function LoadingState() {
|
||||||
|
return (
|
||||||
|
<Card aria-label="正在加载完整性检查" role="status">
|
||||||
|
<CardHeader className="gap-3">
|
||||||
|
<div className="h-6 w-52 animate-pulse rounded-md bg-muted" />
|
||||||
|
<div className="h-4 w-80 max-w-full animate-pulse rounded-md bg-muted" />
|
||||||
|
</CardHeader>
|
||||||
|
<CardContent>
|
||||||
|
<div className="h-20 animate-pulse rounded-md bg-muted" />
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function NoDataState({
|
||||||
|
isPending,
|
||||||
|
onStart,
|
||||||
|
}: {
|
||||||
|
isPending: boolean
|
||||||
|
onStart: () => void
|
||||||
|
}) {
|
||||||
|
return (
|
||||||
|
<Card>
|
||||||
|
<CardHeader>
|
||||||
|
<CardTitle className="flex items-center gap-2 text-xl">
|
||||||
|
<Database className="size-5 text-primary" aria-hidden="true" />
|
||||||
|
暂无检查记录
|
||||||
|
</CardTitle>
|
||||||
|
<CardDescription>
|
||||||
|
当前没有可查看的完整性检查。发起后,系统会比较最近完成同步窗口中的
|
||||||
|
PostgreSQL 与 CSV 数据。
|
||||||
|
</CardDescription>
|
||||||
|
</CardHeader>
|
||||||
|
<CardContent className="space-y-3">
|
||||||
|
<Button disabled={isPending} onClick={onStart}>
|
||||||
|
<ClipboardCheck aria-hidden="true" />
|
||||||
|
{isPending ? "正在发起检查…" : "开始完整性检查"}
|
||||||
|
</Button>
|
||||||
|
<p className="text-xs text-muted-foreground">
|
||||||
|
检查只读数据并生成报告,不会访问 Tushare,也不会自动修复。
|
||||||
|
</p>
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function RunningState({ report }: { report: MarketIntegrityCheck }) {
|
||||||
|
const percentage = progressPercentage(
|
||||||
|
report.checked_count,
|
||||||
|
report.target_count,
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
<Card aria-live="polite">
|
||||||
|
<CardHeader className="gap-3">
|
||||||
|
<div className="flex flex-wrap items-start justify-between gap-3">
|
||||||
|
<div>
|
||||||
|
<CardTitle className="flex items-center gap-2 text-xl">
|
||||||
|
<ClipboardCheck
|
||||||
|
className="size-5 text-primary"
|
||||||
|
aria-hidden="true"
|
||||||
|
/>
|
||||||
|
完整性检查进行中
|
||||||
|
</CardTitle>
|
||||||
|
<CardDescription>
|
||||||
|
检查完成后会自动停止刷新;运行期间不会修改 PostgreSQL 或 CSV
|
||||||
|
数据。
|
||||||
|
</CardDescription>
|
||||||
|
</div>
|
||||||
|
<Badge variant="outline">运行中</Badge>
|
||||||
|
</div>
|
||||||
|
</CardHeader>
|
||||||
|
<CardContent className="space-y-5">
|
||||||
|
<Progress aria-label="完整性检查进度" max={100} value={percentage}>
|
||||||
|
<ProgressLabel>检查进度</ProgressLabel>
|
||||||
|
<ProgressValue>
|
||||||
|
{() => `${report.checked_count}/${report.target_count}`}
|
||||||
|
</ProgressValue>
|
||||||
|
</Progress>
|
||||||
|
<div className="grid gap-3 sm:grid-cols-3">
|
||||||
|
<Metric
|
||||||
|
label="已检查对象"
|
||||||
|
value={`${report.checked_count} / ${report.target_count}`}
|
||||||
|
/>
|
||||||
|
<Metric label="当前问题数" value={String(report.issue_count)} />
|
||||||
|
<Metric label="检查窗口" value={formatWindow(report)} />
|
||||||
|
</div>
|
||||||
|
<div className="flex flex-wrap items-center justify-between gap-3 border-t border-border/60 pt-4">
|
||||||
|
<Timestamp label="开始时间" value={report.created_at} />
|
||||||
|
<Button disabled variant="outline">
|
||||||
|
检查进行中
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function PassedState({
|
||||||
|
isPending,
|
||||||
|
onStart,
|
||||||
|
report,
|
||||||
|
}: {
|
||||||
|
isPending: boolean
|
||||||
|
onStart: () => void
|
||||||
|
report: MarketIntegrityCheck
|
||||||
|
}) {
|
||||||
|
return (
|
||||||
|
<Card>
|
||||||
|
<CardHeader className="gap-3">
|
||||||
|
<div className="flex flex-wrap items-start justify-between gap-3">
|
||||||
|
<div>
|
||||||
|
<CardTitle className="flex items-center gap-2 text-xl">
|
||||||
|
<CheckCircle2
|
||||||
|
className="size-5 text-success"
|
||||||
|
aria-hidden="true"
|
||||||
|
/>
|
||||||
|
检查通过
|
||||||
|
</CardTitle>
|
||||||
|
<CardDescription>
|
||||||
|
未发现 PostgreSQL/CSV 不一致,本次检查只报告结果,没有修改数据。
|
||||||
|
</CardDescription>
|
||||||
|
</div>
|
||||||
|
<Badge className="border-transparent bg-success text-success-foreground">
|
||||||
|
已通过
|
||||||
|
</Badge>
|
||||||
|
</div>
|
||||||
|
</CardHeader>
|
||||||
|
<CardContent className="space-y-4">
|
||||||
|
<div className="rounded-md border border-success/30 bg-success/10 p-4 text-sm">
|
||||||
|
未发现 PostgreSQL/CSV 不一致。
|
||||||
|
</div>
|
||||||
|
<div className="grid gap-3 sm:grid-cols-3">
|
||||||
|
<Metric
|
||||||
|
label="已检查对象"
|
||||||
|
value={`${report.checked_count} / ${report.target_count}`}
|
||||||
|
/>
|
||||||
|
<Metric label="问题数" value={String(report.issue_count)} />
|
||||||
|
<Metric label="检查窗口" value={formatWindow(report)} />
|
||||||
|
</div>
|
||||||
|
<div className="flex flex-wrap items-end justify-between gap-3 border-t border-border/60 pt-4">
|
||||||
|
<div className="flex flex-wrap gap-x-5 gap-y-2">
|
||||||
|
<Timestamp label="开始时间" value={report.created_at} />
|
||||||
|
<Timestamp label="完成时间" value={report.finished_at} />
|
||||||
|
</div>
|
||||||
|
<Button disabled={isPending} onClick={onStart} variant="outline">
|
||||||
|
<RefreshCw aria-hidden="true" />
|
||||||
|
{isPending ? "正在发起…" : "再次检查"}
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function IssuesFoundState({
|
||||||
|
isPending,
|
||||||
|
onPageChange,
|
||||||
|
onPageSizeChange,
|
||||||
|
onStart,
|
||||||
|
report,
|
||||||
|
selectedPage,
|
||||||
|
selectedPageSize,
|
||||||
|
}: {
|
||||||
|
isPending: boolean
|
||||||
|
onPageChange: (page: number) => void
|
||||||
|
onPageSizeChange: (pageSize: number) => void
|
||||||
|
onStart: () => void
|
||||||
|
report: MarketIntegrityCheck
|
||||||
|
selectedPage: number
|
||||||
|
selectedPageSize: number
|
||||||
|
}) {
|
||||||
|
return (
|
||||||
|
<Card>
|
||||||
|
<CardHeader className="gap-3">
|
||||||
|
<div className="flex flex-wrap items-start justify-between gap-3">
|
||||||
|
<div>
|
||||||
|
<CardTitle className="flex items-center gap-2 text-xl">
|
||||||
|
<AlertTriangle
|
||||||
|
className="size-5 text-warning"
|
||||||
|
aria-hidden="true"
|
||||||
|
/>
|
||||||
|
发现数据差异
|
||||||
|
</CardTitle>
|
||||||
|
<CardDescription>
|
||||||
|
检查完成并发现 {report.issue_count}{" "}
|
||||||
|
个问题。以下内容仅供定位,系统不会自动修复。
|
||||||
|
</CardDescription>
|
||||||
|
</div>
|
||||||
|
<Badge className="border-transparent bg-warning text-warning-foreground">
|
||||||
|
{report.issue_count} 个问题
|
||||||
|
</Badge>
|
||||||
|
</div>
|
||||||
|
</CardHeader>
|
||||||
|
<CardContent className="space-y-4">
|
||||||
|
<div className="grid gap-3 sm:grid-cols-3">
|
||||||
|
<Metric
|
||||||
|
label="已检查对象"
|
||||||
|
value={`${report.checked_count} / ${report.target_count}`}
|
||||||
|
/>
|
||||||
|
<Metric label="问题数" value={String(report.issue_count)} />
|
||||||
|
<Metric label="检查窗口" value={formatWindow(report)} />
|
||||||
|
</div>
|
||||||
|
<div className="rounded-md border border-warning/40 bg-warning/10 p-4 text-sm">
|
||||||
|
安全说明:问题报告只展示对象和差异说明,不会执行修复或重新同步。
|
||||||
|
</div>
|
||||||
|
<IssueTable report={report} />
|
||||||
|
<Pagination
|
||||||
|
onPageChange={onPageChange}
|
||||||
|
onPageSizeChange={onPageSizeChange}
|
||||||
|
page={selectedPage}
|
||||||
|
pageSize={selectedPageSize}
|
||||||
|
pageSizeOptions={[10, 20, 50]}
|
||||||
|
total={report.issues_total}
|
||||||
|
/>
|
||||||
|
<div className="flex flex-wrap items-end justify-between gap-3 border-t border-border/60 pt-4">
|
||||||
|
<div className="flex flex-wrap gap-x-5 gap-y-2">
|
||||||
|
<Timestamp label="开始时间" value={report.created_at} />
|
||||||
|
<Timestamp label="完成时间" value={report.finished_at} />
|
||||||
|
</div>
|
||||||
|
<Button disabled={isPending} onClick={onStart} variant="outline">
|
||||||
|
<RefreshCw aria-hidden="true" />
|
||||||
|
{isPending ? "正在发起…" : "再次检查"}
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function IssueTable({ report }: { report: MarketIntegrityCheck }) {
|
||||||
|
if (report.issues.length === 0) {
|
||||||
|
return (
|
||||||
|
<div className="rounded-md border p-4 text-sm text-muted-foreground">
|
||||||
|
当前页没有问题记录。
|
||||||
|
</div>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
return (
|
||||||
|
<div className="overflow-x-auto rounded-md border">
|
||||||
|
<table className="w-full min-w-[680px] text-left text-sm">
|
||||||
|
<caption className="sr-only">完整性检查问题报告</caption>
|
||||||
|
<thead className="bg-muted/50 text-xs text-muted-foreground">
|
||||||
|
<tr>
|
||||||
|
<th className="px-3 py-2.5 font-medium" scope="col">
|
||||||
|
对象类型
|
||||||
|
</th>
|
||||||
|
<th className="px-3 py-2.5 font-medium" scope="col">
|
||||||
|
对象 key
|
||||||
|
</th>
|
||||||
|
<th className="px-3 py-2.5 font-medium" scope="col">
|
||||||
|
问题类型
|
||||||
|
</th>
|
||||||
|
<th className="px-3 py-2.5 font-medium" scope="col">
|
||||||
|
安全说明
|
||||||
|
</th>
|
||||||
|
</tr>
|
||||||
|
</thead>
|
||||||
|
<tbody className="divide-y divide-border/60">
|
||||||
|
{report.issues.map((issue) => (
|
||||||
|
<tr key={issue.issue_key} className="align-top">
|
||||||
|
<td className="px-3 py-3 font-medium">{issue.item_kind}</td>
|
||||||
|
<td className="break-all px-3 py-3 font-mono text-xs">
|
||||||
|
{issue.item_key}
|
||||||
|
</td>
|
||||||
|
<td className="px-3 py-3">{issue.issue_type}</td>
|
||||||
|
<td className="max-w-md px-3 py-3 text-muted-foreground">
|
||||||
|
{issue.message}
|
||||||
|
</td>
|
||||||
|
</tr>
|
||||||
|
))}
|
||||||
|
</tbody>
|
||||||
|
</table>
|
||||||
|
</div>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function FailedState({
|
||||||
|
isPending,
|
||||||
|
onStart,
|
||||||
|
report,
|
||||||
|
}: {
|
||||||
|
isPending: boolean
|
||||||
|
onStart: () => void
|
||||||
|
report: MarketIntegrityCheck
|
||||||
|
}) {
|
||||||
|
return (
|
||||||
|
<Card>
|
||||||
|
<CardHeader className="gap-3">
|
||||||
|
<div className="flex flex-wrap items-start justify-between gap-3">
|
||||||
|
<div>
|
||||||
|
<CardTitle className="flex items-center gap-2 text-xl">
|
||||||
|
<AlertTriangle
|
||||||
|
className="size-5 text-destructive"
|
||||||
|
aria-hidden="true"
|
||||||
|
/>
|
||||||
|
检查未完成
|
||||||
|
</CardTitle>
|
||||||
|
<CardDescription>
|
||||||
|
本次检查遇到安全错误,不能据此判断市场数据已经损坏。
|
||||||
|
</CardDescription>
|
||||||
|
</div>
|
||||||
|
<Badge variant="destructive">失败</Badge>
|
||||||
|
</div>
|
||||||
|
</CardHeader>
|
||||||
|
<CardContent className="space-y-4">
|
||||||
|
<div
|
||||||
|
role="alert"
|
||||||
|
className="rounded-md border border-destructive/30 bg-destructive/10 p-4 text-sm"
|
||||||
|
>
|
||||||
|
<p className="font-medium">未判断市场数据状态</p>
|
||||||
|
<p className="mt-1 text-muted-foreground">
|
||||||
|
{report.error_type
|
||||||
|
? `错误类型:${report.error_type}`
|
||||||
|
: "检查服务暂时无法完成比较。"}
|
||||||
|
</p>
|
||||||
|
{report.error_message ? (
|
||||||
|
<p className="mt-1 text-muted-foreground">{report.error_message}</p>
|
||||||
|
) : null}
|
||||||
|
</div>
|
||||||
|
<div className="grid gap-3 sm:grid-cols-3">
|
||||||
|
<Metric
|
||||||
|
label="已检查对象"
|
||||||
|
value={`${report.checked_count} / ${report.target_count}`}
|
||||||
|
/>
|
||||||
|
<Metric label="问题数" value={String(report.issue_count)} />
|
||||||
|
<Metric label="检查窗口" value={formatWindow(report)} />
|
||||||
|
</div>
|
||||||
|
<div className="flex flex-wrap items-end justify-between gap-3 border-t border-border/60 pt-4">
|
||||||
|
<div className="flex flex-wrap gap-x-5 gap-y-2">
|
||||||
|
<Timestamp label="开始时间" value={report.created_at} />
|
||||||
|
<Timestamp label="结束时间" value={report.finished_at} />
|
||||||
|
</div>
|
||||||
|
<Button disabled={isPending} onClick={onStart}>
|
||||||
|
<RefreshCw aria-hidden="true" />
|
||||||
|
{isPending ? "正在重新发起…" : "重新发起检查"}
|
||||||
|
</Button>
|
||||||
|
</div>
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function RequestError({
|
||||||
|
description,
|
||||||
|
onRetry,
|
||||||
|
title,
|
||||||
|
}: {
|
||||||
|
description: string
|
||||||
|
onRetry: () => void
|
||||||
|
title: string
|
||||||
|
}) {
|
||||||
|
return (
|
||||||
|
<Card role="alert">
|
||||||
|
<CardHeader>
|
||||||
|
<CardTitle className="flex items-center gap-2 text-xl">
|
||||||
|
<AlertTriangle
|
||||||
|
className="size-5 text-destructive"
|
||||||
|
aria-hidden="true"
|
||||||
|
/>
|
||||||
|
{title}
|
||||||
|
</CardTitle>
|
||||||
|
<CardDescription>{description}</CardDescription>
|
||||||
|
</CardHeader>
|
||||||
|
<CardContent>
|
||||||
|
<Button onClick={onRetry} variant="outline">
|
||||||
|
<RefreshCw aria-hidden="true" />
|
||||||
|
重试
|
||||||
|
</Button>
|
||||||
|
</CardContent>
|
||||||
|
</Card>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function Metric({ label, value }: { label: string; value: string }) {
|
||||||
|
return (
|
||||||
|
<div className="rounded-md border bg-muted/30 p-3">
|
||||||
|
<p className="text-xs text-muted-foreground">{label}</p>
|
||||||
|
<p className="mt-1 break-words font-medium tabular-nums">{value}</p>
|
||||||
|
</div>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function Timestamp({ label, value }: { label: string; value: string | null }) {
|
||||||
|
return (
|
||||||
|
<p className="text-xs text-muted-foreground">
|
||||||
|
{label}:
|
||||||
|
<time className="text-foreground" dateTime={value ?? undefined}>
|
||||||
|
{value ?? "—"}
|
||||||
|
</time>
|
||||||
|
</p>
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
function formatWindow(report: MarketIntegrityCheck) {
|
||||||
|
if (!report.window_start || !report.window_end) return "—"
|
||||||
|
return `${report.window_start} 至 ${report.window_end}`
|
||||||
|
}
|
||||||
|
|
||||||
|
function progressPercentage(checked: number, target: number) {
|
||||||
|
if (target <= 0) return 0
|
||||||
|
return Math.min(100, Math.max(0, (checked / target) * 100))
|
||||||
|
}
|
||||||
|
|
||||||
|
function acceptedCheckAsReport(
|
||||||
|
accepted: IntegrityCheckAccepted,
|
||||||
|
): MarketIntegrityCheck {
|
||||||
|
return {
|
||||||
|
check_id: accepted.check_id,
|
||||||
|
status: "running",
|
||||||
|
window_start: accepted.window_start,
|
||||||
|
window_end: accepted.window_end,
|
||||||
|
target_count: 0,
|
||||||
|
checked_count: 0,
|
||||||
|
issue_count: 0,
|
||||||
|
error_type: null,
|
||||||
|
error_message: null,
|
||||||
|
created_at: null,
|
||||||
|
finished_at: null,
|
||||||
|
page: 1,
|
||||||
|
page_size: defaultMarketIntegrityQuery.pageSize,
|
||||||
|
issues_total: 0,
|
||||||
|
issues: [],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function getErrorStatus(error: unknown) {
|
||||||
|
if (error instanceof ApiError) return error.status
|
||||||
|
if (typeof error === "object" && error !== null && "status" in error) {
|
||||||
|
const status = error.status
|
||||||
|
return typeof status === "number" ? status : null
|
||||||
|
}
|
||||||
|
return null
|
||||||
|
}
|
||||||
|
|
||||||
|
function getQueryErrorDescription(error: unknown) {
|
||||||
|
const status = getErrorStatus(error)
|
||||||
|
if (status === 503) {
|
||||||
|
return "检查结果存储暂时不可用,未将这次请求解释为市场数据损坏。"
|
||||||
|
}
|
||||||
|
return "检查结果暂时无法获取,可能是网络或服务异常。请稍后重试。"
|
||||||
|
}
|
||||||
|
|
||||||
|
function getTriggerErrorDescription(error: unknown) {
|
||||||
|
const status = getErrorStatus(error)
|
||||||
|
if (status === 422) {
|
||||||
|
return "当前没有可供比较的已完成市场数据窗口,请先完成一次市场数据同步。"
|
||||||
|
}
|
||||||
|
if (status === 503) {
|
||||||
|
return "检查服务暂时不可用,未发起新的检查。请稍后重试。"
|
||||||
|
}
|
||||||
|
return "检查请求可能没有到达服务,请确认网络后重试。"
|
||||||
|
}
|
||||||
@@ -8,6 +8,7 @@ import {
|
|||||||
type SelectionCategoryFilter,
|
type SelectionCategoryFilter,
|
||||||
} from "@/features/selection/api/selection.types"
|
} from "@/features/selection/api/selection.types"
|
||||||
import { SelectionResultsPage } from "@/features/selection/pages/selection-results-page"
|
import { SelectionResultsPage } from "@/features/selection/pages/selection-results-page"
|
||||||
|
import { SyncPage } from "@/features/sync/pages/sync-page"
|
||||||
|
|
||||||
const rootRoute = createRootRoute({
|
const rootRoute = createRootRoute({
|
||||||
component: () => <Outlet />,
|
component: () => <Outlet />,
|
||||||
@@ -48,6 +49,12 @@ const selectionRoute = createRoute({
|
|||||||
component: SelectionResultsPage,
|
component: SelectionResultsPage,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
const syncRoute = createRoute({
|
||||||
|
getParentRoute: () => workspaceRoute,
|
||||||
|
path: "/sync",
|
||||||
|
component: SyncPage,
|
||||||
|
})
|
||||||
|
|
||||||
const componentsRoute = createRoute({
|
const componentsRoute = createRoute({
|
||||||
getParentRoute: () => workspaceRoute,
|
getParentRoute: () => workspaceRoute,
|
||||||
path: "/components",
|
path: "/components",
|
||||||
@@ -55,5 +62,10 @@ const componentsRoute = createRoute({
|
|||||||
})
|
})
|
||||||
|
|
||||||
export const routeTree = rootRoute.addChildren([
|
export const routeTree = rootRoute.addChildren([
|
||||||
workspaceRoute.addChildren([indexRoute, selectionRoute, componentsRoute]),
|
workspaceRoute.addChildren([
|
||||||
|
indexRoute,
|
||||||
|
selectionRoute,
|
||||||
|
syncRoute,
|
||||||
|
componentsRoute,
|
||||||
|
]),
|
||||||
])
|
])
|
||||||
|
|||||||
Reference in New Issue
Block a user