Merge branch 'develop' into codex/point
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)
|
||||
selection_pattern_scoring_enabled: bool = True
|
||||
|
||||
@@ -5,6 +5,7 @@ from fastapi import APIRouter
|
||||
from zhixing_server.interfaces.http.system import operational_router, system_router
|
||||
from zhixing_server.modules.market_data.presentation.home import home_router
|
||||
from zhixing_server.modules.market_data.presentation.integrity import integrity_router
|
||||
from zhixing_server.modules.sector_radar.presentation.http import sector_radar_router
|
||||
from zhixing_server.modules.selection.presentation.http import selection_router
|
||||
|
||||
api_v1_router = APIRouter(prefix="/api/v1")
|
||||
@@ -16,5 +17,10 @@ api_v1_router.include_router(
|
||||
tags=["market-data"],
|
||||
)
|
||||
api_v1_router.include_router(selection_router, prefix="/selection", tags=["selection"])
|
||||
api_v1_router.include_router(
|
||||
sector_radar_router,
|
||||
prefix="/sector-radar",
|
||||
tags=["sector-radar"],
|
||||
)
|
||||
|
||||
__all__ = ["api_v1_router", "operational_router"]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Independent post-close sector capital radar bounded context."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Application use cases for sector radar production and reads."""
|
||||
@@ -0,0 +1,765 @@
|
||||
"""One-shot, idempotent sector radar publication orchestration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import UTC, date, datetime, time, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
from uuid import uuid4
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from ..domain.facts import aggregate_sector_snapshot
|
||||
from ..domain.metrics import (
|
||||
AmountNetStrategy,
|
||||
MetricStrategy,
|
||||
RatioTurnoverStrategy,
|
||||
SwingEqualThreeToTenStrategy,
|
||||
)
|
||||
from ..domain.models import (
|
||||
MembershipStatus,
|
||||
MetricObservation,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
SectorDailyAggregate,
|
||||
SectorMembershipSnapshot,
|
||||
SectorType,
|
||||
StockDailyFact,
|
||||
StockFactStatus,
|
||||
)
|
||||
from ..domain.normalize import (
|
||||
is_current_listed_stock,
|
||||
normalize_memberships,
|
||||
normalize_stock_facts,
|
||||
)
|
||||
from ..domain.persistence import (
|
||||
DailyAggregateRecord,
|
||||
MembershipRecord,
|
||||
PublicationSourceGroup,
|
||||
PublicationSourceRecord,
|
||||
RankingRecord,
|
||||
SectorRadarRepository,
|
||||
StockFactRecord,
|
||||
)
|
||||
from ..domain.ports import SectorRadarSource
|
||||
from ..domain.ranking import rank_metric_observations, with_rank_changes
|
||||
from ..domain.source import (
|
||||
DailyRow,
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceContractError,
|
||||
SourceResult,
|
||||
SourceScalar,
|
||||
SourceSnapshot,
|
||||
StockBasicRow,
|
||||
SuspendRow,
|
||||
TradeCalendarRow,
|
||||
)
|
||||
|
||||
BuildOutcomeStatus = Literal["success", "partial", "failed", "locked", "unchanged"]
|
||||
SHANGHAI = ZoneInfo("Asia/Shanghai")
|
||||
MARKET_DATA_READY_TIME = time(15, 30)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BuildSectorRadarCommand:
|
||||
"""Select one date, an inclusive range, or a failed publication retry."""
|
||||
|
||||
trade_date: date | None = None
|
||||
start_date: date | None = None
|
||||
end_date: date | None = None
|
||||
retry_publication_id: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject ambiguous build modes before any source or database call."""
|
||||
|
||||
has_range = self.start_date is not None or self.end_date is not None
|
||||
modes = sum(
|
||||
(
|
||||
self.trade_date is not None,
|
||||
has_range,
|
||||
self.retry_publication_id is not None,
|
||||
)
|
||||
)
|
||||
if modes > 1:
|
||||
raise ValueError("trade date, date range, and retry publication are mutually exclusive")
|
||||
if has_range and (self.start_date is None or self.end_date is None):
|
||||
raise ValueError("date range requires both start_date and end_date")
|
||||
if (
|
||||
self.start_date is not None
|
||||
and self.end_date is not None
|
||||
and self.end_date < self.start_date
|
||||
):
|
||||
raise ValueError("end_date must not precede start_date")
|
||||
if self.retry_publication_id is not None and not self.retry_publication_id.strip():
|
||||
raise ValueError("retry_publication_id must not be empty")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BuildDateOutcome:
|
||||
"""Redacted result for one target trade date."""
|
||||
|
||||
target_trade_date: date
|
||||
status: BuildOutcomeStatus
|
||||
publication_id: str | None
|
||||
coverage: Decimal
|
||||
sector_count: int
|
||||
ranking_count: int
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
"""Serialize without raw payloads, credentials, or provider exception text."""
|
||||
|
||||
return {
|
||||
"target_trade_date": self.target_trade_date.isoformat(),
|
||||
"status": self.status,
|
||||
"publication_id": self.publication_id,
|
||||
"coverage": str(self.coverage),
|
||||
"sector_count": self.sector_count,
|
||||
"ranking_count": self.ranking_count,
|
||||
"error_type": self.error_type,
|
||||
"error_message": self.error_message,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BuildSummary:
|
||||
"""Cron-friendly aggregate result for one CLI invocation."""
|
||||
|
||||
outcomes: tuple[BuildDateOutcome, ...]
|
||||
|
||||
@property
|
||||
def status(self) -> str:
|
||||
"""Return the worst invocation state."""
|
||||
|
||||
if not self.outcomes or any(item.status in {"failed", "locked"} for item in self.outcomes):
|
||||
return "failed"
|
||||
if any(item.status == "partial" for item in self.outcomes):
|
||||
return "partial"
|
||||
if all(item.status == "unchanged" for item in self.outcomes):
|
||||
return "unchanged"
|
||||
return "success"
|
||||
|
||||
@property
|
||||
def exit_code(self) -> int:
|
||||
"""Return 0 for usable success, 2 for incomplete input, and 1 for failure."""
|
||||
|
||||
if self.status == "failed":
|
||||
return 1
|
||||
if self.status == "partial":
|
||||
return 2
|
||||
return 0
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
"""Serialize the invocation summary for external schedulers."""
|
||||
|
||||
return {
|
||||
"status": self.status,
|
||||
"exit_code": self.exit_code,
|
||||
"outcomes": [item.as_dict() for item in self.outcomes],
|
||||
}
|
||||
|
||||
|
||||
class BuildSectorRadar:
|
||||
"""Hide target resolution, source replay, metrics, ranking, and publication switching."""
|
||||
|
||||
source_version = "tushare-pro-v1"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
source: SectorRadarSource,
|
||||
repository: SectorRadarRepository,
|
||||
*,
|
||||
coverage_threshold: Decimal = Decimal("0.99"),
|
||||
today: date | None = None,
|
||||
now_fn: Callable[[], datetime] = lambda: datetime.now(UTC),
|
||||
strategies: Sequence[MetricStrategy] | None = None,
|
||||
) -> None:
|
||||
if not Decimal(0) <= coverage_threshold <= Decimal(1):
|
||||
raise ValueError("coverage_threshold must be between 0 and 1")
|
||||
self.now_fn = now_fn
|
||||
self.source = source
|
||||
self.repository = repository
|
||||
self.coverage_threshold = coverage_threshold
|
||||
self.today = today or self.now_fn().astimezone(SHANGHAI).date()
|
||||
self.strategies = tuple(
|
||||
strategies
|
||||
or (
|
||||
AmountNetStrategy(),
|
||||
RatioTurnoverStrategy(),
|
||||
SwingEqualThreeToTenStrategy(),
|
||||
)
|
||||
)
|
||||
|
||||
def execute(self, command: BuildSectorRadarCommand | None = None) -> BuildSummary:
|
||||
"""Build each selected trade date sequentially for deterministic history."""
|
||||
|
||||
command = command or BuildSectorRadarCommand()
|
||||
try:
|
||||
targets = self._resolve_targets(command)
|
||||
except Exception as exc:
|
||||
target = command.trade_date or command.start_date or self.today
|
||||
error_type, message = self._safe_failure(exc)
|
||||
return BuildSummary(
|
||||
(
|
||||
BuildDateOutcome(
|
||||
target,
|
||||
"failed",
|
||||
None,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
error_type,
|
||||
message,
|
||||
),
|
||||
)
|
||||
)
|
||||
return BuildSummary(tuple(self._build_target(target) for target in targets))
|
||||
|
||||
def _resolve_targets(self, command: BuildSectorRadarCommand) -> tuple[_BuildTarget, ...]:
|
||||
if command.retry_publication_id is not None:
|
||||
publication = self.repository.get_publication(command.retry_publication_id)
|
||||
if publication is None:
|
||||
raise ValueError("retry publication does not exist")
|
||||
if publication.status not in {PublicationStatus.PARTIAL, PublicationStatus.FAILED}:
|
||||
raise ValueError("only partial or failed publications can be retried")
|
||||
return (_BuildTarget(publication.target_trade_date, publication.publication_id),)
|
||||
|
||||
if command.trade_date is not None:
|
||||
start = end = command.trade_date
|
||||
elif command.start_date is not None and command.end_date is not None:
|
||||
start, end = command.start_date, command.end_date
|
||||
else:
|
||||
end = self._default_calendar_end()
|
||||
start = end - timedelta(days=14)
|
||||
calendar = self.source.fetch_trade_calendar(start, end)
|
||||
targets = tuple(sorted({row.cal_date for row in calendar.rows if row.is_open}))
|
||||
if command.trade_date is not None and command.trade_date not in targets:
|
||||
raise ValueError("target date is not an open trading day")
|
||||
if not targets:
|
||||
raise ValueError("no open trading date found")
|
||||
selected = targets if command.start_date is not None else (targets[-1],)
|
||||
return tuple(_BuildTarget(target) for target in selected)
|
||||
|
||||
def _default_calendar_end(self) -> date:
|
||||
"""Exclude today's session until Tushare closing facts are expected to be ready."""
|
||||
|
||||
local_now = self.now_fn().astimezone(SHANGHAI)
|
||||
if self.today == local_now.date() and local_now.time() < MARKET_DATA_READY_TIME:
|
||||
return self.today - timedelta(days=1)
|
||||
return self.today
|
||||
|
||||
def _build_target(self, target: _BuildTarget) -> BuildDateOutcome:
|
||||
try:
|
||||
with self.repository.advisory_lock(target.trade_date) as acquired:
|
||||
if not acquired:
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"locked",
|
||||
None,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
"build_locked",
|
||||
"another sector radar build is running for this date",
|
||||
)
|
||||
return self._build_locked(target)
|
||||
except Exception as exc:
|
||||
error_type, message = self._safe_failure(exc)
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"failed",
|
||||
None,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
error_type,
|
||||
message,
|
||||
)
|
||||
|
||||
def _build_locked(self, target: _BuildTarget) -> BuildDateOutcome:
|
||||
started_at = self.now_fn()
|
||||
publication: RadarPublication | None = None
|
||||
publication_created = False
|
||||
try:
|
||||
self.repository.recover_running_publications(
|
||||
target.trade_date,
|
||||
finished_at=started_at,
|
||||
)
|
||||
publication_id = self._running_id(target.trade_date)
|
||||
publication = RadarPublication(
|
||||
publication_id=publication_id,
|
||||
target_trade_date=target.trade_date,
|
||||
status=PublicationStatus.RUNNING,
|
||||
source_version=self.source_version,
|
||||
universe_version="pending",
|
||||
metric_versions=tuple(strategy.metric_version for strategy in self.strategies),
|
||||
input_hash=None,
|
||||
coverage=Decimal(0),
|
||||
started_at=started_at,
|
||||
)
|
||||
self.repository.create_publication(publication)
|
||||
publication_created = True
|
||||
reusable = self._reusable_sources(target.retry_publication_id)
|
||||
collected = self._collect(target.trade_date, publication_id, reusable)
|
||||
input_hash = self._input_hash(collected.snapshots)
|
||||
existing = self.repository.find_reusable_publication(target.trade_date, input_hash)
|
||||
if existing is not None:
|
||||
self.repository.discard_running_publication(publication_id)
|
||||
publication_created = False
|
||||
is_success = existing.status is PublicationStatus.SUCCESS
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"unchanged" if is_success else "partial",
|
||||
existing.publication_id,
|
||||
existing.coverage,
|
||||
0,
|
||||
0,
|
||||
None if is_success else "duplicate_input",
|
||||
None
|
||||
if is_success
|
||||
else "input is unchanged from an existing partial publication",
|
||||
)
|
||||
|
||||
publication = replace(
|
||||
publication,
|
||||
universe_version=self._universe_version(collected.membership_snapshots),
|
||||
input_hash=input_hash,
|
||||
)
|
||||
aggregates = self._aggregate(collected)
|
||||
rankings = self._rank(target.trade_date, aggregates)
|
||||
coverage = self._coverage(collected.stock_facts)
|
||||
membership_complete = all(
|
||||
item.status is MembershipStatus.AVAILABLE for item in collected.memberships
|
||||
)
|
||||
terminal = (
|
||||
PublicationStatus.SUCCESS
|
||||
if membership_complete and coverage >= self.coverage_threshold
|
||||
else PublicationStatus.PARTIAL
|
||||
)
|
||||
finished = RadarPublication(
|
||||
publication_id=publication.publication_id,
|
||||
target_trade_date=target.trade_date,
|
||||
status=terminal,
|
||||
source_version=publication.source_version,
|
||||
universe_version=publication.universe_version,
|
||||
metric_versions=publication.metric_versions,
|
||||
input_hash=input_hash,
|
||||
coverage=coverage,
|
||||
started_at=started_at,
|
||||
finished_at=self.now_fn(),
|
||||
error_summary=(
|
||||
None
|
||||
if terminal is PublicationStatus.SUCCESS
|
||||
else (
|
||||
"membership_unknown"
|
||||
if not membership_complete
|
||||
else "coverage_below_threshold"
|
||||
)
|
||||
),
|
||||
)
|
||||
self.repository.finalize_publication(
|
||||
finished,
|
||||
memberships=collected.memberships,
|
||||
stock_facts=collected.stock_facts,
|
||||
daily_aggregates=(
|
||||
DailyAggregateRecord(publication_id, aggregate) for aggregate in aggregates
|
||||
),
|
||||
rankings=(RankingRecord(publication_id, ranking) for ranking in rankings),
|
||||
retry_source_groups=(
|
||||
self._retry_source_groups(
|
||||
collected.memberships,
|
||||
collected.stock_facts,
|
||||
)
|
||||
if terminal is PublicationStatus.PARTIAL
|
||||
else ()
|
||||
),
|
||||
)
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"success" if terminal is PublicationStatus.SUCCESS else "partial",
|
||||
publication_id,
|
||||
coverage,
|
||||
len(aggregates),
|
||||
len(rankings),
|
||||
)
|
||||
except Exception as exc:
|
||||
error_type, message = self._safe_failure(exc)
|
||||
failed_id = (
|
||||
publication.publication_id
|
||||
if publication is not None
|
||||
else self._failure_id(target.trade_date)
|
||||
)
|
||||
if isinstance(exc, SourceContractError) and exc.claim_diagnostic():
|
||||
logger.error(
|
||||
"sector_radar_build_source_contract_failed "
|
||||
"target_trade_date=%s publication_id=%s validation=%s",
|
||||
target.trade_date.isoformat(),
|
||||
failed_id,
|
||||
exc.operator_message,
|
||||
)
|
||||
if publication is not None and publication_created:
|
||||
self.repository.finish_publication(
|
||||
RadarPublication(
|
||||
publication_id=failed_id,
|
||||
target_trade_date=target.trade_date,
|
||||
status=PublicationStatus.FAILED,
|
||||
source_version=publication.source_version,
|
||||
universe_version=publication.universe_version,
|
||||
metric_versions=publication.metric_versions,
|
||||
input_hash=publication.input_hash,
|
||||
coverage=Decimal(0),
|
||||
started_at=started_at,
|
||||
finished_at=self.now_fn(),
|
||||
error_summary=f"{error_type}:{message}",
|
||||
)
|
||||
)
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"failed",
|
||||
failed_id,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
error_type,
|
||||
message,
|
||||
)
|
||||
|
||||
def _reusable_sources(
|
||||
self, publication_id: str | None
|
||||
) -> dict[PublicationSourceGroup, tuple[SourceSnapshot, ...]]:
|
||||
"""Load successful checkpoints while forcing incomplete coverage facts to refresh."""
|
||||
|
||||
if publication_id is None:
|
||||
return {}
|
||||
publication = self.repository.get_publication(publication_id)
|
||||
if publication is None:
|
||||
raise ValueError("retry publication does not exist")
|
||||
grouped: defaultdict[PublicationSourceGroup, list[PublicationSourceRecord]] = defaultdict(
|
||||
list
|
||||
)
|
||||
for record in self.repository.load_publication_sources(publication_id):
|
||||
if not record.refresh_on_retry:
|
||||
grouped[record.source_group].append(record)
|
||||
result: dict[PublicationSourceGroup, tuple[SourceSnapshot, ...]] = {}
|
||||
for group, records in grouped.items():
|
||||
ordered = sorted(records, key=lambda item: item.source_order)
|
||||
if [item.source_order for item in ordered] != list(range(len(ordered))):
|
||||
raise ValueError("publication source checkpoint order is incomplete")
|
||||
result[group] = tuple(item.snapshot for item in ordered)
|
||||
return result
|
||||
|
||||
def _fetch_group[T](
|
||||
self,
|
||||
publication_id: str,
|
||||
source_group: PublicationSourceGroup,
|
||||
reusable: Mapping[PublicationSourceGroup, tuple[SourceSnapshot, ...]],
|
||||
fetch: Callable[[], SourceResult[T]],
|
||||
parser: Callable[[Mapping[str, SourceScalar]], T],
|
||||
reuse_if: Callable[[SourceResult[T]], bool] | None = None,
|
||||
) -> SourceResult[T]:
|
||||
"""Replay a compatible completed group or fetch and checkpoint it immediately."""
|
||||
|
||||
try:
|
||||
snapshots = reusable.get(source_group)
|
||||
if snapshots is None:
|
||||
result = fetch()
|
||||
else:
|
||||
replayed = SourceResult(
|
||||
snapshots=snapshots,
|
||||
rows=tuple(parser(row) for snapshot in snapshots for row in snapshot.rows),
|
||||
)
|
||||
result = replayed if reuse_if is None or reuse_if(replayed) else fetch()
|
||||
except SourceContractError as exc:
|
||||
if exc.claim_diagnostic():
|
||||
logger.error(
|
||||
"sector_radar_source_group_contract_failed "
|
||||
"publication_id=%s source_group=%s validation=%s",
|
||||
publication_id,
|
||||
source_group.value,
|
||||
exc.operator_message,
|
||||
)
|
||||
raise
|
||||
if not result.snapshots:
|
||||
raise ValueError("source group must include at least one replay snapshot")
|
||||
snapshot_ids = [snapshot.snapshot_id for snapshot in result.snapshots]
|
||||
if len(snapshot_ids) != len(set(snapshot_ids)):
|
||||
raise ValueError("source group contains duplicate snapshots")
|
||||
self.repository.save_source_snapshots(result.snapshots)
|
||||
self.repository.save_publication_sources(
|
||||
PublicationSourceRecord(publication_id, source_group, order, snapshot)
|
||||
for order, snapshot in enumerate(result.snapshots)
|
||||
)
|
||||
return result
|
||||
|
||||
def _collect(
|
||||
self,
|
||||
target: date,
|
||||
publication_id: str,
|
||||
reusable: Mapping[PublicationSourceGroup, tuple[SourceSnapshot, ...]],
|
||||
) -> _CollectedInputs:
|
||||
calendar = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.CALENDAR,
|
||||
reusable,
|
||||
lambda: self.source.fetch_trade_calendar(target, target),
|
||||
TradeCalendarRow.from_mapping,
|
||||
)
|
||||
if target not in {row.cal_date for row in calendar.rows if row.is_open}:
|
||||
raise ValueError("target date is not an open trading day")
|
||||
concepts = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.CONCEPT_INDICES,
|
||||
reusable,
|
||||
lambda: self.source.fetch_sector_indices(target, SectorType.CONCEPT),
|
||||
lambda row: SectorIndexRow.from_mapping(row, SectorType.CONCEPT),
|
||||
)
|
||||
industries = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.INDUSTRY_INDICES,
|
||||
reusable,
|
||||
lambda: self.source.fetch_sector_indices(target, SectorType.INDUSTRY),
|
||||
lambda row: SectorIndexRow.from_mapping(row, SectorType.INDUSTRY),
|
||||
)
|
||||
indices = concepts.rows + industries.rows
|
||||
sector_codes = tuple(row.sector_code for row in indices)
|
||||
members = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.MEMBERS,
|
||||
reusable,
|
||||
lambda: self.source.fetch_sector_members(target, sector_codes),
|
||||
SectorMemberRow.from_mapping,
|
||||
)
|
||||
stock_basics = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.STOCK_BASICS,
|
||||
reusable,
|
||||
self.source.fetch_stock_basics,
|
||||
StockBasicRow.from_mapping,
|
||||
)
|
||||
memberships = normalize_memberships(indices, members)
|
||||
member_codes = tuple(
|
||||
sorted(
|
||||
{
|
||||
item.stock_code
|
||||
for item in memberships
|
||||
if item.status is MembershipStatus.AVAILABLE and item.stock_code is not None
|
||||
}
|
||||
)
|
||||
)
|
||||
current_listed_codes = {
|
||||
row.ts_code for row in stock_basics.rows if is_current_listed_stock(row, target)
|
||||
}
|
||||
moneyflow_candidate_codes = tuple(
|
||||
code for code in member_codes if code in current_listed_codes
|
||||
)
|
||||
suspensions = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.SUSPENSIONS,
|
||||
reusable,
|
||||
lambda: self.source.fetch_suspensions(target),
|
||||
SuspendRow.from_mapping,
|
||||
)
|
||||
daily = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.DAILY,
|
||||
reusable,
|
||||
lambda: self.source.fetch_daily(target),
|
||||
DailyRow.from_mapping,
|
||||
)
|
||||
moneyflow = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.MONEYFLOW_DC,
|
||||
reusable,
|
||||
lambda: self.source.fetch_moneyflow_dc(target, moneyflow_candidate_codes),
|
||||
MoneyflowDcRow.from_mapping,
|
||||
reuse_if=lambda result: set(moneyflow_candidate_codes).issubset(
|
||||
{row.ts_code for row in result.rows}
|
||||
),
|
||||
)
|
||||
|
||||
stock_facts = normalize_stock_facts(
|
||||
target_trade_date=target,
|
||||
candidate_codes=member_codes,
|
||||
stock_basics=stock_basics,
|
||||
suspensions=suspensions,
|
||||
daily=daily,
|
||||
moneyflow=moneyflow,
|
||||
)
|
||||
snapshots = (
|
||||
calendar.snapshots
|
||||
+ concepts.snapshots
|
||||
+ industries.snapshots
|
||||
+ members.snapshots
|
||||
+ stock_basics.snapshots
|
||||
+ suspensions.snapshots
|
||||
+ daily.snapshots
|
||||
+ moneyflow.snapshots
|
||||
)
|
||||
return _CollectedInputs(
|
||||
target_trade_date=target,
|
||||
snapshots=snapshots,
|
||||
membership_snapshots=members.snapshots,
|
||||
memberships=memberships,
|
||||
stock_facts=stock_facts,
|
||||
)
|
||||
|
||||
def _aggregate(self, inputs: _CollectedInputs) -> tuple[SectorDailyAggregate, ...]:
|
||||
facts = tuple(
|
||||
StockDailyFact(
|
||||
trade_date=item.trade_date,
|
||||
ts_code=item.ts_code,
|
||||
status=item.status,
|
||||
turnover_yuan=item.turnover_yuan,
|
||||
net_amount_yuan=item.net_amount_yuan,
|
||||
)
|
||||
for item in inputs.stock_facts
|
||||
)
|
||||
grouped: defaultdict[tuple[SectorType, str, str], list[str]] = defaultdict(list)
|
||||
unknown: set[tuple[SectorType, str, str]] = set()
|
||||
for member in inputs.memberships:
|
||||
key = (member.sector_type, member.sector_code, member.sector_name)
|
||||
if member.status is MembershipStatus.UNKNOWN:
|
||||
unknown.add(key)
|
||||
continue
|
||||
if member.stock_code is None:
|
||||
raise ValueError("available membership requires a stock code")
|
||||
grouped[key].append(member.stock_code)
|
||||
if unknown & set(grouped):
|
||||
raise ValueError("sector cannot have both available and unknown membership")
|
||||
sector_keys = set(grouped) | unknown
|
||||
aggregates = tuple(
|
||||
aggregate_sector_snapshot(
|
||||
SectorMembershipSnapshot(
|
||||
trade_date=inputs.target_trade_date,
|
||||
sector_type=sector_type,
|
||||
sector_code=sector_code,
|
||||
sector_name=sector_name,
|
||||
member_codes=tuple(sorted(grouped.get(key, ()))),
|
||||
status=(
|
||||
MembershipStatus.UNKNOWN if key in unknown else MembershipStatus.AVAILABLE
|
||||
),
|
||||
source_version=self._universe_version(inputs.membership_snapshots),
|
||||
),
|
||||
facts,
|
||||
)
|
||||
for key in sorted(sector_keys, key=lambda item: (str(item[0]), item[1]))
|
||||
for sector_type, sector_code, sector_name in (key,)
|
||||
)
|
||||
if not aggregates:
|
||||
raise ValueError("sector universe produced no aggregates")
|
||||
return aggregates
|
||||
|
||||
def _rank(
|
||||
self, target: date, aggregates: Sequence[SectorDailyAggregate]
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
history = tuple(self.repository.load_daily_aggregate_history(target, limit_dates=9))
|
||||
observations: list[MetricObservation] = []
|
||||
for current in aggregates:
|
||||
sector_history = tuple(
|
||||
item
|
||||
for item in history
|
||||
if (item.sector_type, item.sector_code)
|
||||
== (current.sector_type, current.sector_code)
|
||||
) + (current,)
|
||||
observations.extend(
|
||||
strategy.evaluate(sector_history, target) for strategy in self.strategies
|
||||
)
|
||||
current_rankings = rank_metric_observations(observations)
|
||||
previous = self.repository.load_previous_rankings(target, limit_dates=5)
|
||||
history_by_days = {days: rankings for days, (_, rankings) in enumerate(previous, start=1)}
|
||||
return with_rank_changes(current_rankings, history_by_days)
|
||||
|
||||
@staticmethod
|
||||
def _coverage(stock_facts: Sequence[StockFactRecord]) -> Decimal:
|
||||
expected_statuses = {
|
||||
"available",
|
||||
"missing",
|
||||
"missing_daily",
|
||||
"missing_moneyflow",
|
||||
"null_daily_amount",
|
||||
"null_moneyflow",
|
||||
"low_liquidity",
|
||||
}
|
||||
expected = sum(item.status.value in expected_statuses for item in stock_facts)
|
||||
covered_statuses = {StockFactStatus.AVAILABLE, StockFactStatus.LOW_LIQUIDITY}
|
||||
covered = sum(item.status in covered_statuses for item in stock_facts)
|
||||
return Decimal(covered) / Decimal(expected) if expected else Decimal(0)
|
||||
|
||||
@staticmethod
|
||||
def _retry_source_groups(
|
||||
memberships: Sequence[MembershipRecord],
|
||||
stock_facts: Sequence[StockFactRecord],
|
||||
) -> tuple[PublicationSourceGroup, ...]:
|
||||
groups: list[PublicationSourceGroup] = []
|
||||
if any(item.status is MembershipStatus.UNKNOWN for item in memberships):
|
||||
groups.append(PublicationSourceGroup.MEMBERS)
|
||||
statuses = {item.status for item in stock_facts}
|
||||
if statuses & {
|
||||
StockFactStatus.MISSING,
|
||||
StockFactStatus.MISSING_DAILY,
|
||||
StockFactStatus.NULL_DAILY_AMOUNT,
|
||||
}:
|
||||
groups.append(PublicationSourceGroup.DAILY)
|
||||
if statuses & {
|
||||
StockFactStatus.MISSING,
|
||||
StockFactStatus.MISSING_MONEYFLOW,
|
||||
StockFactStatus.NULL_MONEYFLOW,
|
||||
}:
|
||||
groups.append(PublicationSourceGroup.MONEYFLOW_DC)
|
||||
return tuple(groups)
|
||||
|
||||
def _input_hash(self, snapshots: Sequence[SourceSnapshot]) -> str:
|
||||
payload = json.dumps(
|
||||
{
|
||||
"snapshot_ids": sorted(snapshot.snapshot_id for snapshot in snapshots),
|
||||
"metric_versions": sorted(strategy.metric_version for strategy in self.strategies),
|
||||
"normalizer": "zhixing_stock_fact_v1",
|
||||
},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return hashlib.sha256(payload.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _universe_version(snapshots: Sequence[SourceSnapshot]) -> str:
|
||||
payload = "\n".join(sorted(snapshot.snapshot_id for snapshot in snapshots))
|
||||
return f"eastmoney-dc-{hashlib.sha256(payload.encode()).hexdigest()[:32]}"
|
||||
|
||||
@staticmethod
|
||||
def _failure_id(target: date) -> str:
|
||||
return f"radar-{target:%Y%m%d}-failed-{uuid4().hex[:24]}"
|
||||
|
||||
@staticmethod
|
||||
def _running_id(target: date) -> str:
|
||||
return f"radar-{target:%Y%m%d}-running-{uuid4().hex[:23]}"
|
||||
|
||||
@staticmethod
|
||||
def _safe_failure(error: BaseException) -> tuple[str, str]:
|
||||
if isinstance(error, ValueError):
|
||||
return type(error).__name__, "input or source contract validation failed"
|
||||
return type(error).__name__, "sector radar build failed"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CollectedInputs:
|
||||
target_trade_date: date
|
||||
snapshots: tuple[SourceSnapshot, ...]
|
||||
membership_snapshots: tuple[SourceSnapshot, ...]
|
||||
memberships: tuple[MembershipRecord, ...]
|
||||
stock_facts: tuple[StockFactRecord, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BuildTarget:
|
||||
trade_date: date
|
||||
retry_publication_id: str | None = None
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Stable read model for persisted sector radar publications."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from enum import StrEnum
|
||||
from typing import Literal
|
||||
|
||||
from ..domain.metrics import (
|
||||
AmountNetStrategy,
|
||||
RatioTurnoverStrategy,
|
||||
SwingEqualThreeToTenStrategy,
|
||||
)
|
||||
from ..domain.models import (
|
||||
MetricKind,
|
||||
MetricUnit,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
from ..domain.persistence import SectorRadarRepository
|
||||
from ..domain.ranking import select_percentile_side, select_rank_change_side
|
||||
|
||||
ReadStatus = Literal["success", "no_data"]
|
||||
|
||||
|
||||
class RadarView(StrEnum):
|
||||
"""Supported ranking projections at the HTTP boundary."""
|
||||
|
||||
AMOUNT = "amount"
|
||||
RATIO = "ratio"
|
||||
SWING = "swing"
|
||||
RANK_CHANGE = "rank_change"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RadarMetricDefinition:
|
||||
"""Public definition of one explicitly independent metric implementation."""
|
||||
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
label: str
|
||||
unit: MetricUnit
|
||||
implementation_kind: Literal["independent"] = "independent"
|
||||
disclaimer: str = "知行独立实现,非 OneChartLab 原站公式"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RadarQuery:
|
||||
"""Validated application query for one ranking page."""
|
||||
|
||||
trade_date: date | None = None
|
||||
sector_type: SectorType = SectorType.CONCEPT
|
||||
view: RadarView = RadarView.AMOUNT
|
||||
rank_change_metric: MetricKind = MetricKind.AMOUNT
|
||||
rank_change_days: int = 1
|
||||
side: RankSide = RankSide.ALL
|
||||
search: str | None = None
|
||||
page: int = 1
|
||||
page_size: int = 20
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject invalid pagination and rank-history offsets outside HTTP usage."""
|
||||
|
||||
if not 1 <= self.rank_change_days <= 5:
|
||||
raise ValueError("rank_change_days must be between 1 and 5")
|
||||
if self.page < 1:
|
||||
raise ValueError("page must be positive")
|
||||
if not 1 <= self.page_size <= 100:
|
||||
raise ValueError("page_size must be between 1 and 100")
|
||||
if self.search is not None and len(self.search) > 100:
|
||||
raise ValueError("search must not exceed 100 characters")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RadarDateIndex:
|
||||
"""Available successful dates plus the newest attempt and strict last-good."""
|
||||
|
||||
available_dates: tuple[date, ...]
|
||||
current_attempt: RadarPublication | None
|
||||
last_good: RadarPublication | None
|
||||
|
||||
@property
|
||||
def status(self) -> ReadStatus:
|
||||
"""Return no_data until at least one successful publication exists."""
|
||||
|
||||
return "success" if self.last_good is not None else "no_data"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RankingPage:
|
||||
"""One filtered page without losing publication or metric provenance."""
|
||||
|
||||
status: ReadStatus
|
||||
query: RadarQuery
|
||||
publication: RadarPublication | None
|
||||
definition: RadarMetricDefinition
|
||||
rows: tuple[RankedMetric, ...]
|
||||
total: int
|
||||
|
||||
|
||||
_METRIC_DEFINITIONS = {
|
||||
MetricKind.AMOUNT: RadarMetricDefinition(
|
||||
metric_kind=MetricKind.AMOUNT,
|
||||
metric_version=AmountNetStrategy.metric_version,
|
||||
label="主力净流入(知行独立实现)",
|
||||
unit=MetricUnit.CNY_100M,
|
||||
),
|
||||
MetricKind.RATIO: RadarMetricDefinition(
|
||||
metric_kind=MetricKind.RATIO,
|
||||
metric_version=RatioTurnoverStrategy.metric_version,
|
||||
label="主力净流入/成交额(知行独立实现)",
|
||||
unit=MetricUnit.RATIO,
|
||||
),
|
||||
MetricKind.SWING: RadarMetricDefinition(
|
||||
metric_kind=MetricKind.SWING,
|
||||
metric_version=SwingEqualThreeToTenStrategy.metric_version,
|
||||
label="3—10 日等权资金率(知行独立实现)",
|
||||
unit=MetricUnit.RATIO,
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class ReadSectorRadar:
|
||||
"""Hide last-good selection, ranking filters, search, and pagination."""
|
||||
|
||||
def __init__(self, repository: SectorRadarRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def list_dates(self) -> RadarDateIndex:
|
||||
"""Return successful dates without promoting partial or failed attempts."""
|
||||
|
||||
return RadarDateIndex(
|
||||
available_dates=tuple(self.repository.list_successful_dates()),
|
||||
current_attempt=self.repository.get_latest_publication(),
|
||||
last_good=self.repository.get_last_good_publication(),
|
||||
)
|
||||
|
||||
def query(self, query: RadarQuery) -> RankingPage:
|
||||
"""Return one deterministic page for ordinary or rank-change views."""
|
||||
|
||||
metric_kind = (
|
||||
query.rank_change_metric
|
||||
if query.view is RadarView.RANK_CHANGE
|
||||
else MetricKind(query.view.value)
|
||||
)
|
||||
definition = _METRIC_DEFINITIONS[metric_kind]
|
||||
publication = (
|
||||
self.repository.get_successful_publication(query.trade_date)
|
||||
if query.trade_date is not None
|
||||
else self.repository.get_last_good_publication()
|
||||
)
|
||||
if publication is None:
|
||||
return RankingPage("no_data", query, None, definition, (), 0)
|
||||
|
||||
metric_rows = tuple(
|
||||
row
|
||||
for row in self.repository.load_rankings(publication.publication_id)
|
||||
if row.observation.sector_type is query.sector_type
|
||||
and row.observation.metric_kind is metric_kind
|
||||
and row.observation.metric_version == definition.metric_version
|
||||
)
|
||||
if query.view is RadarView.RANK_CHANGE:
|
||||
if query.side is RankSide.ALL:
|
||||
selected = tuple(
|
||||
sorted(
|
||||
metric_rows,
|
||||
key=lambda row: (
|
||||
row.rank_change(query.rank_change_days) is None,
|
||||
-(row.rank_change(query.rank_change_days) or 0),
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
else:
|
||||
selected = select_rank_change_side(
|
||||
metric_rows,
|
||||
days=query.rank_change_days,
|
||||
side=query.side,
|
||||
)
|
||||
else:
|
||||
selected = select_percentile_side(metric_rows, query.side)
|
||||
if query.side is not RankSide.BOTTOM:
|
||||
selected = tuple(
|
||||
sorted(
|
||||
selected,
|
||||
key=lambda row: (
|
||||
row.rank_position is None,
|
||||
row.rank_position or 0,
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
search = query.search.strip().casefold() if query.search else ""
|
||||
searched = tuple(
|
||||
row
|
||||
for row in selected
|
||||
if not search
|
||||
or search in row.observation.sector_code.casefold()
|
||||
or search in row.observation.sector_name.casefold()
|
||||
)
|
||||
start = (query.page - 1) * query.page_size
|
||||
return RankingPage(
|
||||
status="success",
|
||||
query=query,
|
||||
publication=publication,
|
||||
definition=definition,
|
||||
rows=searched[start : start + query.page_size],
|
||||
total=len(searched),
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RadarDateIndex",
|
||||
"RadarMetricDefinition",
|
||||
"RadarQuery",
|
||||
"RadarView",
|
||||
"RankingPage",
|
||||
"ReadSectorRadar",
|
||||
]
|
||||
@@ -0,0 +1 @@
|
||||
"""Storage-independent sector radar models and calculation rules."""
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Point-in-time stock fact aggregation for sector radar metrics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from decimal import Decimal
|
||||
|
||||
from .models import (
|
||||
MembershipStatus,
|
||||
SectorDailyAggregate,
|
||||
SectorMembershipSnapshot,
|
||||
StockDailyFact,
|
||||
StockFactStatus,
|
||||
)
|
||||
|
||||
|
||||
def aggregate_sector_snapshot(
|
||||
snapshot: SectorMembershipSnapshot,
|
||||
stock_facts: Iterable[StockDailyFact],
|
||||
) -> SectorDailyAggregate:
|
||||
"""Aggregate only the members recorded in one dated membership snapshot.
|
||||
|
||||
Unknown membership returns an unavailable aggregate and deliberately
|
||||
ignores any supplied stock facts. For known membership, suspended and
|
||||
lifecycle-invalid members are excluded from the expected moneyflow
|
||||
denominator; missing facts remain expected and reduce coverage.
|
||||
|
||||
Args:
|
||||
snapshot: Dated sector identity and point-in-time member codes.
|
||||
stock_facts: Normalized facts that may contain records outside the sector.
|
||||
|
||||
Returns:
|
||||
A yuan-denominated aggregate with explicit membership and moneyflow coverage.
|
||||
|
||||
Raises:
|
||||
ValueError: If member facts have a date mismatch or duplicate stock code.
|
||||
"""
|
||||
|
||||
if snapshot.status is MembershipStatus.UNKNOWN:
|
||||
return SectorDailyAggregate(
|
||||
trade_date=snapshot.trade_date,
|
||||
sector_type=snapshot.sector_type,
|
||||
sector_code=snapshot.sector_code,
|
||||
sector_name=snapshot.sector_name,
|
||||
member_count=0,
|
||||
valid_sample_count=0,
|
||||
net_amount_yuan=None,
|
||||
turnover_yuan=None,
|
||||
membership_coverage=Decimal(0),
|
||||
moneyflow_coverage=Decimal(0),
|
||||
)
|
||||
|
||||
members = set(snapshot.member_codes)
|
||||
facts_by_code: dict[str, StockDailyFact] = {}
|
||||
for fact in stock_facts:
|
||||
if fact.ts_code not in members:
|
||||
continue
|
||||
if fact.trade_date != snapshot.trade_date:
|
||||
raise ValueError("member stock facts must match the snapshot trade_date")
|
||||
if fact.ts_code in facts_by_code:
|
||||
raise ValueError("member stock facts must have unique ts_code values")
|
||||
facts_by_code[fact.ts_code] = fact
|
||||
|
||||
net_amount_total = Decimal(0)
|
||||
turnover_total = Decimal(0)
|
||||
valid_count = 0
|
||||
expected_count = 0
|
||||
for member_code in snapshot.member_codes:
|
||||
fact = facts_by_code.get(member_code)
|
||||
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
|
||||
net_amount = fact.net_amount_yuan
|
||||
turnover = fact.turnover_yuan
|
||||
if net_amount is None or turnover is None:
|
||||
raise ValueError("available stock facts require both amounts")
|
||||
net_amount_total += net_amount
|
||||
turnover_total += turnover
|
||||
valid_count += 1
|
||||
|
||||
moneyflow_coverage = (
|
||||
Decimal(valid_count) / Decimal(expected_count) if expected_count else Decimal(1)
|
||||
)
|
||||
|
||||
return SectorDailyAggregate(
|
||||
trade_date=snapshot.trade_date,
|
||||
sector_type=snapshot.sector_type,
|
||||
sector_code=snapshot.sector_code,
|
||||
sector_name=snapshot.sector_name,
|
||||
member_count=len(snapshot.member_codes),
|
||||
valid_sample_count=valid_count,
|
||||
net_amount_yuan=net_amount_total if valid_count else None,
|
||||
turnover_yuan=turnover_total if valid_count else None,
|
||||
membership_coverage=Decimal(1),
|
||||
moneyflow_coverage=moneyflow_coverage,
|
||||
)
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Transparent, versioned metric strategies for the independent radar."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import Protocol
|
||||
|
||||
from .models import (
|
||||
MetricKind,
|
||||
MetricObservation,
|
||||
MetricQuality,
|
||||
MetricUnit,
|
||||
SectorDailyAggregate,
|
||||
)
|
||||
|
||||
|
||||
class MetricStrategy(Protocol):
|
||||
"""Calculate one named metric from a sector's point-in-time daily history."""
|
||||
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
unit: MetricUnit
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Return the target date observation without inventing missing inputs."""
|
||||
|
||||
...
|
||||
|
||||
|
||||
def _target_aggregate(
|
||||
history: Iterable[SectorDailyAggregate], target_trade_date: date
|
||||
) -> SectorDailyAggregate:
|
||||
matches = tuple(row for row in history if row.trade_date == target_trade_date)
|
||||
if len(matches) != 1:
|
||||
raise ValueError("history must contain exactly one target-date aggregate")
|
||||
return matches[0]
|
||||
|
||||
|
||||
def _quality(row: SectorDailyAggregate) -> MetricQuality:
|
||||
if row.valid_sample_count < 5 or row.membership_coverage < 1 or row.moneyflow_coverage < 1:
|
||||
return MetricQuality.AVAILABLE_LIMITED_SAMPLE
|
||||
return MetricQuality.AVAILABLE
|
||||
|
||||
|
||||
def _observation(
|
||||
row: SectorDailyAggregate,
|
||||
*,
|
||||
metric_kind: MetricKind,
|
||||
metric_version: str,
|
||||
unit: MetricUnit,
|
||||
value: Decimal | None,
|
||||
quality: MetricQuality | None = None,
|
||||
) -> MetricObservation:
|
||||
return MetricObservation(
|
||||
trade_date=row.trade_date,
|
||||
sector_type=row.sector_type,
|
||||
sector_code=row.sector_code,
|
||||
sector_name=row.sector_name,
|
||||
metric_kind=metric_kind,
|
||||
metric_version=metric_version,
|
||||
implementation_kind="independent",
|
||||
unit=unit,
|
||||
value=value,
|
||||
quality=(
|
||||
MetricQuality.UNAVAILABLE
|
||||
if value is None
|
||||
else quality
|
||||
if quality is not None
|
||||
else _quality(row)
|
||||
),
|
||||
member_count=row.member_count,
|
||||
valid_sample_count=row.valid_sample_count,
|
||||
membership_coverage=row.membership_coverage,
|
||||
moneyflow_coverage=row.moneyflow_coverage,
|
||||
)
|
||||
|
||||
|
||||
class AmountNetStrategy:
|
||||
"""Aggregate main net amount and expose it in hundred-million yuan."""
|
||||
|
||||
metric_kind = MetricKind.AMOUNT
|
||||
metric_version = "zhixing_amount_net_bn_v1"
|
||||
unit = MetricUnit.CNY_100M
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Return the target net amount; missing moneyflow remains unavailable."""
|
||||
|
||||
row = _target_aggregate(history, target_trade_date)
|
||||
value = None if row.net_amount_yuan is None else row.net_amount_yuan / Decimal("100000000")
|
||||
return _observation(
|
||||
row,
|
||||
metric_kind=self.metric_kind,
|
||||
metric_version=self.metric_version,
|
||||
unit=self.unit,
|
||||
value=value,
|
||||
)
|
||||
|
||||
|
||||
class RatioTurnoverStrategy:
|
||||
"""Divide aggregated main net amount by aggregated daily turnover."""
|
||||
|
||||
metric_kind = MetricKind.RATIO
|
||||
metric_version = "zhixing_ratio_turnover_v1"
|
||||
unit = MetricUnit.RATIO
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Return a ratio only when numerator and positive denominator exist."""
|
||||
|
||||
row = _target_aggregate(history, target_trade_date)
|
||||
value = None
|
||||
if (
|
||||
row.net_amount_yuan is not None
|
||||
and row.turnover_yuan is not None
|
||||
and row.turnover_yuan > 0
|
||||
):
|
||||
value = row.net_amount_yuan / row.turnover_yuan
|
||||
return _observation(
|
||||
row,
|
||||
metric_kind=self.metric_kind,
|
||||
metric_version=self.metric_version,
|
||||
unit=self.unit,
|
||||
value=value,
|
||||
)
|
||||
|
||||
|
||||
class SwingEqualThreeToTenStrategy:
|
||||
"""Average transparent 3-to-10-day aggregate turnover ratios equally.
|
||||
|
||||
This strategy is deliberately named as a Zhixing implementation. It does
|
||||
not reproduce or imply OneChartLab's unpublished window weights or score.
|
||||
"""
|
||||
|
||||
metric_kind = MetricKind.SWING
|
||||
metric_version = "zhixing_swing_equal_3_10_v1"
|
||||
unit = MetricUnit.RATIO
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Calculate eight complete trading-day windows ending at the target."""
|
||||
|
||||
rows = tuple(sorted(history, key=lambda row: row.trade_date))
|
||||
target = _target_aggregate(rows, target_trade_date)
|
||||
eligible = tuple(row for row in rows if row.trade_date <= target_trade_date)
|
||||
if any(
|
||||
(row.sector_type, row.sector_code) != (target.sector_type, target.sector_code)
|
||||
for row in eligible
|
||||
):
|
||||
raise ValueError("history must contain exactly one sector identity")
|
||||
if len({row.trade_date for row in eligible}) != len(eligible):
|
||||
raise ValueError("history must not contain duplicate trade dates")
|
||||
|
||||
value: Decimal | None = None
|
||||
quality: MetricQuality | None = None
|
||||
if len(eligible) >= 10:
|
||||
latest = eligible[-10:]
|
||||
window_ratios: list[Decimal] = []
|
||||
for window_size in range(3, 11):
|
||||
window = latest[-window_size:]
|
||||
net_amount = Decimal(0)
|
||||
turnover = Decimal(0)
|
||||
for row in window:
|
||||
if row.net_amount_yuan is None or row.turnover_yuan is None:
|
||||
break
|
||||
net_amount += row.net_amount_yuan
|
||||
turnover += row.turnover_yuan
|
||||
else:
|
||||
if turnover <= 0:
|
||||
break
|
||||
window_ratios.append(net_amount / turnover)
|
||||
continue
|
||||
break
|
||||
if len(window_ratios) == 8:
|
||||
value = sum(window_ratios, start=Decimal(0)) / Decimal(8)
|
||||
quality = (
|
||||
MetricQuality.AVAILABLE_LIMITED_SAMPLE
|
||||
if any(_quality(row) is not MetricQuality.AVAILABLE for row in latest)
|
||||
else MetricQuality.AVAILABLE
|
||||
)
|
||||
|
||||
return _observation(
|
||||
target,
|
||||
metric_kind=self.metric_kind,
|
||||
metric_version=self.metric_version,
|
||||
unit=self.unit,
|
||||
value=value,
|
||||
quality=quality,
|
||||
)
|
||||
@@ -0,0 +1,319 @@
|
||||
"""Stable domain values for independently produced sector radar metrics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from enum import StrEnum
|
||||
from typing import Literal
|
||||
|
||||
|
||||
class SectorType(StrEnum):
|
||||
"""Independent ranking pools supported by the first radar release."""
|
||||
|
||||
CONCEPT = "concept"
|
||||
INDUSTRY = "industry"
|
||||
|
||||
|
||||
class MembershipStatus(StrEnum):
|
||||
"""Availability of a point-in-time sector membership snapshot."""
|
||||
|
||||
AVAILABLE = "available"
|
||||
UNKNOWN = "membership_unknown"
|
||||
|
||||
|
||||
class StockFactStatus(StrEnum):
|
||||
"""Why one member does or does not contribute to a daily aggregate."""
|
||||
|
||||
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"
|
||||
|
||||
|
||||
class PublicationStatus(StrEnum):
|
||||
"""Immutable build states retained for audit and last-good selection."""
|
||||
|
||||
RUNNING = "running"
|
||||
SUCCESS = "success"
|
||||
PARTIAL = "partial"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class MetricKind(StrEnum):
|
||||
"""User-facing metric families without borrowing private score names."""
|
||||
|
||||
AMOUNT = "amount"
|
||||
RATIO = "ratio"
|
||||
SWING = "swing"
|
||||
|
||||
|
||||
class MetricQuality(StrEnum):
|
||||
"""Whether a metric is usable and whether its sample needs a warning."""
|
||||
|
||||
AVAILABLE = "available"
|
||||
AVAILABLE_LIMITED_SAMPLE = "available_limited_sample"
|
||||
UNAVAILABLE = "unavailable"
|
||||
|
||||
|
||||
class MetricUnit(StrEnum):
|
||||
"""Units exposed by independent metric strategies."""
|
||||
|
||||
CNY_100M = "CNY_100M"
|
||||
RATIO = "ratio"
|
||||
|
||||
|
||||
class RankSide(StrEnum):
|
||||
"""Ordinary percentile views exposed by the ranking read model."""
|
||||
|
||||
TOP = "top"
|
||||
BOTTOM = "bottom"
|
||||
ALL = "all"
|
||||
|
||||
|
||||
def _validate_finite_decimal(value: Decimal | None, field_name: str) -> None:
|
||||
"""Reject non-finite domain values while preserving missing values."""
|
||||
|
||||
if value is not None and not value.is_finite():
|
||||
raise ValueError(f"{field_name} must be finite or None")
|
||||
|
||||
|
||||
def _validate_coverage(value: Decimal, field_name: str) -> None:
|
||||
"""Require a finite fraction in the inclusive zero-to-one range."""
|
||||
|
||||
_validate_finite_decimal(value, field_name)
|
||||
if value < 0 or value > 1:
|
||||
raise ValueError(f"{field_name} must be between 0 and 1")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SectorMembershipSnapshot:
|
||||
"""One sector's membership as observed for exactly one trade date.
|
||||
|
||||
``UNKNOWN`` is an explicit fact: callers must not substitute a current
|
||||
member list when the historical snapshot is unavailable.
|
||||
"""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
member_codes: tuple[str, ...]
|
||||
status: MembershipStatus
|
||||
source_version: str
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate identity, deterministic membership, and unknown semantics."""
|
||||
|
||||
if not self.sector_code.strip():
|
||||
raise ValueError("sector_code must not be empty")
|
||||
if not self.sector_name.strip():
|
||||
raise ValueError("sector_name must not be empty")
|
||||
if not self.source_version.strip():
|
||||
raise ValueError("source_version must not be empty")
|
||||
if any(not code.strip() for code in self.member_codes):
|
||||
raise ValueError("member_codes must not contain empty values")
|
||||
if len(self.member_codes) != len(set(self.member_codes)):
|
||||
raise ValueError("member_codes must be unique")
|
||||
if self.status is MembershipStatus.UNKNOWN and self.member_codes:
|
||||
raise ValueError("unknown membership must not expose member_codes")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StockDailyFact:
|
||||
"""Normalized daily turnover and moneyflow for one member.
|
||||
|
||||
Amounts are expressed in yuan. Available facts require both source
|
||||
values, including an observed zero. Non-available statuses cannot carry
|
||||
amounts because doing so would blur missing, suspended, and lifecycle
|
||||
semantics at the metric boundary.
|
||||
"""
|
||||
|
||||
trade_date: date
|
||||
ts_code: str
|
||||
status: StockFactStatus
|
||||
turnover_yuan: Decimal | None = None
|
||||
net_amount_yuan: Decimal | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject incomplete available facts and hidden non-finite values."""
|
||||
|
||||
if not self.ts_code.strip():
|
||||
raise ValueError("ts_code must not be empty")
|
||||
_validate_finite_decimal(self.turnover_yuan, "turnover_yuan")
|
||||
_validate_finite_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 RadarPublication:
|
||||
"""Traceable identity and lifecycle of one immutable radar build revision."""
|
||||
|
||||
publication_id: str
|
||||
target_trade_date: date
|
||||
status: PublicationStatus
|
||||
source_version: str
|
||||
universe_version: str
|
||||
metric_versions: tuple[str, ...]
|
||||
input_hash: str | None
|
||||
coverage: Decimal
|
||||
started_at: datetime
|
||||
finished_at: datetime | None = None
|
||||
error_summary: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Keep running and terminal lifecycle timestamps internally consistent."""
|
||||
|
||||
if not self.publication_id.strip():
|
||||
raise ValueError("publication_id must not be empty")
|
||||
if not self.source_version.strip() or not self.universe_version.strip():
|
||||
raise ValueError("publication source versions must not be empty")
|
||||
if not self.metric_versions or any(not value.strip() for value in self.metric_versions):
|
||||
raise ValueError("metric_versions must contain named strategies")
|
||||
if len(self.metric_versions) != len(set(self.metric_versions)):
|
||||
raise ValueError("metric_versions must be unique")
|
||||
_validate_coverage(self.coverage, "coverage")
|
||||
if self.started_at.tzinfo is None:
|
||||
raise ValueError("started_at must be timezone-aware")
|
||||
is_running = self.status is PublicationStatus.RUNNING
|
||||
if is_running != (self.finished_at is None):
|
||||
raise ValueError("finished_at must be absent only while publication is running")
|
||||
if self.finished_at is not None:
|
||||
if self.finished_at.tzinfo is None:
|
||||
raise ValueError("finished_at must be timezone-aware")
|
||||
if self.finished_at < self.started_at:
|
||||
raise ValueError("finished_at must not precede started_at")
|
||||
if self.status is PublicationStatus.SUCCESS and self.input_hash is None:
|
||||
raise ValueError("successful publication requires input_hash")
|
||||
if self.input_hash is not None and (
|
||||
len(self.input_hash) != 64
|
||||
or any(character not in "0123456789abcdef" for character in self.input_hash)
|
||||
):
|
||||
raise ValueError("input_hash must be a lowercase SHA-256 hex digest")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SectorDailyAggregate:
|
||||
"""One sector's point-in-time daily facts after source normalization.
|
||||
|
||||
Amounts use yuan so strategies cannot accidentally mix Tushare's
|
||||
``moneyflow_dc.net_amount`` (ten-thousand yuan) with ``daily.amount``
|
||||
(thousand yuan). ``None`` means missing source data; zero remains an
|
||||
observed value.
|
||||
"""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
member_count: int
|
||||
valid_sample_count: int
|
||||
net_amount_yuan: Decimal | None
|
||||
turnover_yuan: Decimal | None
|
||||
membership_coverage: Decimal
|
||||
moneyflow_coverage: Decimal
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate counts, coverage, and finite normalized values."""
|
||||
|
||||
if not self.sector_code.strip():
|
||||
raise ValueError("sector_code must not be empty")
|
||||
if not self.sector_name.strip():
|
||||
raise ValueError("sector_name must not be empty")
|
||||
if self.member_count < 0:
|
||||
raise ValueError("member_count must not be negative")
|
||||
if not 0 <= self.valid_sample_count <= self.member_count:
|
||||
raise ValueError("valid_sample_count must be within member_count")
|
||||
_validate_finite_decimal(self.net_amount_yuan, "net_amount_yuan")
|
||||
_validate_finite_decimal(self.turnover_yuan, "turnover_yuan")
|
||||
_validate_coverage(self.membership_coverage, "membership_coverage")
|
||||
_validate_coverage(self.moneyflow_coverage, "moneyflow_coverage")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MetricObservation:
|
||||
"""One versioned independent metric value ready for cross-sectional ranking."""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
implementation_kind: Literal["independent"]
|
||||
unit: MetricUnit
|
||||
value: Decimal | None
|
||||
quality: MetricQuality
|
||||
member_count: int
|
||||
valid_sample_count: int
|
||||
membership_coverage: Decimal
|
||||
moneyflow_coverage: Decimal
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Keep unavailable and finite-value states internally consistent."""
|
||||
|
||||
_validate_finite_decimal(self.value, "value")
|
||||
if self.value is None and self.quality is not MetricQuality.UNAVAILABLE:
|
||||
raise ValueError("a missing metric value must be unavailable")
|
||||
if self.value is not None and self.quality is MetricQuality.UNAVAILABLE:
|
||||
raise ValueError("an unavailable metric must not expose a value")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RankChange:
|
||||
"""One previous-publication rank delta using past minus current rank."""
|
||||
|
||||
days: int
|
||||
value: int | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Limit the public comparison window to one through five days."""
|
||||
|
||||
if not 1 <= self.days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RankedMetric:
|
||||
"""A metric observation with its position inside one independent pool."""
|
||||
|
||||
observation: MetricObservation
|
||||
rank_position: int | None
|
||||
rank_percentile: Decimal | None
|
||||
rank_changes: tuple[RankChange, ...] = ()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require rank position and percentile to be present or absent together."""
|
||||
|
||||
if (self.rank_position is None) != (self.rank_percentile is None):
|
||||
raise ValueError("rank_position and rank_percentile must be paired")
|
||||
if self.rank_position is not None and self.rank_position < 1:
|
||||
raise ValueError("rank_position must be positive")
|
||||
_validate_finite_decimal(self.rank_percentile, "rank_percentile")
|
||||
if self.rank_percentile is not None and not 0 < self.rank_percentile <= 100:
|
||||
raise ValueError("rank_percentile must be within (0, 100]")
|
||||
days = [change.days for change in self.rank_changes]
|
||||
if len(days) != len(set(days)):
|
||||
raise ValueError("rank change days must be unique")
|
||||
|
||||
def rank_change(self, days: int) -> int | None:
|
||||
"""Return one configured rank delta, or ``None`` when history is absent."""
|
||||
|
||||
if not 1 <= days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
return next(
|
||||
(change.value for change in self.rank_changes if change.days == days),
|
||||
None,
|
||||
)
|
||||
@@ -0,0 +1,245 @@
|
||||
"""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,
|
||||
)
|
||||
members_by_sector: dict[str, list[SectorMemberRow]] = {code: [] for code in index_by_code}
|
||||
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")
|
||||
members_by_sector[member.sector_code].append(member)
|
||||
|
||||
records: list[MembershipRecord] = []
|
||||
for sector_code in sorted(index_by_code):
|
||||
index = index_by_code[sector_code]
|
||||
sector_members = members_by_sector[sector_code]
|
||||
snapshot_id = partition_ids.get(sector_code, all_snapshot_id)
|
||||
if snapshot_id is None:
|
||||
raise SourceContractError("sector membership has no source snapshot")
|
||||
if not sector_members:
|
||||
explicit_partition_id = partition_ids.get(sector_code)
|
||||
if explicit_partition_id is None:
|
||||
raise SourceContractError(
|
||||
"missing sector membership requires an explicit empty partition"
|
||||
)
|
||||
records.append(
|
||||
MembershipRecord(
|
||||
source_snapshot_id=explicit_partition_id,
|
||||
trade_date=index.trade_date,
|
||||
sector_type=index.sector_type,
|
||||
sector_code=index.sector_code,
|
||||
sector_name=index.name,
|
||||
stock_code=None,
|
||||
stock_name=None,
|
||||
status=MembershipStatus.UNKNOWN,
|
||||
)
|
||||
)
|
||||
continue
|
||||
for member in sector_members:
|
||||
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.membership_key)))
|
||||
|
||||
|
||||
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: Current ``list_status=L`` Tushare listings.
|
||||
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_current_listed_stock(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_current_listed_stock(stock: StockBasicRow, target: date) -> bool:
|
||||
"""Return whether one current ``L`` row is an eligible radar security.
|
||||
|
||||
The radar intentionally uses the listings observed at build time rather than
|
||||
reconstructing historical delistings. Code, market, and list-date checks keep
|
||||
the existing Shanghai/Shenzhen A-share boundary intact.
|
||||
|
||||
Args:
|
||||
stock: One validated ``stock_basic`` row.
|
||||
target: Radar date whose list date must already have arrived.
|
||||
|
||||
Returns:
|
||||
Whether the security belongs to the build-time radar universe.
|
||||
"""
|
||||
|
||||
if stock.list_status != "L":
|
||||
return False
|
||||
if not stock.ts_code.endswith((".SH", ".SZ")):
|
||||
return False
|
||||
if stock.symbol.startswith(("200", "900")):
|
||||
return False
|
||||
market = stock.market or ""
|
||||
if "北交" in market or "B股" in market.upper():
|
||||
return False
|
||||
return stock.list_date is not None and stock.list_date <= target
|
||||
|
||||
|
||||
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,251 @@
|
||||
"""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, datetime
|
||||
from decimal import Decimal
|
||||
from enum import StrEnum
|
||||
from typing import Protocol
|
||||
|
||||
from .models import (
|
||||
MembershipStatus,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
SectorDailyAggregate,
|
||||
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 or explicit unknown snapshot."""
|
||||
|
||||
source_snapshot_id: str
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
stock_code: str | None
|
||||
stock_name: str | None
|
||||
status: MembershipStatus = MembershipStatus.AVAILABLE
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate available and unknown membership null semantics."""
|
||||
|
||||
_validate_digest(self.source_snapshot_id, "source_snapshot_id")
|
||||
if not self.sector_code.strip() or not self.sector_name.strip():
|
||||
raise ValueError("membership sector identity fields must not be empty")
|
||||
if self.status is MembershipStatus.AVAILABLE:
|
||||
if self.stock_code is None or self.stock_name is None:
|
||||
raise ValueError("available membership requires stock identity")
|
||||
if not self.stock_code.strip() or not self.stock_name.strip():
|
||||
raise ValueError("available membership stock identity must not be empty")
|
||||
elif self.stock_code is not None or self.stock_name is not None:
|
||||
raise ValueError("unknown membership must not expose stock identity")
|
||||
|
||||
@property
|
||||
def membership_key(self) -> str:
|
||||
"""Return a non-null persistence key without inventing a stock code."""
|
||||
|
||||
return self.stock_code if self.stock_code is not None else "__membership_unknown__"
|
||||
|
||||
|
||||
@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 DailyAggregateRecord:
|
||||
"""One exact daily strategy input owned by a publication revision."""
|
||||
|
||||
publication_id: str
|
||||
aggregate: SectorDailyAggregate
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the publication foreign identity."""
|
||||
|
||||
if not self.publication_id.strip():
|
||||
raise ValueError("publication_id must not be empty")
|
||||
|
||||
|
||||
class PublicationSourceGroup(StrEnum):
|
||||
"""Stable source checkpoints that can be retried independently."""
|
||||
|
||||
CALENDAR = "calendar"
|
||||
CONCEPT_INDICES = "concept_indices"
|
||||
INDUSTRY_INDICES = "industry_indices"
|
||||
MEMBERS = "members"
|
||||
STOCK_BASICS = "stock_basics"
|
||||
SUSPENSIONS = "suspensions"
|
||||
DAILY = "daily"
|
||||
MONEYFLOW_DC = "moneyflow_dc"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PublicationSourceRecord:
|
||||
"""One ordered raw snapshot checkpoint attached to a build attempt."""
|
||||
|
||||
publication_id: str
|
||||
source_group: PublicationSourceGroup
|
||||
source_order: int
|
||||
snapshot: SourceSnapshot
|
||||
refresh_on_retry: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the publication identity and deterministic group ordering."""
|
||||
|
||||
if not self.publication_id.strip():
|
||||
raise ValueError("publication_id must not be empty")
|
||||
if self.source_order < 0:
|
||||
raise ValueError("source_order must not be negative")
|
||||
|
||||
|
||||
@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_publication_sources(
|
||||
self, records: Iterable[PublicationSourceRecord]
|
||||
) -> WriteCounts: ...
|
||||
|
||||
def load_publication_sources(
|
||||
self, publication_id: str
|
||||
) -> Sequence[PublicationSourceRecord]: ...
|
||||
|
||||
def mark_publication_sources_for_retry(
|
||||
self,
|
||||
publication_id: str,
|
||||
source_groups: Sequence[PublicationSourceGroup],
|
||||
) -> None: ...
|
||||
|
||||
def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts: ...
|
||||
|
||||
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ...
|
||||
|
||||
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ...
|
||||
|
||||
def create_publication(self, publication: RadarPublication) -> WriteCounts: ...
|
||||
|
||||
def finish_publication(self, publication: RadarPublication) -> None: ...
|
||||
|
||||
def finalize_publication(
|
||||
self,
|
||||
publication: RadarPublication,
|
||||
*,
|
||||
memberships: Iterable[MembershipRecord],
|
||||
stock_facts: Iterable[StockFactRecord],
|
||||
daily_aggregates: Iterable[DailyAggregateRecord],
|
||||
rankings: Iterable[RankingRecord],
|
||||
retry_source_groups: Sequence[PublicationSourceGroup] = (),
|
||||
) -> None: ...
|
||||
|
||||
def recover_running_publications(
|
||||
self, target_trade_date: date, *, finished_at: datetime
|
||||
) -> Sequence[str]: ...
|
||||
|
||||
def discard_running_publication(self, publication_id: str) -> None: ...
|
||||
|
||||
def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts: ...
|
||||
|
||||
def find_reusable_publication(
|
||||
self, target_trade_date: date, input_hash: str
|
||||
) -> RadarPublication | None: ...
|
||||
|
||||
def get_publication(self, publication_id: str) -> RadarPublication | None: ...
|
||||
|
||||
def get_last_good_publication(
|
||||
self, target_trade_date: date | None = None
|
||||
) -> RadarPublication | None: ...
|
||||
|
||||
def get_successful_publication(self, target_trade_date: date) -> RadarPublication | None: ...
|
||||
|
||||
def get_latest_publication(self) -> RadarPublication | None: ...
|
||||
|
||||
def list_successful_dates(self) -> Sequence[date]: ...
|
||||
|
||||
def load_rankings(self, publication_id: str) -> Sequence[RankedMetric]: ...
|
||||
|
||||
def load_daily_aggregate_history(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[SectorDailyAggregate]: ...
|
||||
|
||||
def load_previous_rankings(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[tuple[date, Sequence[RankedMetric]]]: ...
|
||||
@@ -0,0 +1,57 @@
|
||||
"""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,
|
||||
candidate_codes: Sequence[str],
|
||||
) -> 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,216 @@
|
||||
"""Deterministic cross-sectional ranking for independent sector pools."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import defaultdict
|
||||
from collections.abc import Iterable, Mapping
|
||||
from dataclasses import replace
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from .models import (
|
||||
MetricKind,
|
||||
MetricObservation,
|
||||
RankChange,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
|
||||
PoolKey = tuple[date, SectorType, MetricKind, str]
|
||||
SectorMetricKey = tuple[SectorType, str, MetricKind, str]
|
||||
|
||||
|
||||
def _pool_key(observation: MetricObservation) -> PoolKey:
|
||||
return (
|
||||
observation.trade_date,
|
||||
observation.sector_type,
|
||||
observation.metric_kind,
|
||||
observation.metric_version,
|
||||
)
|
||||
|
||||
|
||||
def _sector_metric_key(observation: MetricObservation) -> SectorMetricKey:
|
||||
return (
|
||||
observation.sector_type,
|
||||
observation.sector_code,
|
||||
observation.metric_kind,
|
||||
observation.metric_version,
|
||||
)
|
||||
|
||||
|
||||
def _available_sort_key(observation: MetricObservation) -> tuple[Decimal, str]:
|
||||
if observation.value is None:
|
||||
raise ValueError("unavailable observations cannot use the ranking sort key")
|
||||
return (-observation.value, observation.sector_code)
|
||||
|
||||
|
||||
def rank_metric_observations(
|
||||
observations: Iterable[MetricObservation],
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Rank observations by value within date, type, metric, and version.
|
||||
|
||||
Concept and industry observations never share a pool. Equal metric values
|
||||
use ascending sector code as the documented Zhixing tie-breaker. Missing
|
||||
values remain visible but do not consume a rank.
|
||||
"""
|
||||
|
||||
pools: defaultdict[PoolKey, list[MetricObservation]] = defaultdict(list)
|
||||
for observation in observations:
|
||||
pools[_pool_key(observation)].append(observation)
|
||||
|
||||
result: list[RankedMetric] = []
|
||||
for pool_key in sorted(
|
||||
pools,
|
||||
key=lambda key: (key[0], key[1].value, key[2].value, key[3]),
|
||||
):
|
||||
pool = pools[pool_key]
|
||||
codes = [observation.sector_code for observation in pool]
|
||||
if len(codes) != len(set(codes)):
|
||||
raise ValueError("a ranking pool must not contain duplicate sector codes")
|
||||
|
||||
available = sorted(
|
||||
(observation for observation in pool if observation.value is not None),
|
||||
key=_available_sort_key,
|
||||
)
|
||||
pool_size = len(available)
|
||||
for rank_position, observation in enumerate(available, start=1):
|
||||
rank_percentile = (
|
||||
Decimal(100) * Decimal(pool_size - rank_position + 1) / Decimal(pool_size)
|
||||
)
|
||||
result.append(
|
||||
RankedMetric(
|
||||
observation=observation,
|
||||
rank_position=rank_position,
|
||||
rank_percentile=rank_percentile,
|
||||
)
|
||||
)
|
||||
|
||||
result.extend(
|
||||
RankedMetric(
|
||||
observation=observation,
|
||||
rank_position=None,
|
||||
rank_percentile=None,
|
||||
)
|
||||
for observation in sorted(
|
||||
(observation for observation in pool if observation.value is None),
|
||||
key=lambda observation: observation.sector_code,
|
||||
)
|
||||
)
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def select_percentile_side(
|
||||
rankings: Iterable[RankedMetric], side: RankSide
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Select confirmed inclusive percentile sides without fixed row counts."""
|
||||
|
||||
rows = tuple(rankings)
|
||||
if side is RankSide.ALL:
|
||||
return rows
|
||||
|
||||
threshold_rows = tuple(
|
||||
row
|
||||
for row in rows
|
||||
if row.rank_percentile is not None
|
||||
and (
|
||||
row.rank_percentile >= Decimal(90)
|
||||
if side is RankSide.TOP
|
||||
else row.rank_percentile <= Decimal(10)
|
||||
)
|
||||
)
|
||||
if side is RankSide.TOP:
|
||||
return threshold_rows
|
||||
|
||||
pools: defaultdict[PoolKey, list[RankedMetric]] = defaultdict(list)
|
||||
for row in threshold_rows:
|
||||
pools[_pool_key(row.observation)].append(row)
|
||||
result: list[RankedMetric] = []
|
||||
for pool_key in sorted(
|
||||
pools,
|
||||
key=lambda key: (key[0], key[1].value, key[2].value, key[3]),
|
||||
):
|
||||
result.extend(
|
||||
sorted(
|
||||
pools[pool_key],
|
||||
key=lambda row: (
|
||||
row.observation.value if row.observation.value is not None else Decimal(0),
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def with_rank_changes(
|
||||
current_rankings: Iterable[RankedMetric],
|
||||
history_by_days: Mapping[int, Iterable[RankedMetric]],
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Attach 1-to-5-day deltas without turning missing history into zero."""
|
||||
|
||||
history_indexes: dict[int, dict[SectorMetricKey, int | None]] = {}
|
||||
for days, historical_rankings in history_by_days.items():
|
||||
if not 1 <= days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
index: dict[SectorMetricKey, int | None] = {}
|
||||
for row in historical_rankings:
|
||||
key = _sector_metric_key(row.observation)
|
||||
if key in index:
|
||||
raise ValueError("historical rankings must have unique sector metrics")
|
||||
index[key] = row.rank_position
|
||||
history_indexes[days] = index
|
||||
|
||||
result: list[RankedMetric] = []
|
||||
for row in current_rankings:
|
||||
key = _sector_metric_key(row.observation)
|
||||
changes: list[RankChange] = []
|
||||
for days in sorted(history_indexes):
|
||||
past_rank = history_indexes[days].get(key)
|
||||
value = (
|
||||
past_rank - row.rank_position
|
||||
if past_rank is not None and row.rank_position is not None
|
||||
else None
|
||||
)
|
||||
changes.append(RankChange(days=days, value=value))
|
||||
result.append(replace(row, rank_changes=tuple(changes)))
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def select_rank_change_side(
|
||||
rankings: Iterable[RankedMetric],
|
||||
*,
|
||||
days: int,
|
||||
side: RankSide,
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Select the strongest or weakest ceiling-ten-percent rank changes per pool."""
|
||||
|
||||
if not 1 <= days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
pools: defaultdict[PoolKey, list[RankedMetric]] = defaultdict(list)
|
||||
for row in rankings:
|
||||
pools[_pool_key(row.observation)].append(row)
|
||||
|
||||
result: list[RankedMetric] = []
|
||||
for pool_key in sorted(
|
||||
pools,
|
||||
key=lambda key: (key[0], key[1].value, key[2].value, key[3]),
|
||||
):
|
||||
pool = pools[pool_key]
|
||||
pool_size = sum(row.rank_position is not None for row in pool)
|
||||
take_count = max(1, (pool_size + 9) // 10) if pool_size else 0
|
||||
candidates = tuple(
|
||||
(change, row) for row in pool if (change := row.rank_change(days)) is not None
|
||||
)
|
||||
if side is RankSide.BOTTOM:
|
||||
ordered = sorted(
|
||||
candidates,
|
||||
key=lambda item: (item[0], item[1].observation.sector_code),
|
||||
)
|
||||
else:
|
||||
ordered = sorted(
|
||||
candidates,
|
||||
key=lambda item: (-item[0], item[1].observation.sector_code),
|
||||
)
|
||||
selected = ordered if side is RankSide.ALL else ordered[:take_count]
|
||||
result.extend(row for _, row in selected)
|
||||
return tuple(result)
|
||||
@@ -0,0 +1,503 @@
|
||||
"""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.
|
||||
|
||||
Messages must remain operator-safe contract descriptions. They may name a
|
||||
field or validation rule, but must never interpolate provider values,
|
||||
credentials, request parameters, or raw payloads because production build
|
||||
diagnostics record this message.
|
||||
"""
|
||||
|
||||
def __init__(self, message: str) -> None:
|
||||
"""Create one violation whose diagnostic can be claimed by the nearest boundary."""
|
||||
|
||||
super().__init__(message)
|
||||
self._diagnostic_claimed = False
|
||||
|
||||
@property
|
||||
def operator_message(self) -> str:
|
||||
"""Return a single-line, bounded diagnostic suitable for production logs."""
|
||||
|
||||
return " ".join(str(self).split())[:200]
|
||||
|
||||
def claim_diagnostic(self) -> bool:
|
||||
"""Return whether this boundary should emit the exception's single diagnostic log."""
|
||||
|
||||
if self._diagnostic_claimed:
|
||||
return False
|
||||
self._diagnostic_claimed = True
|
||||
return True
|
||||
|
||||
|
||||
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")
|
||||
)
|
||||
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)
|
||||
limit_reached = row_limit is not None and row_count >= row_limit
|
||||
canonical_rows = sorted(
|
||||
json.dumps(row, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||
for row in normalized_rows
|
||||
)
|
||||
canonical_content = json.dumps(
|
||||
{
|
||||
"rows": canonical_rows,
|
||||
"returned_fields": fields,
|
||||
"row_limit": row_limit,
|
||||
"limit_reached": limit_reached,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
content_sha256 = hashlib.sha256(canonical_content.encode()).hexdigest()
|
||||
identity = json.dumps(
|
||||
{
|
||||
"api_name": api_name,
|
||||
"params": normalized_params,
|
||||
"partition_key": partition_key,
|
||||
"target_trade_date": (
|
||||
target_trade_date.isoformat() if target_trade_date is not None else None
|
||||
),
|
||||
"content_sha256": content_sha256,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
snapshot_id = hashlib.sha256(identity.encode()).hexdigest()
|
||||
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=limit_reached,
|
||||
)
|
||||
|
||||
|
||||
@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 | None
|
||||
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=_optional_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 | None
|
||||
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=_optional_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,460 @@
|
||||
"""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 dataclasses import replace
|
||||
from datetime import date, datetime
|
||||
|
||||
from ..domain.models import PublicationStatus, RadarPublication, RankedMetric, SectorDailyAggregate
|
||||
from ..domain.persistence import (
|
||||
DailyAggregateRecord,
|
||||
MembershipRecord,
|
||||
PublicationSourceGroup,
|
||||
PublicationSourceRecord,
|
||||
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.publication_sources: dict[tuple[str, str, int], PublicationSourceRecord] = {}
|
||||
self.memberships: dict[tuple[str, str, str], MembershipRecord] = {}
|
||||
self.stock_facts: dict[tuple[str, str], StockFactRecord] = {}
|
||||
self.publications: dict[str, RadarPublication] = {}
|
||||
self.daily_aggregates: dict[tuple[str, str, str], DailyAggregateRecord] = {}
|
||||
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_publication_sources(self, records: Iterable[PublicationSourceRecord]) -> WriteCounts:
|
||||
"""Checkpoint completed source groups under one publication attempt."""
|
||||
|
||||
items = tuple(records)
|
||||
for item in items:
|
||||
if item.publication_id not in self.publications:
|
||||
raise ValueError("publication source publication does not exist")
|
||||
if item.snapshot.snapshot_id not in self.source_snapshots:
|
||||
raise ValueError("publication source snapshot does not exist")
|
||||
return self._insert_immutable(
|
||||
self.publication_sources,
|
||||
items,
|
||||
key=lambda item: (
|
||||
item.publication_id,
|
||||
item.source_group.value,
|
||||
item.source_order,
|
||||
),
|
||||
)
|
||||
|
||||
def load_publication_sources(self, publication_id: str) -> Sequence[PublicationSourceRecord]:
|
||||
"""Load source checkpoints in stable group and request order."""
|
||||
|
||||
return tuple(
|
||||
sorted(
|
||||
(
|
||||
item
|
||||
for item in self.publication_sources.values()
|
||||
if item.publication_id == publication_id
|
||||
),
|
||||
key=lambda item: (item.source_group.value, item.source_order),
|
||||
)
|
||||
)
|
||||
|
||||
def mark_publication_sources_for_retry(
|
||||
self,
|
||||
publication_id: str,
|
||||
source_groups: Sequence[PublicationSourceGroup],
|
||||
) -> None:
|
||||
"""Mark only incomplete source groups for a future partial retry."""
|
||||
|
||||
requested = set(source_groups)
|
||||
for key, record in tuple(self.publication_sources.items()):
|
||||
if record.publication_id == publication_id and record.source_group in requested:
|
||||
self.publication_sources[key] = replace(record, refresh_on_retry=True)
|
||||
|
||||
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.membership_key,
|
||||
),
|
||||
)
|
||||
|
||||
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 save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts:
|
||||
"""Insert publication-owned exact strategy inputs idempotently."""
|
||||
|
||||
items = tuple(records)
|
||||
for item in items:
|
||||
if item.publication_id not in self.publications:
|
||||
raise ValueError("daily aggregate publication does not exist")
|
||||
return self._insert_immutable(
|
||||
self.daily_aggregates,
|
||||
items,
|
||||
key=lambda item: (
|
||||
item.publication_id,
|
||||
item.aggregate.sector_type.value,
|
||||
item.aggregate.sector_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 finalize_publication(
|
||||
self,
|
||||
publication: RadarPublication,
|
||||
*,
|
||||
memberships: Iterable[MembershipRecord],
|
||||
stock_facts: Iterable[StockFactRecord],
|
||||
daily_aggregates: Iterable[DailyAggregateRecord],
|
||||
rankings: Iterable[RankingRecord],
|
||||
retry_source_groups: Sequence[PublicationSourceGroup] = (),
|
||||
) -> None:
|
||||
"""Atomically expose all derived rows and the terminal publication in tests."""
|
||||
|
||||
previous = (
|
||||
self.memberships.copy(),
|
||||
self.stock_facts.copy(),
|
||||
self.daily_aggregates.copy(),
|
||||
self.rankings.copy(),
|
||||
self.publication_sources.copy(),
|
||||
self.publications.copy(),
|
||||
)
|
||||
try:
|
||||
self.save_memberships(memberships)
|
||||
self.save_stock_facts(stock_facts)
|
||||
self.save_daily_aggregates(daily_aggregates)
|
||||
self.save_rankings(rankings)
|
||||
self.mark_publication_sources_for_retry(
|
||||
publication.publication_id,
|
||||
retry_source_groups,
|
||||
)
|
||||
self.finish_publication(publication)
|
||||
except Exception:
|
||||
(
|
||||
self.memberships,
|
||||
self.stock_facts,
|
||||
self.daily_aggregates,
|
||||
self.rankings,
|
||||
self.publication_sources,
|
||||
self.publications,
|
||||
) = previous
|
||||
raise
|
||||
|
||||
def recover_running_publications(
|
||||
self, target_trade_date: date, *, finished_at: datetime
|
||||
) -> Sequence[str]:
|
||||
"""Fail orphaned attempts after the caller has acquired the date lock."""
|
||||
|
||||
recovered: list[str] = []
|
||||
for publication_id, publication in tuple(self.publications.items()):
|
||||
if (
|
||||
publication.target_trade_date == target_trade_date
|
||||
and publication.status is PublicationStatus.RUNNING
|
||||
):
|
||||
self.publications[publication_id] = replace(
|
||||
publication,
|
||||
status=PublicationStatus.FAILED,
|
||||
finished_at=finished_at,
|
||||
error_summary="recovered_stale_running",
|
||||
)
|
||||
recovered.append(publication_id)
|
||||
return tuple(sorted(recovered))
|
||||
|
||||
def discard_running_publication(self, publication_id: str) -> None:
|
||||
"""Remove only a provisional duplicate attempt and its owned projections."""
|
||||
|
||||
publication = self.publications.get(publication_id)
|
||||
if publication is None or publication.status is not PublicationStatus.RUNNING:
|
||||
raise ValueError("discarded publication must exist in running status")
|
||||
del self.publications[publication_id]
|
||||
self.publication_sources = {
|
||||
key: item
|
||||
for key, item in self.publication_sources.items()
|
||||
if item.publication_id != publication_id
|
||||
}
|
||||
self.daily_aggregates = {
|
||||
key: item
|
||||
for key, item in self.daily_aggregates.items()
|
||||
if item.publication_id != publication_id
|
||||
}
|
||||
self.rankings = {
|
||||
key: item
|
||||
for key, item in self.rankings.items()
|
||||
if item.publication_id != publication_id
|
||||
}
|
||||
|
||||
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 find_reusable_publication(
|
||||
self, target_trade_date: date, input_hash: str
|
||||
) -> RadarPublication | None:
|
||||
"""Find an identical success or partial revision without hiding failures."""
|
||||
|
||||
return max(
|
||||
(
|
||||
item
|
||||
for item in self.publications.values()
|
||||
if item.status in {PublicationStatus.SUCCESS, PublicationStatus.PARTIAL}
|
||||
and item.target_trade_date == target_trade_date
|
||||
and item.input_hash == input_hash
|
||||
),
|
||||
key=lambda item: (item.finished_at or item.started_at, item.publication_id),
|
||||
default=None,
|
||||
)
|
||||
|
||||
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,
|
||||
item.publication_id,
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
|
||||
def get_successful_publication(self, target_trade_date: date) -> RadarPublication | None:
|
||||
"""Return the latest successful revision for exactly one date."""
|
||||
|
||||
candidates = tuple(
|
||||
publication
|
||||
for publication in self.publications.values()
|
||||
if publication.status is PublicationStatus.SUCCESS
|
||||
and publication.target_trade_date == target_trade_date
|
||||
)
|
||||
return max(
|
||||
candidates,
|
||||
key=lambda item: (item.finished_at or item.started_at, item.publication_id),
|
||||
default=None,
|
||||
)
|
||||
|
||||
def get_latest_publication(self) -> RadarPublication | None:
|
||||
"""Return the newest build attempt regardless of terminal status."""
|
||||
|
||||
return max(
|
||||
self.publications.values(),
|
||||
key=lambda item: (item.target_trade_date, item.started_at, item.publication_id),
|
||||
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,
|
||||
)
|
||||
)
|
||||
|
||||
def load_rankings(self, publication_id: str) -> Sequence[RankedMetric]:
|
||||
"""Load every ranking projection owned by one publication."""
|
||||
|
||||
return tuple(
|
||||
sorted(
|
||||
(
|
||||
record.ranking
|
||||
for record in self.rankings.values()
|
||||
if record.publication_id == publication_id
|
||||
),
|
||||
key=lambda row: (
|
||||
row.observation.sector_type.value,
|
||||
row.observation.metric_version,
|
||||
row.rank_position is None,
|
||||
row.rank_position or 0,
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
def load_daily_aggregate_history(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[SectorDailyAggregate]:
|
||||
"""Load aggregates from the latest successful revision of prior dates."""
|
||||
|
||||
dates = self._previous_successful_dates(target_trade_date, limit_dates)
|
||||
selected_publications = {
|
||||
self._latest_success_for_date(item).publication_id for item in dates
|
||||
}
|
||||
return tuple(
|
||||
record.aggregate
|
||||
for record in self.daily_aggregates.values()
|
||||
if record.publication_id in selected_publications
|
||||
)
|
||||
|
||||
def load_previous_rankings(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[tuple[date, Sequence[RankedMetric]]]:
|
||||
"""Load prior successful rankings newest first for rank-change attachment."""
|
||||
|
||||
result: list[tuple[date, Sequence[RankedMetric]]] = []
|
||||
for trade_date in self._previous_successful_dates(target_trade_date, limit_dates):
|
||||
publication_id = self._latest_success_for_date(trade_date).publication_id
|
||||
result.append(
|
||||
(
|
||||
trade_date,
|
||||
tuple(
|
||||
record.ranking
|
||||
for record in self.rankings.values()
|
||||
if record.publication_id == publication_id
|
||||
),
|
||||
)
|
||||
)
|
||||
return tuple(result)
|
||||
|
||||
def _previous_successful_dates(self, target_trade_date: date, limit: int) -> tuple[date, ...]:
|
||||
if limit < 1:
|
||||
raise ValueError("limit_dates must be positive")
|
||||
return tuple(
|
||||
sorted(
|
||||
{
|
||||
item.target_trade_date
|
||||
for item in self.publications.values()
|
||||
if item.status is PublicationStatus.SUCCESS
|
||||
and item.target_trade_date < target_trade_date
|
||||
},
|
||||
reverse=True,
|
||||
)[:limit]
|
||||
)
|
||||
|
||||
def _latest_success_for_date(self, trade_date: date) -> RadarPublication:
|
||||
return max(
|
||||
(
|
||||
item
|
||||
for item in self.publications.values()
|
||||
if item.status is PublicationStatus.SUCCESS and item.target_trade_date == trade_date
|
||||
),
|
||||
key=lambda item: (item.finished_at or item.started_at, item.publication_id),
|
||||
)
|
||||
|
||||
@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)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,636 @@
|
||||
"""Tushare adapter for replayable sector radar source facts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import UTC, date, datetime
|
||||
from typing import TypeVar, cast
|
||||
|
||||
from zhixing_server.shared.request_coordinator import (
|
||||
DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
RequestCoordinator,
|
||||
TushareSourceError,
|
||||
)
|
||||
|
||||
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")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
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: "行业板块",
|
||||
}
|
||||
_MONEYFLOW_WORKERS = 2
|
||||
|
||||
|
||||
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._now_fn = now_fn
|
||||
self._coordinator = request_coordinator or RequestCoordinator(
|
||||
max_retries=max_retries,
|
||||
backoff_seconds=backoff_seconds,
|
||||
request_interval_seconds=request_interval_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",
|
||||
)
|
||||
try:
|
||||
initial_rows = tuple(SectorMemberRow.from_mapping(row) for row in initial.rows)
|
||||
self._require_target_date(initial_rows, trade_date, "dc_member")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("dc_member", "all", exc)
|
||||
raise
|
||||
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,
|
||||
)
|
||||
try:
|
||||
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")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("dc_member", sector_code, exc)
|
||||
raise
|
||||
snapshots.append(snapshot)
|
||||
merged_rows.extend(partition_rows)
|
||||
|
||||
try:
|
||||
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}
|
||||
explicitly_observed_codes = {
|
||||
snapshot.partition_key
|
||||
for snapshot in snapshots
|
||||
if snapshot.partition_key not in {None, "all"}
|
||||
}
|
||||
if set(expected_codes) - final_codes - explicitly_observed_codes:
|
||||
raise SourceContractError("dc_member response is missing expected sectors")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("dc_member", "merged", exc)
|
||||
raise
|
||||
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 the build-time current ``L`` listings in one explicit partition."""
|
||||
|
||||
snapshot = self._fetch_snapshot(
|
||||
"stock_basic",
|
||||
{"exchange": "", "list_status": "L"},
|
||||
target_trade_date=None,
|
||||
partition_key="L",
|
||||
)
|
||||
rows = tuple(StockBasicRow.from_mapping(row) for row in snapshot.rows)
|
||||
if any(row.list_status != "L" for row in rows):
|
||||
raise SourceContractError("stock_basic returned an unexpected list_status")
|
||||
self._require_unique(rows, key=lambda row: row.ts_code, api_name="stock_basic")
|
||||
return SourceResult((snapshot,), 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,
|
||||
candidate_codes: Sequence[str],
|
||||
) -> SourceResult[MoneyflowDcRow]:
|
||||
"""Fetch full-market moneyflow and refill uncovered current candidates."""
|
||||
|
||||
expected_codes = tuple(sorted(set(candidate_codes)))
|
||||
if tuple(candidate_codes) != expected_codes or any(
|
||||
not code.strip() for code in expected_codes
|
||||
):
|
||||
raise ValueError("candidate_codes must be sorted unique non-empty values")
|
||||
initial = self._fetch_snapshot(
|
||||
"moneyflow_dc",
|
||||
{"trade_date": trade_date.strftime("%Y%m%d")},
|
||||
target_trade_date=trade_date,
|
||||
partition_key="all",
|
||||
)
|
||||
try:
|
||||
initial_rows = tuple(MoneyflowDcRow.from_mapping(row) for row in initial.rows)
|
||||
self._require_target_date(initial_rows, trade_date, "moneyflow_dc")
|
||||
self._require_unique(
|
||||
initial_rows,
|
||||
key=lambda row: (row.trade_date, row.ts_code),
|
||||
api_name="moneyflow_dc",
|
||||
)
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("moneyflow_dc", "all", exc)
|
||||
raise
|
||||
|
||||
returned_codes = {row.ts_code for row in initial_rows}
|
||||
missing_codes = tuple(code for code in expected_codes if code not in returned_codes)
|
||||
if not missing_codes:
|
||||
return SourceResult(
|
||||
(initial,),
|
||||
tuple(sorted(initial_rows, key=lambda row: row.ts_code)),
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(
|
||||
max_workers=_MONEYFLOW_WORKERS,
|
||||
thread_name_prefix="sector-radar-moneyflow",
|
||||
) as executor:
|
||||
futures = {
|
||||
code: executor.submit(self._fetch_moneyflow_partition, trade_date, code)
|
||||
for code in missing_codes
|
||||
}
|
||||
partition_results = tuple(futures[code].result() for code in missing_codes)
|
||||
|
||||
snapshots = [initial]
|
||||
merged_rows = list(initial_rows)
|
||||
for result in partition_results:
|
||||
if result is None:
|
||||
continue
|
||||
snapshot, rows = result
|
||||
snapshots.append(snapshot)
|
||||
merged_rows.extend(rows)
|
||||
try:
|
||||
self._require_unique(
|
||||
merged_rows,
|
||||
key=lambda row: (row.trade_date, row.ts_code),
|
||||
api_name="moneyflow_dc",
|
||||
)
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("moneyflow_dc", "merged", exc)
|
||||
raise
|
||||
return SourceResult(
|
||||
tuple(snapshots),
|
||||
tuple(sorted(merged_rows, key=lambda row: row.ts_code)),
|
||||
)
|
||||
|
||||
def _fetch_moneyflow_partition(
|
||||
self,
|
||||
trade_date: date,
|
||||
ts_code: str,
|
||||
) -> tuple[SourceSnapshot, tuple[MoneyflowDcRow, ...]] | None:
|
||||
"""Return one validated refill partition or preserve an ordinary gap."""
|
||||
|
||||
try:
|
||||
snapshot = self._fetch_snapshot(
|
||||
"moneyflow_dc",
|
||||
{
|
||||
"trade_date": trade_date.strftime("%Y%m%d"),
|
||||
"ts_code": ts_code,
|
||||
},
|
||||
target_trade_date=trade_date,
|
||||
partition_key=ts_code,
|
||||
)
|
||||
except TushareSourceError:
|
||||
logger.warning(
|
||||
"sector_radar_moneyflow_partition_failed partition_key=%s error_type=%s",
|
||||
self._safe_partition_key(ts_code),
|
||||
TushareSourceError.__name__,
|
||||
)
|
||||
return None
|
||||
if not snapshot.rows:
|
||||
logger.warning(
|
||||
"sector_radar_moneyflow_partition_empty partition_key=%s",
|
||||
self._safe_partition_key(ts_code),
|
||||
)
|
||||
return None
|
||||
try:
|
||||
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.trade_date, row.ts_code),
|
||||
api_name="moneyflow_dc",
|
||||
)
|
||||
if any(row.ts_code != ts_code for row in rows):
|
||||
raise SourceContractError("moneyflow_dc partition returned a different ts_code")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("moneyflow_dc", ts_code, exc)
|
||||
raise
|
||||
return snapshot, rows
|
||||
|
||||
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)
|
||||
try:
|
||||
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:
|
||||
missing = ",".join(sorted(missing_fields))
|
||||
raise SourceContractError(
|
||||
f"{api_name} response is missing requested fields: {missing}"
|
||||
)
|
||||
return snapshot
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure(api_name, partition_key, exc)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _log_contract_failure(
|
||||
api_name: str,
|
||||
partition_key: str | None,
|
||||
error: SourceContractError,
|
||||
) -> None:
|
||||
"""Record only operator-safe contract context, never provider payloads."""
|
||||
|
||||
if not error.claim_diagnostic():
|
||||
return
|
||||
logger.error(
|
||||
"sector_radar_source_contract_failed api_name=%s partition_key=%s validation=%s",
|
||||
api_name,
|
||||
TushareSectorRadarAdapter._safe_partition_key(partition_key),
|
||||
error.operator_message,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _safe_partition_key(partition_key: str | None) -> str:
|
||||
"""Keep expected identifiers readable while preventing log-control injection."""
|
||||
|
||||
if partition_key is None:
|
||||
return "all"
|
||||
sanitized = "".join(
|
||||
character if character.isalnum() or character in {".", "_", "-"} else "_"
|
||||
for character in partition_key
|
||||
)
|
||||
return sanitized[:64] or "unknown"
|
||||
|
||||
@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 @@
|
||||
"""Delivery adapters for sector radar build and query use cases."""
|
||||
@@ -0,0 +1,108 @@
|
||||
"""One-shot ``sector-radar-build`` command for external schedulers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from datetime import date
|
||||
|
||||
from ....bootstrap.config import get_settings
|
||||
from ..application.build import BuildSectorRadar, BuildSectorRadarCommand
|
||||
from ..infrastructure.postgres import PostgresSectorRadarRepository
|
||||
from ..infrastructure.tushare import TushareSectorRadarAdapter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
"""Build mutually exclusive single-date, range, and retry modes."""
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Build independent Tushare sector radar publications"
|
||||
)
|
||||
mode = parser.add_mutually_exclusive_group()
|
||||
mode.add_argument("--trade-date", type=_parse_date, help="target date in YYYY-MM-DD")
|
||||
mode.add_argument(
|
||||
"--retry-publication-id",
|
||||
help="resume failed source groups from a partial or failed publication",
|
||||
)
|
||||
mode.add_argument("--start-date", type=_parse_date, help="inclusive backfill start date")
|
||||
parser.add_argument("--end-date", type=_parse_date, help="inclusive backfill end date")
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
"""Execute one build invocation and print a redacted JSON summary."""
|
||||
|
||||
args = build_parser().parse_args(argv)
|
||||
if (args.start_date is None) != (args.end_date is None):
|
||||
raise SystemExit("--start-date and --end-date must be provided together")
|
||||
command = BuildSectorRadarCommand(
|
||||
trade_date=args.trade_date,
|
||||
start_date=args.start_date,
|
||||
end_date=args.end_date,
|
||||
retry_publication_id=args.retry_publication_id,
|
||||
)
|
||||
try:
|
||||
settings = get_settings()
|
||||
logging.basicConfig(
|
||||
level=settings.log_level.upper(),
|
||||
format="%(asctime)s %(levelname)s %(name)s %(message)s",
|
||||
force=True,
|
||||
)
|
||||
logger.info(
|
||||
"sector_radar_build_cli trade_date=%s start_date=%s end_date=%s retry=%s",
|
||||
command.trade_date or "auto",
|
||||
command.start_date or "none",
|
||||
command.end_date or "none",
|
||||
bool(command.retry_publication_id),
|
||||
)
|
||||
source = TushareSectorRadarAdapter.from_token(
|
||||
settings.tushare_token,
|
||||
max_retries=settings.sector_radar_max_retries,
|
||||
backoff_seconds=settings.sector_radar_retry_backoff_seconds,
|
||||
request_interval_seconds=settings.sector_radar_request_interval_seconds,
|
||||
)
|
||||
repository = PostgresSectorRadarRepository(
|
||||
settings.database_url,
|
||||
advisory_lock_key=settings.sector_radar_advisory_lock_key,
|
||||
)
|
||||
try:
|
||||
summary = BuildSectorRadar(
|
||||
source,
|
||||
repository,
|
||||
coverage_threshold=settings.sector_radar_coverage_threshold,
|
||||
).execute(command)
|
||||
finally:
|
||||
repository.close()
|
||||
except Exception as exc: # noqa: BLE001 - CLI boundary returns a redacted scheduler result
|
||||
logger.error("sector_radar_build_initialization_failed error_type=%s", type(exc).__name__)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"status": "failed",
|
||||
"exit_code": 1,
|
||||
"outcomes": [],
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": "sector radar build initialization failed",
|
||||
},
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
)
|
||||
)
|
||||
return 1
|
||||
print(json.dumps(summary.as_dict(), ensure_ascii=False, sort_keys=True))
|
||||
return summary.exit_code
|
||||
|
||||
|
||||
def _parse_date(value: str) -> date:
|
||||
try:
|
||||
return date.fromisoformat(value)
|
||||
except ValueError as exc:
|
||||
raise argparse.ArgumentTypeError("date must use YYYY-MM-DD") from exc
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,302 @@
|
||||
"""HTTP presentation for persisted sector radar rankings."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import atexit
|
||||
import threading
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ....bootstrap.config import Settings, get_settings
|
||||
from ..application.read import (
|
||||
RadarDateIndex,
|
||||
RadarMetricDefinition,
|
||||
RadarQuery,
|
||||
RadarView,
|
||||
RankingPage,
|
||||
ReadSectorRadar,
|
||||
)
|
||||
from ..domain.models import (
|
||||
MetricKind,
|
||||
MetricQuality,
|
||||
MetricUnit,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
from ..infrastructure.postgres import (
|
||||
PostgresSectorRadarRepository,
|
||||
SectorRadarRepositoryError,
|
||||
)
|
||||
|
||||
sector_radar_router = APIRouter()
|
||||
_REPOSITORY_CACHE_LOCK = threading.Lock()
|
||||
_REPOSITORY_CACHE: dict[tuple[str, int], PostgresSectorRadarRepository] = {}
|
||||
|
||||
|
||||
class RadarPublicationResponse(BaseModel):
|
||||
"""Safe publication provenance and quality metadata."""
|
||||
|
||||
publication_id: str
|
||||
target_trade_date: date
|
||||
status: PublicationStatus
|
||||
source_version: str
|
||||
universe_version: str
|
||||
metric_versions: list[str]
|
||||
input_hash: str | None
|
||||
coverage: Decimal = Field(ge=0, le=1)
|
||||
started_at: datetime
|
||||
finished_at: datetime | None
|
||||
error_summary: str | None
|
||||
|
||||
|
||||
class RadarDatesResponse(BaseModel):
|
||||
"""Successful dates plus newest attempt and strict last-good metadata."""
|
||||
|
||||
status: Literal["success", "no_data"]
|
||||
available_dates: list[date]
|
||||
current_attempt: RadarPublicationResponse | None
|
||||
last_good: RadarPublicationResponse | None
|
||||
|
||||
|
||||
class RadarMetricDefinitionResponse(BaseModel):
|
||||
"""Version and labeling for one independent metric implementation."""
|
||||
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
label: str
|
||||
unit: MetricUnit
|
||||
implementation_kind: Literal["independent"]
|
||||
disclaimer: str
|
||||
|
||||
|
||||
class RadarRankingRowResponse(BaseModel):
|
||||
"""One sector ranking row with explicit null and unit semantics."""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
implementation_kind: Literal["independent"]
|
||||
unit: MetricUnit
|
||||
metric_value: Decimal | None
|
||||
quality: MetricQuality
|
||||
member_count: int = Field(ge=0)
|
||||
valid_sample_count: int = Field(ge=0)
|
||||
membership_coverage: Decimal = Field(ge=0, le=1)
|
||||
moneyflow_coverage: Decimal = Field(ge=0, le=1)
|
||||
rank_position: int | None = Field(default=None, ge=1)
|
||||
rank_percentile: Decimal | None = Field(default=None, gt=0, le=100)
|
||||
rank_change_days: int = Field(ge=1, le=5)
|
||||
rank_change: int | None
|
||||
|
||||
|
||||
def _empty_ranking_rows() -> list[RadarRankingRowResponse]:
|
||||
return []
|
||||
|
||||
|
||||
class RadarRankingsResponse(BaseModel):
|
||||
"""One persisted, filtered ranking page."""
|
||||
|
||||
status: Literal["success", "no_data"]
|
||||
requested_trade_date: date | None
|
||||
sector_type: SectorType
|
||||
view: RadarView
|
||||
rank_change_metric: MetricKind
|
||||
rank_change_days: int = Field(ge=1, le=5)
|
||||
side: RankSide
|
||||
search: str | None
|
||||
publication: RadarPublicationResponse | None
|
||||
definition: RadarMetricDefinitionResponse
|
||||
page: int = Field(ge=1)
|
||||
page_size: int = Field(ge=1, le=100)
|
||||
total: int = Field(ge=0)
|
||||
rows: list[RadarRankingRowResponse] = Field(default_factory=_empty_ranking_rows)
|
||||
|
||||
|
||||
def get_sector_radar_reader(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> ReadSectorRadar:
|
||||
"""Return a reader backed by one process-cached PostgreSQL repository."""
|
||||
|
||||
key = (settings.database_url, settings.sector_radar_advisory_lock_key)
|
||||
with _REPOSITORY_CACHE_LOCK:
|
||||
repository = _REPOSITORY_CACHE.get(key)
|
||||
if repository is None:
|
||||
repository = PostgresSectorRadarRepository(
|
||||
settings.database_url,
|
||||
advisory_lock_key=settings.sector_radar_advisory_lock_key,
|
||||
)
|
||||
_REPOSITORY_CACHE[key] = repository
|
||||
return ReadSectorRadar(repository)
|
||||
|
||||
|
||||
def _close_cached_repositories() -> None:
|
||||
"""Close process-owned radar pools during interpreter shutdown."""
|
||||
|
||||
with _REPOSITORY_CACHE_LOCK:
|
||||
repositories = tuple(_REPOSITORY_CACHE.values())
|
||||
_REPOSITORY_CACHE.clear()
|
||||
for repository in repositories:
|
||||
repository.close()
|
||||
|
||||
|
||||
atexit.register(_close_cached_repositories)
|
||||
|
||||
|
||||
@sector_radar_router.get("/dates", response_model=RadarDatesResponse)
|
||||
def get_sector_radar_dates(
|
||||
reader: Annotated[ReadSectorRadar, Depends(get_sector_radar_reader)],
|
||||
) -> RadarDatesResponse:
|
||||
"""Return persisted availability without invoking Tushare."""
|
||||
|
||||
try:
|
||||
return _dates_response(reader.list_dates())
|
||||
except SectorRadarRepositoryError as exc:
|
||||
raise _storage_error() from exc
|
||||
|
||||
|
||||
@sector_radar_router.get("/rankings", response_model=RadarRankingsResponse)
|
||||
def get_sector_radar_rankings(
|
||||
reader: Annotated[ReadSectorRadar, Depends(get_sector_radar_reader)],
|
||||
trade_date: date | None = None,
|
||||
sector_type: SectorType = SectorType.CONCEPT,
|
||||
view: RadarView = RadarView.AMOUNT,
|
||||
rank_change_metric: MetricKind = MetricKind.AMOUNT,
|
||||
rank_change_days: Annotated[int, Query(ge=1, le=5)] = 1,
|
||||
side: RankSide = RankSide.ALL,
|
||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||
page: Annotated[int, Query(ge=1)] = 1,
|
||||
page_size: Annotated[int, Query(ge=1, le=100)] = 20,
|
||||
) -> RadarRankingsResponse:
|
||||
"""Return one filtered page from a successful publication."""
|
||||
|
||||
query = RadarQuery(
|
||||
trade_date=trade_date,
|
||||
sector_type=sector_type,
|
||||
view=view,
|
||||
rank_change_metric=rank_change_metric,
|
||||
rank_change_days=rank_change_days,
|
||||
side=side,
|
||||
search=search.strip() or None if search else None,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
try:
|
||||
return _rankings_response(reader.query(query))
|
||||
except SectorRadarRepositoryError as exc:
|
||||
raise _storage_error() from exc
|
||||
|
||||
|
||||
def _dates_response(index: RadarDateIndex) -> RadarDatesResponse:
|
||||
return RadarDatesResponse(
|
||||
status=index.status,
|
||||
available_dates=list(index.available_dates),
|
||||
current_attempt=(
|
||||
_publication_response(index.current_attempt)
|
||||
if index.current_attempt is not None
|
||||
else None
|
||||
),
|
||||
last_good=(_publication_response(index.last_good) if index.last_good is not None else None),
|
||||
)
|
||||
|
||||
|
||||
def _rankings_response(page: RankingPage) -> RadarRankingsResponse:
|
||||
query = page.query
|
||||
return RadarRankingsResponse(
|
||||
status=page.status,
|
||||
requested_trade_date=query.trade_date,
|
||||
sector_type=query.sector_type,
|
||||
view=query.view,
|
||||
rank_change_metric=query.rank_change_metric,
|
||||
rank_change_days=query.rank_change_days,
|
||||
side=query.side,
|
||||
search=query.search,
|
||||
publication=(
|
||||
_publication_response(page.publication) if page.publication is not None else None
|
||||
),
|
||||
definition=_definition_response(page.definition),
|
||||
page=query.page,
|
||||
page_size=query.page_size,
|
||||
total=page.total,
|
||||
rows=[_ranking_response(row, query.rank_change_days) for row in page.rows],
|
||||
)
|
||||
|
||||
|
||||
def _publication_response(publication: RadarPublication) -> RadarPublicationResponse:
|
||||
return RadarPublicationResponse(
|
||||
publication_id=publication.publication_id,
|
||||
target_trade_date=publication.target_trade_date,
|
||||
status=publication.status,
|
||||
source_version=publication.source_version,
|
||||
universe_version=publication.universe_version,
|
||||
metric_versions=list(publication.metric_versions),
|
||||
input_hash=publication.input_hash,
|
||||
coverage=publication.coverage,
|
||||
started_at=publication.started_at,
|
||||
finished_at=publication.finished_at,
|
||||
error_summary=publication.error_summary,
|
||||
)
|
||||
|
||||
|
||||
def _definition_response(
|
||||
definition: RadarMetricDefinition,
|
||||
) -> RadarMetricDefinitionResponse:
|
||||
return RadarMetricDefinitionResponse(
|
||||
metric_kind=definition.metric_kind,
|
||||
metric_version=definition.metric_version,
|
||||
label=definition.label,
|
||||
unit=definition.unit,
|
||||
implementation_kind=definition.implementation_kind,
|
||||
disclaimer=definition.disclaimer,
|
||||
)
|
||||
|
||||
|
||||
def _ranking_response(row: RankedMetric, rank_change_days: int) -> RadarRankingRowResponse:
|
||||
observation = row.observation
|
||||
return RadarRankingRowResponse(
|
||||
trade_date=observation.trade_date,
|
||||
sector_type=observation.sector_type,
|
||||
sector_code=observation.sector_code,
|
||||
sector_name=observation.sector_name,
|
||||
metric_kind=observation.metric_kind,
|
||||
metric_version=observation.metric_version,
|
||||
implementation_kind=observation.implementation_kind,
|
||||
unit=observation.unit,
|
||||
metric_value=observation.value,
|
||||
quality=observation.quality,
|
||||
member_count=observation.member_count,
|
||||
valid_sample_count=observation.valid_sample_count,
|
||||
membership_coverage=observation.membership_coverage,
|
||||
moneyflow_coverage=observation.moneyflow_coverage,
|
||||
rank_position=row.rank_position,
|
||||
rank_percentile=row.rank_percentile,
|
||||
rank_change_days=rank_change_days,
|
||||
rank_change=row.rank_change(rank_change_days),
|
||||
)
|
||||
|
||||
|
||||
def _storage_error() -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=503,
|
||||
detail={
|
||||
"code": "sector_radar_storage_unavailable",
|
||||
"message": "sector radar storage is unavailable",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RadarDatesResponse",
|
||||
"RadarRankingsResponse",
|
||||
"get_sector_radar_reader",
|
||||
"sector_radar_router",
|
||||
]
|
||||
@@ -0,0 +1,186 @@
|
||||
"""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, rate-limit cooling, and optional request start spacing.
|
||||
|
||||
Provider calls execute outside the coordinator lock and may overlap. When a
|
||||
positive request interval is configured, only their start times are serialized.
|
||||
Injectable time functions keep waits deterministic in tests without coupling
|
||||
the coordinator to any business bounded context.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
request_interval_seconds: float = 0.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.request_interval_seconds = max(0.0, request_interval_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._next_request_start = 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_request_start(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_request_start(self, method_name: str) -> None:
|
||||
"""Reserve one start slot after both shared wait deadlines have elapsed."""
|
||||
|
||||
while True:
|
||||
with self._condition:
|
||||
now = self.clock()
|
||||
start_at = max(self._cooldown_until, self._next_request_start)
|
||||
delay = start_at - now
|
||||
if delay <= 0:
|
||||
self._next_request_start = now + self.request_interval_seconds
|
||||
return
|
||||
if start_at == self._cooldown_until:
|
||||
logger.info(
|
||||
"provider_rate_limit_wait method=%s wait_seconds=%.1f",
|
||||
method_name,
|
||||
delay,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"provider_request_interval_wait method=%s wait_seconds=%.3f",
|
||||
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