feat(sector-radar): 接入Tushare事实与版本化存储

This commit is contained in:
yuxuanhui
2026-08-29 17:22:39 +08:00
parent 3789008ea6
commit 284c480a90
22 changed files with 3243 additions and 180 deletions
@@ -2,186 +2,28 @@
from __future__ import annotations
import logging
import random
import threading
import time
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import date
from typing import cast
from zhixing_server.shared.request_coordinator import (
DEFAULT_RATE_LIMIT_COOLDOWNS,
RequestCoordinator,
TushareRequestCoordinator,
TushareSourceError,
)
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow, parse_date
from ..domain.rules import filter_current_hs_a_stocks
logger = logging.getLogger(__name__)
DEFAULT_RATE_LIMIT_COOLDOWNS = (60.0, 120.0, 180.0)
_RATE_LIMIT_MESSAGES = (
"访问频繁",
"请稍后",
"超过频率",
"频率限制",
"too many requests",
"rate limit",
"rate_limit",
"http 429",
"status code: 429",
"429",
"http 403",
"status code: 403",
"403",
)
class TushareSourceError(RuntimeError):
"""A vendor request failed after the configured retry budget."""
class RequestCoordinator:
"""Coordinate retry and shared rate-limit cooling for one token client.
Normal requests are deliberately not serialized. Only a provider rate
limit creates a shared cooldown, so independent worker calls can proceed
concurrently during ordinary traffic. ``clock`` and ``wait_fn`` are
injectable to make long cooldown behavior deterministic in unit tests.
"""
def __init__(
self,
*,
max_retries: int = 3,
backoff_seconds: float = 1.0,
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
random_fn: Callable[[], float] = random.random,
clock: Callable[[], float] = time.monotonic,
wait_fn: Callable[[float], None] = time.sleep,
sleep_fn: Callable[[float], None] | None = None,
) -> None:
cooldowns = tuple(float(value) for value in cooldown_seconds)
if not cooldowns or any(value < 0 for value in cooldowns):
raise ValueError("cooldown_seconds must contain non-negative values")
self.max_retries = max(0, max_retries)
self.backoff_seconds = max(0.0, backoff_seconds)
self.cooldown_seconds = cooldowns
self.random_fn = random_fn
self.clock = clock
self.wait_fn = wait_fn
self.sleep_fn = sleep_fn or wait_fn
self._condition = threading.Condition()
self._cooldown_until = 0.0
self._rate_limit_count = 0
@property
def cooldown_until(self) -> float:
"""Return the current monotonic cooldown deadline."""
with self._condition:
return self._cooldown_until
def call(self, method_name: str, request: Callable[[], object]) -> object:
"""Execute one provider request with bounded, shared retry behavior."""
last_error: BaseException | None = None
for attempt in range(self.max_retries + 1):
self._wait_for_cooldown(method_name)
try:
result = request()
except Exception as exc:
last_error = exc
if self.is_rate_limited(exc):
cooldown = self._set_rate_limit_cooldown()
logger.warning(
"tushare_rate_limit method=%s attempt=%d max_attempts=%d "
"cooldown_seconds=%.1f",
method_name,
attempt + 1,
self.max_retries + 1,
cooldown,
)
if attempt < self.max_retries:
continue
break
if not self._is_retryable(exc):
raise
if attempt == self.max_retries:
break
delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn())
logger.warning(
"tushare_request_retry method=%s attempt=%d max_attempts=%d "
"backoff_seconds=%.1f",
method_name,
attempt + 1,
self.max_retries + 1,
delay,
)
self.sleep_fn(delay)
else:
self._clear_rate_limit_after_success()
return result
logger.error(
"tushare_request_failed method=%s attempts=%d",
method_name,
self.max_retries + 1,
)
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
def request(self, method_name: str, operation: Callable[[], object]) -> object:
"""Alias for ``call`` for adapters that model requests as a port."""
return self.call(method_name, operation)
def _wait_for_cooldown(self, method_name: str) -> None:
while True:
with self._condition:
delay = self._cooldown_until - self.clock()
if delay <= 0:
return
logger.info(
"tushare_rate_limit_wait method=%s wait_seconds=%.1f",
method_name,
delay,
)
# A single injected wait hook makes fake-clock tests independent
# from wall time. After waiting, re-check because another worker
# may have extended the shared deadline.
self.wait_fn(delay)
def _set_rate_limit_cooldown(self) -> float:
with self._condition:
self._rate_limit_count += 1
index = min(self._rate_limit_count - 1, len(self.cooldown_seconds) - 1)
duration = self.cooldown_seconds[index]
self._cooldown_until = max(self._cooldown_until, self.clock() + duration)
self._condition.notify_all()
return duration
def _clear_rate_limit_after_success(self) -> None:
with self._condition:
# A request that was already in flight when another worker hit a
# limit may succeed during the shared cooldown. Do not erase the
# escalation history until the cooldown has actually elapsed.
if self.clock() >= self._cooldown_until:
self._rate_limit_count = 0
@staticmethod
def is_rate_limited(error: BaseException) -> bool:
"""Classify stable provider rate-limit signals without logging details."""
for attribute in ("status_code", "status", "code"):
value = getattr(error, attribute, None)
if str(value).strip() in {"403", "429"}:
return True
message = str(error).casefold()
return any(marker.casefold() in message for marker in _RATE_LIMIT_MESSAGES)
@staticmethod
def _is_retryable(error: BaseException) -> bool:
return isinstance(error, (OSError, RuntimeError, TimeoutError))
# The longer name is useful to callers that want to make the infrastructure
# boundary explicit, while the short name remains convenient in unit tests.
TushareRequestCoordinator = RequestCoordinator
__all__ = [
"RequestCoordinator",
"TushareAdapter",
"TushareRequestCoordinator",
"TushareSourceError",
]
class CoordinatedTushareClient:
@@ -67,7 +67,14 @@ def aggregate_sector_snapshot(
expected_count = 0
for member_code in snapshot.member_codes:
fact = facts_by_code.get(member_code)
if fact is None or fact.status is StockFactStatus.MISSING:
if fact is None or fact.status in {
StockFactStatus.MISSING,
StockFactStatus.MISSING_DAILY,
StockFactStatus.MISSING_MONEYFLOW,
StockFactStatus.NULL_DAILY_AMOUNT,
StockFactStatus.NULL_MONEYFLOW,
StockFactStatus.LOW_LIQUIDITY,
}:
expected_count += 1
elif fact.status is StockFactStatus.AVAILABLE:
expected_count += 1
@@ -29,6 +29,10 @@ class StockFactStatus(StrEnum):
AVAILABLE = "available"
SUSPENDED = "suspended"
MISSING = "missing"
MISSING_DAILY = "missing_daily"
MISSING_MONEYFLOW = "missing_moneyflow"
NULL_DAILY_AMOUNT = "null_daily_amount"
NULL_MONEYFLOW = "null_moneyflow"
LIFECYCLE_INVALID = "lifecycle_invalid"
LOW_LIQUIDITY = "low_liquidity"
@@ -0,0 +1,204 @@
"""Normalize typed Tushare rows into point-in-time persisted radar facts."""
from __future__ import annotations
import hashlib
import json
from collections.abc import Callable, Sequence
from datetime import date
from .models import MembershipStatus, StockFactStatus
from .persistence import MembershipRecord, StockFactRecord
from .source import (
DailyRow,
MoneyflowDcRow,
SectorIndexRow,
SectorMemberRow,
SourceContractError,
SourceResult,
StockBasicRow,
SuspendRow,
)
def normalize_memberships(
indices: Sequence[SectorIndexRow],
members: SourceResult[SectorMemberRow],
) -> tuple[MembershipRecord, ...]:
"""Attach each dated member to its sector identity and raw source partition.
Args:
indices: The complete concept or industry universe for one date.
members: Validated membership rows plus all raw request snapshots.
Returns:
Deterministically ordered, source-traceable membership records.
Raises:
SourceContractError: If a member references an unknown sector or lacks a snapshot.
"""
index_by_code = {row.sector_code: row for row in indices}
if len(index_by_code) != len(indices):
raise SourceContractError("sector indices contain duplicate codes")
partition_ids = {
snapshot.partition_key: snapshot.snapshot_id
for snapshot in members.snapshots
if snapshot.partition_key not in {None, "all"}
}
all_snapshot_id = next(
(
snapshot.snapshot_id
for snapshot in members.snapshots
if snapshot.partition_key in {None, "all"}
),
None,
)
records: list[MembershipRecord] = []
for member in members.rows:
index = index_by_code.get(member.sector_code)
if index is None:
raise SourceContractError("dc_member references a sector outside dc_index")
snapshot_id = partition_ids.get(member.sector_code, all_snapshot_id)
if snapshot_id is None:
raise SourceContractError("membership row has no source snapshot")
records.append(
MembershipRecord(
source_snapshot_id=snapshot_id,
trade_date=member.trade_date,
sector_type=index.sector_type,
sector_code=index.sector_code,
sector_name=index.name,
stock_code=member.stock_code,
stock_name=member.stock_name,
status=MembershipStatus.AVAILABLE,
)
)
return tuple(sorted(records, key=lambda item: (item.sector_code, item.stock_code)))
def normalize_stock_facts(
*,
target_trade_date: date,
candidate_codes: Sequence[str],
stock_basics: SourceResult[StockBasicRow],
suspensions: SourceResult[SuspendRow],
daily: SourceResult[DailyRow],
moneyflow: SourceResult[MoneyflowDcRow],
) -> tuple[StockFactRecord, ...]:
"""Build normalized yuan facts without collapsing missing states into zero.
Args:
target_trade_date: Date whose point-in-time lifecycle is evaluated.
candidate_codes: Union of stocks in that date's sector memberships.
stock_basics: All explicit Tushare listing-status partitions.
suspensions: Same-date suspend/resume events.
daily: Same-date stock turnover rows in source units.
moneyflow: Same-date DC main-moneyflow rows in source units.
Returns:
One deterministic fact per candidate code under a content-derived revision.
"""
if len(candidate_codes) != len(set(candidate_codes)):
raise ValueError("candidate_codes must be unique")
basic_by_code = _unique_index(stock_basics.rows, lambda row: row.ts_code, "stock_basic")
daily_by_code = _unique_index(daily.rows, lambda row: row.ts_code, "daily")
moneyflow_by_code = _unique_index(moneyflow.rows, lambda row: row.ts_code, "moneyflow_dc")
suspended_codes = {
row.ts_code
for row in suspensions.rows
if row.trade_date == target_trade_date and _is_suspend_event(row.suspend_type)
}
source_snapshot_ids = tuple(
sorted(
{
snapshot.snapshot_id
for result in (stock_basics, suspensions, daily, moneyflow)
for snapshot in result.snapshots
}
)
)
revision_payload = json.dumps(
{
"target_trade_date": target_trade_date.isoformat(),
"source_snapshot_ids": source_snapshot_ids,
"normalizer": "zhixing_stock_fact_v1",
},
sort_keys=True,
separators=(",", ":"),
)
fact_revision = hashlib.sha256(revision_payload.encode()).hexdigest()
records: list[StockFactRecord] = []
for ts_code in sorted(candidate_codes):
basic = basic_by_code.get(ts_code)
daily_row = daily_by_code.get(ts_code)
moneyflow_row = moneyflow_by_code.get(ts_code)
status = StockFactStatus.AVAILABLE
turnover_yuan = None
net_amount_yuan = None
if basic is None or not _is_lifecycle_candidate(basic, target_trade_date):
status = StockFactStatus.LIFECYCLE_INVALID
elif ts_code in suspended_codes and daily_row is None:
status = StockFactStatus.SUSPENDED
elif daily_row is None:
status = StockFactStatus.MISSING_DAILY
elif daily_row.amount_thousand_yuan is None:
status = StockFactStatus.NULL_DAILY_AMOUNT
elif moneyflow_row is None:
status = StockFactStatus.MISSING_MONEYFLOW
elif moneyflow_row.net_amount_ten_thousand_yuan is None:
status = StockFactStatus.NULL_MONEYFLOW
elif daily_row.turnover_yuan == 0:
status = StockFactStatus.LOW_LIQUIDITY
else:
turnover_yuan = daily_row.turnover_yuan
net_amount_yuan = moneyflow_row.net_amount_yuan
records.append(
StockFactRecord(
fact_revision=fact_revision,
source_snapshot_ids=source_snapshot_ids,
trade_date=target_trade_date,
ts_code=ts_code,
status=status,
turnover_yuan=turnover_yuan,
net_amount_yuan=net_amount_yuan,
)
)
return tuple(records)
def _is_lifecycle_candidate(stock: StockBasicRow, target: date) -> bool:
if not stock.ts_code.endswith((".SH", ".SZ")):
return False
if stock.symbol.startswith(("200", "900")):
return False
if "北交" in stock.market or "B股" in stock.market.upper():
return False
if stock.list_date is None or stock.list_date > target:
return False
return stock.delist_date is None or target <= stock.delist_date
def _is_suspend_event(value: str) -> bool:
normalized = value.strip().casefold()
return normalized in {"s", "suspend", "停牌"} or (
"停牌" in normalized and "复牌" not in normalized
)
def _unique_index[T, K](
rows: Sequence[T],
key: Callable[[T], K],
source_name: str,
) -> dict[K, T]:
result: dict[K, T] = {}
for row in rows:
item_key = key(row)
if item_key in result:
raise SourceContractError(f"{source_name} contains duplicate business keys")
result[item_key] = row
return result
@@ -0,0 +1,144 @@
"""Persistence records and repository port for replayable radar revisions."""
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 typing import Protocol
from .models import (
MembershipStatus,
RadarPublication,
RankedMetric,
SectorType,
StockFactStatus,
)
from .source import SourceSnapshot
def _validate_digest(value: str, field_name: str) -> None:
if len(value) != 64 or any(character not in "0123456789abcdef" for character in value):
raise ValueError(f"{field_name} must be a lowercase SHA-256 digest")
def _validate_optional_decimal(value: Decimal | None, field_name: str) -> None:
if value is not None and not value.is_finite():
raise ValueError(f"{field_name} must be finite or None")
@dataclass(frozen=True, slots=True)
class MembershipRecord:
"""One persisted point-in-time member tied to its raw source revision."""
source_snapshot_id: str
trade_date: date
sector_type: SectorType
sector_code: str
sector_name: str
stock_code: str
stock_name: str
status: MembershipStatus = MembershipStatus.AVAILABLE
def __post_init__(self) -> None:
"""Validate revision identity and member fields."""
_validate_digest(self.source_snapshot_id, "source_snapshot_id")
if self.status is not MembershipStatus.AVAILABLE:
raise ValueError("persisted member rows require available membership")
if any(
not value.strip()
for value in (self.sector_code, self.sector_name, self.stock_code, self.stock_name)
):
raise ValueError("membership identity fields must not be empty")
@dataclass(frozen=True, slots=True)
class StockFactRecord:
"""One normalized stock fact revision with all contributing raw snapshots."""
fact_revision: str
source_snapshot_ids: tuple[str, ...]
trade_date: date
ts_code: str
status: StockFactStatus
turnover_yuan: Decimal | None = None
net_amount_yuan: Decimal | None = None
def __post_init__(self) -> None:
"""Preserve source traceability and stock fact null semantics."""
_validate_digest(self.fact_revision, "fact_revision")
if not self.source_snapshot_ids or len(self.source_snapshot_ids) != len(
set(self.source_snapshot_ids)
):
raise ValueError("source_snapshot_ids must be non-empty and unique")
for value in self.source_snapshot_ids:
_validate_digest(value, "source_snapshot_id")
if not self.ts_code.strip():
raise ValueError("ts_code must not be empty")
_validate_optional_decimal(self.turnover_yuan, "turnover_yuan")
_validate_optional_decimal(self.net_amount_yuan, "net_amount_yuan")
if self.status is StockFactStatus.AVAILABLE:
if self.turnover_yuan is None or self.net_amount_yuan is None:
raise ValueError("available stock facts require both amounts")
if self.turnover_yuan < 0:
raise ValueError("turnover_yuan must not be negative")
elif self.turnover_yuan is not None or self.net_amount_yuan is not None:
raise ValueError("non-available stock facts must not expose amounts")
@dataclass(frozen=True, slots=True)
class RankingRecord:
"""One ranked metric attached to an immutable publication identity."""
publication_id: str
ranking: RankedMetric
def __post_init__(self) -> None:
"""Validate the publication foreign identity."""
if not self.publication_id.strip():
raise ValueError("publication_id must not be empty")
@dataclass(frozen=True, slots=True)
class WriteCounts:
"""Idempotent persistence outcome."""
inserted: int
unchanged: int
def __post_init__(self) -> None:
"""Reject impossible write counts."""
if self.inserted < 0 or self.unchanged < 0:
raise ValueError("write counts must not be negative")
class SectorRadarRepository(Protocol):
"""Persist source revisions, normalized facts, and published rankings."""
def advisory_lock(self, target_trade_date: date) -> AbstractContextManager[bool]: ...
def save_source_snapshots(self, snapshots: Iterable[SourceSnapshot]) -> WriteCounts: ...
def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts: ...
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ...
def create_publication(self, publication: RadarPublication) -> WriteCounts: ...
def finish_publication(self, publication: RadarPublication) -> None: ...
def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts: ...
def get_publication(self, publication_id: str) -> RadarPublication | None: ...
def get_last_good_publication(
self, target_trade_date: date | None = None
) -> RadarPublication | None: ...
def list_successful_dates(self) -> Sequence[date]: ...
@@ -0,0 +1,53 @@
"""Application-facing ports for independent sector radar production."""
from __future__ import annotations
from collections.abc import Sequence
from contextlib import AbstractContextManager
from datetime import date
from typing import Protocol
from .models import SectorType
from .source import (
CapabilityProbeResult,
DailyRow,
MoneyflowDcRow,
SectorIndexRow,
SectorMemberRow,
SourceResult,
StockBasicRow,
SuspendRow,
TradeCalendarRow,
)
class SectorRadarSource(Protocol):
"""Fetch the minimum replayable Tushare facts needed by the MVP."""
def fetch_trade_calendar(self, start: date, end: date) -> SourceResult[TradeCalendarRow]: ...
def fetch_sector_indices(
self, trade_date: date, sector_type: SectorType
) -> SourceResult[SectorIndexRow]: ...
def fetch_sector_members(
self,
trade_date: date,
sector_codes: Sequence[str],
) -> SourceResult[SectorMemberRow]: ...
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]: ...
def fetch_suspensions(self, trade_date: date) -> SourceResult[SuspendRow]: ...
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]: ...
def fetch_moneyflow_dc(self, trade_date: date) -> SourceResult[MoneyflowDcRow]: ...
def probe(self, trade_date: date) -> CapabilityProbeResult: ...
class SectorRadarLock(Protocol):
"""Repository seam for a target-date advisory lock."""
def advisory_lock(self, target_trade_date: date) -> AbstractContextManager[bool]: ...
@@ -0,0 +1,462 @@
"""Typed Tushare input contracts and replayable source snapshot values."""
from __future__ import annotations
import hashlib
import json
import math
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import UTC, date, datetime
from decimal import Decimal, InvalidOperation
from enum import StrEnum
from typing import TypeVar
from .models import SectorType
SourceScalar = str | int | float | bool | None
T = TypeVar("T")
class SourceContractError(ValueError):
"""A provider response violates the replayable input contract."""
class SourceTruncatedError(SourceContractError):
"""A provider response reached its row limit without safe partitioning."""
def normalize_source_scalar(value: object) -> SourceScalar:
"""Normalize flat Tushare cells while distinguishing missing from infinity."""
if value is None:
return None
if isinstance(value, bool):
return value
if isinstance(value, int):
return value
if isinstance(value, float):
if math.isnan(value):
return None
if not math.isfinite(value):
raise SourceContractError("source numeric values must be finite")
return value
if isinstance(value, Decimal):
if value.is_nan():
return None
if not value.is_finite():
raise SourceContractError("source numeric values must be finite")
return str(value)
if isinstance(value, datetime):
return value.isoformat()
if isinstance(value, date):
return value.isoformat()
if isinstance(value, str):
stripped = value.strip()
if not stripped or stripped.casefold() == "nan":
return None
return stripped
raise SourceContractError(f"unsupported source cell type: {type(value).__name__}")
def normalize_source_rows(
rows: Sequence[Mapping[str, object]],
) -> tuple[dict[str, SourceScalar], ...]:
"""Return safe flat rows with deterministic key order."""
return tuple({key: normalize_source_scalar(row[key]) for key in sorted(row)} for row in rows)
@dataclass(frozen=True, slots=True)
class SourceSnapshot:
"""One raw, sanitized provider response identified by safe content hash."""
snapshot_id: str
api_name: str
normalized_params: tuple[tuple[str, str], ...]
target_trade_date: date | None
partition_key: str | None
observed_at: datetime
rows: tuple[dict[str, SourceScalar], ...]
row_count: int
returned_fields: tuple[str, ...]
content_sha256: str
row_limit: int | None
limit_reached: bool
def __post_init__(self) -> None:
"""Validate replay identity and row metadata."""
for field_name, value in (
("snapshot_id", self.snapshot_id),
("content_sha256", self.content_sha256),
):
if len(value) != 64 or any(character not in "0123456789abcdef" for character in value):
raise ValueError(f"{field_name} must be a lowercase SHA-256 digest")
if not self.api_name.strip():
raise ValueError("api_name must not be empty")
if self.observed_at.tzinfo is None:
raise ValueError("observed_at must be timezone-aware")
if self.row_count != len(self.rows):
raise ValueError("row_count must match rows")
if self.row_limit is not None and self.row_limit < 1:
raise ValueError("row_limit must be positive")
if self.limit_reached != (self.row_limit is not None and self.row_count >= self.row_limit):
raise ValueError("limit_reached must match row_count and row_limit")
def build_source_snapshot(
*,
api_name: str,
params: Mapping[str, object],
rows: Sequence[Mapping[str, object]],
target_trade_date: date | None,
partition_key: str | None = None,
observed_at: datetime | None = None,
row_limit: int | None = None,
returned_fields: Sequence[str] | None = None,
) -> SourceSnapshot:
"""Build an order-stable, token-free raw response snapshot."""
normalized_rows = normalize_source_rows(rows)
normalized_params = tuple(
sorted((key, str(value)) for key, value in params.items() if key != "token")
)
canonical_rows = sorted(
json.dumps(row, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
for row in normalized_rows
)
content_sha256 = hashlib.sha256("\n".join(canonical_rows).encode()).hexdigest()
identity = json.dumps(
{
"api_name": api_name,
"params": normalized_params,
"partition_key": partition_key,
"content_sha256": content_sha256,
},
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
)
snapshot_id = hashlib.sha256(identity.encode()).hexdigest()
fields = tuple(
sorted(
set(returned_fields)
if returned_fields is not None
else {key for row in normalized_rows for key in row}
)
)
row_count = len(normalized_rows)
return SourceSnapshot(
snapshot_id=snapshot_id,
api_name=api_name,
normalized_params=normalized_params,
target_trade_date=target_trade_date,
partition_key=partition_key,
observed_at=observed_at or datetime.now(UTC),
rows=normalized_rows,
row_count=row_count,
returned_fields=fields,
content_sha256=content_sha256,
row_limit=row_limit,
limit_reached=row_limit is not None and row_count >= row_limit,
)
@dataclass(frozen=True, slots=True)
class SourceResult[T]:
"""Typed rows accompanied by every raw request needed to produce them."""
snapshots: tuple[SourceSnapshot, ...]
rows: tuple[T, ...]
def _required_text(row: Mapping[str, SourceScalar], key: str) -> str:
value = row.get(key)
if not isinstance(value, str) or not value.strip():
raise SourceContractError(f"{key} must be a non-empty string")
return value.strip()
def _optional_text(row: Mapping[str, SourceScalar], key: str) -> str | None:
value = row.get(key)
if value is None:
return None
return str(value).strip() or None
def _source_date(
row: Mapping[str, SourceScalar], key: str, *, required: bool = True
) -> date | None:
value = row.get(key)
if value is None:
if required:
raise SourceContractError(f"{key} is required")
return None
text = str(value).strip().replace("-", "")
try:
return datetime.strptime(text, "%Y%m%d").date()
except ValueError as exc:
raise SourceContractError(f"{key} must use YYYYMMDD") from exc
def _decimal(row: Mapping[str, SourceScalar], key: str) -> Decimal | None:
value = row.get(key)
if value is None:
return None
try:
result = Decimal(str(value))
except InvalidOperation as exc:
raise SourceContractError(f"{key} must be numeric or missing") from exc
if result.is_nan():
return None
if not result.is_finite():
raise SourceContractError(f"{key} must be finite")
return result
@dataclass(frozen=True, slots=True)
class TradeCalendarRow:
"""One exchange calendar observation."""
exchange: str
cal_date: date
is_open: bool
pretrade_date: date | None
@classmethod
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> TradeCalendarRow:
"""Parse one Tushare ``trade_cal`` row."""
cal_date = _source_date(row, "cal_date")
assert cal_date is not None
return cls(
exchange=_optional_text(row, "exchange") or "",
cal_date=cal_date,
is_open=str(row.get("is_open")).strip().casefold() in {"1", "true"},
pretrade_date=_source_date(row, "pretrade_date", required=False),
)
@dataclass(frozen=True, slots=True)
class SectorIndexRow:
"""One Eastmoney concept or industry identity on a trade date."""
trade_date: date
sector_type: SectorType
sector_code: str
name: str
level: str | None
pct_change: Decimal | None
leading_code: str | None
@classmethod
def from_mapping(
cls,
row: Mapping[str, SourceScalar],
sector_type: SectorType,
) -> SectorIndexRow:
"""Parse and validate one ``dc_index`` row."""
trade_date = _source_date(row, "trade_date")
assert trade_date is not None
return cls(
trade_date=trade_date,
sector_type=sector_type,
sector_code=_required_text(row, "ts_code"),
name=_required_text(row, "name"),
level=_optional_text(row, "level"),
pct_change=_decimal(row, "pct_change"),
leading_code=_optional_text(row, "leading_code"),
)
@dataclass(frozen=True, slots=True)
class SectorMemberRow:
"""One point-in-time sector member returned by ``dc_member``."""
trade_date: date
sector_code: str
stock_code: str
stock_name: str
@classmethod
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> SectorMemberRow:
"""Parse one dated membership row."""
trade_date = _source_date(row, "trade_date")
assert trade_date is not None
return cls(
trade_date=trade_date,
sector_code=_required_text(row, "ts_code"),
stock_code=_required_text(row, "con_code"),
stock_name=_required_text(row, "name"),
)
@dataclass(frozen=True, slots=True)
class StockBasicRow:
"""Lifecycle and market identity from one explicit listing-status query."""
ts_code: str
symbol: str
name: str
market: str
exchange: str
list_status: str
list_date: date | None
delist_date: date | None
@classmethod
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> StockBasicRow:
"""Parse one ``stock_basic`` row without applying ST filtering."""
return cls(
ts_code=_required_text(row, "ts_code"),
symbol=_required_text(row, "symbol"),
name=_required_text(row, "name"),
market=_required_text(row, "market"),
exchange=_required_text(row, "exchange"),
list_status=_required_text(row, "list_status"),
list_date=_source_date(row, "list_date", required=False),
delist_date=_source_date(row, "delist_date", required=False),
)
@dataclass(frozen=True, slots=True)
class SuspendRow:
"""One daily suspend/resume event."""
ts_code: str
trade_date: date
suspend_timing: str
suspend_type: str
@classmethod
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> SuspendRow:
"""Parse one ``suspend_d`` row."""
trade_date = _source_date(row, "trade_date")
assert trade_date is not None
return cls(
ts_code=_required_text(row, "ts_code"),
trade_date=trade_date,
suspend_timing=_required_text(row, "suspend_timing"),
suspend_type=_required_text(row, "suspend_type"),
)
@dataclass(frozen=True, slots=True)
class DailyRow:
"""One stock daily row retaining Tushare's thousand-yuan amount."""
ts_code: str
trade_date: date
close: Decimal | None
pre_close: Decimal | None
pct_chg: Decimal | None
volume: Decimal | None
amount_thousand_yuan: Decimal | None
@property
def turnover_yuan(self) -> Decimal | None:
"""Convert observed turnover to yuan without inventing missing values."""
return (
None if self.amount_thousand_yuan is None else self.amount_thousand_yuan * Decimal(1000)
)
@classmethod
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> DailyRow:
"""Parse one ``daily`` row."""
trade_date = _source_date(row, "trade_date")
assert trade_date is not None
return cls(
ts_code=_required_text(row, "ts_code"),
trade_date=trade_date,
close=_decimal(row, "close"),
pre_close=_decimal(row, "pre_close"),
pct_chg=_decimal(row, "pct_chg"),
volume=_decimal(row, "vol"),
amount_thousand_yuan=_decimal(row, "amount"),
)
@dataclass(frozen=True, slots=True)
class MoneyflowDcRow:
"""One stock main-moneyflow row retaining Tushare's ten-thousand-yuan amount."""
trade_date: date
ts_code: str
name: str
net_amount_ten_thousand_yuan: Decimal | None
net_amount_rate: Decimal | None
pct_change: Decimal | None
close: Decimal | None
@property
def net_amount_yuan(self) -> Decimal | None:
"""Convert observed main net amount to yuan without filling NULL as zero."""
return (
None
if self.net_amount_ten_thousand_yuan is None
else self.net_amount_ten_thousand_yuan * Decimal(10_000)
)
@classmethod
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> MoneyflowDcRow:
"""Parse one ``moneyflow_dc`` row."""
trade_date = _source_date(row, "trade_date")
assert trade_date is not None
return cls(
trade_date=trade_date,
ts_code=_required_text(row, "ts_code"),
name=_required_text(row, "name"),
net_amount_ten_thousand_yuan=_decimal(row, "net_amount"),
net_amount_rate=_decimal(row, "net_amount_rate"),
pct_change=_decimal(row, "pct_change"),
close=_decimal(row, "close"),
)
class CapabilityStatus(StrEnum):
"""Safe capability outcomes that never expose provider error text."""
OK = "ok"
FORBIDDEN = "forbidden"
RATE_LIMITED = "rate_limited"
SERVER_ERROR = "server_error"
SCHEMA_ERROR = "schema_error"
TRUNCATED = "truncated"
@dataclass(frozen=True, slots=True)
class CapabilityInterfaceResult:
"""Safe, credential-free observation for one required interface."""
api_name: str
requested_fields: tuple[str, ...]
returned_fields: tuple[str, ...]
status: CapabilityStatus
row_count: int
row_limit: int | None
retryable: bool
@dataclass(frozen=True, slots=True)
class CapabilityProbeResult:
"""Read-only account capability report for the seven MVP interfaces."""
observed_at: datetime
interfaces: tuple[CapabilityInterfaceResult, ...]
@property
def succeeded(self) -> bool:
"""Return whether every required interface passed its probe."""
return bool(self.interfaces) and all(
result.status is CapabilityStatus.OK for result in self.interfaces
)
@@ -0,0 +1 @@
"""Infrastructure adapters for the sector radar bounded context."""
@@ -0,0 +1,179 @@
"""Deterministic in-memory repository used by application and contract tests."""
from __future__ import annotations
from collections.abc import Callable, Generator, Iterable, Sequence
from contextlib import contextmanager
from datetime import date
from ..domain.models import PublicationStatus, RadarPublication
from ..domain.persistence import (
MembershipRecord,
RankingRecord,
StockFactRecord,
WriteCounts,
)
from ..domain.source import SourceSnapshot
class InMemorySectorRadarRepository:
"""Keep immutable radar revisions in dictionaries without hiding overwrites."""
def __init__(self) -> None:
self.source_snapshots: dict[str, SourceSnapshot] = {}
self.memberships: dict[tuple[str, str, str], MembershipRecord] = {}
self.stock_facts: dict[tuple[str, str], StockFactRecord] = {}
self.publications: dict[str, RadarPublication] = {}
self.rankings: dict[tuple[str, str, str, str], RankingRecord] = {}
self.lock_available = True
@contextmanager
def advisory_lock(self, target_trade_date: date) -> Generator[bool]:
"""Expose a controllable lock result for build orchestration tests."""
del target_trade_date
yield self.lock_available
def save_source_snapshots(self, snapshots: Iterable[SourceSnapshot]) -> WriteCounts:
"""Insert new content-addressed snapshots and count identical replays."""
inserted = 0
unchanged = 0
seen: set[str] = set()
for snapshot in snapshots:
if snapshot.snapshot_id in seen:
raise ValueError("one write batch must not contain duplicate business keys")
seen.add(snapshot.snapshot_id)
if snapshot.snapshot_id in self.source_snapshots:
unchanged += 1
else:
self.source_snapshots[snapshot.snapshot_id] = snapshot
inserted += 1
return WriteCounts(inserted, unchanged)
def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts:
"""Insert membership rows without overwriting an earlier source revision."""
return self._insert_immutable(
self.memberships,
records,
key=lambda item: (item.source_snapshot_id, item.sector_code, item.stock_code),
)
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
"""Insert normalized fact revisions idempotently."""
return self._insert_immutable(
self.stock_facts,
records,
key=lambda item: (item.fact_revision, item.ts_code),
)
def create_publication(self, publication: RadarPublication) -> WriteCounts:
"""Create one running publication without replacing an existing identity."""
if publication.status is not PublicationStatus.RUNNING:
raise ValueError("new publications must start in running status")
if any(
item.status is PublicationStatus.RUNNING
and item.target_trade_date == publication.target_trade_date
and item.publication_id != publication.publication_id
for item in self.publications.values()
):
raise ValueError("target date already has a running publication")
return self._insert_immutable(
self.publications,
(publication,),
key=lambda item: item.publication_id,
)
def finish_publication(self, publication: RadarPublication) -> None:
"""Apply the sole allowed mutation: running to one terminal audit state."""
if publication.status is PublicationStatus.RUNNING:
raise ValueError("finished publication must use a terminal status")
current = self.publications.get(publication.publication_id)
if current is None or current.status is not PublicationStatus.RUNNING:
raise ValueError("publication must exist in running status")
if current.target_trade_date != publication.target_trade_date:
raise ValueError("publication target_trade_date cannot change")
self.publications[publication.publication_id] = publication
def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts:
"""Insert publication-owned rankings idempotently."""
items = tuple(records)
for item in items:
if item.publication_id not in self.publications:
raise ValueError("ranking publication does not exist")
return self._insert_immutable(
self.rankings,
items,
key=lambda item: (
item.publication_id,
item.ranking.observation.sector_type.value,
item.ranking.observation.sector_code,
item.ranking.observation.metric_version,
),
)
def get_publication(self, publication_id: str) -> RadarPublication | None:
"""Return one publication revision by identity."""
return self.publications.get(publication_id)
def get_last_good_publication(
self, target_trade_date: date | None = None
) -> RadarPublication | None:
"""Return only a successful publication; partial and failed never qualify."""
candidates = tuple(
publication
for publication in self.publications.values()
if publication.status is PublicationStatus.SUCCESS
and (target_trade_date is None or publication.target_trade_date <= target_trade_date)
)
return max(
candidates,
key=lambda item: (item.target_trade_date, item.finished_at or item.started_at),
default=None,
)
def list_successful_dates(self) -> Sequence[date]:
"""Return distinct successful dates newest first."""
return tuple(
sorted(
{
item.target_trade_date
for item in self.publications.values()
if item.status is PublicationStatus.SUCCESS
},
reverse=True,
)
)
@staticmethod
def _insert_immutable[K, V](
target: dict[K, V],
values: Iterable[V],
*,
key: Callable[[V], K],
) -> WriteCounts:
inserted = 0
unchanged = 0
seen: set[K] = set()
for value in values:
item_key = key(value)
if item_key in seen:
raise ValueError("one write batch must not contain duplicate business keys")
seen.add(item_key)
existing = target.get(item_key)
if existing is None:
target[item_key] = value
inserted += 1
elif existing == value:
unchanged += 1
else:
raise ValueError("immutable revision identity cannot change content")
return WriteCounts(inserted=inserted, unchanged=unchanged)
@@ -0,0 +1,475 @@
"""Psycopg repository for replayable sector radar inputs and publications."""
from __future__ import annotations
import threading
from collections.abc import Callable, Generator, Iterable, Sequence
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from typing import Any
from psycopg.types.json import Jsonb
from psycopg_pool import ConnectionPool
from ..domain.models import PublicationStatus, RadarPublication
from ..domain.persistence import (
MembershipRecord,
RankingRecord,
StockFactRecord,
WriteCounts,
)
from ..domain.source import SourceSnapshot
class SectorRadarRepositoryError(RuntimeError):
"""PostgreSQL could not complete a radar repository operation safely."""
class PostgresSectorRadarRepository:
"""Persist immutable input revisions with COPY staging and strict last-good reads."""
def __init__(
self,
database_url: str,
*,
advisory_lock_key: int = 7_380_522,
max_connections: int = 4,
pool: ConnectionPool[Any] | None = None,
) -> None:
if max_connections < 1:
raise ValueError("max_connections must be at least 1")
self.database_url = database_url
self.advisory_lock_key = advisory_lock_key
self.pool = pool or ConnectionPool(
conninfo=database_url,
min_size=1,
max_size=max_connections,
open=False,
)
self._owns_pool = pool is None
self._pool_open = False
self._pool_state_lock = threading.Lock()
def open(self) -> None:
"""Open the owned or injected pool exactly once."""
with self._pool_state_lock:
if self._pool_open:
return
if bool(getattr(self.pool, "_opened", False)):
self._pool_open = True
return
self.pool.open(wait=True)
self._pool_open = True
def close(self) -> None:
"""Close only a pool owned by this repository."""
with self._pool_state_lock:
if self._owns_pool and (self._pool_open or bool(getattr(self.pool, "_opened", False))):
self.pool.close()
self._pool_open = False
def __enter__(self) -> PostgresSectorRadarRepository:
self.open()
return self
def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None:
self.close()
@contextmanager
def advisory_lock(self, target_trade_date: date) -> Generator[bool]:
"""Hold a session advisory lock for one target date and build lifetime."""
lock_name = f"sector-radar:{self.advisory_lock_key}:{target_trade_date.isoformat()}"
with self._connection() as connection:
row = connection.execute(
"SELECT pg_try_advisory_lock(hashtext(%s))",
(lock_name,),
).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(hashtext(%s))",
(lock_name,),
)
def save_source_snapshots(self, snapshots: Iterable[SourceSnapshot]) -> WriteCounts:
"""Insert content-addressed raw snapshots without replacing old payloads."""
items = tuple(snapshots)
self._require_unique(items, lambda item: item.snapshot_id)
if not items:
return WriteCounts(0, 0)
with self._connection() as connection, connection.transaction():
existing = {
str(row[0])
for row in connection.execute(
"SELECT id FROM sector_radar_source_snapshot WHERE id = ANY(%s)",
([item.snapshot_id for item in items],),
).fetchall()
}
connection.cursor().executemany(
"""
INSERT INTO sector_radar_source_snapshot (
id, api_name, normalized_params, target_trade_date, partition_key,
observed_at, payload, row_count, returned_fields, content_sha256,
row_limit, limit_reached
) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
ON CONFLICT (id) DO NOTHING
""",
tuple(
(
item.snapshot_id,
item.api_name,
Jsonb(dict(item.normalized_params)),
item.target_trade_date,
item.partition_key,
item.observed_at,
Jsonb(list(item.rows)),
item.row_count,
Jsonb(list(item.returned_fields)),
item.content_sha256,
item.row_limit,
item.limit_reached,
)
for item in items
),
)
unchanged = len(existing)
return WriteCounts(inserted=len(items) - unchanged, unchanged=unchanged)
def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts:
"""COPY point-in-time members into an immutable revision key."""
items = tuple(records)
self._require_unique(
items,
lambda item: (item.source_snapshot_id, item.sector_code, item.stock_code),
)
rows = tuple(
(
item.source_snapshot_id,
item.trade_date,
item.sector_type.value,
item.sector_code,
item.sector_name,
item.stock_code,
item.stock_name,
item.status.value,
)
for item in items
)
return self._copy_immutable(
"sector_radar_membership",
(
"source_snapshot_id",
"trade_date",
"sector_type",
"sector_code",
"sector_name",
"stock_code",
"stock_name",
"membership_status",
),
("source_snapshot_id", "sector_code", "stock_code"),
rows,
)
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
"""COPY normalized stock facts while preserving contributing source ids."""
items = tuple(records)
self._require_unique(items, lambda item: (item.fact_revision, item.ts_code))
rows = tuple(
(
item.fact_revision,
item.trade_date,
item.ts_code,
Jsonb(list(item.source_snapshot_ids)),
item.status.value,
item.turnover_yuan,
item.net_amount_yuan,
)
for item in items
)
return self._copy_immutable(
"sector_radar_stock_fact",
(
"fact_revision",
"trade_date",
"ts_code",
"source_snapshot_ids",
"status",
"turnover_yuan",
"net_amount_yuan",
),
("fact_revision", "ts_code"),
rows,
)
def create_publication(self, publication: RadarPublication) -> WriteCounts:
"""Insert a new running publication identity idempotently."""
if publication.status is not PublicationStatus.RUNNING:
raise ValueError("new publications must start in running status")
with self._connection() as connection, connection.transaction():
existing = connection.execute(
"SELECT status, target_trade_date FROM sector_radar_publication WHERE id = %s",
(publication.publication_id,),
).fetchone()
if existing is not None:
if (
str(existing[0]) != publication.status.value
or existing[1] != publication.target_trade_date
):
raise SectorRadarRepositoryError("publication identity has conflicting content")
return WriteCounts(0, 1)
connection.execute(
"""
INSERT INTO sector_radar_publication (
id, target_trade_date, status, source_version, universe_version,
metric_versions, input_hash, coverage, started_at, finished_at, error_summary
) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""",
self._publication_values(publication),
)
return WriteCounts(1, 0)
def finish_publication(self, publication: RadarPublication) -> None:
"""Transition one running publication to a terminal audit state."""
if publication.status is PublicationStatus.RUNNING:
raise ValueError("finished publication must use a terminal status")
with self._connection() as connection, connection.transaction():
result = connection.execute(
"""
UPDATE sector_radar_publication
SET status = %s, source_version = %s, universe_version = %s,
metric_versions = %s, input_hash = %s, coverage = %s,
finished_at = %s, error_summary = %s
WHERE id = %s AND target_trade_date = %s AND status = 'running'
""",
(
publication.status.value,
publication.source_version,
publication.universe_version,
Jsonb(list(publication.metric_versions)),
publication.input_hash,
publication.coverage,
publication.finished_at,
self._safe_error(publication.error_summary),
publication.publication_id,
publication.target_trade_date,
),
)
if result.rowcount != 1:
raise SectorRadarRepositoryError("publication is not in running status")
def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts:
"""COPY versioned ranking projections under one publication revision."""
items = tuple(records)
self._require_unique(
items,
lambda item: (
item.publication_id,
item.ranking.observation.sector_type,
item.ranking.observation.sector_code,
item.ranking.observation.metric_version,
),
)
rows: list[tuple[object, ...]] = []
for item in items:
ranking = item.ranking
observation = ranking.observation
rows.append(
(
item.publication_id,
observation.trade_date,
observation.sector_type.value,
observation.sector_code,
observation.sector_name,
observation.metric_kind.value,
observation.metric_version,
observation.implementation_kind,
observation.unit.value,
observation.value,
observation.quality.value,
observation.member_count,
observation.valid_sample_count,
observation.membership_coverage,
observation.moneyflow_coverage,
ranking.rank_position,
ranking.rank_percentile,
Jsonb({str(change.days): change.value for change in ranking.rank_changes}),
)
)
return self._copy_immutable(
"sector_radar_ranking",
(
"publication_id",
"trade_date",
"sector_type",
"sector_code",
"sector_name",
"metric_kind",
"metric_version",
"implementation_kind",
"unit",
"metric_value",
"quality",
"member_count",
"valid_sample_count",
"membership_coverage",
"moneyflow_coverage",
"rank_position",
"rank_percentile",
"rank_changes",
),
("publication_id", "sector_type", "sector_code", "metric_version"),
tuple(rows),
)
def get_publication(self, publication_id: str) -> RadarPublication | None:
"""Read one publication by immutable identity."""
with self._connection() as connection:
row = connection.execute(
self._publication_select() + " WHERE id = %s",
(publication_id,),
).fetchone()
return None if row is None else self._publication_from_row(row)
def get_last_good_publication(
self, target_trade_date: date | None = None
) -> RadarPublication | None:
"""Read only status=success, optionally bounded by a requested date."""
where = " WHERE status = 'success'"
parameters: tuple[object, ...] = ()
if target_trade_date is not None:
where += " AND target_trade_date <= %s"
parameters = (target_trade_date,)
query = (
self._publication_select()
+ where
+ " ORDER BY target_trade_date DESC, finished_at DESC, created_at DESC, id DESC LIMIT 1"
)
with self._connection() as connection:
row = connection.execute(query, parameters).fetchone()
return None if row is None else self._publication_from_row(row)
def list_successful_dates(self) -> Sequence[date]:
"""List distinct successful target dates newest first."""
with self._connection() as connection:
rows = connection.execute(
"""
SELECT DISTINCT target_trade_date
FROM sector_radar_publication
WHERE status = 'success'
ORDER BY target_trade_date DESC
"""
).fetchall()
return tuple(row[0] for row in rows)
def _copy_immutable(
self,
table: str,
columns: tuple[str, ...],
conflict_columns: tuple[str, ...],
rows: tuple[tuple[object, ...], ...],
) -> WriteCounts:
if not rows:
return WriteCounts(0, 0)
if table not in {
"sector_radar_membership",
"sector_radar_stock_fact",
"sector_radar_ranking",
}:
raise ValueError("unsupported radar staging table")
stage = f"{table}_stage"
column_sql = ", ".join(columns)
conflict_sql = ", ".join(conflict_columns)
with self._connection() as connection, connection.transaction():
cursor = connection.cursor()
cursor.execute(
f"CREATE TEMP TABLE {stage} (LIKE {table} INCLUDING DEFAULTS) ON COMMIT DROP"
)
with cursor.copy(f"COPY {stage} ({column_sql}) FROM STDIN") as copy:
for row in rows:
copy.write_row(row)
inserted_rows = cursor.execute(
f"INSERT INTO {table} ({column_sql}) SELECT {column_sql} FROM {stage} "
f"ON CONFLICT ({conflict_sql}) DO NOTHING RETURNING 1"
).fetchall()
inserted = len(inserted_rows)
return WriteCounts(inserted=inserted, unchanged=len(rows) - inserted)
@contextmanager
def _connection(self) -> Generator[Any]:
try:
self.open()
with self.pool.connection() as connection:
yield connection
except SectorRadarRepositoryError:
raise
except Exception as exc: # noqa: BLE001 - database errors are redacted at this boundary
raise SectorRadarRepositoryError("sector radar database operation failed") from exc
@staticmethod
def _require_unique[T, K](items: Sequence[T], key: Callable[[T], K]) -> None:
keys = [key(item) for item in items]
if len(keys) != len(set(keys)):
raise ValueError("one write batch must not contain duplicate business keys")
@staticmethod
def _publication_values(publication: RadarPublication) -> tuple[object, ...]:
return (
publication.publication_id,
publication.target_trade_date,
publication.status.value,
publication.source_version,
publication.universe_version,
Jsonb(list(publication.metric_versions)),
publication.input_hash,
publication.coverage,
publication.started_at,
publication.finished_at,
PostgresSectorRadarRepository._safe_error(publication.error_summary),
)
@staticmethod
def _publication_select() -> str:
return (
"SELECT id, target_trade_date, status, source_version, universe_version, "
"metric_versions, input_hash, coverage, started_at, finished_at, error_summary "
"FROM sector_radar_publication"
)
@staticmethod
def _publication_from_row(row: tuple[Any, ...]) -> RadarPublication:
return RadarPublication(
publication_id=str(row[0]),
target_trade_date=row[1],
status=PublicationStatus(str(row[2])),
source_version=str(row[3]),
universe_version=str(row[4]),
metric_versions=tuple(str(value) for value in row[5]),
input_hash=None if row[6] is None else str(row[6]),
coverage=Decimal(str(row[7])),
started_at=row[8],
finished_at=row[9],
error_summary=None if row[10] is None else str(row[10]),
)
@staticmethod
def _safe_error(message: str | None) -> str | None:
return None if message is None else " ".join(message.split())[:500]
@@ -0,0 +1,485 @@
"""Tushare adapter for replayable sector radar source facts."""
from __future__ import annotations
import time
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import UTC, date, datetime
from typing import TypeVar, cast
from zhixing_server.shared.request_coordinator import (
DEFAULT_RATE_LIMIT_COOLDOWNS,
RequestCoordinator,
)
from ..domain.models import SectorType
from ..domain.source import (
CapabilityInterfaceResult,
CapabilityProbeResult,
CapabilityStatus,
DailyRow,
MoneyflowDcRow,
SectorIndexRow,
SectorMemberRow,
SourceContractError,
SourceResult,
SourceSnapshot,
SourceTruncatedError,
StockBasicRow,
SuspendRow,
TradeCalendarRow,
build_source_snapshot,
)
T = TypeVar("T")
FIELDS: dict[str, tuple[str, ...]] = {
"trade_cal": ("exchange", "cal_date", "is_open", "pretrade_date"),
"dc_index": (
"ts_code",
"trade_date",
"name",
"idx_type",
"level",
"pct_change",
"leading_code",
),
"dc_member": ("trade_date", "ts_code", "con_code", "name"),
"stock_basic": (
"ts_code",
"symbol",
"name",
"market",
"exchange",
"list_status",
"list_date",
"delist_date",
),
"suspend_d": ("ts_code", "trade_date", "suspend_timing", "suspend_type"),
"daily": ("ts_code", "trade_date", "close", "pre_close", "pct_chg", "vol", "amount"),
"moneyflow_dc": (
"trade_date",
"ts_code",
"name",
"net_amount",
"net_amount_rate",
"pct_change",
"close",
),
}
ROW_LIMITS: dict[str, int | None] = {
"trade_cal": None,
"dc_index": 5_000,
"dc_member": 5_000,
"stock_basic": None,
"suspend_d": None,
"daily": 6_000,
"moneyflow_dc": 6_000,
}
_SECTOR_TYPE_PARAM = {
SectorType.CONCEPT: "概念板块",
SectorType.INDUSTRY: "行业板块",
}
class TushareSectorRadarAdapter:
"""Fetch seven Tushare interfaces with schema, limit, and replay metadata."""
def __init__(
self,
client: object,
*,
request_coordinator: RequestCoordinator | None = None,
max_retries: int = 3,
backoff_seconds: float = 1.0,
request_interval_seconds: float = 0.2,
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
sleep_fn: Callable[[float], None] = time.sleep,
now_fn: Callable[[], datetime] = lambda: datetime.now(UTC),
) -> None:
"""Create an adapter around one already-authenticated SDK client."""
self._client = client
self._sleep_fn = sleep_fn
self._request_interval_seconds = max(0.0, request_interval_seconds)
self._now_fn = now_fn
self._coordinator = request_coordinator or RequestCoordinator(
max_retries=max_retries,
backoff_seconds=backoff_seconds,
cooldown_seconds=cooldown_seconds,
wait_fn=sleep_fn,
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,
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
) -> TushareSectorRadarAdapter:
"""Create a production client without calling ``set_token`` or retaining the token."""
if not token.strip():
raise ValueError("ZHIXING_TUSHARE_TOKEN is required for sector radar")
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
return cls(
cast(object, ts.pro_api(token)),
max_retries=max_retries,
backoff_seconds=backoff_seconds,
request_interval_seconds=request_interval_seconds,
cooldown_seconds=cooldown_seconds,
)
def fetch_trade_calendar(self, start: date, end: date) -> SourceResult[TradeCalendarRow]:
"""Fetch and validate an inclusive exchange calendar range."""
if end < start:
raise ValueError("end must not precede start")
snapshot = self._fetch_snapshot(
"trade_cal",
{
"exchange": "",
"start_date": start.strftime("%Y%m%d"),
"end_date": end.strftime("%Y%m%d"),
},
target_trade_date=end,
)
rows = tuple(TradeCalendarRow.from_mapping(row) for row in snapshot.rows)
self._require_unique(
rows, key=lambda row: (row.exchange, row.cal_date), api_name="trade_cal"
)
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.cal_date)))
def fetch_sector_indices(
self,
trade_date: date,
sector_type: SectorType,
) -> SourceResult[SectorIndexRow]:
"""Fetch one independent concept or industry universe."""
idx_type = _SECTOR_TYPE_PARAM[sector_type]
snapshot = self._fetch_snapshot(
"dc_index",
{"trade_date": trade_date.strftime("%Y%m%d"), "idx_type": idx_type},
target_trade_date=trade_date,
partition_key=sector_type.value,
)
if any(str(row.get("idx_type")) != idx_type for row in snapshot.rows):
raise SourceContractError("dc_index returned a different idx_type")
self._reject_limit(snapshot)
rows = tuple(SectorIndexRow.from_mapping(row, sector_type) for row in snapshot.rows)
self._require_target_date(rows, trade_date, "dc_index")
self._require_unique(rows, key=lambda row: row.sector_code, api_name="dc_index")
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.sector_code)))
def fetch_sector_members(
self,
trade_date: date,
sector_codes: Sequence[str],
) -> SourceResult[SectorMemberRow]:
"""Fetch dated members and partition when the all-market result is incomplete."""
expected_codes = tuple(sorted(set(sector_codes)))
if len(expected_codes) != len(sector_codes) or any(
not code.strip() for code in expected_codes
):
raise ValueError("sector_codes must contain unique non-empty values")
initial = self._fetch_snapshot(
"dc_member",
{"trade_date": trade_date.strftime("%Y%m%d")},
target_trade_date=trade_date,
partition_key="all",
)
initial_rows = tuple(SectorMemberRow.from_mapping(row) for row in initial.rows)
self._require_target_date(initial_rows, trade_date, "dc_member")
returned_codes = {row.sector_code for row in initial_rows}
missing_codes = tuple(code for code in expected_codes if code not in returned_codes)
if initial.limit_reached:
partition_codes = expected_codes
merged_rows: list[SectorMemberRow] = []
snapshots: list[SourceSnapshot] = [initial]
else:
partition_codes = missing_codes
merged_rows = list(initial_rows)
snapshots = [initial]
if initial.limit_reached and not partition_codes:
raise SourceTruncatedError("dc_member reached its limit without sector partitions")
for sector_code in partition_codes:
snapshot = self._fetch_snapshot(
"dc_member",
{
"trade_date": trade_date.strftime("%Y%m%d"),
"ts_code": sector_code,
},
target_trade_date=trade_date,
partition_key=sector_code,
)
self._reject_limit(snapshot)
partition_rows = tuple(SectorMemberRow.from_mapping(row) for row in snapshot.rows)
self._require_target_date(partition_rows, trade_date, "dc_member")
if any(row.sector_code != sector_code for row in partition_rows):
raise SourceContractError("dc_member partition returned a different sector")
if not partition_rows:
raise SourceContractError("dc_member cannot prove complete membership for a sector")
snapshots.append(snapshot)
merged_rows.extend(partition_rows)
self._require_unique(
merged_rows,
key=lambda row: (row.trade_date, row.sector_code, row.stock_code),
api_name="dc_member",
)
final_codes = {row.sector_code for row in merged_rows}
if set(expected_codes) - final_codes:
raise SourceContractError("dc_member response is missing expected sectors")
return SourceResult(
tuple(snapshots),
tuple(sorted(merged_rows, key=lambda row: (row.sector_code, row.stock_code))),
)
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
"""Fetch every documented listing status instead of relying on the L default."""
snapshots: list[SourceSnapshot] = []
rows: list[StockBasicRow] = []
for status in ("L", "D", "P", "G", "UN"):
snapshot = self._fetch_snapshot(
"stock_basic",
{"exchange": "", "list_status": status},
target_trade_date=None,
partition_key=status,
)
snapshots.append(snapshot)
parsed = tuple(StockBasicRow.from_mapping(row) for row in snapshot.rows)
if any(row.list_status != status for row in parsed):
raise SourceContractError("stock_basic returned an unexpected list_status")
rows.extend(parsed)
self._require_unique(rows, key=lambda row: row.ts_code, api_name="stock_basic")
return SourceResult(tuple(snapshots), tuple(sorted(rows, key=lambda row: row.ts_code)))
def fetch_suspensions(self, trade_date: date) -> SourceResult[SuspendRow]:
"""Fetch explicit suspend/resume events for one date."""
snapshot = self._fetch_snapshot(
"suspend_d",
{"trade_date": trade_date.strftime("%Y%m%d")},
target_trade_date=trade_date,
)
rows = tuple(SuspendRow.from_mapping(row) for row in snapshot.rows)
self._require_target_date(rows, trade_date, "suspend_d")
self._require_unique(
rows,
key=lambda row: (row.ts_code, row.trade_date, row.suspend_type, row.suspend_timing),
api_name="suspend_d",
)
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
"""Fetch a full-market daily snapshot in its documented source unit."""
snapshot = self._fetch_snapshot(
"daily",
{"trade_date": trade_date.strftime("%Y%m%d")},
target_trade_date=trade_date,
)
self._reject_limit(snapshot)
rows = tuple(DailyRow.from_mapping(row) for row in snapshot.rows)
self._require_target_date(rows, trade_date, "daily")
self._require_unique(rows, key=lambda row: row.ts_code, api_name="daily")
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
def fetch_moneyflow_dc(self, trade_date: date) -> SourceResult[MoneyflowDcRow]:
"""Fetch a full-market DC moneyflow snapshot in its documented source unit."""
snapshot = self._fetch_snapshot(
"moneyflow_dc",
{"trade_date": trade_date.strftime("%Y%m%d")},
target_trade_date=trade_date,
)
self._reject_limit(snapshot)
rows = tuple(MoneyflowDcRow.from_mapping(row) for row in snapshot.rows)
self._require_target_date(rows, trade_date, "moneyflow_dc")
self._require_unique(rows, key=lambda row: row.ts_code, api_name="moneyflow_dc")
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
def probe(self, trade_date: date) -> CapabilityProbeResult:
"""Probe required interfaces while returning only safe classifications."""
results: list[CapabilityInterfaceResult] = []
concept_codes: tuple[str, ...] = ()
calendar = self._probe_call(
"trade_cal", lambda: self.fetch_trade_calendar(trade_date, trade_date)
)
results.append(calendar[0])
try:
concept = self.fetch_sector_indices(trade_date, SectorType.CONCEPT)
industry = self.fetch_sector_indices(trade_date, SectorType.INDUSTRY)
combined = SourceResult(
concept.snapshots + industry.snapshots,
concept.rows + industry.rows,
)
concept_codes = tuple(row.sector_code for row in combined.rows)
results.append(self._capability_success("dc_index", combined.snapshots))
except Exception as exc:
results.append(self._capability_failure("dc_index", exc))
member = self._probe_call(
"dc_member", lambda: self.fetch_sector_members(trade_date, concept_codes)
)
results.append(member[0])
for api_name, operation in (
("stock_basic", self.fetch_stock_basics),
("suspend_d", lambda: self.fetch_suspensions(trade_date)),
("daily", lambda: self.fetch_daily(trade_date)),
("moneyflow_dc", lambda: self.fetch_moneyflow_dc(trade_date)),
):
results.append(self._probe_call(api_name, operation)[0])
return CapabilityProbeResult(observed_at=self._now_fn(), interfaces=tuple(results))
def _fetch_snapshot(
self,
api_name: str,
params: Mapping[str, object],
*,
target_trade_date: date | None,
partition_key: str | None = None,
) -> SourceSnapshot:
fields = ",".join(FIELDS[api_name])
def request() -> object:
query = getattr(self._client, "query", None)
if callable(query):
return query(api_name, fields=fields, **params)
method = getattr(self._client, api_name, None)
if not callable(method):
raise TypeError(f"Tushare client has no callable {api_name}")
return method(fields=fields, **params)
result = self._coordinator.call(api_name, request)
self._sleep_fn(self._request_interval_seconds)
columns = getattr(result, "columns", None)
returned_fields = (
tuple(str(column) for column in cast(Iterable[object], columns))
if isinstance(columns, Iterable) and not isinstance(columns, (str, bytes))
else None
)
rows = self._as_records(result)
snapshot = build_source_snapshot(
api_name=api_name,
params={**params, "fields": fields},
rows=rows,
target_trade_date=target_trade_date,
partition_key=partition_key,
observed_at=self._now_fn(),
row_limit=ROW_LIMITS[api_name],
returned_fields=returned_fields,
)
missing_fields = set(FIELDS[api_name]) - set(snapshot.returned_fields)
if snapshot.returned_fields and missing_fields:
raise SourceContractError(f"{api_name} response is missing requested fields")
return snapshot
@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 SourceContractError("Tushare rows must be mappings")
records.append(cast(Mapping[str, object], row))
return tuple(records)
raise SourceContractError("unsupported Tushare tabular response")
@staticmethod
def _reject_limit(snapshot: SourceSnapshot) -> None:
if snapshot.limit_reached:
raise SourceTruncatedError(f"{snapshot.api_name} reached its provider row limit")
@staticmethod
def _require_target_date(rows: Sequence[object], target: date, api_name: str) -> None:
if any(getattr(row, "trade_date", None) != target for row in rows):
raise SourceContractError(f"{api_name} returned a different trade_date")
@staticmethod
def _require_unique(
rows: Sequence[T],
*,
key: Callable[[T], object],
api_name: str,
) -> None:
keys = [key(row) for row in rows]
if len(keys) != len(set(keys)):
raise SourceContractError(f"{api_name} returned duplicate business keys")
def _probe_call(
self,
api_name: str,
operation: Callable[[], SourceResult[object]],
) -> tuple[CapabilityInterfaceResult, SourceResult[object] | None]:
try:
result = operation()
except Exception as exc:
return self._capability_failure(api_name, exc), None
return self._capability_success(api_name, result.snapshots), result
@staticmethod
def _capability_success(
api_name: str,
snapshots: Sequence[SourceSnapshot],
) -> CapabilityInterfaceResult:
return CapabilityInterfaceResult(
api_name=api_name,
requested_fields=FIELDS[api_name],
returned_fields=tuple(
sorted({field for item in snapshots for field in item.returned_fields})
),
status=CapabilityStatus.OK,
row_count=sum(item.row_count for item in snapshots),
row_limit=ROW_LIMITS[api_name],
retryable=False,
)
@staticmethod
def _capability_failure(api_name: str, error: BaseException) -> CapabilityInterfaceResult:
classified_error = error.__cause__ if error.__cause__ is not None else error
message = str(classified_error).casefold()
if isinstance(error, SourceTruncatedError):
status = CapabilityStatus.TRUNCATED
elif RequestCoordinator.is_rate_limited(error) or RequestCoordinator.is_rate_limited(
classified_error
):
status = CapabilityStatus.RATE_LIMITED
elif "权限" in message or "forbidden" in message or "permission" in message:
status = CapabilityStatus.FORBIDDEN
elif isinstance(error, (SourceContractError, ValueError, TypeError)):
status = CapabilityStatus.SCHEMA_ERROR
else:
status = CapabilityStatus.SERVER_ERROR
return CapabilityInterfaceResult(
api_name=api_name,
requested_fields=FIELDS[api_name],
returned_fields=(),
status=status,
row_count=0,
row_limit=ROW_LIMITS[api_name],
retryable=status in {CapabilityStatus.RATE_LIMITED, CapabilityStatus.SERVER_ERROR},
)