diff --git a/.trellis/tasks/08-28-sector-capital-radar/implement.md b/.trellis/tasks/08-28-sector-capital-radar/implement.md index e031c97..75d3f38 100644 --- a/.trellis/tasks/08-28-sector-capital-radar/implement.md +++ b/.trellis/tasks/08-28-sector-capital-radar/implement.md @@ -18,13 +18,15 @@ ## 2. Tushare 输入与持久化 -- [ ] 把已有 RequestCoordinator 提升到 shared 基础设施,保持 market-data 适配器及测试行为不变。 -- [ ] 定义 `SectorRadarSource` 与 Tushare adapter,显式请求七类接口及 fields;token 仅由 `Settings` 注入。 -- [ ] 实现服务端错误分类、有限重试、行数上限检测、`dc_member` 分片和账号 capability probe;输出不得包含 token。 -- [ ] 新增 Alembic 表、约束、索引和 downgrade,保存原始 JSONB/hash、成员快照、股票事实、publication 与 ranking。 -- [ ] 实现 PostgreSQL staging/COPY、幂等重跑、同日多修订、advisory lock 和 last-good 查询。 -- [ ] 为 repository fake、Tushare fake、迁移和 PostgreSQL 集成补测试;仅在 `ZHIXING_TEST_DATABASE_URL` 存在时执行数据库集成测试。 -- [ ] 若运行环境存在 `ZHIXING_TUSHARE_TOKEN`,执行只读 capability probe 并记录接口成功、字段和行数,不打印原始凭据;否则明确记录 live 验证未执行。 +- [x] 把已有 RequestCoordinator 提升到 shared 基础设施,保持 market-data 适配器及测试行为不变。 +- [x] 定义 `SectorRadarSource` 与 Tushare adapter,显式请求七类接口及 fields;token 仅由 `Settings` 注入。 +- [x] 实现服务端错误分类、有限重试、行数上限检测、`dc_member` 分片和账号 capability probe;输出不得包含 token。 +- [x] 新增 Alembic 表、约束、索引和 downgrade,保存原始 JSONB/hash、成员快照、股票事实、publication 与 ranking。 +- [x] 实现 PostgreSQL staging/COPY、幂等重跑、同日多修订、advisory lock 和 last-good 查询。 +- [x] 为 repository fake、Tushare fake、迁移和 PostgreSQL 集成补测试;仅在 `ZHIXING_TEST_DATABASE_URL` 存在时执行数据库集成测试。 +- [x] 若运行环境存在 `ZHIXING_TUSHARE_TOKEN`,执行只读 capability probe 并记录接口成功、字段和行数,不打印原始凭据;否则明确记录 live 验证未执行。 + +阶段结果(2026-08-29):七接口 source 契约、共享限流协调、源快照 hash、point-in-time 规范化、五张 PostgreSQL 表、COPY staging、同日修订与严格 `success` last-good 已落地。完整后端门禁为 109 passed、3 skipped;当前环境未设置 `ZHIXING_TUSHARE_TOKEN` 和 `ZHIXING_TEST_DATABASE_URL`,因此 live capability probe 与三项 PostgreSQL 集成测试未执行,未将其误报为通过。 ## 3. 构建 Job diff --git a/zhixing-server/migrations/versions/0004_sector_radar.py b/zhixing-server/migrations/versions/0004_sector_radar.py new file mode 100644 index 0000000..7d0c3f3 --- /dev/null +++ b/zhixing-server/migrations/versions/0004_sector_radar.py @@ -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") diff --git a/zhixing-server/src/zhixing_server/bootstrap/config.py b/zhixing-server/src/zhixing_server/bootstrap/config.py index afa7572..5374a4b 100644 --- a/zhixing-server/src/zhixing_server/bootstrap/config.py +++ b/zhixing-server/src/zhixing_server/bootstrap/config.py @@ -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) diff --git a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py index 9ccf80e..d835904 100644 --- a/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py +++ b/zhixing-server/src/zhixing_server/modules/market_data/infrastructure/tushare.py @@ -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: diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/facts.py b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/facts.py index 0f4104b..0b1f300 100644 --- a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/facts.py +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/facts.py @@ -67,7 +67,14 @@ def aggregate_sector_snapshot( expected_count = 0 for member_code in snapshot.member_codes: fact = facts_by_code.get(member_code) - if fact is None or fact.status is StockFactStatus.MISSING: + 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 diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/models.py b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/models.py index 7ebf4c1..2944604 100644 --- a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/models.py +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/models.py @@ -29,6 +29,10 @@ class StockFactStatus(StrEnum): 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" diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/normalize.py b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/normalize.py new file mode 100644 index 0000000..4e3abe7 --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/normalize.py @@ -0,0 +1,204 @@ +"""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, + ) + records: list[MembershipRecord] = [] + 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") + snapshot_id = partition_ids.get(member.sector_code, all_snapshot_id) + if snapshot_id is None: + raise SourceContractError("membership row has no source snapshot") + 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.stock_code))) + + +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: All explicit Tushare listing-status partitions. + 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_lifecycle_candidate(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_lifecycle_candidate(stock: StockBasicRow, target: date) -> bool: + if not stock.ts_code.endswith((".SH", ".SZ")): + return False + if stock.symbol.startswith(("200", "900")): + return False + if "北交" in stock.market or "B股" in stock.market.upper(): + return False + if stock.list_date is None or stock.list_date > target: + return False + return stock.delist_date is None or target <= stock.delist_date + + +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 diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/persistence.py b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/persistence.py new file mode 100644 index 0000000..2b9b2d4 --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/persistence.py @@ -0,0 +1,144 @@ +"""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 +from decimal import Decimal +from typing import Protocol + +from .models import ( + MembershipStatus, + RadarPublication, + RankedMetric, + 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 tied to its raw source revision.""" + + source_snapshot_id: str + trade_date: date + sector_type: SectorType + sector_code: str + sector_name: str + stock_code: str + stock_name: str + status: MembershipStatus = MembershipStatus.AVAILABLE + + def __post_init__(self) -> None: + """Validate revision identity and member fields.""" + + _validate_digest(self.source_snapshot_id, "source_snapshot_id") + if self.status is not MembershipStatus.AVAILABLE: + raise ValueError("persisted member rows require available membership") + if any( + not value.strip() + for value in (self.sector_code, self.sector_name, self.stock_code, self.stock_name) + ): + raise ValueError("membership identity fields must not be empty") + + +@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 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_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts: ... + + def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ... + + def create_publication(self, publication: RadarPublication) -> WriteCounts: ... + + def finish_publication(self, publication: RadarPublication) -> None: ... + + def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts: ... + + def get_publication(self, publication_id: str) -> RadarPublication | None: ... + + def get_last_good_publication( + self, target_trade_date: date | None = None + ) -> RadarPublication | None: ... + + def list_successful_dates(self) -> Sequence[date]: ... diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/ports.py b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/ports.py new file mode 100644 index 0000000..06f8ce6 --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/ports.py @@ -0,0 +1,53 @@ +"""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) -> 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]: ... diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/domain/source.py b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/source.py new file mode 100644 index 0000000..a8b8c59 --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/domain/source.py @@ -0,0 +1,462 @@ +"""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.""" + + +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") + ) + canonical_rows = sorted( + json.dumps(row, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + for row in normalized_rows + ) + content_sha256 = hashlib.sha256("\n".join(canonical_rows).encode()).hexdigest() + identity = json.dumps( + { + "api_name": api_name, + "params": normalized_params, + "partition_key": partition_key, + "content_sha256": content_sha256, + }, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + snapshot_id = hashlib.sha256(identity.encode()).hexdigest() + 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) + 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=row_limit is not None and row_count >= row_limit, + ) + + +@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 + 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=_required_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 + 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=_required_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 + ) diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/__init__.py b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/__init__.py new file mode 100644 index 0000000..7b7497f --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/__init__.py @@ -0,0 +1 @@ +"""Infrastructure adapters for the sector radar bounded context.""" diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/memory.py b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/memory.py new file mode 100644 index 0000000..e735c9f --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/memory.py @@ -0,0 +1,179 @@ +"""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 datetime import date + +from ..domain.models import PublicationStatus, RadarPublication +from ..domain.persistence import ( + MembershipRecord, + 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.memberships: dict[tuple[str, str, str], MembershipRecord] = {} + self.stock_facts: dict[tuple[str, str], StockFactRecord] = {} + self.publications: dict[str, RadarPublication] = {} + 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_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.stock_code), + ) + + 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 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 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 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), + 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, + ) + ) + + @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) diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/postgres.py b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/postgres.py new file mode 100644 index 0000000..79b30dc --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/postgres.py @@ -0,0 +1,475 @@ +"""Psycopg repository for replayable sector radar inputs and publications.""" + +from __future__ import annotations + +import threading +from collections.abc import Callable, Generator, Iterable, Sequence +from contextlib import contextmanager +from datetime import date +from decimal import Decimal +from typing import Any + +from psycopg.types.json import Jsonb +from psycopg_pool import ConnectionPool + +from ..domain.models import PublicationStatus, RadarPublication +from ..domain.persistence import ( + MembershipRecord, + RankingRecord, + StockFactRecord, + WriteCounts, +) +from ..domain.source import SourceSnapshot + + +class SectorRadarRepositoryError(RuntimeError): + """PostgreSQL could not complete a radar repository operation safely.""" + + +class PostgresSectorRadarRepository: + """Persist immutable input revisions with COPY staging and strict last-good reads.""" + + def __init__( + self, + database_url: str, + *, + advisory_lock_key: int = 7_380_522, + max_connections: int = 4, + pool: ConnectionPool[Any] | None = None, + ) -> None: + if max_connections < 1: + raise ValueError("max_connections must be at least 1") + self.database_url = database_url + self.advisory_lock_key = advisory_lock_key + self.pool = pool or ConnectionPool( + conninfo=database_url, + min_size=1, + max_size=max_connections, + open=False, + ) + self._owns_pool = pool is None + self._pool_open = False + self._pool_state_lock = threading.Lock() + + def open(self) -> None: + """Open the owned or injected pool exactly once.""" + + with self._pool_state_lock: + if self._pool_open: + return + if bool(getattr(self.pool, "_opened", False)): + self._pool_open = True + return + self.pool.open(wait=True) + self._pool_open = True + + def close(self) -> None: + """Close only a pool owned by this repository.""" + + with self._pool_state_lock: + if self._owns_pool and (self._pool_open or bool(getattr(self.pool, "_opened", False))): + self.pool.close() + self._pool_open = False + + def __enter__(self) -> PostgresSectorRadarRepository: + self.open() + return self + + def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None: + self.close() + + @contextmanager + def advisory_lock(self, target_trade_date: date) -> Generator[bool]: + """Hold a session advisory lock for one target date and build lifetime.""" + + lock_name = f"sector-radar:{self.advisory_lock_key}:{target_trade_date.isoformat()}" + with self._connection() as connection: + row = connection.execute( + "SELECT pg_try_advisory_lock(hashtext(%s))", + (lock_name,), + ).fetchone() + acquired = bool(row[0]) if row is not None else False + if not acquired: + yield False + return + try: + yield True + finally: + connection.execute( + "SELECT pg_advisory_unlock(hashtext(%s))", + (lock_name,), + ) + + def save_source_snapshots(self, snapshots: Iterable[SourceSnapshot]) -> WriteCounts: + """Insert content-addressed raw snapshots without replacing old payloads.""" + + items = tuple(snapshots) + self._require_unique(items, lambda item: item.snapshot_id) + if not items: + return WriteCounts(0, 0) + with self._connection() as connection, connection.transaction(): + existing = { + str(row[0]) + for row in connection.execute( + "SELECT id FROM sector_radar_source_snapshot WHERE id = ANY(%s)", + ([item.snapshot_id for item in items],), + ).fetchall() + } + connection.cursor().executemany( + """ + INSERT INTO sector_radar_source_snapshot ( + id, api_name, normalized_params, target_trade_date, partition_key, + observed_at, payload, row_count, returned_fields, content_sha256, + row_limit, limit_reached + ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + ON CONFLICT (id) DO NOTHING + """, + tuple( + ( + item.snapshot_id, + item.api_name, + Jsonb(dict(item.normalized_params)), + item.target_trade_date, + item.partition_key, + item.observed_at, + Jsonb(list(item.rows)), + item.row_count, + Jsonb(list(item.returned_fields)), + item.content_sha256, + item.row_limit, + item.limit_reached, + ) + for item in items + ), + ) + unchanged = len(existing) + return WriteCounts(inserted=len(items) - unchanged, unchanged=unchanged) + + def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts: + """COPY point-in-time members into an immutable revision key.""" + + items = tuple(records) + self._require_unique( + items, + lambda item: (item.source_snapshot_id, item.sector_code, item.stock_code), + ) + rows = tuple( + ( + item.source_snapshot_id, + item.trade_date, + item.sector_type.value, + item.sector_code, + item.sector_name, + item.stock_code, + item.stock_name, + item.status.value, + ) + for item in items + ) + return self._copy_immutable( + "sector_radar_membership", + ( + "source_snapshot_id", + "trade_date", + "sector_type", + "sector_code", + "sector_name", + "stock_code", + "stock_name", + "membership_status", + ), + ("source_snapshot_id", "sector_code", "stock_code"), + rows, + ) + + def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: + """COPY normalized stock facts while preserving contributing source ids.""" + + items = tuple(records) + self._require_unique(items, lambda item: (item.fact_revision, item.ts_code)) + rows = tuple( + ( + item.fact_revision, + item.trade_date, + item.ts_code, + Jsonb(list(item.source_snapshot_ids)), + item.status.value, + item.turnover_yuan, + item.net_amount_yuan, + ) + for item in items + ) + return self._copy_immutable( + "sector_radar_stock_fact", + ( + "fact_revision", + "trade_date", + "ts_code", + "source_snapshot_ids", + "status", + "turnover_yuan", + "net_amount_yuan", + ), + ("fact_revision", "ts_code"), + rows, + ) + + def create_publication(self, publication: RadarPublication) -> WriteCounts: + """Insert a new running publication identity idempotently.""" + + if publication.status is not PublicationStatus.RUNNING: + raise ValueError("new publications must start in running status") + with self._connection() as connection, connection.transaction(): + existing = connection.execute( + "SELECT status, target_trade_date FROM sector_radar_publication WHERE id = %s", + (publication.publication_id,), + ).fetchone() + if existing is not None: + if ( + str(existing[0]) != publication.status.value + or existing[1] != publication.target_trade_date + ): + raise SectorRadarRepositoryError("publication identity has conflicting content") + return WriteCounts(0, 1) + connection.execute( + """ + INSERT INTO sector_radar_publication ( + id, target_trade_date, status, source_version, universe_version, + metric_versions, input_hash, coverage, started_at, finished_at, error_summary + ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """, + self._publication_values(publication), + ) + return WriteCounts(1, 0) + + def finish_publication(self, publication: RadarPublication) -> None: + """Transition one running publication to a terminal audit state.""" + + if publication.status is PublicationStatus.RUNNING: + raise ValueError("finished publication must use a terminal status") + with self._connection() as connection, connection.transaction(): + result = connection.execute( + """ + UPDATE sector_radar_publication + SET status = %s, source_version = %s, universe_version = %s, + metric_versions = %s, input_hash = %s, coverage = %s, + finished_at = %s, error_summary = %s + WHERE id = %s AND target_trade_date = %s AND status = 'running' + """, + ( + publication.status.value, + publication.source_version, + publication.universe_version, + Jsonb(list(publication.metric_versions)), + publication.input_hash, + publication.coverage, + publication.finished_at, + self._safe_error(publication.error_summary), + publication.publication_id, + publication.target_trade_date, + ), + ) + if result.rowcount != 1: + raise SectorRadarRepositoryError("publication is not in running status") + + def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts: + """COPY versioned ranking projections under one publication revision.""" + + items = tuple(records) + self._require_unique( + items, + lambda item: ( + item.publication_id, + item.ranking.observation.sector_type, + item.ranking.observation.sector_code, + item.ranking.observation.metric_version, + ), + ) + rows: list[tuple[object, ...]] = [] + for item in items: + ranking = item.ranking + observation = ranking.observation + rows.append( + ( + item.publication_id, + observation.trade_date, + observation.sector_type.value, + observation.sector_code, + observation.sector_name, + observation.metric_kind.value, + observation.metric_version, + observation.implementation_kind, + observation.unit.value, + observation.value, + observation.quality.value, + observation.member_count, + observation.valid_sample_count, + observation.membership_coverage, + observation.moneyflow_coverage, + ranking.rank_position, + ranking.rank_percentile, + Jsonb({str(change.days): change.value for change in ranking.rank_changes}), + ) + ) + return self._copy_immutable( + "sector_radar_ranking", + ( + "publication_id", + "trade_date", + "sector_type", + "sector_code", + "sector_name", + "metric_kind", + "metric_version", + "implementation_kind", + "unit", + "metric_value", + "quality", + "member_count", + "valid_sample_count", + "membership_coverage", + "moneyflow_coverage", + "rank_position", + "rank_percentile", + "rank_changes", + ), + ("publication_id", "sector_type", "sector_code", "metric_version"), + tuple(rows), + ) + + def get_publication(self, publication_id: str) -> RadarPublication | None: + """Read one publication by immutable identity.""" + + with self._connection() as connection: + row = connection.execute( + self._publication_select() + " WHERE id = %s", + (publication_id,), + ).fetchone() + return None if row is None else self._publication_from_row(row) + + def get_last_good_publication( + self, target_trade_date: date | None = None + ) -> RadarPublication | None: + """Read only status=success, optionally bounded by a requested date.""" + + where = " WHERE status = 'success'" + parameters: tuple[object, ...] = () + if target_trade_date is not None: + where += " AND target_trade_date <= %s" + parameters = (target_trade_date,) + query = ( + self._publication_select() + + where + + " ORDER BY target_trade_date DESC, finished_at DESC, created_at DESC, id DESC LIMIT 1" + ) + with self._connection() as connection: + row = connection.execute(query, parameters).fetchone() + return None if row is None else self._publication_from_row(row) + + def list_successful_dates(self) -> Sequence[date]: + """List distinct successful target dates newest first.""" + + with self._connection() as connection: + rows = connection.execute( + """ + SELECT DISTINCT target_trade_date + FROM sector_radar_publication + WHERE status = 'success' + ORDER BY target_trade_date DESC + """ + ).fetchall() + return tuple(row[0] for row in rows) + + def _copy_immutable( + self, + table: str, + columns: tuple[str, ...], + conflict_columns: tuple[str, ...], + rows: tuple[tuple[object, ...], ...], + ) -> WriteCounts: + if not rows: + return WriteCounts(0, 0) + if table not in { + "sector_radar_membership", + "sector_radar_stock_fact", + "sector_radar_ranking", + }: + raise ValueError("unsupported radar staging table") + stage = f"{table}_stage" + column_sql = ", ".join(columns) + conflict_sql = ", ".join(conflict_columns) + with self._connection() as connection, connection.transaction(): + cursor = connection.cursor() + cursor.execute( + f"CREATE TEMP TABLE {stage} (LIKE {table} INCLUDING DEFAULTS) ON COMMIT DROP" + ) + with cursor.copy(f"COPY {stage} ({column_sql}) FROM STDIN") as copy: + for row in rows: + copy.write_row(row) + inserted_rows = cursor.execute( + f"INSERT INTO {table} ({column_sql}) SELECT {column_sql} FROM {stage} " + f"ON CONFLICT ({conflict_sql}) DO NOTHING RETURNING 1" + ).fetchall() + inserted = len(inserted_rows) + return WriteCounts(inserted=inserted, unchanged=len(rows) - inserted) + + @contextmanager + def _connection(self) -> Generator[Any]: + try: + self.open() + with self.pool.connection() as connection: + yield connection + except SectorRadarRepositoryError: + raise + except Exception as exc: # noqa: BLE001 - database errors are redacted at this boundary + raise SectorRadarRepositoryError("sector radar database operation failed") from exc + + @staticmethod + def _require_unique[T, K](items: Sequence[T], key: Callable[[T], K]) -> None: + keys = [key(item) for item in items] + if len(keys) != len(set(keys)): + raise ValueError("one write batch must not contain duplicate business keys") + + @staticmethod + def _publication_values(publication: RadarPublication) -> tuple[object, ...]: + return ( + publication.publication_id, + publication.target_trade_date, + publication.status.value, + publication.source_version, + publication.universe_version, + Jsonb(list(publication.metric_versions)), + publication.input_hash, + publication.coverage, + publication.started_at, + publication.finished_at, + PostgresSectorRadarRepository._safe_error(publication.error_summary), + ) + + @staticmethod + def _publication_select() -> str: + return ( + "SELECT id, target_trade_date, status, source_version, universe_version, " + "metric_versions, input_hash, coverage, started_at, finished_at, error_summary " + "FROM sector_radar_publication" + ) + + @staticmethod + def _publication_from_row(row: tuple[Any, ...]) -> RadarPublication: + return RadarPublication( + publication_id=str(row[0]), + target_trade_date=row[1], + status=PublicationStatus(str(row[2])), + source_version=str(row[3]), + universe_version=str(row[4]), + metric_versions=tuple(str(value) for value in row[5]), + input_hash=None if row[6] is None else str(row[6]), + coverage=Decimal(str(row[7])), + started_at=row[8], + finished_at=row[9], + error_summary=None if row[10] is None else str(row[10]), + ) + + @staticmethod + def _safe_error(message: str | None) -> str | None: + return None if message is None else " ".join(message.split())[:500] diff --git a/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/tushare.py b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/tushare.py new file mode 100644 index 0000000..4e48a83 --- /dev/null +++ b/zhixing-server/src/zhixing_server/modules/sector_radar/infrastructure/tushare.py @@ -0,0 +1,485 @@ +"""Tushare adapter for replayable sector radar source facts.""" + +from __future__ import annotations + +import time +from collections.abc import Callable, Iterable, Mapping, Sequence +from datetime import UTC, date, datetime +from typing import TypeVar, cast + +from zhixing_server.shared.request_coordinator import ( + DEFAULT_RATE_LIMIT_COOLDOWNS, + RequestCoordinator, +) + +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") + +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: "行业板块", +} + + +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._sleep_fn = sleep_fn + self._request_interval_seconds = max(0.0, request_interval_seconds) + self._now_fn = now_fn + self._coordinator = request_coordinator or RequestCoordinator( + max_retries=max_retries, + backoff_seconds=backoff_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", + ) + initial_rows = tuple(SectorMemberRow.from_mapping(row) for row in initial.rows) + self._require_target_date(initial_rows, trade_date, "dc_member") + 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, + ) + 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") + if not partition_rows: + raise SourceContractError("dc_member cannot prove complete membership for a sector") + snapshots.append(snapshot) + merged_rows.extend(partition_rows) + + self._require_unique( + merged_rows, + key=lambda row: (row.trade_date, row.sector_code, row.stock_code), + api_name="dc_member", + ) + final_codes = {row.sector_code for row in merged_rows} + if set(expected_codes) - final_codes: + raise SourceContractError("dc_member response is missing expected sectors") + 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 every documented listing status instead of relying on the L default.""" + + snapshots: list[SourceSnapshot] = [] + rows: list[StockBasicRow] = [] + for status in ("L", "D", "P", "G", "UN"): + snapshot = self._fetch_snapshot( + "stock_basic", + {"exchange": "", "list_status": status}, + target_trade_date=None, + partition_key=status, + ) + snapshots.append(snapshot) + parsed = tuple(StockBasicRow.from_mapping(row) for row in snapshot.rows) + if any(row.list_status != status for row in parsed): + raise SourceContractError("stock_basic returned an unexpected list_status") + rows.extend(parsed) + self._require_unique(rows, key=lambda row: row.ts_code, api_name="stock_basic") + return SourceResult(tuple(snapshots), 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) -> SourceResult[MoneyflowDcRow]: + """Fetch a full-market DC moneyflow snapshot in its documented source unit.""" + + snapshot = self._fetch_snapshot( + "moneyflow_dc", + {"trade_date": trade_date.strftime("%Y%m%d")}, + target_trade_date=trade_date, + ) + 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.ts_code, api_name="moneyflow_dc") + return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code))) + + 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) + self._sleep_fn(self._request_interval_seconds) + columns = getattr(result, "columns", None) + returned_fields = ( + tuple(str(column) for column in cast(Iterable[object], columns)) + if isinstance(columns, Iterable) and not isinstance(columns, (str, bytes)) + else None + ) + 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: + raise SourceContractError(f"{api_name} response is missing requested fields") + return snapshot + + @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}, + ) diff --git a/zhixing-server/src/zhixing_server/shared/request_coordinator.py b/zhixing-server/src/zhixing_server/shared/request_coordinator.py new file mode 100644 index 0000000..5491d23 --- /dev/null +++ b/zhixing-server/src/zhixing_server/shared/request_coordinator.py @@ -0,0 +1,170 @@ +"""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 and shared rate-limit cooling for one provider client. + + Normal requests are not serialized. Only a classified provider limit creates + a shared cooldown. Injectable time functions keep long cooldowns deterministic + in tests without coupling the coordinator to any business bounded context. + """ + + 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( + "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_cooldown(self, method_name: str) -> None: + while True: + with self._condition: + delay = self._cooldown_until - self.clock() + if delay <= 0: + return + logger.info( + "provider_rate_limit_wait method=%s wait_seconds=%.1f", + 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 diff --git a/zhixing-server/tests/integration/test_market_data_migration.py b/zhixing-server/tests/integration/test_market_data_migration.py index 12b93fa..cc6954a 100644 --- a/zhixing-server/tests/integration/test_market_data_migration.py +++ b/zhixing-server/tests/integration/test_market_data_migration.py @@ -38,6 +38,11 @@ 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", } <= tables finally: engine.dispose() diff --git a/zhixing-server/tests/integration/test_sector_radar_repository.py b/zhixing-server/tests/integration/test_sector_radar_repository.py new file mode 100644 index 0000000..3d0e163 --- /dev/null +++ b/zhixing-server/tests/integration/test_sector_radar_repository.py @@ -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,), + ) diff --git a/zhixing-server/tests/unit/sector_radar/test_facts.py b/zhixing-server/tests/unit/sector_radar/test_facts.py index 6e88ab9..6f5af77 100644 --- a/zhixing-server/tests/unit/sector_radar/test_facts.py +++ b/zhixing-server/tests/unit/sector_radar/test_facts.py @@ -50,7 +50,7 @@ def test_point_in_time_aggregation_distinguishes_suspension_missing_and_zero() - StockDailyFact( trade_date=TARGET_DATE, ts_code="000004.SZ", - status=StockFactStatus.MISSING, + status=StockFactStatus.MISSING_MONEYFLOW, ), ) diff --git a/zhixing-server/tests/unit/sector_radar/test_normalize.py b/zhixing-server/tests/unit/sector_radar/test_normalize.py new file mode 100644 index 0000000..d580249 --- /dev/null +++ b/zhixing-server/tests/unit/sector_radar/test_normalize.py @@ -0,0 +1,116 @@ +from datetime import UTC, date, datetime +from decimal import Decimal + +from zhixing_server.modules.sector_radar.domain.models import StockFactStatus +from zhixing_server.modules.sector_radar.domain.normalize import normalize_stock_facts +from zhixing_server.modules.sector_radar.domain.source import ( + DailyRow, + MoneyflowDcRow, + 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)) -> StockBasicRow: + return StockBasicRow( + ts_code=ts_code, + symbol=ts_code.split(".")[0], + name=ts_code, + market="主板", + exchange="SZSE", + list_status="L", + list_date=list_date, + delist_date=None, + ) + + +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_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="09:30", + 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 diff --git a/zhixing-server/tests/unit/sector_radar/test_postgres_repository.py b/zhixing-server/tests/unit/sector_radar/test_postgres_repository.py new file mode 100644 index 0000000..d96698e --- /dev/null +++ b/zhixing-server/tests/unit/sector_radar/test_postgres_repository.py @@ -0,0 +1,102 @@ +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) -> None: + self.row = row + + def fetchone(self) -> tuple[object, ...] | None: + return self.row + + +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_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] diff --git a/zhixing-server/tests/unit/sector_radar/test_repository.py b/zhixing-server/tests/unit/sector_radar/test_repository.py new file mode 100644 index 0000000..7440e5a --- /dev/null +++ b/zhixing-server/tests/unit/sector_radar/test_repository.py @@ -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)) diff --git a/zhixing-server/tests/unit/sector_radar/test_tushare_source.py b/zhixing-server/tests/unit/sector_radar/test_tushare_source.py new file mode 100644 index 0000000..7b8fccf --- /dev/null +++ b/zhixing-server/tests/unit/sector_radar/test_tushare_source.py @@ -0,0 +1,277 @@ +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]]] = [] + + def query(self, api_name: str, **kwargs: object) -> object: + self.calls.append((api_name, kwargs)) + partition = str(kwargs.get("ts_code") or kwargs.get("list_status") or "") + 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 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) + + 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_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_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_stock_basic_explicitly_requests_all_lifecycle_statuses() -> None: + responses = { + ( + "stock_basic", + status, + ): ( + { + "ts_code": f"00000{index}.SZ", + "symbol": f"00000{index}", + "name": status, + "market": "主板", + "exchange": "SZSE", + "list_status": status, + "list_date": "20200101", + "delist_date": None, + }, + ) + for index, status in enumerate(("L", "D", "P", "G", "UN"), start=1) + } + client = QueryClient(responses) + + result = make_adapter(client).fetch_stock_basics() + + assert {row.list_status for row in result.rows} == {"L", "D", "P", "G", "UN"} + assert [call[1]["list_status"] for call in client.calls] == ["L", "D", "P", "G", "UN"] + + +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_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"] == "概念板块"