fix(sector-radar): 增加安全源契约诊断 #16

Merged
sakibcc merged 1 commits from develop into main 2026-08-30 00:15:01 +08:00
5 changed files with 262 additions and 48 deletions
Showing only changes of commit 6e217ac70f - Show all commits
@@ -4,6 +4,7 @@ from __future__ import annotations
import hashlib import hashlib
import json import json
import logging
from collections import defaultdict from collections import defaultdict
from collections.abc import Callable, Mapping, Sequence from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass, replace from dataclasses import dataclass, replace
@@ -49,6 +50,7 @@ from ..domain.source import (
MoneyflowDcRow, MoneyflowDcRow,
SectorIndexRow, SectorIndexRow,
SectorMemberRow, SectorMemberRow,
SourceContractError,
SourceResult, SourceResult,
SourceScalar, SourceScalar,
SourceSnapshot, SourceSnapshot,
@@ -60,6 +62,7 @@ from ..domain.source import (
BuildOutcomeStatus = Literal["success", "partial", "failed", "locked", "unchanged"] BuildOutcomeStatus = Literal["success", "partial", "failed", "locked", "unchanged"]
SHANGHAI = ZoneInfo("Asia/Shanghai") SHANGHAI = ZoneInfo("Asia/Shanghai")
MARKET_DATA_READY_TIME = time(15, 30) MARKET_DATA_READY_TIME = time(15, 30)
logger = logging.getLogger(__name__)
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -392,6 +395,14 @@ class BuildSectorRadar:
if publication is not None if publication is not None
else self._failure_id(target.trade_date) 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: if publication is not None and publication_created:
self.repository.finish_publication( self.repository.finish_publication(
RadarPublication( RadarPublication(
@@ -453,6 +464,7 @@ class BuildSectorRadar:
) -> SourceResult[T]: ) -> SourceResult[T]:
"""Replay a completed group or fetch and checkpoint it immediately.""" """Replay a completed group or fetch and checkpoint it immediately."""
try:
snapshots = reusable.get(source_group) snapshots = reusable.get(source_group)
if snapshots is None: if snapshots is None:
result = fetch() result = fetch()
@@ -461,6 +473,16 @@ class BuildSectorRadar:
snapshots=snapshots, snapshots=snapshots,
rows=tuple(parser(row) for snapshot in snapshots for row in snapshot.rows), 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: if not result.snapshots:
raise ValueError("source group must include at least one replay snapshot") raise ValueError("source group must include at least one replay snapshot")
snapshot_ids = [snapshot.snapshot_id for snapshot in result.snapshots] snapshot_ids = [snapshot.snapshot_id for snapshot in result.snapshots]
@@ -19,7 +19,33 @@ T = TypeVar("T")
class SourceContractError(ValueError): 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): class SourceTruncatedError(SourceContractError):
@@ -2,6 +2,7 @@
from __future__ import annotations from __future__ import annotations
import logging
import time import time
from collections.abc import Callable, Iterable, Mapping, Sequence from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import UTC, date, datetime from datetime import UTC, date, datetime
@@ -32,6 +33,7 @@ from ..domain.source import (
) )
T = TypeVar("T") T = TypeVar("T")
logger = logging.getLogger(__name__)
FIELDS: dict[str, tuple[str, ...]] = { FIELDS: dict[str, tuple[str, ...]] = {
"trade_cal": ("exchange", "cal_date", "is_open", "pretrade_date"), "trade_cal": ("exchange", "cal_date", "is_open", "pretrade_date"),
@@ -197,8 +199,12 @@ class TushareSectorRadarAdapter:
target_trade_date=trade_date, target_trade_date=trade_date,
partition_key="all", partition_key="all",
) )
try:
initial_rows = tuple(SectorMemberRow.from_mapping(row) for row in initial.rows) initial_rows = tuple(SectorMemberRow.from_mapping(row) for row in initial.rows)
self._require_target_date(initial_rows, trade_date, "dc_member") 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} 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) missing_codes = tuple(code for code in expected_codes if code not in returned_codes)
@@ -224,14 +230,19 @@ class TushareSectorRadarAdapter:
target_trade_date=trade_date, target_trade_date=trade_date,
partition_key=sector_code, partition_key=sector_code,
) )
try:
self._reject_limit(snapshot) self._reject_limit(snapshot)
partition_rows = tuple(SectorMemberRow.from_mapping(row) for row in snapshot.rows) partition_rows = tuple(SectorMemberRow.from_mapping(row) for row in snapshot.rows)
self._require_target_date(partition_rows, trade_date, "dc_member") self._require_target_date(partition_rows, trade_date, "dc_member")
if any(row.sector_code != sector_code for row in partition_rows): if any(row.sector_code != sector_code for row in partition_rows):
raise SourceContractError("dc_member partition returned a different sector") 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) snapshots.append(snapshot)
merged_rows.extend(partition_rows) merged_rows.extend(partition_rows)
try:
self._require_unique( self._require_unique(
merged_rows, merged_rows,
key=lambda row: (row.trade_date, row.sector_code, row.stock_code), key=lambda row: (row.trade_date, row.sector_code, row.stock_code),
@@ -245,6 +256,9 @@ class TushareSectorRadarAdapter:
} }
if set(expected_codes) - final_codes - explicitly_observed_codes: if set(expected_codes) - final_codes - explicitly_observed_codes:
raise SourceContractError("dc_member response is missing expected sectors") raise SourceContractError("dc_member response is missing expected sectors")
except SourceContractError as exc:
self._log_contract_failure("dc_member", "merged", exc)
raise
return SourceResult( return SourceResult(
tuple(snapshots), tuple(snapshots),
tuple(sorted(merged_rows, key=lambda row: (row.sector_code, row.stock_code))), tuple(sorted(merged_rows, key=lambda row: (row.sector_code, row.stock_code))),
@@ -372,6 +386,7 @@ class TushareSectorRadarAdapter:
result = self._coordinator.call(api_name, request) result = self._coordinator.call(api_name, request)
self._sleep_fn(self._request_interval_seconds) self._sleep_fn(self._request_interval_seconds)
try:
columns = getattr(result, "columns", None) columns = getattr(result, "columns", None)
returned_fields = ( returned_fields = (
tuple(str(column) for column in cast(Iterable[object], columns)) tuple(str(column) for column in cast(Iterable[object], columns))
@@ -391,8 +406,43 @@ class TushareSectorRadarAdapter:
) )
missing_fields = set(FIELDS[api_name]) - set(snapshot.returned_fields) missing_fields = set(FIELDS[api_name]) - set(snapshot.returned_fields)
if snapshot.returned_fields and missing_fields: if snapshot.returned_fields and missing_fields:
raise SourceContractError(f"{api_name} response is missing requested fields") missing = ",".join(sorted(missing_fields))
raise SourceContractError(
f"{api_name} response is missing requested fields: {missing}"
)
return snapshot 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 @staticmethod
def _as_records(result: object) -> tuple[Mapping[str, object], ...]: def _as_records(result: object) -> tuple[Mapping[str, object], ...]:
@@ -1,7 +1,10 @@
import logging
from collections.abc import Sequence from collections.abc import Sequence
from datetime import UTC, date, datetime, timedelta from datetime import UTC, date, datetime, timedelta
from decimal import Decimal from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.application.build import ( from zhixing_server.modules.sector_radar.application.build import (
BuildSectorRadar, BuildSectorRadar,
BuildSectorRadarCommand, BuildSectorRadarCommand,
@@ -19,6 +22,7 @@ from zhixing_server.modules.sector_radar.domain.source import (
MoneyflowDcRow, MoneyflowDcRow,
SectorIndexRow, SectorIndexRow,
SectorMemberRow, SectorMemberRow,
SourceContractError,
SourceResult, SourceResult,
SourceSnapshot, SourceSnapshot,
StockBasicRow, StockBasicRow,
@@ -284,6 +288,37 @@ def test_successful_build_is_idempotent_and_failed_retry_preserves_last_good() -
assert any(item.status is PublicationStatus.FAILED for item in repository.publications.values()) assert any(item.status is PublicationStatus.FAILED for item in repository.publications.values())
def test_source_contract_failure_is_logged_with_safe_build_context(
caplog: pytest.LogCaptureFixture,
) -> None:
class InvalidDailySource(FakeRadarSource):
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
self.calls.append("daily")
raise SourceContractError("daily returned duplicate business keys")
repository = InMemorySectorRadarRepository()
caplog.set_level(
logging.ERROR,
logger="zhixing_server.modules.sector_radar.application.build",
)
failed = BuildSectorRadar(
InvalidDailySource(),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
messages = "\n".join(record.getMessage() for record in caplog.records)
assert failed.status == "failed"
assert failed.outcomes[0].error_message == "input or source contract validation failed"
assert "sector_radar_source_group_contract_failed" in messages
assert "source_group=daily" in messages
assert "publication_id=radar-20260828-running-" in messages
assert "validation=daily returned duplicate business keys" in messages
assert len(caplog.records) == 1
def test_partial_coverage_and_lock_have_distinct_exit_codes() -> None: def test_partial_coverage_and_lock_have_distinct_exit_codes() -> None:
repository = InMemorySectorRadarRepository() repository = InMemorySectorRadarRepository()
partial = BuildSectorRadar( partial = BuildSectorRadar(
@@ -1,3 +1,4 @@
import logging
from collections.abc import Mapping from collections.abc import Mapping
from datetime import UTC, date, datetime from datetime import UTC, date, datetime
from decimal import Decimal from decimal import Decimal
@@ -132,6 +133,86 @@ def test_non_finite_source_values_are_rejected() -> None:
make_adapter(client).fetch_daily(TARGET_DATE) make_adapter(client).fetch_daily(TARGET_DATE)
def test_contract_failure_log_identifies_member_partition_without_payload(
caplog: pytest.LogCaptureFixture,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "dc_member", 2)
client = QueryClient(
{
(
"dc_member",
"",
): (
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "private-payload-marker",
},
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000002.SZ",
"name": "private-payload-marker",
},
),
(
"dc_member",
"BK0001.DC",
): (
{
"trade_date": "20260828",
"ts_code": "BK9999.DC",
"con_code": "000001.SZ",
"name": "private-payload-marker",
},
),
}
)
caplog.set_level(
logging.ERROR,
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
)
with pytest.raises(SourceContractError, match="different sector"):
make_adapter(client).fetch_sector_members(TARGET_DATE, ("BK0001.DC",))
messages = "\n".join(record.getMessage() for record in caplog.records)
assert "sector_radar_source_contract_failed" in messages
assert "api_name=dc_member" in messages
assert "partition_key=BK0001.DC" in messages
assert "validation=dc_member partition returned a different sector" in messages
assert "private-payload-marker" not in messages
assert len(caplog.records) == 1
def test_merged_member_contract_failure_has_one_interface_level_log(
caplog: pytest.LogCaptureFixture,
) -> None:
duplicate = {
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "private-payload-marker",
}
client = QueryClient({("dc_member", ""): (duplicate, duplicate)})
caplog.set_level(
logging.ERROR,
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
)
with pytest.raises(SourceContractError, match="duplicate business keys"):
make_adapter(client).fetch_sector_members(TARGET_DATE, ("BK0001.DC",))
messages = "\n".join(record.getMessage() for record in caplog.records)
assert "api_name=dc_member" in messages
assert "partition_key=merged" in messages
assert "validation=dc_member returned duplicate business keys" in messages
assert "private-payload-marker" not in messages
assert len(caplog.records) == 1
def test_dc_member_reloads_by_sector_when_the_all_market_call_hits_limit( def test_dc_member_reloads_by_sector_when_the_all_market_call_hits_limit(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None: