feat(sector-radar): 完成可恢复构建任务

This commit is contained in:
yuxuanhui
2026-08-29 18:35:16 +08:00
parent 284c480a90
commit d9bae722d8
17 changed files with 2515 additions and 65 deletions
@@ -0,0 +1 @@
"""Application use cases for sector radar production and reads."""
@@ -0,0 +1,696 @@
"""One-shot, idempotent sector radar publication orchestration."""
from __future__ import annotations
import hashlib
import json
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 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,
SourceResult,
SourceScalar,
SourceSnapshot,
StockBasicRow,
SuspendRow,
TradeCalendarRow,
)
BuildOutcomeStatus = Literal["success", "partial", "failed", "locked", "unchanged"]
SHANGHAI = ZoneInfo("Asia/Shanghai")
MARKET_DATA_READY_TIME = time(15, 30)
@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)
terminal = (
PublicationStatus.SUCCESS
if 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 "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.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 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],
) -> SourceResult[T]:
"""Replay a completed group or fetch and checkpoint it immediately."""
snapshots = reusable.get(source_group)
if snapshots is None:
result = fetch()
else:
result = SourceResult(
snapshots=snapshots,
rows=tuple(parser(row) for snapshot in snapshots for row in snapshot.rows),
)
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,
)
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),
MoneyflowDcRow.from_mapping,
)
memberships = normalize_memberships(indices, members)
candidate_codes = tuple(sorted({item.stock_code for item in memberships}))
stock_facts = normalize_stock_facts(
target_trade_date=target,
candidate_codes=candidate_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)
for member in inputs.memberships:
grouped[(member.sector_type, member.sector_code, member.sector_name)].append(
member.stock_code
)
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(member_codes)),
status=MembershipStatus.AVAILABLE,
source_version=self._universe_version(inputs.membership_snapshots),
),
facts,
)
for (sector_type, sector_code, sector_name), member_codes in sorted(
grouped.items(), key=lambda item: (str(item[0][0]), item[0][1])
)
)
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(
stock_facts: Sequence[StockFactRecord],
) -> tuple[PublicationSourceGroup, ...]:
groups: list[PublicationSourceGroup] = []
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
@@ -5,14 +5,16 @@ from __future__ import annotations
from collections.abc import Iterable, Sequence
from contextlib import AbstractContextManager
from dataclasses import dataclass
from datetime import date
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,
)
@@ -104,6 +106,52 @@ class RankingRecord:
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."""
@@ -125,16 +173,53 @@ class SectorRadarRepository(Protocol):
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(
@@ -142,3 +227,11 @@ class SectorRadarRepository(Protocol):
) -> RadarPublication | None: ...
def list_successful_dates(self) -> Sequence[date]: ...
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]]]: ...
@@ -4,11 +4,15 @@ from __future__ import annotations
from collections.abc import Callable, Generator, Iterable, Sequence
from contextlib import contextmanager
from datetime import date
from dataclasses import replace
from datetime import date, datetime
from ..domain.models import PublicationStatus, RadarPublication
from ..domain.models import PublicationStatus, RadarPublication, RankedMetric, SectorDailyAggregate
from ..domain.persistence import (
DailyAggregateRecord,
MembershipRecord,
PublicationSourceGroup,
PublicationSourceRecord,
RankingRecord,
StockFactRecord,
WriteCounts,
@@ -21,9 +25,11 @@ class InMemorySectorRadarRepository:
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
@@ -51,6 +57,51 @@ class InMemorySectorRadarRepository:
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."""
@@ -69,6 +120,23 @@ class InMemorySectorRadarRepository:
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."""
@@ -99,6 +167,90 @@ class InMemorySectorRadarRepository:
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."""
@@ -122,6 +274,22 @@ class InMemorySectorRadarRepository:
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 next(
(
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
),
None,
)
def get_last_good_publication(
self, target_trade_date: date | None = None
) -> RadarPublication | None:
@@ -153,6 +321,66 @@ class InMemorySectorRadarRepository:
)
)
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,
)
@staticmethod
def _insert_immutable[K, V](
target: dict[K, V],
@@ -5,16 +5,30 @@ from __future__ import annotations
import threading
from collections.abc import Callable, Generator, Iterable, Sequence
from contextlib import contextmanager
from datetime import date
from datetime import date, datetime
from decimal import Decimal
from typing import Any
from psycopg.types.json import Jsonb
from psycopg_pool import ConnectionPool
from ..domain.models import PublicationStatus, RadarPublication
from ..domain.models import (
MetricKind,
MetricObservation,
MetricQuality,
MetricUnit,
PublicationStatus,
RadarPublication,
RankChange,
RankedMetric,
SectorDailyAggregate,
SectorType,
)
from ..domain.persistence import (
DailyAggregateRecord,
MembershipRecord,
PublicationSourceGroup,
PublicationSourceRecord,
RankingRecord,
StockFactRecord,
WriteCounts,
@@ -145,6 +159,87 @@ class PostgresSectorRadarRepository:
unchanged = len(existing)
return WriteCounts(inserted=len(items) - unchanged, unchanged=unchanged)
def save_publication_sources(self, records: Iterable[PublicationSourceRecord]) -> WriteCounts:
"""Checkpoint ordered source snapshots under a running publication."""
items = tuple(records)
self._require_unique(
items,
lambda item: (item.publication_id, item.source_group, item.source_order),
)
return self._copy_immutable(
"sector_radar_publication_source",
(
"publication_id",
"source_group",
"source_order",
"source_snapshot_id",
"refresh_on_retry",
),
("publication_id", "source_group", "source_order"),
tuple(
(
item.publication_id,
item.source_group.value,
item.source_order,
item.snapshot.snapshot_id,
item.refresh_on_retry,
)
for item in items
),
)
def load_publication_sources(self, publication_id: str) -> Sequence[PublicationSourceRecord]:
"""Load source checkpoints with their sanitized raw snapshots."""
with self._connection() as connection:
rows = connection.execute(
"""
SELECT link.publication_id, link.source_group, link.source_order,
snapshot.id, snapshot.api_name, snapshot.normalized_params,
snapshot.target_trade_date, snapshot.partition_key,
snapshot.observed_at, snapshot.payload, snapshot.row_count,
snapshot.returned_fields, snapshot.content_sha256,
snapshot.row_limit, snapshot.limit_reached,
link.refresh_on_retry
FROM sector_radar_publication_source AS link
JOIN sector_radar_source_snapshot AS snapshot
ON snapshot.id = link.source_snapshot_id
WHERE link.publication_id = %s
ORDER BY link.source_group, link.source_order
""",
(publication_id,),
).fetchall()
return tuple(
PublicationSourceRecord(
publication_id=str(row[0]),
source_group=PublicationSourceGroup(str(row[1])),
source_order=int(row[2]),
snapshot=self._snapshot_from_row(row[3:]),
refresh_on_retry=bool(row[15]),
)
for row in rows
)
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."""
if not source_groups:
return
with self._connection() as connection, connection.transaction():
connection.execute(
"""
UPDATE sector_radar_publication_source
SET refresh_on_retry = TRUE
WHERE publication_id = %s AND source_group = ANY(%s)
""",
(publication_id, [group.value for group in source_groups]),
)
def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts:
"""COPY point-in-time members into an immutable revision key."""
@@ -214,6 +309,53 @@ class PostgresSectorRadarRepository:
rows,
)
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts:
"""COPY exact daily metric inputs under their publication revision."""
items = tuple(records)
self._require_unique(
items,
lambda item: (
item.publication_id,
item.aggregate.sector_type,
item.aggregate.sector_code,
),
)
rows = tuple(
(
item.publication_id,
item.aggregate.trade_date,
item.aggregate.sector_type.value,
item.aggregate.sector_code,
item.aggregate.sector_name,
item.aggregate.member_count,
item.aggregate.valid_sample_count,
item.aggregate.net_amount_yuan,
item.aggregate.turnover_yuan,
item.aggregate.membership_coverage,
item.aggregate.moneyflow_coverage,
)
for item in items
)
return self._copy_immutable(
"sector_radar_daily_aggregate",
(
"publication_id",
"trade_date",
"sector_type",
"sector_code",
"sector_name",
"member_count",
"valid_sample_count",
"net_amount_yuan",
"turnover_yuan",
"membership_coverage",
"moneyflow_coverage",
),
("publication_id", "sector_type", "sector_code"),
rows,
)
def create_publication(self, publication: RadarPublication) -> WriteCounts:
"""Insert a new running publication identity idempotently."""
@@ -248,26 +390,199 @@ class PostgresSectorRadarRepository:
if publication.status is PublicationStatus.RUNNING:
raise ValueError("finished publication must use a terminal status")
with self._connection() as connection, connection.transaction():
result = connection.execute(
self._finish_publication_on_connection(connection, 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:
"""Commit all derived projections and the terminal state in one transaction."""
if publication.status is PublicationStatus.RUNNING:
raise ValueError("finalized publication must use a terminal status")
membership_items = tuple(memberships)
stock_items = tuple(stock_facts)
aggregate_items = tuple(daily_aggregates)
ranking_items = tuple(rankings)
self._require_unique(
membership_items,
lambda item: (item.source_snapshot_id, item.sector_code, item.stock_code),
)
self._require_unique(stock_items, lambda item: (item.fact_revision, item.ts_code))
self._require_unique(
aggregate_items,
lambda item: (
item.publication_id,
item.aggregate.sector_type,
item.aggregate.sector_code,
),
)
self._require_unique(
ranking_items,
lambda item: (
item.publication_id,
item.ranking.observation.sector_type,
item.ranking.observation.sector_code,
item.ranking.observation.metric_version,
),
)
with self._connection() as connection, connection.transaction():
self._copy_immutable_on_connection(
connection,
"sector_radar_membership",
(
"source_snapshot_id",
"trade_date",
"sector_type",
"sector_code",
"sector_name",
"stock_code",
"stock_name",
"membership_status",
),
("source_snapshot_id", "sector_code", "stock_code"),
tuple(
(
item.source_snapshot_id,
item.trade_date,
item.sector_type.value,
item.sector_code,
item.sector_name,
item.stock_code,
item.stock_name,
item.status.value,
)
for item in membership_items
),
)
self._copy_immutable_on_connection(
connection,
"sector_radar_stock_fact",
(
"fact_revision",
"trade_date",
"ts_code",
"source_snapshot_ids",
"status",
"turnover_yuan",
"net_amount_yuan",
),
("fact_revision", "ts_code"),
tuple(
(
item.fact_revision,
item.trade_date,
item.ts_code,
Jsonb(list(item.source_snapshot_ids)),
item.status.value,
item.turnover_yuan,
item.net_amount_yuan,
)
for item in stock_items
),
)
self._copy_immutable_on_connection(
connection,
"sector_radar_daily_aggregate",
(
"publication_id",
"trade_date",
"sector_type",
"sector_code",
"sector_name",
"member_count",
"valid_sample_count",
"net_amount_yuan",
"turnover_yuan",
"membership_coverage",
"moneyflow_coverage",
),
("publication_id", "sector_type", "sector_code"),
tuple(
(
item.publication_id,
item.aggregate.trade_date,
item.aggregate.sector_type.value,
item.aggregate.sector_code,
item.aggregate.sector_name,
item.aggregate.member_count,
item.aggregate.valid_sample_count,
item.aggregate.net_amount_yuan,
item.aggregate.turnover_yuan,
item.aggregate.membership_coverage,
item.aggregate.moneyflow_coverage,
)
for item in aggregate_items
),
)
self._copy_immutable_on_connection(
connection,
"sector_radar_ranking",
(
"publication_id",
"trade_date",
"sector_type",
"sector_code",
"sector_name",
"metric_kind",
"metric_version",
"implementation_kind",
"unit",
"metric_value",
"quality",
"member_count",
"valid_sample_count",
"membership_coverage",
"moneyflow_coverage",
"rank_position",
"rank_percentile",
"rank_changes",
),
("publication_id", "sector_type", "sector_code", "metric_version"),
self._ranking_rows(ranking_items),
)
if retry_source_groups:
connection.execute(
"UPDATE sector_radar_publication_source SET refresh_on_retry = TRUE "
"WHERE publication_id = %s AND source_group = ANY(%s)",
(
publication.publication_id,
[group.value for group in retry_source_groups],
),
)
self._finish_publication_on_connection(connection, publication)
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."""
with self._connection() as connection, connection.transaction():
rows = connection.execute(
"""
UPDATE sector_radar_publication
SET status = %s, source_version = %s, universe_version = %s,
metric_versions = %s, input_hash = %s, coverage = %s,
finished_at = %s, error_summary = %s
WHERE id = %s AND target_trade_date = %s AND status = 'running'
SET status = 'failed', finished_at = %s,
error_summary = 'recovered_stale_running'
WHERE target_trade_date = %s AND status = 'running'
RETURNING id
""",
(
publication.status.value,
publication.source_version,
publication.universe_version,
Jsonb(list(publication.metric_versions)),
publication.input_hash,
publication.coverage,
publication.finished_at,
self._safe_error(publication.error_summary),
publication.publication_id,
publication.target_trade_date,
),
(finished_at, target_trade_date),
).fetchall()
return tuple(sorted(str(row[0]) for row in rows))
def discard_running_publication(self, publication_id: str) -> None:
"""Delete only a provisional duplicate publication and owned projections."""
with self._connection() as connection, connection.transaction():
result = connection.execute(
"DELETE FROM sector_radar_publication WHERE id = %s AND status = 'running'",
(publication_id,),
)
if result.rowcount != 1:
raise SectorRadarRepositoryError("publication is not in running status")
@@ -285,32 +600,6 @@ class PostgresSectorRadarRepository:
item.ranking.observation.metric_version,
),
)
rows: list[tuple[object, ...]] = []
for item in items:
ranking = item.ranking
observation = ranking.observation
rows.append(
(
item.publication_id,
observation.trade_date,
observation.sector_type.value,
observation.sector_code,
observation.sector_name,
observation.metric_kind.value,
observation.metric_version,
observation.implementation_kind,
observation.unit.value,
observation.value,
observation.quality.value,
observation.member_count,
observation.valid_sample_count,
observation.membership_coverage,
observation.moneyflow_coverage,
ranking.rank_position,
ranking.rank_percentile,
Jsonb({str(change.days): change.value for change in ranking.rank_changes}),
)
)
return self._copy_immutable(
"sector_radar_ranking",
(
@@ -334,7 +623,7 @@ class PostgresSectorRadarRepository:
"rank_changes",
),
("publication_id", "sector_type", "sector_code", "metric_version"),
tuple(rows),
self._ranking_rows(items),
)
def get_publication(self, publication_id: str) -> RadarPublication | None:
@@ -347,6 +636,20 @@ class PostgresSectorRadarRepository:
).fetchone()
return None if row is None else self._publication_from_row(row)
def find_reusable_publication(
self, target_trade_date: date, input_hash: str
) -> RadarPublication | None:
"""Find an identical success or partial input revision for idempotent reruns."""
with self._connection() as connection:
row = connection.execute(
self._publication_select() + " WHERE target_trade_date = %s AND input_hash = %s "
"AND status IN ('success', 'partial') "
"ORDER BY finished_at DESC, created_at DESC, id DESC LIMIT 1",
(target_trade_date, input_hash),
).fetchone()
return None if row is None else self._publication_from_row(row)
def get_last_good_publication(
self, target_trade_date: date | None = None
) -> RadarPublication | None:
@@ -380,6 +683,75 @@ class PostgresSectorRadarRepository:
).fetchall()
return tuple(row[0] for row in rows)
def load_daily_aggregate_history(
self, target_trade_date: date, *, limit_dates: int
) -> Sequence[SectorDailyAggregate]:
"""Load prior aggregates from the latest successful revision per date."""
if limit_dates < 1:
raise ValueError("limit_dates must be positive")
with self._connection() as connection:
rows = connection.execute(
"""
WITH selected AS (
SELECT DISTINCT ON (target_trade_date) id, target_trade_date
FROM sector_radar_publication
WHERE status = 'success' AND target_trade_date < %s
ORDER BY target_trade_date DESC, finished_at DESC, created_at DESC, id DESC
LIMIT %s
)
SELECT aggregate.trade_date, aggregate.sector_type, aggregate.sector_code,
aggregate.sector_name, aggregate.member_count,
aggregate.valid_sample_count, aggregate.net_amount_yuan,
aggregate.turnover_yuan, aggregate.membership_coverage,
aggregate.moneyflow_coverage
FROM sector_radar_daily_aggregate AS aggregate
JOIN selected ON selected.id = aggregate.publication_id
ORDER BY aggregate.trade_date, aggregate.sector_type, aggregate.sector_code
""",
(target_trade_date, limit_dates),
).fetchall()
return tuple(self._aggregate_from_row(row) for row in rows)
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 calculation."""
if limit_dates < 1:
raise ValueError("limit_dates must be positive")
with self._connection() as connection:
rows = connection.execute(
"""
WITH selected AS (
SELECT DISTINCT ON (target_trade_date) id, target_trade_date
FROM sector_radar_publication
WHERE status = 'success' AND target_trade_date < %s
ORDER BY target_trade_date DESC, finished_at DESC, created_at DESC, id DESC
LIMIT %s
)
SELECT selected.target_trade_date, ranking.trade_date, ranking.sector_type,
ranking.sector_code, ranking.sector_name, ranking.metric_kind,
ranking.metric_version, ranking.implementation_kind, ranking.unit,
ranking.metric_value, ranking.quality, ranking.member_count,
ranking.valid_sample_count, ranking.membership_coverage,
ranking.moneyflow_coverage, ranking.rank_position,
ranking.rank_percentile, ranking.rank_changes
FROM sector_radar_ranking AS ranking
JOIN selected ON selected.id = ranking.publication_id
ORDER BY selected.target_trade_date DESC, ranking.sector_type,
ranking.metric_version, ranking.rank_position NULLS LAST,
ranking.sector_code
""",
(target_trade_date, limit_dates),
).fetchall()
grouped: dict[date, list[RankedMetric]] = {}
for row in rows:
grouped.setdefault(row[0], []).append(self._ranking_from_row(row[1:]))
return tuple(
(trade_date, tuple(grouped[trade_date])) for trade_date in sorted(grouped, reverse=True)
)
def _copy_immutable(
self,
table: str,
@@ -390,26 +762,46 @@ class PostgresSectorRadarRepository:
if not rows:
return WriteCounts(0, 0)
if table not in {
"sector_radar_publication_source",
"sector_radar_membership",
"sector_radar_stock_fact",
"sector_radar_daily_aggregate",
"sector_radar_ranking",
}:
raise ValueError("unsupported radar staging table")
with self._connection() as connection, connection.transaction():
return self._copy_immutable_on_connection(
connection,
table,
columns,
conflict_columns,
rows,
)
@staticmethod
def _copy_immutable_on_connection(
connection: Any,
table: str,
columns: tuple[str, ...],
conflict_columns: tuple[str, ...],
rows: tuple[tuple[object, ...], ...],
) -> WriteCounts:
if not rows:
return WriteCounts(0, 0)
stage = f"{table}_stage"
column_sql = ", ".join(columns)
conflict_sql = ", ".join(conflict_columns)
with self._connection() as connection, connection.transaction():
cursor = connection.cursor()
cursor.execute(
f"CREATE TEMP TABLE {stage} (LIKE {table} INCLUDING DEFAULTS) ON COMMIT DROP"
)
with cursor.copy(f"COPY {stage} ({column_sql}) FROM STDIN") as copy:
for row in rows:
copy.write_row(row)
inserted_rows = cursor.execute(
f"INSERT INTO {table} ({column_sql}) SELECT {column_sql} FROM {stage} "
f"ON CONFLICT ({conflict_sql}) DO NOTHING RETURNING 1"
).fetchall()
cursor = connection.cursor()
cursor.execute(
f"CREATE TEMP TABLE {stage} (LIKE {table} INCLUDING DEFAULTS) ON COMMIT DROP"
)
with cursor.copy(f"COPY {stage} ({column_sql}) FROM STDIN") as copy:
for row in rows:
copy.write_row(row)
inserted_rows = cursor.execute(
f"INSERT INTO {table} ({column_sql}) SELECT {column_sql} FROM {stage} "
f"ON CONFLICT ({conflict_sql}) DO NOTHING RETURNING 1"
).fetchall()
inserted = len(inserted_rows)
return WriteCounts(inserted=inserted, unchanged=len(rows) - inserted)
@@ -470,6 +862,130 @@ class PostgresSectorRadarRepository:
error_summary=None if row[10] is None else str(row[10]),
)
@staticmethod
def _ranking_rows(items: Sequence[RankingRecord]) -> tuple[tuple[object, ...], ...]:
rows: list[tuple[object, ...]] = []
for item in items:
ranking = item.ranking
observation = ranking.observation
rows.append(
(
item.publication_id,
observation.trade_date,
observation.sector_type.value,
observation.sector_code,
observation.sector_name,
observation.metric_kind.value,
observation.metric_version,
observation.implementation_kind,
observation.unit.value,
observation.value,
observation.quality.value,
observation.member_count,
observation.valid_sample_count,
observation.membership_coverage,
observation.moneyflow_coverage,
ranking.rank_position,
ranking.rank_percentile,
Jsonb({str(change.days): change.value for change in ranking.rank_changes}),
)
)
return tuple(rows)
@staticmethod
def _finish_publication_on_connection(
connection: Any,
publication: RadarPublication,
) -> None:
result = connection.execute(
"""
UPDATE sector_radar_publication
SET status = %s, source_version = %s, universe_version = %s,
metric_versions = %s, input_hash = %s, coverage = %s,
finished_at = %s, error_summary = %s
WHERE id = %s AND target_trade_date = %s AND status = 'running'
""",
(
publication.status.value,
publication.source_version,
publication.universe_version,
Jsonb(list(publication.metric_versions)),
publication.input_hash,
publication.coverage,
publication.finished_at,
PostgresSectorRadarRepository._safe_error(publication.error_summary),
publication.publication_id,
publication.target_trade_date,
),
)
if result.rowcount != 1:
raise SectorRadarRepositoryError("publication is not in running status")
@staticmethod
def _snapshot_from_row(row: tuple[Any, ...]) -> SourceSnapshot:
normalized_params = row[2]
return SourceSnapshot(
snapshot_id=str(row[0]),
api_name=str(row[1]),
normalized_params=tuple(
sorted((str(key), str(value)) for key, value in normalized_params.items())
),
target_trade_date=row[3],
partition_key=None if row[4] is None else str(row[4]),
observed_at=row[5],
rows=tuple(dict(item) for item in row[6]),
row_count=int(row[7]),
returned_fields=tuple(str(value) for value in row[8]),
content_sha256=str(row[9]),
row_limit=None if row[10] is None else int(row[10]),
limit_reached=bool(row[11]),
)
@staticmethod
def _aggregate_from_row(row: tuple[Any, ...]) -> SectorDailyAggregate:
return SectorDailyAggregate(
trade_date=row[0],
sector_type=SectorType(str(row[1])),
sector_code=str(row[2]),
sector_name=str(row[3]),
member_count=int(row[4]),
valid_sample_count=int(row[5]),
net_amount_yuan=None if row[6] is None else Decimal(str(row[6])),
turnover_yuan=None if row[7] is None else Decimal(str(row[7])),
membership_coverage=Decimal(str(row[8])),
moneyflow_coverage=Decimal(str(row[9])),
)
@staticmethod
def _ranking_from_row(row: tuple[Any, ...]) -> RankedMetric:
raw_changes = row[16]
changes = tuple(
RankChange(days=int(days), value=None if value is None else int(value))
for days, value in sorted(raw_changes.items(), key=lambda item: int(item[0]))
)
observation = MetricObservation(
trade_date=row[0],
sector_type=SectorType(str(row[1])),
sector_code=str(row[2]),
sector_name=str(row[3]),
metric_kind=MetricKind(str(row[4])),
metric_version=str(row[5]),
implementation_kind="independent",
unit=MetricUnit(str(row[7])),
value=None if row[8] is None else Decimal(str(row[8])),
quality=MetricQuality(str(row[9])),
member_count=int(row[10]),
valid_sample_count=int(row[11]),
membership_coverage=Decimal(str(row[12])),
moneyflow_coverage=Decimal(str(row[13])),
)
return RankedMetric(
observation=observation,
rank_position=None if row[14] is None else int(row[14]),
rank_percentile=None if row[15] is None else Decimal(str(row[15])),
rank_changes=changes,
)
@staticmethod
def _safe_error(message: str | None) -> str | None:
return None if message is None else " ".join(message.split())[:500]
@@ -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())