feat(sector-radar): 接入Tushare事实与版本化存储
This commit is contained in:
@@ -24,6 +24,11 @@ class Settings(BaseSettings):
|
||||
market_data_max_retries: int = 3
|
||||
market_data_retry_backoff_seconds: float = 1.0
|
||||
market_data_advisory_lock_key: int = 7_380_521
|
||||
sector_radar_coverage_threshold: Decimal = Decimal("0.99")
|
||||
sector_radar_request_interval_seconds: float = 0.2
|
||||
sector_radar_max_retries: int = 3
|
||||
sector_radar_retry_backoff_seconds: float = 1.0
|
||||
sector_radar_advisory_lock_key: int = 7_380_522
|
||||
selection_max_workers: int = Field(default=4, ge=1)
|
||||
selection_batch_size: int = Field(default=200, ge=1)
|
||||
|
||||
|
||||
@@ -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},
|
||||
)
|
||||
@@ -0,0 +1,170 @@
|
||||
"""Shared bounded retry and provider rate-limit coordination."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Sequence
|
||||
|
||||
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 Tushare request failed after the configured retry budget."""
|
||||
|
||||
|
||||
class RequestCoordinator:
|
||||
"""Coordinate retries and shared rate-limit cooling for one provider client.
|
||||
|
||||
Normal requests are not serialized. Only a classified provider limit creates
|
||||
a shared cooldown. Injectable time functions keep long cooldowns deterministic
|
||||
in tests without coupling the coordinator to any business bounded context.
|
||||
"""
|
||||
|
||||
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(
|
||||
"provider_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(
|
||||
"provider_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(
|
||||
"provider_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(
|
||||
"provider_rate_limit_wait method=%s wait_seconds=%.1f",
|
||||
method_name,
|
||||
delay,
|
||||
)
|
||||
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:
|
||||
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))
|
||||
|
||||
|
||||
TushareRequestCoordinator = RequestCoordinator
|
||||
Reference in New Issue
Block a user