fix(sector-radar): 增加安全源契约诊断
This commit is contained in:
@@ -4,6 +4,7 @@ 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
|
||||
@@ -49,6 +50,7 @@ from ..domain.source import (
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceContractError,
|
||||
SourceResult,
|
||||
SourceScalar,
|
||||
SourceSnapshot,
|
||||
@@ -60,6 +62,7 @@ from ..domain.source import (
|
||||
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)
|
||||
@@ -392,6 +395,14 @@ class BuildSectorRadar:
|
||||
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(
|
||||
@@ -453,14 +464,25 @@ class BuildSectorRadar:
|
||||
) -> 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),
|
||||
)
|
||||
try:
|
||||
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),
|
||||
)
|
||||
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]
|
||||
|
||||
@@ -19,7 +19,33 @@ T = TypeVar("T")
|
||||
|
||||
|
||||
class SourceContractError(ValueError):
|
||||
"""A provider response violates the replayable input contract."""
|
||||
"""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):
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from datetime import UTC, date, datetime
|
||||
@@ -32,6 +33,7 @@ from ..domain.source import (
|
||||
)
|
||||
|
||||
T = TypeVar("T")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
FIELDS: dict[str, tuple[str, ...]] = {
|
||||
"trade_cal": ("exchange", "cal_date", "is_open", "pretrade_date"),
|
||||
@@ -197,8 +199,12 @@ class TushareSectorRadarAdapter:
|
||||
target_trade_date=trade_date,
|
||||
partition_key="all",
|
||||
)
|
||||
initial_rows = tuple(SectorMemberRow.from_mapping(row) for row in initial.rows)
|
||||
self._require_target_date(initial_rows, trade_date, "dc_member")
|
||||
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)
|
||||
|
||||
@@ -224,27 +230,35 @@ class TushareSectorRadarAdapter:
|
||||
target_trade_date=trade_date,
|
||||
partition_key=sector_code,
|
||||
)
|
||||
self._reject_limit(snapshot)
|
||||
partition_rows = tuple(SectorMemberRow.from_mapping(row) for row in snapshot.rows)
|
||||
self._require_target_date(partition_rows, trade_date, "dc_member")
|
||||
if any(row.sector_code != sector_code for row in partition_rows):
|
||||
raise SourceContractError("dc_member partition returned a different sector")
|
||||
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)
|
||||
|
||||
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")
|
||||
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))),
|
||||
@@ -372,27 +386,63 @@ class TushareSectorRadarAdapter:
|
||||
|
||||
result = self._coordinator.call(api_name, request)
|
||||
self._sleep_fn(self._request_interval_seconds)
|
||||
columns = getattr(result, "columns", None)
|
||||
returned_fields = (
|
||||
tuple(str(column) for column in cast(Iterable[object], columns))
|
||||
if isinstance(columns, Iterable) and not isinstance(columns, (str, bytes))
|
||||
else None
|
||||
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,
|
||||
)
|
||||
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,
|
||||
|
||||
@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
|
||||
)
|
||||
missing_fields = set(FIELDS[api_name]) - set(snapshot.returned_fields)
|
||||
if snapshot.returned_fields and missing_fields:
|
||||
raise SourceContractError(f"{api_name} response is missing requested fields")
|
||||
return snapshot
|
||||
return sanitized[:64] or "unknown"
|
||||
|
||||
@staticmethod
|
||||
def _as_records(result: object) -> tuple[Mapping[str, object], ...]:
|
||||
|
||||
Reference in New Issue
Block a user