feat(market-data): 迁移 Tushare 同步到 PostgreSQL

This commit is contained in:
yuxuanhui
2026-08-05 17:57:32 +08:00
parent e69bcbe821
commit d7d5ed8979
34 changed files with 3230 additions and 3 deletions
@@ -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())