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
@@ -0,0 +1,234 @@
"""Create replayable independent sector radar tables."""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects.postgresql import JSONB
revision: str = "0004_sector_radar"
down_revision: str | None = "0003_market_integrity_checks"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Create source, point-in-time fact, publication, and ranking tables."""
op.create_table(
"sector_radar_source_snapshot",
sa.Column("id", sa.String(64), primary_key=True),
sa.Column("api_name", sa.String(32), nullable=False),
sa.Column("normalized_params", JSONB, nullable=False),
sa.Column("target_trade_date", sa.Date()),
sa.Column("partition_key", sa.String(64)),
sa.Column("observed_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("payload", JSONB, nullable=False),
sa.Column("row_count", sa.Integer(), nullable=False),
sa.Column("returned_fields", JSONB, nullable=False),
sa.Column("content_sha256", sa.String(64), nullable=False),
sa.Column("row_limit", sa.Integer()),
sa.Column("limit_reached", sa.Boolean(), nullable=False),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.CheckConstraint("row_count >= 0", name="ck_sector_radar_source_row_count"),
sa.CheckConstraint(
"row_limit IS NULL OR row_limit > 0", name="ck_sector_radar_source_limit"
),
)
op.create_index(
"ix_sector_radar_source_api_date",
"sector_radar_source_snapshot",
["api_name", "target_trade_date", "observed_at"],
)
op.create_table(
"sector_radar_membership",
sa.Column(
"source_snapshot_id",
sa.String(64),
sa.ForeignKey("sector_radar_source_snapshot.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("trade_date", sa.Date(), nullable=False),
sa.Column("sector_type", sa.String(16), nullable=False),
sa.Column("sector_code", sa.String(16), nullable=False),
sa.Column("sector_name", sa.String(128), nullable=False),
sa.Column("stock_code", sa.String(12), nullable=False),
sa.Column("stock_name", sa.String(128), nullable=False),
sa.Column("membership_status", sa.String(32), nullable=False),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.PrimaryKeyConstraint("source_snapshot_id", "sector_code", "stock_code"),
sa.CheckConstraint(
"sector_type IN ('concept', 'industry')",
name="ck_sector_radar_membership_type",
),
sa.CheckConstraint(
"membership_status = 'available'",
name="ck_sector_radar_membership_status",
),
)
op.create_index(
"ix_sector_radar_membership_date_sector",
"sector_radar_membership",
["trade_date", "sector_type", "sector_code"],
)
op.create_table(
"sector_radar_stock_fact",
sa.Column("fact_revision", sa.String(64), nullable=False),
sa.Column("trade_date", sa.Date(), nullable=False),
sa.Column("ts_code", sa.String(12), nullable=False),
sa.Column("source_snapshot_ids", JSONB, nullable=False),
sa.Column("status", sa.String(32), nullable=False),
sa.Column("turnover_yuan", sa.Numeric(28, 6)),
sa.Column("net_amount_yuan", sa.Numeric(28, 6)),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.PrimaryKeyConstraint("fact_revision", "ts_code"),
sa.CheckConstraint(
"turnover_yuan IS NULL OR turnover_yuan >= 0",
name="ck_sector_radar_stock_turnover",
),
)
op.create_index(
"ix_sector_radar_stock_fact_date",
"sector_radar_stock_fact",
["trade_date", "ts_code"],
)
op.create_table(
"sector_radar_publication",
sa.Column("id", sa.String(64), primary_key=True),
sa.Column("target_trade_date", sa.Date(), nullable=False),
sa.Column("status", sa.String(16), nullable=False),
sa.Column("source_version", sa.String(128), nullable=False),
sa.Column("universe_version", sa.String(128), nullable=False),
sa.Column("metric_versions", JSONB, nullable=False),
sa.Column("input_hash", sa.String(64)),
sa.Column("coverage", sa.Numeric(8, 6), nullable=False),
sa.Column("started_at", sa.DateTime(timezone=True), nullable=False),
sa.Column("finished_at", sa.DateTime(timezone=True)),
sa.Column("error_summary", sa.String(500)),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.CheckConstraint(
"status IN ('running', 'success', 'partial', 'failed')",
name="ck_sector_radar_publication_status",
),
sa.CheckConstraint(
"coverage >= 0 AND coverage <= 1",
name="ck_sector_radar_publication_coverage",
),
sa.CheckConstraint(
"(status = 'running' AND finished_at IS NULL) OR "
"(status <> 'running' AND finished_at IS NOT NULL)",
name="ck_sector_radar_publication_finished",
),
)
op.create_index(
"ix_sector_radar_publication_status_date",
"sector_radar_publication",
["status", "target_trade_date", "finished_at"],
)
op.create_index(
"uq_sector_radar_publication_running_date",
"sector_radar_publication",
["target_trade_date"],
unique=True,
postgresql_where=sa.text("status = 'running'"),
)
op.create_table(
"sector_radar_ranking",
sa.Column(
"publication_id",
sa.String(64),
sa.ForeignKey("sector_radar_publication.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("trade_date", sa.Date(), nullable=False),
sa.Column("sector_type", sa.String(16), nullable=False),
sa.Column("sector_code", sa.String(16), nullable=False),
sa.Column("sector_name", sa.String(128), nullable=False),
sa.Column("metric_kind", sa.String(16), nullable=False),
sa.Column("metric_version", sa.String(128), nullable=False),
sa.Column("implementation_kind", sa.String(16), nullable=False),
sa.Column("unit", sa.String(16), nullable=False),
sa.Column("metric_value", sa.Numeric(28, 12)),
sa.Column("quality", sa.String(32), nullable=False),
sa.Column("member_count", sa.Integer(), nullable=False),
sa.Column("valid_sample_count", sa.Integer(), nullable=False),
sa.Column("membership_coverage", sa.Numeric(8, 6), nullable=False),
sa.Column("moneyflow_coverage", sa.Numeric(8, 6), nullable=False),
sa.Column("rank_position", sa.Integer()),
sa.Column("rank_percentile", sa.Numeric(18, 12)),
sa.Column("rank_changes", JSONB, nullable=False),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.PrimaryKeyConstraint(
"publication_id",
"sector_type",
"sector_code",
"metric_version",
),
sa.CheckConstraint(
"sector_type IN ('concept', 'industry')",
name="ck_sector_radar_ranking_type",
),
sa.CheckConstraint(
"implementation_kind = 'independent'",
name="ck_sector_radar_ranking_implementation",
),
)
op.create_index(
"ix_sector_radar_ranking_query",
"sector_radar_ranking",
["publication_id", "sector_type", "metric_version", "rank_position"],
)
def downgrade() -> None:
"""Drop only sector radar tables in dependency-safe order."""
op.drop_index("ix_sector_radar_ranking_query", table_name="sector_radar_ranking")
op.drop_table("sector_radar_ranking")
op.drop_index(
"uq_sector_radar_publication_running_date",
table_name="sector_radar_publication",
)
op.drop_index(
"ix_sector_radar_publication_status_date",
table_name="sector_radar_publication",
)
op.drop_table("sector_radar_publication")
op.drop_index("ix_sector_radar_stock_fact_date", table_name="sector_radar_stock_fact")
op.drop_table("sector_radar_stock_fact")
op.drop_index(
"ix_sector_radar_membership_date_sector",
table_name="sector_radar_membership",
)
op.drop_table("sector_radar_membership")
op.drop_index("ix_sector_radar_source_api_date", table_name="sector_radar_source_snapshot")
op.drop_table("sector_radar_source_snapshot")
@@ -0,0 +1,155 @@
"""Persist publication-owned sector daily aggregates for exact metric replay."""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "0005_radar_daily_aggregate"
down_revision: str | None = "0004_sector_radar"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Create source recovery links and exact multi-day metric inputs."""
op.create_table(
"sector_radar_publication_source",
sa.Column(
"publication_id",
sa.String(64),
sa.ForeignKey("sector_radar_publication.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("source_group", sa.String(32), nullable=False),
sa.Column("source_order", sa.Integer(), nullable=False),
sa.Column(
"refresh_on_retry",
sa.Boolean(),
nullable=False,
server_default=sa.false(),
),
sa.Column(
"source_snapshot_id",
sa.String(64),
sa.ForeignKey("sector_radar_source_snapshot.id", ondelete="RESTRICT"),
nullable=False,
),
sa.PrimaryKeyConstraint("publication_id", "source_group", "source_order"),
sa.UniqueConstraint(
"publication_id",
"source_group",
"source_snapshot_id",
name="uq_sector_radar_publication_source_snapshot",
),
sa.CheckConstraint(
"source_group IN ('calendar', 'concept_indices', 'industry_indices', "
"'members', 'stock_basics', 'suspensions', 'daily', 'moneyflow_dc')",
name="ck_sector_radar_publication_source_group",
),
sa.CheckConstraint(
"source_order >= 0",
name="ck_sector_radar_publication_source_order",
),
)
op.create_index(
"ix_sector_radar_publication_source_snapshot",
"sector_radar_publication_source",
["source_snapshot_id"],
)
op.create_unique_constraint(
"uq_sector_radar_publication_id_date",
"sector_radar_publication",
["id", "target_trade_date"],
)
op.create_foreign_key(
"fk_sector_radar_ranking_publication_date",
"sector_radar_ranking",
"sector_radar_publication",
["publication_id", "trade_date"],
["id", "target_trade_date"],
ondelete="CASCADE",
)
op.create_table(
"sector_radar_daily_aggregate",
sa.Column("publication_id", sa.String(64), nullable=False),
sa.Column("trade_date", sa.Date(), nullable=False),
sa.Column("sector_type", sa.String(16), nullable=False),
sa.Column("sector_code", sa.String(16), nullable=False),
sa.Column("sector_name", sa.String(128), nullable=False),
sa.Column("member_count", sa.Integer(), nullable=False),
sa.Column("valid_sample_count", sa.Integer(), nullable=False),
sa.Column("net_amount_yuan", sa.Numeric(28, 6)),
sa.Column("turnover_yuan", sa.Numeric(28, 6)),
sa.Column("membership_coverage", sa.Numeric(8, 6), nullable=False),
sa.Column("moneyflow_coverage", sa.Numeric(8, 6), nullable=False),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.text("now()"),
),
sa.PrimaryKeyConstraint("publication_id", "sector_type", "sector_code"),
sa.ForeignKeyConstraint(
["publication_id", "trade_date"],
["sector_radar_publication.id", "sector_radar_publication.target_trade_date"],
name="fk_sector_radar_daily_aggregate_publication_date",
ondelete="CASCADE",
),
sa.CheckConstraint(
"sector_type IN ('concept', 'industry')",
name="ck_sector_radar_daily_aggregate_type",
),
sa.CheckConstraint(
"member_count >= 0 AND valid_sample_count >= 0 AND valid_sample_count <= member_count",
name="ck_sector_radar_daily_aggregate_counts",
),
sa.CheckConstraint(
"membership_coverage >= 0 AND membership_coverage <= 1 "
"AND moneyflow_coverage >= 0 AND moneyflow_coverage <= 1",
name="ck_sector_radar_daily_aggregate_coverage",
),
sa.CheckConstraint(
"turnover_yuan IS NULL OR (turnover_yuan >= 0 AND "
"turnover_yuan NOT IN ('NaN'::numeric, 'Infinity'::numeric))",
name="ck_sector_radar_daily_aggregate_turnover",
),
sa.CheckConstraint(
"net_amount_yuan IS NULL OR net_amount_yuan NOT IN "
"('NaN'::numeric, 'Infinity'::numeric, '-Infinity'::numeric)",
name="ck_sector_radar_daily_aggregate_net_amount",
),
)
op.create_index(
"ix_sector_radar_daily_aggregate_history",
"sector_radar_daily_aggregate",
["trade_date", "sector_type", "sector_code"],
)
def downgrade() -> None:
"""Drop only the replay aggregate extension."""
op.drop_index(
"ix_sector_radar_daily_aggregate_history",
table_name="sector_radar_daily_aggregate",
)
op.drop_table("sector_radar_daily_aggregate")
op.drop_constraint(
"fk_sector_radar_ranking_publication_date",
"sector_radar_ranking",
type_="foreignkey",
)
op.drop_constraint(
"uq_sector_radar_publication_id_date",
"sector_radar_publication",
type_="unique",
)
op.drop_index(
"ix_sector_radar_publication_source_snapshot",
table_name="sector_radar_publication_source",
)
op.drop_table("sector_radar_publication_source")
@@ -0,0 +1,103 @@
"""Persist explicit unknown point-in-time sector membership snapshots."""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "0006_membership_unknown"
down_revision: str | None = "0005_radar_daily_aggregate"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Allow one null-stock marker for an explicitly empty sector partition."""
op.add_column(
"sector_radar_membership",
sa.Column("membership_key", sa.String(32), nullable=True),
)
op.execute("UPDATE sector_radar_membership SET membership_key = stock_code")
op.drop_constraint(
"sector_radar_membership_pkey",
"sector_radar_membership",
type_="primary",
)
op.drop_constraint(
"ck_sector_radar_membership_status",
"sector_radar_membership",
type_="check",
)
op.alter_column(
"sector_radar_membership",
"membership_key",
existing_type=sa.String(32),
nullable=False,
)
op.alter_column(
"sector_radar_membership",
"stock_code",
existing_type=sa.String(12),
nullable=True,
)
op.alter_column(
"sector_radar_membership",
"stock_name",
existing_type=sa.String(128),
nullable=True,
)
op.create_primary_key(
"sector_radar_membership_pkey",
"sector_radar_membership",
["source_snapshot_id", "sector_code", "membership_key"],
)
op.create_check_constraint(
"ck_sector_radar_membership_status",
"sector_radar_membership",
"(membership_status = 'available' "
"AND stock_code IS NOT NULL AND stock_name IS NOT NULL "
"AND membership_key = stock_code) OR "
"(membership_status = 'membership_unknown' "
"AND stock_code IS NULL AND stock_name IS NULL "
"AND membership_key = '__membership_unknown__')",
)
def downgrade() -> None:
"""Discard unknown markers and restore the available-member-only schema."""
op.execute("DELETE FROM sector_radar_membership WHERE membership_status = 'membership_unknown'")
op.drop_constraint(
"ck_sector_radar_membership_status",
"sector_radar_membership",
type_="check",
)
op.drop_constraint(
"sector_radar_membership_pkey",
"sector_radar_membership",
type_="primary",
)
op.alter_column(
"sector_radar_membership",
"stock_code",
existing_type=sa.String(12),
nullable=False,
)
op.alter_column(
"sector_radar_membership",
"stock_name",
existing_type=sa.String(128),
nullable=False,
)
op.drop_column("sector_radar_membership", "membership_key")
op.create_primary_key(
"sector_radar_membership_pkey",
"sector_radar_membership",
["source_snapshot_id", "sector_code", "stock_code"],
)
op.create_check_constraint(
"ck_sector_radar_membership_status",
"sector_radar_membership",
"membership_status = 'available'",
)
+1
View File
@@ -36,6 +36,7 @@ packages = ["src/zhixing_server"]
[project.scripts]
market-data-sync = "zhixing_server.modules.market_data.presentation.cli:main"
sector-radar-build = "zhixing_server.modules.sector_radar.presentation.cli:main"
[tool.pytest.ini_options]
addopts = "-ra --strict-config --strict-markers"
@@ -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
@@ -39,6 +39,13 @@ def test_postgres_migration_creates_market_data_contract(
"selection_run",
"selection_run_item",
"selection_signal",
"sector_radar_source_snapshot",
"sector_radar_membership",
"sector_radar_stock_fact",
"sector_radar_publication",
"sector_radar_ranking",
"sector_radar_daily_aggregate",
"sector_radar_publication_source",
} <= tables
item_columns = {column["name"] for column in inspector.get_columns("selection_run_item")}
assert {
@@ -0,0 +1,175 @@
import os
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
from pathlib import Path
import psycopg
import pytest
from alembic import command
from alembic.config import Config
from zhixing_server.bootstrap.config import sqlalchemy_database_url
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
PublicationStatus,
RadarPublication,
SectorType,
StockFactStatus,
)
from zhixing_server.modules.sector_radar.domain.persistence import (
MembershipRecord,
StockFactRecord,
)
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
PostgresSectorRadarRepository,
)
TARGET_DATE = date(2099, 1, 4)
STARTED_AT = datetime(2099, 1, 4, 17, 30, tzinfo=UTC)
def prepare_database(database_url: str) -> None:
server_root = Path(__file__).parents[2]
config = Config(str(server_root / "alembic.ini"))
sqlalchemy_url = sqlalchemy_database_url(database_url)
config.set_main_option("sqlalchemy.url", sqlalchemy_url.replace("%", "%%"))
command.upgrade(config, "head")
@pytest.mark.integration
def test_postgres_sector_radar_revisions_and_last_good() -> None:
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
if not database_url:
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
prepare_database(database_url)
snapshot = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20990104"},
rows=(
{
"trade_date": "20990104",
"ts_code": "BKTEST.DC",
"con_code": "000001.SZ",
"name": "测试股票",
},
),
target_trade_date=TARGET_DATE,
observed_at=STARTED_AT,
)
fact_revision = "b" * 64
publication_ids = ("test-sector-radar-success", "test-sector-radar-failed")
with psycopg.connect(database_url) as connection, connection.transaction():
connection.execute(
"DELETE FROM sector_radar_publication WHERE id = ANY(%s)",
(list(publication_ids),),
)
connection.execute(
"DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s",
(fact_revision,),
)
connection.execute(
"DELETE FROM sector_radar_source_snapshot WHERE id = %s",
(snapshot.snapshot_id,),
)
repository = PostgresSectorRadarRepository(database_url, max_connections=2)
try:
assert repository.save_source_snapshots((snapshot,)).inserted == 1
assert (
repository.save_source_snapshots(
(replace(snapshot, observed_at=STARTED_AT + timedelta(minutes=1)),)
).unchanged
== 1
)
assert (
repository.save_memberships(
(
MembershipRecord(
source_snapshot_id=snapshot.snapshot_id,
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BKTEST.DC",
sector_name="测试概念",
stock_code="000001.SZ",
stock_name="测试股票",
status=MembershipStatus.AVAILABLE,
),
)
).inserted
== 1
)
assert (
repository.save_stock_facts(
(
StockFactRecord(
fact_revision=fact_revision,
source_snapshot_ids=(snapshot.snapshot_id,),
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("1000"),
net_amount_yuan=Decimal("100"),
),
)
).inserted
== 1
)
running = RadarPublication(
publication_id=publication_ids[0],
target_trade_date=TARGET_DATE,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version=snapshot.content_sha256,
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal(0),
started_at=STARTED_AT,
)
repository.create_publication(running)
repository.finish_publication(
replace(
running,
status=PublicationStatus.SUCCESS,
input_hash="a" * 64,
coverage=Decimal(1),
finished_at=STARTED_AT + timedelta(minutes=5),
)
)
failed = replace(
running,
publication_id=publication_ids[1],
started_at=STARTED_AT + timedelta(minutes=6),
)
repository.create_publication(failed)
repository.finish_publication(
replace(
failed,
status=PublicationStatus.FAILED,
coverage=Decimal("0.8"),
finished_at=STARTED_AT + timedelta(minutes=7),
error_summary="safe_error",
)
)
last_good = repository.get_last_good_publication(TARGET_DATE)
assert last_good is not None
assert last_good.publication_id == publication_ids[0]
finally:
repository.close()
with psycopg.connect(database_url) as connection, connection.transaction():
connection.execute(
"DELETE FROM sector_radar_publication WHERE id = ANY(%s)",
(list(publication_ids),),
)
connection.execute(
"DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s",
(fact_revision,),
)
connection.execute(
"DELETE FROM sector_radar_source_snapshot WHERE id = %s",
(snapshot.snapshot_id,),
)
@@ -0,0 +1,230 @@
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
import pytest
from fastapi.testclient import TestClient
from pydantic import ValidationError
from zhixing_server.bootstrap.app import create_app
from zhixing_server.modules.sector_radar.application.read import (
RadarDateIndex,
RadarMetricDefinition,
RadarQuery,
RadarView,
RankingPage,
)
from zhixing_server.modules.sector_radar.domain.metrics import AmountNetStrategy
from zhixing_server.modules.sector_radar.domain.models import (
MetricKind,
MetricObservation,
MetricQuality,
MetricUnit,
PublicationStatus,
RadarPublication,
RankChange,
RankedMetric,
RankSide,
SectorType,
)
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
SectorRadarRepositoryError,
)
from zhixing_server.modules.sector_radar.presentation.http import (
RadarRankingRowResponse,
get_sector_radar_reader,
)
TARGET_DATE = date(2026, 8, 28)
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
def _publication(
publication_id: str,
status: PublicationStatus,
*,
trade_date: date = TARGET_DATE,
) -> RadarPublication:
return RadarPublication(
publication_id=publication_id,
target_trade_date=trade_date,
status=status,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-v1",
metric_versions=(AmountNetStrategy.metric_version,),
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
coverage=Decimal(1) if status is PublicationStatus.SUCCESS else Decimal("0.8"),
started_at=NOW,
finished_at=None if status is PublicationStatus.RUNNING else NOW + timedelta(minutes=5),
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
)
def _ranking() -> RankedMetric:
return RankedMetric(
observation=MetricObservation(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="机器人",
metric_kind=MetricKind.AMOUNT,
metric_version=AmountNetStrategy.metric_version,
implementation_kind="independent",
unit=MetricUnit.CNY_100M,
value=Decimal("12.5"),
quality=MetricQuality.AVAILABLE,
member_count=20,
valid_sample_count=19,
membership_coverage=Decimal(1),
moneyflow_coverage=Decimal("0.95"),
),
rank_position=1,
rank_percentile=Decimal(100),
rank_changes=(RankChange(days=5, value=3),),
)
class FakeReader:
def __init__(self, *, no_data: bool = False, fail: bool = False) -> None:
self.fail = fail
self.last_query: RadarQuery | None = None
success = _publication("publication-success", PublicationStatus.SUCCESS)
current = _publication(
"publication-partial",
PublicationStatus.PARTIAL,
trade_date=TARGET_DATE + timedelta(days=1),
)
self.date_index = RadarDateIndex(
available_dates=() if no_data else (TARGET_DATE,),
current_attempt=None if no_data else current,
last_good=None if no_data else success,
)
query = RadarQuery()
self.page = RankingPage(
status="no_data" if no_data else "success",
query=query,
publication=None if no_data else success,
definition=RadarMetricDefinition(
metric_kind=MetricKind.AMOUNT,
metric_version=AmountNetStrategy.metric_version,
label="主力净流入(知行独立实现)",
unit=MetricUnit.CNY_100M,
),
rows=() if no_data else (_ranking(),),
total=0 if no_data else 1,
)
def list_dates(self) -> RadarDateIndex:
if self.fail:
raise SectorRadarRepositoryError("private database detail")
return self.date_index
def query(self, query: RadarQuery) -> RankingPage:
if self.fail:
raise SectorRadarRepositoryError("private database detail")
self.last_query = query
return replace(self.page, query=query)
def _client(reader: FakeReader) -> TestClient:
application = create_app()
application.dependency_overrides[get_sector_radar_reader] = lambda: reader
return TestClient(application)
def test_dates_exposes_partial_attempt_without_replacing_last_good() -> None:
response = _client(FakeReader()).get("/api/v1/sector-radar/dates")
assert response.status_code == 200
payload = response.json()
assert payload["status"] == "success"
assert payload["available_dates"] == ["2026-08-28"]
assert payload["current_attempt"]["status"] == "partial"
assert payload["last_good"]["status"] == "success"
assert payload["last_good"]["coverage"] == "1"
def test_rankings_maps_filters_and_independent_metric_contract() -> None:
reader = FakeReader()
response = _client(reader).get(
"/api/v1/sector-radar/rankings",
params={
"trade_date": "2026-08-28",
"sector_type": "concept",
"view": "rank_change",
"rank_change_metric": "amount",
"rank_change_days": 5,
"side": "top",
"search": " 机器人 ",
"page": 2,
"page_size": 10,
},
)
assert response.status_code == 200
assert reader.last_query == RadarQuery(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
view=RadarView.RANK_CHANGE,
rank_change_metric=MetricKind.AMOUNT,
rank_change_days=5,
side=RankSide.TOP,
search="机器人",
page=2,
page_size=10,
)
payload = response.json()
assert payload["definition"]["metric_version"] == "zhixing_amount_net_bn_v1"
assert payload["definition"]["implementation_kind"] == "independent"
assert "知行独立实现" in payload["definition"]["disclaimer"]
assert payload["rows"][0]["rank_change"] == 3
assert payload["rows"][0]["unit"] == "CNY_100M"
def test_no_data_is_a_stable_200_response() -> None:
client = _client(FakeReader(no_data=True))
dates = client.get("/api/v1/sector-radar/dates")
rankings = client.get("/api/v1/sector-radar/rankings")
assert dates.status_code == 200
assert dates.json()["status"] == "no_data"
assert rankings.status_code == 200
assert rankings.json()["status"] == "no_data"
assert rankings.json()["rows"] == []
def test_http_contract_rejects_zero_rank_percentile() -> None:
payload = _client(FakeReader()).get("/api/v1/sector-radar/rankings").json()["rows"][0]
payload["rank_percentile"] = "0"
with pytest.raises(ValidationError):
RadarRankingRowResponse.model_validate(payload)
def test_invalid_query_values_return_422() -> None:
client = _client(FakeReader())
for params in (
{"rank_change_days": 0},
{"rank_change_days": 6},
{"page": 0},
{"page_size": 101},
{"sector_type": "region"},
{"view": "unknown"},
{"side": "unknown"},
):
assert client.get("/api/v1/sector-radar/rankings", params=params).status_code == 422
def test_repository_error_maps_to_redacted_503() -> None:
response = _client(FakeReader(fail=True)).get("/api/v1/sector-radar/rankings")
assert response.status_code == 503
assert response.json() == {
"detail": {
"code": "sector_radar_storage_unavailable",
"message": "sector radar storage is unavailable",
}
}
assert "private database detail" not in response.text
@@ -1,3 +1,4 @@
import threading
from datetime import date
import pytest
@@ -87,6 +88,109 @@ def test_rate_limit_cooldown_is_shared_by_following_requests() -> None:
assert waits == [60]
def test_request_start_interval_allows_overlapping_provider_calls() -> None:
current = [0.0]
state_lock = threading.Lock()
first_started = threading.Event()
release_first = threading.Event()
waits: list[float] = []
starts: list[tuple[str, float]] = []
errors: list[BaseException] = []
def clock() -> float:
with state_lock:
return current[0]
def wait(seconds: float) -> None:
with state_lock:
waits.append(seconds)
current[0] += seconds
coordinator = RequestCoordinator(
max_retries=0,
request_interval_seconds=0.2,
clock=clock,
wait_fn=wait,
sleep_fn=wait,
)
def first_request() -> object:
starts.append(("first", clock()))
first_started.set()
if not release_first.wait(timeout=2):
raise AssertionError("first provider call was not released")
return "first"
def run_first() -> None:
try:
coordinator.call("first", first_request)
except BaseException as exc: # pragma: no cover - surfaced by the assertion below
errors.append(exc)
first_thread = threading.Thread(target=run_first)
first_thread.start()
assert first_started.wait(timeout=2)
second = coordinator.call(
"second",
lambda: starts.append(("second", clock())) or "second",
)
assert second == "second"
assert first_thread.is_alive()
release_first.set()
first_thread.join(timeout=2)
assert not first_thread.is_alive()
assert errors == []
assert starts == [("first", 0.0), ("second", 0.2)]
assert waits == [0.2]
def test_request_start_interval_is_disabled_by_default() -> None:
waits: list[float] = []
starts: list[str] = []
coordinator = RequestCoordinator(
max_retries=0,
clock=lambda: 0.0,
wait_fn=waits.append,
)
coordinator.call("first", lambda: starts.append("first"))
coordinator.call("second", lambda: starts.append("second"))
assert starts == ["first", "second"]
assert waits == []
def test_request_start_interval_applies_to_retry_attempts() -> None:
current = [0.0]
waits: list[float] = []
starts: list[float] = []
def wait(seconds: float) -> None:
waits.append(seconds)
current[0] += seconds
coordinator = RequestCoordinator(
max_retries=1,
backoff_seconds=0,
request_interval_seconds=0.2,
clock=lambda: current[0],
wait_fn=wait,
sleep_fn=wait,
)
def request() -> object:
starts.append(current[0])
if len(starts) == 1:
raise RuntimeError("transient provider failure")
return "ok"
assert coordinator.call("daily", request) == "ok"
assert starts == [0.0, 0.2]
assert waits == [0.0, 0.2]
def test_pro_bar_qfq_calls_are_bound_to_the_shared_coordinator(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -0,0 +1,641 @@
import logging
from collections.abc import Sequence
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.application.build import (
BuildSectorRadar,
BuildSectorRadarCommand,
)
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
PublicationStatus,
RadarPublication,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.persistence import PublicationSourceGroup
from zhixing_server.modules.sector_radar.domain.source import (
CapabilityProbeResult,
DailyRow,
MoneyflowDcRow,
SectorIndexRow,
SectorMemberRow,
SourceContractError,
SourceResult,
SourceSnapshot,
StockBasicRow,
SuspendRow,
TradeCalendarRow,
build_source_snapshot,
)
from zhixing_server.modules.sector_radar.infrastructure.memory import (
InMemorySectorRadarRepository,
)
TARGET_DATE = date(2026, 8, 28)
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
class FakeRadarSource:
def __init__(
self,
*,
missing_moneyflow: bool = False,
missing_membership: bool = False,
net_scale: Decimal = Decimal(1),
) -> None:
self.missing_moneyflow = missing_moneyflow
self.missing_membership = missing_membership
self.net_scale = net_scale
self.fail_daily = False
self.calls: list[str] = []
self.moneyflow_candidate_codes: list[tuple[str, ...]] = []
def _result[T](
self, api_name: str, target: date | None, rows: tuple[T, ...]
) -> SourceResult[T]:
snapshot = build_source_snapshot(
api_name=api_name,
params={
"trade_date": target.isoformat() if target is not None else "all",
"fixture_fingerprint": repr(rows),
},
rows=tuple(self._raw_row(row) for row in rows),
target_trade_date=target,
partition_key="all" if api_name == "dc_member" else None,
observed_at=NOW,
)
return SourceResult((snapshot,), rows)
@staticmethod
def _raw_row(row: object) -> dict[str, object]:
if isinstance(row, TradeCalendarRow):
return {
"exchange": row.exchange,
"cal_date": row.cal_date,
"is_open": int(row.is_open),
"pretrade_date": row.pretrade_date,
}
if isinstance(row, SectorIndexRow):
return {
"trade_date": row.trade_date,
"ts_code": row.sector_code,
"name": row.name,
"level": row.level,
"pct_change": row.pct_change,
"leading_code": row.leading_code,
}
if isinstance(row, SectorMemberRow):
return {
"trade_date": row.trade_date,
"ts_code": row.sector_code,
"con_code": row.stock_code,
"name": row.stock_name,
}
if isinstance(row, StockBasicRow):
return {
"ts_code": row.ts_code,
"symbol": row.symbol,
"name": row.name,
"market": row.market,
"exchange": row.exchange,
"list_status": row.list_status,
"list_date": row.list_date,
"delist_date": row.delist_date,
}
if isinstance(row, SuspendRow):
return {
"ts_code": row.ts_code,
"trade_date": row.trade_date,
"suspend_timing": row.suspend_timing,
"suspend_type": row.suspend_type,
}
if isinstance(row, DailyRow):
return {
"ts_code": row.ts_code,
"trade_date": row.trade_date,
"close": row.close,
"pre_close": row.pre_close,
"pct_chg": row.pct_chg,
"vol": row.volume,
"amount": row.amount_thousand_yuan,
}
if isinstance(row, MoneyflowDcRow):
return {
"trade_date": row.trade_date,
"ts_code": row.ts_code,
"name": row.name,
"net_amount": row.net_amount_ten_thousand_yuan,
"net_amount_rate": row.net_amount_rate,
"pct_change": row.pct_change,
"close": row.close,
}
raise TypeError(f"unsupported fake source row: {type(row).__name__}")
def fetch_trade_calendar(self, start: date, end: date) -> SourceResult[TradeCalendarRow]:
self.calls.append("calendar")
rows = tuple(
TradeCalendarRow("SSE", start + timedelta(days=offset), True, None)
for offset in range((end - start).days + 1)
)
return self._result("trade_cal", end, rows)
def fetch_sector_indices(
self, trade_date: date, sector_type: SectorType
) -> SourceResult[SectorIndexRow]:
self.calls.append(f"{sector_type.value}_indices")
prefix = "BK0" if sector_type is SectorType.CONCEPT else "BK1"
row = SectorIndexRow(
trade_date,
sector_type,
f"{prefix}001.DC",
"示例概念" if sector_type is SectorType.CONCEPT else "示例行业",
"一级",
Decimal(1),
"000001.SZ",
)
return self._result(f"dc_index_{sector_type.value}", trade_date, (row,))
def fetch_sector_members(
self, trade_date: date, sector_codes: Sequence[str]
) -> SourceResult[SectorMemberRow]:
self.calls.append("members")
if self.missing_membership:
snapshots: list[SourceSnapshot] = []
member_rows: list[SectorMemberRow] = []
for sector_code in sector_codes:
sector_rows = (
()
if sector_code == sector_codes[-1]
else tuple(
SectorMemberRow(
trade_date,
sector_code,
f"00000{index}.SZ",
f"股票{index}",
)
for index in range(1, 6)
)
)
snapshots.append(
build_source_snapshot(
api_name="dc_member",
params={
"trade_date": trade_date.isoformat(),
"ts_code": sector_code,
},
rows=tuple(self._raw_row(row) for row in sector_rows),
target_trade_date=trade_date,
partition_key=sector_code,
observed_at=NOW,
)
)
member_rows.extend(sector_rows)
return SourceResult(tuple(snapshots), tuple(member_rows))
rows = tuple(
SectorMemberRow(trade_date, sector_code, f"00000{index}.SZ", f"股票{index}")
for sector_code in sector_codes
for index in range(1, 6)
)
return self._result("dc_member", trade_date, rows)
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
self.calls.append("stock_basics")
rows = tuple(
StockBasicRow(
f"00000{index}.SZ",
f"00000{index}",
f"股票{index}",
"主板",
"SZSE",
"L",
date(2020, 1, 1),
None,
)
for index in range(1, 6)
)
return self._result("stock_basic", None, rows)
def fetch_suspensions(self, trade_date: date) -> SourceResult[SuspendRow]:
self.calls.append("suspensions")
return self._result("suspend_d", trade_date, ())
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
self.calls.append("daily")
if self.fail_daily:
raise RuntimeError("private provider detail")
rows = tuple(
DailyRow(
f"00000{index}.SZ",
trade_date,
Decimal(10),
Decimal(10),
Decimal(0),
Decimal(100),
Decimal(1000),
)
for index in range(1, 6)
)
return self._result("daily", trade_date, rows)
def fetch_moneyflow_dc(
self,
trade_date: date,
candidate_codes: Sequence[str],
) -> SourceResult[MoneyflowDcRow]:
self.calls.append("moneyflow_dc")
self.moneyflow_candidate_codes.append(tuple(candidate_codes))
count = 4 if self.missing_moneyflow else 5
rows = tuple(
MoneyflowDcRow(
trade_date,
f"00000{index}.SZ",
f"股票{index}",
Decimal(index) * self.net_scale,
Decimal(0),
Decimal(0),
Decimal(10),
)
for index in range(1, count + 1)
)
return self._result("moneyflow_dc", trade_date, rows)
def probe(self, trade_date: date) -> CapabilityProbeResult:
return CapabilityProbeResult(NOW, ())
def test_successful_build_is_idempotent_and_failed_retry_preserves_last_good() -> None:
source = FakeRadarSource()
repository = InMemorySectorRadarRepository()
use_case = BuildSectorRadar(
source,
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
)
first = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
repeated = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
source.fail_daily = True
failed = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
assert first.status == "success"
assert first.exit_code == 0
assert first.outcomes[0].ranking_count == 6
assert repeated.status == "unchanged"
assert repeated.outcomes[0].publication_id == first.outcomes[0].publication_id
assert failed.status == "failed"
assert "private provider detail" not in str(failed.as_dict())
last_good = repository.get_last_good_publication()
assert last_good is not None
assert last_good.publication_id == first.outcomes[0].publication_id
assert any(item.status is PublicationStatus.FAILED for item in repository.publications.values())
def test_build_passes_stable_current_listing_member_intersection_to_moneyflow() -> None:
class FutureListingSource(FakeRadarSource):
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
result = super().fetch_stock_basics()
rows = result.rows[:-1] + (replace(result.rows[-1], list_date=date(2027, 1, 1)),)
return self._result("stock_basic", None, rows)
source = FutureListingSource()
summary = BuildSectorRadar(
source,
InMemorySectorRadarRepository(),
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
assert summary.status == "success"
assert source.moneyflow_candidate_codes == [
("000001.SZ", "000002.SZ", "000003.SZ", "000004.SZ")
]
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:
repository = InMemorySectorRadarRepository()
partial = BuildSectorRadar(
FakeRadarSource(missing_moneyflow=True),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
publication_count = len(repository.publications)
partial_repeated = BuildSectorRadar(
FakeRadarSource(missing_moneyflow=True),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
repository.lock_available = False
locked = BuildSectorRadar(
FakeRadarSource(),
repository,
today=TARGET_DATE,
now_fn=lambda: NOW,
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
assert partial.status == "partial"
assert partial.exit_code == 2
assert partial.outcomes[0].coverage == Decimal("0.8")
assert partial_repeated.status == "partial"
assert len(repository.publications) == publication_count
assert repository.get_last_good_publication() is None
assert {
record.source_group
for record in repository.load_publication_sources(partial.outcomes[0].publication_id or "")
if record.refresh_on_retry
} == {PublicationSourceGroup.MONEYFLOW_DC}
assert locked.status == "failed"
assert locked.exit_code == 1
assert locked.outcomes[0].status == "locked"
def test_unknown_membership_is_persisted_as_partial_and_retried_independently() -> None:
repository = InMemorySectorRadarRepository()
source = FakeRadarSource(missing_membership=True)
use_case = BuildSectorRadar(source, repository, today=TARGET_DATE, now_fn=lambda: NOW)
partial = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
partial_id = partial.outcomes[0].publication_id
assert partial_id is not None
publication = repository.get_publication(partial_id)
assert partial.status == "partial"
assert publication is not None
assert publication.error_summary == "membership_unknown"
assert any(item.status is MembershipStatus.UNKNOWN for item in repository.memberships.values())
assert {
record.source_group
for record in repository.load_publication_sources(partial_id)
if record.refresh_on_retry
} == {PublicationSourceGroup.MEMBERS}
source.missing_membership = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
assert retried.status == "success"
assert source.calls == ["members"]
def test_membership_retry_refreshes_moneyflow_when_replay_misses_new_candidates() -> None:
class ExpandingMembershipSource(FakeRadarSource):
def fetch_sector_members(
self,
trade_date: date,
sector_codes: Sequence[str],
) -> SourceResult[SectorMemberRow]:
self.calls.append("members")
rows: list[SectorMemberRow] = []
snapshots: list[SourceSnapshot] = []
for index, sector_code in enumerate(sector_codes, start=1):
sector_rows = (
()
if self.missing_membership and index == len(sector_codes)
else (
SectorMemberRow(
trade_date,
sector_code,
f"00000{index}.SZ",
f"股票{index}",
),
)
)
snapshots.append(
build_source_snapshot(
api_name="dc_member",
params={
"trade_date": trade_date.isoformat(),
"ts_code": sector_code,
},
rows=tuple(self._raw_row(row) for row in sector_rows),
target_trade_date=trade_date,
partition_key=sector_code,
observed_at=NOW,
)
)
rows.extend(sector_rows)
return SourceResult(tuple(snapshots), tuple(rows))
def fetch_moneyflow_dc(
self,
trade_date: date,
candidate_codes: Sequence[str],
) -> SourceResult[MoneyflowDcRow]:
self.calls.append("moneyflow_dc")
self.moneyflow_candidate_codes.append(tuple(candidate_codes))
rows = tuple(
MoneyflowDcRow(
trade_date,
code,
code,
Decimal(1),
Decimal(0),
Decimal(0),
Decimal(10),
)
for code in candidate_codes
)
return self._result("moneyflow_dc", trade_date, rows)
repository = InMemorySectorRadarRepository()
source = ExpandingMembershipSource(missing_membership=True)
use_case = BuildSectorRadar(source, repository, today=TARGET_DATE, now_fn=lambda: NOW)
partial = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
partial_id = partial.outcomes[0].publication_id
assert partial.status == "partial"
assert partial_id is not None
assert source.moneyflow_candidate_codes == [("000001.SZ",)]
source.missing_membership = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
assert retried.status == "success"
assert source.calls == ["members", "moneyflow_dc"]
assert source.moneyflow_candidate_codes[-1] == ("000001.SZ", "000002.SZ")
def test_range_builds_dates_in_order_and_retry_uses_old_target() -> None:
repository = InMemorySectorRadarRepository()
source = FakeRadarSource(missing_moneyflow=True)
use_case = BuildSectorRadar(source, repository, now_fn=lambda: NOW)
end = TARGET_DATE + timedelta(days=1)
summary = use_case.execute(BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end))
partial_id = summary.outcomes[0].publication_id
assert partial_id is not None
source.missing_moneyflow = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
assert [item.target_trade_date for item in summary.outcomes] == [TARGET_DATE, end]
assert all(item.status == "partial" for item in summary.outcomes)
assert retried.outcomes[0].target_trade_date == TARGET_DATE
assert retried.status == "success"
assert source.calls == ["moneyflow_dc"]
def test_failed_retry_reuses_every_completed_source_group() -> None:
repository = InMemorySectorRadarRepository()
source = FakeRadarSource()
source.fail_daily = True
use_case = BuildSectorRadar(source, repository, now_fn=lambda: NOW)
failed = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
failed_id = failed.outcomes[0].publication_id
assert failed_id is not None
source.fail_daily = False
source.calls.clear()
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=failed_id))
assert retried.status == "success"
assert source.calls == ["daily", "moneyflow_dc"]
old_groups = {record.source_group for record in repository.load_publication_sources(failed_id)}
assert len(old_groups) == 6
def test_date_lock_recovers_an_orphaned_running_publication() -> None:
repository = InMemorySectorRadarRepository()
stale = RadarPublication(
publication_id="stale-running",
target_trade_date=TARGET_DATE,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="pending",
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal(0),
started_at=NOW - timedelta(hours=1),
)
repository.create_publication(stale)
summary = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
recovered = repository.get_publication("stale-running")
assert summary.status == "success"
assert recovered is not None
assert recovered.status is PublicationStatus.FAILED
assert recovered.error_summary == "recovered_stale_running"
def test_default_target_excludes_today_before_closing_data_is_ready() -> None:
before_close = datetime(2026, 8, 28, 6, 0, tzinfo=UTC)
after_close = datetime(2026, 8, 28, 8, 0, tzinfo=UTC)
before = BuildSectorRadar(
FakeRadarSource(),
InMemorySectorRadarRepository(),
today=TARGET_DATE,
now_fn=lambda: before_close,
).execute()
after = BuildSectorRadar(
FakeRadarSource(),
InMemorySectorRadarRepository(),
today=TARGET_DATE,
now_fn=lambda: after_close,
).execute()
assert before.outcomes[0].target_trade_date == TARGET_DATE - timedelta(days=1)
assert after.outcomes[0].target_trade_date == TARGET_DATE
def test_tenth_trading_day_publishes_swing_and_five_rank_changes() -> None:
repository = InMemorySectorRadarRepository()
end = TARGET_DATE + timedelta(days=9)
summary = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end)
)
publication = repository.get_last_good_publication(end)
assert summary.status == "success"
assert publication is not None
current = tuple(
record.ranking
for record in repository.rankings.values()
if record.publication_id == publication.publication_id
)
swing = tuple(
ranking
for ranking in current
if ranking.observation.metric_version == "zhixing_swing_equal_3_10_v1"
)
assert len(swing) == 2
assert all(ranking.observation.value == Decimal("0.03") for ranking in swing)
assert all(
tuple(change.value for change in ranking.rank_changes) == (None, None, None, None, None)
for ranking in swing
)
amount = tuple(
ranking
for ranking in current
if ranking.observation.metric_version == "zhixing_amount_net_bn_v1"
)
assert all(
tuple(change.value for change in ranking.rank_changes) == (0, 0, 0, 0, 0)
for ranking in amount
)
def test_history_uses_latest_successful_input_revision_for_a_date() -> None:
repository = InMemorySectorRadarRepository()
clock = [NOW]
first_source = FakeRadarSource(net_scale=Decimal(1))
second_source = FakeRadarSource(net_scale=Decimal(2))
first = BuildSectorRadar(first_source, repository, now_fn=lambda: clock[0]).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
clock[0] = NOW + timedelta(minutes=5)
second = BuildSectorRadar(second_source, repository, now_fn=lambda: clock[0]).execute(
BuildSectorRadarCommand(trade_date=TARGET_DATE)
)
history = repository.load_daily_aggregate_history(
TARGET_DATE + timedelta(days=1), limit_dates=1
)
assert first.status == "success"
assert second.status == "success"
assert len(history) == 2
assert all(item.net_amount_yuan == Decimal(300_000) for item in history)
@@ -0,0 +1,129 @@
import json
from datetime import date
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.application.build import (
BuildDateOutcome,
BuildOutcomeStatus,
BuildSummary,
)
from zhixing_server.modules.sector_radar.presentation import cli
from zhixing_server.modules.sector_radar.presentation.cli import build_parser
def test_sector_radar_cli_parses_single_range_and_retry_modes() -> None:
parser = build_parser()
single = parser.parse_args(["--trade-date", "2026-08-28"])
date_range = parser.parse_args(["--start-date", "2026-08-18", "--end-date", "2026-08-28"])
retry = parser.parse_args(["--retry-publication-id", "publication-a"])
assert single.trade_date == date(2026, 8, 28)
assert date_range.start_date == date(2026, 8, 18)
assert date_range.end_date == date(2026, 8, 28)
assert retry.retry_publication_id == "publication-a"
class FakeSettings:
log_level = "INFO"
tushare_token = "secret-token"
database_url = "postgresql://unused"
sector_radar_max_retries = 3
sector_radar_retry_backoff_seconds = 1.0
sector_radar_request_interval_seconds = 0.2
sector_radar_advisory_lock_key = 7_380_522
sector_radar_coverage_threshold = Decimal("0.99")
class FakeRepository:
closed = False
def __init__(self, database_url: str, *, advisory_lock_key: int) -> None:
assert database_url == "postgresql://unused"
assert advisory_lock_key == 7_380_522
def close(self) -> None:
self.closed = True
@pytest.mark.parametrize(
("outcome_status", "coverage", "expected_code"),
(("success", Decimal(1), 0), ("partial", Decimal("0.8"), 2), ("failed", Decimal(0), 1)),
)
def test_cli_main_returns_summary_exit_code_and_json(
monkeypatch: pytest.MonkeyPatch,
capsys: pytest.CaptureFixture[str],
outcome_status: BuildOutcomeStatus,
coverage: Decimal,
expected_code: int,
) -> None:
summary = BuildSummary(
(
BuildDateOutcome(
date(2026, 8, 28),
outcome_status,
"publication-a",
coverage,
2,
6,
),
)
)
class FakeSourceFactory:
@staticmethod
def from_token(token: str, **kwargs: object) -> object:
assert token == "secret-token"
assert kwargs["max_retries"] == 3
assert kwargs["backoff_seconds"] == 1.0
assert kwargs["request_interval_seconds"] == 0.2
return object()
class FakeBuild:
def __init__(self, source: object, repository: object, **kwargs: object) -> None:
assert source is not None
assert repository is not None
assert kwargs
def execute(self, command: object) -> BuildSummary:
assert command is not None
return summary
monkeypatch.setattr(cli, "get_settings", FakeSettings)
monkeypatch.setattr(cli, "TushareSectorRadarAdapter", FakeSourceFactory)
monkeypatch.setattr(cli, "PostgresSectorRadarRepository", FakeRepository)
monkeypatch.setattr(cli, "BuildSectorRadar", FakeBuild)
exit_code = cli.main(["--trade-date", "2026-08-28"])
output = json.loads(capsys.readouterr().out)
assert exit_code == expected_code
assert output["status"] == summary.status
assert output["exit_code"] == expected_code
assert "secret-token" not in str(output)
def test_cli_initialization_failure_is_redacted(
monkeypatch: pytest.MonkeyPatch,
capsys: pytest.CaptureFixture[str],
) -> None:
class FailingSourceFactory:
@staticmethod
def from_token(token: str, **kwargs: object) -> object:
del token, kwargs
raise RuntimeError("private provider detail secret-token")
monkeypatch.setattr(cli, "get_settings", FakeSettings)
monkeypatch.setattr(cli, "TushareSectorRadarAdapter", FailingSourceFactory)
exit_code = cli.main(["--trade-date", "2026-08-28"])
captured = capsys.readouterr()
output = json.loads(captured.out)
assert exit_code == 1
assert output["status"] == "failed"
assert output["error_type"] == "RuntimeError"
assert "private provider detail" not in captured.out
assert "secret-token" not in captured.out
@@ -0,0 +1,146 @@
from datetime import UTC, date, datetime
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.domain.facts import aggregate_sector_snapshot
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
PublicationStatus,
RadarPublication,
SectorMembershipSnapshot,
SectorType,
StockDailyFact,
StockFactStatus,
)
TARGET_DATE = date(2026, 8, 28)
def test_point_in_time_aggregation_distinguishes_suspension_missing_and_zero() -> None:
snapshot = SectorMembershipSnapshot(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
member_codes=("000001.SZ", "000002.SZ", "000003.SZ", "000004.SZ"),
status=MembershipStatus.AVAILABLE,
source_version="dc-member-20260828-a",
)
facts = (
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("1000"),
net_amount_yuan=Decimal("100"),
),
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000002.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("2000"),
net_amount_yuan=Decimal("0"),
),
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000003.SZ",
status=StockFactStatus.SUSPENDED,
),
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000004.SZ",
status=StockFactStatus.MISSING_MONEYFLOW,
),
)
aggregate = aggregate_sector_snapshot(snapshot, facts)
assert aggregate.member_count == 4
assert aggregate.valid_sample_count == 2
assert aggregate.net_amount_yuan == Decimal("100")
assert aggregate.turnover_yuan == Decimal("3000")
assert aggregate.membership_coverage == Decimal("1")
assert aggregate.moneyflow_coverage == Decimal("2") / Decimal("3")
def test_unknown_membership_never_falls_back_to_available_stock_facts() -> None:
snapshot = SectorMembershipSnapshot(
trade_date=TARGET_DATE,
sector_type=SectorType.INDUSTRY,
sector_code="BK1001.DC",
sector_name="示例行业",
member_codes=(),
status=MembershipStatus.UNKNOWN,
source_version="dc-member-missing",
)
fact = StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("1000"),
net_amount_yuan=Decimal("100"),
)
aggregate = aggregate_sector_snapshot(snapshot, (fact,))
assert aggregate.member_count == 0
assert aggregate.net_amount_yuan is None
assert aggregate.turnover_yuan is None
assert aggregate.membership_coverage == Decimal("0")
def test_stock_fact_rejects_non_finite_values_and_invalid_status_payloads() -> None:
with pytest.raises(ValueError, match="finite"):
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.AVAILABLE,
turnover_yuan=Decimal("Infinity"),
net_amount_yuan=Decimal("1"),
)
with pytest.raises(ValueError, match="must not expose amounts"):
StockDailyFact(
trade_date=TARGET_DATE,
ts_code="000001.SZ",
status=StockFactStatus.SUSPENDED,
turnover_yuan=Decimal("0"),
)
def test_publication_requires_terminal_completion_and_replay_identity() -> None:
started_at = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
publication = RadarPublication(
publication_id="radar-20260828-a",
target_trade_date=TARGET_DATE,
status=PublicationStatus.SUCCESS,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-20260828-a",
metric_versions=(
"zhixing_amount_net_bn_v1",
"zhixing_ratio_turnover_v1",
"zhixing_swing_equal_3_10_v1",
),
input_hash="a" * 64,
coverage=Decimal("0.995"),
started_at=started_at,
finished_at=datetime(2026, 8, 28, 17, 35, tzinfo=UTC),
)
assert publication.status is PublicationStatus.SUCCESS
with pytest.raises(ValueError, match="finished_at"):
RadarPublication(
publication_id="radar-20260828-running",
target_trade_date=TARGET_DATE,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-20260828-a",
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal("0"),
started_at=started_at,
finished_at=started_at,
)
@@ -0,0 +1,112 @@
from datetime import date
from decimal import Decimal
from zhixing_server.modules.sector_radar.domain.metrics import (
AmountNetStrategy,
RatioTurnoverStrategy,
SwingEqualThreeToTenStrategy,
)
from zhixing_server.modules.sector_radar.domain.models import (
MetricQuality,
SectorDailyAggregate,
SectorType,
)
TARGET_DATE = date(2026, 8, 28)
def make_aggregate(
*,
net_amount_yuan: Decimal | None = Decimal("125000000"),
turnover_yuan: Decimal | None = Decimal("5000000000"),
) -> SectorDailyAggregate:
return SectorDailyAggregate(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
member_count=10,
valid_sample_count=10,
net_amount_yuan=net_amount_yuan,
turnover_yuan=turnover_yuan,
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
def test_amount_and_ratio_strategies_expose_independent_versioned_values() -> None:
aggregate = make_aggregate()
amount = AmountNetStrategy().evaluate((aggregate,), TARGET_DATE)
ratio = RatioTurnoverStrategy().evaluate((aggregate,), TARGET_DATE)
assert amount.value == Decimal("1.25")
assert amount.metric_version == "zhixing_amount_net_bn_v1"
assert amount.implementation_kind == "independent"
assert amount.unit == "CNY_100M"
assert amount.quality is MetricQuality.AVAILABLE
assert ratio.value == Decimal("0.025")
assert ratio.metric_version == "zhixing_ratio_turnover_v1"
assert ratio.implementation_kind == "independent"
assert ratio.unit == "ratio"
def test_missing_moneyflow_is_unavailable_but_zero_remains_a_real_value() -> None:
missing = AmountNetStrategy().evaluate((make_aggregate(net_amount_yuan=None),), TARGET_DATE)
zero = AmountNetStrategy().evaluate(
(make_aggregate(net_amount_yuan=Decimal("0")),), TARGET_DATE
)
assert missing.value is None
assert missing.quality is MetricQuality.UNAVAILABLE
assert zero.value == Decimal("0")
assert zero.quality is MetricQuality.AVAILABLE
def test_swing_strategy_uses_each_days_point_in_time_aggregate() -> None:
history = tuple(
SectorDailyAggregate(
trade_date=date(2026, 8, 18 + offset),
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
member_count=6 + offset,
valid_sample_count=6 + offset,
net_amount_yuan=Decimal(str(offset + 1)),
turnover_yuan=Decimal("100"),
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
for offset in range(10)
)
result = SwingEqualThreeToTenStrategy().evaluate(history, date(2026, 8, 27))
# The worked 3..10-day window ratios average to exactly 0.0725.
assert result.value == Decimal("0.0725")
assert result.metric_version == "zhixing_swing_equal_3_10_v1"
assert result.member_count == 15
def test_swing_strategy_carries_forward_limited_historical_sample_quality() -> None:
history = tuple(
SectorDailyAggregate(
trade_date=date(2026, 8, 18 + offset),
sector_type=SectorType.INDUSTRY,
sector_code="BK1001.DC",
sector_name="示例行业",
member_count=10,
valid_sample_count=4 if offset == 0 else 10,
net_amount_yuan=Decimal("10"),
turnover_yuan=Decimal("100"),
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
for offset in range(10)
)
result = SwingEqualThreeToTenStrategy().evaluate(history, date(2026, 8, 27))
assert result.value == Decimal("0.1")
assert result.quality is MetricQuality.AVAILABLE_LIMITED_SAMPLE
@@ -0,0 +1,215 @@
from datetime import UTC, date, datetime
from decimal import Decimal
from zhixing_server.modules.sector_radar.domain.models import (
MembershipStatus,
SectorType,
StockFactStatus,
)
from zhixing_server.modules.sector_radar.domain.normalize import (
normalize_memberships,
normalize_stock_facts,
)
from zhixing_server.modules.sector_radar.domain.source import (
DailyRow,
MoneyflowDcRow,
SectorIndexRow,
SectorMemberRow,
SourceResult,
StockBasicRow,
SuspendRow,
build_source_snapshot,
)
TARGET_DATE = date(2026, 8, 28)
OBSERVED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
def result[T](api_name: str, rows: tuple[T, ...]) -> SourceResult[T]:
snapshot = build_source_snapshot(
api_name=api_name,
params={"trade_date": "20260828"},
rows=(),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
)
return SourceResult((snapshot,), rows)
def basic(
ts_code: str,
*,
list_date: date = date(2020, 1, 1),
market: str | None = "主板",
) -> StockBasicRow:
return StockBasicRow(
ts_code=ts_code,
symbol=ts_code.split(".")[0],
name=ts_code,
market=market,
exchange="SZSE",
list_status="L",
list_date=list_date,
delist_date=None,
)
def test_missing_market_keeps_hs_a_stock_but_code_rules_still_exclude_bse_and_b_shares() -> None:
codes = ("000001.SZ", "920001.BJ", "200001.SZ", "900001.SH")
basics = tuple(basic(code, market=None) for code in codes)
daily_rows = tuple(daily(code, Decimal("1")) for code in codes)
moneyflow_rows = tuple(moneyflow(code, Decimal("1")) for code in codes)
facts = normalize_stock_facts(
target_trade_date=TARGET_DATE,
candidate_codes=codes,
stock_basics=result("stock_basic", basics),
suspensions=result("suspend_d", ()),
daily=result("daily", daily_rows),
moneyflow=result("moneyflow_dc", moneyflow_rows),
)
by_code = {fact.ts_code: fact for fact in facts}
assert by_code["000001.SZ"].status is StockFactStatus.AVAILABLE
assert by_code["920001.BJ"].status is StockFactStatus.LIFECYCLE_INVALID
assert by_code["200001.SZ"].status is StockFactStatus.LIFECYCLE_INVALID
assert by_code["900001.SH"].status is StockFactStatus.LIFECYCLE_INVALID
def daily(ts_code: str, amount: Decimal | None) -> DailyRow:
return DailyRow(
ts_code=ts_code,
trade_date=TARGET_DATE,
close=Decimal("10"),
pre_close=Decimal("10"),
pct_chg=Decimal(0),
volume=Decimal(0),
amount_thousand_yuan=amount,
)
def moneyflow(ts_code: str, amount: Decimal | None) -> MoneyflowDcRow:
return MoneyflowDcRow(
trade_date=TARGET_DATE,
ts_code=ts_code,
name=ts_code,
net_amount_ten_thousand_yuan=amount,
net_amount_rate=Decimal(0),
pct_change=Decimal(0),
close=Decimal("10"),
)
def test_membership_normalization_persists_an_explicit_unknown_sector() -> None:
indices = (
SectorIndexRow(
TARGET_DATE,
SectorType.CONCEPT,
"BK0001.DC",
"机器人",
"一级",
Decimal(1),
None,
),
SectorIndexRow(
TARGET_DATE,
SectorType.CONCEPT,
"BK0002.DC",
"低空经济",
"一级",
Decimal(1),
None,
),
)
member = SectorMemberRow(
TARGET_DATE,
"BK0001.DC",
"000001.SZ",
"平安银行",
)
all_snapshot = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20260828"},
rows=(
{
"trade_date": "20260828",
"ts_code": member.sector_code,
"con_code": member.stock_code,
"name": member.stock_name,
},
),
target_trade_date=TARGET_DATE,
partition_key="all",
observed_at=OBSERVED_AT,
)
empty_partition = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20260828", "ts_code": "BK0002.DC"},
rows=(),
target_trade_date=TARGET_DATE,
partition_key="BK0002.DC",
observed_at=OBSERVED_AT,
)
records = normalize_memberships(
indices,
SourceResult((all_snapshot, empty_partition), (member,)),
)
assert records[0].status is MembershipStatus.AVAILABLE
assert records[0].stock_code == "000001.SZ"
assert records[1].status is MembershipStatus.UNKNOWN
assert records[1].stock_code is None
assert records[1].membership_key == "__membership_unknown__"
def test_stock_fact_normalization_preserves_all_missing_and_zero_states() -> None:
codes = tuple(f"00000{index}.SZ" for index in range(1, 9))
basics = tuple(
basic(code, list_date=date(2027, 1, 1) if code == codes[7] else date(2020, 1, 1))
for code in codes
)
daily_rows = (
daily(codes[0], Decimal("1")),
daily(codes[3], None),
daily(codes[4], Decimal("1")),
daily(codes[5], Decimal("1")),
daily(codes[6], Decimal("0")),
daily(codes[7], Decimal("1")),
)
moneyflow_rows = (
moneyflow(codes[0], Decimal("0")),
moneyflow(codes[3], Decimal("1")),
moneyflow(codes[5], None),
moneyflow(codes[6], Decimal("0")),
moneyflow(codes[7], Decimal("1")),
)
suspensions = (
SuspendRow(
ts_code=codes[1],
trade_date=TARGET_DATE,
suspend_timing=None,
suspend_type="停牌",
),
)
facts = normalize_stock_facts(
target_trade_date=TARGET_DATE,
candidate_codes=codes,
stock_basics=result("stock_basic", basics),
suspensions=result("suspend_d", suspensions),
daily=result("daily", daily_rows),
moneyflow=result("moneyflow_dc", moneyflow_rows),
)
by_code = {fact.ts_code: fact for fact in facts}
assert by_code[codes[0]].status is StockFactStatus.AVAILABLE
assert by_code[codes[0]].turnover_yuan == Decimal("1000")
assert by_code[codes[0]].net_amount_yuan == Decimal("0")
assert by_code[codes[1]].status is StockFactStatus.SUSPENDED
assert by_code[codes[2]].status is StockFactStatus.MISSING_DAILY
assert by_code[codes[3]].status is StockFactStatus.NULL_DAILY_AMOUNT
assert by_code[codes[4]].status is StockFactStatus.MISSING_MONEYFLOW
assert by_code[codes[5]].status is StockFactStatus.NULL_MONEYFLOW
assert by_code[codes[6]].status is StockFactStatus.LOW_LIQUIDITY
assert by_code[codes[7]].status is StockFactStatus.LIFECYCLE_INVALID
@@ -0,0 +1,170 @@
from collections.abc import Generator
from contextlib import contextmanager
from datetime import UTC, date, datetime
from decimal import Decimal
from typing import Any, cast
from psycopg_pool import ConnectionPool
from zhixing_server.modules.sector_radar.domain.models import PublicationStatus
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
PostgresSectorRadarRepository,
)
TARGET_DATE = date(2026, 8, 28)
class FakeResult:
def __init__(
self,
row: tuple[object, ...] | None = None,
rows: tuple[tuple[object, ...], ...] | None = None,
) -> None:
self.row = row
self.rows = rows or (() if row is None else (row,))
def fetchone(self) -> tuple[object, ...] | None:
return self.row
def fetchall(self) -> tuple[tuple[object, ...], ...]:
return self.rows
class FakeConnection:
def __init__(self) -> None:
self.statements: list[tuple[str, tuple[object, ...]]] = []
def execute(
self,
query: str,
parameters: tuple[object, ...] = (),
) -> FakeResult:
self.statements.append((query, parameters))
if "FROM sector_radar_ranking" in query:
return FakeResult(
rows=(
(
TARGET_DATE,
"concept",
"BK0001.DC",
"机器人",
"amount",
"zhixing_amount_net_bn_v1",
"independent",
"CNY_100M",
Decimal("12.5"),
"available",
20,
19,
Decimal(1),
Decimal("0.95"),
1,
Decimal(100),
{"1": 3, "2": None},
),
)
)
if "FROM sector_radar_publication" in query:
return FakeResult(
(
"publication-a",
TARGET_DATE,
"success",
"tushare-pro-v1",
"eastmoney-dc-v1",
["zhixing_amount_net_bn_v1"],
"a" * 64,
Decimal("1"),
datetime(2026, 8, 28, 17, 30, tzinfo=UTC),
datetime(2026, 8, 28, 17, 35, tzinfo=UTC),
None,
)
)
if "pg_try_advisory_lock" in query:
return FakeResult((True,))
return FakeResult((True,))
class FakePool:
def __init__(self, connection: FakeConnection) -> None:
self._connection = connection
self._opened = False
def open(self, *, wait: bool) -> None:
assert wait
self._opened = True
def close(self) -> None:
self._opened = False
@contextmanager
def connection(self) -> Generator[FakeConnection]:
yield self._connection
def make_repository(connection: FakeConnection) -> PostgresSectorRadarRepository:
pool = cast(ConnectionPool[Any], cast(object, FakePool(connection)))
return PostgresSectorRadarRepository("postgresql://unused", pool=pool)
def test_last_good_query_strictly_filters_success_and_date() -> None:
connection = FakeConnection()
publication = make_repository(connection).get_last_good_publication(TARGET_DATE)
assert publication is not None
assert publication.status is PublicationStatus.SUCCESS
query, parameters = connection.statements[0]
assert "status = 'success'" in query
assert "partial" not in query
assert "target_trade_date <= %s" in query
assert parameters == (TARGET_DATE,)
def test_advisory_lock_uses_target_date_and_releases_same_key() -> None:
connection = FakeConnection()
with make_repository(connection).advisory_lock(TARGET_DATE) as acquired:
assert acquired
assert len(connection.statements) == 2
assert "pg_try_advisory_lock" in connection.statements[0][0]
assert "2026-08-28" in str(connection.statements[0][1][0])
assert "pg_advisory_unlock" in connection.statements[1][0]
assert connection.statements[0][1] == connection.statements[1][1]
def test_exact_success_and_latest_attempt_queries_use_distinct_semantics() -> None:
connection = FakeConnection()
repository = make_repository(connection)
exact = repository.get_successful_publication(TARGET_DATE)
latest = repository.get_latest_publication()
assert exact is not None
assert latest is not None
exact_query, exact_parameters = connection.statements[0]
latest_query, latest_parameters = connection.statements[1]
assert "status = 'success' AND target_trade_date = %s" in exact_query
assert exact_parameters == (TARGET_DATE,)
assert "status = 'success'" not in latest_query
assert "started_at DESC" in latest_query
assert latest_parameters == ()
def test_load_rankings_reconstructs_values_and_rank_changes() -> None:
connection = FakeConnection()
rankings = make_repository(connection).load_rankings("publication-a")
assert len(rankings) == 1
ranking = rankings[0]
assert ranking.observation.metric_version == "zhixing_amount_net_bn_v1"
assert ranking.observation.value == Decimal("12.5")
assert ranking.rank_position == 1
assert ranking.rank_change(1) == 3
assert ranking.rank_change(2) is None
query, parameters = connection.statements[0]
assert "WHERE publication_id = %s" in query
assert "rank_position NULLS LAST" in query
assert parameters == ("publication-a",)
@@ -0,0 +1,147 @@
from datetime import date
from decimal import Decimal
from zhixing_server.modules.sector_radar.domain.models import (
MetricKind,
MetricObservation,
MetricQuality,
MetricUnit,
RankSide,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.ranking import (
rank_metric_observations,
select_percentile_side,
select_rank_change_side,
with_rank_changes,
)
TARGET_DATE = date(2026, 8, 28)
def make_observation(
sector_code: str,
sector_type: SectorType,
value: str | None,
*,
trade_date: date = TARGET_DATE,
) -> MetricObservation:
metric_value = Decimal(value) if value is not None else None
return MetricObservation(
trade_date=trade_date,
sector_type=sector_type,
sector_code=sector_code,
sector_name=sector_code,
metric_kind=MetricKind.AMOUNT,
metric_version="zhixing_amount_net_bn_v1",
implementation_kind="independent",
unit=MetricUnit.CNY_100M,
value=metric_value,
quality=(
MetricQuality.AVAILABLE if metric_value is not None else MetricQuality.UNAVAILABLE
),
member_count=10,
valid_sample_count=10 if metric_value is not None else 0,
membership_coverage=Decimal("1"),
moneyflow_coverage=Decimal("1"),
)
def test_ranking_separates_types_and_uses_code_as_stable_tie_breaker() -> None:
observations = (
make_observation("BK2002.DC", SectorType.INDUSTRY, "20"),
make_observation("BK1002.DC", SectorType.CONCEPT, "30"),
make_observation("BK2001.DC", SectorType.INDUSTRY, "20"),
make_observation("BK1001.DC", SectorType.CONCEPT, "10"),
)
ranked = rank_metric_observations(tuple(reversed(observations)))
by_code = {row.observation.sector_code: row for row in ranked}
assert by_code["BK1002.DC"].rank_position == 1
assert by_code["BK1002.DC"].rank_percentile == Decimal("100")
assert by_code["BK1001.DC"].rank_position == 2
assert by_code["BK1001.DC"].rank_percentile == Decimal("50")
assert by_code["BK2001.DC"].rank_position == 1
assert by_code["BK2002.DC"].rank_position == 2
def test_ranking_handles_empty_and_single_element_pools() -> None:
assert rank_metric_observations(()) == ()
[single] = rank_metric_observations((make_observation("BK0001.DC", SectorType.CONCEPT, "0"),))
assert single.rank_position == 1
assert single.rank_percentile == Decimal("100")
def test_percentile_sides_use_confirmed_inclusive_thresholds() -> None:
ranked = rank_metric_observations(
make_observation(f"BK{position:04d}.DC", SectorType.CONCEPT, str(11 - position))
for position in range(1, 11)
)
top = select_percentile_side(ranked, RankSide.TOP)
bottom = select_percentile_side(ranked, RankSide.BOTTOM)
assert [row.observation.sector_code for row in top] == ["BK0001.DC", "BK0002.DC"]
assert [row.observation.sector_code for row in bottom] == ["BK0010.DC"]
def test_rank_change_is_past_rank_minus_current_and_preserves_missing_history() -> None:
current = rank_metric_observations(
(
make_observation("BK0001.DC", SectorType.CONCEPT, "30"),
make_observation("BK0002.DC", SectorType.CONCEPT, "20"),
)
)
previous = rank_metric_observations(
(
make_observation(
"BK0001.DC",
SectorType.CONCEPT,
"10",
trade_date=date(2026, 8, 27),
),
make_observation(
"BK0002.DC",
SectorType.CONCEPT,
"40",
trade_date=date(2026, 8, 27),
),
)
)
changed = with_rank_changes(current, {1: previous, 5: ()})
by_code = {row.observation.sector_code: row for row in changed}
assert by_code["BK0001.DC"].rank_change(1) == 1
assert by_code["BK0002.DC"].rank_change(1) == -1
assert by_code["BK0001.DC"].rank_change(5) is None
def test_rank_change_sides_take_ceiling_ten_percent_per_pool() -> None:
current = rank_metric_observations(
make_observation(f"BK{position:04d}.DC", SectorType.CONCEPT, str(12 - position))
for position in range(1, 12)
)
previous = rank_metric_observations(
make_observation(
f"BK{position:04d}.DC",
SectorType.CONCEPT,
str(position),
trade_date=date(2026, 8, 27),
)
for position in range(1, 12)
)
changed = with_rank_changes(current, {1: previous})
top = select_rank_change_side(changed, days=1, side=RankSide.TOP)
bottom = select_rank_change_side(changed, days=1, side=RankSide.BOTTOM)
assert [row.observation.sector_code for row in top] == ["BK0001.DC", "BK0002.DC"]
assert [row.observation.sector_code for row in bottom] == [
"BK0011.DC",
"BK0010.DC",
]
@@ -0,0 +1,189 @@
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
from zhixing_server.modules.sector_radar.application.read import (
RadarQuery,
RadarView,
ReadSectorRadar,
)
from zhixing_server.modules.sector_radar.domain.metrics import AmountNetStrategy
from zhixing_server.modules.sector_radar.domain.models import (
MetricKind,
MetricObservation,
MetricQuality,
MetricUnit,
PublicationStatus,
RadarPublication,
RankChange,
RankedMetric,
RankSide,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.persistence import RankingRecord
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
from zhixing_server.modules.sector_radar.infrastructure.memory import (
InMemorySectorRadarRepository,
)
TARGET_DATE = date(2026, 8, 28)
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
def _running(publication_id: str, trade_date: date) -> RadarPublication:
return RadarPublication(
publication_id=publication_id,
target_trade_date=trade_date,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-v1",
metric_versions=(AmountNetStrategy.metric_version,),
input_hash=None,
coverage=Decimal(0),
started_at=NOW,
)
def _finish(
publication: RadarPublication,
status: PublicationStatus,
) -> RadarPublication:
return replace(
publication,
status=status,
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
coverage=Decimal(1) if status is PublicationStatus.SUCCESS else Decimal("0.8"),
finished_at=publication.started_at + timedelta(minutes=5),
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
)
def _amount_rankings() -> tuple[RankedMetric, ...]:
observations = tuple(
MetricObservation(
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code=f"BK{index:04d}.DC",
sector_name=f"概念{index}",
metric_kind=MetricKind.AMOUNT,
metric_version=AmountNetStrategy.metric_version,
implementation_kind="independent",
unit=MetricUnit.CNY_100M,
value=Decimal(11 - index),
quality=MetricQuality.AVAILABLE,
member_count=5,
valid_sample_count=5,
membership_coverage=Decimal(1),
moneyflow_coverage=Decimal(1),
)
for index in range(1, 11)
)
rankings = rank_metric_observations(observations)
return tuple(
replace(
row,
rank_changes=tuple(
RankChange(
days=days,
value=(
None
if row.observation.sector_code == "BK0005.DC" and days == 5
else (row.rank_position or 0) - 5
),
)
for days in range(1, 6)
),
)
for row in rankings
)
def _published_repository() -> InMemorySectorRadarRepository:
repository = InMemorySectorRadarRepository()
publication = _running("publication-success", TARGET_DATE)
repository.create_publication(publication)
repository.finish_publication(_finish(publication, PublicationStatus.SUCCESS))
repository.save_rankings(
RankingRecord(publication.publication_id, ranking) for ranking in _amount_rankings()
)
return repository
def test_no_successful_publication_returns_stable_no_data() -> None:
reader = ReadSectorRadar(InMemorySectorRadarRepository())
dates = reader.list_dates()
rankings = reader.query(RadarQuery())
assert dates.status == "no_data"
assert dates.available_dates == ()
assert rankings.status == "no_data"
assert rankings.publication is None
assert rankings.total == 0
assert rankings.definition.metric_version == AmountNetStrategy.metric_version
def test_explicit_date_never_falls_back_to_an_earlier_last_good() -> None:
reader = ReadSectorRadar(_published_repository())
missing = reader.query(RadarQuery(trade_date=TARGET_DATE + timedelta(days=1)))
assert missing.status == "no_data"
assert missing.publication is None
def test_percentile_side_is_selected_before_search_and_pagination() -> None:
reader = ReadSectorRadar(_published_repository())
top = reader.query(RadarQuery(side=RankSide.TOP, page_size=1))
second_page = reader.query(RadarQuery(side=RankSide.TOP, page=2, page_size=1))
searched = reader.query(RadarQuery(side=RankSide.TOP, search="概念2"))
bottom = reader.query(RadarQuery(side=RankSide.BOTTOM))
assert top.total == 2
assert top.rows[0].observation.sector_code == "BK0001.DC"
assert second_page.rows[0].observation.sector_code == "BK0002.DC"
assert searched.total == 1
assert searched.rows[0].observation.sector_name == "概念2"
assert bottom.total == 1
assert bottom.rows[0].observation.sector_code == "BK0010.DC"
def test_rank_change_uses_selected_metric_days_and_pool_sides() -> None:
reader = ReadSectorRadar(_published_repository())
query = RadarQuery(
view=RadarView.RANK_CHANGE,
rank_change_metric=MetricKind.AMOUNT,
rank_change_days=5,
)
top = reader.query(replace(query, side=RankSide.TOP))
bottom = reader.query(replace(query, side=RankSide.BOTTOM))
all_rows = reader.query(query)
assert top.total == 1
assert top.rows[0].rank_change(5) == 5
assert bottom.total == 1
assert bottom.rows[0].rank_change(5) == -4
assert all_rows.total == 10
assert all_rows.rows[-1].observation.sector_code == "BK0005.DC"
assert all_rows.rows[-1].rank_change(5) is None
def test_latest_partial_attempt_is_visible_but_does_not_replace_last_good() -> None:
repository = _published_repository()
partial = replace(
_running("publication-partial", TARGET_DATE + timedelta(days=1)),
started_at=NOW + timedelta(days=1),
)
repository.create_publication(partial)
repository.finish_publication(_finish(partial, PublicationStatus.PARTIAL))
index = ReadSectorRadar(repository).list_dates()
assert index.status == "success"
assert index.current_attempt is not None
assert index.current_attempt.status is PublicationStatus.PARTIAL
assert index.last_good is not None
assert index.last_good.publication_id == "publication-success"
assert index.available_dates == (TARGET_DATE,)
@@ -0,0 +1,121 @@
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.domain.models import (
PublicationStatus,
RadarPublication,
SectorType,
)
from zhixing_server.modules.sector_radar.domain.persistence import MembershipRecord
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
from zhixing_server.modules.sector_radar.infrastructure.memory import (
InMemorySectorRadarRepository,
)
TARGET_DATE = date(2026, 8, 28)
STARTED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
def make_running(publication_id: str, target_trade_date: date = TARGET_DATE) -> RadarPublication:
return RadarPublication(
publication_id=publication_id,
target_trade_date=target_trade_date,
status=PublicationStatus.RUNNING,
source_version="tushare-pro-v1",
universe_version="eastmoney-dc-v1",
metric_versions=("zhixing_amount_net_bn_v1",),
input_hash=None,
coverage=Decimal(0),
started_at=STARTED_AT,
)
def finish(
publication: RadarPublication,
status: PublicationStatus,
*,
offset_minutes: int = 5,
) -> RadarPublication:
return replace(
publication,
status=status,
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
coverage=Decimal("1") if status is PublicationStatus.SUCCESS else Decimal("0.8"),
finished_at=publication.started_at + timedelta(minutes=offset_minutes),
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
)
def test_source_and_membership_revisions_are_idempotent_but_not_overwritable() -> None:
repository = InMemorySectorRadarRepository()
snapshot = build_source_snapshot(
api_name="dc_member",
params={"trade_date": "20260828"},
rows=(
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "平安银行",
},
),
target_trade_date=TARGET_DATE,
observed_at=STARTED_AT,
)
member = MembershipRecord(
source_snapshot_id=snapshot.snapshot_id,
trade_date=TARGET_DATE,
sector_type=SectorType.CONCEPT,
sector_code="BK0001.DC",
sector_name="示例概念",
stock_code="000001.SZ",
stock_name="平安银行",
)
assert repository.save_source_snapshots((snapshot,)).inserted == 1
assert repository.save_source_snapshots((snapshot,)).unchanged == 1
assert repository.save_memberships((member,)).inserted == 1
assert repository.save_memberships((member,)).unchanged == 1
with pytest.raises(ValueError, match="cannot change content"):
repository.save_memberships((replace(member, stock_name="已改变"),))
def test_partial_and_failed_revisions_never_replace_last_good() -> None:
repository = InMemorySectorRadarRepository()
successful = make_running("success-a")
partial = make_running("partial-b")
failed = make_running("failed-c", TARGET_DATE + timedelta(days=1))
repository.create_publication(successful)
repository.finish_publication(finish(successful, PublicationStatus.SUCCESS))
repository.create_publication(partial)
repository.finish_publication(finish(partial, PublicationStatus.PARTIAL, offset_minutes=6))
repository.create_publication(failed)
repository.finish_publication(finish(failed, PublicationStatus.FAILED, offset_minutes=7))
last_good = repository.get_last_good_publication()
assert last_good is not None
assert last_good.publication_id == "success-a"
assert repository.list_successful_dates() == (TARGET_DATE,)
def test_publication_identity_allows_sequential_same_date_revisions() -> None:
repository = InMemorySectorRadarRepository()
first = make_running("revision-a")
second = make_running("revision-b")
assert repository.create_publication(first).inserted == 1
with pytest.raises(ValueError, match="already has a running"):
repository.create_publication(second)
with pytest.raises(ValueError, match="terminal"):
repository.finish_publication(first)
repository.finish_publication(finish(first, PublicationStatus.FAILED))
assert repository.create_publication(second).inserted == 1
with pytest.raises(ValueError, match="running status"):
repository.finish_publication(finish(first, PublicationStatus.SUCCESS))
@@ -0,0 +1,669 @@
import logging
import threading
from collections.abc import Mapping
from datetime import UTC, date, datetime
from decimal import Decimal
import pytest
from zhixing_server.modules.sector_radar.domain.models import SectorType
from zhixing_server.modules.sector_radar.domain.source import (
CapabilityStatus,
SourceContractError,
build_source_snapshot,
)
from zhixing_server.modules.sector_radar.infrastructure import tushare as source_module
from zhixing_server.modules.sector_radar.infrastructure.tushare import (
TushareSectorRadarAdapter,
)
TARGET_DATE = date(2026, 8, 28)
OBSERVED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
class QueryClient:
def __init__(self, responses: Mapping[tuple[str, str], object]) -> None:
self.responses = dict(responses)
self.calls: list[tuple[str, dict[str, object]]] = []
self._lock = threading.Lock()
def query(self, api_name: str, **kwargs: object) -> object:
partition = str(kwargs.get("ts_code") or kwargs.get("list_status") or "")
with self._lock:
self.calls.append((api_name, kwargs))
response = self.responses.get((api_name, partition), ())
if isinstance(response, BaseException):
raise response
return response
def make_adapter(client: object) -> TushareSectorRadarAdapter:
return TushareSectorRadarAdapter(
client,
max_retries=0,
request_interval_seconds=0,
sleep_fn=lambda _: None,
now_fn=lambda: OBSERVED_AT,
)
def moneyflow_record(
ts_code: str,
*,
trade_date: str = "20260828",
) -> dict[str, object]:
return {
"trade_date": trade_date,
"ts_code": ts_code,
"name": ts_code,
"net_amount": "1",
"net_amount_rate": "0.1",
"pct_change": "1",
"close": "10",
}
def test_daily_and_moneyflow_keep_source_units_and_distinguish_missing_from_zero() -> None:
client = QueryClient(
{
(
"daily",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"close": "10",
"pre_close": "9.5",
"pct_chg": "1.5",
"vol": "100",
"amount": "12.5",
},
{
"ts_code": "000002.SZ",
"trade_date": "20260828",
"close": "20",
"pre_close": "20",
"pct_chg": "0",
"vol": "0",
"amount": float("nan"),
},
),
(
"moneyflow_dc",
"",
): (
{
"trade_date": "20260828",
"ts_code": "000001.SZ",
"name": "平安银行",
"net_amount": "2.5",
"net_amount_rate": "0.2",
"pct_change": "1.5",
"close": "10",
},
{
"trade_date": "20260828",
"ts_code": "000002.SZ",
"name": "示例股票",
"net_amount": "0",
"net_amount_rate": "0",
"pct_change": "0",
"close": "20",
},
),
}
)
adapter = make_adapter(client)
daily = adapter.fetch_daily(TARGET_DATE)
moneyflow = adapter.fetch_moneyflow_dc(TARGET_DATE, ("000001.SZ", "000002.SZ"))
assert daily.rows[0].amount_thousand_yuan == Decimal("12.5")
assert daily.rows[0].turnover_yuan == Decimal("12500.0")
assert daily.rows[1].amount_thousand_yuan is None
assert moneyflow.rows[0].net_amount_ten_thousand_yuan == Decimal("2.5")
assert moneyflow.rows[0].net_amount_yuan == Decimal("25000.0")
assert moneyflow.rows[1].net_amount_yuan == Decimal("0")
assert client.calls[0][1]["fields"] == ",".join(source_module.FIELDS["daily"])
def test_moneyflow_accepts_a_full_initial_snapshot_at_the_provider_limit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", 2)
client = QueryClient(
{
("moneyflow_dc", ""): (
moneyflow_record("000001.SZ"),
moneyflow_record("000002.SZ"),
)
}
)
result = make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
("000001.SZ", "000002.SZ"),
)
assert result.snapshots[0].limit_reached is True
assert [row.ts_code for row in result.rows] == ["000001.SZ", "000002.SZ"]
assert len(client.calls) == 1
@pytest.mark.parametrize(
("initial_rows", "message"),
(
((moneyflow_record("000001.SZ", trade_date="20260827"),), "trade_date"),
(
(moneyflow_record("000001.SZ"), moneyflow_record("000001.SZ")),
"duplicate business keys",
),
),
)
def test_moneyflow_initial_contract_errors_fail_closed(
initial_rows: tuple[dict[str, object], ...],
message: str,
) -> None:
client = QueryClient({("moneyflow_dc", ""): initial_rows})
with pytest.raises(SourceContractError, match=message):
make_adapter(client).fetch_moneyflow_dc(TARGET_DATE, ())
def test_moneyflow_refills_only_missing_codes_in_stable_snapshot_order(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", 3)
third_finished = threading.Event()
completion_order: list[str] = []
completion_lock = threading.Lock()
class ReverseCompletionClient(QueryClient):
def query(self, api_name: str, **kwargs: object) -> object:
response = super().query(api_name, **kwargs)
ts_code = str(kwargs.get("ts_code") or "")
if ts_code == "000004.SZ":
if not third_finished.wait(timeout=2):
raise AssertionError("second moneyflow worker did not start")
elif ts_code == "000005.SZ":
third_finished.set()
if ts_code:
with completion_lock:
completion_order.append(ts_code)
return response
client = ReverseCompletionClient(
{
("moneyflow_dc", ""): tuple(
moneyflow_record(f"00000{index}.SZ") for index in range(1, 4)
),
("moneyflow_dc", "000004.SZ"): (moneyflow_record("000004.SZ"),),
("moneyflow_dc", "000005.SZ"): (moneyflow_record("000005.SZ"),),
}
)
result = make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
tuple(f"00000{index}.SZ" for index in range(1, 6)),
)
assert completion_order == ["000005.SZ", "000004.SZ"]
assert [snapshot.partition_key for snapshot in result.snapshots] == [
"all",
"000004.SZ",
"000005.SZ",
]
assert [row.ts_code for row in result.rows] == [
"000001.SZ",
"000002.SZ",
"000003.SZ",
"000004.SZ",
"000005.SZ",
]
assert len(client.calls) == 3
def test_moneyflow_empty_and_exhausted_refills_remain_real_gaps(
caplog: pytest.LogCaptureFixture,
) -> None:
client = QueryClient(
{
("moneyflow_dc", ""): (moneyflow_record("000001.SZ"),),
("moneyflow_dc", "000002.SZ"): (),
("moneyflow_dc", "000003.SZ"): RuntimeError("private provider payload"),
}
)
caplog.set_level(
logging.WARNING,
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
)
result = make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
("000001.SZ", "000002.SZ", "000003.SZ"),
)
assert [row.ts_code for row in result.rows] == ["000001.SZ"]
assert [snapshot.partition_key for snapshot in result.snapshots] == ["all"]
messages = "\n".join(record.getMessage() for record in caplog.records)
assert "partition_empty partition_key=000002.SZ" in messages
assert "partition_failed partition_key=000003.SZ" in messages
assert "private provider payload" not in messages
@pytest.mark.parametrize(
("partition_rows", "row_limit", "message"),
(
((moneyflow_record("000002.SZ", trade_date="20260827"),), 6_000, "trade_date"),
((moneyflow_record("000099.SZ"),), 6_000, "different ts_code"),
(
(moneyflow_record("000002.SZ"), moneyflow_record("000002.SZ")),
6_000,
"duplicate business keys",
),
(
(moneyflow_record("000002.SZ"), moneyflow_record("000002.SZ")),
2,
"provider row limit",
),
),
)
def test_moneyflow_partition_contract_errors_fail_closed(
monkeypatch: pytest.MonkeyPatch,
partition_rows: tuple[dict[str, object], ...],
row_limit: int,
message: str,
) -> None:
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", row_limit)
client = QueryClient(
{
("moneyflow_dc", ""): (moneyflow_record("000001.SZ"),),
("moneyflow_dc", "000002.SZ"): partition_rows,
}
)
with pytest.raises(SourceContractError, match=message):
make_adapter(client).fetch_moneyflow_dc(
TARGET_DATE,
("000001.SZ", "000002.SZ"),
)
def test_non_finite_source_values_are_rejected() -> None:
client = QueryClient(
{
(
"daily",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"close": "Infinity",
"pre_close": "9.5",
"pct_chg": "1.5",
"vol": "100",
"amount": "12.5",
},
)
}
)
with pytest.raises(SourceContractError, match="finite"):
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(
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": "A",
},
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000002.SZ",
"name": "B",
},
),
(
"dc_member",
"BK0001.DC",
): (
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "A",
},
),
(
"dc_member",
"BK0002.DC",
): (
{
"trade_date": "20260828",
"ts_code": "BK0002.DC",
"con_code": "600000.SH",
"name": "C",
},
),
}
)
result = make_adapter(client).fetch_sector_members(
TARGET_DATE,
("BK0001.DC", "BK0002.DC"),
)
assert [row.stock_code for row in result.rows] == ["000001.SZ", "600000.SH"]
assert [snapshot.partition_key for snapshot in result.snapshots] == [
"all",
"BK0001.DC",
"BK0002.DC",
]
def test_dc_member_preserves_an_explicit_empty_partition() -> None:
client = QueryClient(
{
(
"dc_member",
"",
): (
{
"trade_date": "20260828",
"ts_code": "BK0001.DC",
"con_code": "000001.SZ",
"name": "A",
},
),
("dc_member", "BK0002.DC"): (),
}
)
result = make_adapter(client).fetch_sector_members(
TARGET_DATE,
("BK0001.DC", "BK0002.DC"),
)
assert [row.sector_code for row in result.rows] == ["BK0001.DC"]
assert [snapshot.partition_key for snapshot in result.snapshots] == [
"all",
"BK0002.DC",
]
assert result.snapshots[1].row_count == 0
def test_stock_basic_requests_only_current_listings() -> None:
client = QueryClient(
{
(
"stock_basic",
"L",
): (
{
"ts_code": "000001.SZ",
"symbol": "000001",
"name": "L",
"market": "主板",
"exchange": "SZSE",
"list_status": "L",
"list_date": "20200101",
"delist_date": None,
},
)
}
)
result = make_adapter(client).fetch_stock_basics()
assert {row.list_status for row in result.rows} == {"L"}
assert [snapshot.partition_key for snapshot in result.snapshots] == ["L"]
assert [call[1]["list_status"] for call in client.calls] == ["L"]
def test_stock_basic_rejects_a_non_listed_row_from_the_l_partition() -> None:
client = QueryClient(
{
("stock_basic", "L"): (
{
"ts_code": "000001.SZ",
"symbol": "000001",
"name": "unexpected",
"market": "主板",
"exchange": "SZSE",
"list_status": "D",
"list_date": "20200101",
"delist_date": "20260828",
},
)
}
)
with pytest.raises(SourceContractError, match="unexpected list_status"):
make_adapter(client).fetch_stock_basics()
def test_suspend_timing_may_be_missing_while_suspend_type_remains_required() -> None:
client = QueryClient(
{
(
"suspend_d",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"suspend_timing": None,
"suspend_type": "S",
},
)
}
)
result = make_adapter(client).fetch_suspensions(TARGET_DATE)
assert result.rows[0].suspend_timing is None
assert result.rows[0].suspend_type == "S"
missing_type = QueryClient(
{
(
"suspend_d",
"",
): (
{
"ts_code": "000001.SZ",
"trade_date": "20260828",
"suspend_timing": None,
"suspend_type": None,
},
)
}
)
with pytest.raises(SourceContractError, match="suspend_type must be a non-empty string"):
make_adapter(missing_type).fetch_suspensions(TARGET_DATE)
def test_source_snapshot_hash_is_order_stable_and_excludes_token_params() -> None:
first = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828", "token": "secret"},
rows=({"ts_code": "2"}, {"ts_code": "1"}),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
)
second = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "1"}, {"ts_code": "2"}),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
)
assert first.snapshot_id == second.snapshot_id
assert "secret" not in repr(first)
def test_source_snapshot_identity_includes_schema_and_limit_metadata() -> None:
first = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "000001.SZ"},),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
returned_fields=("ts_code",),
row_limit=1,
)
changed_schema = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "000001.SZ"},),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
returned_fields=("name", "ts_code"),
row_limit=1,
)
changed_limit = build_source_snapshot(
api_name="daily",
params={"trade_date": "20260828"},
rows=({"ts_code": "000001.SZ"},),
target_trade_date=TARGET_DATE,
observed_at=OBSERVED_AT,
returned_fields=("ts_code",),
row_limit=2,
)
assert first.content_sha256 != changed_schema.content_sha256
assert first.snapshot_id != changed_schema.snapshot_id
assert first.content_sha256 != changed_limit.content_sha256
assert first.snapshot_id != changed_limit.snapshot_id
def test_capability_probe_classifies_errors_without_exposing_provider_text() -> None:
client = QueryClient({("daily", ""): RuntimeError("权限不足 private-detail")})
probe = make_adapter(client).probe(TARGET_DATE)
by_name = {result.api_name: result for result in probe.interfaces}
assert by_name["daily"].status is CapabilityStatus.FORBIDDEN
assert "private-detail" not in repr(probe)
assert len(probe.interfaces) == 7
def test_sector_index_uses_independent_concept_and_industry_params() -> None:
client = QueryClient(
{
(
"dc_index",
"",
): (
{
"ts_code": "BK0001.DC",
"trade_date": "20260828",
"name": "示例",
"idx_type": "概念板块",
"level": "一级",
"pct_change": "1",
"leading_code": "000001.SZ",
},
)
}
)
result = make_adapter(client).fetch_sector_indices(TARGET_DATE, SectorType.CONCEPT)
assert result.rows[0].sector_type is SectorType.CONCEPT
assert client.calls[0][1]["idx_type"] == "概念板块"