feat(market-data): 迁移 Tushare 同步到 PostgreSQL
This commit is contained in:
@@ -1,6 +1,8 @@
|
||||
"""Environment-backed application configuration."""
|
||||
|
||||
from decimal import Decimal
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
@@ -12,6 +14,15 @@ class Settings(BaseSettings):
|
||||
app_name: str = "Zhixing Server"
|
||||
app_env: Literal["development", "test", "production"] = "development"
|
||||
log_level: str = "INFO"
|
||||
database_url: str = "postgresql://zhixing:zhixing@localhost:5432/zhixing"
|
||||
tushare_token: str = ""
|
||||
market_data_csv_root: Path = Path("./data/market-data")
|
||||
market_data_coverage_threshold: Decimal = Decimal("0.99")
|
||||
market_data_max_workers: int = 4
|
||||
market_data_request_interval_seconds: float = 0.2
|
||||
market_data_max_retries: int = 3
|
||||
market_data_retry_backoff_seconds: float = 1.0
|
||||
market_data_advisory_lock_key: int = 7_380_521
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Market data bounded context.
|
||||
|
||||
The context owns the current A-share universe, daily qfq bars, daily basic
|
||||
metrics, and the externally-triggered synchronization use case.
|
||||
"""
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Market data synchronization use cases."""
|
||||
|
||||
from .sync import SyncBatchSummary, SyncMarketData, SyncMarketDataCommand
|
||||
|
||||
__all__ = ["SyncBatchSummary", "SyncMarketData", "SyncMarketDataCommand"]
|
||||
@@ -0,0 +1,403 @@
|
||||
"""One-shot market data synchronization orchestration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
|
||||
from ..domain.fingerprint import SnapshotChange, compare_snapshots
|
||||
from ..domain.models import Bar, Stock, SyncWindow
|
||||
from ..domain.ports import MarketDataRepository, MarketDataSource, SnapshotStore, WriteResult
|
||||
from ..domain.rules import filter_current_hs_a_stocks
|
||||
|
||||
SyncMode = Literal["daily", "initialize", "retry"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SyncMarketDataCommand:
|
||||
"""Input for a daily, initial, or failed-item retry synchronization."""
|
||||
|
||||
mode: SyncMode = "daily"
|
||||
target_trade_date: date | None = None
|
||||
parent_batch_id: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SyncFailure:
|
||||
"""A redacted, actionable item failure."""
|
||||
|
||||
item_kind: str
|
||||
item_key: str
|
||||
error_type: str
|
||||
message: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SyncBatchSummary:
|
||||
"""Stable output contract for CLI, cron, and later strategy callers."""
|
||||
|
||||
batch_id: str | None
|
||||
target_trade_date: date | None
|
||||
window: SyncWindow | None
|
||||
status: str
|
||||
target_count: int
|
||||
valid_count: int
|
||||
coverage: Decimal
|
||||
strategy_eligible: bool
|
||||
inserted_count: int = 0
|
||||
updated_count: int = 0
|
||||
unchanged_count: int = 0
|
||||
failures: tuple[SyncFailure, ...] = field(default_factory=tuple)
|
||||
|
||||
@property
|
||||
def exit_code(self) -> int:
|
||||
"""Return a cron-friendly code: success, incomplete, or infrastructure failure."""
|
||||
|
||||
if self.status == "failed":
|
||||
return 1
|
||||
return 0 if self.strategy_eligible else 2
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
"""Serialize the summary without credentials or raw vendor responses."""
|
||||
|
||||
return {
|
||||
"batch_id": self.batch_id,
|
||||
"target_trade_date": self.target_trade_date.isoformat()
|
||||
if self.target_trade_date
|
||||
else None,
|
||||
"window_start": self.window.start.isoformat() if self.window else None,
|
||||
"window_end": self.window.end.isoformat() if self.window else None,
|
||||
"status": self.status,
|
||||
"target_count": self.target_count,
|
||||
"valid_count": self.valid_count,
|
||||
"coverage": str(self.coverage),
|
||||
"strategy_eligible": self.strategy_eligible,
|
||||
"inserted_count": self.inserted_count,
|
||||
"updated_count": self.updated_count,
|
||||
"unchanged_count": self.unchanged_count,
|
||||
"failures": [
|
||||
{
|
||||
"item_kind": failure.item_kind,
|
||||
"item_key": failure.item_key,
|
||||
"error_type": failure.error_type,
|
||||
"message": failure.message,
|
||||
}
|
||||
for failure in self.failures
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class SyncMarketData:
|
||||
"""Coordinate source, snapshot, and repository ports at item boundaries."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
source: MarketDataSource,
|
||||
snapshots: SnapshotStore,
|
||||
repository: MarketDataRepository,
|
||||
*,
|
||||
coverage_threshold: Decimal = Decimal("0.99"),
|
||||
lock_key: int = 7_380_521,
|
||||
today: date | None = None,
|
||||
) -> None:
|
||||
if not Decimal("0") <= coverage_threshold <= Decimal("1"):
|
||||
raise ValueError("coverage_threshold must be between 0 and 1")
|
||||
self.source = source
|
||||
self.snapshots = snapshots
|
||||
self.repository = repository
|
||||
self.coverage_threshold = coverage_threshold
|
||||
self.lock_key = lock_key
|
||||
self.today = today or date.today()
|
||||
|
||||
def execute(self, command: SyncMarketDataCommand | None = None) -> SyncBatchSummary:
|
||||
"""Run one synchronization and retain successful items on partial failure."""
|
||||
|
||||
command = command or SyncMarketDataCommand()
|
||||
with self.repository.advisory_lock(self.lock_key) as acquired:
|
||||
if not acquired:
|
||||
return SyncBatchSummary(
|
||||
batch_id=None,
|
||||
target_trade_date=None,
|
||||
window=None,
|
||||
status="failed",
|
||||
target_count=0,
|
||||
valid_count=0,
|
||||
coverage=Decimal("0"),
|
||||
strategy_eligible=False,
|
||||
failures=(
|
||||
SyncFailure("batch", "lock", "sync_locked", "another sync is running"),
|
||||
),
|
||||
)
|
||||
return self._execute_locked(command)
|
||||
|
||||
def _execute_locked(self, command: SyncMarketDataCommand) -> SyncBatchSummary:
|
||||
target_trade_date = self._resolve_target(command.target_trade_date)
|
||||
window = SyncWindow.from_target(target_trade_date)
|
||||
all_stocks = filter_current_hs_a_stocks(self.source.fetch_stocks())
|
||||
if not all_stocks:
|
||||
return SyncBatchSummary(
|
||||
None,
|
||||
target_trade_date,
|
||||
window,
|
||||
"failed",
|
||||
0,
|
||||
0,
|
||||
Decimal("0"),
|
||||
False,
|
||||
failures=(
|
||||
SyncFailure("stock", "universe", "empty_universe", "no eligible stocks"),
|
||||
),
|
||||
)
|
||||
batch_id = self.repository.create_batch(
|
||||
target_trade_date,
|
||||
window,
|
||||
command.mode,
|
||||
command.parent_batch_id,
|
||||
len(all_stocks),
|
||||
)
|
||||
failures: list[SyncFailure] = []
|
||||
totals = [0, 0, 0]
|
||||
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)
|
||||
dates = self._dates_to_process(window, target_trade_date, command.mode, retry_items)
|
||||
for trade_date in dates:
|
||||
if (
|
||||
command.mode == "retry"
|
||||
and ("daily_basic", trade_date.isoformat()) not in retry_items
|
||||
):
|
||||
continue
|
||||
self._process_daily_basic(batch_id, trade_date, stock_codes, window, failures, totals)
|
||||
for stock in all_stocks:
|
||||
if command.mode == "retry" and ("bar", stock.ts_code) not in retry_items:
|
||||
continue
|
||||
self._process_bar(batch_id, stock, window, failures, totals)
|
||||
|
||||
if not failures:
|
||||
try:
|
||||
self.repository.purge_before(window)
|
||||
self.snapshots.clean_daily_basic_before(window.start)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
failure = self._failure("batch", "retention", exc)
|
||||
failures.append(failure)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
"batch",
|
||||
"retention",
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
)
|
||||
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)
|
||||
)
|
||||
coverage = Decimal(valid_count) / Decimal(len(all_stocks))
|
||||
status = "success" if not failures else "partial_success" if valid_count else "failed"
|
||||
eligible = coverage >= self.coverage_threshold
|
||||
self.repository.record_batch(batch_id, status, valid_count, coverage, eligible)
|
||||
return SyncBatchSummary(
|
||||
batch_id,
|
||||
target_trade_date,
|
||||
window,
|
||||
status,
|
||||
len(all_stocks),
|
||||
valid_count,
|
||||
coverage,
|
||||
eligible,
|
||||
totals[0],
|
||||
totals[1],
|
||||
totals[2],
|
||||
tuple(failures),
|
||||
)
|
||||
|
||||
def _resolve_target(self, requested: date | None) -> date:
|
||||
if requested is not None:
|
||||
open_dates = self.source.fetch_open_dates(requested, requested)
|
||||
if requested not in open_dates:
|
||||
raise ValueError(f"target date is not an open trading day: {requested}")
|
||||
return requested
|
||||
start = self.today - timedelta(days=14)
|
||||
open_dates = self.source.fetch_open_dates(start, self.today)
|
||||
if not open_dates:
|
||||
raise ValueError("no open trading day found before today")
|
||||
return max(open_dates)
|
||||
|
||||
def _dates_to_process(
|
||||
self,
|
||||
window: SyncWindow,
|
||||
target: date,
|
||||
mode: SyncMode,
|
||||
retry_items: set[tuple[str, str]],
|
||||
) -> tuple[date, ...]:
|
||||
if mode == "initialize":
|
||||
dates = self.source.fetch_open_dates(window.start, target)
|
||||
return tuple(sorted(set(dates)))
|
||||
if mode == "retry":
|
||||
return tuple(
|
||||
sorted(
|
||||
date.fromisoformat(key) for kind, key in retry_items if kind == "daily_basic"
|
||||
)
|
||||
)
|
||||
return (target,)
|
||||
|
||||
def _retry_items(self, parent_batch_id: str | None) -> set[tuple[str, str]]:
|
||||
if not parent_batch_id:
|
||||
raise ValueError("retry mode requires parent_batch_id")
|
||||
return set(self.repository.failed_items(parent_batch_id))
|
||||
|
||||
def _process_stock_master(
|
||||
self,
|
||||
batch_id: str,
|
||||
stocks: Sequence[Stock],
|
||||
failures: list[SyncFailure],
|
||||
totals: list[int],
|
||||
) -> None:
|
||||
staged = None
|
||||
try:
|
||||
staged = self.snapshots.stage_stocks(stocks)
|
||||
result = self.repository.upsert_stocks(stocks)
|
||||
self.snapshots.publish(staged)
|
||||
totals[0] += result.inserted
|
||||
totals[1] += result.updated
|
||||
totals[2] += result.unchanged
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
"stock",
|
||||
"current",
|
||||
"success",
|
||||
result,
|
||||
staged.fingerprint,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
if staged is not None:
|
||||
self.snapshots.discard(staged)
|
||||
failure = self._failure("stock", "current", exc)
|
||||
failures.append(failure)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
"stock",
|
||||
"current",
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
)
|
||||
|
||||
def _process_daily_basic(
|
||||
self,
|
||||
batch_id: str,
|
||||
trade_date: date,
|
||||
stock_codes: set[str],
|
||||
window: SyncWindow,
|
||||
failures: list[SyncFailure],
|
||||
totals: list[int],
|
||||
) -> None:
|
||||
staged = None
|
||||
key = trade_date.isoformat()
|
||||
try:
|
||||
rows = tuple(
|
||||
row
|
||||
for row in self.source.fetch_daily_basic(trade_date)
|
||||
if row.ts_code in stock_codes
|
||||
)
|
||||
if not rows:
|
||||
raise ValueError(f"daily-basic returned no target rows for {trade_date}")
|
||||
staged = self.snapshots.stage_daily_basic(trade_date, rows)
|
||||
result = self.repository.upsert_daily_basic(rows, window)
|
||||
self.snapshots.publish(staged)
|
||||
self._add_counts(totals, result)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
"daily_basic",
|
||||
key,
|
||||
"success",
|
||||
result,
|
||||
staged.fingerprint,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
if staged is not None:
|
||||
self.snapshots.discard(staged)
|
||||
failure = self._failure("daily_basic", key, exc)
|
||||
failures.append(failure)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
"daily_basic",
|
||||
key,
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
)
|
||||
|
||||
def _process_bar(
|
||||
self,
|
||||
batch_id: str,
|
||||
stock: Stock,
|
||||
window: SyncWindow,
|
||||
failures: list[SyncFailure],
|
||||
totals: list[int],
|
||||
) -> None:
|
||||
staged = None
|
||||
try:
|
||||
rows = tuple(self.source.fetch_bars(stock.ts_code, window))
|
||||
old_rows = self.snapshots.read_bars(stock.ts_code)
|
||||
comparison = compare_snapshots(old_rows, rows)
|
||||
staged = self.snapshots.stage_bars(stock.ts_code, rows)
|
||||
if comparison.change is SnapshotChange.UNCHANGED:
|
||||
result = WriteResult(unchanged=len(rows))
|
||||
else:
|
||||
to_write: Iterable[Bar] = rows
|
||||
if comparison.change is SnapshotChange.NEW_DATES:
|
||||
old_end = max((row.trade_date for row in old_rows or ()), default=window.start)
|
||||
to_write = tuple(row for row in rows if row.trade_date > old_end)
|
||||
result = self.repository.upsert_bars(
|
||||
to_write,
|
||||
window,
|
||||
full_snapshot=comparison.change
|
||||
in {SnapshotChange.INITIAL, SnapshotChange.CHANGED},
|
||||
)
|
||||
self.snapshots.publish(staged)
|
||||
self._add_counts(totals, result)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
"bar",
|
||||
stock.ts_code,
|
||||
"success",
|
||||
result,
|
||||
staged.fingerprint,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError, ValueError) as exc:
|
||||
if staged is not None:
|
||||
self.snapshots.discard(staged)
|
||||
failure = self._failure("bar", stock.ts_code, exc)
|
||||
failures.append(failure)
|
||||
self.repository.record_item(
|
||||
batch_id,
|
||||
"bar",
|
||||
stock.ts_code,
|
||||
"failed",
|
||||
WriteResult(),
|
||||
error_type=failure.error_type,
|
||||
error_message=failure.message,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _add_counts(totals: list[int], result: WriteResult) -> None:
|
||||
totals[0] += result.inserted
|
||||
totals[1] += result.updated
|
||||
totals[2] += result.unchanged
|
||||
|
||||
@staticmethod
|
||||
def _failure(item_kind: str, item_key: str, error: BaseException) -> SyncFailure:
|
||||
message = " ".join(str(error).split())[:500] or "synchronization item failed"
|
||||
return SyncFailure(item_kind, item_key, type(error).__name__, message)
|
||||
@@ -0,0 +1,18 @@
|
||||
"""Pure market data domain types and ports."""
|
||||
|
||||
from .fingerprint import SnapshotChange, SnapshotComparison, compare_snapshots, snapshot_fingerprint
|
||||
from .models import Bar, DailyBasic, Stock, SyncWindow
|
||||
from .rules import filter_current_hs_a_stocks, is_current_hs_a_stock
|
||||
|
||||
__all__ = [
|
||||
"Bar",
|
||||
"DailyBasic",
|
||||
"SnapshotChange",
|
||||
"SnapshotComparison",
|
||||
"Stock",
|
||||
"SyncWindow",
|
||||
"compare_snapshots",
|
||||
"filter_current_hs_a_stocks",
|
||||
"is_current_hs_a_stock",
|
||||
"snapshot_fingerprint",
|
||||
]
|
||||
@@ -0,0 +1,150 @@
|
||||
"""Stable qfq snapshot fingerprinting and change classification."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from enum import StrEnum
|
||||
from hashlib import sha256
|
||||
|
||||
from .models import Bar, decimal_text
|
||||
|
||||
|
||||
class SnapshotChange(StrEnum):
|
||||
"""Action required after comparing a previous and current bar snapshot."""
|
||||
|
||||
INITIAL = "initial"
|
||||
UNCHANGED = "unchanged"
|
||||
NEW_DATES = "new_dates"
|
||||
CHANGED = "changed"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SnapshotComparison:
|
||||
"""Result of comparing normalized snapshots for one stock."""
|
||||
|
||||
change: SnapshotChange
|
||||
old_fingerprint: str | None
|
||||
new_fingerprint: str
|
||||
new_dates: tuple[date, ...]
|
||||
|
||||
|
||||
FINGERPRINT_FIELDS = (
|
||||
"trade_date",
|
||||
"open",
|
||||
"high",
|
||||
"low",
|
||||
"close",
|
||||
"pre_close",
|
||||
"change",
|
||||
"pct_chg",
|
||||
"vol",
|
||||
"amount",
|
||||
)
|
||||
|
||||
|
||||
def _validate(rows: Iterable[Bar]) -> tuple[Bar, ...]:
|
||||
"""Sort rows and reject duplicate identity keys before hashing or writing."""
|
||||
|
||||
sorted_rows = tuple(sorted(rows, key=lambda row: (row.ts_code, row.trade_date)))
|
||||
seen: set[tuple[str, date]] = set()
|
||||
for row in sorted_rows:
|
||||
key = (row.ts_code, row.trade_date)
|
||||
if not row.ts_code or key in seen:
|
||||
raise ValueError(f"duplicate or empty bar key: {key!r}")
|
||||
seen.add(key)
|
||||
for value in (row.open, row.high, row.low, row.close):
|
||||
if value is not None and value < 0:
|
||||
raise ValueError(f"OHLC values must be non-negative: {key!r}")
|
||||
if row.high is not None and row.low is not None and row.high < row.low:
|
||||
raise ValueError(f"high must not be lower than low: {key!r}")
|
||||
return sorted_rows
|
||||
|
||||
|
||||
def _canonical(row: Bar) -> str:
|
||||
values = (
|
||||
row.trade_date.isoformat(),
|
||||
decimal_text(row.open),
|
||||
decimal_text(row.high),
|
||||
decimal_text(row.low),
|
||||
decimal_text(row.close),
|
||||
decimal_text(row.pre_close),
|
||||
decimal_text(row.change),
|
||||
decimal_text(row.pct_chg),
|
||||
decimal_text(row.vol),
|
||||
decimal_text(row.amount),
|
||||
)
|
||||
return ",".join(values)
|
||||
|
||||
|
||||
def snapshot_fingerprint(rows: Iterable[Bar]) -> str:
|
||||
"""Hash sorted, normalized fixed bar columns, excluding runtime metadata."""
|
||||
|
||||
normalized = _validate(rows)
|
||||
payload = "\n".join(_canonical(row) for row in normalized).encode("utf-8")
|
||||
return sha256(payload).hexdigest()
|
||||
|
||||
|
||||
def _overlap_fingerprint(rows: dict[date, Bar], dates: set[date]) -> str:
|
||||
payload = "\n".join(_canonical(rows[day]) for day in sorted(dates)).encode("utf-8")
|
||||
return sha256(payload).hexdigest()
|
||||
|
||||
|
||||
def compare_snapshots(
|
||||
old_rows: Iterable[Bar] | None,
|
||||
new_rows: Iterable[Bar],
|
||||
) -> SnapshotComparison:
|
||||
"""Compare a complete old/new six-year snapshot.
|
||||
|
||||
A shifted rolling start is expected and remains unchanged. Missing rows
|
||||
or changed values inside the common historical range are repairs, while
|
||||
rows strictly after the old upper bound are ordinary new dates.
|
||||
"""
|
||||
|
||||
current = _validate(new_rows)
|
||||
if not current:
|
||||
raise ValueError("a bar snapshot must contain at least one row")
|
||||
new_by_date = {row.trade_date: row for row in current}
|
||||
new_dates = set(new_by_date)
|
||||
new_fingerprint = snapshot_fingerprint(current)
|
||||
if old_rows is None:
|
||||
return SnapshotComparison(
|
||||
SnapshotChange.INITIAL,
|
||||
None,
|
||||
new_fingerprint,
|
||||
tuple(sorted(new_dates)),
|
||||
)
|
||||
|
||||
previous = _validate(old_rows)
|
||||
if not previous:
|
||||
return SnapshotComparison(
|
||||
SnapshotChange.INITIAL,
|
||||
None,
|
||||
new_fingerprint,
|
||||
tuple(sorted(new_dates)),
|
||||
)
|
||||
old_by_date = {row.trade_date: row for row in previous}
|
||||
old_dates = set(old_by_date)
|
||||
overlap_start = max(min(old_dates), min(new_dates))
|
||||
overlap_end = min(max(old_dates), max(new_dates))
|
||||
overlap = {day for day in new_dates if overlap_start <= day <= overlap_end}
|
||||
old_overlap = {day for day in old_dates if overlap_start <= day <= overlap_end}
|
||||
old_fp = _overlap_fingerprint(old_by_date, old_overlap)
|
||||
new_fp = _overlap_fingerprint(new_by_date, overlap)
|
||||
|
||||
starts_earlier = min(new_dates) < min(old_dates)
|
||||
changed = (
|
||||
old_overlap != overlap
|
||||
or old_fp != new_fp
|
||||
or max(new_dates) < max(old_dates)
|
||||
or starts_earlier
|
||||
)
|
||||
if changed:
|
||||
kind = SnapshotChange.CHANGED
|
||||
elif max(new_dates) > max(old_dates):
|
||||
kind = SnapshotChange.NEW_DATES
|
||||
else:
|
||||
kind = SnapshotChange.UNCHANGED
|
||||
additions = tuple(sorted(day for day in new_dates if day > max(old_dates)))
|
||||
return SnapshotComparison(kind, old_fp, new_fp, additions)
|
||||
@@ -0,0 +1,199 @@
|
||||
"""Market data entities and stable scalar normalization helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from decimal import Decimal, InvalidOperation
|
||||
|
||||
|
||||
def parse_date(value: object) -> date:
|
||||
"""Parse a Tushare date value without accepting locale-dependent formats."""
|
||||
|
||||
if isinstance(value, date):
|
||||
return value
|
||||
text = str(value).strip()
|
||||
if len(text) == 8 and text.isdigit():
|
||||
return date.fromisoformat(f"{text[:4]}-{text[4:6]}-{text[6:]}")
|
||||
return date.fromisoformat(text[:10])
|
||||
|
||||
|
||||
def normalize_decimal(value: object | None) -> Decimal | None:
|
||||
"""Convert vendor numbers to finite, text-stable decimals.
|
||||
|
||||
Tushare may return a float, a decimal, a string, or ``None`` depending on
|
||||
the transport. Decimal constructed from the textual representation keeps
|
||||
those variants from changing snapshot fingerprints.
|
||||
"""
|
||||
|
||||
if value is None or str(value).strip() == "":
|
||||
return None
|
||||
try:
|
||||
result = Decimal(str(value).strip())
|
||||
except (InvalidOperation, ValueError) as exc:
|
||||
raise ValueError(f"invalid numeric value: {value!r}") from exc
|
||||
if not result.is_finite():
|
||||
raise ValueError(f"numeric value must be finite: {value!r}")
|
||||
return Decimal(0) if result == 0 else result.normalize()
|
||||
|
||||
|
||||
def decimal_text(value: Decimal | None) -> str:
|
||||
"""Serialize a decimal in a locale-independent, non-exponent form."""
|
||||
|
||||
if value is None:
|
||||
return ""
|
||||
text = format(value, "f")
|
||||
if "." in text:
|
||||
text = text.rstrip("0").rstrip(".")
|
||||
return text or "0"
|
||||
|
||||
|
||||
def _text(row: Mapping[str, object], key: str, default: str = "") -> str:
|
||||
value = row.get(key, default)
|
||||
return default if value is None else str(value).strip()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Stock:
|
||||
"""Current stock master data used to construct the target universe."""
|
||||
|
||||
ts_code: str
|
||||
name: str
|
||||
market: str = ""
|
||||
exchange: str = ""
|
||||
list_status: str = "L"
|
||||
list_date: date | None = None
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, object]) -> Stock:
|
||||
"""Build a stock from a ``stock_basic`` row."""
|
||||
|
||||
raw_date = row.get("list_date")
|
||||
return cls(
|
||||
ts_code=_text(row, "ts_code"),
|
||||
name=_text(row, "name"),
|
||||
market=_text(row, "market"),
|
||||
exchange=_text(row, "exchange"),
|
||||
list_status=_text(row, "list_status", "L").upper(),
|
||||
list_date=parse_date(raw_date) if raw_date else None,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Bar:
|
||||
"""One qfq daily OHLCV record."""
|
||||
|
||||
ts_code: str
|
||||
trade_date: date
|
||||
open: Decimal | None = None
|
||||
high: Decimal | None = None
|
||||
low: Decimal | None = None
|
||||
close: Decimal | None = None
|
||||
pre_close: Decimal | None = None
|
||||
change: Decimal | None = None
|
||||
pct_chg: Decimal | None = None
|
||||
vol: Decimal | None = None
|
||||
amount: Decimal | None = None
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, object]) -> Bar:
|
||||
"""Build a normalized bar from a Tushare ``pro_bar`` row."""
|
||||
|
||||
return cls(
|
||||
ts_code=_text(row, "ts_code"),
|
||||
trade_date=parse_date(row["trade_date"]),
|
||||
open=normalize_decimal(row.get("open")),
|
||||
high=normalize_decimal(row.get("high")),
|
||||
low=normalize_decimal(row.get("low")),
|
||||
close=normalize_decimal(row.get("close")),
|
||||
pre_close=normalize_decimal(row.get("pre_close")),
|
||||
change=normalize_decimal(row.get("change")),
|
||||
pct_chg=normalize_decimal(row.get("pct_chg")),
|
||||
vol=normalize_decimal(row.get("vol")),
|
||||
amount=normalize_decimal(row.get("amount")),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DailyBasic:
|
||||
"""One trading-day valuation and liquidity snapshot."""
|
||||
|
||||
ts_code: str
|
||||
trade_date: date
|
||||
close: Decimal | None = None
|
||||
turnover_rate: Decimal | None = None
|
||||
turnover_rate_f: Decimal | None = None
|
||||
volume_ratio: Decimal | None = None
|
||||
pe: Decimal | None = None
|
||||
pe_ttm: Decimal | None = None
|
||||
pb: Decimal | None = None
|
||||
ps: Decimal | None = None
|
||||
ps_ttm: Decimal | None = None
|
||||
dv_ratio: Decimal | None = None
|
||||
dv_ttm: Decimal | None = None
|
||||
total_share: Decimal | None = None
|
||||
float_share: Decimal | None = None
|
||||
free_share: Decimal | None = None
|
||||
total_mv: Decimal | None = None
|
||||
circ_mv: Decimal | None = None
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, object]) -> DailyBasic:
|
||||
"""Build a normalized row from a Tushare ``daily_basic`` response."""
|
||||
|
||||
decimal_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: normalize_decimal(row.get(field)) for field in decimal_fields}
|
||||
return cls(
|
||||
ts_code=_text(row, "ts_code"),
|
||||
trade_date=parse_date(row["trade_date"]),
|
||||
**values,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SyncWindow:
|
||||
"""Inclusive rolling window used by both the database and CSV retention."""
|
||||
|
||||
start: date
|
||||
end: date
|
||||
|
||||
@classmethod
|
||||
def from_target(cls, target_trade_date: date, years: int = 6) -> SyncWindow:
|
||||
"""Return a six-calendar-year window including the target date."""
|
||||
|
||||
if years < 1:
|
||||
raise ValueError("years must be positive")
|
||||
try:
|
||||
start = target_trade_date.replace(year=target_trade_date.year - years)
|
||||
except ValueError:
|
||||
# 29 February has no counterpart in a non-leap year. The last
|
||||
# day of February is the deterministic calendar boundary.
|
||||
start = target_trade_date.replace(
|
||||
year=target_trade_date.year - years,
|
||||
month=2,
|
||||
day=28,
|
||||
)
|
||||
return cls(start=start, end=target_trade_date)
|
||||
|
||||
def contains(self, value: date) -> bool:
|
||||
"""Return whether a date is within the inclusive window."""
|
||||
|
||||
return self.start <= value <= self.end
|
||||
@@ -0,0 +1,126 @@
|
||||
"""Ports used by the market data application layer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Sequence
|
||||
from contextlib import AbstractContextManager
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
from typing import Protocol
|
||||
|
||||
from .models import Bar, DailyBasic, Stock, SyncWindow
|
||||
|
||||
|
||||
class MarketDataRepositoryError(RuntimeError):
|
||||
"""A storage adapter failed without exposing vendor or credential details."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WriteResult:
|
||||
"""Counts returned by one idempotent repository write."""
|
||||
|
||||
inserted: int = 0
|
||||
updated: int = 0
|
||||
unchanged: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StagedSnapshot:
|
||||
"""A temporary file that may be published after its database transaction."""
|
||||
|
||||
temporary_path: Path
|
||||
final_path: Path
|
||||
fingerprint: str
|
||||
|
||||
|
||||
class MarketDataSource(Protocol):
|
||||
"""External source port for stock, calendar, bar, and metric data."""
|
||||
|
||||
def fetch_stocks(self) -> Sequence[Stock]: ...
|
||||
|
||||
def fetch_open_dates(self, start: date, end: date) -> Sequence[date]: ...
|
||||
|
||||
def fetch_bars(self, ts_code: str, window: SyncWindow) -> Sequence[Bar]: ...
|
||||
|
||||
def fetch_daily_basic(self, trade_date: date) -> Sequence[DailyBasic]: ...
|
||||
|
||||
|
||||
class SnapshotStore(Protocol):
|
||||
"""Snapshot and atomic-file publication port."""
|
||||
|
||||
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None: ...
|
||||
|
||||
def stage_bars(self, ts_code: str, rows: Iterable[Bar]) -> StagedSnapshot: ...
|
||||
|
||||
def stage_daily_basic(self, trade_date: date, rows: Iterable[DailyBasic]) -> StagedSnapshot: ...
|
||||
|
||||
def stage_stocks(self, rows: Iterable[Stock]) -> StagedSnapshot: ...
|
||||
|
||||
def publish(self, staged: StagedSnapshot) -> None: ...
|
||||
|
||||
def discard(self, staged: StagedSnapshot) -> None: ...
|
||||
|
||||
def clean_daily_basic_before(self, boundary: date) -> int: ...
|
||||
|
||||
|
||||
class MarketDataRepository(Protocol):
|
||||
"""PostgreSQL repository port.
|
||||
|
||||
``full_snapshot`` controls whether rows missing from the current six-year
|
||||
window are deleted for the affected stock. Incremental writes never
|
||||
remove rows before the database transaction has succeeded.
|
||||
"""
|
||||
|
||||
def upsert_stocks(self, rows: Iterable[Stock]) -> WriteResult: ...
|
||||
|
||||
def upsert_bars(
|
||||
self,
|
||||
rows: Iterable[Bar],
|
||||
window: SyncWindow,
|
||||
*,
|
||||
full_snapshot: bool,
|
||||
) -> WriteResult: ...
|
||||
|
||||
def upsert_daily_basic(self, rows: Iterable[DailyBasic], window: SyncWindow) -> WriteResult: ...
|
||||
|
||||
def purge_before(self, window: SyncWindow) -> None: ...
|
||||
|
||||
def has_bar(self, ts_code: str, trade_date: date) -> bool: ...
|
||||
|
||||
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool: ...
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
target_trade_date: date,
|
||||
window: SyncWindow,
|
||||
mode: str,
|
||||
parent_batch_id: str | None,
|
||||
target_count: int,
|
||||
) -> str: ...
|
||||
|
||||
def record_batch(
|
||||
self,
|
||||
batch_id: str,
|
||||
status: str,
|
||||
valid_count: int,
|
||||
coverage: Decimal,
|
||||
strategy_eligible: bool,
|
||||
) -> None: ...
|
||||
|
||||
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: ...
|
||||
|
||||
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]: ...
|
||||
|
||||
def advisory_lock(self, key: int) -> AbstractContextManager[bool]: ...
|
||||
@@ -0,0 +1,32 @@
|
||||
"""Target universe rules for the current listed Shanghai/Shenzhen A shares."""
|
||||
|
||||
from collections.abc import Iterable
|
||||
|
||||
from .models import Stock
|
||||
|
||||
|
||||
def is_current_hs_a_stock(stock: Stock) -> bool:
|
||||
"""Return whether a stock is a current non-ST Shanghai/Shenzhen A share."""
|
||||
|
||||
code, separator, suffix = stock.ts_code.upper().partition(".")
|
||||
if not separator or suffix not in {"SH", "SZ"} or not code.isdigit():
|
||||
return False
|
||||
if stock.list_status.upper() != "L":
|
||||
return False
|
||||
if stock.exchange and stock.exchange.upper() not in {"SSE", "SZSE", "SH", "SZ"}:
|
||||
return False
|
||||
market = stock.market.upper().replace(" ", "")
|
||||
if "北交所" in market or market in {"B", "B股", "BEIJING"}:
|
||||
return False
|
||||
return "ST" not in stock.name.upper().replace(" ", "") and "退" not in stock.name
|
||||
|
||||
|
||||
def filter_current_hs_a_stocks(stocks: Iterable[Stock]) -> tuple[Stock, ...]:
|
||||
"""Filter and deterministically sort the current target universe."""
|
||||
|
||||
return tuple(
|
||||
sorted(
|
||||
(stock for stock in stocks if is_current_hs_a_stock(stock)),
|
||||
key=lambda stock: stock.ts_code,
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Compatibility import for the market data rolling window."""
|
||||
|
||||
from .models import SyncWindow
|
||||
|
||||
__all__ = ["SyncWindow"]
|
||||
@@ -0,0 +1 @@
|
||||
"""Market data infrastructure adapters."""
|
||||
@@ -0,0 +1,256 @@
|
||||
"""Fixed-layout CSV snapshots with atomic publication semantics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import os
|
||||
import re
|
||||
import tempfile
|
||||
from collections.abc import Iterable
|
||||
from contextlib import suppress
|
||||
from datetime import date
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
|
||||
from ..domain.fingerprint import snapshot_fingerprint
|
||||
from ..domain.models import Bar, DailyBasic, Stock, decimal_text
|
||||
from ..domain.ports import StagedSnapshot
|
||||
|
||||
_SAFE_CODE = re.compile(r"^[0-9]{6}\.(SH|SZ)$")
|
||||
BAR_COLUMNS = (
|
||||
"ts_code",
|
||||
"trade_date",
|
||||
"open",
|
||||
"high",
|
||||
"low",
|
||||
"close",
|
||||
"pre_close",
|
||||
"change",
|
||||
"pct_chg",
|
||||
"vol",
|
||||
"amount",
|
||||
)
|
||||
DAILY_BASIC_COLUMNS = (
|
||||
"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",
|
||||
)
|
||||
STOCK_COLUMNS = ("ts_code", "name", "market", "exchange", "list_status", "list_date")
|
||||
_DAILY_DECIMAL_FIELDS = DAILY_BASIC_COLUMNS[2:]
|
||||
|
||||
|
||||
def _code_path(ts_code: str) -> str:
|
||||
if not _SAFE_CODE.fullmatch(ts_code):
|
||||
raise ValueError(f"unsupported stock code: {ts_code!r}")
|
||||
return ts_code
|
||||
|
||||
|
||||
def _bar_row(row: Bar) -> dict[str, str]:
|
||||
return {
|
||||
"ts_code": row.ts_code,
|
||||
"trade_date": row.trade_date.isoformat(),
|
||||
"open": decimal_text(row.open),
|
||||
"high": decimal_text(row.high),
|
||||
"low": decimal_text(row.low),
|
||||
"close": decimal_text(row.close),
|
||||
"pre_close": decimal_text(row.pre_close),
|
||||
"change": decimal_text(row.change),
|
||||
"pct_chg": decimal_text(row.pct_chg),
|
||||
"vol": decimal_text(row.vol),
|
||||
"amount": decimal_text(row.amount),
|
||||
}
|
||||
|
||||
|
||||
def _daily_basic_row(row: DailyBasic) -> dict[str, str]:
|
||||
values: dict[str, str] = {
|
||||
"ts_code": row.ts_code,
|
||||
"trade_date": row.trade_date.isoformat(),
|
||||
}
|
||||
for field in _DAILY_DECIMAL_FIELDS:
|
||||
values[field] = decimal_text(getattr(row, field))
|
||||
return values
|
||||
|
||||
|
||||
def _stock_row(row: Stock) -> dict[str, str]:
|
||||
return {
|
||||
"ts_code": row.ts_code,
|
||||
"name": row.name,
|
||||
"market": row.market,
|
||||
"exchange": row.exchange,
|
||||
"list_status": row.list_status,
|
||||
"list_date": row.list_date.isoformat() if row.list_date else "",
|
||||
}
|
||||
|
||||
|
||||
class CsvSnapshotStore:
|
||||
"""Store and publish the three market-data CSV layouts.
|
||||
|
||||
Temporary files are created beside their final file. Therefore ``os.replace``
|
||||
is a same-filesystem operation and a database failure can leave the previous
|
||||
published snapshot untouched.
|
||||
"""
|
||||
|
||||
def __init__(self, root: Path) -> None:
|
||||
self.root = root
|
||||
|
||||
def bars_path(self, ts_code: str) -> Path:
|
||||
"""Return the formal per-stock qfq path."""
|
||||
|
||||
return self.root / "bars" / f"{_code_path(ts_code)}.csv"
|
||||
|
||||
def daily_basic_path(self, trade_date: date) -> Path:
|
||||
"""Return the formal per-trading-day metrics path."""
|
||||
|
||||
return self.root / "daily-basic" / f"{trade_date:%Y}" / f"{trade_date:%Y%m%d}.csv"
|
||||
|
||||
@property
|
||||
def stock_basic_path(self) -> Path:
|
||||
"""Return the formal current stock-master path."""
|
||||
|
||||
return self.root / "stock-basic" / "current.csv"
|
||||
|
||||
def read_bars(self, ts_code: str) -> tuple[Bar, ...] | None:
|
||||
"""Read the current formal snapshot, returning ``None`` if absent."""
|
||||
|
||||
path = self.bars_path(ts_code)
|
||||
if not path.exists():
|
||||
return None
|
||||
with path.open(newline="", encoding="utf-8") as handle:
|
||||
reader = csv.DictReader(handle)
|
||||
if tuple(reader.fieldnames or ()) != BAR_COLUMNS:
|
||||
raise ValueError(f"unexpected bar CSV header: {path}")
|
||||
rows = tuple(Bar.from_mapping(row) for row in reader)
|
||||
return tuple(sorted(rows, key=lambda row: row.trade_date))
|
||||
|
||||
def stage_bars(self, ts_code: str, rows: Iterable[Bar]) -> StagedSnapshot:
|
||||
"""Validate and write a temporary bar snapshot."""
|
||||
|
||||
normalized = tuple(sorted(rows, key=lambda row: row.trade_date))
|
||||
if not normalized or any(row.ts_code != ts_code for row in normalized):
|
||||
raise ValueError("bar snapshot must be non-empty and belong to one stock")
|
||||
final_path = self.bars_path(ts_code)
|
||||
temp_path = self._write_csv(final_path, BAR_COLUMNS, (_bar_row(row) for row in normalized))
|
||||
return StagedSnapshot(temp_path, final_path, snapshot_fingerprint(normalized))
|
||||
|
||||
def stage_daily_basic(self, trade_date: date, rows: Iterable[DailyBasic]) -> StagedSnapshot:
|
||||
"""Validate and stage one daily-basic date snapshot."""
|
||||
|
||||
normalized = tuple(sorted(rows, key=lambda row: row.ts_code))
|
||||
if any(row.trade_date != trade_date for row in normalized):
|
||||
raise ValueError("daily-basic snapshot contains an unexpected date")
|
||||
keys = [row.ts_code for row in normalized]
|
||||
if len(keys) != len(set(keys)):
|
||||
raise ValueError("daily-basic snapshot contains duplicate stock codes")
|
||||
final_path = self.daily_basic_path(trade_date)
|
||||
payload = "\n".join(
|
||||
",".join(_daily_basic_row(row)[column] for column in DAILY_BASIC_COLUMNS)
|
||||
for row in normalized
|
||||
).encode("utf-8")
|
||||
fingerprint = sha256(payload).hexdigest()
|
||||
temp_path = self._write_csv(
|
||||
final_path,
|
||||
DAILY_BASIC_COLUMNS,
|
||||
(_daily_basic_row(row) for row in normalized),
|
||||
)
|
||||
return StagedSnapshot(temp_path, final_path, fingerprint)
|
||||
|
||||
def stage_stocks(self, rows: Iterable[Stock]) -> StagedSnapshot:
|
||||
"""Validate and stage the current stock-master snapshot."""
|
||||
|
||||
normalized = tuple(sorted(rows, key=lambda row: row.ts_code))
|
||||
if not normalized:
|
||||
raise ValueError("stock-master snapshot must not be empty")
|
||||
keys = [row.ts_code for row in normalized]
|
||||
if len(keys) != len(set(keys)):
|
||||
raise ValueError("stock-master snapshot contains duplicate stock codes")
|
||||
payload = "\n".join(
|
||||
",".join(_stock_row(row)[field] for field in STOCK_COLUMNS) for row in normalized
|
||||
)
|
||||
final_path = self.stock_basic_path
|
||||
temp_path = self._write_csv(
|
||||
final_path,
|
||||
STOCK_COLUMNS,
|
||||
(_stock_row(row) for row in normalized),
|
||||
)
|
||||
return StagedSnapshot(temp_path, final_path, sha256(payload.encode("utf-8")).hexdigest())
|
||||
|
||||
def publish(self, staged: StagedSnapshot) -> None:
|
||||
"""Atomically promote a staged snapshot and sync the parent directory."""
|
||||
|
||||
if not staged.temporary_path.exists():
|
||||
raise FileNotFoundError(staged.temporary_path)
|
||||
os.replace(staged.temporary_path, staged.final_path)
|
||||
try:
|
||||
directory_fd = os.open(staged.final_path.parent, os.O_RDONLY)
|
||||
except OSError:
|
||||
return
|
||||
try:
|
||||
os.fsync(directory_fd)
|
||||
finally:
|
||||
os.close(directory_fd)
|
||||
|
||||
def discard(self, staged: StagedSnapshot) -> None:
|
||||
"""Remove a staged file without touching its formal counterpart."""
|
||||
|
||||
with suppress(FileNotFoundError):
|
||||
staged.temporary_path.unlink()
|
||||
|
||||
def clean_daily_basic_before(self, boundary: date) -> int:
|
||||
"""Remove only formal daily-basic files older than the retention boundary."""
|
||||
|
||||
removed = 0
|
||||
root = self.root / "daily-basic"
|
||||
if not root.exists():
|
||||
return removed
|
||||
for path in root.glob("*/????????.csv"):
|
||||
try:
|
||||
file_date = date.fromisoformat(
|
||||
path.stem[:4] + "-" + path.stem[4:6] + "-" + path.stem[6:]
|
||||
)
|
||||
except ValueError:
|
||||
continue
|
||||
if file_date < boundary:
|
||||
path.unlink()
|
||||
removed += 1
|
||||
return removed
|
||||
|
||||
def _write_csv(
|
||||
self,
|
||||
final_path: Path,
|
||||
columns: tuple[str, ...],
|
||||
rows: Iterable[dict[str, str]],
|
||||
) -> Path:
|
||||
final_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
descriptor, raw_path = tempfile.mkstemp(
|
||||
prefix=f".{final_path.name}.{uuid4().hex}.",
|
||||
suffix=".tmp",
|
||||
dir=final_path.parent,
|
||||
)
|
||||
temp_path = Path(raw_path)
|
||||
try:
|
||||
with os.fdopen(descriptor, "w", encoding="utf-8", newline="") as handle:
|
||||
writer = csv.DictWriter(handle, fieldnames=columns, lineterminator="\n")
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
handle.flush()
|
||||
os.fsync(handle.fileno())
|
||||
except BaseException:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
return temp_path
|
||||
@@ -0,0 +1,465 @@
|
||||
"""Psycopg 3 PostgreSQL adapter using COPY staging and set-based upserts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator, Iterable
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
import psycopg
|
||||
|
||||
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow
|
||||
from ..domain.ports import MarketDataRepositoryError, WriteResult
|
||||
|
||||
|
||||
class PostgresMarketDataRepository:
|
||||
"""Persist market data without an ORM identity map.
|
||||
|
||||
Each public write opens a short transaction. The caller publishes the
|
||||
matching CSV only after this method returns successfully.
|
||||
"""
|
||||
|
||||
def __init__(self, database_url: str) -> None:
|
||||
self.database_url = database_url
|
||||
|
||||
def upsert_stocks(self, rows: Iterable[Stock]) -> WriteResult:
|
||||
records = tuple(rows)
|
||||
if not records:
|
||||
return WriteResult()
|
||||
with self._connection() as connection, connection.transaction():
|
||||
cursor = connection.cursor()
|
||||
codes = [row.ts_code for row in records]
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE market_stock
|
||||
SET is_active = false, updated_at = now()
|
||||
WHERE is_active AND NOT (ts_code = ANY(%s))
|
||||
""",
|
||||
(codes,),
|
||||
)
|
||||
inserted = 0
|
||||
updated = 0
|
||||
unchanged = 0
|
||||
for row in records:
|
||||
result = cursor.execute(
|
||||
"""
|
||||
INSERT INTO market_stock
|
||||
(
|
||||
ts_code, name, market, exchange, list_status, list_date,
|
||||
is_active, updated_at
|
||||
)
|
||||
VALUES (%s, %s, %s, %s, %s, %s, true, now())
|
||||
ON CONFLICT (ts_code) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
market = EXCLUDED.market,
|
||||
exchange = EXCLUDED.exchange,
|
||||
list_status = EXCLUDED.list_status,
|
||||
list_date = EXCLUDED.list_date,
|
||||
is_active = true,
|
||||
updated_at = now()
|
||||
WHERE (market_stock.name, market_stock.market, market_stock.exchange,
|
||||
market_stock.list_status, market_stock.list_date, market_stock.is_active)
|
||||
IS DISTINCT FROM (EXCLUDED.name, EXCLUDED.market, EXCLUDED.exchange,
|
||||
EXCLUDED.list_status, EXCLUDED.list_date, true)
|
||||
RETURNING (xmax = 0) AS inserted
|
||||
""",
|
||||
(
|
||||
row.ts_code,
|
||||
row.name,
|
||||
row.market,
|
||||
row.exchange,
|
||||
row.list_status,
|
||||
row.list_date,
|
||||
),
|
||||
).fetchone()
|
||||
if result is None:
|
||||
unchanged += 1
|
||||
elif bool(result[0]):
|
||||
inserted += 1
|
||||
else:
|
||||
updated += 1
|
||||
return WriteResult(inserted=inserted, updated=updated, unchanged=unchanged)
|
||||
|
||||
def upsert_bars(
|
||||
self,
|
||||
rows: Iterable[Bar],
|
||||
window: SyncWindow,
|
||||
*,
|
||||
full_snapshot: bool,
|
||||
) -> WriteResult:
|
||||
records = tuple(rows)
|
||||
if records and len({row.ts_code for row in records}) != 1:
|
||||
raise ValueError("a bar write must contain one stock")
|
||||
if not records:
|
||||
return WriteResult()
|
||||
code = records[0].ts_code
|
||||
with self._connection() as connection, connection.transaction():
|
||||
cursor = connection.cursor()
|
||||
self._create_bar_stage(cursor)
|
||||
self._copy_bars(cursor, records)
|
||||
before = self._count_stage_matches(cursor, "market_daily_bar")
|
||||
inserted, updated = self._upsert_bars_from_stage(cursor)
|
||||
if full_snapshot:
|
||||
cursor.execute(
|
||||
"""
|
||||
DELETE FROM market_daily_bar AS target
|
||||
WHERE target.ts_code = %s
|
||||
AND target.trade_date BETWEEN %s AND %s
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM market_daily_bar_stage AS stage
|
||||
WHERE stage.ts_code = target.ts_code
|
||||
AND stage.trade_date = target.trade_date
|
||||
)
|
||||
""",
|
||||
(code, window.start, window.end),
|
||||
)
|
||||
return WriteResult(inserted=inserted, updated=updated, unchanged=max(0, before - updated))
|
||||
|
||||
def upsert_daily_basic(self, rows: Iterable[DailyBasic], window: SyncWindow) -> WriteResult:
|
||||
records = tuple(rows)
|
||||
if not records:
|
||||
return WriteResult()
|
||||
if len({row.trade_date for row in records}) != 1:
|
||||
raise ValueError("a daily-basic write must contain one trade date")
|
||||
with self._connection() as connection, connection.transaction():
|
||||
cursor = connection.cursor()
|
||||
self._create_basic_stage(cursor)
|
||||
self._copy_daily_basic(cursor, records)
|
||||
before = self._count_stage_matches(cursor, "market_daily_basic")
|
||||
inserted, updated = self._upsert_daily_basic_from_stage(cursor)
|
||||
return WriteResult(inserted=inserted, updated=updated, unchanged=max(0, before - updated))
|
||||
|
||||
def purge_before(self, window: SyncWindow) -> None:
|
||||
"""Apply rolling retention only after the caller has finished a batch."""
|
||||
|
||||
with self._connection() as connection, connection.transaction():
|
||||
cursor = connection.cursor()
|
||||
cursor.execute("DELETE FROM market_daily_bar WHERE trade_date < %s", (window.start,))
|
||||
cursor.execute("DELETE FROM market_daily_basic WHERE trade_date < %s", (window.start,))
|
||||
|
||||
def has_bar(self, ts_code: str, trade_date: date) -> bool:
|
||||
return self._exists("market_daily_bar", ts_code, trade_date)
|
||||
|
||||
def has_daily_basic(self, ts_code: str, trade_date: date) -> bool:
|
||||
return self._exists("market_daily_basic", ts_code, trade_date)
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
target_trade_date: date,
|
||||
window: SyncWindow,
|
||||
mode: str,
|
||||
parent_batch_id: str | None,
|
||||
target_count: int,
|
||||
) -> str:
|
||||
batch_id = str(uuid4())
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO market_sync_batch
|
||||
(id, target_trade_date, window_start, mode, status, target_count)
|
||||
VALUES (%s, %s, %s, %s, 'running', %s)
|
||||
""",
|
||||
(batch_id, target_trade_date, window.start, mode, target_count),
|
||||
)
|
||||
if parent_batch_id:
|
||||
connection.execute(
|
||||
"UPDATE market_sync_batch SET parent_batch_id = %s WHERE id = %s",
|
||||
(parent_batch_id, batch_id),
|
||||
)
|
||||
return batch_id
|
||||
|
||||
def record_batch(
|
||||
self,
|
||||
batch_id: str,
|
||||
status: str,
|
||||
valid_count: int,
|
||||
coverage: Decimal,
|
||||
strategy_eligible: bool,
|
||||
) -> None:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
UPDATE market_sync_batch
|
||||
SET status = %s, valid_count = %s, coverage = %s,
|
||||
strategy_eligible = %s, finished_at = now()
|
||||
WHERE id = %s
|
||||
""",
|
||||
(status, valid_count, coverage, strategy_eligible, batch_id),
|
||||
)
|
||||
|
||||
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:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
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
|
||||
""",
|
||||
(
|
||||
batch_id,
|
||||
item_kind,
|
||||
item_key,
|
||||
status,
|
||||
result.inserted,
|
||||
result.updated,
|
||||
result.unchanged,
|
||||
fingerprint,
|
||||
error_type,
|
||||
self._safe_error(error_message),
|
||||
),
|
||||
)
|
||||
|
||||
def failed_items(self, parent_batch_id: str) -> tuple[tuple[str, str], ...]:
|
||||
with self._connection() as connection:
|
||||
rows = connection.execute(
|
||||
"""
|
||||
SELECT item_kind, item_key
|
||||
FROM market_sync_item
|
||||
WHERE batch_id = %s AND status = 'failed'
|
||||
ORDER BY item_kind, item_key
|
||||
""",
|
||||
(parent_batch_id,),
|
||||
).fetchall()
|
||||
return tuple((str(row[0]), str(row[1])) for row in rows)
|
||||
|
||||
@contextmanager
|
||||
def advisory_lock(self, key: int) -> Generator[bool, None, None]:
|
||||
"""Hold a PostgreSQL advisory lock for the lifetime of one sync run."""
|
||||
|
||||
with self._connection() as connection:
|
||||
row = connection.execute(
|
||||
"SELECT pg_try_advisory_lock(%s)",
|
||||
(key,),
|
||||
).fetchone()
|
||||
acquired = bool(row[0]) if row is not None else False
|
||||
if not acquired:
|
||||
yield False
|
||||
return
|
||||
try:
|
||||
yield True
|
||||
finally:
|
||||
connection.execute("SELECT pg_advisory_unlock(%s)", (key,))
|
||||
|
||||
def _exists(self, table: str, ts_code: str, trade_date: date) -> bool:
|
||||
if table == "market_daily_bar":
|
||||
query = "SELECT 1 FROM market_daily_bar WHERE ts_code = %s AND trade_date = %s LIMIT 1"
|
||||
elif table == "market_daily_basic":
|
||||
query = (
|
||||
"SELECT 1 FROM market_daily_basic WHERE ts_code = %s AND trade_date = %s LIMIT 1"
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"unsupported existence table: {table}")
|
||||
with self._connection() as connection:
|
||||
result = connection.execute(
|
||||
query,
|
||||
(ts_code, trade_date),
|
||||
).fetchone()
|
||||
return result is not None
|
||||
|
||||
@contextmanager
|
||||
def _connection(self) -> Generator[Any, None, None]:
|
||||
"""Translate driver failures into a safe application-level error."""
|
||||
|
||||
try:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
yield connection
|
||||
except psycopg.Error as exc:
|
||||
raise MarketDataRepositoryError("market data database operation failed") from exc
|
||||
|
||||
@staticmethod
|
||||
def _create_bar_stage(cursor: Any) -> None:
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TEMP TABLE market_daily_bar_stage (
|
||||
ts_code varchar(12) NOT NULL,
|
||||
trade_date date NOT NULL,
|
||||
open numeric(20,6), high numeric(20,6), low numeric(20,6), close numeric(20,6),
|
||||
pre_close numeric(20,6), change numeric(20,6), pct_chg numeric(20,6),
|
||||
vol numeric(24,6), amount numeric(24,6)
|
||||
) ON COMMIT DROP
|
||||
"""
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _create_basic_stage(cursor: Any) -> None:
|
||||
cursor.execute(
|
||||
"""
|
||||
CREATE TEMP TABLE market_daily_basic_stage (
|
||||
ts_code varchar(12) NOT NULL,
|
||||
trade_date date NOT NULL,
|
||||
close numeric(24,6), turnover_rate numeric(24,6),
|
||||
turnover_rate_f numeric(24,6), volume_ratio numeric(24,6),
|
||||
pe numeric(24,6), pe_ttm numeric(24,6), pb numeric(24,6),
|
||||
ps numeric(24,6), ps_ttm numeric(24,6), dv_ratio numeric(24,6),
|
||||
dv_ttm numeric(24,6), total_share numeric(24,6),
|
||||
float_share numeric(24,6), free_share numeric(24,6),
|
||||
total_mv numeric(24,6), circ_mv numeric(24,6)
|
||||
) ON COMMIT DROP
|
||||
"""
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _copy_bars(cursor: Any, rows: Iterable[Bar]) -> None:
|
||||
copy_sql = (
|
||||
"COPY market_daily_bar_stage "
|
||||
"(ts_code, trade_date, open, high, low, close, pre_close, "
|
||||
"change, pct_chg, vol, amount) FROM STDIN"
|
||||
)
|
||||
with cursor.copy(copy_sql) as copy:
|
||||
for row in rows:
|
||||
copy.write_row(
|
||||
(
|
||||
row.ts_code,
|
||||
row.trade_date,
|
||||
row.open,
|
||||
row.high,
|
||||
row.low,
|
||||
row.close,
|
||||
row.pre_close,
|
||||
row.change,
|
||||
row.pct_chg,
|
||||
row.vol,
|
||||
row.amount,
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _copy_daily_basic(cursor: Any, rows: Iterable[DailyBasic]) -> None:
|
||||
copy_sql = (
|
||||
"COPY market_daily_basic_stage "
|
||||
"(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 STDIN"
|
||||
)
|
||||
with cursor.copy(copy_sql) as copy:
|
||||
for row in rows:
|
||||
copy.write_row(
|
||||
(
|
||||
row.ts_code,
|
||||
row.trade_date,
|
||||
row.close,
|
||||
row.turnover_rate,
|
||||
row.turnover_rate_f,
|
||||
row.volume_ratio,
|
||||
row.pe,
|
||||
row.pe_ttm,
|
||||
row.pb,
|
||||
row.ps,
|
||||
row.ps_ttm,
|
||||
row.dv_ratio,
|
||||
row.dv_ttm,
|
||||
row.total_share,
|
||||
row.float_share,
|
||||
row.free_share,
|
||||
row.total_mv,
|
||||
row.circ_mv,
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _upsert_bars_from_stage(cursor: Any) -> tuple[int, int]:
|
||||
rows = cursor.execute(
|
||||
"""
|
||||
INSERT INTO market_daily_bar
|
||||
(
|
||||
ts_code, trade_date, open, high, low, close, pre_close,
|
||||
change, pct_chg, vol, amount, source_adj, updated_at
|
||||
)
|
||||
SELECT
|
||||
ts_code, trade_date, open, high, low, close, pre_close,
|
||||
change, pct_chg, vol, amount, 'qfq', now()
|
||||
FROM market_daily_bar_stage
|
||||
ON CONFLICT (ts_code, trade_date) DO UPDATE SET
|
||||
open = EXCLUDED.open, high = EXCLUDED.high, low = EXCLUDED.low,
|
||||
close = EXCLUDED.close, pre_close = EXCLUDED.pre_close,
|
||||
change = EXCLUDED.change, pct_chg = EXCLUDED.pct_chg,
|
||||
vol = EXCLUDED.vol, amount = EXCLUDED.amount, source_adj = 'qfq', updated_at = now()
|
||||
WHERE (
|
||||
market_daily_bar.open, market_daily_bar.high, market_daily_bar.low,
|
||||
market_daily_bar.close, market_daily_bar.pre_close,
|
||||
market_daily_bar.change, market_daily_bar.pct_chg,
|
||||
market_daily_bar.vol, market_daily_bar.amount
|
||||
) IS DISTINCT FROM (
|
||||
EXCLUDED.open, EXCLUDED.high, EXCLUDED.low, EXCLUDED.close,
|
||||
EXCLUDED.pre_close, EXCLUDED.change, EXCLUDED.pct_chg,
|
||||
EXCLUDED.vol, EXCLUDED.amount
|
||||
)
|
||||
RETURNING (xmax = 0) AS inserted
|
||||
"""
|
||||
).fetchall()
|
||||
inserted = sum(1 for row in rows if bool(row[0]))
|
||||
return inserted, len(rows) - inserted
|
||||
|
||||
@staticmethod
|
||||
def _upsert_daily_basic_from_stage(cursor: Any) -> tuple[int, int]:
|
||||
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",
|
||||
)
|
||||
assignments = ", ".join(f"{field} = EXCLUDED.{field}" for field in fields)
|
||||
comparison = ", ".join(f"market_daily_basic.{field}" for field in fields)
|
||||
excluded = ", ".join(f"EXCLUDED.{field}" for field in fields)
|
||||
query = f"""
|
||||
INSERT INTO market_daily_basic (ts_code, trade_date, {", ".join(fields)}, updated_at)
|
||||
SELECT ts_code, trade_date, {", ".join(fields)}, now() FROM market_daily_basic_stage
|
||||
ON CONFLICT (ts_code, trade_date) DO UPDATE SET {assignments}, updated_at = now()
|
||||
WHERE ({comparison}) IS DISTINCT FROM ({excluded})
|
||||
RETURNING (xmax = 0) AS inserted
|
||||
"""
|
||||
rows = cursor.execute(query).fetchall()
|
||||
inserted = sum(1 for row in rows if bool(row[0]))
|
||||
return inserted, len(rows) - inserted
|
||||
|
||||
@staticmethod
|
||||
def _count_stage_matches(cursor: Any, table: str) -> int:
|
||||
stage = (
|
||||
"market_daily_bar_stage" if table == "market_daily_bar" else "market_daily_basic_stage"
|
||||
)
|
||||
return int(
|
||||
cursor.execute(
|
||||
f"SELECT count(*) FROM {stage} AS stage JOIN {table} AS target "
|
||||
"USING (ts_code, trade_date)"
|
||||
).fetchone()[0]
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _safe_error(message: str | None) -> str | None:
|
||||
if message is None:
|
||||
return None
|
||||
return " ".join(message.split())[:500]
|
||||
@@ -0,0 +1,116 @@
|
||||
"""SQLAlchemy metadata for the market-data migration contract."""
|
||||
|
||||
from sqlalchemy import (
|
||||
Boolean,
|
||||
Column,
|
||||
Date,
|
||||
DateTime,
|
||||
Integer,
|
||||
MetaData,
|
||||
Numeric,
|
||||
PrimaryKeyConstraint,
|
||||
String,
|
||||
Table,
|
||||
Text,
|
||||
func,
|
||||
)
|
||||
|
||||
metadata = MetaData()
|
||||
|
||||
market_stock = Table(
|
||||
"market_stock",
|
||||
metadata,
|
||||
Column("ts_code", String(12), primary_key=True),
|
||||
Column("name", String(128), nullable=False),
|
||||
Column("market", String(32), nullable=False),
|
||||
Column("exchange", String(16), nullable=False),
|
||||
Column("list_status", String(2), nullable=False),
|
||||
Column("list_date", Date),
|
||||
Column("is_active", Boolean, nullable=False, server_default="true"),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
)
|
||||
|
||||
market_daily_bar = Table(
|
||||
"market_daily_bar",
|
||||
metadata,
|
||||
Column("ts_code", String(12), nullable=False),
|
||||
Column("trade_date", Date, nullable=False),
|
||||
Column("open", Numeric(20, 6)),
|
||||
Column("high", Numeric(20, 6)),
|
||||
Column("low", Numeric(20, 6)),
|
||||
Column("close", Numeric(20, 6)),
|
||||
Column("pre_close", Numeric(20, 6)),
|
||||
Column("change", Numeric(20, 6)),
|
||||
Column("pct_chg", Numeric(20, 6)),
|
||||
Column("vol", Numeric(24, 6)),
|
||||
Column("amount", Numeric(24, 6)),
|
||||
Column("source_adj", String(8), nullable=False, server_default="qfq"),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
PrimaryKeyConstraint("ts_code", "trade_date"),
|
||||
)
|
||||
|
||||
market_daily_basic = Table(
|
||||
"market_daily_basic",
|
||||
metadata,
|
||||
Column("ts_code", String(12), nullable=False),
|
||||
Column("trade_date", Date, nullable=False),
|
||||
*(
|
||||
[
|
||||
Column(field, Numeric(24, 6))
|
||||
for field in (
|
||||
"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",
|
||||
)
|
||||
]
|
||||
),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
PrimaryKeyConstraint("ts_code", "trade_date"),
|
||||
)
|
||||
|
||||
market_sync_batch = Table(
|
||||
"market_sync_batch",
|
||||
metadata,
|
||||
Column("id", String(36), primary_key=True),
|
||||
Column("target_trade_date", Date, nullable=False),
|
||||
Column("window_start", Date, nullable=False),
|
||||
Column("mode", String(16), nullable=False),
|
||||
Column("status", String(24), nullable=False),
|
||||
Column("target_count", Integer, nullable=False),
|
||||
Column("valid_count", Integer, nullable=False, server_default="0"),
|
||||
Column("coverage", Numeric(8, 6), nullable=False, server_default="0"),
|
||||
Column("strategy_eligible", Boolean, nullable=False, server_default="false"),
|
||||
Column("parent_batch_id", String(36)),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("finished_at", DateTime(timezone=True)),
|
||||
)
|
||||
|
||||
market_sync_item = Table(
|
||||
"market_sync_item",
|
||||
metadata,
|
||||
Column("batch_id", String(36), nullable=False),
|
||||
Column("item_kind", String(24), nullable=False),
|
||||
Column("item_key", String(64), nullable=False),
|
||||
Column("status", String(24), nullable=False),
|
||||
Column("inserted_count", Integer, nullable=False, server_default="0"),
|
||||
Column("updated_count", Integer, nullable=False, server_default="0"),
|
||||
Column("unchanged_count", Integer, nullable=False, server_default="0"),
|
||||
Column("fingerprint", String(64)),
|
||||
Column("error_type", String(64)),
|
||||
Column("error_message", Text),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
PrimaryKeyConstraint("batch_id", "item_kind", "item_key"),
|
||||
)
|
||||
@@ -0,0 +1,175 @@
|
||||
"""Tushare source adapter and current-universe filtering."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
import time
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from datetime import date
|
||||
from typing import cast
|
||||
|
||||
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow, parse_date
|
||||
from ..domain.rules import filter_current_hs_a_stocks
|
||||
|
||||
|
||||
class TushareSourceError(RuntimeError):
|
||||
"""A vendor request failed after the configured retry budget."""
|
||||
|
||||
|
||||
class TushareAdapter:
|
||||
"""Translate Tushare SDK responses into domain records.
|
||||
|
||||
The SDK is kept behind this adapter so ordinary domain/application tests
|
||||
can inject a tiny fake client and never need a network token.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client: object,
|
||||
*,
|
||||
pro_bar: Callable[..., object] | None = None,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
request_interval_seconds: float = 0.2,
|
||||
random_fn: Callable[[], float] = random.random,
|
||||
sleep_fn: Callable[[float], None] = time.sleep,
|
||||
) -> None:
|
||||
self.client = client
|
||||
self.pro_bar = pro_bar
|
||||
self.max_retries = max(0, max_retries)
|
||||
self.backoff_seconds = max(0.0, backoff_seconds)
|
||||
self.request_interval_seconds = max(0.0, request_interval_seconds)
|
||||
self.random_fn = random_fn
|
||||
self.sleep_fn = sleep_fn
|
||||
|
||||
@classmethod
|
||||
def from_token(
|
||||
cls,
|
||||
token: str,
|
||||
*,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
request_interval_seconds: float = 0.2,
|
||||
) -> TushareAdapter:
|
||||
"""Create a production adapter from a token without exposing it."""
|
||||
|
||||
if not token.strip():
|
||||
raise ValueError("ZHIXING_TUSHARE_TOKEN is required for market sync")
|
||||
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
||||
|
||||
pro_bar_function = cast(
|
||||
Callable[..., object],
|
||||
ts.pro_bar, # pyright: ignore[reportUnknownMemberType]
|
||||
)
|
||||
return cls(
|
||||
cast(object, ts.pro_api(token)),
|
||||
pro_bar=pro_bar_function,
|
||||
max_retries=max_retries,
|
||||
backoff_seconds=backoff_seconds,
|
||||
request_interval_seconds=request_interval_seconds,
|
||||
)
|
||||
|
||||
def fetch_stocks(self) -> tuple[Stock, ...]:
|
||||
"""Fetch current listings and apply the target-universe rules."""
|
||||
|
||||
rows = self._records(
|
||||
"stock_basic",
|
||||
exchange="",
|
||||
list_status="L",
|
||||
fields="ts_code,name,market,exchange,list_status,list_date",
|
||||
)
|
||||
return filter_current_hs_a_stocks(Stock.from_mapping(row) for row in rows)
|
||||
|
||||
def fetch_open_dates(self, start: date, end: date) -> tuple[date, ...]:
|
||||
"""Fetch open trading dates in an inclusive range."""
|
||||
|
||||
rows = self._records(
|
||||
"trade_cal",
|
||||
exchange="",
|
||||
start_date=start.strftime("%Y%m%d"),
|
||||
end_date=end.strftime("%Y%m%d"),
|
||||
is_open=1,
|
||||
)
|
||||
dates: list[date] = []
|
||||
for row in rows:
|
||||
is_open = row.get("is_open")
|
||||
if str(is_open).strip() not in {"1", "True", "true"}:
|
||||
continue
|
||||
dates.append(parse_date(row.get("cal_date")))
|
||||
return tuple(sorted(set(dates)))
|
||||
|
||||
def fetch_bars(self, ts_code: str, window: SyncWindow) -> tuple[Bar, ...]:
|
||||
"""Fetch a complete six-year qfq snapshot for one stock."""
|
||||
|
||||
rows = self._records(
|
||||
"pro_bar",
|
||||
ts_code=ts_code,
|
||||
adj="qfq",
|
||||
start_date=window.start.strftime("%Y%m%d"),
|
||||
end_date=window.end.strftime("%Y%m%d"),
|
||||
freq="D",
|
||||
)
|
||||
bars = tuple(Bar.from_mapping(row) for row in rows)
|
||||
if any(row.ts_code != ts_code for row in bars):
|
||||
raise ValueError(f"Tushare returned a different stock for {ts_code}")
|
||||
return tuple(
|
||||
sorted(
|
||||
(row for row in bars if window.contains(row.trade_date)),
|
||||
key=lambda row: row.trade_date,
|
||||
)
|
||||
)
|
||||
|
||||
def fetch_daily_basic(self, trade_date: date) -> tuple[DailyBasic, ...]:
|
||||
"""Fetch one full-market daily-basic snapshot by trading date."""
|
||||
|
||||
rows = self._records("daily_basic", trade_date=trade_date.strftime("%Y%m%d"))
|
||||
metrics = tuple(DailyBasic.from_mapping(row) for row in rows)
|
||||
if any(row.trade_date != trade_date for row in metrics):
|
||||
raise ValueError(f"Tushare returned a different metric date for {trade_date}")
|
||||
return tuple(sorted(metrics, key=lambda row: row.ts_code))
|
||||
|
||||
def _records(self, method_name: str, **kwargs: object) -> tuple[Mapping[str, object], ...]:
|
||||
"""Call an SDK method with bounded retry and normalize its tabular output."""
|
||||
|
||||
def request() -> object:
|
||||
if method_name == "pro_bar":
|
||||
if self.pro_bar is not None:
|
||||
return self.pro_bar(**kwargs)
|
||||
method = getattr(self.client, method_name, None)
|
||||
else:
|
||||
method = getattr(self.client, method_name, None)
|
||||
if not callable(method):
|
||||
raise TypeError(f"Tushare client has no callable {method_name}")
|
||||
return method(**kwargs)
|
||||
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
result = request()
|
||||
self.sleep_fn(self.request_interval_seconds)
|
||||
return self._as_records(result)
|
||||
except (OSError, RuntimeError, TimeoutError) as exc:
|
||||
last_error = exc
|
||||
if attempt == self.max_retries:
|
||||
break
|
||||
delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn())
|
||||
self.sleep_fn(delay)
|
||||
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
|
||||
|
||||
@staticmethod
|
||||
def _as_records(result: object) -> tuple[Mapping[str, object], ...]:
|
||||
if result is None:
|
||||
return ()
|
||||
to_dict = getattr(result, "to_dict", None)
|
||||
if callable(to_dict):
|
||||
result = to_dict("records")
|
||||
if isinstance(result, Mapping):
|
||||
return (cast(Mapping[str, object], result),)
|
||||
if isinstance(result, Iterable) and not isinstance(result, (str, bytes)):
|
||||
records: list[Mapping[str, object]] = []
|
||||
for row in cast(Iterable[object], result):
|
||||
if not isinstance(row, Mapping):
|
||||
raise TypeError("Tushare rows must be mappings")
|
||||
records.append(cast(Mapping[str, object], row))
|
||||
return tuple(records)
|
||||
raise TypeError("unsupported Tushare tabular response")
|
||||
@@ -0,0 +1 @@
|
||||
"""CLI delivery adapter for the market-data use case."""
|
||||
@@ -0,0 +1,80 @@
|
||||
"""One-shot ``market-data-sync`` command used by Compose and cron."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from datetime import date
|
||||
|
||||
from ....bootstrap.config import get_settings
|
||||
from ..application.sync import SyncMarketData, SyncMarketDataCommand
|
||||
from ..infrastructure.csv_snapshot import CsvSnapshotStore
|
||||
from ..infrastructure.postgres import PostgresMarketDataRepository
|
||||
from ..infrastructure.tushare import TushareAdapter
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
"""Build the explicit, repeatable synchronization CLI."""
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Synchronize Tushare qfq market data into PostgreSQL"
|
||||
)
|
||||
mode = parser.add_mutually_exclusive_group()
|
||||
mode.add_argument(
|
||||
"--initialize",
|
||||
action="store_true",
|
||||
help="backfill daily-basic for the six-year window",
|
||||
)
|
||||
mode.add_argument("--retry-batch-id", help="retry only failed items from a prior batch")
|
||||
parser.add_argument(
|
||||
"--trade-date",
|
||||
type=_parse_date,
|
||||
help="target open date in YYYY-MM-DD format",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
"""Execute one synchronization and print its redacted JSON summary."""
|
||||
|
||||
args = build_parser().parse_args(argv)
|
||||
settings = get_settings()
|
||||
if args.retry_batch_id:
|
||||
command = SyncMarketDataCommand(
|
||||
mode="retry",
|
||||
target_trade_date=args.trade_date,
|
||||
parent_batch_id=args.retry_batch_id,
|
||||
)
|
||||
else:
|
||||
command = SyncMarketDataCommand(
|
||||
mode="initialize" if args.initialize else "daily",
|
||||
target_trade_date=args.trade_date,
|
||||
)
|
||||
source = TushareAdapter.from_token(
|
||||
settings.tushare_token,
|
||||
max_retries=settings.market_data_max_retries,
|
||||
backoff_seconds=settings.market_data_retry_backoff_seconds,
|
||||
request_interval_seconds=settings.market_data_request_interval_seconds,
|
||||
)
|
||||
use_case = SyncMarketData(
|
||||
source,
|
||||
CsvSnapshotStore(settings.market_data_csv_root),
|
||||
PostgresMarketDataRepository(settings.database_url),
|
||||
coverage_threshold=settings.market_data_coverage_threshold,
|
||||
lock_key=settings.market_data_advisory_lock_key,
|
||||
)
|
||||
summary = use_case.execute(command)
|
||||
print(json.dumps(summary.as_dict(), ensure_ascii=False, sort_keys=True))
|
||||
return summary.exit_code
|
||||
|
||||
|
||||
def _parse_date(value: str) -> date:
|
||||
try:
|
||||
return date.fromisoformat(value)
|
||||
except ValueError as exc:
|
||||
raise argparse.ArgumentTypeError("trade date must use YYYY-MM-DD") from exc
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user