Merge branch 'develop' into codex/point

This commit is contained in:
yuxuanhui
2026-08-31 16:14:35 +08:00
86 changed files with 13384 additions and 217 deletions
@@ -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)
@@ -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