Merge branch 'develop' into codex/point
This commit is contained in:
@@ -0,0 +1,234 @@
|
||||
"""Create replayable independent sector radar tables."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
revision: str = "0004_sector_radar"
|
||||
down_revision: str | None = "0003_market_integrity_checks"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create source, point-in-time fact, publication, and ranking tables."""
|
||||
|
||||
op.create_table(
|
||||
"sector_radar_source_snapshot",
|
||||
sa.Column("id", sa.String(64), primary_key=True),
|
||||
sa.Column("api_name", sa.String(32), nullable=False),
|
||||
sa.Column("normalized_params", JSONB, nullable=False),
|
||||
sa.Column("target_trade_date", sa.Date()),
|
||||
sa.Column("partition_key", sa.String(64)),
|
||||
sa.Column("observed_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("payload", JSONB, nullable=False),
|
||||
sa.Column("row_count", sa.Integer(), nullable=False),
|
||||
sa.Column("returned_fields", JSONB, nullable=False),
|
||||
sa.Column("content_sha256", sa.String(64), nullable=False),
|
||||
sa.Column("row_limit", sa.Integer()),
|
||||
sa.Column("limit_reached", sa.Boolean(), nullable=False),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.CheckConstraint("row_count >= 0", name="ck_sector_radar_source_row_count"),
|
||||
sa.CheckConstraint(
|
||||
"row_limit IS NULL OR row_limit > 0", name="ck_sector_radar_source_limit"
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_sector_radar_source_api_date",
|
||||
"sector_radar_source_snapshot",
|
||||
["api_name", "target_trade_date", "observed_at"],
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"sector_radar_membership",
|
||||
sa.Column(
|
||||
"source_snapshot_id",
|
||||
sa.String(64),
|
||||
sa.ForeignKey("sector_radar_source_snapshot.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("sector_type", sa.String(16), nullable=False),
|
||||
sa.Column("sector_code", sa.String(16), nullable=False),
|
||||
sa.Column("sector_name", sa.String(128), nullable=False),
|
||||
sa.Column("stock_code", sa.String(12), nullable=False),
|
||||
sa.Column("stock_name", sa.String(128), nullable=False),
|
||||
sa.Column("membership_status", sa.String(32), nullable=False),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("source_snapshot_id", "sector_code", "stock_code"),
|
||||
sa.CheckConstraint(
|
||||
"sector_type IN ('concept', 'industry')",
|
||||
name="ck_sector_radar_membership_type",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"membership_status = 'available'",
|
||||
name="ck_sector_radar_membership_status",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_sector_radar_membership_date_sector",
|
||||
"sector_radar_membership",
|
||||
["trade_date", "sector_type", "sector_code"],
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"sector_radar_stock_fact",
|
||||
sa.Column("fact_revision", sa.String(64), nullable=False),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("ts_code", sa.String(12), nullable=False),
|
||||
sa.Column("source_snapshot_ids", JSONB, nullable=False),
|
||||
sa.Column("status", sa.String(32), nullable=False),
|
||||
sa.Column("turnover_yuan", sa.Numeric(28, 6)),
|
||||
sa.Column("net_amount_yuan", sa.Numeric(28, 6)),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("fact_revision", "ts_code"),
|
||||
sa.CheckConstraint(
|
||||
"turnover_yuan IS NULL OR turnover_yuan >= 0",
|
||||
name="ck_sector_radar_stock_turnover",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_sector_radar_stock_fact_date",
|
||||
"sector_radar_stock_fact",
|
||||
["trade_date", "ts_code"],
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"sector_radar_publication",
|
||||
sa.Column("id", sa.String(64), primary_key=True),
|
||||
sa.Column("target_trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("status", sa.String(16), nullable=False),
|
||||
sa.Column("source_version", sa.String(128), nullable=False),
|
||||
sa.Column("universe_version", sa.String(128), nullable=False),
|
||||
sa.Column("metric_versions", JSONB, nullable=False),
|
||||
sa.Column("input_hash", sa.String(64)),
|
||||
sa.Column("coverage", sa.Numeric(8, 6), nullable=False),
|
||||
sa.Column("started_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("finished_at", sa.DateTime(timezone=True)),
|
||||
sa.Column("error_summary", sa.String(500)),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"status IN ('running', 'success', 'partial', 'failed')",
|
||||
name="ck_sector_radar_publication_status",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"coverage >= 0 AND coverage <= 1",
|
||||
name="ck_sector_radar_publication_coverage",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(status = 'running' AND finished_at IS NULL) OR "
|
||||
"(status <> 'running' AND finished_at IS NOT NULL)",
|
||||
name="ck_sector_radar_publication_finished",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_sector_radar_publication_status_date",
|
||||
"sector_radar_publication",
|
||||
["status", "target_trade_date", "finished_at"],
|
||||
)
|
||||
op.create_index(
|
||||
"uq_sector_radar_publication_running_date",
|
||||
"sector_radar_publication",
|
||||
["target_trade_date"],
|
||||
unique=True,
|
||||
postgresql_where=sa.text("status = 'running'"),
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"sector_radar_ranking",
|
||||
sa.Column(
|
||||
"publication_id",
|
||||
sa.String(64),
|
||||
sa.ForeignKey("sector_radar_publication.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("sector_type", sa.String(16), nullable=False),
|
||||
sa.Column("sector_code", sa.String(16), nullable=False),
|
||||
sa.Column("sector_name", sa.String(128), nullable=False),
|
||||
sa.Column("metric_kind", sa.String(16), nullable=False),
|
||||
sa.Column("metric_version", sa.String(128), nullable=False),
|
||||
sa.Column("implementation_kind", sa.String(16), nullable=False),
|
||||
sa.Column("unit", sa.String(16), nullable=False),
|
||||
sa.Column("metric_value", sa.Numeric(28, 12)),
|
||||
sa.Column("quality", sa.String(32), nullable=False),
|
||||
sa.Column("member_count", sa.Integer(), nullable=False),
|
||||
sa.Column("valid_sample_count", sa.Integer(), nullable=False),
|
||||
sa.Column("membership_coverage", sa.Numeric(8, 6), nullable=False),
|
||||
sa.Column("moneyflow_coverage", sa.Numeric(8, 6), nullable=False),
|
||||
sa.Column("rank_position", sa.Integer()),
|
||||
sa.Column("rank_percentile", sa.Numeric(18, 12)),
|
||||
sa.Column("rank_changes", JSONB, nullable=False),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint(
|
||||
"publication_id",
|
||||
"sector_type",
|
||||
"sector_code",
|
||||
"metric_version",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"sector_type IN ('concept', 'industry')",
|
||||
name="ck_sector_radar_ranking_type",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"implementation_kind = 'independent'",
|
||||
name="ck_sector_radar_ranking_implementation",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_sector_radar_ranking_query",
|
||||
"sector_radar_ranking",
|
||||
["publication_id", "sector_type", "metric_version", "rank_position"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop only sector radar tables in dependency-safe order."""
|
||||
|
||||
op.drop_index("ix_sector_radar_ranking_query", table_name="sector_radar_ranking")
|
||||
op.drop_table("sector_radar_ranking")
|
||||
op.drop_index(
|
||||
"uq_sector_radar_publication_running_date",
|
||||
table_name="sector_radar_publication",
|
||||
)
|
||||
op.drop_index(
|
||||
"ix_sector_radar_publication_status_date",
|
||||
table_name="sector_radar_publication",
|
||||
)
|
||||
op.drop_table("sector_radar_publication")
|
||||
op.drop_index("ix_sector_radar_stock_fact_date", table_name="sector_radar_stock_fact")
|
||||
op.drop_table("sector_radar_stock_fact")
|
||||
op.drop_index(
|
||||
"ix_sector_radar_membership_date_sector",
|
||||
table_name="sector_radar_membership",
|
||||
)
|
||||
op.drop_table("sector_radar_membership")
|
||||
op.drop_index("ix_sector_radar_source_api_date", table_name="sector_radar_source_snapshot")
|
||||
op.drop_table("sector_radar_source_snapshot")
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Persist publication-owned sector daily aggregates for exact metric replay."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "0005_radar_daily_aggregate"
|
||||
down_revision: str | None = "0004_sector_radar"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create source recovery links and exact multi-day metric inputs."""
|
||||
|
||||
op.create_table(
|
||||
"sector_radar_publication_source",
|
||||
sa.Column(
|
||||
"publication_id",
|
||||
sa.String(64),
|
||||
sa.ForeignKey("sector_radar_publication.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("source_group", sa.String(32), nullable=False),
|
||||
sa.Column("source_order", sa.Integer(), nullable=False),
|
||||
sa.Column(
|
||||
"refresh_on_retry",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.false(),
|
||||
),
|
||||
sa.Column(
|
||||
"source_snapshot_id",
|
||||
sa.String(64),
|
||||
sa.ForeignKey("sector_radar_source_snapshot.id", ondelete="RESTRICT"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.PrimaryKeyConstraint("publication_id", "source_group", "source_order"),
|
||||
sa.UniqueConstraint(
|
||||
"publication_id",
|
||||
"source_group",
|
||||
"source_snapshot_id",
|
||||
name="uq_sector_radar_publication_source_snapshot",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"source_group IN ('calendar', 'concept_indices', 'industry_indices', "
|
||||
"'members', 'stock_basics', 'suspensions', 'daily', 'moneyflow_dc')",
|
||||
name="ck_sector_radar_publication_source_group",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"source_order >= 0",
|
||||
name="ck_sector_radar_publication_source_order",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_sector_radar_publication_source_snapshot",
|
||||
"sector_radar_publication_source",
|
||||
["source_snapshot_id"],
|
||||
)
|
||||
|
||||
op.create_unique_constraint(
|
||||
"uq_sector_radar_publication_id_date",
|
||||
"sector_radar_publication",
|
||||
["id", "target_trade_date"],
|
||||
)
|
||||
op.create_foreign_key(
|
||||
"fk_sector_radar_ranking_publication_date",
|
||||
"sector_radar_ranking",
|
||||
"sector_radar_publication",
|
||||
["publication_id", "trade_date"],
|
||||
["id", "target_trade_date"],
|
||||
ondelete="CASCADE",
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"sector_radar_daily_aggregate",
|
||||
sa.Column("publication_id", sa.String(64), nullable=False),
|
||||
sa.Column("trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("sector_type", sa.String(16), nullable=False),
|
||||
sa.Column("sector_code", sa.String(16), nullable=False),
|
||||
sa.Column("sector_name", sa.String(128), nullable=False),
|
||||
sa.Column("member_count", sa.Integer(), nullable=False),
|
||||
sa.Column("valid_sample_count", sa.Integer(), nullable=False),
|
||||
sa.Column("net_amount_yuan", sa.Numeric(28, 6)),
|
||||
sa.Column("turnover_yuan", sa.Numeric(28, 6)),
|
||||
sa.Column("membership_coverage", sa.Numeric(8, 6), nullable=False),
|
||||
sa.Column("moneyflow_coverage", sa.Numeric(8, 6), nullable=False),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("publication_id", "sector_type", "sector_code"),
|
||||
sa.ForeignKeyConstraint(
|
||||
["publication_id", "trade_date"],
|
||||
["sector_radar_publication.id", "sector_radar_publication.target_trade_date"],
|
||||
name="fk_sector_radar_daily_aggregate_publication_date",
|
||||
ondelete="CASCADE",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"sector_type IN ('concept', 'industry')",
|
||||
name="ck_sector_radar_daily_aggregate_type",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"member_count >= 0 AND valid_sample_count >= 0 AND valid_sample_count <= member_count",
|
||||
name="ck_sector_radar_daily_aggregate_counts",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"membership_coverage >= 0 AND membership_coverage <= 1 "
|
||||
"AND moneyflow_coverage >= 0 AND moneyflow_coverage <= 1",
|
||||
name="ck_sector_radar_daily_aggregate_coverage",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"turnover_yuan IS NULL OR (turnover_yuan >= 0 AND "
|
||||
"turnover_yuan NOT IN ('NaN'::numeric, 'Infinity'::numeric))",
|
||||
name="ck_sector_radar_daily_aggregate_turnover",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"net_amount_yuan IS NULL OR net_amount_yuan NOT IN "
|
||||
"('NaN'::numeric, 'Infinity'::numeric, '-Infinity'::numeric)",
|
||||
name="ck_sector_radar_daily_aggregate_net_amount",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_sector_radar_daily_aggregate_history",
|
||||
"sector_radar_daily_aggregate",
|
||||
["trade_date", "sector_type", "sector_code"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop only the replay aggregate extension."""
|
||||
|
||||
op.drop_index(
|
||||
"ix_sector_radar_daily_aggregate_history",
|
||||
table_name="sector_radar_daily_aggregate",
|
||||
)
|
||||
op.drop_table("sector_radar_daily_aggregate")
|
||||
op.drop_constraint(
|
||||
"fk_sector_radar_ranking_publication_date",
|
||||
"sector_radar_ranking",
|
||||
type_="foreignkey",
|
||||
)
|
||||
op.drop_constraint(
|
||||
"uq_sector_radar_publication_id_date",
|
||||
"sector_radar_publication",
|
||||
type_="unique",
|
||||
)
|
||||
op.drop_index(
|
||||
"ix_sector_radar_publication_source_snapshot",
|
||||
table_name="sector_radar_publication_source",
|
||||
)
|
||||
op.drop_table("sector_radar_publication_source")
|
||||
@@ -0,0 +1,103 @@
|
||||
"""Persist explicit unknown point-in-time sector membership snapshots."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "0006_membership_unknown"
|
||||
down_revision: str | None = "0005_radar_daily_aggregate"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Allow one null-stock marker for an explicitly empty sector partition."""
|
||||
|
||||
op.add_column(
|
||||
"sector_radar_membership",
|
||||
sa.Column("membership_key", sa.String(32), nullable=True),
|
||||
)
|
||||
op.execute("UPDATE sector_radar_membership SET membership_key = stock_code")
|
||||
op.drop_constraint(
|
||||
"sector_radar_membership_pkey",
|
||||
"sector_radar_membership",
|
||||
type_="primary",
|
||||
)
|
||||
op.drop_constraint(
|
||||
"ck_sector_radar_membership_status",
|
||||
"sector_radar_membership",
|
||||
type_="check",
|
||||
)
|
||||
op.alter_column(
|
||||
"sector_radar_membership",
|
||||
"membership_key",
|
||||
existing_type=sa.String(32),
|
||||
nullable=False,
|
||||
)
|
||||
op.alter_column(
|
||||
"sector_radar_membership",
|
||||
"stock_code",
|
||||
existing_type=sa.String(12),
|
||||
nullable=True,
|
||||
)
|
||||
op.alter_column(
|
||||
"sector_radar_membership",
|
||||
"stock_name",
|
||||
existing_type=sa.String(128),
|
||||
nullable=True,
|
||||
)
|
||||
op.create_primary_key(
|
||||
"sector_radar_membership_pkey",
|
||||
"sector_radar_membership",
|
||||
["source_snapshot_id", "sector_code", "membership_key"],
|
||||
)
|
||||
op.create_check_constraint(
|
||||
"ck_sector_radar_membership_status",
|
||||
"sector_radar_membership",
|
||||
"(membership_status = 'available' "
|
||||
"AND stock_code IS NOT NULL AND stock_name IS NOT NULL "
|
||||
"AND membership_key = stock_code) OR "
|
||||
"(membership_status = 'membership_unknown' "
|
||||
"AND stock_code IS NULL AND stock_name IS NULL "
|
||||
"AND membership_key = '__membership_unknown__')",
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Discard unknown markers and restore the available-member-only schema."""
|
||||
|
||||
op.execute("DELETE FROM sector_radar_membership WHERE membership_status = 'membership_unknown'")
|
||||
op.drop_constraint(
|
||||
"ck_sector_radar_membership_status",
|
||||
"sector_radar_membership",
|
||||
type_="check",
|
||||
)
|
||||
op.drop_constraint(
|
||||
"sector_radar_membership_pkey",
|
||||
"sector_radar_membership",
|
||||
type_="primary",
|
||||
)
|
||||
op.alter_column(
|
||||
"sector_radar_membership",
|
||||
"stock_code",
|
||||
existing_type=sa.String(12),
|
||||
nullable=False,
|
||||
)
|
||||
op.alter_column(
|
||||
"sector_radar_membership",
|
||||
"stock_name",
|
||||
existing_type=sa.String(128),
|
||||
nullable=False,
|
||||
)
|
||||
op.drop_column("sector_radar_membership", "membership_key")
|
||||
op.create_primary_key(
|
||||
"sector_radar_membership_pkey",
|
||||
"sector_radar_membership",
|
||||
["source_snapshot_id", "sector_code", "stock_code"],
|
||||
)
|
||||
op.create_check_constraint(
|
||||
"ck_sector_radar_membership_status",
|
||||
"sector_radar_membership",
|
||||
"membership_status = 'available'",
|
||||
)
|
||||
@@ -36,6 +36,7 @@ packages = ["src/zhixing_server"]
|
||||
|
||||
[project.scripts]
|
||||
market-data-sync = "zhixing_server.modules.market_data.presentation.cli:main"
|
||||
sector-radar-build = "zhixing_server.modules.sector_radar.presentation.cli:main"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
addopts = "-ra --strict-config --strict-markers"
|
||||
|
||||
@@ -24,6 +24,11 @@ class Settings(BaseSettings):
|
||||
market_data_max_retries: int = 3
|
||||
market_data_retry_backoff_seconds: float = 1.0
|
||||
market_data_advisory_lock_key: int = 7_380_521
|
||||
sector_radar_coverage_threshold: Decimal = Decimal("0.99")
|
||||
sector_radar_request_interval_seconds: float = 0.2
|
||||
sector_radar_max_retries: int = 3
|
||||
sector_radar_retry_backoff_seconds: float = 1.0
|
||||
sector_radar_advisory_lock_key: int = 7_380_522
|
||||
selection_max_workers: int = Field(default=4, ge=1)
|
||||
selection_batch_size: int = Field(default=200, ge=1)
|
||||
selection_pattern_scoring_enabled: bool = True
|
||||
|
||||
@@ -5,6 +5,7 @@ from fastapi import APIRouter
|
||||
from zhixing_server.interfaces.http.system import operational_router, system_router
|
||||
from zhixing_server.modules.market_data.presentation.home import home_router
|
||||
from zhixing_server.modules.market_data.presentation.integrity import integrity_router
|
||||
from zhixing_server.modules.sector_radar.presentation.http import sector_radar_router
|
||||
from zhixing_server.modules.selection.presentation.http import selection_router
|
||||
|
||||
api_v1_router = APIRouter(prefix="/api/v1")
|
||||
@@ -16,5 +17,10 @@ api_v1_router.include_router(
|
||||
tags=["market-data"],
|
||||
)
|
||||
api_v1_router.include_router(selection_router, prefix="/selection", tags=["selection"])
|
||||
api_v1_router.include_router(
|
||||
sector_radar_router,
|
||||
prefix="/sector-radar",
|
||||
tags=["sector-radar"],
|
||||
)
|
||||
|
||||
__all__ = ["api_v1_router", "operational_router"]
|
||||
|
||||
@@ -2,186 +2,28 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from datetime import date
|
||||
from typing import cast
|
||||
|
||||
from zhixing_server.shared.request_coordinator import (
|
||||
DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
RequestCoordinator,
|
||||
TushareRequestCoordinator,
|
||||
TushareSourceError,
|
||||
)
|
||||
|
||||
from ..domain.models import Bar, DailyBasic, Stock, SyncWindow, parse_date
|
||||
from ..domain.rules import filter_current_hs_a_stocks
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_RATE_LIMIT_COOLDOWNS = (60.0, 120.0, 180.0)
|
||||
_RATE_LIMIT_MESSAGES = (
|
||||
"访问频繁",
|
||||
"请稍后",
|
||||
"超过频率",
|
||||
"频率限制",
|
||||
"too many requests",
|
||||
"rate limit",
|
||||
"rate_limit",
|
||||
"http 429",
|
||||
"status code: 429",
|
||||
"429",
|
||||
"http 403",
|
||||
"status code: 403",
|
||||
"403",
|
||||
)
|
||||
|
||||
|
||||
class TushareSourceError(RuntimeError):
|
||||
"""A vendor request failed after the configured retry budget."""
|
||||
|
||||
|
||||
class RequestCoordinator:
|
||||
"""Coordinate retry and shared rate-limit cooling for one token client.
|
||||
|
||||
Normal requests are deliberately not serialized. Only a provider rate
|
||||
limit creates a shared cooldown, so independent worker calls can proceed
|
||||
concurrently during ordinary traffic. ``clock`` and ``wait_fn`` are
|
||||
injectable to make long cooldown behavior deterministic in unit tests.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
random_fn: Callable[[], float] = random.random,
|
||||
clock: Callable[[], float] = time.monotonic,
|
||||
wait_fn: Callable[[float], None] = time.sleep,
|
||||
sleep_fn: Callable[[float], None] | None = None,
|
||||
) -> None:
|
||||
cooldowns = tuple(float(value) for value in cooldown_seconds)
|
||||
if not cooldowns or any(value < 0 for value in cooldowns):
|
||||
raise ValueError("cooldown_seconds must contain non-negative values")
|
||||
self.max_retries = max(0, max_retries)
|
||||
self.backoff_seconds = max(0.0, backoff_seconds)
|
||||
self.cooldown_seconds = cooldowns
|
||||
self.random_fn = random_fn
|
||||
self.clock = clock
|
||||
self.wait_fn = wait_fn
|
||||
self.sleep_fn = sleep_fn or wait_fn
|
||||
self._condition = threading.Condition()
|
||||
self._cooldown_until = 0.0
|
||||
self._rate_limit_count = 0
|
||||
|
||||
@property
|
||||
def cooldown_until(self) -> float:
|
||||
"""Return the current monotonic cooldown deadline."""
|
||||
|
||||
with self._condition:
|
||||
return self._cooldown_until
|
||||
|
||||
def call(self, method_name: str, request: Callable[[], object]) -> object:
|
||||
"""Execute one provider request with bounded, shared retry behavior."""
|
||||
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
self._wait_for_cooldown(method_name)
|
||||
try:
|
||||
result = request()
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
if self.is_rate_limited(exc):
|
||||
cooldown = self._set_rate_limit_cooldown()
|
||||
logger.warning(
|
||||
"tushare_rate_limit method=%s attempt=%d max_attempts=%d "
|
||||
"cooldown_seconds=%.1f",
|
||||
method_name,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
cooldown,
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
continue
|
||||
break
|
||||
if not self._is_retryable(exc):
|
||||
raise
|
||||
if attempt == self.max_retries:
|
||||
break
|
||||
delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn())
|
||||
logger.warning(
|
||||
"tushare_request_retry method=%s attempt=%d max_attempts=%d "
|
||||
"backoff_seconds=%.1f",
|
||||
method_name,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
delay,
|
||||
)
|
||||
self.sleep_fn(delay)
|
||||
else:
|
||||
self._clear_rate_limit_after_success()
|
||||
return result
|
||||
logger.error(
|
||||
"tushare_request_failed method=%s attempts=%d",
|
||||
method_name,
|
||||
self.max_retries + 1,
|
||||
)
|
||||
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
|
||||
|
||||
def request(self, method_name: str, operation: Callable[[], object]) -> object:
|
||||
"""Alias for ``call`` for adapters that model requests as a port."""
|
||||
|
||||
return self.call(method_name, operation)
|
||||
|
||||
def _wait_for_cooldown(self, method_name: str) -> None:
|
||||
while True:
|
||||
with self._condition:
|
||||
delay = self._cooldown_until - self.clock()
|
||||
if delay <= 0:
|
||||
return
|
||||
logger.info(
|
||||
"tushare_rate_limit_wait method=%s wait_seconds=%.1f",
|
||||
method_name,
|
||||
delay,
|
||||
)
|
||||
# A single injected wait hook makes fake-clock tests independent
|
||||
# from wall time. After waiting, re-check because another worker
|
||||
# may have extended the shared deadline.
|
||||
self.wait_fn(delay)
|
||||
|
||||
def _set_rate_limit_cooldown(self) -> float:
|
||||
with self._condition:
|
||||
self._rate_limit_count += 1
|
||||
index = min(self._rate_limit_count - 1, len(self.cooldown_seconds) - 1)
|
||||
duration = self.cooldown_seconds[index]
|
||||
self._cooldown_until = max(self._cooldown_until, self.clock() + duration)
|
||||
self._condition.notify_all()
|
||||
return duration
|
||||
|
||||
def _clear_rate_limit_after_success(self) -> None:
|
||||
with self._condition:
|
||||
# A request that was already in flight when another worker hit a
|
||||
# limit may succeed during the shared cooldown. Do not erase the
|
||||
# escalation history until the cooldown has actually elapsed.
|
||||
if self.clock() >= self._cooldown_until:
|
||||
self._rate_limit_count = 0
|
||||
|
||||
@staticmethod
|
||||
def is_rate_limited(error: BaseException) -> bool:
|
||||
"""Classify stable provider rate-limit signals without logging details."""
|
||||
|
||||
for attribute in ("status_code", "status", "code"):
|
||||
value = getattr(error, attribute, None)
|
||||
if str(value).strip() in {"403", "429"}:
|
||||
return True
|
||||
message = str(error).casefold()
|
||||
return any(marker.casefold() in message for marker in _RATE_LIMIT_MESSAGES)
|
||||
|
||||
@staticmethod
|
||||
def _is_retryable(error: BaseException) -> bool:
|
||||
return isinstance(error, (OSError, RuntimeError, TimeoutError))
|
||||
|
||||
|
||||
# The longer name is useful to callers that want to make the infrastructure
|
||||
# boundary explicit, while the short name remains convenient in unit tests.
|
||||
TushareRequestCoordinator = RequestCoordinator
|
||||
__all__ = [
|
||||
"RequestCoordinator",
|
||||
"TushareAdapter",
|
||||
"TushareRequestCoordinator",
|
||||
"TushareSourceError",
|
||||
]
|
||||
|
||||
|
||||
class CoordinatedTushareClient:
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Independent post-close sector capital radar bounded context."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Application use cases for sector radar production and reads."""
|
||||
@@ -0,0 +1,765 @@
|
||||
"""One-shot, idempotent sector radar publication orchestration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import UTC, date, datetime, time, timedelta
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
from uuid import uuid4
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from ..domain.facts import aggregate_sector_snapshot
|
||||
from ..domain.metrics import (
|
||||
AmountNetStrategy,
|
||||
MetricStrategy,
|
||||
RatioTurnoverStrategy,
|
||||
SwingEqualThreeToTenStrategy,
|
||||
)
|
||||
from ..domain.models import (
|
||||
MembershipStatus,
|
||||
MetricObservation,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
SectorDailyAggregate,
|
||||
SectorMembershipSnapshot,
|
||||
SectorType,
|
||||
StockDailyFact,
|
||||
StockFactStatus,
|
||||
)
|
||||
from ..domain.normalize import (
|
||||
is_current_listed_stock,
|
||||
normalize_memberships,
|
||||
normalize_stock_facts,
|
||||
)
|
||||
from ..domain.persistence import (
|
||||
DailyAggregateRecord,
|
||||
MembershipRecord,
|
||||
PublicationSourceGroup,
|
||||
PublicationSourceRecord,
|
||||
RankingRecord,
|
||||
SectorRadarRepository,
|
||||
StockFactRecord,
|
||||
)
|
||||
from ..domain.ports import SectorRadarSource
|
||||
from ..domain.ranking import rank_metric_observations, with_rank_changes
|
||||
from ..domain.source import (
|
||||
DailyRow,
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceContractError,
|
||||
SourceResult,
|
||||
SourceScalar,
|
||||
SourceSnapshot,
|
||||
StockBasicRow,
|
||||
SuspendRow,
|
||||
TradeCalendarRow,
|
||||
)
|
||||
|
||||
BuildOutcomeStatus = Literal["success", "partial", "failed", "locked", "unchanged"]
|
||||
SHANGHAI = ZoneInfo("Asia/Shanghai")
|
||||
MARKET_DATA_READY_TIME = time(15, 30)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BuildSectorRadarCommand:
|
||||
"""Select one date, an inclusive range, or a failed publication retry."""
|
||||
|
||||
trade_date: date | None = None
|
||||
start_date: date | None = None
|
||||
end_date: date | None = None
|
||||
retry_publication_id: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject ambiguous build modes before any source or database call."""
|
||||
|
||||
has_range = self.start_date is not None or self.end_date is not None
|
||||
modes = sum(
|
||||
(
|
||||
self.trade_date is not None,
|
||||
has_range,
|
||||
self.retry_publication_id is not None,
|
||||
)
|
||||
)
|
||||
if modes > 1:
|
||||
raise ValueError("trade date, date range, and retry publication are mutually exclusive")
|
||||
if has_range and (self.start_date is None or self.end_date is None):
|
||||
raise ValueError("date range requires both start_date and end_date")
|
||||
if (
|
||||
self.start_date is not None
|
||||
and self.end_date is not None
|
||||
and self.end_date < self.start_date
|
||||
):
|
||||
raise ValueError("end_date must not precede start_date")
|
||||
if self.retry_publication_id is not None and not self.retry_publication_id.strip():
|
||||
raise ValueError("retry_publication_id must not be empty")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BuildDateOutcome:
|
||||
"""Redacted result for one target trade date."""
|
||||
|
||||
target_trade_date: date
|
||||
status: BuildOutcomeStatus
|
||||
publication_id: str | None
|
||||
coverage: Decimal
|
||||
sector_count: int
|
||||
ranking_count: int
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
"""Serialize without raw payloads, credentials, or provider exception text."""
|
||||
|
||||
return {
|
||||
"target_trade_date": self.target_trade_date.isoformat(),
|
||||
"status": self.status,
|
||||
"publication_id": self.publication_id,
|
||||
"coverage": str(self.coverage),
|
||||
"sector_count": self.sector_count,
|
||||
"ranking_count": self.ranking_count,
|
||||
"error_type": self.error_type,
|
||||
"error_message": self.error_message,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BuildSummary:
|
||||
"""Cron-friendly aggregate result for one CLI invocation."""
|
||||
|
||||
outcomes: tuple[BuildDateOutcome, ...]
|
||||
|
||||
@property
|
||||
def status(self) -> str:
|
||||
"""Return the worst invocation state."""
|
||||
|
||||
if not self.outcomes or any(item.status in {"failed", "locked"} for item in self.outcomes):
|
||||
return "failed"
|
||||
if any(item.status == "partial" for item in self.outcomes):
|
||||
return "partial"
|
||||
if all(item.status == "unchanged" for item in self.outcomes):
|
||||
return "unchanged"
|
||||
return "success"
|
||||
|
||||
@property
|
||||
def exit_code(self) -> int:
|
||||
"""Return 0 for usable success, 2 for incomplete input, and 1 for failure."""
|
||||
|
||||
if self.status == "failed":
|
||||
return 1
|
||||
if self.status == "partial":
|
||||
return 2
|
||||
return 0
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
"""Serialize the invocation summary for external schedulers."""
|
||||
|
||||
return {
|
||||
"status": self.status,
|
||||
"exit_code": self.exit_code,
|
||||
"outcomes": [item.as_dict() for item in self.outcomes],
|
||||
}
|
||||
|
||||
|
||||
class BuildSectorRadar:
|
||||
"""Hide target resolution, source replay, metrics, ranking, and publication switching."""
|
||||
|
||||
source_version = "tushare-pro-v1"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
source: SectorRadarSource,
|
||||
repository: SectorRadarRepository,
|
||||
*,
|
||||
coverage_threshold: Decimal = Decimal("0.99"),
|
||||
today: date | None = None,
|
||||
now_fn: Callable[[], datetime] = lambda: datetime.now(UTC),
|
||||
strategies: Sequence[MetricStrategy] | None = None,
|
||||
) -> None:
|
||||
if not Decimal(0) <= coverage_threshold <= Decimal(1):
|
||||
raise ValueError("coverage_threshold must be between 0 and 1")
|
||||
self.now_fn = now_fn
|
||||
self.source = source
|
||||
self.repository = repository
|
||||
self.coverage_threshold = coverage_threshold
|
||||
self.today = today or self.now_fn().astimezone(SHANGHAI).date()
|
||||
self.strategies = tuple(
|
||||
strategies
|
||||
or (
|
||||
AmountNetStrategy(),
|
||||
RatioTurnoverStrategy(),
|
||||
SwingEqualThreeToTenStrategy(),
|
||||
)
|
||||
)
|
||||
|
||||
def execute(self, command: BuildSectorRadarCommand | None = None) -> BuildSummary:
|
||||
"""Build each selected trade date sequentially for deterministic history."""
|
||||
|
||||
command = command or BuildSectorRadarCommand()
|
||||
try:
|
||||
targets = self._resolve_targets(command)
|
||||
except Exception as exc:
|
||||
target = command.trade_date or command.start_date or self.today
|
||||
error_type, message = self._safe_failure(exc)
|
||||
return BuildSummary(
|
||||
(
|
||||
BuildDateOutcome(
|
||||
target,
|
||||
"failed",
|
||||
None,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
error_type,
|
||||
message,
|
||||
),
|
||||
)
|
||||
)
|
||||
return BuildSummary(tuple(self._build_target(target) for target in targets))
|
||||
|
||||
def _resolve_targets(self, command: BuildSectorRadarCommand) -> tuple[_BuildTarget, ...]:
|
||||
if command.retry_publication_id is not None:
|
||||
publication = self.repository.get_publication(command.retry_publication_id)
|
||||
if publication is None:
|
||||
raise ValueError("retry publication does not exist")
|
||||
if publication.status not in {PublicationStatus.PARTIAL, PublicationStatus.FAILED}:
|
||||
raise ValueError("only partial or failed publications can be retried")
|
||||
return (_BuildTarget(publication.target_trade_date, publication.publication_id),)
|
||||
|
||||
if command.trade_date is not None:
|
||||
start = end = command.trade_date
|
||||
elif command.start_date is not None and command.end_date is not None:
|
||||
start, end = command.start_date, command.end_date
|
||||
else:
|
||||
end = self._default_calendar_end()
|
||||
start = end - timedelta(days=14)
|
||||
calendar = self.source.fetch_trade_calendar(start, end)
|
||||
targets = tuple(sorted({row.cal_date for row in calendar.rows if row.is_open}))
|
||||
if command.trade_date is not None and command.trade_date not in targets:
|
||||
raise ValueError("target date is not an open trading day")
|
||||
if not targets:
|
||||
raise ValueError("no open trading date found")
|
||||
selected = targets if command.start_date is not None else (targets[-1],)
|
||||
return tuple(_BuildTarget(target) for target in selected)
|
||||
|
||||
def _default_calendar_end(self) -> date:
|
||||
"""Exclude today's session until Tushare closing facts are expected to be ready."""
|
||||
|
||||
local_now = self.now_fn().astimezone(SHANGHAI)
|
||||
if self.today == local_now.date() and local_now.time() < MARKET_DATA_READY_TIME:
|
||||
return self.today - timedelta(days=1)
|
||||
return self.today
|
||||
|
||||
def _build_target(self, target: _BuildTarget) -> BuildDateOutcome:
|
||||
try:
|
||||
with self.repository.advisory_lock(target.trade_date) as acquired:
|
||||
if not acquired:
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"locked",
|
||||
None,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
"build_locked",
|
||||
"another sector radar build is running for this date",
|
||||
)
|
||||
return self._build_locked(target)
|
||||
except Exception as exc:
|
||||
error_type, message = self._safe_failure(exc)
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"failed",
|
||||
None,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
error_type,
|
||||
message,
|
||||
)
|
||||
|
||||
def _build_locked(self, target: _BuildTarget) -> BuildDateOutcome:
|
||||
started_at = self.now_fn()
|
||||
publication: RadarPublication | None = None
|
||||
publication_created = False
|
||||
try:
|
||||
self.repository.recover_running_publications(
|
||||
target.trade_date,
|
||||
finished_at=started_at,
|
||||
)
|
||||
publication_id = self._running_id(target.trade_date)
|
||||
publication = RadarPublication(
|
||||
publication_id=publication_id,
|
||||
target_trade_date=target.trade_date,
|
||||
status=PublicationStatus.RUNNING,
|
||||
source_version=self.source_version,
|
||||
universe_version="pending",
|
||||
metric_versions=tuple(strategy.metric_version for strategy in self.strategies),
|
||||
input_hash=None,
|
||||
coverage=Decimal(0),
|
||||
started_at=started_at,
|
||||
)
|
||||
self.repository.create_publication(publication)
|
||||
publication_created = True
|
||||
reusable = self._reusable_sources(target.retry_publication_id)
|
||||
collected = self._collect(target.trade_date, publication_id, reusable)
|
||||
input_hash = self._input_hash(collected.snapshots)
|
||||
existing = self.repository.find_reusable_publication(target.trade_date, input_hash)
|
||||
if existing is not None:
|
||||
self.repository.discard_running_publication(publication_id)
|
||||
publication_created = False
|
||||
is_success = existing.status is PublicationStatus.SUCCESS
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"unchanged" if is_success else "partial",
|
||||
existing.publication_id,
|
||||
existing.coverage,
|
||||
0,
|
||||
0,
|
||||
None if is_success else "duplicate_input",
|
||||
None
|
||||
if is_success
|
||||
else "input is unchanged from an existing partial publication",
|
||||
)
|
||||
|
||||
publication = replace(
|
||||
publication,
|
||||
universe_version=self._universe_version(collected.membership_snapshots),
|
||||
input_hash=input_hash,
|
||||
)
|
||||
aggregates = self._aggregate(collected)
|
||||
rankings = self._rank(target.trade_date, aggregates)
|
||||
coverage = self._coverage(collected.stock_facts)
|
||||
membership_complete = all(
|
||||
item.status is MembershipStatus.AVAILABLE for item in collected.memberships
|
||||
)
|
||||
terminal = (
|
||||
PublicationStatus.SUCCESS
|
||||
if membership_complete and coverage >= self.coverage_threshold
|
||||
else PublicationStatus.PARTIAL
|
||||
)
|
||||
finished = RadarPublication(
|
||||
publication_id=publication.publication_id,
|
||||
target_trade_date=target.trade_date,
|
||||
status=terminal,
|
||||
source_version=publication.source_version,
|
||||
universe_version=publication.universe_version,
|
||||
metric_versions=publication.metric_versions,
|
||||
input_hash=input_hash,
|
||||
coverage=coverage,
|
||||
started_at=started_at,
|
||||
finished_at=self.now_fn(),
|
||||
error_summary=(
|
||||
None
|
||||
if terminal is PublicationStatus.SUCCESS
|
||||
else (
|
||||
"membership_unknown"
|
||||
if not membership_complete
|
||||
else "coverage_below_threshold"
|
||||
)
|
||||
),
|
||||
)
|
||||
self.repository.finalize_publication(
|
||||
finished,
|
||||
memberships=collected.memberships,
|
||||
stock_facts=collected.stock_facts,
|
||||
daily_aggregates=(
|
||||
DailyAggregateRecord(publication_id, aggregate) for aggregate in aggregates
|
||||
),
|
||||
rankings=(RankingRecord(publication_id, ranking) for ranking in rankings),
|
||||
retry_source_groups=(
|
||||
self._retry_source_groups(
|
||||
collected.memberships,
|
||||
collected.stock_facts,
|
||||
)
|
||||
if terminal is PublicationStatus.PARTIAL
|
||||
else ()
|
||||
),
|
||||
)
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"success" if terminal is PublicationStatus.SUCCESS else "partial",
|
||||
publication_id,
|
||||
coverage,
|
||||
len(aggregates),
|
||||
len(rankings),
|
||||
)
|
||||
except Exception as exc:
|
||||
error_type, message = self._safe_failure(exc)
|
||||
failed_id = (
|
||||
publication.publication_id
|
||||
if publication is not None
|
||||
else self._failure_id(target.trade_date)
|
||||
)
|
||||
if isinstance(exc, SourceContractError) and exc.claim_diagnostic():
|
||||
logger.error(
|
||||
"sector_radar_build_source_contract_failed "
|
||||
"target_trade_date=%s publication_id=%s validation=%s",
|
||||
target.trade_date.isoformat(),
|
||||
failed_id,
|
||||
exc.operator_message,
|
||||
)
|
||||
if publication is not None and publication_created:
|
||||
self.repository.finish_publication(
|
||||
RadarPublication(
|
||||
publication_id=failed_id,
|
||||
target_trade_date=target.trade_date,
|
||||
status=PublicationStatus.FAILED,
|
||||
source_version=publication.source_version,
|
||||
universe_version=publication.universe_version,
|
||||
metric_versions=publication.metric_versions,
|
||||
input_hash=publication.input_hash,
|
||||
coverage=Decimal(0),
|
||||
started_at=started_at,
|
||||
finished_at=self.now_fn(),
|
||||
error_summary=f"{error_type}:{message}",
|
||||
)
|
||||
)
|
||||
return BuildDateOutcome(
|
||||
target.trade_date,
|
||||
"failed",
|
||||
failed_id,
|
||||
Decimal(0),
|
||||
0,
|
||||
0,
|
||||
error_type,
|
||||
message,
|
||||
)
|
||||
|
||||
def _reusable_sources(
|
||||
self, publication_id: str | None
|
||||
) -> dict[PublicationSourceGroup, tuple[SourceSnapshot, ...]]:
|
||||
"""Load successful checkpoints while forcing incomplete coverage facts to refresh."""
|
||||
|
||||
if publication_id is None:
|
||||
return {}
|
||||
publication = self.repository.get_publication(publication_id)
|
||||
if publication is None:
|
||||
raise ValueError("retry publication does not exist")
|
||||
grouped: defaultdict[PublicationSourceGroup, list[PublicationSourceRecord]] = defaultdict(
|
||||
list
|
||||
)
|
||||
for record in self.repository.load_publication_sources(publication_id):
|
||||
if not record.refresh_on_retry:
|
||||
grouped[record.source_group].append(record)
|
||||
result: dict[PublicationSourceGroup, tuple[SourceSnapshot, ...]] = {}
|
||||
for group, records in grouped.items():
|
||||
ordered = sorted(records, key=lambda item: item.source_order)
|
||||
if [item.source_order for item in ordered] != list(range(len(ordered))):
|
||||
raise ValueError("publication source checkpoint order is incomplete")
|
||||
result[group] = tuple(item.snapshot for item in ordered)
|
||||
return result
|
||||
|
||||
def _fetch_group[T](
|
||||
self,
|
||||
publication_id: str,
|
||||
source_group: PublicationSourceGroup,
|
||||
reusable: Mapping[PublicationSourceGroup, tuple[SourceSnapshot, ...]],
|
||||
fetch: Callable[[], SourceResult[T]],
|
||||
parser: Callable[[Mapping[str, SourceScalar]], T],
|
||||
reuse_if: Callable[[SourceResult[T]], bool] | None = None,
|
||||
) -> SourceResult[T]:
|
||||
"""Replay a compatible completed group or fetch and checkpoint it immediately."""
|
||||
|
||||
try:
|
||||
snapshots = reusable.get(source_group)
|
||||
if snapshots is None:
|
||||
result = fetch()
|
||||
else:
|
||||
replayed = SourceResult(
|
||||
snapshots=snapshots,
|
||||
rows=tuple(parser(row) for snapshot in snapshots for row in snapshot.rows),
|
||||
)
|
||||
result = replayed if reuse_if is None or reuse_if(replayed) else fetch()
|
||||
except SourceContractError as exc:
|
||||
if exc.claim_diagnostic():
|
||||
logger.error(
|
||||
"sector_radar_source_group_contract_failed "
|
||||
"publication_id=%s source_group=%s validation=%s",
|
||||
publication_id,
|
||||
source_group.value,
|
||||
exc.operator_message,
|
||||
)
|
||||
raise
|
||||
if not result.snapshots:
|
||||
raise ValueError("source group must include at least one replay snapshot")
|
||||
snapshot_ids = [snapshot.snapshot_id for snapshot in result.snapshots]
|
||||
if len(snapshot_ids) != len(set(snapshot_ids)):
|
||||
raise ValueError("source group contains duplicate snapshots")
|
||||
self.repository.save_source_snapshots(result.snapshots)
|
||||
self.repository.save_publication_sources(
|
||||
PublicationSourceRecord(publication_id, source_group, order, snapshot)
|
||||
for order, snapshot in enumerate(result.snapshots)
|
||||
)
|
||||
return result
|
||||
|
||||
def _collect(
|
||||
self,
|
||||
target: date,
|
||||
publication_id: str,
|
||||
reusable: Mapping[PublicationSourceGroup, tuple[SourceSnapshot, ...]],
|
||||
) -> _CollectedInputs:
|
||||
calendar = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.CALENDAR,
|
||||
reusable,
|
||||
lambda: self.source.fetch_trade_calendar(target, target),
|
||||
TradeCalendarRow.from_mapping,
|
||||
)
|
||||
if target not in {row.cal_date for row in calendar.rows if row.is_open}:
|
||||
raise ValueError("target date is not an open trading day")
|
||||
concepts = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.CONCEPT_INDICES,
|
||||
reusable,
|
||||
lambda: self.source.fetch_sector_indices(target, SectorType.CONCEPT),
|
||||
lambda row: SectorIndexRow.from_mapping(row, SectorType.CONCEPT),
|
||||
)
|
||||
industries = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.INDUSTRY_INDICES,
|
||||
reusable,
|
||||
lambda: self.source.fetch_sector_indices(target, SectorType.INDUSTRY),
|
||||
lambda row: SectorIndexRow.from_mapping(row, SectorType.INDUSTRY),
|
||||
)
|
||||
indices = concepts.rows + industries.rows
|
||||
sector_codes = tuple(row.sector_code for row in indices)
|
||||
members = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.MEMBERS,
|
||||
reusable,
|
||||
lambda: self.source.fetch_sector_members(target, sector_codes),
|
||||
SectorMemberRow.from_mapping,
|
||||
)
|
||||
stock_basics = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.STOCK_BASICS,
|
||||
reusable,
|
||||
self.source.fetch_stock_basics,
|
||||
StockBasicRow.from_mapping,
|
||||
)
|
||||
memberships = normalize_memberships(indices, members)
|
||||
member_codes = tuple(
|
||||
sorted(
|
||||
{
|
||||
item.stock_code
|
||||
for item in memberships
|
||||
if item.status is MembershipStatus.AVAILABLE and item.stock_code is not None
|
||||
}
|
||||
)
|
||||
)
|
||||
current_listed_codes = {
|
||||
row.ts_code for row in stock_basics.rows if is_current_listed_stock(row, target)
|
||||
}
|
||||
moneyflow_candidate_codes = tuple(
|
||||
code for code in member_codes if code in current_listed_codes
|
||||
)
|
||||
suspensions = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.SUSPENSIONS,
|
||||
reusable,
|
||||
lambda: self.source.fetch_suspensions(target),
|
||||
SuspendRow.from_mapping,
|
||||
)
|
||||
daily = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.DAILY,
|
||||
reusable,
|
||||
lambda: self.source.fetch_daily(target),
|
||||
DailyRow.from_mapping,
|
||||
)
|
||||
moneyflow = self._fetch_group(
|
||||
publication_id,
|
||||
PublicationSourceGroup.MONEYFLOW_DC,
|
||||
reusable,
|
||||
lambda: self.source.fetch_moneyflow_dc(target, moneyflow_candidate_codes),
|
||||
MoneyflowDcRow.from_mapping,
|
||||
reuse_if=lambda result: set(moneyflow_candidate_codes).issubset(
|
||||
{row.ts_code for row in result.rows}
|
||||
),
|
||||
)
|
||||
|
||||
stock_facts = normalize_stock_facts(
|
||||
target_trade_date=target,
|
||||
candidate_codes=member_codes,
|
||||
stock_basics=stock_basics,
|
||||
suspensions=suspensions,
|
||||
daily=daily,
|
||||
moneyflow=moneyflow,
|
||||
)
|
||||
snapshots = (
|
||||
calendar.snapshots
|
||||
+ concepts.snapshots
|
||||
+ industries.snapshots
|
||||
+ members.snapshots
|
||||
+ stock_basics.snapshots
|
||||
+ suspensions.snapshots
|
||||
+ daily.snapshots
|
||||
+ moneyflow.snapshots
|
||||
)
|
||||
return _CollectedInputs(
|
||||
target_trade_date=target,
|
||||
snapshots=snapshots,
|
||||
membership_snapshots=members.snapshots,
|
||||
memberships=memberships,
|
||||
stock_facts=stock_facts,
|
||||
)
|
||||
|
||||
def _aggregate(self, inputs: _CollectedInputs) -> tuple[SectorDailyAggregate, ...]:
|
||||
facts = tuple(
|
||||
StockDailyFact(
|
||||
trade_date=item.trade_date,
|
||||
ts_code=item.ts_code,
|
||||
status=item.status,
|
||||
turnover_yuan=item.turnover_yuan,
|
||||
net_amount_yuan=item.net_amount_yuan,
|
||||
)
|
||||
for item in inputs.stock_facts
|
||||
)
|
||||
grouped: defaultdict[tuple[SectorType, str, str], list[str]] = defaultdict(list)
|
||||
unknown: set[tuple[SectorType, str, str]] = set()
|
||||
for member in inputs.memberships:
|
||||
key = (member.sector_type, member.sector_code, member.sector_name)
|
||||
if member.status is MembershipStatus.UNKNOWN:
|
||||
unknown.add(key)
|
||||
continue
|
||||
if member.stock_code is None:
|
||||
raise ValueError("available membership requires a stock code")
|
||||
grouped[key].append(member.stock_code)
|
||||
if unknown & set(grouped):
|
||||
raise ValueError("sector cannot have both available and unknown membership")
|
||||
sector_keys = set(grouped) | unknown
|
||||
aggregates = tuple(
|
||||
aggregate_sector_snapshot(
|
||||
SectorMembershipSnapshot(
|
||||
trade_date=inputs.target_trade_date,
|
||||
sector_type=sector_type,
|
||||
sector_code=sector_code,
|
||||
sector_name=sector_name,
|
||||
member_codes=tuple(sorted(grouped.get(key, ()))),
|
||||
status=(
|
||||
MembershipStatus.UNKNOWN if key in unknown else MembershipStatus.AVAILABLE
|
||||
),
|
||||
source_version=self._universe_version(inputs.membership_snapshots),
|
||||
),
|
||||
facts,
|
||||
)
|
||||
for key in sorted(sector_keys, key=lambda item: (str(item[0]), item[1]))
|
||||
for sector_type, sector_code, sector_name in (key,)
|
||||
)
|
||||
if not aggregates:
|
||||
raise ValueError("sector universe produced no aggregates")
|
||||
return aggregates
|
||||
|
||||
def _rank(
|
||||
self, target: date, aggregates: Sequence[SectorDailyAggregate]
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
history = tuple(self.repository.load_daily_aggregate_history(target, limit_dates=9))
|
||||
observations: list[MetricObservation] = []
|
||||
for current in aggregates:
|
||||
sector_history = tuple(
|
||||
item
|
||||
for item in history
|
||||
if (item.sector_type, item.sector_code)
|
||||
== (current.sector_type, current.sector_code)
|
||||
) + (current,)
|
||||
observations.extend(
|
||||
strategy.evaluate(sector_history, target) for strategy in self.strategies
|
||||
)
|
||||
current_rankings = rank_metric_observations(observations)
|
||||
previous = self.repository.load_previous_rankings(target, limit_dates=5)
|
||||
history_by_days = {days: rankings for days, (_, rankings) in enumerate(previous, start=1)}
|
||||
return with_rank_changes(current_rankings, history_by_days)
|
||||
|
||||
@staticmethod
|
||||
def _coverage(stock_facts: Sequence[StockFactRecord]) -> Decimal:
|
||||
expected_statuses = {
|
||||
"available",
|
||||
"missing",
|
||||
"missing_daily",
|
||||
"missing_moneyflow",
|
||||
"null_daily_amount",
|
||||
"null_moneyflow",
|
||||
"low_liquidity",
|
||||
}
|
||||
expected = sum(item.status.value in expected_statuses for item in stock_facts)
|
||||
covered_statuses = {StockFactStatus.AVAILABLE, StockFactStatus.LOW_LIQUIDITY}
|
||||
covered = sum(item.status in covered_statuses for item in stock_facts)
|
||||
return Decimal(covered) / Decimal(expected) if expected else Decimal(0)
|
||||
|
||||
@staticmethod
|
||||
def _retry_source_groups(
|
||||
memberships: Sequence[MembershipRecord],
|
||||
stock_facts: Sequence[StockFactRecord],
|
||||
) -> tuple[PublicationSourceGroup, ...]:
|
||||
groups: list[PublicationSourceGroup] = []
|
||||
if any(item.status is MembershipStatus.UNKNOWN for item in memberships):
|
||||
groups.append(PublicationSourceGroup.MEMBERS)
|
||||
statuses = {item.status for item in stock_facts}
|
||||
if statuses & {
|
||||
StockFactStatus.MISSING,
|
||||
StockFactStatus.MISSING_DAILY,
|
||||
StockFactStatus.NULL_DAILY_AMOUNT,
|
||||
}:
|
||||
groups.append(PublicationSourceGroup.DAILY)
|
||||
if statuses & {
|
||||
StockFactStatus.MISSING,
|
||||
StockFactStatus.MISSING_MONEYFLOW,
|
||||
StockFactStatus.NULL_MONEYFLOW,
|
||||
}:
|
||||
groups.append(PublicationSourceGroup.MONEYFLOW_DC)
|
||||
return tuple(groups)
|
||||
|
||||
def _input_hash(self, snapshots: Sequence[SourceSnapshot]) -> str:
|
||||
payload = json.dumps(
|
||||
{
|
||||
"snapshot_ids": sorted(snapshot.snapshot_id for snapshot in snapshots),
|
||||
"metric_versions": sorted(strategy.metric_version for strategy in self.strategies),
|
||||
"normalizer": "zhixing_stock_fact_v1",
|
||||
},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return hashlib.sha256(payload.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _universe_version(snapshots: Sequence[SourceSnapshot]) -> str:
|
||||
payload = "\n".join(sorted(snapshot.snapshot_id for snapshot in snapshots))
|
||||
return f"eastmoney-dc-{hashlib.sha256(payload.encode()).hexdigest()[:32]}"
|
||||
|
||||
@staticmethod
|
||||
def _failure_id(target: date) -> str:
|
||||
return f"radar-{target:%Y%m%d}-failed-{uuid4().hex[:24]}"
|
||||
|
||||
@staticmethod
|
||||
def _running_id(target: date) -> str:
|
||||
return f"radar-{target:%Y%m%d}-running-{uuid4().hex[:23]}"
|
||||
|
||||
@staticmethod
|
||||
def _safe_failure(error: BaseException) -> tuple[str, str]:
|
||||
if isinstance(error, ValueError):
|
||||
return type(error).__name__, "input or source contract validation failed"
|
||||
return type(error).__name__, "sector radar build failed"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CollectedInputs:
|
||||
target_trade_date: date
|
||||
snapshots: tuple[SourceSnapshot, ...]
|
||||
membership_snapshots: tuple[SourceSnapshot, ...]
|
||||
memberships: tuple[MembershipRecord, ...]
|
||||
stock_facts: tuple[StockFactRecord, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BuildTarget:
|
||||
trade_date: date
|
||||
retry_publication_id: str | None = None
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Stable read model for persisted sector radar publications."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from enum import StrEnum
|
||||
from typing import Literal
|
||||
|
||||
from ..domain.metrics import (
|
||||
AmountNetStrategy,
|
||||
RatioTurnoverStrategy,
|
||||
SwingEqualThreeToTenStrategy,
|
||||
)
|
||||
from ..domain.models import (
|
||||
MetricKind,
|
||||
MetricUnit,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
from ..domain.persistence import SectorRadarRepository
|
||||
from ..domain.ranking import select_percentile_side, select_rank_change_side
|
||||
|
||||
ReadStatus = Literal["success", "no_data"]
|
||||
|
||||
|
||||
class RadarView(StrEnum):
|
||||
"""Supported ranking projections at the HTTP boundary."""
|
||||
|
||||
AMOUNT = "amount"
|
||||
RATIO = "ratio"
|
||||
SWING = "swing"
|
||||
RANK_CHANGE = "rank_change"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RadarMetricDefinition:
|
||||
"""Public definition of one explicitly independent metric implementation."""
|
||||
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
label: str
|
||||
unit: MetricUnit
|
||||
implementation_kind: Literal["independent"] = "independent"
|
||||
disclaimer: str = "知行独立实现,非 OneChartLab 原站公式"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RadarQuery:
|
||||
"""Validated application query for one ranking page."""
|
||||
|
||||
trade_date: date | None = None
|
||||
sector_type: SectorType = SectorType.CONCEPT
|
||||
view: RadarView = RadarView.AMOUNT
|
||||
rank_change_metric: MetricKind = MetricKind.AMOUNT
|
||||
rank_change_days: int = 1
|
||||
side: RankSide = RankSide.ALL
|
||||
search: str | None = None
|
||||
page: int = 1
|
||||
page_size: int = 20
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject invalid pagination and rank-history offsets outside HTTP usage."""
|
||||
|
||||
if not 1 <= self.rank_change_days <= 5:
|
||||
raise ValueError("rank_change_days must be between 1 and 5")
|
||||
if self.page < 1:
|
||||
raise ValueError("page must be positive")
|
||||
if not 1 <= self.page_size <= 100:
|
||||
raise ValueError("page_size must be between 1 and 100")
|
||||
if self.search is not None and len(self.search) > 100:
|
||||
raise ValueError("search must not exceed 100 characters")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RadarDateIndex:
|
||||
"""Available successful dates plus the newest attempt and strict last-good."""
|
||||
|
||||
available_dates: tuple[date, ...]
|
||||
current_attempt: RadarPublication | None
|
||||
last_good: RadarPublication | None
|
||||
|
||||
@property
|
||||
def status(self) -> ReadStatus:
|
||||
"""Return no_data until at least one successful publication exists."""
|
||||
|
||||
return "success" if self.last_good is not None else "no_data"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RankingPage:
|
||||
"""One filtered page without losing publication or metric provenance."""
|
||||
|
||||
status: ReadStatus
|
||||
query: RadarQuery
|
||||
publication: RadarPublication | None
|
||||
definition: RadarMetricDefinition
|
||||
rows: tuple[RankedMetric, ...]
|
||||
total: int
|
||||
|
||||
|
||||
_METRIC_DEFINITIONS = {
|
||||
MetricKind.AMOUNT: RadarMetricDefinition(
|
||||
metric_kind=MetricKind.AMOUNT,
|
||||
metric_version=AmountNetStrategy.metric_version,
|
||||
label="主力净流入(知行独立实现)",
|
||||
unit=MetricUnit.CNY_100M,
|
||||
),
|
||||
MetricKind.RATIO: RadarMetricDefinition(
|
||||
metric_kind=MetricKind.RATIO,
|
||||
metric_version=RatioTurnoverStrategy.metric_version,
|
||||
label="主力净流入/成交额(知行独立实现)",
|
||||
unit=MetricUnit.RATIO,
|
||||
),
|
||||
MetricKind.SWING: RadarMetricDefinition(
|
||||
metric_kind=MetricKind.SWING,
|
||||
metric_version=SwingEqualThreeToTenStrategy.metric_version,
|
||||
label="3—10 日等权资金率(知行独立实现)",
|
||||
unit=MetricUnit.RATIO,
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class ReadSectorRadar:
|
||||
"""Hide last-good selection, ranking filters, search, and pagination."""
|
||||
|
||||
def __init__(self, repository: SectorRadarRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def list_dates(self) -> RadarDateIndex:
|
||||
"""Return successful dates without promoting partial or failed attempts."""
|
||||
|
||||
return RadarDateIndex(
|
||||
available_dates=tuple(self.repository.list_successful_dates()),
|
||||
current_attempt=self.repository.get_latest_publication(),
|
||||
last_good=self.repository.get_last_good_publication(),
|
||||
)
|
||||
|
||||
def query(self, query: RadarQuery) -> RankingPage:
|
||||
"""Return one deterministic page for ordinary or rank-change views."""
|
||||
|
||||
metric_kind = (
|
||||
query.rank_change_metric
|
||||
if query.view is RadarView.RANK_CHANGE
|
||||
else MetricKind(query.view.value)
|
||||
)
|
||||
definition = _METRIC_DEFINITIONS[metric_kind]
|
||||
publication = (
|
||||
self.repository.get_successful_publication(query.trade_date)
|
||||
if query.trade_date is not None
|
||||
else self.repository.get_last_good_publication()
|
||||
)
|
||||
if publication is None:
|
||||
return RankingPage("no_data", query, None, definition, (), 0)
|
||||
|
||||
metric_rows = tuple(
|
||||
row
|
||||
for row in self.repository.load_rankings(publication.publication_id)
|
||||
if row.observation.sector_type is query.sector_type
|
||||
and row.observation.metric_kind is metric_kind
|
||||
and row.observation.metric_version == definition.metric_version
|
||||
)
|
||||
if query.view is RadarView.RANK_CHANGE:
|
||||
if query.side is RankSide.ALL:
|
||||
selected = tuple(
|
||||
sorted(
|
||||
metric_rows,
|
||||
key=lambda row: (
|
||||
row.rank_change(query.rank_change_days) is None,
|
||||
-(row.rank_change(query.rank_change_days) or 0),
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
else:
|
||||
selected = select_rank_change_side(
|
||||
metric_rows,
|
||||
days=query.rank_change_days,
|
||||
side=query.side,
|
||||
)
|
||||
else:
|
||||
selected = select_percentile_side(metric_rows, query.side)
|
||||
if query.side is not RankSide.BOTTOM:
|
||||
selected = tuple(
|
||||
sorted(
|
||||
selected,
|
||||
key=lambda row: (
|
||||
row.rank_position is None,
|
||||
row.rank_position or 0,
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
search = query.search.strip().casefold() if query.search else ""
|
||||
searched = tuple(
|
||||
row
|
||||
for row in selected
|
||||
if not search
|
||||
or search in row.observation.sector_code.casefold()
|
||||
or search in row.observation.sector_name.casefold()
|
||||
)
|
||||
start = (query.page - 1) * query.page_size
|
||||
return RankingPage(
|
||||
status="success",
|
||||
query=query,
|
||||
publication=publication,
|
||||
definition=definition,
|
||||
rows=searched[start : start + query.page_size],
|
||||
total=len(searched),
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RadarDateIndex",
|
||||
"RadarMetricDefinition",
|
||||
"RadarQuery",
|
||||
"RadarView",
|
||||
"RankingPage",
|
||||
"ReadSectorRadar",
|
||||
]
|
||||
@@ -0,0 +1 @@
|
||||
"""Storage-independent sector radar models and calculation rules."""
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Point-in-time stock fact aggregation for sector radar metrics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from decimal import Decimal
|
||||
|
||||
from .models import (
|
||||
MembershipStatus,
|
||||
SectorDailyAggregate,
|
||||
SectorMembershipSnapshot,
|
||||
StockDailyFact,
|
||||
StockFactStatus,
|
||||
)
|
||||
|
||||
|
||||
def aggregate_sector_snapshot(
|
||||
snapshot: SectorMembershipSnapshot,
|
||||
stock_facts: Iterable[StockDailyFact],
|
||||
) -> SectorDailyAggregate:
|
||||
"""Aggregate only the members recorded in one dated membership snapshot.
|
||||
|
||||
Unknown membership returns an unavailable aggregate and deliberately
|
||||
ignores any supplied stock facts. For known membership, suspended and
|
||||
lifecycle-invalid members are excluded from the expected moneyflow
|
||||
denominator; missing facts remain expected and reduce coverage.
|
||||
|
||||
Args:
|
||||
snapshot: Dated sector identity and point-in-time member codes.
|
||||
stock_facts: Normalized facts that may contain records outside the sector.
|
||||
|
||||
Returns:
|
||||
A yuan-denominated aggregate with explicit membership and moneyflow coverage.
|
||||
|
||||
Raises:
|
||||
ValueError: If member facts have a date mismatch or duplicate stock code.
|
||||
"""
|
||||
|
||||
if snapshot.status is MembershipStatus.UNKNOWN:
|
||||
return SectorDailyAggregate(
|
||||
trade_date=snapshot.trade_date,
|
||||
sector_type=snapshot.sector_type,
|
||||
sector_code=snapshot.sector_code,
|
||||
sector_name=snapshot.sector_name,
|
||||
member_count=0,
|
||||
valid_sample_count=0,
|
||||
net_amount_yuan=None,
|
||||
turnover_yuan=None,
|
||||
membership_coverage=Decimal(0),
|
||||
moneyflow_coverage=Decimal(0),
|
||||
)
|
||||
|
||||
members = set(snapshot.member_codes)
|
||||
facts_by_code: dict[str, StockDailyFact] = {}
|
||||
for fact in stock_facts:
|
||||
if fact.ts_code not in members:
|
||||
continue
|
||||
if fact.trade_date != snapshot.trade_date:
|
||||
raise ValueError("member stock facts must match the snapshot trade_date")
|
||||
if fact.ts_code in facts_by_code:
|
||||
raise ValueError("member stock facts must have unique ts_code values")
|
||||
facts_by_code[fact.ts_code] = fact
|
||||
|
||||
net_amount_total = Decimal(0)
|
||||
turnover_total = Decimal(0)
|
||||
valid_count = 0
|
||||
expected_count = 0
|
||||
for member_code in snapshot.member_codes:
|
||||
fact = facts_by_code.get(member_code)
|
||||
if fact is None or fact.status in {
|
||||
StockFactStatus.MISSING,
|
||||
StockFactStatus.MISSING_DAILY,
|
||||
StockFactStatus.MISSING_MONEYFLOW,
|
||||
StockFactStatus.NULL_DAILY_AMOUNT,
|
||||
StockFactStatus.NULL_MONEYFLOW,
|
||||
StockFactStatus.LOW_LIQUIDITY,
|
||||
}:
|
||||
expected_count += 1
|
||||
elif fact.status is StockFactStatus.AVAILABLE:
|
||||
expected_count += 1
|
||||
net_amount = fact.net_amount_yuan
|
||||
turnover = fact.turnover_yuan
|
||||
if net_amount is None or turnover is None:
|
||||
raise ValueError("available stock facts require both amounts")
|
||||
net_amount_total += net_amount
|
||||
turnover_total += turnover
|
||||
valid_count += 1
|
||||
|
||||
moneyflow_coverage = (
|
||||
Decimal(valid_count) / Decimal(expected_count) if expected_count else Decimal(1)
|
||||
)
|
||||
|
||||
return SectorDailyAggregate(
|
||||
trade_date=snapshot.trade_date,
|
||||
sector_type=snapshot.sector_type,
|
||||
sector_code=snapshot.sector_code,
|
||||
sector_name=snapshot.sector_name,
|
||||
member_count=len(snapshot.member_codes),
|
||||
valid_sample_count=valid_count,
|
||||
net_amount_yuan=net_amount_total if valid_count else None,
|
||||
turnover_yuan=turnover_total if valid_count else None,
|
||||
membership_coverage=Decimal(1),
|
||||
moneyflow_coverage=moneyflow_coverage,
|
||||
)
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Transparent, versioned metric strategies for the independent radar."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import Protocol
|
||||
|
||||
from .models import (
|
||||
MetricKind,
|
||||
MetricObservation,
|
||||
MetricQuality,
|
||||
MetricUnit,
|
||||
SectorDailyAggregate,
|
||||
)
|
||||
|
||||
|
||||
class MetricStrategy(Protocol):
|
||||
"""Calculate one named metric from a sector's point-in-time daily history."""
|
||||
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
unit: MetricUnit
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Return the target date observation without inventing missing inputs."""
|
||||
|
||||
...
|
||||
|
||||
|
||||
def _target_aggregate(
|
||||
history: Iterable[SectorDailyAggregate], target_trade_date: date
|
||||
) -> SectorDailyAggregate:
|
||||
matches = tuple(row for row in history if row.trade_date == target_trade_date)
|
||||
if len(matches) != 1:
|
||||
raise ValueError("history must contain exactly one target-date aggregate")
|
||||
return matches[0]
|
||||
|
||||
|
||||
def _quality(row: SectorDailyAggregate) -> MetricQuality:
|
||||
if row.valid_sample_count < 5 or row.membership_coverage < 1 or row.moneyflow_coverage < 1:
|
||||
return MetricQuality.AVAILABLE_LIMITED_SAMPLE
|
||||
return MetricQuality.AVAILABLE
|
||||
|
||||
|
||||
def _observation(
|
||||
row: SectorDailyAggregate,
|
||||
*,
|
||||
metric_kind: MetricKind,
|
||||
metric_version: str,
|
||||
unit: MetricUnit,
|
||||
value: Decimal | None,
|
||||
quality: MetricQuality | None = None,
|
||||
) -> MetricObservation:
|
||||
return MetricObservation(
|
||||
trade_date=row.trade_date,
|
||||
sector_type=row.sector_type,
|
||||
sector_code=row.sector_code,
|
||||
sector_name=row.sector_name,
|
||||
metric_kind=metric_kind,
|
||||
metric_version=metric_version,
|
||||
implementation_kind="independent",
|
||||
unit=unit,
|
||||
value=value,
|
||||
quality=(
|
||||
MetricQuality.UNAVAILABLE
|
||||
if value is None
|
||||
else quality
|
||||
if quality is not None
|
||||
else _quality(row)
|
||||
),
|
||||
member_count=row.member_count,
|
||||
valid_sample_count=row.valid_sample_count,
|
||||
membership_coverage=row.membership_coverage,
|
||||
moneyflow_coverage=row.moneyflow_coverage,
|
||||
)
|
||||
|
||||
|
||||
class AmountNetStrategy:
|
||||
"""Aggregate main net amount and expose it in hundred-million yuan."""
|
||||
|
||||
metric_kind = MetricKind.AMOUNT
|
||||
metric_version = "zhixing_amount_net_bn_v1"
|
||||
unit = MetricUnit.CNY_100M
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Return the target net amount; missing moneyflow remains unavailable."""
|
||||
|
||||
row = _target_aggregate(history, target_trade_date)
|
||||
value = None if row.net_amount_yuan is None else row.net_amount_yuan / Decimal("100000000")
|
||||
return _observation(
|
||||
row,
|
||||
metric_kind=self.metric_kind,
|
||||
metric_version=self.metric_version,
|
||||
unit=self.unit,
|
||||
value=value,
|
||||
)
|
||||
|
||||
|
||||
class RatioTurnoverStrategy:
|
||||
"""Divide aggregated main net amount by aggregated daily turnover."""
|
||||
|
||||
metric_kind = MetricKind.RATIO
|
||||
metric_version = "zhixing_ratio_turnover_v1"
|
||||
unit = MetricUnit.RATIO
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Return a ratio only when numerator and positive denominator exist."""
|
||||
|
||||
row = _target_aggregate(history, target_trade_date)
|
||||
value = None
|
||||
if (
|
||||
row.net_amount_yuan is not None
|
||||
and row.turnover_yuan is not None
|
||||
and row.turnover_yuan > 0
|
||||
):
|
||||
value = row.net_amount_yuan / row.turnover_yuan
|
||||
return _observation(
|
||||
row,
|
||||
metric_kind=self.metric_kind,
|
||||
metric_version=self.metric_version,
|
||||
unit=self.unit,
|
||||
value=value,
|
||||
)
|
||||
|
||||
|
||||
class SwingEqualThreeToTenStrategy:
|
||||
"""Average transparent 3-to-10-day aggregate turnover ratios equally.
|
||||
|
||||
This strategy is deliberately named as a Zhixing implementation. It does
|
||||
not reproduce or imply OneChartLab's unpublished window weights or score.
|
||||
"""
|
||||
|
||||
metric_kind = MetricKind.SWING
|
||||
metric_version = "zhixing_swing_equal_3_10_v1"
|
||||
unit = MetricUnit.RATIO
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
history: Iterable[SectorDailyAggregate],
|
||||
target_trade_date: date,
|
||||
) -> MetricObservation:
|
||||
"""Calculate eight complete trading-day windows ending at the target."""
|
||||
|
||||
rows = tuple(sorted(history, key=lambda row: row.trade_date))
|
||||
target = _target_aggregate(rows, target_trade_date)
|
||||
eligible = tuple(row for row in rows if row.trade_date <= target_trade_date)
|
||||
if any(
|
||||
(row.sector_type, row.sector_code) != (target.sector_type, target.sector_code)
|
||||
for row in eligible
|
||||
):
|
||||
raise ValueError("history must contain exactly one sector identity")
|
||||
if len({row.trade_date for row in eligible}) != len(eligible):
|
||||
raise ValueError("history must not contain duplicate trade dates")
|
||||
|
||||
value: Decimal | None = None
|
||||
quality: MetricQuality | None = None
|
||||
if len(eligible) >= 10:
|
||||
latest = eligible[-10:]
|
||||
window_ratios: list[Decimal] = []
|
||||
for window_size in range(3, 11):
|
||||
window = latest[-window_size:]
|
||||
net_amount = Decimal(0)
|
||||
turnover = Decimal(0)
|
||||
for row in window:
|
||||
if row.net_amount_yuan is None or row.turnover_yuan is None:
|
||||
break
|
||||
net_amount += row.net_amount_yuan
|
||||
turnover += row.turnover_yuan
|
||||
else:
|
||||
if turnover <= 0:
|
||||
break
|
||||
window_ratios.append(net_amount / turnover)
|
||||
continue
|
||||
break
|
||||
if len(window_ratios) == 8:
|
||||
value = sum(window_ratios, start=Decimal(0)) / Decimal(8)
|
||||
quality = (
|
||||
MetricQuality.AVAILABLE_LIMITED_SAMPLE
|
||||
if any(_quality(row) is not MetricQuality.AVAILABLE for row in latest)
|
||||
else MetricQuality.AVAILABLE
|
||||
)
|
||||
|
||||
return _observation(
|
||||
target,
|
||||
metric_kind=self.metric_kind,
|
||||
metric_version=self.metric_version,
|
||||
unit=self.unit,
|
||||
value=value,
|
||||
quality=quality,
|
||||
)
|
||||
@@ -0,0 +1,319 @@
|
||||
"""Stable domain values for independently produced sector radar metrics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from enum import StrEnum
|
||||
from typing import Literal
|
||||
|
||||
|
||||
class SectorType(StrEnum):
|
||||
"""Independent ranking pools supported by the first radar release."""
|
||||
|
||||
CONCEPT = "concept"
|
||||
INDUSTRY = "industry"
|
||||
|
||||
|
||||
class MembershipStatus(StrEnum):
|
||||
"""Availability of a point-in-time sector membership snapshot."""
|
||||
|
||||
AVAILABLE = "available"
|
||||
UNKNOWN = "membership_unknown"
|
||||
|
||||
|
||||
class StockFactStatus(StrEnum):
|
||||
"""Why one member does or does not contribute to a daily aggregate."""
|
||||
|
||||
AVAILABLE = "available"
|
||||
SUSPENDED = "suspended"
|
||||
MISSING = "missing"
|
||||
MISSING_DAILY = "missing_daily"
|
||||
MISSING_MONEYFLOW = "missing_moneyflow"
|
||||
NULL_DAILY_AMOUNT = "null_daily_amount"
|
||||
NULL_MONEYFLOW = "null_moneyflow"
|
||||
LIFECYCLE_INVALID = "lifecycle_invalid"
|
||||
LOW_LIQUIDITY = "low_liquidity"
|
||||
|
||||
|
||||
class PublicationStatus(StrEnum):
|
||||
"""Immutable build states retained for audit and last-good selection."""
|
||||
|
||||
RUNNING = "running"
|
||||
SUCCESS = "success"
|
||||
PARTIAL = "partial"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class MetricKind(StrEnum):
|
||||
"""User-facing metric families without borrowing private score names."""
|
||||
|
||||
AMOUNT = "amount"
|
||||
RATIO = "ratio"
|
||||
SWING = "swing"
|
||||
|
||||
|
||||
class MetricQuality(StrEnum):
|
||||
"""Whether a metric is usable and whether its sample needs a warning."""
|
||||
|
||||
AVAILABLE = "available"
|
||||
AVAILABLE_LIMITED_SAMPLE = "available_limited_sample"
|
||||
UNAVAILABLE = "unavailable"
|
||||
|
||||
|
||||
class MetricUnit(StrEnum):
|
||||
"""Units exposed by independent metric strategies."""
|
||||
|
||||
CNY_100M = "CNY_100M"
|
||||
RATIO = "ratio"
|
||||
|
||||
|
||||
class RankSide(StrEnum):
|
||||
"""Ordinary percentile views exposed by the ranking read model."""
|
||||
|
||||
TOP = "top"
|
||||
BOTTOM = "bottom"
|
||||
ALL = "all"
|
||||
|
||||
|
||||
def _validate_finite_decimal(value: Decimal | None, field_name: str) -> None:
|
||||
"""Reject non-finite domain values while preserving missing values."""
|
||||
|
||||
if value is not None and not value.is_finite():
|
||||
raise ValueError(f"{field_name} must be finite or None")
|
||||
|
||||
|
||||
def _validate_coverage(value: Decimal, field_name: str) -> None:
|
||||
"""Require a finite fraction in the inclusive zero-to-one range."""
|
||||
|
||||
_validate_finite_decimal(value, field_name)
|
||||
if value < 0 or value > 1:
|
||||
raise ValueError(f"{field_name} must be between 0 and 1")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SectorMembershipSnapshot:
|
||||
"""One sector's membership as observed for exactly one trade date.
|
||||
|
||||
``UNKNOWN`` is an explicit fact: callers must not substitute a current
|
||||
member list when the historical snapshot is unavailable.
|
||||
"""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
member_codes: tuple[str, ...]
|
||||
status: MembershipStatus
|
||||
source_version: str
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate identity, deterministic membership, and unknown semantics."""
|
||||
|
||||
if not self.sector_code.strip():
|
||||
raise ValueError("sector_code must not be empty")
|
||||
if not self.sector_name.strip():
|
||||
raise ValueError("sector_name must not be empty")
|
||||
if not self.source_version.strip():
|
||||
raise ValueError("source_version must not be empty")
|
||||
if any(not code.strip() for code in self.member_codes):
|
||||
raise ValueError("member_codes must not contain empty values")
|
||||
if len(self.member_codes) != len(set(self.member_codes)):
|
||||
raise ValueError("member_codes must be unique")
|
||||
if self.status is MembershipStatus.UNKNOWN and self.member_codes:
|
||||
raise ValueError("unknown membership must not expose member_codes")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StockDailyFact:
|
||||
"""Normalized daily turnover and moneyflow for one member.
|
||||
|
||||
Amounts are expressed in yuan. Available facts require both source
|
||||
values, including an observed zero. Non-available statuses cannot carry
|
||||
amounts because doing so would blur missing, suspended, and lifecycle
|
||||
semantics at the metric boundary.
|
||||
"""
|
||||
|
||||
trade_date: date
|
||||
ts_code: str
|
||||
status: StockFactStatus
|
||||
turnover_yuan: Decimal | None = None
|
||||
net_amount_yuan: Decimal | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject incomplete available facts and hidden non-finite values."""
|
||||
|
||||
if not self.ts_code.strip():
|
||||
raise ValueError("ts_code must not be empty")
|
||||
_validate_finite_decimal(self.turnover_yuan, "turnover_yuan")
|
||||
_validate_finite_decimal(self.net_amount_yuan, "net_amount_yuan")
|
||||
if self.status is StockFactStatus.AVAILABLE:
|
||||
if self.turnover_yuan is None or self.net_amount_yuan is None:
|
||||
raise ValueError("available stock facts require both amounts")
|
||||
if self.turnover_yuan < 0:
|
||||
raise ValueError("turnover_yuan must not be negative")
|
||||
elif self.turnover_yuan is not None or self.net_amount_yuan is not None:
|
||||
raise ValueError("non-available stock facts must not expose amounts")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RadarPublication:
|
||||
"""Traceable identity and lifecycle of one immutable radar build revision."""
|
||||
|
||||
publication_id: str
|
||||
target_trade_date: date
|
||||
status: PublicationStatus
|
||||
source_version: str
|
||||
universe_version: str
|
||||
metric_versions: tuple[str, ...]
|
||||
input_hash: str | None
|
||||
coverage: Decimal
|
||||
started_at: datetime
|
||||
finished_at: datetime | None = None
|
||||
error_summary: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Keep running and terminal lifecycle timestamps internally consistent."""
|
||||
|
||||
if not self.publication_id.strip():
|
||||
raise ValueError("publication_id must not be empty")
|
||||
if not self.source_version.strip() or not self.universe_version.strip():
|
||||
raise ValueError("publication source versions must not be empty")
|
||||
if not self.metric_versions or any(not value.strip() for value in self.metric_versions):
|
||||
raise ValueError("metric_versions must contain named strategies")
|
||||
if len(self.metric_versions) != len(set(self.metric_versions)):
|
||||
raise ValueError("metric_versions must be unique")
|
||||
_validate_coverage(self.coverage, "coverage")
|
||||
if self.started_at.tzinfo is None:
|
||||
raise ValueError("started_at must be timezone-aware")
|
||||
is_running = self.status is PublicationStatus.RUNNING
|
||||
if is_running != (self.finished_at is None):
|
||||
raise ValueError("finished_at must be absent only while publication is running")
|
||||
if self.finished_at is not None:
|
||||
if self.finished_at.tzinfo is None:
|
||||
raise ValueError("finished_at must be timezone-aware")
|
||||
if self.finished_at < self.started_at:
|
||||
raise ValueError("finished_at must not precede started_at")
|
||||
if self.status is PublicationStatus.SUCCESS and self.input_hash is None:
|
||||
raise ValueError("successful publication requires input_hash")
|
||||
if self.input_hash is not None and (
|
||||
len(self.input_hash) != 64
|
||||
or any(character not in "0123456789abcdef" for character in self.input_hash)
|
||||
):
|
||||
raise ValueError("input_hash must be a lowercase SHA-256 hex digest")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SectorDailyAggregate:
|
||||
"""One sector's point-in-time daily facts after source normalization.
|
||||
|
||||
Amounts use yuan so strategies cannot accidentally mix Tushare's
|
||||
``moneyflow_dc.net_amount`` (ten-thousand yuan) with ``daily.amount``
|
||||
(thousand yuan). ``None`` means missing source data; zero remains an
|
||||
observed value.
|
||||
"""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
member_count: int
|
||||
valid_sample_count: int
|
||||
net_amount_yuan: Decimal | None
|
||||
turnover_yuan: Decimal | None
|
||||
membership_coverage: Decimal
|
||||
moneyflow_coverage: Decimal
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate counts, coverage, and finite normalized values."""
|
||||
|
||||
if not self.sector_code.strip():
|
||||
raise ValueError("sector_code must not be empty")
|
||||
if not self.sector_name.strip():
|
||||
raise ValueError("sector_name must not be empty")
|
||||
if self.member_count < 0:
|
||||
raise ValueError("member_count must not be negative")
|
||||
if not 0 <= self.valid_sample_count <= self.member_count:
|
||||
raise ValueError("valid_sample_count must be within member_count")
|
||||
_validate_finite_decimal(self.net_amount_yuan, "net_amount_yuan")
|
||||
_validate_finite_decimal(self.turnover_yuan, "turnover_yuan")
|
||||
_validate_coverage(self.membership_coverage, "membership_coverage")
|
||||
_validate_coverage(self.moneyflow_coverage, "moneyflow_coverage")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MetricObservation:
|
||||
"""One versioned independent metric value ready for cross-sectional ranking."""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
implementation_kind: Literal["independent"]
|
||||
unit: MetricUnit
|
||||
value: Decimal | None
|
||||
quality: MetricQuality
|
||||
member_count: int
|
||||
valid_sample_count: int
|
||||
membership_coverage: Decimal
|
||||
moneyflow_coverage: Decimal
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Keep unavailable and finite-value states internally consistent."""
|
||||
|
||||
_validate_finite_decimal(self.value, "value")
|
||||
if self.value is None and self.quality is not MetricQuality.UNAVAILABLE:
|
||||
raise ValueError("a missing metric value must be unavailable")
|
||||
if self.value is not None and self.quality is MetricQuality.UNAVAILABLE:
|
||||
raise ValueError("an unavailable metric must not expose a value")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RankChange:
|
||||
"""One previous-publication rank delta using past minus current rank."""
|
||||
|
||||
days: int
|
||||
value: int | None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Limit the public comparison window to one through five days."""
|
||||
|
||||
if not 1 <= self.days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RankedMetric:
|
||||
"""A metric observation with its position inside one independent pool."""
|
||||
|
||||
observation: MetricObservation
|
||||
rank_position: int | None
|
||||
rank_percentile: Decimal | None
|
||||
rank_changes: tuple[RankChange, ...] = ()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Require rank position and percentile to be present or absent together."""
|
||||
|
||||
if (self.rank_position is None) != (self.rank_percentile is None):
|
||||
raise ValueError("rank_position and rank_percentile must be paired")
|
||||
if self.rank_position is not None and self.rank_position < 1:
|
||||
raise ValueError("rank_position must be positive")
|
||||
_validate_finite_decimal(self.rank_percentile, "rank_percentile")
|
||||
if self.rank_percentile is not None and not 0 < self.rank_percentile <= 100:
|
||||
raise ValueError("rank_percentile must be within (0, 100]")
|
||||
days = [change.days for change in self.rank_changes]
|
||||
if len(days) != len(set(days)):
|
||||
raise ValueError("rank change days must be unique")
|
||||
|
||||
def rank_change(self, days: int) -> int | None:
|
||||
"""Return one configured rank delta, or ``None`` when history is absent."""
|
||||
|
||||
if not 1 <= days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
return next(
|
||||
(change.value for change in self.rank_changes if change.days == days),
|
||||
None,
|
||||
)
|
||||
@@ -0,0 +1,245 @@
|
||||
"""Normalize typed Tushare rows into point-in-time persisted radar facts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Callable, Sequence
|
||||
from datetime import date
|
||||
|
||||
from .models import MembershipStatus, StockFactStatus
|
||||
from .persistence import MembershipRecord, StockFactRecord
|
||||
from .source import (
|
||||
DailyRow,
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceContractError,
|
||||
SourceResult,
|
||||
StockBasicRow,
|
||||
SuspendRow,
|
||||
)
|
||||
|
||||
|
||||
def normalize_memberships(
|
||||
indices: Sequence[SectorIndexRow],
|
||||
members: SourceResult[SectorMemberRow],
|
||||
) -> tuple[MembershipRecord, ...]:
|
||||
"""Attach each dated member to its sector identity and raw source partition.
|
||||
|
||||
Args:
|
||||
indices: The complete concept or industry universe for one date.
|
||||
members: Validated membership rows plus all raw request snapshots.
|
||||
|
||||
Returns:
|
||||
Deterministically ordered, source-traceable membership records.
|
||||
|
||||
Raises:
|
||||
SourceContractError: If a member references an unknown sector or lacks a snapshot.
|
||||
"""
|
||||
|
||||
index_by_code = {row.sector_code: row for row in indices}
|
||||
if len(index_by_code) != len(indices):
|
||||
raise SourceContractError("sector indices contain duplicate codes")
|
||||
partition_ids = {
|
||||
snapshot.partition_key: snapshot.snapshot_id
|
||||
for snapshot in members.snapshots
|
||||
if snapshot.partition_key not in {None, "all"}
|
||||
}
|
||||
all_snapshot_id = next(
|
||||
(
|
||||
snapshot.snapshot_id
|
||||
for snapshot in members.snapshots
|
||||
if snapshot.partition_key in {None, "all"}
|
||||
),
|
||||
None,
|
||||
)
|
||||
members_by_sector: dict[str, list[SectorMemberRow]] = {code: [] for code in index_by_code}
|
||||
for member in members.rows:
|
||||
index = index_by_code.get(member.sector_code)
|
||||
if index is None:
|
||||
raise SourceContractError("dc_member references a sector outside dc_index")
|
||||
members_by_sector[member.sector_code].append(member)
|
||||
|
||||
records: list[MembershipRecord] = []
|
||||
for sector_code in sorted(index_by_code):
|
||||
index = index_by_code[sector_code]
|
||||
sector_members = members_by_sector[sector_code]
|
||||
snapshot_id = partition_ids.get(sector_code, all_snapshot_id)
|
||||
if snapshot_id is None:
|
||||
raise SourceContractError("sector membership has no source snapshot")
|
||||
if not sector_members:
|
||||
explicit_partition_id = partition_ids.get(sector_code)
|
||||
if explicit_partition_id is None:
|
||||
raise SourceContractError(
|
||||
"missing sector membership requires an explicit empty partition"
|
||||
)
|
||||
records.append(
|
||||
MembershipRecord(
|
||||
source_snapshot_id=explicit_partition_id,
|
||||
trade_date=index.trade_date,
|
||||
sector_type=index.sector_type,
|
||||
sector_code=index.sector_code,
|
||||
sector_name=index.name,
|
||||
stock_code=None,
|
||||
stock_name=None,
|
||||
status=MembershipStatus.UNKNOWN,
|
||||
)
|
||||
)
|
||||
continue
|
||||
for member in sector_members:
|
||||
records.append(
|
||||
MembershipRecord(
|
||||
source_snapshot_id=snapshot_id,
|
||||
trade_date=member.trade_date,
|
||||
sector_type=index.sector_type,
|
||||
sector_code=index.sector_code,
|
||||
sector_name=index.name,
|
||||
stock_code=member.stock_code,
|
||||
stock_name=member.stock_name,
|
||||
status=MembershipStatus.AVAILABLE,
|
||||
)
|
||||
)
|
||||
return tuple(sorted(records, key=lambda item: (item.sector_code, item.membership_key)))
|
||||
|
||||
|
||||
def normalize_stock_facts(
|
||||
*,
|
||||
target_trade_date: date,
|
||||
candidate_codes: Sequence[str],
|
||||
stock_basics: SourceResult[StockBasicRow],
|
||||
suspensions: SourceResult[SuspendRow],
|
||||
daily: SourceResult[DailyRow],
|
||||
moneyflow: SourceResult[MoneyflowDcRow],
|
||||
) -> tuple[StockFactRecord, ...]:
|
||||
"""Build normalized yuan facts without collapsing missing states into zero.
|
||||
|
||||
Args:
|
||||
target_trade_date: Date whose point-in-time lifecycle is evaluated.
|
||||
candidate_codes: Union of stocks in that date's sector memberships.
|
||||
stock_basics: Current ``list_status=L`` Tushare listings.
|
||||
suspensions: Same-date suspend/resume events.
|
||||
daily: Same-date stock turnover rows in source units.
|
||||
moneyflow: Same-date DC main-moneyflow rows in source units.
|
||||
|
||||
Returns:
|
||||
One deterministic fact per candidate code under a content-derived revision.
|
||||
"""
|
||||
|
||||
if len(candidate_codes) != len(set(candidate_codes)):
|
||||
raise ValueError("candidate_codes must be unique")
|
||||
basic_by_code = _unique_index(stock_basics.rows, lambda row: row.ts_code, "stock_basic")
|
||||
daily_by_code = _unique_index(daily.rows, lambda row: row.ts_code, "daily")
|
||||
moneyflow_by_code = _unique_index(moneyflow.rows, lambda row: row.ts_code, "moneyflow_dc")
|
||||
suspended_codes = {
|
||||
row.ts_code
|
||||
for row in suspensions.rows
|
||||
if row.trade_date == target_trade_date and _is_suspend_event(row.suspend_type)
|
||||
}
|
||||
source_snapshot_ids = tuple(
|
||||
sorted(
|
||||
{
|
||||
snapshot.snapshot_id
|
||||
for result in (stock_basics, suspensions, daily, moneyflow)
|
||||
for snapshot in result.snapshots
|
||||
}
|
||||
)
|
||||
)
|
||||
revision_payload = json.dumps(
|
||||
{
|
||||
"target_trade_date": target_trade_date.isoformat(),
|
||||
"source_snapshot_ids": source_snapshot_ids,
|
||||
"normalizer": "zhixing_stock_fact_v1",
|
||||
},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
fact_revision = hashlib.sha256(revision_payload.encode()).hexdigest()
|
||||
|
||||
records: list[StockFactRecord] = []
|
||||
for ts_code in sorted(candidate_codes):
|
||||
basic = basic_by_code.get(ts_code)
|
||||
daily_row = daily_by_code.get(ts_code)
|
||||
moneyflow_row = moneyflow_by_code.get(ts_code)
|
||||
status = StockFactStatus.AVAILABLE
|
||||
turnover_yuan = None
|
||||
net_amount_yuan = None
|
||||
|
||||
if basic is None or not is_current_listed_stock(basic, target_trade_date):
|
||||
status = StockFactStatus.LIFECYCLE_INVALID
|
||||
elif ts_code in suspended_codes and daily_row is None:
|
||||
status = StockFactStatus.SUSPENDED
|
||||
elif daily_row is None:
|
||||
status = StockFactStatus.MISSING_DAILY
|
||||
elif daily_row.amount_thousand_yuan is None:
|
||||
status = StockFactStatus.NULL_DAILY_AMOUNT
|
||||
elif moneyflow_row is None:
|
||||
status = StockFactStatus.MISSING_MONEYFLOW
|
||||
elif moneyflow_row.net_amount_ten_thousand_yuan is None:
|
||||
status = StockFactStatus.NULL_MONEYFLOW
|
||||
elif daily_row.turnover_yuan == 0:
|
||||
status = StockFactStatus.LOW_LIQUIDITY
|
||||
else:
|
||||
turnover_yuan = daily_row.turnover_yuan
|
||||
net_amount_yuan = moneyflow_row.net_amount_yuan
|
||||
|
||||
records.append(
|
||||
StockFactRecord(
|
||||
fact_revision=fact_revision,
|
||||
source_snapshot_ids=source_snapshot_ids,
|
||||
trade_date=target_trade_date,
|
||||
ts_code=ts_code,
|
||||
status=status,
|
||||
turnover_yuan=turnover_yuan,
|
||||
net_amount_yuan=net_amount_yuan,
|
||||
)
|
||||
)
|
||||
return tuple(records)
|
||||
|
||||
|
||||
def is_current_listed_stock(stock: StockBasicRow, target: date) -> bool:
|
||||
"""Return whether one current ``L`` row is an eligible radar security.
|
||||
|
||||
The radar intentionally uses the listings observed at build time rather than
|
||||
reconstructing historical delistings. Code, market, and list-date checks keep
|
||||
the existing Shanghai/Shenzhen A-share boundary intact.
|
||||
|
||||
Args:
|
||||
stock: One validated ``stock_basic`` row.
|
||||
target: Radar date whose list date must already have arrived.
|
||||
|
||||
Returns:
|
||||
Whether the security belongs to the build-time radar universe.
|
||||
"""
|
||||
|
||||
if stock.list_status != "L":
|
||||
return False
|
||||
if not stock.ts_code.endswith((".SH", ".SZ")):
|
||||
return False
|
||||
if stock.symbol.startswith(("200", "900")):
|
||||
return False
|
||||
market = stock.market or ""
|
||||
if "北交" in market or "B股" in market.upper():
|
||||
return False
|
||||
return stock.list_date is not None and stock.list_date <= target
|
||||
|
||||
|
||||
def _is_suspend_event(value: str) -> bool:
|
||||
normalized = value.strip().casefold()
|
||||
return normalized in {"s", "suspend", "停牌"} or (
|
||||
"停牌" in normalized and "复牌" not in normalized
|
||||
)
|
||||
|
||||
|
||||
def _unique_index[T, K](
|
||||
rows: Sequence[T],
|
||||
key: Callable[[T], K],
|
||||
source_name: str,
|
||||
) -> dict[K, T]:
|
||||
result: dict[K, T] = {}
|
||||
for row in rows:
|
||||
item_key = key(row)
|
||||
if item_key in result:
|
||||
raise SourceContractError(f"{source_name} contains duplicate business keys")
|
||||
result[item_key] = row
|
||||
return result
|
||||
@@ -0,0 +1,251 @@
|
||||
"""Persistence records and repository port for replayable radar revisions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Sequence
|
||||
from contextlib import AbstractContextManager
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from enum import StrEnum
|
||||
from typing import Protocol
|
||||
|
||||
from .models import (
|
||||
MembershipStatus,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
SectorDailyAggregate,
|
||||
SectorType,
|
||||
StockFactStatus,
|
||||
)
|
||||
from .source import SourceSnapshot
|
||||
|
||||
|
||||
def _validate_digest(value: str, field_name: str) -> None:
|
||||
if len(value) != 64 or any(character not in "0123456789abcdef" for character in value):
|
||||
raise ValueError(f"{field_name} must be a lowercase SHA-256 digest")
|
||||
|
||||
|
||||
def _validate_optional_decimal(value: Decimal | None, field_name: str) -> None:
|
||||
if value is not None and not value.is_finite():
|
||||
raise ValueError(f"{field_name} must be finite or None")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MembershipRecord:
|
||||
"""One persisted point-in-time member or explicit unknown snapshot."""
|
||||
|
||||
source_snapshot_id: str
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
stock_code: str | None
|
||||
stock_name: str | None
|
||||
status: MembershipStatus = MembershipStatus.AVAILABLE
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate available and unknown membership null semantics."""
|
||||
|
||||
_validate_digest(self.source_snapshot_id, "source_snapshot_id")
|
||||
if not self.sector_code.strip() or not self.sector_name.strip():
|
||||
raise ValueError("membership sector identity fields must not be empty")
|
||||
if self.status is MembershipStatus.AVAILABLE:
|
||||
if self.stock_code is None or self.stock_name is None:
|
||||
raise ValueError("available membership requires stock identity")
|
||||
if not self.stock_code.strip() or not self.stock_name.strip():
|
||||
raise ValueError("available membership stock identity must not be empty")
|
||||
elif self.stock_code is not None or self.stock_name is not None:
|
||||
raise ValueError("unknown membership must not expose stock identity")
|
||||
|
||||
@property
|
||||
def membership_key(self) -> str:
|
||||
"""Return a non-null persistence key without inventing a stock code."""
|
||||
|
||||
return self.stock_code if self.stock_code is not None else "__membership_unknown__"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StockFactRecord:
|
||||
"""One normalized stock fact revision with all contributing raw snapshots."""
|
||||
|
||||
fact_revision: str
|
||||
source_snapshot_ids: tuple[str, ...]
|
||||
trade_date: date
|
||||
ts_code: str
|
||||
status: StockFactStatus
|
||||
turnover_yuan: Decimal | None = None
|
||||
net_amount_yuan: Decimal | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Preserve source traceability and stock fact null semantics."""
|
||||
|
||||
_validate_digest(self.fact_revision, "fact_revision")
|
||||
if not self.source_snapshot_ids or len(self.source_snapshot_ids) != len(
|
||||
set(self.source_snapshot_ids)
|
||||
):
|
||||
raise ValueError("source_snapshot_ids must be non-empty and unique")
|
||||
for value in self.source_snapshot_ids:
|
||||
_validate_digest(value, "source_snapshot_id")
|
||||
if not self.ts_code.strip():
|
||||
raise ValueError("ts_code must not be empty")
|
||||
_validate_optional_decimal(self.turnover_yuan, "turnover_yuan")
|
||||
_validate_optional_decimal(self.net_amount_yuan, "net_amount_yuan")
|
||||
if self.status is StockFactStatus.AVAILABLE:
|
||||
if self.turnover_yuan is None or self.net_amount_yuan is None:
|
||||
raise ValueError("available stock facts require both amounts")
|
||||
if self.turnover_yuan < 0:
|
||||
raise ValueError("turnover_yuan must not be negative")
|
||||
elif self.turnover_yuan is not None or self.net_amount_yuan is not None:
|
||||
raise ValueError("non-available stock facts must not expose amounts")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RankingRecord:
|
||||
"""One ranked metric attached to an immutable publication identity."""
|
||||
|
||||
publication_id: str
|
||||
ranking: RankedMetric
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the publication foreign identity."""
|
||||
|
||||
if not self.publication_id.strip():
|
||||
raise ValueError("publication_id must not be empty")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DailyAggregateRecord:
|
||||
"""One exact daily strategy input owned by a publication revision."""
|
||||
|
||||
publication_id: str
|
||||
aggregate: SectorDailyAggregate
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the publication foreign identity."""
|
||||
|
||||
if not self.publication_id.strip():
|
||||
raise ValueError("publication_id must not be empty")
|
||||
|
||||
|
||||
class PublicationSourceGroup(StrEnum):
|
||||
"""Stable source checkpoints that can be retried independently."""
|
||||
|
||||
CALENDAR = "calendar"
|
||||
CONCEPT_INDICES = "concept_indices"
|
||||
INDUSTRY_INDICES = "industry_indices"
|
||||
MEMBERS = "members"
|
||||
STOCK_BASICS = "stock_basics"
|
||||
SUSPENSIONS = "suspensions"
|
||||
DAILY = "daily"
|
||||
MONEYFLOW_DC = "moneyflow_dc"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PublicationSourceRecord:
|
||||
"""One ordered raw snapshot checkpoint attached to a build attempt."""
|
||||
|
||||
publication_id: str
|
||||
source_group: PublicationSourceGroup
|
||||
source_order: int
|
||||
snapshot: SourceSnapshot
|
||||
refresh_on_retry: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate the publication identity and deterministic group ordering."""
|
||||
|
||||
if not self.publication_id.strip():
|
||||
raise ValueError("publication_id must not be empty")
|
||||
if self.source_order < 0:
|
||||
raise ValueError("source_order must not be negative")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WriteCounts:
|
||||
"""Idempotent persistence outcome."""
|
||||
|
||||
inserted: int
|
||||
unchanged: int
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject impossible write counts."""
|
||||
|
||||
if self.inserted < 0 or self.unchanged < 0:
|
||||
raise ValueError("write counts must not be negative")
|
||||
|
||||
|
||||
class SectorRadarRepository(Protocol):
|
||||
"""Persist source revisions, normalized facts, and published rankings."""
|
||||
|
||||
def advisory_lock(self, target_trade_date: date) -> AbstractContextManager[bool]: ...
|
||||
|
||||
def save_source_snapshots(self, snapshots: Iterable[SourceSnapshot]) -> WriteCounts: ...
|
||||
|
||||
def save_publication_sources(
|
||||
self, records: Iterable[PublicationSourceRecord]
|
||||
) -> WriteCounts: ...
|
||||
|
||||
def load_publication_sources(
|
||||
self, publication_id: str
|
||||
) -> Sequence[PublicationSourceRecord]: ...
|
||||
|
||||
def mark_publication_sources_for_retry(
|
||||
self,
|
||||
publication_id: str,
|
||||
source_groups: Sequence[PublicationSourceGroup],
|
||||
) -> None: ...
|
||||
|
||||
def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts: ...
|
||||
|
||||
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts: ...
|
||||
|
||||
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts: ...
|
||||
|
||||
def create_publication(self, publication: RadarPublication) -> WriteCounts: ...
|
||||
|
||||
def finish_publication(self, publication: RadarPublication) -> None: ...
|
||||
|
||||
def finalize_publication(
|
||||
self,
|
||||
publication: RadarPublication,
|
||||
*,
|
||||
memberships: Iterable[MembershipRecord],
|
||||
stock_facts: Iterable[StockFactRecord],
|
||||
daily_aggregates: Iterable[DailyAggregateRecord],
|
||||
rankings: Iterable[RankingRecord],
|
||||
retry_source_groups: Sequence[PublicationSourceGroup] = (),
|
||||
) -> None: ...
|
||||
|
||||
def recover_running_publications(
|
||||
self, target_trade_date: date, *, finished_at: datetime
|
||||
) -> Sequence[str]: ...
|
||||
|
||||
def discard_running_publication(self, publication_id: str) -> None: ...
|
||||
|
||||
def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts: ...
|
||||
|
||||
def find_reusable_publication(
|
||||
self, target_trade_date: date, input_hash: str
|
||||
) -> RadarPublication | None: ...
|
||||
|
||||
def get_publication(self, publication_id: str) -> RadarPublication | None: ...
|
||||
|
||||
def get_last_good_publication(
|
||||
self, target_trade_date: date | None = None
|
||||
) -> RadarPublication | None: ...
|
||||
|
||||
def get_successful_publication(self, target_trade_date: date) -> RadarPublication | None: ...
|
||||
|
||||
def get_latest_publication(self) -> RadarPublication | None: ...
|
||||
|
||||
def list_successful_dates(self) -> Sequence[date]: ...
|
||||
|
||||
def load_rankings(self, publication_id: str) -> Sequence[RankedMetric]: ...
|
||||
|
||||
def load_daily_aggregate_history(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[SectorDailyAggregate]: ...
|
||||
|
||||
def load_previous_rankings(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[tuple[date, Sequence[RankedMetric]]]: ...
|
||||
@@ -0,0 +1,57 @@
|
||||
"""Application-facing ports for independent sector radar production."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from contextlib import AbstractContextManager
|
||||
from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from .models import SectorType
|
||||
from .source import (
|
||||
CapabilityProbeResult,
|
||||
DailyRow,
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceResult,
|
||||
StockBasicRow,
|
||||
SuspendRow,
|
||||
TradeCalendarRow,
|
||||
)
|
||||
|
||||
|
||||
class SectorRadarSource(Protocol):
|
||||
"""Fetch the minimum replayable Tushare facts needed by the MVP."""
|
||||
|
||||
def fetch_trade_calendar(self, start: date, end: date) -> SourceResult[TradeCalendarRow]: ...
|
||||
|
||||
def fetch_sector_indices(
|
||||
self, trade_date: date, sector_type: SectorType
|
||||
) -> SourceResult[SectorIndexRow]: ...
|
||||
|
||||
def fetch_sector_members(
|
||||
self,
|
||||
trade_date: date,
|
||||
sector_codes: Sequence[str],
|
||||
) -> SourceResult[SectorMemberRow]: ...
|
||||
|
||||
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]: ...
|
||||
|
||||
def fetch_suspensions(self, trade_date: date) -> SourceResult[SuspendRow]: ...
|
||||
|
||||
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]: ...
|
||||
|
||||
def fetch_moneyflow_dc(
|
||||
self,
|
||||
trade_date: date,
|
||||
candidate_codes: Sequence[str],
|
||||
) -> SourceResult[MoneyflowDcRow]: ...
|
||||
|
||||
def probe(self, trade_date: date) -> CapabilityProbeResult: ...
|
||||
|
||||
|
||||
class SectorRadarLock(Protocol):
|
||||
"""Repository seam for a target-date advisory lock."""
|
||||
|
||||
def advisory_lock(self, target_trade_date: date) -> AbstractContextManager[bool]: ...
|
||||
@@ -0,0 +1,216 @@
|
||||
"""Deterministic cross-sectional ranking for independent sector pools."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import defaultdict
|
||||
from collections.abc import Iterable, Mapping
|
||||
from dataclasses import replace
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from .models import (
|
||||
MetricKind,
|
||||
MetricObservation,
|
||||
RankChange,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
|
||||
PoolKey = tuple[date, SectorType, MetricKind, str]
|
||||
SectorMetricKey = tuple[SectorType, str, MetricKind, str]
|
||||
|
||||
|
||||
def _pool_key(observation: MetricObservation) -> PoolKey:
|
||||
return (
|
||||
observation.trade_date,
|
||||
observation.sector_type,
|
||||
observation.metric_kind,
|
||||
observation.metric_version,
|
||||
)
|
||||
|
||||
|
||||
def _sector_metric_key(observation: MetricObservation) -> SectorMetricKey:
|
||||
return (
|
||||
observation.sector_type,
|
||||
observation.sector_code,
|
||||
observation.metric_kind,
|
||||
observation.metric_version,
|
||||
)
|
||||
|
||||
|
||||
def _available_sort_key(observation: MetricObservation) -> tuple[Decimal, str]:
|
||||
if observation.value is None:
|
||||
raise ValueError("unavailable observations cannot use the ranking sort key")
|
||||
return (-observation.value, observation.sector_code)
|
||||
|
||||
|
||||
def rank_metric_observations(
|
||||
observations: Iterable[MetricObservation],
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Rank observations by value within date, type, metric, and version.
|
||||
|
||||
Concept and industry observations never share a pool. Equal metric values
|
||||
use ascending sector code as the documented Zhixing tie-breaker. Missing
|
||||
values remain visible but do not consume a rank.
|
||||
"""
|
||||
|
||||
pools: defaultdict[PoolKey, list[MetricObservation]] = defaultdict(list)
|
||||
for observation in observations:
|
||||
pools[_pool_key(observation)].append(observation)
|
||||
|
||||
result: list[RankedMetric] = []
|
||||
for pool_key in sorted(
|
||||
pools,
|
||||
key=lambda key: (key[0], key[1].value, key[2].value, key[3]),
|
||||
):
|
||||
pool = pools[pool_key]
|
||||
codes = [observation.sector_code for observation in pool]
|
||||
if len(codes) != len(set(codes)):
|
||||
raise ValueError("a ranking pool must not contain duplicate sector codes")
|
||||
|
||||
available = sorted(
|
||||
(observation for observation in pool if observation.value is not None),
|
||||
key=_available_sort_key,
|
||||
)
|
||||
pool_size = len(available)
|
||||
for rank_position, observation in enumerate(available, start=1):
|
||||
rank_percentile = (
|
||||
Decimal(100) * Decimal(pool_size - rank_position + 1) / Decimal(pool_size)
|
||||
)
|
||||
result.append(
|
||||
RankedMetric(
|
||||
observation=observation,
|
||||
rank_position=rank_position,
|
||||
rank_percentile=rank_percentile,
|
||||
)
|
||||
)
|
||||
|
||||
result.extend(
|
||||
RankedMetric(
|
||||
observation=observation,
|
||||
rank_position=None,
|
||||
rank_percentile=None,
|
||||
)
|
||||
for observation in sorted(
|
||||
(observation for observation in pool if observation.value is None),
|
||||
key=lambda observation: observation.sector_code,
|
||||
)
|
||||
)
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def select_percentile_side(
|
||||
rankings: Iterable[RankedMetric], side: RankSide
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Select confirmed inclusive percentile sides without fixed row counts."""
|
||||
|
||||
rows = tuple(rankings)
|
||||
if side is RankSide.ALL:
|
||||
return rows
|
||||
|
||||
threshold_rows = tuple(
|
||||
row
|
||||
for row in rows
|
||||
if row.rank_percentile is not None
|
||||
and (
|
||||
row.rank_percentile >= Decimal(90)
|
||||
if side is RankSide.TOP
|
||||
else row.rank_percentile <= Decimal(10)
|
||||
)
|
||||
)
|
||||
if side is RankSide.TOP:
|
||||
return threshold_rows
|
||||
|
||||
pools: defaultdict[PoolKey, list[RankedMetric]] = defaultdict(list)
|
||||
for row in threshold_rows:
|
||||
pools[_pool_key(row.observation)].append(row)
|
||||
result: list[RankedMetric] = []
|
||||
for pool_key in sorted(
|
||||
pools,
|
||||
key=lambda key: (key[0], key[1].value, key[2].value, key[3]),
|
||||
):
|
||||
result.extend(
|
||||
sorted(
|
||||
pools[pool_key],
|
||||
key=lambda row: (
|
||||
row.observation.value if row.observation.value is not None else Decimal(0),
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def with_rank_changes(
|
||||
current_rankings: Iterable[RankedMetric],
|
||||
history_by_days: Mapping[int, Iterable[RankedMetric]],
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Attach 1-to-5-day deltas without turning missing history into zero."""
|
||||
|
||||
history_indexes: dict[int, dict[SectorMetricKey, int | None]] = {}
|
||||
for days, historical_rankings in history_by_days.items():
|
||||
if not 1 <= days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
index: dict[SectorMetricKey, int | None] = {}
|
||||
for row in historical_rankings:
|
||||
key = _sector_metric_key(row.observation)
|
||||
if key in index:
|
||||
raise ValueError("historical rankings must have unique sector metrics")
|
||||
index[key] = row.rank_position
|
||||
history_indexes[days] = index
|
||||
|
||||
result: list[RankedMetric] = []
|
||||
for row in current_rankings:
|
||||
key = _sector_metric_key(row.observation)
|
||||
changes: list[RankChange] = []
|
||||
for days in sorted(history_indexes):
|
||||
past_rank = history_indexes[days].get(key)
|
||||
value = (
|
||||
past_rank - row.rank_position
|
||||
if past_rank is not None and row.rank_position is not None
|
||||
else None
|
||||
)
|
||||
changes.append(RankChange(days=days, value=value))
|
||||
result.append(replace(row, rank_changes=tuple(changes)))
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def select_rank_change_side(
|
||||
rankings: Iterable[RankedMetric],
|
||||
*,
|
||||
days: int,
|
||||
side: RankSide,
|
||||
) -> tuple[RankedMetric, ...]:
|
||||
"""Select the strongest or weakest ceiling-ten-percent rank changes per pool."""
|
||||
|
||||
if not 1 <= days <= 5:
|
||||
raise ValueError("rank change days must be between 1 and 5")
|
||||
pools: defaultdict[PoolKey, list[RankedMetric]] = defaultdict(list)
|
||||
for row in rankings:
|
||||
pools[_pool_key(row.observation)].append(row)
|
||||
|
||||
result: list[RankedMetric] = []
|
||||
for pool_key in sorted(
|
||||
pools,
|
||||
key=lambda key: (key[0], key[1].value, key[2].value, key[3]),
|
||||
):
|
||||
pool = pools[pool_key]
|
||||
pool_size = sum(row.rank_position is not None for row in pool)
|
||||
take_count = max(1, (pool_size + 9) // 10) if pool_size else 0
|
||||
candidates = tuple(
|
||||
(change, row) for row in pool if (change := row.rank_change(days)) is not None
|
||||
)
|
||||
if side is RankSide.BOTTOM:
|
||||
ordered = sorted(
|
||||
candidates,
|
||||
key=lambda item: (item[0], item[1].observation.sector_code),
|
||||
)
|
||||
else:
|
||||
ordered = sorted(
|
||||
candidates,
|
||||
key=lambda item: (-item[0], item[1].observation.sector_code),
|
||||
)
|
||||
selected = ordered if side is RankSide.ALL else ordered[:take_count]
|
||||
result.extend(row for _, row in selected)
|
||||
return tuple(result)
|
||||
@@ -0,0 +1,503 @@
|
||||
"""Typed Tushare input contracts and replayable source snapshot values."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, date, datetime
|
||||
from decimal import Decimal, InvalidOperation
|
||||
from enum import StrEnum
|
||||
from typing import TypeVar
|
||||
|
||||
from .models import SectorType
|
||||
|
||||
SourceScalar = str | int | float | bool | None
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class SourceContractError(ValueError):
|
||||
"""A provider response violates the replayable input contract.
|
||||
|
||||
Messages must remain operator-safe contract descriptions. They may name a
|
||||
field or validation rule, but must never interpolate provider values,
|
||||
credentials, request parameters, or raw payloads because production build
|
||||
diagnostics record this message.
|
||||
"""
|
||||
|
||||
def __init__(self, message: str) -> None:
|
||||
"""Create one violation whose diagnostic can be claimed by the nearest boundary."""
|
||||
|
||||
super().__init__(message)
|
||||
self._diagnostic_claimed = False
|
||||
|
||||
@property
|
||||
def operator_message(self) -> str:
|
||||
"""Return a single-line, bounded diagnostic suitable for production logs."""
|
||||
|
||||
return " ".join(str(self).split())[:200]
|
||||
|
||||
def claim_diagnostic(self) -> bool:
|
||||
"""Return whether this boundary should emit the exception's single diagnostic log."""
|
||||
|
||||
if self._diagnostic_claimed:
|
||||
return False
|
||||
self._diagnostic_claimed = True
|
||||
return True
|
||||
|
||||
|
||||
class SourceTruncatedError(SourceContractError):
|
||||
"""A provider response reached its row limit without safe partitioning."""
|
||||
|
||||
|
||||
def normalize_source_scalar(value: object) -> SourceScalar:
|
||||
"""Normalize flat Tushare cells while distinguishing missing from infinity."""
|
||||
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
if isinstance(value, float):
|
||||
if math.isnan(value):
|
||||
return None
|
||||
if not math.isfinite(value):
|
||||
raise SourceContractError("source numeric values must be finite")
|
||||
return value
|
||||
if isinstance(value, Decimal):
|
||||
if value.is_nan():
|
||||
return None
|
||||
if not value.is_finite():
|
||||
raise SourceContractError("source numeric values must be finite")
|
||||
return str(value)
|
||||
if isinstance(value, datetime):
|
||||
return value.isoformat()
|
||||
if isinstance(value, date):
|
||||
return value.isoformat()
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if not stripped or stripped.casefold() == "nan":
|
||||
return None
|
||||
return stripped
|
||||
raise SourceContractError(f"unsupported source cell type: {type(value).__name__}")
|
||||
|
||||
|
||||
def normalize_source_rows(
|
||||
rows: Sequence[Mapping[str, object]],
|
||||
) -> tuple[dict[str, SourceScalar], ...]:
|
||||
"""Return safe flat rows with deterministic key order."""
|
||||
|
||||
return tuple({key: normalize_source_scalar(row[key]) for key in sorted(row)} for row in rows)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SourceSnapshot:
|
||||
"""One raw, sanitized provider response identified by safe content hash."""
|
||||
|
||||
snapshot_id: str
|
||||
api_name: str
|
||||
normalized_params: tuple[tuple[str, str], ...]
|
||||
target_trade_date: date | None
|
||||
partition_key: str | None
|
||||
observed_at: datetime
|
||||
rows: tuple[dict[str, SourceScalar], ...]
|
||||
row_count: int
|
||||
returned_fields: tuple[str, ...]
|
||||
content_sha256: str
|
||||
row_limit: int | None
|
||||
limit_reached: bool
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate replay identity and row metadata."""
|
||||
|
||||
for field_name, value in (
|
||||
("snapshot_id", self.snapshot_id),
|
||||
("content_sha256", self.content_sha256),
|
||||
):
|
||||
if len(value) != 64 or any(character not in "0123456789abcdef" for character in value):
|
||||
raise ValueError(f"{field_name} must be a lowercase SHA-256 digest")
|
||||
if not self.api_name.strip():
|
||||
raise ValueError("api_name must not be empty")
|
||||
if self.observed_at.tzinfo is None:
|
||||
raise ValueError("observed_at must be timezone-aware")
|
||||
if self.row_count != len(self.rows):
|
||||
raise ValueError("row_count must match rows")
|
||||
if self.row_limit is not None and self.row_limit < 1:
|
||||
raise ValueError("row_limit must be positive")
|
||||
if self.limit_reached != (self.row_limit is not None and self.row_count >= self.row_limit):
|
||||
raise ValueError("limit_reached must match row_count and row_limit")
|
||||
|
||||
|
||||
def build_source_snapshot(
|
||||
*,
|
||||
api_name: str,
|
||||
params: Mapping[str, object],
|
||||
rows: Sequence[Mapping[str, object]],
|
||||
target_trade_date: date | None,
|
||||
partition_key: str | None = None,
|
||||
observed_at: datetime | None = None,
|
||||
row_limit: int | None = None,
|
||||
returned_fields: Sequence[str] | None = None,
|
||||
) -> SourceSnapshot:
|
||||
"""Build an order-stable, token-free raw response snapshot."""
|
||||
|
||||
normalized_rows = normalize_source_rows(rows)
|
||||
normalized_params = tuple(
|
||||
sorted((key, str(value)) for key, value in params.items() if key != "token")
|
||||
)
|
||||
fields = tuple(
|
||||
sorted(
|
||||
set(returned_fields)
|
||||
if returned_fields is not None
|
||||
else {key for row in normalized_rows for key in row}
|
||||
)
|
||||
)
|
||||
row_count = len(normalized_rows)
|
||||
limit_reached = row_limit is not None and row_count >= row_limit
|
||||
canonical_rows = sorted(
|
||||
json.dumps(row, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||
for row in normalized_rows
|
||||
)
|
||||
canonical_content = json.dumps(
|
||||
{
|
||||
"rows": canonical_rows,
|
||||
"returned_fields": fields,
|
||||
"row_limit": row_limit,
|
||||
"limit_reached": limit_reached,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
content_sha256 = hashlib.sha256(canonical_content.encode()).hexdigest()
|
||||
identity = json.dumps(
|
||||
{
|
||||
"api_name": api_name,
|
||||
"params": normalized_params,
|
||||
"partition_key": partition_key,
|
||||
"target_trade_date": (
|
||||
target_trade_date.isoformat() if target_trade_date is not None else None
|
||||
),
|
||||
"content_sha256": content_sha256,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
snapshot_id = hashlib.sha256(identity.encode()).hexdigest()
|
||||
return SourceSnapshot(
|
||||
snapshot_id=snapshot_id,
|
||||
api_name=api_name,
|
||||
normalized_params=normalized_params,
|
||||
target_trade_date=target_trade_date,
|
||||
partition_key=partition_key,
|
||||
observed_at=observed_at or datetime.now(UTC),
|
||||
rows=normalized_rows,
|
||||
row_count=row_count,
|
||||
returned_fields=fields,
|
||||
content_sha256=content_sha256,
|
||||
row_limit=row_limit,
|
||||
limit_reached=limit_reached,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SourceResult[T]:
|
||||
"""Typed rows accompanied by every raw request needed to produce them."""
|
||||
|
||||
snapshots: tuple[SourceSnapshot, ...]
|
||||
rows: tuple[T, ...]
|
||||
|
||||
|
||||
def _required_text(row: Mapping[str, SourceScalar], key: str) -> str:
|
||||
value = row.get(key)
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise SourceContractError(f"{key} must be a non-empty string")
|
||||
return value.strip()
|
||||
|
||||
|
||||
def _optional_text(row: Mapping[str, SourceScalar], key: str) -> str | None:
|
||||
value = row.get(key)
|
||||
if value is None:
|
||||
return None
|
||||
return str(value).strip() or None
|
||||
|
||||
|
||||
def _source_date(
|
||||
row: Mapping[str, SourceScalar], key: str, *, required: bool = True
|
||||
) -> date | None:
|
||||
value = row.get(key)
|
||||
if value is None:
|
||||
if required:
|
||||
raise SourceContractError(f"{key} is required")
|
||||
return None
|
||||
text = str(value).strip().replace("-", "")
|
||||
try:
|
||||
return datetime.strptime(text, "%Y%m%d").date()
|
||||
except ValueError as exc:
|
||||
raise SourceContractError(f"{key} must use YYYYMMDD") from exc
|
||||
|
||||
|
||||
def _decimal(row: Mapping[str, SourceScalar], key: str) -> Decimal | None:
|
||||
value = row.get(key)
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
result = Decimal(str(value))
|
||||
except InvalidOperation as exc:
|
||||
raise SourceContractError(f"{key} must be numeric or missing") from exc
|
||||
if result.is_nan():
|
||||
return None
|
||||
if not result.is_finite():
|
||||
raise SourceContractError(f"{key} must be finite")
|
||||
return result
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TradeCalendarRow:
|
||||
"""One exchange calendar observation."""
|
||||
|
||||
exchange: str
|
||||
cal_date: date
|
||||
is_open: bool
|
||||
pretrade_date: date | None
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> TradeCalendarRow:
|
||||
"""Parse one Tushare ``trade_cal`` row."""
|
||||
|
||||
cal_date = _source_date(row, "cal_date")
|
||||
assert cal_date is not None
|
||||
return cls(
|
||||
exchange=_optional_text(row, "exchange") or "",
|
||||
cal_date=cal_date,
|
||||
is_open=str(row.get("is_open")).strip().casefold() in {"1", "true"},
|
||||
pretrade_date=_source_date(row, "pretrade_date", required=False),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SectorIndexRow:
|
||||
"""One Eastmoney concept or industry identity on a trade date."""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
name: str
|
||||
level: str | None
|
||||
pct_change: Decimal | None
|
||||
leading_code: str | None
|
||||
|
||||
@classmethod
|
||||
def from_mapping(
|
||||
cls,
|
||||
row: Mapping[str, SourceScalar],
|
||||
sector_type: SectorType,
|
||||
) -> SectorIndexRow:
|
||||
"""Parse and validate one ``dc_index`` row."""
|
||||
|
||||
trade_date = _source_date(row, "trade_date")
|
||||
assert trade_date is not None
|
||||
return cls(
|
||||
trade_date=trade_date,
|
||||
sector_type=sector_type,
|
||||
sector_code=_required_text(row, "ts_code"),
|
||||
name=_required_text(row, "name"),
|
||||
level=_optional_text(row, "level"),
|
||||
pct_change=_decimal(row, "pct_change"),
|
||||
leading_code=_optional_text(row, "leading_code"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SectorMemberRow:
|
||||
"""One point-in-time sector member returned by ``dc_member``."""
|
||||
|
||||
trade_date: date
|
||||
sector_code: str
|
||||
stock_code: str
|
||||
stock_name: str
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> SectorMemberRow:
|
||||
"""Parse one dated membership row."""
|
||||
|
||||
trade_date = _source_date(row, "trade_date")
|
||||
assert trade_date is not None
|
||||
return cls(
|
||||
trade_date=trade_date,
|
||||
sector_code=_required_text(row, "ts_code"),
|
||||
stock_code=_required_text(row, "con_code"),
|
||||
stock_name=_required_text(row, "name"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StockBasicRow:
|
||||
"""Lifecycle and market identity from one explicit listing-status query."""
|
||||
|
||||
ts_code: str
|
||||
symbol: str
|
||||
name: str
|
||||
market: str | None
|
||||
exchange: str
|
||||
list_status: str
|
||||
list_date: date | None
|
||||
delist_date: date | None
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> StockBasicRow:
|
||||
"""Parse one ``stock_basic`` row without applying ST filtering."""
|
||||
|
||||
return cls(
|
||||
ts_code=_required_text(row, "ts_code"),
|
||||
symbol=_required_text(row, "symbol"),
|
||||
name=_required_text(row, "name"),
|
||||
market=_optional_text(row, "market"),
|
||||
exchange=_required_text(row, "exchange"),
|
||||
list_status=_required_text(row, "list_status"),
|
||||
list_date=_source_date(row, "list_date", required=False),
|
||||
delist_date=_source_date(row, "delist_date", required=False),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SuspendRow:
|
||||
"""One daily suspend/resume event."""
|
||||
|
||||
ts_code: str
|
||||
trade_date: date
|
||||
suspend_timing: str | None
|
||||
suspend_type: str
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> SuspendRow:
|
||||
"""Parse one ``suspend_d`` row."""
|
||||
|
||||
trade_date = _source_date(row, "trade_date")
|
||||
assert trade_date is not None
|
||||
return cls(
|
||||
ts_code=_required_text(row, "ts_code"),
|
||||
trade_date=trade_date,
|
||||
suspend_timing=_optional_text(row, "suspend_timing"),
|
||||
suspend_type=_required_text(row, "suspend_type"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DailyRow:
|
||||
"""One stock daily row retaining Tushare's thousand-yuan amount."""
|
||||
|
||||
ts_code: str
|
||||
trade_date: date
|
||||
close: Decimal | None
|
||||
pre_close: Decimal | None
|
||||
pct_chg: Decimal | None
|
||||
volume: Decimal | None
|
||||
amount_thousand_yuan: Decimal | None
|
||||
|
||||
@property
|
||||
def turnover_yuan(self) -> Decimal | None:
|
||||
"""Convert observed turnover to yuan without inventing missing values."""
|
||||
|
||||
return (
|
||||
None if self.amount_thousand_yuan is None else self.amount_thousand_yuan * Decimal(1000)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> DailyRow:
|
||||
"""Parse one ``daily`` row."""
|
||||
|
||||
trade_date = _source_date(row, "trade_date")
|
||||
assert trade_date is not None
|
||||
return cls(
|
||||
ts_code=_required_text(row, "ts_code"),
|
||||
trade_date=trade_date,
|
||||
close=_decimal(row, "close"),
|
||||
pre_close=_decimal(row, "pre_close"),
|
||||
pct_chg=_decimal(row, "pct_chg"),
|
||||
volume=_decimal(row, "vol"),
|
||||
amount_thousand_yuan=_decimal(row, "amount"),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MoneyflowDcRow:
|
||||
"""One stock main-moneyflow row retaining Tushare's ten-thousand-yuan amount."""
|
||||
|
||||
trade_date: date
|
||||
ts_code: str
|
||||
name: str
|
||||
net_amount_ten_thousand_yuan: Decimal | None
|
||||
net_amount_rate: Decimal | None
|
||||
pct_change: Decimal | None
|
||||
close: Decimal | None
|
||||
|
||||
@property
|
||||
def net_amount_yuan(self) -> Decimal | None:
|
||||
"""Convert observed main net amount to yuan without filling NULL as zero."""
|
||||
|
||||
return (
|
||||
None
|
||||
if self.net_amount_ten_thousand_yuan is None
|
||||
else self.net_amount_ten_thousand_yuan * Decimal(10_000)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_mapping(cls, row: Mapping[str, SourceScalar]) -> MoneyflowDcRow:
|
||||
"""Parse one ``moneyflow_dc`` row."""
|
||||
|
||||
trade_date = _source_date(row, "trade_date")
|
||||
assert trade_date is not None
|
||||
return cls(
|
||||
trade_date=trade_date,
|
||||
ts_code=_required_text(row, "ts_code"),
|
||||
name=_required_text(row, "name"),
|
||||
net_amount_ten_thousand_yuan=_decimal(row, "net_amount"),
|
||||
net_amount_rate=_decimal(row, "net_amount_rate"),
|
||||
pct_change=_decimal(row, "pct_change"),
|
||||
close=_decimal(row, "close"),
|
||||
)
|
||||
|
||||
|
||||
class CapabilityStatus(StrEnum):
|
||||
"""Safe capability outcomes that never expose provider error text."""
|
||||
|
||||
OK = "ok"
|
||||
FORBIDDEN = "forbidden"
|
||||
RATE_LIMITED = "rate_limited"
|
||||
SERVER_ERROR = "server_error"
|
||||
SCHEMA_ERROR = "schema_error"
|
||||
TRUNCATED = "truncated"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CapabilityInterfaceResult:
|
||||
"""Safe, credential-free observation for one required interface."""
|
||||
|
||||
api_name: str
|
||||
requested_fields: tuple[str, ...]
|
||||
returned_fields: tuple[str, ...]
|
||||
status: CapabilityStatus
|
||||
row_count: int
|
||||
row_limit: int | None
|
||||
retryable: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CapabilityProbeResult:
|
||||
"""Read-only account capability report for the seven MVP interfaces."""
|
||||
|
||||
observed_at: datetime
|
||||
interfaces: tuple[CapabilityInterfaceResult, ...]
|
||||
|
||||
@property
|
||||
def succeeded(self) -> bool:
|
||||
"""Return whether every required interface passed its probe."""
|
||||
|
||||
return bool(self.interfaces) and all(
|
||||
result.status is CapabilityStatus.OK for result in self.interfaces
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
"""Infrastructure adapters for the sector radar bounded context."""
|
||||
@@ -0,0 +1,460 @@
|
||||
"""Deterministic in-memory repository used by application and contract tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Generator, Iterable, Sequence
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import replace
|
||||
from datetime import date, datetime
|
||||
|
||||
from ..domain.models import PublicationStatus, RadarPublication, RankedMetric, SectorDailyAggregate
|
||||
from ..domain.persistence import (
|
||||
DailyAggregateRecord,
|
||||
MembershipRecord,
|
||||
PublicationSourceGroup,
|
||||
PublicationSourceRecord,
|
||||
RankingRecord,
|
||||
StockFactRecord,
|
||||
WriteCounts,
|
||||
)
|
||||
from ..domain.source import SourceSnapshot
|
||||
|
||||
|
||||
class InMemorySectorRadarRepository:
|
||||
"""Keep immutable radar revisions in dictionaries without hiding overwrites."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.source_snapshots: dict[str, SourceSnapshot] = {}
|
||||
self.publication_sources: dict[tuple[str, str, int], PublicationSourceRecord] = {}
|
||||
self.memberships: dict[tuple[str, str, str], MembershipRecord] = {}
|
||||
self.stock_facts: dict[tuple[str, str], StockFactRecord] = {}
|
||||
self.publications: dict[str, RadarPublication] = {}
|
||||
self.daily_aggregates: dict[tuple[str, str, str], DailyAggregateRecord] = {}
|
||||
self.rankings: dict[tuple[str, str, str, str], RankingRecord] = {}
|
||||
self.lock_available = True
|
||||
|
||||
@contextmanager
|
||||
def advisory_lock(self, target_trade_date: date) -> Generator[bool]:
|
||||
"""Expose a controllable lock result for build orchestration tests."""
|
||||
|
||||
del target_trade_date
|
||||
yield self.lock_available
|
||||
|
||||
def save_source_snapshots(self, snapshots: Iterable[SourceSnapshot]) -> WriteCounts:
|
||||
"""Insert new content-addressed snapshots and count identical replays."""
|
||||
|
||||
inserted = 0
|
||||
unchanged = 0
|
||||
seen: set[str] = set()
|
||||
for snapshot in snapshots:
|
||||
if snapshot.snapshot_id in seen:
|
||||
raise ValueError("one write batch must not contain duplicate business keys")
|
||||
seen.add(snapshot.snapshot_id)
|
||||
if snapshot.snapshot_id in self.source_snapshots:
|
||||
unchanged += 1
|
||||
else:
|
||||
self.source_snapshots[snapshot.snapshot_id] = snapshot
|
||||
inserted += 1
|
||||
return WriteCounts(inserted, unchanged)
|
||||
|
||||
def save_publication_sources(self, records: Iterable[PublicationSourceRecord]) -> WriteCounts:
|
||||
"""Checkpoint completed source groups under one publication attempt."""
|
||||
|
||||
items = tuple(records)
|
||||
for item in items:
|
||||
if item.publication_id not in self.publications:
|
||||
raise ValueError("publication source publication does not exist")
|
||||
if item.snapshot.snapshot_id not in self.source_snapshots:
|
||||
raise ValueError("publication source snapshot does not exist")
|
||||
return self._insert_immutable(
|
||||
self.publication_sources,
|
||||
items,
|
||||
key=lambda item: (
|
||||
item.publication_id,
|
||||
item.source_group.value,
|
||||
item.source_order,
|
||||
),
|
||||
)
|
||||
|
||||
def load_publication_sources(self, publication_id: str) -> Sequence[PublicationSourceRecord]:
|
||||
"""Load source checkpoints in stable group and request order."""
|
||||
|
||||
return tuple(
|
||||
sorted(
|
||||
(
|
||||
item
|
||||
for item in self.publication_sources.values()
|
||||
if item.publication_id == publication_id
|
||||
),
|
||||
key=lambda item: (item.source_group.value, item.source_order),
|
||||
)
|
||||
)
|
||||
|
||||
def mark_publication_sources_for_retry(
|
||||
self,
|
||||
publication_id: str,
|
||||
source_groups: Sequence[PublicationSourceGroup],
|
||||
) -> None:
|
||||
"""Mark only incomplete source groups for a future partial retry."""
|
||||
|
||||
requested = set(source_groups)
|
||||
for key, record in tuple(self.publication_sources.items()):
|
||||
if record.publication_id == publication_id and record.source_group in requested:
|
||||
self.publication_sources[key] = replace(record, refresh_on_retry=True)
|
||||
|
||||
def save_memberships(self, records: Iterable[MembershipRecord]) -> WriteCounts:
|
||||
"""Insert membership rows without overwriting an earlier source revision."""
|
||||
|
||||
return self._insert_immutable(
|
||||
self.memberships,
|
||||
records,
|
||||
key=lambda item: (
|
||||
item.source_snapshot_id,
|
||||
item.sector_code,
|
||||
item.membership_key,
|
||||
),
|
||||
)
|
||||
|
||||
def save_stock_facts(self, records: Iterable[StockFactRecord]) -> WriteCounts:
|
||||
"""Insert normalized fact revisions idempotently."""
|
||||
|
||||
return self._insert_immutable(
|
||||
self.stock_facts,
|
||||
records,
|
||||
key=lambda item: (item.fact_revision, item.ts_code),
|
||||
)
|
||||
|
||||
def save_daily_aggregates(self, records: Iterable[DailyAggregateRecord]) -> WriteCounts:
|
||||
"""Insert publication-owned exact strategy inputs idempotently."""
|
||||
|
||||
items = tuple(records)
|
||||
for item in items:
|
||||
if item.publication_id not in self.publications:
|
||||
raise ValueError("daily aggregate publication does not exist")
|
||||
return self._insert_immutable(
|
||||
self.daily_aggregates,
|
||||
items,
|
||||
key=lambda item: (
|
||||
item.publication_id,
|
||||
item.aggregate.sector_type.value,
|
||||
item.aggregate.sector_code,
|
||||
),
|
||||
)
|
||||
|
||||
def create_publication(self, publication: RadarPublication) -> WriteCounts:
|
||||
"""Create one running publication without replacing an existing identity."""
|
||||
|
||||
if publication.status is not PublicationStatus.RUNNING:
|
||||
raise ValueError("new publications must start in running status")
|
||||
if any(
|
||||
item.status is PublicationStatus.RUNNING
|
||||
and item.target_trade_date == publication.target_trade_date
|
||||
and item.publication_id != publication.publication_id
|
||||
for item in self.publications.values()
|
||||
):
|
||||
raise ValueError("target date already has a running publication")
|
||||
return self._insert_immutable(
|
||||
self.publications,
|
||||
(publication,),
|
||||
key=lambda item: item.publication_id,
|
||||
)
|
||||
|
||||
def finish_publication(self, publication: RadarPublication) -> None:
|
||||
"""Apply the sole allowed mutation: running to one terminal audit state."""
|
||||
|
||||
if publication.status is PublicationStatus.RUNNING:
|
||||
raise ValueError("finished publication must use a terminal status")
|
||||
current = self.publications.get(publication.publication_id)
|
||||
if current is None or current.status is not PublicationStatus.RUNNING:
|
||||
raise ValueError("publication must exist in running status")
|
||||
if current.target_trade_date != publication.target_trade_date:
|
||||
raise ValueError("publication target_trade_date cannot change")
|
||||
self.publications[publication.publication_id] = publication
|
||||
|
||||
def finalize_publication(
|
||||
self,
|
||||
publication: RadarPublication,
|
||||
*,
|
||||
memberships: Iterable[MembershipRecord],
|
||||
stock_facts: Iterable[StockFactRecord],
|
||||
daily_aggregates: Iterable[DailyAggregateRecord],
|
||||
rankings: Iterable[RankingRecord],
|
||||
retry_source_groups: Sequence[PublicationSourceGroup] = (),
|
||||
) -> None:
|
||||
"""Atomically expose all derived rows and the terminal publication in tests."""
|
||||
|
||||
previous = (
|
||||
self.memberships.copy(),
|
||||
self.stock_facts.copy(),
|
||||
self.daily_aggregates.copy(),
|
||||
self.rankings.copy(),
|
||||
self.publication_sources.copy(),
|
||||
self.publications.copy(),
|
||||
)
|
||||
try:
|
||||
self.save_memberships(memberships)
|
||||
self.save_stock_facts(stock_facts)
|
||||
self.save_daily_aggregates(daily_aggregates)
|
||||
self.save_rankings(rankings)
|
||||
self.mark_publication_sources_for_retry(
|
||||
publication.publication_id,
|
||||
retry_source_groups,
|
||||
)
|
||||
self.finish_publication(publication)
|
||||
except Exception:
|
||||
(
|
||||
self.memberships,
|
||||
self.stock_facts,
|
||||
self.daily_aggregates,
|
||||
self.rankings,
|
||||
self.publication_sources,
|
||||
self.publications,
|
||||
) = previous
|
||||
raise
|
||||
|
||||
def recover_running_publications(
|
||||
self, target_trade_date: date, *, finished_at: datetime
|
||||
) -> Sequence[str]:
|
||||
"""Fail orphaned attempts after the caller has acquired the date lock."""
|
||||
|
||||
recovered: list[str] = []
|
||||
for publication_id, publication in tuple(self.publications.items()):
|
||||
if (
|
||||
publication.target_trade_date == target_trade_date
|
||||
and publication.status is PublicationStatus.RUNNING
|
||||
):
|
||||
self.publications[publication_id] = replace(
|
||||
publication,
|
||||
status=PublicationStatus.FAILED,
|
||||
finished_at=finished_at,
|
||||
error_summary="recovered_stale_running",
|
||||
)
|
||||
recovered.append(publication_id)
|
||||
return tuple(sorted(recovered))
|
||||
|
||||
def discard_running_publication(self, publication_id: str) -> None:
|
||||
"""Remove only a provisional duplicate attempt and its owned projections."""
|
||||
|
||||
publication = self.publications.get(publication_id)
|
||||
if publication is None or publication.status is not PublicationStatus.RUNNING:
|
||||
raise ValueError("discarded publication must exist in running status")
|
||||
del self.publications[publication_id]
|
||||
self.publication_sources = {
|
||||
key: item
|
||||
for key, item in self.publication_sources.items()
|
||||
if item.publication_id != publication_id
|
||||
}
|
||||
self.daily_aggregates = {
|
||||
key: item
|
||||
for key, item in self.daily_aggregates.items()
|
||||
if item.publication_id != publication_id
|
||||
}
|
||||
self.rankings = {
|
||||
key: item
|
||||
for key, item in self.rankings.items()
|
||||
if item.publication_id != publication_id
|
||||
}
|
||||
|
||||
def save_rankings(self, records: Iterable[RankingRecord]) -> WriteCounts:
|
||||
"""Insert publication-owned rankings idempotently."""
|
||||
|
||||
items = tuple(records)
|
||||
for item in items:
|
||||
if item.publication_id not in self.publications:
|
||||
raise ValueError("ranking publication does not exist")
|
||||
return self._insert_immutable(
|
||||
self.rankings,
|
||||
items,
|
||||
key=lambda item: (
|
||||
item.publication_id,
|
||||
item.ranking.observation.sector_type.value,
|
||||
item.ranking.observation.sector_code,
|
||||
item.ranking.observation.metric_version,
|
||||
),
|
||||
)
|
||||
|
||||
def get_publication(self, publication_id: str) -> RadarPublication | None:
|
||||
"""Return one publication revision by identity."""
|
||||
|
||||
return self.publications.get(publication_id)
|
||||
|
||||
def find_reusable_publication(
|
||||
self, target_trade_date: date, input_hash: str
|
||||
) -> RadarPublication | None:
|
||||
"""Find an identical success or partial revision without hiding failures."""
|
||||
|
||||
return max(
|
||||
(
|
||||
item
|
||||
for item in self.publications.values()
|
||||
if item.status in {PublicationStatus.SUCCESS, PublicationStatus.PARTIAL}
|
||||
and item.target_trade_date == target_trade_date
|
||||
and item.input_hash == input_hash
|
||||
),
|
||||
key=lambda item: (item.finished_at or item.started_at, item.publication_id),
|
||||
default=None,
|
||||
)
|
||||
|
||||
def get_last_good_publication(
|
||||
self, target_trade_date: date | None = None
|
||||
) -> RadarPublication | None:
|
||||
"""Return only a successful publication; partial and failed never qualify."""
|
||||
|
||||
candidates = tuple(
|
||||
publication
|
||||
for publication in self.publications.values()
|
||||
if publication.status is PublicationStatus.SUCCESS
|
||||
and (target_trade_date is None or publication.target_trade_date <= target_trade_date)
|
||||
)
|
||||
return max(
|
||||
candidates,
|
||||
key=lambda item: (
|
||||
item.target_trade_date,
|
||||
item.finished_at or item.started_at,
|
||||
item.publication_id,
|
||||
),
|
||||
default=None,
|
||||
)
|
||||
|
||||
def get_successful_publication(self, target_trade_date: date) -> RadarPublication | None:
|
||||
"""Return the latest successful revision for exactly one date."""
|
||||
|
||||
candidates = tuple(
|
||||
publication
|
||||
for publication in self.publications.values()
|
||||
if publication.status is PublicationStatus.SUCCESS
|
||||
and publication.target_trade_date == target_trade_date
|
||||
)
|
||||
return max(
|
||||
candidates,
|
||||
key=lambda item: (item.finished_at or item.started_at, item.publication_id),
|
||||
default=None,
|
||||
)
|
||||
|
||||
def get_latest_publication(self) -> RadarPublication | None:
|
||||
"""Return the newest build attempt regardless of terminal status."""
|
||||
|
||||
return max(
|
||||
self.publications.values(),
|
||||
key=lambda item: (item.target_trade_date, item.started_at, item.publication_id),
|
||||
default=None,
|
||||
)
|
||||
|
||||
def list_successful_dates(self) -> Sequence[date]:
|
||||
"""Return distinct successful dates newest first."""
|
||||
|
||||
return tuple(
|
||||
sorted(
|
||||
{
|
||||
item.target_trade_date
|
||||
for item in self.publications.values()
|
||||
if item.status is PublicationStatus.SUCCESS
|
||||
},
|
||||
reverse=True,
|
||||
)
|
||||
)
|
||||
|
||||
def load_rankings(self, publication_id: str) -> Sequence[RankedMetric]:
|
||||
"""Load every ranking projection owned by one publication."""
|
||||
|
||||
return tuple(
|
||||
sorted(
|
||||
(
|
||||
record.ranking
|
||||
for record in self.rankings.values()
|
||||
if record.publication_id == publication_id
|
||||
),
|
||||
key=lambda row: (
|
||||
row.observation.sector_type.value,
|
||||
row.observation.metric_version,
|
||||
row.rank_position is None,
|
||||
row.rank_position or 0,
|
||||
row.observation.sector_code,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
def load_daily_aggregate_history(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[SectorDailyAggregate]:
|
||||
"""Load aggregates from the latest successful revision of prior dates."""
|
||||
|
||||
dates = self._previous_successful_dates(target_trade_date, limit_dates)
|
||||
selected_publications = {
|
||||
self._latest_success_for_date(item).publication_id for item in dates
|
||||
}
|
||||
return tuple(
|
||||
record.aggregate
|
||||
for record in self.daily_aggregates.values()
|
||||
if record.publication_id in selected_publications
|
||||
)
|
||||
|
||||
def load_previous_rankings(
|
||||
self, target_trade_date: date, *, limit_dates: int
|
||||
) -> Sequence[tuple[date, Sequence[RankedMetric]]]:
|
||||
"""Load prior successful rankings newest first for rank-change attachment."""
|
||||
|
||||
result: list[tuple[date, Sequence[RankedMetric]]] = []
|
||||
for trade_date in self._previous_successful_dates(target_trade_date, limit_dates):
|
||||
publication_id = self._latest_success_for_date(trade_date).publication_id
|
||||
result.append(
|
||||
(
|
||||
trade_date,
|
||||
tuple(
|
||||
record.ranking
|
||||
for record in self.rankings.values()
|
||||
if record.publication_id == publication_id
|
||||
),
|
||||
)
|
||||
)
|
||||
return tuple(result)
|
||||
|
||||
def _previous_successful_dates(self, target_trade_date: date, limit: int) -> tuple[date, ...]:
|
||||
if limit < 1:
|
||||
raise ValueError("limit_dates must be positive")
|
||||
return tuple(
|
||||
sorted(
|
||||
{
|
||||
item.target_trade_date
|
||||
for item in self.publications.values()
|
||||
if item.status is PublicationStatus.SUCCESS
|
||||
and item.target_trade_date < target_trade_date
|
||||
},
|
||||
reverse=True,
|
||||
)[:limit]
|
||||
)
|
||||
|
||||
def _latest_success_for_date(self, trade_date: date) -> RadarPublication:
|
||||
return max(
|
||||
(
|
||||
item
|
||||
for item in self.publications.values()
|
||||
if item.status is PublicationStatus.SUCCESS and item.target_trade_date == trade_date
|
||||
),
|
||||
key=lambda item: (item.finished_at or item.started_at, item.publication_id),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _insert_immutable[K, V](
|
||||
target: dict[K, V],
|
||||
values: Iterable[V],
|
||||
*,
|
||||
key: Callable[[V], K],
|
||||
) -> WriteCounts:
|
||||
inserted = 0
|
||||
unchanged = 0
|
||||
seen: set[K] = set()
|
||||
for value in values:
|
||||
item_key = key(value)
|
||||
if item_key in seen:
|
||||
raise ValueError("one write batch must not contain duplicate business keys")
|
||||
seen.add(item_key)
|
||||
existing = target.get(item_key)
|
||||
if existing is None:
|
||||
target[item_key] = value
|
||||
inserted += 1
|
||||
elif existing == value:
|
||||
unchanged += 1
|
||||
else:
|
||||
raise ValueError("immutable revision identity cannot change content")
|
||||
return WriteCounts(inserted=inserted, unchanged=unchanged)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,636 @@
|
||||
"""Tushare adapter for replayable sector radar source facts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import UTC, date, datetime
|
||||
from typing import TypeVar, cast
|
||||
|
||||
from zhixing_server.shared.request_coordinator import (
|
||||
DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
RequestCoordinator,
|
||||
TushareSourceError,
|
||||
)
|
||||
|
||||
from ..domain.models import SectorType
|
||||
from ..domain.source import (
|
||||
CapabilityInterfaceResult,
|
||||
CapabilityProbeResult,
|
||||
CapabilityStatus,
|
||||
DailyRow,
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceContractError,
|
||||
SourceResult,
|
||||
SourceSnapshot,
|
||||
SourceTruncatedError,
|
||||
StockBasicRow,
|
||||
SuspendRow,
|
||||
TradeCalendarRow,
|
||||
build_source_snapshot,
|
||||
)
|
||||
|
||||
T = TypeVar("T")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
FIELDS: dict[str, tuple[str, ...]] = {
|
||||
"trade_cal": ("exchange", "cal_date", "is_open", "pretrade_date"),
|
||||
"dc_index": (
|
||||
"ts_code",
|
||||
"trade_date",
|
||||
"name",
|
||||
"idx_type",
|
||||
"level",
|
||||
"pct_change",
|
||||
"leading_code",
|
||||
),
|
||||
"dc_member": ("trade_date", "ts_code", "con_code", "name"),
|
||||
"stock_basic": (
|
||||
"ts_code",
|
||||
"symbol",
|
||||
"name",
|
||||
"market",
|
||||
"exchange",
|
||||
"list_status",
|
||||
"list_date",
|
||||
"delist_date",
|
||||
),
|
||||
"suspend_d": ("ts_code", "trade_date", "suspend_timing", "suspend_type"),
|
||||
"daily": ("ts_code", "trade_date", "close", "pre_close", "pct_chg", "vol", "amount"),
|
||||
"moneyflow_dc": (
|
||||
"trade_date",
|
||||
"ts_code",
|
||||
"name",
|
||||
"net_amount",
|
||||
"net_amount_rate",
|
||||
"pct_change",
|
||||
"close",
|
||||
),
|
||||
}
|
||||
|
||||
ROW_LIMITS: dict[str, int | None] = {
|
||||
"trade_cal": None,
|
||||
"dc_index": 5_000,
|
||||
"dc_member": 5_000,
|
||||
"stock_basic": None,
|
||||
"suspend_d": None,
|
||||
"daily": 6_000,
|
||||
"moneyflow_dc": 6_000,
|
||||
}
|
||||
|
||||
_SECTOR_TYPE_PARAM = {
|
||||
SectorType.CONCEPT: "概念板块",
|
||||
SectorType.INDUSTRY: "行业板块",
|
||||
}
|
||||
_MONEYFLOW_WORKERS = 2
|
||||
|
||||
|
||||
class TushareSectorRadarAdapter:
|
||||
"""Fetch seven Tushare interfaces with schema, limit, and replay metadata."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client: object,
|
||||
*,
|
||||
request_coordinator: RequestCoordinator | None = None,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
request_interval_seconds: float = 0.2,
|
||||
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
sleep_fn: Callable[[float], None] = time.sleep,
|
||||
now_fn: Callable[[], datetime] = lambda: datetime.now(UTC),
|
||||
) -> None:
|
||||
"""Create an adapter around one already-authenticated SDK client."""
|
||||
|
||||
self._client = client
|
||||
self._now_fn = now_fn
|
||||
self._coordinator = request_coordinator or RequestCoordinator(
|
||||
max_retries=max_retries,
|
||||
backoff_seconds=backoff_seconds,
|
||||
request_interval_seconds=request_interval_seconds,
|
||||
cooldown_seconds=cooldown_seconds,
|
||||
wait_fn=sleep_fn,
|
||||
sleep_fn=sleep_fn,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_token(
|
||||
cls,
|
||||
token: str,
|
||||
*,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
request_interval_seconds: float = 0.2,
|
||||
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
) -> TushareSectorRadarAdapter:
|
||||
"""Create a production client without calling ``set_token`` or retaining the token."""
|
||||
|
||||
if not token.strip():
|
||||
raise ValueError("ZHIXING_TUSHARE_TOKEN is required for sector radar")
|
||||
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
||||
|
||||
return cls(
|
||||
cast(object, ts.pro_api(token)),
|
||||
max_retries=max_retries,
|
||||
backoff_seconds=backoff_seconds,
|
||||
request_interval_seconds=request_interval_seconds,
|
||||
cooldown_seconds=cooldown_seconds,
|
||||
)
|
||||
|
||||
def fetch_trade_calendar(self, start: date, end: date) -> SourceResult[TradeCalendarRow]:
|
||||
"""Fetch and validate an inclusive exchange calendar range."""
|
||||
|
||||
if end < start:
|
||||
raise ValueError("end must not precede start")
|
||||
snapshot = self._fetch_snapshot(
|
||||
"trade_cal",
|
||||
{
|
||||
"exchange": "",
|
||||
"start_date": start.strftime("%Y%m%d"),
|
||||
"end_date": end.strftime("%Y%m%d"),
|
||||
},
|
||||
target_trade_date=end,
|
||||
)
|
||||
rows = tuple(TradeCalendarRow.from_mapping(row) for row in snapshot.rows)
|
||||
self._require_unique(
|
||||
rows, key=lambda row: (row.exchange, row.cal_date), api_name="trade_cal"
|
||||
)
|
||||
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.cal_date)))
|
||||
|
||||
def fetch_sector_indices(
|
||||
self,
|
||||
trade_date: date,
|
||||
sector_type: SectorType,
|
||||
) -> SourceResult[SectorIndexRow]:
|
||||
"""Fetch one independent concept or industry universe."""
|
||||
|
||||
idx_type = _SECTOR_TYPE_PARAM[sector_type]
|
||||
snapshot = self._fetch_snapshot(
|
||||
"dc_index",
|
||||
{"trade_date": trade_date.strftime("%Y%m%d"), "idx_type": idx_type},
|
||||
target_trade_date=trade_date,
|
||||
partition_key=sector_type.value,
|
||||
)
|
||||
if any(str(row.get("idx_type")) != idx_type for row in snapshot.rows):
|
||||
raise SourceContractError("dc_index returned a different idx_type")
|
||||
self._reject_limit(snapshot)
|
||||
rows = tuple(SectorIndexRow.from_mapping(row, sector_type) for row in snapshot.rows)
|
||||
self._require_target_date(rows, trade_date, "dc_index")
|
||||
self._require_unique(rows, key=lambda row: row.sector_code, api_name="dc_index")
|
||||
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.sector_code)))
|
||||
|
||||
def fetch_sector_members(
|
||||
self,
|
||||
trade_date: date,
|
||||
sector_codes: Sequence[str],
|
||||
) -> SourceResult[SectorMemberRow]:
|
||||
"""Fetch dated members and partition when the all-market result is incomplete."""
|
||||
|
||||
expected_codes = tuple(sorted(set(sector_codes)))
|
||||
if len(expected_codes) != len(sector_codes) or any(
|
||||
not code.strip() for code in expected_codes
|
||||
):
|
||||
raise ValueError("sector_codes must contain unique non-empty values")
|
||||
initial = self._fetch_snapshot(
|
||||
"dc_member",
|
||||
{"trade_date": trade_date.strftime("%Y%m%d")},
|
||||
target_trade_date=trade_date,
|
||||
partition_key="all",
|
||||
)
|
||||
try:
|
||||
initial_rows = tuple(SectorMemberRow.from_mapping(row) for row in initial.rows)
|
||||
self._require_target_date(initial_rows, trade_date, "dc_member")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("dc_member", "all", exc)
|
||||
raise
|
||||
returned_codes = {row.sector_code for row in initial_rows}
|
||||
missing_codes = tuple(code for code in expected_codes if code not in returned_codes)
|
||||
|
||||
if initial.limit_reached:
|
||||
partition_codes = expected_codes
|
||||
merged_rows: list[SectorMemberRow] = []
|
||||
snapshots: list[SourceSnapshot] = [initial]
|
||||
else:
|
||||
partition_codes = missing_codes
|
||||
merged_rows = list(initial_rows)
|
||||
snapshots = [initial]
|
||||
|
||||
if initial.limit_reached and not partition_codes:
|
||||
raise SourceTruncatedError("dc_member reached its limit without sector partitions")
|
||||
|
||||
for sector_code in partition_codes:
|
||||
snapshot = self._fetch_snapshot(
|
||||
"dc_member",
|
||||
{
|
||||
"trade_date": trade_date.strftime("%Y%m%d"),
|
||||
"ts_code": sector_code,
|
||||
},
|
||||
target_trade_date=trade_date,
|
||||
partition_key=sector_code,
|
||||
)
|
||||
try:
|
||||
self._reject_limit(snapshot)
|
||||
partition_rows = tuple(SectorMemberRow.from_mapping(row) for row in snapshot.rows)
|
||||
self._require_target_date(partition_rows, trade_date, "dc_member")
|
||||
if any(row.sector_code != sector_code for row in partition_rows):
|
||||
raise SourceContractError("dc_member partition returned a different sector")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("dc_member", sector_code, exc)
|
||||
raise
|
||||
snapshots.append(snapshot)
|
||||
merged_rows.extend(partition_rows)
|
||||
|
||||
try:
|
||||
self._require_unique(
|
||||
merged_rows,
|
||||
key=lambda row: (row.trade_date, row.sector_code, row.stock_code),
|
||||
api_name="dc_member",
|
||||
)
|
||||
final_codes = {row.sector_code for row in merged_rows}
|
||||
explicitly_observed_codes = {
|
||||
snapshot.partition_key
|
||||
for snapshot in snapshots
|
||||
if snapshot.partition_key not in {None, "all"}
|
||||
}
|
||||
if set(expected_codes) - final_codes - explicitly_observed_codes:
|
||||
raise SourceContractError("dc_member response is missing expected sectors")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("dc_member", "merged", exc)
|
||||
raise
|
||||
return SourceResult(
|
||||
tuple(snapshots),
|
||||
tuple(sorted(merged_rows, key=lambda row: (row.sector_code, row.stock_code))),
|
||||
)
|
||||
|
||||
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
|
||||
"""Fetch the build-time current ``L`` listings in one explicit partition."""
|
||||
|
||||
snapshot = self._fetch_snapshot(
|
||||
"stock_basic",
|
||||
{"exchange": "", "list_status": "L"},
|
||||
target_trade_date=None,
|
||||
partition_key="L",
|
||||
)
|
||||
rows = tuple(StockBasicRow.from_mapping(row) for row in snapshot.rows)
|
||||
if any(row.list_status != "L" for row in rows):
|
||||
raise SourceContractError("stock_basic returned an unexpected list_status")
|
||||
self._require_unique(rows, key=lambda row: row.ts_code, api_name="stock_basic")
|
||||
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
|
||||
|
||||
def fetch_suspensions(self, trade_date: date) -> SourceResult[SuspendRow]:
|
||||
"""Fetch explicit suspend/resume events for one date."""
|
||||
|
||||
snapshot = self._fetch_snapshot(
|
||||
"suspend_d",
|
||||
{"trade_date": trade_date.strftime("%Y%m%d")},
|
||||
target_trade_date=trade_date,
|
||||
)
|
||||
rows = tuple(SuspendRow.from_mapping(row) for row in snapshot.rows)
|
||||
self._require_target_date(rows, trade_date, "suspend_d")
|
||||
self._require_unique(
|
||||
rows,
|
||||
key=lambda row: (row.ts_code, row.trade_date, row.suspend_type, row.suspend_timing),
|
||||
api_name="suspend_d",
|
||||
)
|
||||
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
|
||||
|
||||
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
|
||||
"""Fetch a full-market daily snapshot in its documented source unit."""
|
||||
|
||||
snapshot = self._fetch_snapshot(
|
||||
"daily",
|
||||
{"trade_date": trade_date.strftime("%Y%m%d")},
|
||||
target_trade_date=trade_date,
|
||||
)
|
||||
self._reject_limit(snapshot)
|
||||
rows = tuple(DailyRow.from_mapping(row) for row in snapshot.rows)
|
||||
self._require_target_date(rows, trade_date, "daily")
|
||||
self._require_unique(rows, key=lambda row: row.ts_code, api_name="daily")
|
||||
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
|
||||
|
||||
def fetch_moneyflow_dc(
|
||||
self,
|
||||
trade_date: date,
|
||||
candidate_codes: Sequence[str],
|
||||
) -> SourceResult[MoneyflowDcRow]:
|
||||
"""Fetch full-market moneyflow and refill uncovered current candidates."""
|
||||
|
||||
expected_codes = tuple(sorted(set(candidate_codes)))
|
||||
if tuple(candidate_codes) != expected_codes or any(
|
||||
not code.strip() for code in expected_codes
|
||||
):
|
||||
raise ValueError("candidate_codes must be sorted unique non-empty values")
|
||||
initial = self._fetch_snapshot(
|
||||
"moneyflow_dc",
|
||||
{"trade_date": trade_date.strftime("%Y%m%d")},
|
||||
target_trade_date=trade_date,
|
||||
partition_key="all",
|
||||
)
|
||||
try:
|
||||
initial_rows = tuple(MoneyflowDcRow.from_mapping(row) for row in initial.rows)
|
||||
self._require_target_date(initial_rows, trade_date, "moneyflow_dc")
|
||||
self._require_unique(
|
||||
initial_rows,
|
||||
key=lambda row: (row.trade_date, row.ts_code),
|
||||
api_name="moneyflow_dc",
|
||||
)
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("moneyflow_dc", "all", exc)
|
||||
raise
|
||||
|
||||
returned_codes = {row.ts_code for row in initial_rows}
|
||||
missing_codes = tuple(code for code in expected_codes if code not in returned_codes)
|
||||
if not missing_codes:
|
||||
return SourceResult(
|
||||
(initial,),
|
||||
tuple(sorted(initial_rows, key=lambda row: row.ts_code)),
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(
|
||||
max_workers=_MONEYFLOW_WORKERS,
|
||||
thread_name_prefix="sector-radar-moneyflow",
|
||||
) as executor:
|
||||
futures = {
|
||||
code: executor.submit(self._fetch_moneyflow_partition, trade_date, code)
|
||||
for code in missing_codes
|
||||
}
|
||||
partition_results = tuple(futures[code].result() for code in missing_codes)
|
||||
|
||||
snapshots = [initial]
|
||||
merged_rows = list(initial_rows)
|
||||
for result in partition_results:
|
||||
if result is None:
|
||||
continue
|
||||
snapshot, rows = result
|
||||
snapshots.append(snapshot)
|
||||
merged_rows.extend(rows)
|
||||
try:
|
||||
self._require_unique(
|
||||
merged_rows,
|
||||
key=lambda row: (row.trade_date, row.ts_code),
|
||||
api_name="moneyflow_dc",
|
||||
)
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("moneyflow_dc", "merged", exc)
|
||||
raise
|
||||
return SourceResult(
|
||||
tuple(snapshots),
|
||||
tuple(sorted(merged_rows, key=lambda row: row.ts_code)),
|
||||
)
|
||||
|
||||
def _fetch_moneyflow_partition(
|
||||
self,
|
||||
trade_date: date,
|
||||
ts_code: str,
|
||||
) -> tuple[SourceSnapshot, tuple[MoneyflowDcRow, ...]] | None:
|
||||
"""Return one validated refill partition or preserve an ordinary gap."""
|
||||
|
||||
try:
|
||||
snapshot = self._fetch_snapshot(
|
||||
"moneyflow_dc",
|
||||
{
|
||||
"trade_date": trade_date.strftime("%Y%m%d"),
|
||||
"ts_code": ts_code,
|
||||
},
|
||||
target_trade_date=trade_date,
|
||||
partition_key=ts_code,
|
||||
)
|
||||
except TushareSourceError:
|
||||
logger.warning(
|
||||
"sector_radar_moneyflow_partition_failed partition_key=%s error_type=%s",
|
||||
self._safe_partition_key(ts_code),
|
||||
TushareSourceError.__name__,
|
||||
)
|
||||
return None
|
||||
if not snapshot.rows:
|
||||
logger.warning(
|
||||
"sector_radar_moneyflow_partition_empty partition_key=%s",
|
||||
self._safe_partition_key(ts_code),
|
||||
)
|
||||
return None
|
||||
try:
|
||||
self._reject_limit(snapshot)
|
||||
rows = tuple(MoneyflowDcRow.from_mapping(row) for row in snapshot.rows)
|
||||
self._require_target_date(rows, trade_date, "moneyflow_dc")
|
||||
self._require_unique(
|
||||
rows,
|
||||
key=lambda row: (row.trade_date, row.ts_code),
|
||||
api_name="moneyflow_dc",
|
||||
)
|
||||
if any(row.ts_code != ts_code for row in rows):
|
||||
raise SourceContractError("moneyflow_dc partition returned a different ts_code")
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure("moneyflow_dc", ts_code, exc)
|
||||
raise
|
||||
return snapshot, rows
|
||||
|
||||
def probe(self, trade_date: date) -> CapabilityProbeResult:
|
||||
"""Probe required interfaces while returning only safe classifications."""
|
||||
|
||||
results: list[CapabilityInterfaceResult] = []
|
||||
concept_codes: tuple[str, ...] = ()
|
||||
|
||||
calendar = self._probe_call(
|
||||
"trade_cal", lambda: self.fetch_trade_calendar(trade_date, trade_date)
|
||||
)
|
||||
results.append(calendar[0])
|
||||
|
||||
try:
|
||||
concept = self.fetch_sector_indices(trade_date, SectorType.CONCEPT)
|
||||
industry = self.fetch_sector_indices(trade_date, SectorType.INDUSTRY)
|
||||
combined = SourceResult(
|
||||
concept.snapshots + industry.snapshots,
|
||||
concept.rows + industry.rows,
|
||||
)
|
||||
concept_codes = tuple(row.sector_code for row in combined.rows)
|
||||
results.append(self._capability_success("dc_index", combined.snapshots))
|
||||
except Exception as exc:
|
||||
results.append(self._capability_failure("dc_index", exc))
|
||||
|
||||
member = self._probe_call(
|
||||
"dc_member", lambda: self.fetch_sector_members(trade_date, concept_codes)
|
||||
)
|
||||
results.append(member[0])
|
||||
for api_name, operation in (
|
||||
("stock_basic", self.fetch_stock_basics),
|
||||
("suspend_d", lambda: self.fetch_suspensions(trade_date)),
|
||||
("daily", lambda: self.fetch_daily(trade_date)),
|
||||
("moneyflow_dc", lambda: self.fetch_moneyflow_dc(trade_date, ())),
|
||||
):
|
||||
results.append(self._probe_call(api_name, operation)[0])
|
||||
return CapabilityProbeResult(observed_at=self._now_fn(), interfaces=tuple(results))
|
||||
|
||||
def _fetch_snapshot(
|
||||
self,
|
||||
api_name: str,
|
||||
params: Mapping[str, object],
|
||||
*,
|
||||
target_trade_date: date | None,
|
||||
partition_key: str | None = None,
|
||||
) -> SourceSnapshot:
|
||||
fields = ",".join(FIELDS[api_name])
|
||||
|
||||
def request() -> object:
|
||||
query = getattr(self._client, "query", None)
|
||||
if callable(query):
|
||||
return query(api_name, fields=fields, **params)
|
||||
method = getattr(self._client, api_name, None)
|
||||
if not callable(method):
|
||||
raise TypeError(f"Tushare client has no callable {api_name}")
|
||||
return method(fields=fields, **params)
|
||||
|
||||
result = self._coordinator.call(api_name, request)
|
||||
try:
|
||||
columns = getattr(result, "columns", None)
|
||||
returned_fields = (
|
||||
tuple(str(column) for column in cast(Iterable[object], columns))
|
||||
if isinstance(columns, Iterable) and not isinstance(columns, (str, bytes))
|
||||
else None
|
||||
)
|
||||
rows = self._as_records(result)
|
||||
snapshot = build_source_snapshot(
|
||||
api_name=api_name,
|
||||
params={**params, "fields": fields},
|
||||
rows=rows,
|
||||
target_trade_date=target_trade_date,
|
||||
partition_key=partition_key,
|
||||
observed_at=self._now_fn(),
|
||||
row_limit=ROW_LIMITS[api_name],
|
||||
returned_fields=returned_fields,
|
||||
)
|
||||
missing_fields = set(FIELDS[api_name]) - set(snapshot.returned_fields)
|
||||
if snapshot.returned_fields and missing_fields:
|
||||
missing = ",".join(sorted(missing_fields))
|
||||
raise SourceContractError(
|
||||
f"{api_name} response is missing requested fields: {missing}"
|
||||
)
|
||||
return snapshot
|
||||
except SourceContractError as exc:
|
||||
self._log_contract_failure(api_name, partition_key, exc)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _log_contract_failure(
|
||||
api_name: str,
|
||||
partition_key: str | None,
|
||||
error: SourceContractError,
|
||||
) -> None:
|
||||
"""Record only operator-safe contract context, never provider payloads."""
|
||||
|
||||
if not error.claim_diagnostic():
|
||||
return
|
||||
logger.error(
|
||||
"sector_radar_source_contract_failed api_name=%s partition_key=%s validation=%s",
|
||||
api_name,
|
||||
TushareSectorRadarAdapter._safe_partition_key(partition_key),
|
||||
error.operator_message,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _safe_partition_key(partition_key: str | None) -> str:
|
||||
"""Keep expected identifiers readable while preventing log-control injection."""
|
||||
|
||||
if partition_key is None:
|
||||
return "all"
|
||||
sanitized = "".join(
|
||||
character if character.isalnum() or character in {".", "_", "-"} else "_"
|
||||
for character in partition_key
|
||||
)
|
||||
return sanitized[:64] or "unknown"
|
||||
|
||||
@staticmethod
|
||||
def _as_records(result: object) -> tuple[Mapping[str, object], ...]:
|
||||
if result is None:
|
||||
return ()
|
||||
to_dict = getattr(result, "to_dict", None)
|
||||
if callable(to_dict):
|
||||
result = to_dict("records")
|
||||
if isinstance(result, Mapping):
|
||||
return (cast(Mapping[str, object], result),)
|
||||
if isinstance(result, Iterable) and not isinstance(result, (str, bytes)):
|
||||
records: list[Mapping[str, object]] = []
|
||||
for row in cast(Iterable[object], result):
|
||||
if not isinstance(row, Mapping):
|
||||
raise SourceContractError("Tushare rows must be mappings")
|
||||
records.append(cast(Mapping[str, object], row))
|
||||
return tuple(records)
|
||||
raise SourceContractError("unsupported Tushare tabular response")
|
||||
|
||||
@staticmethod
|
||||
def _reject_limit(snapshot: SourceSnapshot) -> None:
|
||||
if snapshot.limit_reached:
|
||||
raise SourceTruncatedError(f"{snapshot.api_name} reached its provider row limit")
|
||||
|
||||
@staticmethod
|
||||
def _require_target_date(rows: Sequence[object], target: date, api_name: str) -> None:
|
||||
if any(getattr(row, "trade_date", None) != target for row in rows):
|
||||
raise SourceContractError(f"{api_name} returned a different trade_date")
|
||||
|
||||
@staticmethod
|
||||
def _require_unique(
|
||||
rows: Sequence[T],
|
||||
*,
|
||||
key: Callable[[T], object],
|
||||
api_name: str,
|
||||
) -> None:
|
||||
keys = [key(row) for row in rows]
|
||||
if len(keys) != len(set(keys)):
|
||||
raise SourceContractError(f"{api_name} returned duplicate business keys")
|
||||
|
||||
def _probe_call(
|
||||
self,
|
||||
api_name: str,
|
||||
operation: Callable[[], SourceResult[object]],
|
||||
) -> tuple[CapabilityInterfaceResult, SourceResult[object] | None]:
|
||||
try:
|
||||
result = operation()
|
||||
except Exception as exc:
|
||||
return self._capability_failure(api_name, exc), None
|
||||
return self._capability_success(api_name, result.snapshots), result
|
||||
|
||||
@staticmethod
|
||||
def _capability_success(
|
||||
api_name: str,
|
||||
snapshots: Sequence[SourceSnapshot],
|
||||
) -> CapabilityInterfaceResult:
|
||||
return CapabilityInterfaceResult(
|
||||
api_name=api_name,
|
||||
requested_fields=FIELDS[api_name],
|
||||
returned_fields=tuple(
|
||||
sorted({field for item in snapshots for field in item.returned_fields})
|
||||
),
|
||||
status=CapabilityStatus.OK,
|
||||
row_count=sum(item.row_count for item in snapshots),
|
||||
row_limit=ROW_LIMITS[api_name],
|
||||
retryable=False,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _capability_failure(api_name: str, error: BaseException) -> CapabilityInterfaceResult:
|
||||
classified_error = error.__cause__ if error.__cause__ is not None else error
|
||||
message = str(classified_error).casefold()
|
||||
if isinstance(error, SourceTruncatedError):
|
||||
status = CapabilityStatus.TRUNCATED
|
||||
elif RequestCoordinator.is_rate_limited(error) or RequestCoordinator.is_rate_limited(
|
||||
classified_error
|
||||
):
|
||||
status = CapabilityStatus.RATE_LIMITED
|
||||
elif "权限" in message or "forbidden" in message or "permission" in message:
|
||||
status = CapabilityStatus.FORBIDDEN
|
||||
elif isinstance(error, (SourceContractError, ValueError, TypeError)):
|
||||
status = CapabilityStatus.SCHEMA_ERROR
|
||||
else:
|
||||
status = CapabilityStatus.SERVER_ERROR
|
||||
return CapabilityInterfaceResult(
|
||||
api_name=api_name,
|
||||
requested_fields=FIELDS[api_name],
|
||||
returned_fields=(),
|
||||
status=status,
|
||||
row_count=0,
|
||||
row_limit=ROW_LIMITS[api_name],
|
||||
retryable=status in {CapabilityStatus.RATE_LIMITED, CapabilityStatus.SERVER_ERROR},
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
"""Delivery adapters for sector radar build and query use cases."""
|
||||
@@ -0,0 +1,108 @@
|
||||
"""One-shot ``sector-radar-build`` command for external schedulers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from datetime import date
|
||||
|
||||
from ....bootstrap.config import get_settings
|
||||
from ..application.build import BuildSectorRadar, BuildSectorRadarCommand
|
||||
from ..infrastructure.postgres import PostgresSectorRadarRepository
|
||||
from ..infrastructure.tushare import TushareSectorRadarAdapter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
"""Build mutually exclusive single-date, range, and retry modes."""
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Build independent Tushare sector radar publications"
|
||||
)
|
||||
mode = parser.add_mutually_exclusive_group()
|
||||
mode.add_argument("--trade-date", type=_parse_date, help="target date in YYYY-MM-DD")
|
||||
mode.add_argument(
|
||||
"--retry-publication-id",
|
||||
help="resume failed source groups from a partial or failed publication",
|
||||
)
|
||||
mode.add_argument("--start-date", type=_parse_date, help="inclusive backfill start date")
|
||||
parser.add_argument("--end-date", type=_parse_date, help="inclusive backfill end date")
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
"""Execute one build invocation and print a redacted JSON summary."""
|
||||
|
||||
args = build_parser().parse_args(argv)
|
||||
if (args.start_date is None) != (args.end_date is None):
|
||||
raise SystemExit("--start-date and --end-date must be provided together")
|
||||
command = BuildSectorRadarCommand(
|
||||
trade_date=args.trade_date,
|
||||
start_date=args.start_date,
|
||||
end_date=args.end_date,
|
||||
retry_publication_id=args.retry_publication_id,
|
||||
)
|
||||
try:
|
||||
settings = get_settings()
|
||||
logging.basicConfig(
|
||||
level=settings.log_level.upper(),
|
||||
format="%(asctime)s %(levelname)s %(name)s %(message)s",
|
||||
force=True,
|
||||
)
|
||||
logger.info(
|
||||
"sector_radar_build_cli trade_date=%s start_date=%s end_date=%s retry=%s",
|
||||
command.trade_date or "auto",
|
||||
command.start_date or "none",
|
||||
command.end_date or "none",
|
||||
bool(command.retry_publication_id),
|
||||
)
|
||||
source = TushareSectorRadarAdapter.from_token(
|
||||
settings.tushare_token,
|
||||
max_retries=settings.sector_radar_max_retries,
|
||||
backoff_seconds=settings.sector_radar_retry_backoff_seconds,
|
||||
request_interval_seconds=settings.sector_radar_request_interval_seconds,
|
||||
)
|
||||
repository = PostgresSectorRadarRepository(
|
||||
settings.database_url,
|
||||
advisory_lock_key=settings.sector_radar_advisory_lock_key,
|
||||
)
|
||||
try:
|
||||
summary = BuildSectorRadar(
|
||||
source,
|
||||
repository,
|
||||
coverage_threshold=settings.sector_radar_coverage_threshold,
|
||||
).execute(command)
|
||||
finally:
|
||||
repository.close()
|
||||
except Exception as exc: # noqa: BLE001 - CLI boundary returns a redacted scheduler result
|
||||
logger.error("sector_radar_build_initialization_failed error_type=%s", type(exc).__name__)
|
||||
print(
|
||||
json.dumps(
|
||||
{
|
||||
"status": "failed",
|
||||
"exit_code": 1,
|
||||
"outcomes": [],
|
||||
"error_type": type(exc).__name__,
|
||||
"error_message": "sector radar build initialization failed",
|
||||
},
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
)
|
||||
)
|
||||
return 1
|
||||
print(json.dumps(summary.as_dict(), ensure_ascii=False, sort_keys=True))
|
||||
return summary.exit_code
|
||||
|
||||
|
||||
def _parse_date(value: str) -> date:
|
||||
try:
|
||||
return date.fromisoformat(value)
|
||||
except ValueError as exc:
|
||||
raise argparse.ArgumentTypeError("date must use YYYY-MM-DD") from exc
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,302 @@
|
||||
"""HTTP presentation for persisted sector radar rankings."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import atexit
|
||||
import threading
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from ....bootstrap.config import Settings, get_settings
|
||||
from ..application.read import (
|
||||
RadarDateIndex,
|
||||
RadarMetricDefinition,
|
||||
RadarQuery,
|
||||
RadarView,
|
||||
RankingPage,
|
||||
ReadSectorRadar,
|
||||
)
|
||||
from ..domain.models import (
|
||||
MetricKind,
|
||||
MetricQuality,
|
||||
MetricUnit,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
from ..infrastructure.postgres import (
|
||||
PostgresSectorRadarRepository,
|
||||
SectorRadarRepositoryError,
|
||||
)
|
||||
|
||||
sector_radar_router = APIRouter()
|
||||
_REPOSITORY_CACHE_LOCK = threading.Lock()
|
||||
_REPOSITORY_CACHE: dict[tuple[str, int], PostgresSectorRadarRepository] = {}
|
||||
|
||||
|
||||
class RadarPublicationResponse(BaseModel):
|
||||
"""Safe publication provenance and quality metadata."""
|
||||
|
||||
publication_id: str
|
||||
target_trade_date: date
|
||||
status: PublicationStatus
|
||||
source_version: str
|
||||
universe_version: str
|
||||
metric_versions: list[str]
|
||||
input_hash: str | None
|
||||
coverage: Decimal = Field(ge=0, le=1)
|
||||
started_at: datetime
|
||||
finished_at: datetime | None
|
||||
error_summary: str | None
|
||||
|
||||
|
||||
class RadarDatesResponse(BaseModel):
|
||||
"""Successful dates plus newest attempt and strict last-good metadata."""
|
||||
|
||||
status: Literal["success", "no_data"]
|
||||
available_dates: list[date]
|
||||
current_attempt: RadarPublicationResponse | None
|
||||
last_good: RadarPublicationResponse | None
|
||||
|
||||
|
||||
class RadarMetricDefinitionResponse(BaseModel):
|
||||
"""Version and labeling for one independent metric implementation."""
|
||||
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
label: str
|
||||
unit: MetricUnit
|
||||
implementation_kind: Literal["independent"]
|
||||
disclaimer: str
|
||||
|
||||
|
||||
class RadarRankingRowResponse(BaseModel):
|
||||
"""One sector ranking row with explicit null and unit semantics."""
|
||||
|
||||
trade_date: date
|
||||
sector_type: SectorType
|
||||
sector_code: str
|
||||
sector_name: str
|
||||
metric_kind: MetricKind
|
||||
metric_version: str
|
||||
implementation_kind: Literal["independent"]
|
||||
unit: MetricUnit
|
||||
metric_value: Decimal | None
|
||||
quality: MetricQuality
|
||||
member_count: int = Field(ge=0)
|
||||
valid_sample_count: int = Field(ge=0)
|
||||
membership_coverage: Decimal = Field(ge=0, le=1)
|
||||
moneyflow_coverage: Decimal = Field(ge=0, le=1)
|
||||
rank_position: int | None = Field(default=None, ge=1)
|
||||
rank_percentile: Decimal | None = Field(default=None, gt=0, le=100)
|
||||
rank_change_days: int = Field(ge=1, le=5)
|
||||
rank_change: int | None
|
||||
|
||||
|
||||
def _empty_ranking_rows() -> list[RadarRankingRowResponse]:
|
||||
return []
|
||||
|
||||
|
||||
class RadarRankingsResponse(BaseModel):
|
||||
"""One persisted, filtered ranking page."""
|
||||
|
||||
status: Literal["success", "no_data"]
|
||||
requested_trade_date: date | None
|
||||
sector_type: SectorType
|
||||
view: RadarView
|
||||
rank_change_metric: MetricKind
|
||||
rank_change_days: int = Field(ge=1, le=5)
|
||||
side: RankSide
|
||||
search: str | None
|
||||
publication: RadarPublicationResponse | None
|
||||
definition: RadarMetricDefinitionResponse
|
||||
page: int = Field(ge=1)
|
||||
page_size: int = Field(ge=1, le=100)
|
||||
total: int = Field(ge=0)
|
||||
rows: list[RadarRankingRowResponse] = Field(default_factory=_empty_ranking_rows)
|
||||
|
||||
|
||||
def get_sector_radar_reader(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> ReadSectorRadar:
|
||||
"""Return a reader backed by one process-cached PostgreSQL repository."""
|
||||
|
||||
key = (settings.database_url, settings.sector_radar_advisory_lock_key)
|
||||
with _REPOSITORY_CACHE_LOCK:
|
||||
repository = _REPOSITORY_CACHE.get(key)
|
||||
if repository is None:
|
||||
repository = PostgresSectorRadarRepository(
|
||||
settings.database_url,
|
||||
advisory_lock_key=settings.sector_radar_advisory_lock_key,
|
||||
)
|
||||
_REPOSITORY_CACHE[key] = repository
|
||||
return ReadSectorRadar(repository)
|
||||
|
||||
|
||||
def _close_cached_repositories() -> None:
|
||||
"""Close process-owned radar pools during interpreter shutdown."""
|
||||
|
||||
with _REPOSITORY_CACHE_LOCK:
|
||||
repositories = tuple(_REPOSITORY_CACHE.values())
|
||||
_REPOSITORY_CACHE.clear()
|
||||
for repository in repositories:
|
||||
repository.close()
|
||||
|
||||
|
||||
atexit.register(_close_cached_repositories)
|
||||
|
||||
|
||||
@sector_radar_router.get("/dates", response_model=RadarDatesResponse)
|
||||
def get_sector_radar_dates(
|
||||
reader: Annotated[ReadSectorRadar, Depends(get_sector_radar_reader)],
|
||||
) -> RadarDatesResponse:
|
||||
"""Return persisted availability without invoking Tushare."""
|
||||
|
||||
try:
|
||||
return _dates_response(reader.list_dates())
|
||||
except SectorRadarRepositoryError as exc:
|
||||
raise _storage_error() from exc
|
||||
|
||||
|
||||
@sector_radar_router.get("/rankings", response_model=RadarRankingsResponse)
|
||||
def get_sector_radar_rankings(
|
||||
reader: Annotated[ReadSectorRadar, Depends(get_sector_radar_reader)],
|
||||
trade_date: date | None = None,
|
||||
sector_type: SectorType = SectorType.CONCEPT,
|
||||
view: RadarView = RadarView.AMOUNT,
|
||||
rank_change_metric: MetricKind = MetricKind.AMOUNT,
|
||||
rank_change_days: Annotated[int, Query(ge=1, le=5)] = 1,
|
||||
side: RankSide = RankSide.ALL,
|
||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||
page: Annotated[int, Query(ge=1)] = 1,
|
||||
page_size: Annotated[int, Query(ge=1, le=100)] = 20,
|
||||
) -> RadarRankingsResponse:
|
||||
"""Return one filtered page from a successful publication."""
|
||||
|
||||
query = RadarQuery(
|
||||
trade_date=trade_date,
|
||||
sector_type=sector_type,
|
||||
view=view,
|
||||
rank_change_metric=rank_change_metric,
|
||||
rank_change_days=rank_change_days,
|
||||
side=side,
|
||||
search=search.strip() or None if search else None,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
try:
|
||||
return _rankings_response(reader.query(query))
|
||||
except SectorRadarRepositoryError as exc:
|
||||
raise _storage_error() from exc
|
||||
|
||||
|
||||
def _dates_response(index: RadarDateIndex) -> RadarDatesResponse:
|
||||
return RadarDatesResponse(
|
||||
status=index.status,
|
||||
available_dates=list(index.available_dates),
|
||||
current_attempt=(
|
||||
_publication_response(index.current_attempt)
|
||||
if index.current_attempt is not None
|
||||
else None
|
||||
),
|
||||
last_good=(_publication_response(index.last_good) if index.last_good is not None else None),
|
||||
)
|
||||
|
||||
|
||||
def _rankings_response(page: RankingPage) -> RadarRankingsResponse:
|
||||
query = page.query
|
||||
return RadarRankingsResponse(
|
||||
status=page.status,
|
||||
requested_trade_date=query.trade_date,
|
||||
sector_type=query.sector_type,
|
||||
view=query.view,
|
||||
rank_change_metric=query.rank_change_metric,
|
||||
rank_change_days=query.rank_change_days,
|
||||
side=query.side,
|
||||
search=query.search,
|
||||
publication=(
|
||||
_publication_response(page.publication) if page.publication is not None else None
|
||||
),
|
||||
definition=_definition_response(page.definition),
|
||||
page=query.page,
|
||||
page_size=query.page_size,
|
||||
total=page.total,
|
||||
rows=[_ranking_response(row, query.rank_change_days) for row in page.rows],
|
||||
)
|
||||
|
||||
|
||||
def _publication_response(publication: RadarPublication) -> RadarPublicationResponse:
|
||||
return RadarPublicationResponse(
|
||||
publication_id=publication.publication_id,
|
||||
target_trade_date=publication.target_trade_date,
|
||||
status=publication.status,
|
||||
source_version=publication.source_version,
|
||||
universe_version=publication.universe_version,
|
||||
metric_versions=list(publication.metric_versions),
|
||||
input_hash=publication.input_hash,
|
||||
coverage=publication.coverage,
|
||||
started_at=publication.started_at,
|
||||
finished_at=publication.finished_at,
|
||||
error_summary=publication.error_summary,
|
||||
)
|
||||
|
||||
|
||||
def _definition_response(
|
||||
definition: RadarMetricDefinition,
|
||||
) -> RadarMetricDefinitionResponse:
|
||||
return RadarMetricDefinitionResponse(
|
||||
metric_kind=definition.metric_kind,
|
||||
metric_version=definition.metric_version,
|
||||
label=definition.label,
|
||||
unit=definition.unit,
|
||||
implementation_kind=definition.implementation_kind,
|
||||
disclaimer=definition.disclaimer,
|
||||
)
|
||||
|
||||
|
||||
def _ranking_response(row: RankedMetric, rank_change_days: int) -> RadarRankingRowResponse:
|
||||
observation = row.observation
|
||||
return RadarRankingRowResponse(
|
||||
trade_date=observation.trade_date,
|
||||
sector_type=observation.sector_type,
|
||||
sector_code=observation.sector_code,
|
||||
sector_name=observation.sector_name,
|
||||
metric_kind=observation.metric_kind,
|
||||
metric_version=observation.metric_version,
|
||||
implementation_kind=observation.implementation_kind,
|
||||
unit=observation.unit,
|
||||
metric_value=observation.value,
|
||||
quality=observation.quality,
|
||||
member_count=observation.member_count,
|
||||
valid_sample_count=observation.valid_sample_count,
|
||||
membership_coverage=observation.membership_coverage,
|
||||
moneyflow_coverage=observation.moneyflow_coverage,
|
||||
rank_position=row.rank_position,
|
||||
rank_percentile=row.rank_percentile,
|
||||
rank_change_days=rank_change_days,
|
||||
rank_change=row.rank_change(rank_change_days),
|
||||
)
|
||||
|
||||
|
||||
def _storage_error() -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=503,
|
||||
detail={
|
||||
"code": "sector_radar_storage_unavailable",
|
||||
"message": "sector radar storage is unavailable",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RadarDatesResponse",
|
||||
"RadarRankingsResponse",
|
||||
"get_sector_radar_reader",
|
||||
"sector_radar_router",
|
||||
]
|
||||
@@ -0,0 +1,186 @@
|
||||
"""Shared bounded retry and provider rate-limit coordination."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Sequence
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_RATE_LIMIT_COOLDOWNS = (60.0, 120.0, 180.0)
|
||||
_RATE_LIMIT_MESSAGES = (
|
||||
"访问频繁",
|
||||
"请稍后",
|
||||
"超过频率",
|
||||
"频率限制",
|
||||
"too many requests",
|
||||
"rate limit",
|
||||
"rate_limit",
|
||||
"http 429",
|
||||
"status code: 429",
|
||||
"429",
|
||||
"http 403",
|
||||
"status code: 403",
|
||||
"403",
|
||||
)
|
||||
|
||||
|
||||
class TushareSourceError(RuntimeError):
|
||||
"""A Tushare request failed after the configured retry budget."""
|
||||
|
||||
|
||||
class RequestCoordinator:
|
||||
"""Coordinate retries, rate-limit cooling, and optional request start spacing.
|
||||
|
||||
Provider calls execute outside the coordinator lock and may overlap. When a
|
||||
positive request interval is configured, only their start times are serialized.
|
||||
Injectable time functions keep waits deterministic in tests without coupling
|
||||
the coordinator to any business bounded context.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
max_retries: int = 3,
|
||||
backoff_seconds: float = 1.0,
|
||||
request_interval_seconds: float = 0.0,
|
||||
cooldown_seconds: Sequence[float] = DEFAULT_RATE_LIMIT_COOLDOWNS,
|
||||
random_fn: Callable[[], float] = random.random,
|
||||
clock: Callable[[], float] = time.monotonic,
|
||||
wait_fn: Callable[[float], None] = time.sleep,
|
||||
sleep_fn: Callable[[float], None] | None = None,
|
||||
) -> None:
|
||||
cooldowns = tuple(float(value) for value in cooldown_seconds)
|
||||
if not cooldowns or any(value < 0 for value in cooldowns):
|
||||
raise ValueError("cooldown_seconds must contain non-negative values")
|
||||
self.max_retries = max(0, max_retries)
|
||||
self.backoff_seconds = max(0.0, backoff_seconds)
|
||||
self.request_interval_seconds = max(0.0, request_interval_seconds)
|
||||
self.cooldown_seconds = cooldowns
|
||||
self.random_fn = random_fn
|
||||
self.clock = clock
|
||||
self.wait_fn = wait_fn
|
||||
self.sleep_fn = sleep_fn or wait_fn
|
||||
self._condition = threading.Condition()
|
||||
self._cooldown_until = 0.0
|
||||
self._next_request_start = 0.0
|
||||
self._rate_limit_count = 0
|
||||
|
||||
@property
|
||||
def cooldown_until(self) -> float:
|
||||
"""Return the current monotonic cooldown deadline."""
|
||||
|
||||
with self._condition:
|
||||
return self._cooldown_until
|
||||
|
||||
def call(self, method_name: str, request: Callable[[], object]) -> object:
|
||||
"""Execute one provider request with bounded, shared retry behavior."""
|
||||
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
self._wait_for_request_start(method_name)
|
||||
try:
|
||||
result = request()
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
if self.is_rate_limited(exc):
|
||||
cooldown = self._set_rate_limit_cooldown()
|
||||
logger.warning(
|
||||
"provider_rate_limit method=%s attempt=%d max_attempts=%d "
|
||||
"cooldown_seconds=%.1f",
|
||||
method_name,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
cooldown,
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
continue
|
||||
break
|
||||
if not self._is_retryable(exc):
|
||||
raise
|
||||
if attempt == self.max_retries:
|
||||
break
|
||||
delay = self.backoff_seconds * (2**attempt) * (0.5 + self.random_fn())
|
||||
logger.warning(
|
||||
"provider_request_retry method=%s attempt=%d max_attempts=%d "
|
||||
"backoff_seconds=%.1f",
|
||||
method_name,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
delay,
|
||||
)
|
||||
self.sleep_fn(delay)
|
||||
else:
|
||||
self._clear_rate_limit_after_success()
|
||||
return result
|
||||
logger.error(
|
||||
"provider_request_failed method=%s attempts=%d",
|
||||
method_name,
|
||||
self.max_retries + 1,
|
||||
)
|
||||
raise TushareSourceError(f"Tushare request failed: {method_name}") from last_error
|
||||
|
||||
def request(self, method_name: str, operation: Callable[[], object]) -> object:
|
||||
"""Alias for ``call`` for adapters that model requests as a port."""
|
||||
|
||||
return self.call(method_name, operation)
|
||||
|
||||
def _wait_for_request_start(self, method_name: str) -> None:
|
||||
"""Reserve one start slot after both shared wait deadlines have elapsed."""
|
||||
|
||||
while True:
|
||||
with self._condition:
|
||||
now = self.clock()
|
||||
start_at = max(self._cooldown_until, self._next_request_start)
|
||||
delay = start_at - now
|
||||
if delay <= 0:
|
||||
self._next_request_start = now + self.request_interval_seconds
|
||||
return
|
||||
if start_at == self._cooldown_until:
|
||||
logger.info(
|
||||
"provider_rate_limit_wait method=%s wait_seconds=%.1f",
|
||||
method_name,
|
||||
delay,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"provider_request_interval_wait method=%s wait_seconds=%.3f",
|
||||
method_name,
|
||||
delay,
|
||||
)
|
||||
self.wait_fn(delay)
|
||||
|
||||
def _set_rate_limit_cooldown(self) -> float:
|
||||
with self._condition:
|
||||
self._rate_limit_count += 1
|
||||
index = min(self._rate_limit_count - 1, len(self.cooldown_seconds) - 1)
|
||||
duration = self.cooldown_seconds[index]
|
||||
self._cooldown_until = max(self._cooldown_until, self.clock() + duration)
|
||||
self._condition.notify_all()
|
||||
return duration
|
||||
|
||||
def _clear_rate_limit_after_success(self) -> None:
|
||||
with self._condition:
|
||||
if self.clock() >= self._cooldown_until:
|
||||
self._rate_limit_count = 0
|
||||
|
||||
@staticmethod
|
||||
def is_rate_limited(error: BaseException) -> bool:
|
||||
"""Classify stable provider rate-limit signals without logging details."""
|
||||
|
||||
for attribute in ("status_code", "status", "code"):
|
||||
value = getattr(error, attribute, None)
|
||||
if str(value).strip() in {"403", "429"}:
|
||||
return True
|
||||
message = str(error).casefold()
|
||||
return any(marker.casefold() in message for marker in _RATE_LIMIT_MESSAGES)
|
||||
|
||||
@staticmethod
|
||||
def _is_retryable(error: BaseException) -> bool:
|
||||
return isinstance(error, (OSError, RuntimeError, TimeoutError))
|
||||
|
||||
|
||||
TushareRequestCoordinator = RequestCoordinator
|
||||
@@ -39,6 +39,13 @@ def test_postgres_migration_creates_market_data_contract(
|
||||
"selection_run",
|
||||
"selection_run_item",
|
||||
"selection_signal",
|
||||
"sector_radar_source_snapshot",
|
||||
"sector_radar_membership",
|
||||
"sector_radar_stock_fact",
|
||||
"sector_radar_publication",
|
||||
"sector_radar_ranking",
|
||||
"sector_radar_daily_aggregate",
|
||||
"sector_radar_publication_source",
|
||||
} <= tables
|
||||
item_columns = {column["name"] for column in inspector.get_columns("selection_run_item")}
|
||||
assert {
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
import os
|
||||
from dataclasses import replace
|
||||
from datetime import UTC, date, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
|
||||
import psycopg
|
||||
import pytest
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
|
||||
from zhixing_server.bootstrap.config import sqlalchemy_database_url
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MembershipStatus,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
SectorType,
|
||||
StockFactStatus,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.persistence import (
|
||||
MembershipRecord,
|
||||
StockFactRecord,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
|
||||
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
|
||||
PostgresSectorRadarRepository,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2099, 1, 4)
|
||||
STARTED_AT = datetime(2099, 1, 4, 17, 30, tzinfo=UTC)
|
||||
|
||||
|
||||
def prepare_database(database_url: str) -> None:
|
||||
server_root = Path(__file__).parents[2]
|
||||
config = Config(str(server_root / "alembic.ini"))
|
||||
sqlalchemy_url = sqlalchemy_database_url(database_url)
|
||||
config.set_main_option("sqlalchemy.url", sqlalchemy_url.replace("%", "%%"))
|
||||
command.upgrade(config, "head")
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
def test_postgres_sector_radar_revisions_and_last_good() -> None:
|
||||
database_url = os.getenv("ZHIXING_TEST_DATABASE_URL")
|
||||
if not database_url:
|
||||
pytest.skip("set ZHIXING_TEST_DATABASE_URL to run PostgreSQL integration tests")
|
||||
prepare_database(database_url)
|
||||
snapshot = build_source_snapshot(
|
||||
api_name="dc_member",
|
||||
params={"trade_date": "20990104"},
|
||||
rows=(
|
||||
{
|
||||
"trade_date": "20990104",
|
||||
"ts_code": "BKTEST.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "测试股票",
|
||||
},
|
||||
),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=STARTED_AT,
|
||||
)
|
||||
fact_revision = "b" * 64
|
||||
publication_ids = ("test-sector-radar-success", "test-sector-radar-failed")
|
||||
with psycopg.connect(database_url) as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"DELETE FROM sector_radar_publication WHERE id = ANY(%s)",
|
||||
(list(publication_ids),),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s",
|
||||
(fact_revision,),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM sector_radar_source_snapshot WHERE id = %s",
|
||||
(snapshot.snapshot_id,),
|
||||
)
|
||||
|
||||
repository = PostgresSectorRadarRepository(database_url, max_connections=2)
|
||||
try:
|
||||
assert repository.save_source_snapshots((snapshot,)).inserted == 1
|
||||
assert (
|
||||
repository.save_source_snapshots(
|
||||
(replace(snapshot, observed_at=STARTED_AT + timedelta(minutes=1)),)
|
||||
).unchanged
|
||||
== 1
|
||||
)
|
||||
assert (
|
||||
repository.save_memberships(
|
||||
(
|
||||
MembershipRecord(
|
||||
source_snapshot_id=snapshot.snapshot_id,
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.CONCEPT,
|
||||
sector_code="BKTEST.DC",
|
||||
sector_name="测试概念",
|
||||
stock_code="000001.SZ",
|
||||
stock_name="测试股票",
|
||||
status=MembershipStatus.AVAILABLE,
|
||||
),
|
||||
)
|
||||
).inserted
|
||||
== 1
|
||||
)
|
||||
assert (
|
||||
repository.save_stock_facts(
|
||||
(
|
||||
StockFactRecord(
|
||||
fact_revision=fact_revision,
|
||||
source_snapshot_ids=(snapshot.snapshot_id,),
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000001.SZ",
|
||||
status=StockFactStatus.AVAILABLE,
|
||||
turnover_yuan=Decimal("1000"),
|
||||
net_amount_yuan=Decimal("100"),
|
||||
),
|
||||
)
|
||||
).inserted
|
||||
== 1
|
||||
)
|
||||
|
||||
running = RadarPublication(
|
||||
publication_id=publication_ids[0],
|
||||
target_trade_date=TARGET_DATE,
|
||||
status=PublicationStatus.RUNNING,
|
||||
source_version="tushare-pro-v1",
|
||||
universe_version=snapshot.content_sha256,
|
||||
metric_versions=("zhixing_amount_net_bn_v1",),
|
||||
input_hash=None,
|
||||
coverage=Decimal(0),
|
||||
started_at=STARTED_AT,
|
||||
)
|
||||
repository.create_publication(running)
|
||||
repository.finish_publication(
|
||||
replace(
|
||||
running,
|
||||
status=PublicationStatus.SUCCESS,
|
||||
input_hash="a" * 64,
|
||||
coverage=Decimal(1),
|
||||
finished_at=STARTED_AT + timedelta(minutes=5),
|
||||
)
|
||||
)
|
||||
failed = replace(
|
||||
running,
|
||||
publication_id=publication_ids[1],
|
||||
started_at=STARTED_AT + timedelta(minutes=6),
|
||||
)
|
||||
repository.create_publication(failed)
|
||||
repository.finish_publication(
|
||||
replace(
|
||||
failed,
|
||||
status=PublicationStatus.FAILED,
|
||||
coverage=Decimal("0.8"),
|
||||
finished_at=STARTED_AT + timedelta(minutes=7),
|
||||
error_summary="safe_error",
|
||||
)
|
||||
)
|
||||
|
||||
last_good = repository.get_last_good_publication(TARGET_DATE)
|
||||
|
||||
assert last_good is not None
|
||||
assert last_good.publication_id == publication_ids[0]
|
||||
finally:
|
||||
repository.close()
|
||||
with psycopg.connect(database_url) as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"DELETE FROM sector_radar_publication WHERE id = ANY(%s)",
|
||||
(list(publication_ids),),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM sector_radar_stock_fact WHERE fact_revision = %s",
|
||||
(fact_revision,),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM sector_radar_source_snapshot WHERE id = %s",
|
||||
(snapshot.snapshot_id,),
|
||||
)
|
||||
@@ -0,0 +1,230 @@
|
||||
from dataclasses import replace
|
||||
from datetime import UTC, date, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from pydantic import ValidationError
|
||||
|
||||
from zhixing_server.bootstrap.app import create_app
|
||||
from zhixing_server.modules.sector_radar.application.read import (
|
||||
RadarDateIndex,
|
||||
RadarMetricDefinition,
|
||||
RadarQuery,
|
||||
RadarView,
|
||||
RankingPage,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.metrics import AmountNetStrategy
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MetricKind,
|
||||
MetricObservation,
|
||||
MetricQuality,
|
||||
MetricUnit,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
RankChange,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
|
||||
SectorRadarRepositoryError,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.presentation.http import (
|
||||
RadarRankingRowResponse,
|
||||
get_sector_radar_reader,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
||||
|
||||
|
||||
def _publication(
|
||||
publication_id: str,
|
||||
status: PublicationStatus,
|
||||
*,
|
||||
trade_date: date = TARGET_DATE,
|
||||
) -> RadarPublication:
|
||||
return RadarPublication(
|
||||
publication_id=publication_id,
|
||||
target_trade_date=trade_date,
|
||||
status=status,
|
||||
source_version="tushare-pro-v1",
|
||||
universe_version="eastmoney-dc-v1",
|
||||
metric_versions=(AmountNetStrategy.metric_version,),
|
||||
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
|
||||
coverage=Decimal(1) if status is PublicationStatus.SUCCESS else Decimal("0.8"),
|
||||
started_at=NOW,
|
||||
finished_at=None if status is PublicationStatus.RUNNING else NOW + timedelta(minutes=5),
|
||||
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
|
||||
)
|
||||
|
||||
|
||||
def _ranking() -> RankedMetric:
|
||||
return RankedMetric(
|
||||
observation=MetricObservation(
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.CONCEPT,
|
||||
sector_code="BK0001.DC",
|
||||
sector_name="机器人",
|
||||
metric_kind=MetricKind.AMOUNT,
|
||||
metric_version=AmountNetStrategy.metric_version,
|
||||
implementation_kind="independent",
|
||||
unit=MetricUnit.CNY_100M,
|
||||
value=Decimal("12.5"),
|
||||
quality=MetricQuality.AVAILABLE,
|
||||
member_count=20,
|
||||
valid_sample_count=19,
|
||||
membership_coverage=Decimal(1),
|
||||
moneyflow_coverage=Decimal("0.95"),
|
||||
),
|
||||
rank_position=1,
|
||||
rank_percentile=Decimal(100),
|
||||
rank_changes=(RankChange(days=5, value=3),),
|
||||
)
|
||||
|
||||
|
||||
class FakeReader:
|
||||
def __init__(self, *, no_data: bool = False, fail: bool = False) -> None:
|
||||
self.fail = fail
|
||||
self.last_query: RadarQuery | None = None
|
||||
success = _publication("publication-success", PublicationStatus.SUCCESS)
|
||||
current = _publication(
|
||||
"publication-partial",
|
||||
PublicationStatus.PARTIAL,
|
||||
trade_date=TARGET_DATE + timedelta(days=1),
|
||||
)
|
||||
self.date_index = RadarDateIndex(
|
||||
available_dates=() if no_data else (TARGET_DATE,),
|
||||
current_attempt=None if no_data else current,
|
||||
last_good=None if no_data else success,
|
||||
)
|
||||
query = RadarQuery()
|
||||
self.page = RankingPage(
|
||||
status="no_data" if no_data else "success",
|
||||
query=query,
|
||||
publication=None if no_data else success,
|
||||
definition=RadarMetricDefinition(
|
||||
metric_kind=MetricKind.AMOUNT,
|
||||
metric_version=AmountNetStrategy.metric_version,
|
||||
label="主力净流入(知行独立实现)",
|
||||
unit=MetricUnit.CNY_100M,
|
||||
),
|
||||
rows=() if no_data else (_ranking(),),
|
||||
total=0 if no_data else 1,
|
||||
)
|
||||
|
||||
def list_dates(self) -> RadarDateIndex:
|
||||
if self.fail:
|
||||
raise SectorRadarRepositoryError("private database detail")
|
||||
return self.date_index
|
||||
|
||||
def query(self, query: RadarQuery) -> RankingPage:
|
||||
if self.fail:
|
||||
raise SectorRadarRepositoryError("private database detail")
|
||||
self.last_query = query
|
||||
return replace(self.page, query=query)
|
||||
|
||||
|
||||
def _client(reader: FakeReader) -> TestClient:
|
||||
application = create_app()
|
||||
application.dependency_overrides[get_sector_radar_reader] = lambda: reader
|
||||
return TestClient(application)
|
||||
|
||||
|
||||
def test_dates_exposes_partial_attempt_without_replacing_last_good() -> None:
|
||||
response = _client(FakeReader()).get("/api/v1/sector-radar/dates")
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = response.json()
|
||||
assert payload["status"] == "success"
|
||||
assert payload["available_dates"] == ["2026-08-28"]
|
||||
assert payload["current_attempt"]["status"] == "partial"
|
||||
assert payload["last_good"]["status"] == "success"
|
||||
assert payload["last_good"]["coverage"] == "1"
|
||||
|
||||
|
||||
def test_rankings_maps_filters_and_independent_metric_contract() -> None:
|
||||
reader = FakeReader()
|
||||
response = _client(reader).get(
|
||||
"/api/v1/sector-radar/rankings",
|
||||
params={
|
||||
"trade_date": "2026-08-28",
|
||||
"sector_type": "concept",
|
||||
"view": "rank_change",
|
||||
"rank_change_metric": "amount",
|
||||
"rank_change_days": 5,
|
||||
"side": "top",
|
||||
"search": " 机器人 ",
|
||||
"page": 2,
|
||||
"page_size": 10,
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert reader.last_query == RadarQuery(
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.CONCEPT,
|
||||
view=RadarView.RANK_CHANGE,
|
||||
rank_change_metric=MetricKind.AMOUNT,
|
||||
rank_change_days=5,
|
||||
side=RankSide.TOP,
|
||||
search="机器人",
|
||||
page=2,
|
||||
page_size=10,
|
||||
)
|
||||
payload = response.json()
|
||||
assert payload["definition"]["metric_version"] == "zhixing_amount_net_bn_v1"
|
||||
assert payload["definition"]["implementation_kind"] == "independent"
|
||||
assert "知行独立实现" in payload["definition"]["disclaimer"]
|
||||
assert payload["rows"][0]["rank_change"] == 3
|
||||
assert payload["rows"][0]["unit"] == "CNY_100M"
|
||||
|
||||
|
||||
def test_no_data_is_a_stable_200_response() -> None:
|
||||
client = _client(FakeReader(no_data=True))
|
||||
|
||||
dates = client.get("/api/v1/sector-radar/dates")
|
||||
rankings = client.get("/api/v1/sector-radar/rankings")
|
||||
|
||||
assert dates.status_code == 200
|
||||
assert dates.json()["status"] == "no_data"
|
||||
assert rankings.status_code == 200
|
||||
assert rankings.json()["status"] == "no_data"
|
||||
assert rankings.json()["rows"] == []
|
||||
|
||||
|
||||
def test_http_contract_rejects_zero_rank_percentile() -> None:
|
||||
payload = _client(FakeReader()).get("/api/v1/sector-radar/rankings").json()["rows"][0]
|
||||
payload["rank_percentile"] = "0"
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
RadarRankingRowResponse.model_validate(payload)
|
||||
|
||||
|
||||
def test_invalid_query_values_return_422() -> None:
|
||||
client = _client(FakeReader())
|
||||
|
||||
for params in (
|
||||
{"rank_change_days": 0},
|
||||
{"rank_change_days": 6},
|
||||
{"page": 0},
|
||||
{"page_size": 101},
|
||||
{"sector_type": "region"},
|
||||
{"view": "unknown"},
|
||||
{"side": "unknown"},
|
||||
):
|
||||
assert client.get("/api/v1/sector-radar/rankings", params=params).status_code == 422
|
||||
|
||||
|
||||
def test_repository_error_maps_to_redacted_503() -> None:
|
||||
response = _client(FakeReader(fail=True)).get("/api/v1/sector-radar/rankings")
|
||||
|
||||
assert response.status_code == 503
|
||||
assert response.json() == {
|
||||
"detail": {
|
||||
"code": "sector_radar_storage_unavailable",
|
||||
"message": "sector radar storage is unavailable",
|
||||
}
|
||||
}
|
||||
assert "private database detail" not in response.text
|
||||
@@ -1,3 +1,4 @@
|
||||
import threading
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
@@ -87,6 +88,109 @@ def test_rate_limit_cooldown_is_shared_by_following_requests() -> None:
|
||||
assert waits == [60]
|
||||
|
||||
|
||||
def test_request_start_interval_allows_overlapping_provider_calls() -> None:
|
||||
current = [0.0]
|
||||
state_lock = threading.Lock()
|
||||
first_started = threading.Event()
|
||||
release_first = threading.Event()
|
||||
waits: list[float] = []
|
||||
starts: list[tuple[str, float]] = []
|
||||
errors: list[BaseException] = []
|
||||
|
||||
def clock() -> float:
|
||||
with state_lock:
|
||||
return current[0]
|
||||
|
||||
def wait(seconds: float) -> None:
|
||||
with state_lock:
|
||||
waits.append(seconds)
|
||||
current[0] += seconds
|
||||
|
||||
coordinator = RequestCoordinator(
|
||||
max_retries=0,
|
||||
request_interval_seconds=0.2,
|
||||
clock=clock,
|
||||
wait_fn=wait,
|
||||
sleep_fn=wait,
|
||||
)
|
||||
|
||||
def first_request() -> object:
|
||||
starts.append(("first", clock()))
|
||||
first_started.set()
|
||||
if not release_first.wait(timeout=2):
|
||||
raise AssertionError("first provider call was not released")
|
||||
return "first"
|
||||
|
||||
def run_first() -> None:
|
||||
try:
|
||||
coordinator.call("first", first_request)
|
||||
except BaseException as exc: # pragma: no cover - surfaced by the assertion below
|
||||
errors.append(exc)
|
||||
|
||||
first_thread = threading.Thread(target=run_first)
|
||||
first_thread.start()
|
||||
assert first_started.wait(timeout=2)
|
||||
|
||||
second = coordinator.call(
|
||||
"second",
|
||||
lambda: starts.append(("second", clock())) or "second",
|
||||
)
|
||||
|
||||
assert second == "second"
|
||||
assert first_thread.is_alive()
|
||||
release_first.set()
|
||||
first_thread.join(timeout=2)
|
||||
assert not first_thread.is_alive()
|
||||
assert errors == []
|
||||
assert starts == [("first", 0.0), ("second", 0.2)]
|
||||
assert waits == [0.2]
|
||||
|
||||
|
||||
def test_request_start_interval_is_disabled_by_default() -> None:
|
||||
waits: list[float] = []
|
||||
starts: list[str] = []
|
||||
coordinator = RequestCoordinator(
|
||||
max_retries=0,
|
||||
clock=lambda: 0.0,
|
||||
wait_fn=waits.append,
|
||||
)
|
||||
|
||||
coordinator.call("first", lambda: starts.append("first"))
|
||||
coordinator.call("second", lambda: starts.append("second"))
|
||||
|
||||
assert starts == ["first", "second"]
|
||||
assert waits == []
|
||||
|
||||
|
||||
def test_request_start_interval_applies_to_retry_attempts() -> None:
|
||||
current = [0.0]
|
||||
waits: list[float] = []
|
||||
starts: list[float] = []
|
||||
|
||||
def wait(seconds: float) -> None:
|
||||
waits.append(seconds)
|
||||
current[0] += seconds
|
||||
|
||||
coordinator = RequestCoordinator(
|
||||
max_retries=1,
|
||||
backoff_seconds=0,
|
||||
request_interval_seconds=0.2,
|
||||
clock=lambda: current[0],
|
||||
wait_fn=wait,
|
||||
sleep_fn=wait,
|
||||
)
|
||||
|
||||
def request() -> object:
|
||||
starts.append(current[0])
|
||||
if len(starts) == 1:
|
||||
raise RuntimeError("transient provider failure")
|
||||
return "ok"
|
||||
|
||||
assert coordinator.call("daily", request) == "ok"
|
||||
assert starts == [0.0, 0.2]
|
||||
assert waits == [0.0, 0.2]
|
||||
|
||||
|
||||
def test_pro_bar_qfq_calls_are_bound_to_the_shared_coordinator(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
|
||||
@@ -0,0 +1,641 @@
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import replace
|
||||
from datetime import UTC, date, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.sector_radar.application.build import (
|
||||
BuildSectorRadar,
|
||||
BuildSectorRadarCommand,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MembershipStatus,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
SectorType,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.persistence import PublicationSourceGroup
|
||||
from zhixing_server.modules.sector_radar.domain.source import (
|
||||
CapabilityProbeResult,
|
||||
DailyRow,
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceContractError,
|
||||
SourceResult,
|
||||
SourceSnapshot,
|
||||
StockBasicRow,
|
||||
SuspendRow,
|
||||
TradeCalendarRow,
|
||||
build_source_snapshot,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.infrastructure.memory import (
|
||||
InMemorySectorRadarRepository,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
||||
|
||||
|
||||
class FakeRadarSource:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
missing_moneyflow: bool = False,
|
||||
missing_membership: bool = False,
|
||||
net_scale: Decimal = Decimal(1),
|
||||
) -> None:
|
||||
self.missing_moneyflow = missing_moneyflow
|
||||
self.missing_membership = missing_membership
|
||||
self.net_scale = net_scale
|
||||
self.fail_daily = False
|
||||
self.calls: list[str] = []
|
||||
self.moneyflow_candidate_codes: list[tuple[str, ...]] = []
|
||||
|
||||
def _result[T](
|
||||
self, api_name: str, target: date | None, rows: tuple[T, ...]
|
||||
) -> SourceResult[T]:
|
||||
snapshot = build_source_snapshot(
|
||||
api_name=api_name,
|
||||
params={
|
||||
"trade_date": target.isoformat() if target is not None else "all",
|
||||
"fixture_fingerprint": repr(rows),
|
||||
},
|
||||
rows=tuple(self._raw_row(row) for row in rows),
|
||||
target_trade_date=target,
|
||||
partition_key="all" if api_name == "dc_member" else None,
|
||||
observed_at=NOW,
|
||||
)
|
||||
return SourceResult((snapshot,), rows)
|
||||
|
||||
@staticmethod
|
||||
def _raw_row(row: object) -> dict[str, object]:
|
||||
if isinstance(row, TradeCalendarRow):
|
||||
return {
|
||||
"exchange": row.exchange,
|
||||
"cal_date": row.cal_date,
|
||||
"is_open": int(row.is_open),
|
||||
"pretrade_date": row.pretrade_date,
|
||||
}
|
||||
if isinstance(row, SectorIndexRow):
|
||||
return {
|
||||
"trade_date": row.trade_date,
|
||||
"ts_code": row.sector_code,
|
||||
"name": row.name,
|
||||
"level": row.level,
|
||||
"pct_change": row.pct_change,
|
||||
"leading_code": row.leading_code,
|
||||
}
|
||||
if isinstance(row, SectorMemberRow):
|
||||
return {
|
||||
"trade_date": row.trade_date,
|
||||
"ts_code": row.sector_code,
|
||||
"con_code": row.stock_code,
|
||||
"name": row.stock_name,
|
||||
}
|
||||
if isinstance(row, StockBasicRow):
|
||||
return {
|
||||
"ts_code": row.ts_code,
|
||||
"symbol": row.symbol,
|
||||
"name": row.name,
|
||||
"market": row.market,
|
||||
"exchange": row.exchange,
|
||||
"list_status": row.list_status,
|
||||
"list_date": row.list_date,
|
||||
"delist_date": row.delist_date,
|
||||
}
|
||||
if isinstance(row, SuspendRow):
|
||||
return {
|
||||
"ts_code": row.ts_code,
|
||||
"trade_date": row.trade_date,
|
||||
"suspend_timing": row.suspend_timing,
|
||||
"suspend_type": row.suspend_type,
|
||||
}
|
||||
if isinstance(row, DailyRow):
|
||||
return {
|
||||
"ts_code": row.ts_code,
|
||||
"trade_date": row.trade_date,
|
||||
"close": row.close,
|
||||
"pre_close": row.pre_close,
|
||||
"pct_chg": row.pct_chg,
|
||||
"vol": row.volume,
|
||||
"amount": row.amount_thousand_yuan,
|
||||
}
|
||||
if isinstance(row, MoneyflowDcRow):
|
||||
return {
|
||||
"trade_date": row.trade_date,
|
||||
"ts_code": row.ts_code,
|
||||
"name": row.name,
|
||||
"net_amount": row.net_amount_ten_thousand_yuan,
|
||||
"net_amount_rate": row.net_amount_rate,
|
||||
"pct_change": row.pct_change,
|
||||
"close": row.close,
|
||||
}
|
||||
raise TypeError(f"unsupported fake source row: {type(row).__name__}")
|
||||
|
||||
def fetch_trade_calendar(self, start: date, end: date) -> SourceResult[TradeCalendarRow]:
|
||||
self.calls.append("calendar")
|
||||
rows = tuple(
|
||||
TradeCalendarRow("SSE", start + timedelta(days=offset), True, None)
|
||||
for offset in range((end - start).days + 1)
|
||||
)
|
||||
return self._result("trade_cal", end, rows)
|
||||
|
||||
def fetch_sector_indices(
|
||||
self, trade_date: date, sector_type: SectorType
|
||||
) -> SourceResult[SectorIndexRow]:
|
||||
self.calls.append(f"{sector_type.value}_indices")
|
||||
prefix = "BK0" if sector_type is SectorType.CONCEPT else "BK1"
|
||||
row = SectorIndexRow(
|
||||
trade_date,
|
||||
sector_type,
|
||||
f"{prefix}001.DC",
|
||||
"示例概念" if sector_type is SectorType.CONCEPT else "示例行业",
|
||||
"一级",
|
||||
Decimal(1),
|
||||
"000001.SZ",
|
||||
)
|
||||
return self._result(f"dc_index_{sector_type.value}", trade_date, (row,))
|
||||
|
||||
def fetch_sector_members(
|
||||
self, trade_date: date, sector_codes: Sequence[str]
|
||||
) -> SourceResult[SectorMemberRow]:
|
||||
self.calls.append("members")
|
||||
if self.missing_membership:
|
||||
snapshots: list[SourceSnapshot] = []
|
||||
member_rows: list[SectorMemberRow] = []
|
||||
for sector_code in sector_codes:
|
||||
sector_rows = (
|
||||
()
|
||||
if sector_code == sector_codes[-1]
|
||||
else tuple(
|
||||
SectorMemberRow(
|
||||
trade_date,
|
||||
sector_code,
|
||||
f"00000{index}.SZ",
|
||||
f"股票{index}",
|
||||
)
|
||||
for index in range(1, 6)
|
||||
)
|
||||
)
|
||||
snapshots.append(
|
||||
build_source_snapshot(
|
||||
api_name="dc_member",
|
||||
params={
|
||||
"trade_date": trade_date.isoformat(),
|
||||
"ts_code": sector_code,
|
||||
},
|
||||
rows=tuple(self._raw_row(row) for row in sector_rows),
|
||||
target_trade_date=trade_date,
|
||||
partition_key=sector_code,
|
||||
observed_at=NOW,
|
||||
)
|
||||
)
|
||||
member_rows.extend(sector_rows)
|
||||
return SourceResult(tuple(snapshots), tuple(member_rows))
|
||||
rows = tuple(
|
||||
SectorMemberRow(trade_date, sector_code, f"00000{index}.SZ", f"股票{index}")
|
||||
for sector_code in sector_codes
|
||||
for index in range(1, 6)
|
||||
)
|
||||
return self._result("dc_member", trade_date, rows)
|
||||
|
||||
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
|
||||
self.calls.append("stock_basics")
|
||||
rows = tuple(
|
||||
StockBasicRow(
|
||||
f"00000{index}.SZ",
|
||||
f"00000{index}",
|
||||
f"股票{index}",
|
||||
"主板",
|
||||
"SZSE",
|
||||
"L",
|
||||
date(2020, 1, 1),
|
||||
None,
|
||||
)
|
||||
for index in range(1, 6)
|
||||
)
|
||||
return self._result("stock_basic", None, rows)
|
||||
|
||||
def fetch_suspensions(self, trade_date: date) -> SourceResult[SuspendRow]:
|
||||
self.calls.append("suspensions")
|
||||
return self._result("suspend_d", trade_date, ())
|
||||
|
||||
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
|
||||
self.calls.append("daily")
|
||||
if self.fail_daily:
|
||||
raise RuntimeError("private provider detail")
|
||||
rows = tuple(
|
||||
DailyRow(
|
||||
f"00000{index}.SZ",
|
||||
trade_date,
|
||||
Decimal(10),
|
||||
Decimal(10),
|
||||
Decimal(0),
|
||||
Decimal(100),
|
||||
Decimal(1000),
|
||||
)
|
||||
for index in range(1, 6)
|
||||
)
|
||||
return self._result("daily", trade_date, rows)
|
||||
|
||||
def fetch_moneyflow_dc(
|
||||
self,
|
||||
trade_date: date,
|
||||
candidate_codes: Sequence[str],
|
||||
) -> SourceResult[MoneyflowDcRow]:
|
||||
self.calls.append("moneyflow_dc")
|
||||
self.moneyflow_candidate_codes.append(tuple(candidate_codes))
|
||||
count = 4 if self.missing_moneyflow else 5
|
||||
rows = tuple(
|
||||
MoneyflowDcRow(
|
||||
trade_date,
|
||||
f"00000{index}.SZ",
|
||||
f"股票{index}",
|
||||
Decimal(index) * self.net_scale,
|
||||
Decimal(0),
|
||||
Decimal(0),
|
||||
Decimal(10),
|
||||
)
|
||||
for index in range(1, count + 1)
|
||||
)
|
||||
return self._result("moneyflow_dc", trade_date, rows)
|
||||
|
||||
def probe(self, trade_date: date) -> CapabilityProbeResult:
|
||||
return CapabilityProbeResult(NOW, ())
|
||||
|
||||
|
||||
def test_successful_build_is_idempotent_and_failed_retry_preserves_last_good() -> None:
|
||||
source = FakeRadarSource()
|
||||
repository = InMemorySectorRadarRepository()
|
||||
use_case = BuildSectorRadar(
|
||||
source,
|
||||
repository,
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: NOW,
|
||||
)
|
||||
|
||||
first = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
repeated = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
source.fail_daily = True
|
||||
failed = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
|
||||
assert first.status == "success"
|
||||
assert first.exit_code == 0
|
||||
assert first.outcomes[0].ranking_count == 6
|
||||
assert repeated.status == "unchanged"
|
||||
assert repeated.outcomes[0].publication_id == first.outcomes[0].publication_id
|
||||
assert failed.status == "failed"
|
||||
assert "private provider detail" not in str(failed.as_dict())
|
||||
last_good = repository.get_last_good_publication()
|
||||
assert last_good is not None
|
||||
assert last_good.publication_id == first.outcomes[0].publication_id
|
||||
assert any(item.status is PublicationStatus.FAILED for item in repository.publications.values())
|
||||
|
||||
|
||||
def test_build_passes_stable_current_listing_member_intersection_to_moneyflow() -> None:
|
||||
class FutureListingSource(FakeRadarSource):
|
||||
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
|
||||
result = super().fetch_stock_basics()
|
||||
rows = result.rows[:-1] + (replace(result.rows[-1], list_date=date(2027, 1, 1)),)
|
||||
return self._result("stock_basic", None, rows)
|
||||
|
||||
source = FutureListingSource()
|
||||
|
||||
summary = BuildSectorRadar(
|
||||
source,
|
||||
InMemorySectorRadarRepository(),
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: NOW,
|
||||
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
|
||||
assert summary.status == "success"
|
||||
assert source.moneyflow_candidate_codes == [
|
||||
("000001.SZ", "000002.SZ", "000003.SZ", "000004.SZ")
|
||||
]
|
||||
|
||||
|
||||
def test_source_contract_failure_is_logged_with_safe_build_context(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
class InvalidDailySource(FakeRadarSource):
|
||||
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]:
|
||||
self.calls.append("daily")
|
||||
raise SourceContractError("daily returned duplicate business keys")
|
||||
|
||||
repository = InMemorySectorRadarRepository()
|
||||
caplog.set_level(
|
||||
logging.ERROR,
|
||||
logger="zhixing_server.modules.sector_radar.application.build",
|
||||
)
|
||||
|
||||
failed = BuildSectorRadar(
|
||||
InvalidDailySource(),
|
||||
repository,
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: NOW,
|
||||
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
|
||||
messages = "\n".join(record.getMessage() for record in caplog.records)
|
||||
assert failed.status == "failed"
|
||||
assert failed.outcomes[0].error_message == "input or source contract validation failed"
|
||||
assert "sector_radar_source_group_contract_failed" in messages
|
||||
assert "source_group=daily" in messages
|
||||
assert "publication_id=radar-20260828-running-" in messages
|
||||
assert "validation=daily returned duplicate business keys" in messages
|
||||
assert len(caplog.records) == 1
|
||||
|
||||
|
||||
def test_partial_coverage_and_lock_have_distinct_exit_codes() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
partial = BuildSectorRadar(
|
||||
FakeRadarSource(missing_moneyflow=True),
|
||||
repository,
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: NOW,
|
||||
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
publication_count = len(repository.publications)
|
||||
partial_repeated = BuildSectorRadar(
|
||||
FakeRadarSource(missing_moneyflow=True),
|
||||
repository,
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: NOW,
|
||||
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
|
||||
repository.lock_available = False
|
||||
locked = BuildSectorRadar(
|
||||
FakeRadarSource(),
|
||||
repository,
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: NOW,
|
||||
).execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
|
||||
assert partial.status == "partial"
|
||||
assert partial.exit_code == 2
|
||||
assert partial.outcomes[0].coverage == Decimal("0.8")
|
||||
assert partial_repeated.status == "partial"
|
||||
assert len(repository.publications) == publication_count
|
||||
assert repository.get_last_good_publication() is None
|
||||
assert {
|
||||
record.source_group
|
||||
for record in repository.load_publication_sources(partial.outcomes[0].publication_id or "")
|
||||
if record.refresh_on_retry
|
||||
} == {PublicationSourceGroup.MONEYFLOW_DC}
|
||||
assert locked.status == "failed"
|
||||
assert locked.exit_code == 1
|
||||
assert locked.outcomes[0].status == "locked"
|
||||
|
||||
|
||||
def test_unknown_membership_is_persisted_as_partial_and_retried_independently() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
source = FakeRadarSource(missing_membership=True)
|
||||
use_case = BuildSectorRadar(source, repository, today=TARGET_DATE, now_fn=lambda: NOW)
|
||||
|
||||
partial = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
partial_id = partial.outcomes[0].publication_id
|
||||
assert partial_id is not None
|
||||
publication = repository.get_publication(partial_id)
|
||||
assert partial.status == "partial"
|
||||
assert publication is not None
|
||||
assert publication.error_summary == "membership_unknown"
|
||||
assert any(item.status is MembershipStatus.UNKNOWN for item in repository.memberships.values())
|
||||
assert {
|
||||
record.source_group
|
||||
for record in repository.load_publication_sources(partial_id)
|
||||
if record.refresh_on_retry
|
||||
} == {PublicationSourceGroup.MEMBERS}
|
||||
|
||||
source.missing_membership = False
|
||||
source.calls.clear()
|
||||
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
|
||||
|
||||
assert retried.status == "success"
|
||||
assert source.calls == ["members"]
|
||||
|
||||
|
||||
def test_membership_retry_refreshes_moneyflow_when_replay_misses_new_candidates() -> None:
|
||||
class ExpandingMembershipSource(FakeRadarSource):
|
||||
def fetch_sector_members(
|
||||
self,
|
||||
trade_date: date,
|
||||
sector_codes: Sequence[str],
|
||||
) -> SourceResult[SectorMemberRow]:
|
||||
self.calls.append("members")
|
||||
rows: list[SectorMemberRow] = []
|
||||
snapshots: list[SourceSnapshot] = []
|
||||
for index, sector_code in enumerate(sector_codes, start=1):
|
||||
sector_rows = (
|
||||
()
|
||||
if self.missing_membership and index == len(sector_codes)
|
||||
else (
|
||||
SectorMemberRow(
|
||||
trade_date,
|
||||
sector_code,
|
||||
f"00000{index}.SZ",
|
||||
f"股票{index}",
|
||||
),
|
||||
)
|
||||
)
|
||||
snapshots.append(
|
||||
build_source_snapshot(
|
||||
api_name="dc_member",
|
||||
params={
|
||||
"trade_date": trade_date.isoformat(),
|
||||
"ts_code": sector_code,
|
||||
},
|
||||
rows=tuple(self._raw_row(row) for row in sector_rows),
|
||||
target_trade_date=trade_date,
|
||||
partition_key=sector_code,
|
||||
observed_at=NOW,
|
||||
)
|
||||
)
|
||||
rows.extend(sector_rows)
|
||||
return SourceResult(tuple(snapshots), tuple(rows))
|
||||
|
||||
def fetch_moneyflow_dc(
|
||||
self,
|
||||
trade_date: date,
|
||||
candidate_codes: Sequence[str],
|
||||
) -> SourceResult[MoneyflowDcRow]:
|
||||
self.calls.append("moneyflow_dc")
|
||||
self.moneyflow_candidate_codes.append(tuple(candidate_codes))
|
||||
rows = tuple(
|
||||
MoneyflowDcRow(
|
||||
trade_date,
|
||||
code,
|
||||
code,
|
||||
Decimal(1),
|
||||
Decimal(0),
|
||||
Decimal(0),
|
||||
Decimal(10),
|
||||
)
|
||||
for code in candidate_codes
|
||||
)
|
||||
return self._result("moneyflow_dc", trade_date, rows)
|
||||
|
||||
repository = InMemorySectorRadarRepository()
|
||||
source = ExpandingMembershipSource(missing_membership=True)
|
||||
use_case = BuildSectorRadar(source, repository, today=TARGET_DATE, now_fn=lambda: NOW)
|
||||
|
||||
partial = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
partial_id = partial.outcomes[0].publication_id
|
||||
assert partial.status == "partial"
|
||||
assert partial_id is not None
|
||||
assert source.moneyflow_candidate_codes == [("000001.SZ",)]
|
||||
|
||||
source.missing_membership = False
|
||||
source.calls.clear()
|
||||
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
|
||||
|
||||
assert retried.status == "success"
|
||||
assert source.calls == ["members", "moneyflow_dc"]
|
||||
assert source.moneyflow_candidate_codes[-1] == ("000001.SZ", "000002.SZ")
|
||||
|
||||
|
||||
def test_range_builds_dates_in_order_and_retry_uses_old_target() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
source = FakeRadarSource(missing_moneyflow=True)
|
||||
use_case = BuildSectorRadar(source, repository, now_fn=lambda: NOW)
|
||||
end = TARGET_DATE + timedelta(days=1)
|
||||
|
||||
summary = use_case.execute(BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end))
|
||||
partial_id = summary.outcomes[0].publication_id
|
||||
assert partial_id is not None
|
||||
source.missing_moneyflow = False
|
||||
source.calls.clear()
|
||||
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=partial_id))
|
||||
|
||||
assert [item.target_trade_date for item in summary.outcomes] == [TARGET_DATE, end]
|
||||
assert all(item.status == "partial" for item in summary.outcomes)
|
||||
assert retried.outcomes[0].target_trade_date == TARGET_DATE
|
||||
assert retried.status == "success"
|
||||
assert source.calls == ["moneyflow_dc"]
|
||||
|
||||
|
||||
def test_failed_retry_reuses_every_completed_source_group() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
source = FakeRadarSource()
|
||||
source.fail_daily = True
|
||||
use_case = BuildSectorRadar(source, repository, now_fn=lambda: NOW)
|
||||
|
||||
failed = use_case.execute(BuildSectorRadarCommand(trade_date=TARGET_DATE))
|
||||
failed_id = failed.outcomes[0].publication_id
|
||||
assert failed_id is not None
|
||||
source.fail_daily = False
|
||||
source.calls.clear()
|
||||
|
||||
retried = use_case.execute(BuildSectorRadarCommand(retry_publication_id=failed_id))
|
||||
|
||||
assert retried.status == "success"
|
||||
assert source.calls == ["daily", "moneyflow_dc"]
|
||||
old_groups = {record.source_group for record in repository.load_publication_sources(failed_id)}
|
||||
assert len(old_groups) == 6
|
||||
|
||||
|
||||
def test_date_lock_recovers_an_orphaned_running_publication() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
stale = RadarPublication(
|
||||
publication_id="stale-running",
|
||||
target_trade_date=TARGET_DATE,
|
||||
status=PublicationStatus.RUNNING,
|
||||
source_version="tushare-pro-v1",
|
||||
universe_version="pending",
|
||||
metric_versions=("zhixing_amount_net_bn_v1",),
|
||||
input_hash=None,
|
||||
coverage=Decimal(0),
|
||||
started_at=NOW - timedelta(hours=1),
|
||||
)
|
||||
repository.create_publication(stale)
|
||||
|
||||
summary = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
|
||||
BuildSectorRadarCommand(trade_date=TARGET_DATE)
|
||||
)
|
||||
|
||||
recovered = repository.get_publication("stale-running")
|
||||
assert summary.status == "success"
|
||||
assert recovered is not None
|
||||
assert recovered.status is PublicationStatus.FAILED
|
||||
assert recovered.error_summary == "recovered_stale_running"
|
||||
|
||||
|
||||
def test_default_target_excludes_today_before_closing_data_is_ready() -> None:
|
||||
before_close = datetime(2026, 8, 28, 6, 0, tzinfo=UTC)
|
||||
after_close = datetime(2026, 8, 28, 8, 0, tzinfo=UTC)
|
||||
|
||||
before = BuildSectorRadar(
|
||||
FakeRadarSource(),
|
||||
InMemorySectorRadarRepository(),
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: before_close,
|
||||
).execute()
|
||||
after = BuildSectorRadar(
|
||||
FakeRadarSource(),
|
||||
InMemorySectorRadarRepository(),
|
||||
today=TARGET_DATE,
|
||||
now_fn=lambda: after_close,
|
||||
).execute()
|
||||
|
||||
assert before.outcomes[0].target_trade_date == TARGET_DATE - timedelta(days=1)
|
||||
assert after.outcomes[0].target_trade_date == TARGET_DATE
|
||||
|
||||
|
||||
def test_tenth_trading_day_publishes_swing_and_five_rank_changes() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
end = TARGET_DATE + timedelta(days=9)
|
||||
summary = BuildSectorRadar(FakeRadarSource(), repository, now_fn=lambda: NOW).execute(
|
||||
BuildSectorRadarCommand(start_date=TARGET_DATE, end_date=end)
|
||||
)
|
||||
|
||||
publication = repository.get_last_good_publication(end)
|
||||
assert summary.status == "success"
|
||||
assert publication is not None
|
||||
current = tuple(
|
||||
record.ranking
|
||||
for record in repository.rankings.values()
|
||||
if record.publication_id == publication.publication_id
|
||||
)
|
||||
swing = tuple(
|
||||
ranking
|
||||
for ranking in current
|
||||
if ranking.observation.metric_version == "zhixing_swing_equal_3_10_v1"
|
||||
)
|
||||
assert len(swing) == 2
|
||||
assert all(ranking.observation.value == Decimal("0.03") for ranking in swing)
|
||||
assert all(
|
||||
tuple(change.value for change in ranking.rank_changes) == (None, None, None, None, None)
|
||||
for ranking in swing
|
||||
)
|
||||
amount = tuple(
|
||||
ranking
|
||||
for ranking in current
|
||||
if ranking.observation.metric_version == "zhixing_amount_net_bn_v1"
|
||||
)
|
||||
assert all(
|
||||
tuple(change.value for change in ranking.rank_changes) == (0, 0, 0, 0, 0)
|
||||
for ranking in amount
|
||||
)
|
||||
|
||||
|
||||
def test_history_uses_latest_successful_input_revision_for_a_date() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
clock = [NOW]
|
||||
first_source = FakeRadarSource(net_scale=Decimal(1))
|
||||
second_source = FakeRadarSource(net_scale=Decimal(2))
|
||||
|
||||
first = BuildSectorRadar(first_source, repository, now_fn=lambda: clock[0]).execute(
|
||||
BuildSectorRadarCommand(trade_date=TARGET_DATE)
|
||||
)
|
||||
clock[0] = NOW + timedelta(minutes=5)
|
||||
second = BuildSectorRadar(second_source, repository, now_fn=lambda: clock[0]).execute(
|
||||
BuildSectorRadarCommand(trade_date=TARGET_DATE)
|
||||
)
|
||||
|
||||
history = repository.load_daily_aggregate_history(
|
||||
TARGET_DATE + timedelta(days=1), limit_dates=1
|
||||
)
|
||||
assert first.status == "success"
|
||||
assert second.status == "success"
|
||||
assert len(history) == 2
|
||||
assert all(item.net_amount_yuan == Decimal(300_000) for item in history)
|
||||
@@ -0,0 +1,129 @@
|
||||
import json
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.sector_radar.application.build import (
|
||||
BuildDateOutcome,
|
||||
BuildOutcomeStatus,
|
||||
BuildSummary,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.presentation import cli
|
||||
from zhixing_server.modules.sector_radar.presentation.cli import build_parser
|
||||
|
||||
|
||||
def test_sector_radar_cli_parses_single_range_and_retry_modes() -> None:
|
||||
parser = build_parser()
|
||||
|
||||
single = parser.parse_args(["--trade-date", "2026-08-28"])
|
||||
date_range = parser.parse_args(["--start-date", "2026-08-18", "--end-date", "2026-08-28"])
|
||||
retry = parser.parse_args(["--retry-publication-id", "publication-a"])
|
||||
|
||||
assert single.trade_date == date(2026, 8, 28)
|
||||
assert date_range.start_date == date(2026, 8, 18)
|
||||
assert date_range.end_date == date(2026, 8, 28)
|
||||
assert retry.retry_publication_id == "publication-a"
|
||||
|
||||
|
||||
class FakeSettings:
|
||||
log_level = "INFO"
|
||||
tushare_token = "secret-token"
|
||||
database_url = "postgresql://unused"
|
||||
sector_radar_max_retries = 3
|
||||
sector_radar_retry_backoff_seconds = 1.0
|
||||
sector_radar_request_interval_seconds = 0.2
|
||||
sector_radar_advisory_lock_key = 7_380_522
|
||||
sector_radar_coverage_threshold = Decimal("0.99")
|
||||
|
||||
|
||||
class FakeRepository:
|
||||
closed = False
|
||||
|
||||
def __init__(self, database_url: str, *, advisory_lock_key: int) -> None:
|
||||
assert database_url == "postgresql://unused"
|
||||
assert advisory_lock_key == 7_380_522
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("outcome_status", "coverage", "expected_code"),
|
||||
(("success", Decimal(1), 0), ("partial", Decimal("0.8"), 2), ("failed", Decimal(0), 1)),
|
||||
)
|
||||
def test_cli_main_returns_summary_exit_code_and_json(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: pytest.CaptureFixture[str],
|
||||
outcome_status: BuildOutcomeStatus,
|
||||
coverage: Decimal,
|
||||
expected_code: int,
|
||||
) -> None:
|
||||
summary = BuildSummary(
|
||||
(
|
||||
BuildDateOutcome(
|
||||
date(2026, 8, 28),
|
||||
outcome_status,
|
||||
"publication-a",
|
||||
coverage,
|
||||
2,
|
||||
6,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
class FakeSourceFactory:
|
||||
@staticmethod
|
||||
def from_token(token: str, **kwargs: object) -> object:
|
||||
assert token == "secret-token"
|
||||
assert kwargs["max_retries"] == 3
|
||||
assert kwargs["backoff_seconds"] == 1.0
|
||||
assert kwargs["request_interval_seconds"] == 0.2
|
||||
return object()
|
||||
|
||||
class FakeBuild:
|
||||
def __init__(self, source: object, repository: object, **kwargs: object) -> None:
|
||||
assert source is not None
|
||||
assert repository is not None
|
||||
assert kwargs
|
||||
|
||||
def execute(self, command: object) -> BuildSummary:
|
||||
assert command is not None
|
||||
return summary
|
||||
|
||||
monkeypatch.setattr(cli, "get_settings", FakeSettings)
|
||||
monkeypatch.setattr(cli, "TushareSectorRadarAdapter", FakeSourceFactory)
|
||||
monkeypatch.setattr(cli, "PostgresSectorRadarRepository", FakeRepository)
|
||||
monkeypatch.setattr(cli, "BuildSectorRadar", FakeBuild)
|
||||
|
||||
exit_code = cli.main(["--trade-date", "2026-08-28"])
|
||||
output = json.loads(capsys.readouterr().out)
|
||||
|
||||
assert exit_code == expected_code
|
||||
assert output["status"] == summary.status
|
||||
assert output["exit_code"] == expected_code
|
||||
assert "secret-token" not in str(output)
|
||||
|
||||
|
||||
def test_cli_initialization_failure_is_redacted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: pytest.CaptureFixture[str],
|
||||
) -> None:
|
||||
class FailingSourceFactory:
|
||||
@staticmethod
|
||||
def from_token(token: str, **kwargs: object) -> object:
|
||||
del token, kwargs
|
||||
raise RuntimeError("private provider detail secret-token")
|
||||
|
||||
monkeypatch.setattr(cli, "get_settings", FakeSettings)
|
||||
monkeypatch.setattr(cli, "TushareSectorRadarAdapter", FailingSourceFactory)
|
||||
|
||||
exit_code = cli.main(["--trade-date", "2026-08-28"])
|
||||
captured = capsys.readouterr()
|
||||
output = json.loads(captured.out)
|
||||
|
||||
assert exit_code == 1
|
||||
assert output["status"] == "failed"
|
||||
assert output["error_type"] == "RuntimeError"
|
||||
assert "private provider detail" not in captured.out
|
||||
assert "secret-token" not in captured.out
|
||||
@@ -0,0 +1,146 @@
|
||||
from datetime import UTC, date, datetime
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.sector_radar.domain.facts import aggregate_sector_snapshot
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MembershipStatus,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
SectorMembershipSnapshot,
|
||||
SectorType,
|
||||
StockDailyFact,
|
||||
StockFactStatus,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
|
||||
|
||||
def test_point_in_time_aggregation_distinguishes_suspension_missing_and_zero() -> None:
|
||||
snapshot = SectorMembershipSnapshot(
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.CONCEPT,
|
||||
sector_code="BK0001.DC",
|
||||
sector_name="示例概念",
|
||||
member_codes=("000001.SZ", "000002.SZ", "000003.SZ", "000004.SZ"),
|
||||
status=MembershipStatus.AVAILABLE,
|
||||
source_version="dc-member-20260828-a",
|
||||
)
|
||||
facts = (
|
||||
StockDailyFact(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000001.SZ",
|
||||
status=StockFactStatus.AVAILABLE,
|
||||
turnover_yuan=Decimal("1000"),
|
||||
net_amount_yuan=Decimal("100"),
|
||||
),
|
||||
StockDailyFact(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000002.SZ",
|
||||
status=StockFactStatus.AVAILABLE,
|
||||
turnover_yuan=Decimal("2000"),
|
||||
net_amount_yuan=Decimal("0"),
|
||||
),
|
||||
StockDailyFact(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000003.SZ",
|
||||
status=StockFactStatus.SUSPENDED,
|
||||
),
|
||||
StockDailyFact(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000004.SZ",
|
||||
status=StockFactStatus.MISSING_MONEYFLOW,
|
||||
),
|
||||
)
|
||||
|
||||
aggregate = aggregate_sector_snapshot(snapshot, facts)
|
||||
|
||||
assert aggregate.member_count == 4
|
||||
assert aggregate.valid_sample_count == 2
|
||||
assert aggregate.net_amount_yuan == Decimal("100")
|
||||
assert aggregate.turnover_yuan == Decimal("3000")
|
||||
assert aggregate.membership_coverage == Decimal("1")
|
||||
assert aggregate.moneyflow_coverage == Decimal("2") / Decimal("3")
|
||||
|
||||
|
||||
def test_unknown_membership_never_falls_back_to_available_stock_facts() -> None:
|
||||
snapshot = SectorMembershipSnapshot(
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.INDUSTRY,
|
||||
sector_code="BK1001.DC",
|
||||
sector_name="示例行业",
|
||||
member_codes=(),
|
||||
status=MembershipStatus.UNKNOWN,
|
||||
source_version="dc-member-missing",
|
||||
)
|
||||
fact = StockDailyFact(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000001.SZ",
|
||||
status=StockFactStatus.AVAILABLE,
|
||||
turnover_yuan=Decimal("1000"),
|
||||
net_amount_yuan=Decimal("100"),
|
||||
)
|
||||
|
||||
aggregate = aggregate_sector_snapshot(snapshot, (fact,))
|
||||
|
||||
assert aggregate.member_count == 0
|
||||
assert aggregate.net_amount_yuan is None
|
||||
assert aggregate.turnover_yuan is None
|
||||
assert aggregate.membership_coverage == Decimal("0")
|
||||
|
||||
|
||||
def test_stock_fact_rejects_non_finite_values_and_invalid_status_payloads() -> None:
|
||||
with pytest.raises(ValueError, match="finite"):
|
||||
StockDailyFact(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000001.SZ",
|
||||
status=StockFactStatus.AVAILABLE,
|
||||
turnover_yuan=Decimal("Infinity"),
|
||||
net_amount_yuan=Decimal("1"),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="must not expose amounts"):
|
||||
StockDailyFact(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code="000001.SZ",
|
||||
status=StockFactStatus.SUSPENDED,
|
||||
turnover_yuan=Decimal("0"),
|
||||
)
|
||||
|
||||
|
||||
def test_publication_requires_terminal_completion_and_replay_identity() -> None:
|
||||
started_at = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
||||
|
||||
publication = RadarPublication(
|
||||
publication_id="radar-20260828-a",
|
||||
target_trade_date=TARGET_DATE,
|
||||
status=PublicationStatus.SUCCESS,
|
||||
source_version="tushare-pro-v1",
|
||||
universe_version="eastmoney-dc-20260828-a",
|
||||
metric_versions=(
|
||||
"zhixing_amount_net_bn_v1",
|
||||
"zhixing_ratio_turnover_v1",
|
||||
"zhixing_swing_equal_3_10_v1",
|
||||
),
|
||||
input_hash="a" * 64,
|
||||
coverage=Decimal("0.995"),
|
||||
started_at=started_at,
|
||||
finished_at=datetime(2026, 8, 28, 17, 35, tzinfo=UTC),
|
||||
)
|
||||
|
||||
assert publication.status is PublicationStatus.SUCCESS
|
||||
|
||||
with pytest.raises(ValueError, match="finished_at"):
|
||||
RadarPublication(
|
||||
publication_id="radar-20260828-running",
|
||||
target_trade_date=TARGET_DATE,
|
||||
status=PublicationStatus.RUNNING,
|
||||
source_version="tushare-pro-v1",
|
||||
universe_version="eastmoney-dc-20260828-a",
|
||||
metric_versions=("zhixing_amount_net_bn_v1",),
|
||||
input_hash=None,
|
||||
coverage=Decimal("0"),
|
||||
started_at=started_at,
|
||||
finished_at=started_at,
|
||||
)
|
||||
@@ -0,0 +1,112 @@
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from zhixing_server.modules.sector_radar.domain.metrics import (
|
||||
AmountNetStrategy,
|
||||
RatioTurnoverStrategy,
|
||||
SwingEqualThreeToTenStrategy,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MetricQuality,
|
||||
SectorDailyAggregate,
|
||||
SectorType,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
|
||||
|
||||
def make_aggregate(
|
||||
*,
|
||||
net_amount_yuan: Decimal | None = Decimal("125000000"),
|
||||
turnover_yuan: Decimal | None = Decimal("5000000000"),
|
||||
) -> SectorDailyAggregate:
|
||||
return SectorDailyAggregate(
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.CONCEPT,
|
||||
sector_code="BK0001.DC",
|
||||
sector_name="示例概念",
|
||||
member_count=10,
|
||||
valid_sample_count=10,
|
||||
net_amount_yuan=net_amount_yuan,
|
||||
turnover_yuan=turnover_yuan,
|
||||
membership_coverage=Decimal("1"),
|
||||
moneyflow_coverage=Decimal("1"),
|
||||
)
|
||||
|
||||
|
||||
def test_amount_and_ratio_strategies_expose_independent_versioned_values() -> None:
|
||||
aggregate = make_aggregate()
|
||||
|
||||
amount = AmountNetStrategy().evaluate((aggregate,), TARGET_DATE)
|
||||
ratio = RatioTurnoverStrategy().evaluate((aggregate,), TARGET_DATE)
|
||||
|
||||
assert amount.value == Decimal("1.25")
|
||||
assert amount.metric_version == "zhixing_amount_net_bn_v1"
|
||||
assert amount.implementation_kind == "independent"
|
||||
assert amount.unit == "CNY_100M"
|
||||
assert amount.quality is MetricQuality.AVAILABLE
|
||||
|
||||
assert ratio.value == Decimal("0.025")
|
||||
assert ratio.metric_version == "zhixing_ratio_turnover_v1"
|
||||
assert ratio.implementation_kind == "independent"
|
||||
assert ratio.unit == "ratio"
|
||||
|
||||
|
||||
def test_missing_moneyflow_is_unavailable_but_zero_remains_a_real_value() -> None:
|
||||
missing = AmountNetStrategy().evaluate((make_aggregate(net_amount_yuan=None),), TARGET_DATE)
|
||||
zero = AmountNetStrategy().evaluate(
|
||||
(make_aggregate(net_amount_yuan=Decimal("0")),), TARGET_DATE
|
||||
)
|
||||
|
||||
assert missing.value is None
|
||||
assert missing.quality is MetricQuality.UNAVAILABLE
|
||||
assert zero.value == Decimal("0")
|
||||
assert zero.quality is MetricQuality.AVAILABLE
|
||||
|
||||
|
||||
def test_swing_strategy_uses_each_days_point_in_time_aggregate() -> None:
|
||||
history = tuple(
|
||||
SectorDailyAggregate(
|
||||
trade_date=date(2026, 8, 18 + offset),
|
||||
sector_type=SectorType.CONCEPT,
|
||||
sector_code="BK0001.DC",
|
||||
sector_name="示例概念",
|
||||
member_count=6 + offset,
|
||||
valid_sample_count=6 + offset,
|
||||
net_amount_yuan=Decimal(str(offset + 1)),
|
||||
turnover_yuan=Decimal("100"),
|
||||
membership_coverage=Decimal("1"),
|
||||
moneyflow_coverage=Decimal("1"),
|
||||
)
|
||||
for offset in range(10)
|
||||
)
|
||||
|
||||
result = SwingEqualThreeToTenStrategy().evaluate(history, date(2026, 8, 27))
|
||||
|
||||
# The worked 3..10-day window ratios average to exactly 0.0725.
|
||||
assert result.value == Decimal("0.0725")
|
||||
assert result.metric_version == "zhixing_swing_equal_3_10_v1"
|
||||
assert result.member_count == 15
|
||||
|
||||
|
||||
def test_swing_strategy_carries_forward_limited_historical_sample_quality() -> None:
|
||||
history = tuple(
|
||||
SectorDailyAggregate(
|
||||
trade_date=date(2026, 8, 18 + offset),
|
||||
sector_type=SectorType.INDUSTRY,
|
||||
sector_code="BK1001.DC",
|
||||
sector_name="示例行业",
|
||||
member_count=10,
|
||||
valid_sample_count=4 if offset == 0 else 10,
|
||||
net_amount_yuan=Decimal("10"),
|
||||
turnover_yuan=Decimal("100"),
|
||||
membership_coverage=Decimal("1"),
|
||||
moneyflow_coverage=Decimal("1"),
|
||||
)
|
||||
for offset in range(10)
|
||||
)
|
||||
|
||||
result = SwingEqualThreeToTenStrategy().evaluate(history, date(2026, 8, 27))
|
||||
|
||||
assert result.value == Decimal("0.1")
|
||||
assert result.quality is MetricQuality.AVAILABLE_LIMITED_SAMPLE
|
||||
@@ -0,0 +1,215 @@
|
||||
from datetime import UTC, date, datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MembershipStatus,
|
||||
SectorType,
|
||||
StockFactStatus,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.normalize import (
|
||||
normalize_memberships,
|
||||
normalize_stock_facts,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.source import (
|
||||
DailyRow,
|
||||
MoneyflowDcRow,
|
||||
SectorIndexRow,
|
||||
SectorMemberRow,
|
||||
SourceResult,
|
||||
StockBasicRow,
|
||||
SuspendRow,
|
||||
build_source_snapshot,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
OBSERVED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
||||
|
||||
|
||||
def result[T](api_name: str, rows: tuple[T, ...]) -> SourceResult[T]:
|
||||
snapshot = build_source_snapshot(
|
||||
api_name=api_name,
|
||||
params={"trade_date": "20260828"},
|
||||
rows=(),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
return SourceResult((snapshot,), rows)
|
||||
|
||||
|
||||
def basic(
|
||||
ts_code: str,
|
||||
*,
|
||||
list_date: date = date(2020, 1, 1),
|
||||
market: str | None = "主板",
|
||||
) -> StockBasicRow:
|
||||
return StockBasicRow(
|
||||
ts_code=ts_code,
|
||||
symbol=ts_code.split(".")[0],
|
||||
name=ts_code,
|
||||
market=market,
|
||||
exchange="SZSE",
|
||||
list_status="L",
|
||||
list_date=list_date,
|
||||
delist_date=None,
|
||||
)
|
||||
|
||||
|
||||
def test_missing_market_keeps_hs_a_stock_but_code_rules_still_exclude_bse_and_b_shares() -> None:
|
||||
codes = ("000001.SZ", "920001.BJ", "200001.SZ", "900001.SH")
|
||||
basics = tuple(basic(code, market=None) for code in codes)
|
||||
daily_rows = tuple(daily(code, Decimal("1")) for code in codes)
|
||||
moneyflow_rows = tuple(moneyflow(code, Decimal("1")) for code in codes)
|
||||
|
||||
facts = normalize_stock_facts(
|
||||
target_trade_date=TARGET_DATE,
|
||||
candidate_codes=codes,
|
||||
stock_basics=result("stock_basic", basics),
|
||||
suspensions=result("suspend_d", ()),
|
||||
daily=result("daily", daily_rows),
|
||||
moneyflow=result("moneyflow_dc", moneyflow_rows),
|
||||
)
|
||||
by_code = {fact.ts_code: fact for fact in facts}
|
||||
|
||||
assert by_code["000001.SZ"].status is StockFactStatus.AVAILABLE
|
||||
assert by_code["920001.BJ"].status is StockFactStatus.LIFECYCLE_INVALID
|
||||
assert by_code["200001.SZ"].status is StockFactStatus.LIFECYCLE_INVALID
|
||||
assert by_code["900001.SH"].status is StockFactStatus.LIFECYCLE_INVALID
|
||||
|
||||
|
||||
def daily(ts_code: str, amount: Decimal | None) -> DailyRow:
|
||||
return DailyRow(
|
||||
ts_code=ts_code,
|
||||
trade_date=TARGET_DATE,
|
||||
close=Decimal("10"),
|
||||
pre_close=Decimal("10"),
|
||||
pct_chg=Decimal(0),
|
||||
volume=Decimal(0),
|
||||
amount_thousand_yuan=amount,
|
||||
)
|
||||
|
||||
|
||||
def moneyflow(ts_code: str, amount: Decimal | None) -> MoneyflowDcRow:
|
||||
return MoneyflowDcRow(
|
||||
trade_date=TARGET_DATE,
|
||||
ts_code=ts_code,
|
||||
name=ts_code,
|
||||
net_amount_ten_thousand_yuan=amount,
|
||||
net_amount_rate=Decimal(0),
|
||||
pct_change=Decimal(0),
|
||||
close=Decimal("10"),
|
||||
)
|
||||
|
||||
|
||||
def test_membership_normalization_persists_an_explicit_unknown_sector() -> None:
|
||||
indices = (
|
||||
SectorIndexRow(
|
||||
TARGET_DATE,
|
||||
SectorType.CONCEPT,
|
||||
"BK0001.DC",
|
||||
"机器人",
|
||||
"一级",
|
||||
Decimal(1),
|
||||
None,
|
||||
),
|
||||
SectorIndexRow(
|
||||
TARGET_DATE,
|
||||
SectorType.CONCEPT,
|
||||
"BK0002.DC",
|
||||
"低空经济",
|
||||
"一级",
|
||||
Decimal(1),
|
||||
None,
|
||||
),
|
||||
)
|
||||
member = SectorMemberRow(
|
||||
TARGET_DATE,
|
||||
"BK0001.DC",
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
)
|
||||
all_snapshot = build_source_snapshot(
|
||||
api_name="dc_member",
|
||||
params={"trade_date": "20260828"},
|
||||
rows=(
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": member.sector_code,
|
||||
"con_code": member.stock_code,
|
||||
"name": member.stock_name,
|
||||
},
|
||||
),
|
||||
target_trade_date=TARGET_DATE,
|
||||
partition_key="all",
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
empty_partition = build_source_snapshot(
|
||||
api_name="dc_member",
|
||||
params={"trade_date": "20260828", "ts_code": "BK0002.DC"},
|
||||
rows=(),
|
||||
target_trade_date=TARGET_DATE,
|
||||
partition_key="BK0002.DC",
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
records = normalize_memberships(
|
||||
indices,
|
||||
SourceResult((all_snapshot, empty_partition), (member,)),
|
||||
)
|
||||
|
||||
assert records[0].status is MembershipStatus.AVAILABLE
|
||||
assert records[0].stock_code == "000001.SZ"
|
||||
assert records[1].status is MembershipStatus.UNKNOWN
|
||||
assert records[1].stock_code is None
|
||||
assert records[1].membership_key == "__membership_unknown__"
|
||||
|
||||
|
||||
def test_stock_fact_normalization_preserves_all_missing_and_zero_states() -> None:
|
||||
codes = tuple(f"00000{index}.SZ" for index in range(1, 9))
|
||||
basics = tuple(
|
||||
basic(code, list_date=date(2027, 1, 1) if code == codes[7] else date(2020, 1, 1))
|
||||
for code in codes
|
||||
)
|
||||
daily_rows = (
|
||||
daily(codes[0], Decimal("1")),
|
||||
daily(codes[3], None),
|
||||
daily(codes[4], Decimal("1")),
|
||||
daily(codes[5], Decimal("1")),
|
||||
daily(codes[6], Decimal("0")),
|
||||
daily(codes[7], Decimal("1")),
|
||||
)
|
||||
moneyflow_rows = (
|
||||
moneyflow(codes[0], Decimal("0")),
|
||||
moneyflow(codes[3], Decimal("1")),
|
||||
moneyflow(codes[5], None),
|
||||
moneyflow(codes[6], Decimal("0")),
|
||||
moneyflow(codes[7], Decimal("1")),
|
||||
)
|
||||
suspensions = (
|
||||
SuspendRow(
|
||||
ts_code=codes[1],
|
||||
trade_date=TARGET_DATE,
|
||||
suspend_timing=None,
|
||||
suspend_type="停牌",
|
||||
),
|
||||
)
|
||||
|
||||
facts = normalize_stock_facts(
|
||||
target_trade_date=TARGET_DATE,
|
||||
candidate_codes=codes,
|
||||
stock_basics=result("stock_basic", basics),
|
||||
suspensions=result("suspend_d", suspensions),
|
||||
daily=result("daily", daily_rows),
|
||||
moneyflow=result("moneyflow_dc", moneyflow_rows),
|
||||
)
|
||||
by_code = {fact.ts_code: fact for fact in facts}
|
||||
|
||||
assert by_code[codes[0]].status is StockFactStatus.AVAILABLE
|
||||
assert by_code[codes[0]].turnover_yuan == Decimal("1000")
|
||||
assert by_code[codes[0]].net_amount_yuan == Decimal("0")
|
||||
assert by_code[codes[1]].status is StockFactStatus.SUSPENDED
|
||||
assert by_code[codes[2]].status is StockFactStatus.MISSING_DAILY
|
||||
assert by_code[codes[3]].status is StockFactStatus.NULL_DAILY_AMOUNT
|
||||
assert by_code[codes[4]].status is StockFactStatus.MISSING_MONEYFLOW
|
||||
assert by_code[codes[5]].status is StockFactStatus.NULL_MONEYFLOW
|
||||
assert by_code[codes[6]].status is StockFactStatus.LOW_LIQUIDITY
|
||||
assert by_code[codes[7]].status is StockFactStatus.LIFECYCLE_INVALID
|
||||
@@ -0,0 +1,170 @@
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any, cast
|
||||
|
||||
from psycopg_pool import ConnectionPool
|
||||
|
||||
from zhixing_server.modules.sector_radar.domain.models import PublicationStatus
|
||||
from zhixing_server.modules.sector_radar.infrastructure.postgres import (
|
||||
PostgresSectorRadarRepository,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
|
||||
|
||||
class FakeResult:
|
||||
def __init__(
|
||||
self,
|
||||
row: tuple[object, ...] | None = None,
|
||||
rows: tuple[tuple[object, ...], ...] | None = None,
|
||||
) -> None:
|
||||
self.row = row
|
||||
self.rows = rows or (() if row is None else (row,))
|
||||
|
||||
def fetchone(self) -> tuple[object, ...] | None:
|
||||
return self.row
|
||||
|
||||
def fetchall(self) -> tuple[tuple[object, ...], ...]:
|
||||
return self.rows
|
||||
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self) -> None:
|
||||
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
||||
|
||||
def execute(
|
||||
self,
|
||||
query: str,
|
||||
parameters: tuple[object, ...] = (),
|
||||
) -> FakeResult:
|
||||
self.statements.append((query, parameters))
|
||||
if "FROM sector_radar_ranking" in query:
|
||||
return FakeResult(
|
||||
rows=(
|
||||
(
|
||||
TARGET_DATE,
|
||||
"concept",
|
||||
"BK0001.DC",
|
||||
"机器人",
|
||||
"amount",
|
||||
"zhixing_amount_net_bn_v1",
|
||||
"independent",
|
||||
"CNY_100M",
|
||||
Decimal("12.5"),
|
||||
"available",
|
||||
20,
|
||||
19,
|
||||
Decimal(1),
|
||||
Decimal("0.95"),
|
||||
1,
|
||||
Decimal(100),
|
||||
{"1": 3, "2": None},
|
||||
),
|
||||
)
|
||||
)
|
||||
if "FROM sector_radar_publication" in query:
|
||||
return FakeResult(
|
||||
(
|
||||
"publication-a",
|
||||
TARGET_DATE,
|
||||
"success",
|
||||
"tushare-pro-v1",
|
||||
"eastmoney-dc-v1",
|
||||
["zhixing_amount_net_bn_v1"],
|
||||
"a" * 64,
|
||||
Decimal("1"),
|
||||
datetime(2026, 8, 28, 17, 30, tzinfo=UTC),
|
||||
datetime(2026, 8, 28, 17, 35, tzinfo=UTC),
|
||||
None,
|
||||
)
|
||||
)
|
||||
if "pg_try_advisory_lock" in query:
|
||||
return FakeResult((True,))
|
||||
return FakeResult((True,))
|
||||
|
||||
|
||||
class FakePool:
|
||||
def __init__(self, connection: FakeConnection) -> None:
|
||||
self._connection = connection
|
||||
self._opened = False
|
||||
|
||||
def open(self, *, wait: bool) -> None:
|
||||
assert wait
|
||||
self._opened = True
|
||||
|
||||
def close(self) -> None:
|
||||
self._opened = False
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[FakeConnection]:
|
||||
yield self._connection
|
||||
|
||||
|
||||
def make_repository(connection: FakeConnection) -> PostgresSectorRadarRepository:
|
||||
pool = cast(ConnectionPool[Any], cast(object, FakePool(connection)))
|
||||
return PostgresSectorRadarRepository("postgresql://unused", pool=pool)
|
||||
|
||||
|
||||
def test_last_good_query_strictly_filters_success_and_date() -> None:
|
||||
connection = FakeConnection()
|
||||
|
||||
publication = make_repository(connection).get_last_good_publication(TARGET_DATE)
|
||||
|
||||
assert publication is not None
|
||||
assert publication.status is PublicationStatus.SUCCESS
|
||||
query, parameters = connection.statements[0]
|
||||
assert "status = 'success'" in query
|
||||
assert "partial" not in query
|
||||
assert "target_trade_date <= %s" in query
|
||||
assert parameters == (TARGET_DATE,)
|
||||
|
||||
|
||||
def test_advisory_lock_uses_target_date_and_releases_same_key() -> None:
|
||||
connection = FakeConnection()
|
||||
|
||||
with make_repository(connection).advisory_lock(TARGET_DATE) as acquired:
|
||||
assert acquired
|
||||
|
||||
assert len(connection.statements) == 2
|
||||
assert "pg_try_advisory_lock" in connection.statements[0][0]
|
||||
assert "2026-08-28" in str(connection.statements[0][1][0])
|
||||
assert "pg_advisory_unlock" in connection.statements[1][0]
|
||||
assert connection.statements[0][1] == connection.statements[1][1]
|
||||
|
||||
|
||||
def test_exact_success_and_latest_attempt_queries_use_distinct_semantics() -> None:
|
||||
connection = FakeConnection()
|
||||
repository = make_repository(connection)
|
||||
|
||||
exact = repository.get_successful_publication(TARGET_DATE)
|
||||
latest = repository.get_latest_publication()
|
||||
|
||||
assert exact is not None
|
||||
assert latest is not None
|
||||
exact_query, exact_parameters = connection.statements[0]
|
||||
latest_query, latest_parameters = connection.statements[1]
|
||||
assert "status = 'success' AND target_trade_date = %s" in exact_query
|
||||
assert exact_parameters == (TARGET_DATE,)
|
||||
assert "status = 'success'" not in latest_query
|
||||
assert "started_at DESC" in latest_query
|
||||
assert latest_parameters == ()
|
||||
|
||||
|
||||
def test_load_rankings_reconstructs_values_and_rank_changes() -> None:
|
||||
connection = FakeConnection()
|
||||
|
||||
rankings = make_repository(connection).load_rankings("publication-a")
|
||||
|
||||
assert len(rankings) == 1
|
||||
ranking = rankings[0]
|
||||
assert ranking.observation.metric_version == "zhixing_amount_net_bn_v1"
|
||||
assert ranking.observation.value == Decimal("12.5")
|
||||
assert ranking.rank_position == 1
|
||||
assert ranking.rank_change(1) == 3
|
||||
assert ranking.rank_change(2) is None
|
||||
query, parameters = connection.statements[0]
|
||||
assert "WHERE publication_id = %s" in query
|
||||
assert "rank_position NULLS LAST" in query
|
||||
assert parameters == ("publication-a",)
|
||||
@@ -0,0 +1,147 @@
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MetricKind,
|
||||
MetricObservation,
|
||||
MetricQuality,
|
||||
MetricUnit,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.ranking import (
|
||||
rank_metric_observations,
|
||||
select_percentile_side,
|
||||
select_rank_change_side,
|
||||
with_rank_changes,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
|
||||
|
||||
def make_observation(
|
||||
sector_code: str,
|
||||
sector_type: SectorType,
|
||||
value: str | None,
|
||||
*,
|
||||
trade_date: date = TARGET_DATE,
|
||||
) -> MetricObservation:
|
||||
metric_value = Decimal(value) if value is not None else None
|
||||
return MetricObservation(
|
||||
trade_date=trade_date,
|
||||
sector_type=sector_type,
|
||||
sector_code=sector_code,
|
||||
sector_name=sector_code,
|
||||
metric_kind=MetricKind.AMOUNT,
|
||||
metric_version="zhixing_amount_net_bn_v1",
|
||||
implementation_kind="independent",
|
||||
unit=MetricUnit.CNY_100M,
|
||||
value=metric_value,
|
||||
quality=(
|
||||
MetricQuality.AVAILABLE if metric_value is not None else MetricQuality.UNAVAILABLE
|
||||
),
|
||||
member_count=10,
|
||||
valid_sample_count=10 if metric_value is not None else 0,
|
||||
membership_coverage=Decimal("1"),
|
||||
moneyflow_coverage=Decimal("1"),
|
||||
)
|
||||
|
||||
|
||||
def test_ranking_separates_types_and_uses_code_as_stable_tie_breaker() -> None:
|
||||
observations = (
|
||||
make_observation("BK2002.DC", SectorType.INDUSTRY, "20"),
|
||||
make_observation("BK1002.DC", SectorType.CONCEPT, "30"),
|
||||
make_observation("BK2001.DC", SectorType.INDUSTRY, "20"),
|
||||
make_observation("BK1001.DC", SectorType.CONCEPT, "10"),
|
||||
)
|
||||
|
||||
ranked = rank_metric_observations(tuple(reversed(observations)))
|
||||
by_code = {row.observation.sector_code: row for row in ranked}
|
||||
|
||||
assert by_code["BK1002.DC"].rank_position == 1
|
||||
assert by_code["BK1002.DC"].rank_percentile == Decimal("100")
|
||||
assert by_code["BK1001.DC"].rank_position == 2
|
||||
assert by_code["BK1001.DC"].rank_percentile == Decimal("50")
|
||||
|
||||
assert by_code["BK2001.DC"].rank_position == 1
|
||||
assert by_code["BK2002.DC"].rank_position == 2
|
||||
|
||||
|
||||
def test_ranking_handles_empty_and_single_element_pools() -> None:
|
||||
assert rank_metric_observations(()) == ()
|
||||
|
||||
[single] = rank_metric_observations((make_observation("BK0001.DC", SectorType.CONCEPT, "0"),))
|
||||
|
||||
assert single.rank_position == 1
|
||||
assert single.rank_percentile == Decimal("100")
|
||||
|
||||
|
||||
def test_percentile_sides_use_confirmed_inclusive_thresholds() -> None:
|
||||
ranked = rank_metric_observations(
|
||||
make_observation(f"BK{position:04d}.DC", SectorType.CONCEPT, str(11 - position))
|
||||
for position in range(1, 11)
|
||||
)
|
||||
|
||||
top = select_percentile_side(ranked, RankSide.TOP)
|
||||
bottom = select_percentile_side(ranked, RankSide.BOTTOM)
|
||||
|
||||
assert [row.observation.sector_code for row in top] == ["BK0001.DC", "BK0002.DC"]
|
||||
assert [row.observation.sector_code for row in bottom] == ["BK0010.DC"]
|
||||
|
||||
|
||||
def test_rank_change_is_past_rank_minus_current_and_preserves_missing_history() -> None:
|
||||
current = rank_metric_observations(
|
||||
(
|
||||
make_observation("BK0001.DC", SectorType.CONCEPT, "30"),
|
||||
make_observation("BK0002.DC", SectorType.CONCEPT, "20"),
|
||||
)
|
||||
)
|
||||
previous = rank_metric_observations(
|
||||
(
|
||||
make_observation(
|
||||
"BK0001.DC",
|
||||
SectorType.CONCEPT,
|
||||
"10",
|
||||
trade_date=date(2026, 8, 27),
|
||||
),
|
||||
make_observation(
|
||||
"BK0002.DC",
|
||||
SectorType.CONCEPT,
|
||||
"40",
|
||||
trade_date=date(2026, 8, 27),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
changed = with_rank_changes(current, {1: previous, 5: ()})
|
||||
by_code = {row.observation.sector_code: row for row in changed}
|
||||
|
||||
assert by_code["BK0001.DC"].rank_change(1) == 1
|
||||
assert by_code["BK0002.DC"].rank_change(1) == -1
|
||||
assert by_code["BK0001.DC"].rank_change(5) is None
|
||||
|
||||
|
||||
def test_rank_change_sides_take_ceiling_ten_percent_per_pool() -> None:
|
||||
current = rank_metric_observations(
|
||||
make_observation(f"BK{position:04d}.DC", SectorType.CONCEPT, str(12 - position))
|
||||
for position in range(1, 12)
|
||||
)
|
||||
previous = rank_metric_observations(
|
||||
make_observation(
|
||||
f"BK{position:04d}.DC",
|
||||
SectorType.CONCEPT,
|
||||
str(position),
|
||||
trade_date=date(2026, 8, 27),
|
||||
)
|
||||
for position in range(1, 12)
|
||||
)
|
||||
changed = with_rank_changes(current, {1: previous})
|
||||
|
||||
top = select_rank_change_side(changed, days=1, side=RankSide.TOP)
|
||||
bottom = select_rank_change_side(changed, days=1, side=RankSide.BOTTOM)
|
||||
|
||||
assert [row.observation.sector_code for row in top] == ["BK0001.DC", "BK0002.DC"]
|
||||
assert [row.observation.sector_code for row in bottom] == [
|
||||
"BK0011.DC",
|
||||
"BK0010.DC",
|
||||
]
|
||||
@@ -0,0 +1,189 @@
|
||||
from dataclasses import replace
|
||||
from datetime import UTC, date, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
|
||||
from zhixing_server.modules.sector_radar.application.read import (
|
||||
RadarQuery,
|
||||
RadarView,
|
||||
ReadSectorRadar,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.metrics import AmountNetStrategy
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
MetricKind,
|
||||
MetricObservation,
|
||||
MetricQuality,
|
||||
MetricUnit,
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
RankChange,
|
||||
RankedMetric,
|
||||
RankSide,
|
||||
SectorType,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.persistence import RankingRecord
|
||||
from zhixing_server.modules.sector_radar.domain.ranking import rank_metric_observations
|
||||
from zhixing_server.modules.sector_radar.infrastructure.memory import (
|
||||
InMemorySectorRadarRepository,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
NOW = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
||||
|
||||
|
||||
def _running(publication_id: str, trade_date: date) -> RadarPublication:
|
||||
return RadarPublication(
|
||||
publication_id=publication_id,
|
||||
target_trade_date=trade_date,
|
||||
status=PublicationStatus.RUNNING,
|
||||
source_version="tushare-pro-v1",
|
||||
universe_version="eastmoney-dc-v1",
|
||||
metric_versions=(AmountNetStrategy.metric_version,),
|
||||
input_hash=None,
|
||||
coverage=Decimal(0),
|
||||
started_at=NOW,
|
||||
)
|
||||
|
||||
|
||||
def _finish(
|
||||
publication: RadarPublication,
|
||||
status: PublicationStatus,
|
||||
) -> RadarPublication:
|
||||
return replace(
|
||||
publication,
|
||||
status=status,
|
||||
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
|
||||
coverage=Decimal(1) if status is PublicationStatus.SUCCESS else Decimal("0.8"),
|
||||
finished_at=publication.started_at + timedelta(minutes=5),
|
||||
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
|
||||
)
|
||||
|
||||
|
||||
def _amount_rankings() -> tuple[RankedMetric, ...]:
|
||||
observations = tuple(
|
||||
MetricObservation(
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.CONCEPT,
|
||||
sector_code=f"BK{index:04d}.DC",
|
||||
sector_name=f"概念{index}",
|
||||
metric_kind=MetricKind.AMOUNT,
|
||||
metric_version=AmountNetStrategy.metric_version,
|
||||
implementation_kind="independent",
|
||||
unit=MetricUnit.CNY_100M,
|
||||
value=Decimal(11 - index),
|
||||
quality=MetricQuality.AVAILABLE,
|
||||
member_count=5,
|
||||
valid_sample_count=5,
|
||||
membership_coverage=Decimal(1),
|
||||
moneyflow_coverage=Decimal(1),
|
||||
)
|
||||
for index in range(1, 11)
|
||||
)
|
||||
rankings = rank_metric_observations(observations)
|
||||
return tuple(
|
||||
replace(
|
||||
row,
|
||||
rank_changes=tuple(
|
||||
RankChange(
|
||||
days=days,
|
||||
value=(
|
||||
None
|
||||
if row.observation.sector_code == "BK0005.DC" and days == 5
|
||||
else (row.rank_position or 0) - 5
|
||||
),
|
||||
)
|
||||
for days in range(1, 6)
|
||||
),
|
||||
)
|
||||
for row in rankings
|
||||
)
|
||||
|
||||
|
||||
def _published_repository() -> InMemorySectorRadarRepository:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
publication = _running("publication-success", TARGET_DATE)
|
||||
repository.create_publication(publication)
|
||||
repository.finish_publication(_finish(publication, PublicationStatus.SUCCESS))
|
||||
repository.save_rankings(
|
||||
RankingRecord(publication.publication_id, ranking) for ranking in _amount_rankings()
|
||||
)
|
||||
return repository
|
||||
|
||||
|
||||
def test_no_successful_publication_returns_stable_no_data() -> None:
|
||||
reader = ReadSectorRadar(InMemorySectorRadarRepository())
|
||||
|
||||
dates = reader.list_dates()
|
||||
rankings = reader.query(RadarQuery())
|
||||
|
||||
assert dates.status == "no_data"
|
||||
assert dates.available_dates == ()
|
||||
assert rankings.status == "no_data"
|
||||
assert rankings.publication is None
|
||||
assert rankings.total == 0
|
||||
assert rankings.definition.metric_version == AmountNetStrategy.metric_version
|
||||
|
||||
|
||||
def test_explicit_date_never_falls_back_to_an_earlier_last_good() -> None:
|
||||
reader = ReadSectorRadar(_published_repository())
|
||||
|
||||
missing = reader.query(RadarQuery(trade_date=TARGET_DATE + timedelta(days=1)))
|
||||
|
||||
assert missing.status == "no_data"
|
||||
assert missing.publication is None
|
||||
|
||||
|
||||
def test_percentile_side_is_selected_before_search_and_pagination() -> None:
|
||||
reader = ReadSectorRadar(_published_repository())
|
||||
|
||||
top = reader.query(RadarQuery(side=RankSide.TOP, page_size=1))
|
||||
second_page = reader.query(RadarQuery(side=RankSide.TOP, page=2, page_size=1))
|
||||
searched = reader.query(RadarQuery(side=RankSide.TOP, search="概念2"))
|
||||
bottom = reader.query(RadarQuery(side=RankSide.BOTTOM))
|
||||
|
||||
assert top.total == 2
|
||||
assert top.rows[0].observation.sector_code == "BK0001.DC"
|
||||
assert second_page.rows[0].observation.sector_code == "BK0002.DC"
|
||||
assert searched.total == 1
|
||||
assert searched.rows[0].observation.sector_name == "概念2"
|
||||
assert bottom.total == 1
|
||||
assert bottom.rows[0].observation.sector_code == "BK0010.DC"
|
||||
|
||||
|
||||
def test_rank_change_uses_selected_metric_days_and_pool_sides() -> None:
|
||||
reader = ReadSectorRadar(_published_repository())
|
||||
query = RadarQuery(
|
||||
view=RadarView.RANK_CHANGE,
|
||||
rank_change_metric=MetricKind.AMOUNT,
|
||||
rank_change_days=5,
|
||||
)
|
||||
|
||||
top = reader.query(replace(query, side=RankSide.TOP))
|
||||
bottom = reader.query(replace(query, side=RankSide.BOTTOM))
|
||||
all_rows = reader.query(query)
|
||||
|
||||
assert top.total == 1
|
||||
assert top.rows[0].rank_change(5) == 5
|
||||
assert bottom.total == 1
|
||||
assert bottom.rows[0].rank_change(5) == -4
|
||||
assert all_rows.total == 10
|
||||
assert all_rows.rows[-1].observation.sector_code == "BK0005.DC"
|
||||
assert all_rows.rows[-1].rank_change(5) is None
|
||||
|
||||
|
||||
def test_latest_partial_attempt_is_visible_but_does_not_replace_last_good() -> None:
|
||||
repository = _published_repository()
|
||||
partial = replace(
|
||||
_running("publication-partial", TARGET_DATE + timedelta(days=1)),
|
||||
started_at=NOW + timedelta(days=1),
|
||||
)
|
||||
repository.create_publication(partial)
|
||||
repository.finish_publication(_finish(partial, PublicationStatus.PARTIAL))
|
||||
|
||||
index = ReadSectorRadar(repository).list_dates()
|
||||
|
||||
assert index.status == "success"
|
||||
assert index.current_attempt is not None
|
||||
assert index.current_attempt.status is PublicationStatus.PARTIAL
|
||||
assert index.last_good is not None
|
||||
assert index.last_good.publication_id == "publication-success"
|
||||
assert index.available_dates == (TARGET_DATE,)
|
||||
@@ -0,0 +1,121 @@
|
||||
from dataclasses import replace
|
||||
from datetime import UTC, date, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.sector_radar.domain.models import (
|
||||
PublicationStatus,
|
||||
RadarPublication,
|
||||
SectorType,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.domain.persistence import MembershipRecord
|
||||
from zhixing_server.modules.sector_radar.domain.source import build_source_snapshot
|
||||
from zhixing_server.modules.sector_radar.infrastructure.memory import (
|
||||
InMemorySectorRadarRepository,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
STARTED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
||||
|
||||
|
||||
def make_running(publication_id: str, target_trade_date: date = TARGET_DATE) -> RadarPublication:
|
||||
return RadarPublication(
|
||||
publication_id=publication_id,
|
||||
target_trade_date=target_trade_date,
|
||||
status=PublicationStatus.RUNNING,
|
||||
source_version="tushare-pro-v1",
|
||||
universe_version="eastmoney-dc-v1",
|
||||
metric_versions=("zhixing_amount_net_bn_v1",),
|
||||
input_hash=None,
|
||||
coverage=Decimal(0),
|
||||
started_at=STARTED_AT,
|
||||
)
|
||||
|
||||
|
||||
def finish(
|
||||
publication: RadarPublication,
|
||||
status: PublicationStatus,
|
||||
*,
|
||||
offset_minutes: int = 5,
|
||||
) -> RadarPublication:
|
||||
return replace(
|
||||
publication,
|
||||
status=status,
|
||||
input_hash="a" * 64 if status is PublicationStatus.SUCCESS else None,
|
||||
coverage=Decimal("1") if status is PublicationStatus.SUCCESS else Decimal("0.8"),
|
||||
finished_at=publication.started_at + timedelta(minutes=offset_minutes),
|
||||
error_summary=None if status is PublicationStatus.SUCCESS else "safe_error",
|
||||
)
|
||||
|
||||
|
||||
def test_source_and_membership_revisions_are_idempotent_but_not_overwritable() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
snapshot = build_source_snapshot(
|
||||
api_name="dc_member",
|
||||
params={"trade_date": "20260828"},
|
||||
rows=(
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "平安银行",
|
||||
},
|
||||
),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=STARTED_AT,
|
||||
)
|
||||
member = MembershipRecord(
|
||||
source_snapshot_id=snapshot.snapshot_id,
|
||||
trade_date=TARGET_DATE,
|
||||
sector_type=SectorType.CONCEPT,
|
||||
sector_code="BK0001.DC",
|
||||
sector_name="示例概念",
|
||||
stock_code="000001.SZ",
|
||||
stock_name="平安银行",
|
||||
)
|
||||
|
||||
assert repository.save_source_snapshots((snapshot,)).inserted == 1
|
||||
assert repository.save_source_snapshots((snapshot,)).unchanged == 1
|
||||
assert repository.save_memberships((member,)).inserted == 1
|
||||
assert repository.save_memberships((member,)).unchanged == 1
|
||||
|
||||
with pytest.raises(ValueError, match="cannot change content"):
|
||||
repository.save_memberships((replace(member, stock_name="已改变"),))
|
||||
|
||||
|
||||
def test_partial_and_failed_revisions_never_replace_last_good() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
successful = make_running("success-a")
|
||||
partial = make_running("partial-b")
|
||||
failed = make_running("failed-c", TARGET_DATE + timedelta(days=1))
|
||||
repository.create_publication(successful)
|
||||
repository.finish_publication(finish(successful, PublicationStatus.SUCCESS))
|
||||
repository.create_publication(partial)
|
||||
repository.finish_publication(finish(partial, PublicationStatus.PARTIAL, offset_minutes=6))
|
||||
repository.create_publication(failed)
|
||||
repository.finish_publication(finish(failed, PublicationStatus.FAILED, offset_minutes=7))
|
||||
|
||||
last_good = repository.get_last_good_publication()
|
||||
|
||||
assert last_good is not None
|
||||
assert last_good.publication_id == "success-a"
|
||||
assert repository.list_successful_dates() == (TARGET_DATE,)
|
||||
|
||||
|
||||
def test_publication_identity_allows_sequential_same_date_revisions() -> None:
|
||||
repository = InMemorySectorRadarRepository()
|
||||
first = make_running("revision-a")
|
||||
second = make_running("revision-b")
|
||||
|
||||
assert repository.create_publication(first).inserted == 1
|
||||
with pytest.raises(ValueError, match="already has a running"):
|
||||
repository.create_publication(second)
|
||||
|
||||
with pytest.raises(ValueError, match="terminal"):
|
||||
repository.finish_publication(first)
|
||||
|
||||
repository.finish_publication(finish(first, PublicationStatus.FAILED))
|
||||
assert repository.create_publication(second).inserted == 1
|
||||
with pytest.raises(ValueError, match="running status"):
|
||||
repository.finish_publication(finish(first, PublicationStatus.SUCCESS))
|
||||
@@ -0,0 +1,669 @@
|
||||
import logging
|
||||
import threading
|
||||
from collections.abc import Mapping
|
||||
from datetime import UTC, date, datetime
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.sector_radar.domain.models import SectorType
|
||||
from zhixing_server.modules.sector_radar.domain.source import (
|
||||
CapabilityStatus,
|
||||
SourceContractError,
|
||||
build_source_snapshot,
|
||||
)
|
||||
from zhixing_server.modules.sector_radar.infrastructure import tushare as source_module
|
||||
from zhixing_server.modules.sector_radar.infrastructure.tushare import (
|
||||
TushareSectorRadarAdapter,
|
||||
)
|
||||
|
||||
TARGET_DATE = date(2026, 8, 28)
|
||||
OBSERVED_AT = datetime(2026, 8, 28, 17, 30, tzinfo=UTC)
|
||||
|
||||
|
||||
class QueryClient:
|
||||
def __init__(self, responses: Mapping[tuple[str, str], object]) -> None:
|
||||
self.responses = dict(responses)
|
||||
self.calls: list[tuple[str, dict[str, object]]] = []
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def query(self, api_name: str, **kwargs: object) -> object:
|
||||
partition = str(kwargs.get("ts_code") or kwargs.get("list_status") or "")
|
||||
with self._lock:
|
||||
self.calls.append((api_name, kwargs))
|
||||
response = self.responses.get((api_name, partition), ())
|
||||
if isinstance(response, BaseException):
|
||||
raise response
|
||||
return response
|
||||
|
||||
|
||||
def make_adapter(client: object) -> TushareSectorRadarAdapter:
|
||||
return TushareSectorRadarAdapter(
|
||||
client,
|
||||
max_retries=0,
|
||||
request_interval_seconds=0,
|
||||
sleep_fn=lambda _: None,
|
||||
now_fn=lambda: OBSERVED_AT,
|
||||
)
|
||||
|
||||
|
||||
def moneyflow_record(
|
||||
ts_code: str,
|
||||
*,
|
||||
trade_date: str = "20260828",
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"trade_date": trade_date,
|
||||
"ts_code": ts_code,
|
||||
"name": ts_code,
|
||||
"net_amount": "1",
|
||||
"net_amount_rate": "0.1",
|
||||
"pct_change": "1",
|
||||
"close": "10",
|
||||
}
|
||||
|
||||
|
||||
def test_daily_and_moneyflow_keep_source_units_and_distinguish_missing_from_zero() -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"daily",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"ts_code": "000001.SZ",
|
||||
"trade_date": "20260828",
|
||||
"close": "10",
|
||||
"pre_close": "9.5",
|
||||
"pct_chg": "1.5",
|
||||
"vol": "100",
|
||||
"amount": "12.5",
|
||||
},
|
||||
{
|
||||
"ts_code": "000002.SZ",
|
||||
"trade_date": "20260828",
|
||||
"close": "20",
|
||||
"pre_close": "20",
|
||||
"pct_chg": "0",
|
||||
"vol": "0",
|
||||
"amount": float("nan"),
|
||||
},
|
||||
),
|
||||
(
|
||||
"moneyflow_dc",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "000001.SZ",
|
||||
"name": "平安银行",
|
||||
"net_amount": "2.5",
|
||||
"net_amount_rate": "0.2",
|
||||
"pct_change": "1.5",
|
||||
"close": "10",
|
||||
},
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "000002.SZ",
|
||||
"name": "示例股票",
|
||||
"net_amount": "0",
|
||||
"net_amount_rate": "0",
|
||||
"pct_change": "0",
|
||||
"close": "20",
|
||||
},
|
||||
),
|
||||
}
|
||||
)
|
||||
adapter = make_adapter(client)
|
||||
|
||||
daily = adapter.fetch_daily(TARGET_DATE)
|
||||
moneyflow = adapter.fetch_moneyflow_dc(TARGET_DATE, ("000001.SZ", "000002.SZ"))
|
||||
|
||||
assert daily.rows[0].amount_thousand_yuan == Decimal("12.5")
|
||||
assert daily.rows[0].turnover_yuan == Decimal("12500.0")
|
||||
assert daily.rows[1].amount_thousand_yuan is None
|
||||
assert moneyflow.rows[0].net_amount_ten_thousand_yuan == Decimal("2.5")
|
||||
assert moneyflow.rows[0].net_amount_yuan == Decimal("25000.0")
|
||||
assert moneyflow.rows[1].net_amount_yuan == Decimal("0")
|
||||
assert client.calls[0][1]["fields"] == ",".join(source_module.FIELDS["daily"])
|
||||
|
||||
|
||||
def test_moneyflow_accepts_a_full_initial_snapshot_at_the_provider_limit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", 2)
|
||||
client = QueryClient(
|
||||
{
|
||||
("moneyflow_dc", ""): (
|
||||
moneyflow_record("000001.SZ"),
|
||||
moneyflow_record("000002.SZ"),
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_moneyflow_dc(
|
||||
TARGET_DATE,
|
||||
("000001.SZ", "000002.SZ"),
|
||||
)
|
||||
|
||||
assert result.snapshots[0].limit_reached is True
|
||||
assert [row.ts_code for row in result.rows] == ["000001.SZ", "000002.SZ"]
|
||||
assert len(client.calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("initial_rows", "message"),
|
||||
(
|
||||
((moneyflow_record("000001.SZ", trade_date="20260827"),), "trade_date"),
|
||||
(
|
||||
(moneyflow_record("000001.SZ"), moneyflow_record("000001.SZ")),
|
||||
"duplicate business keys",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_moneyflow_initial_contract_errors_fail_closed(
|
||||
initial_rows: tuple[dict[str, object], ...],
|
||||
message: str,
|
||||
) -> None:
|
||||
client = QueryClient({("moneyflow_dc", ""): initial_rows})
|
||||
|
||||
with pytest.raises(SourceContractError, match=message):
|
||||
make_adapter(client).fetch_moneyflow_dc(TARGET_DATE, ())
|
||||
|
||||
|
||||
def test_moneyflow_refills_only_missing_codes_in_stable_snapshot_order(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", 3)
|
||||
third_finished = threading.Event()
|
||||
completion_order: list[str] = []
|
||||
completion_lock = threading.Lock()
|
||||
|
||||
class ReverseCompletionClient(QueryClient):
|
||||
def query(self, api_name: str, **kwargs: object) -> object:
|
||||
response = super().query(api_name, **kwargs)
|
||||
ts_code = str(kwargs.get("ts_code") or "")
|
||||
if ts_code == "000004.SZ":
|
||||
if not third_finished.wait(timeout=2):
|
||||
raise AssertionError("second moneyflow worker did not start")
|
||||
elif ts_code == "000005.SZ":
|
||||
third_finished.set()
|
||||
if ts_code:
|
||||
with completion_lock:
|
||||
completion_order.append(ts_code)
|
||||
return response
|
||||
|
||||
client = ReverseCompletionClient(
|
||||
{
|
||||
("moneyflow_dc", ""): tuple(
|
||||
moneyflow_record(f"00000{index}.SZ") for index in range(1, 4)
|
||||
),
|
||||
("moneyflow_dc", "000004.SZ"): (moneyflow_record("000004.SZ"),),
|
||||
("moneyflow_dc", "000005.SZ"): (moneyflow_record("000005.SZ"),),
|
||||
}
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_moneyflow_dc(
|
||||
TARGET_DATE,
|
||||
tuple(f"00000{index}.SZ" for index in range(1, 6)),
|
||||
)
|
||||
|
||||
assert completion_order == ["000005.SZ", "000004.SZ"]
|
||||
assert [snapshot.partition_key for snapshot in result.snapshots] == [
|
||||
"all",
|
||||
"000004.SZ",
|
||||
"000005.SZ",
|
||||
]
|
||||
assert [row.ts_code for row in result.rows] == [
|
||||
"000001.SZ",
|
||||
"000002.SZ",
|
||||
"000003.SZ",
|
||||
"000004.SZ",
|
||||
"000005.SZ",
|
||||
]
|
||||
assert len(client.calls) == 3
|
||||
|
||||
|
||||
def test_moneyflow_empty_and_exhausted_refills_remain_real_gaps(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
("moneyflow_dc", ""): (moneyflow_record("000001.SZ"),),
|
||||
("moneyflow_dc", "000002.SZ"): (),
|
||||
("moneyflow_dc", "000003.SZ"): RuntimeError("private provider payload"),
|
||||
}
|
||||
)
|
||||
caplog.set_level(
|
||||
logging.WARNING,
|
||||
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_moneyflow_dc(
|
||||
TARGET_DATE,
|
||||
("000001.SZ", "000002.SZ", "000003.SZ"),
|
||||
)
|
||||
|
||||
assert [row.ts_code for row in result.rows] == ["000001.SZ"]
|
||||
assert [snapshot.partition_key for snapshot in result.snapshots] == ["all"]
|
||||
messages = "\n".join(record.getMessage() for record in caplog.records)
|
||||
assert "partition_empty partition_key=000002.SZ" in messages
|
||||
assert "partition_failed partition_key=000003.SZ" in messages
|
||||
assert "private provider payload" not in messages
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("partition_rows", "row_limit", "message"),
|
||||
(
|
||||
((moneyflow_record("000002.SZ", trade_date="20260827"),), 6_000, "trade_date"),
|
||||
((moneyflow_record("000099.SZ"),), 6_000, "different ts_code"),
|
||||
(
|
||||
(moneyflow_record("000002.SZ"), moneyflow_record("000002.SZ")),
|
||||
6_000,
|
||||
"duplicate business keys",
|
||||
),
|
||||
(
|
||||
(moneyflow_record("000002.SZ"), moneyflow_record("000002.SZ")),
|
||||
2,
|
||||
"provider row limit",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_moneyflow_partition_contract_errors_fail_closed(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
partition_rows: tuple[dict[str, object], ...],
|
||||
row_limit: int,
|
||||
message: str,
|
||||
) -> None:
|
||||
monkeypatch.setitem(source_module.ROW_LIMITS, "moneyflow_dc", row_limit)
|
||||
client = QueryClient(
|
||||
{
|
||||
("moneyflow_dc", ""): (moneyflow_record("000001.SZ"),),
|
||||
("moneyflow_dc", "000002.SZ"): partition_rows,
|
||||
}
|
||||
)
|
||||
|
||||
with pytest.raises(SourceContractError, match=message):
|
||||
make_adapter(client).fetch_moneyflow_dc(
|
||||
TARGET_DATE,
|
||||
("000001.SZ", "000002.SZ"),
|
||||
)
|
||||
|
||||
|
||||
def test_non_finite_source_values_are_rejected() -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"daily",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"ts_code": "000001.SZ",
|
||||
"trade_date": "20260828",
|
||||
"close": "Infinity",
|
||||
"pre_close": "9.5",
|
||||
"pct_chg": "1.5",
|
||||
"vol": "100",
|
||||
"amount": "12.5",
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
with pytest.raises(SourceContractError, match="finite"):
|
||||
make_adapter(client).fetch_daily(TARGET_DATE)
|
||||
|
||||
|
||||
def test_contract_failure_log_identifies_member_partition_without_payload(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(source_module.ROW_LIMITS, "dc_member", 2)
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"dc_member",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "private-payload-marker",
|
||||
},
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000002.SZ",
|
||||
"name": "private-payload-marker",
|
||||
},
|
||||
),
|
||||
(
|
||||
"dc_member",
|
||||
"BK0001.DC",
|
||||
): (
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK9999.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "private-payload-marker",
|
||||
},
|
||||
),
|
||||
}
|
||||
)
|
||||
caplog.set_level(
|
||||
logging.ERROR,
|
||||
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
|
||||
)
|
||||
|
||||
with pytest.raises(SourceContractError, match="different sector"):
|
||||
make_adapter(client).fetch_sector_members(TARGET_DATE, ("BK0001.DC",))
|
||||
|
||||
messages = "\n".join(record.getMessage() for record in caplog.records)
|
||||
assert "sector_radar_source_contract_failed" in messages
|
||||
assert "api_name=dc_member" in messages
|
||||
assert "partition_key=BK0001.DC" in messages
|
||||
assert "validation=dc_member partition returned a different sector" in messages
|
||||
assert "private-payload-marker" not in messages
|
||||
assert len(caplog.records) == 1
|
||||
|
||||
|
||||
def test_merged_member_contract_failure_has_one_interface_level_log(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
duplicate = {
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "private-payload-marker",
|
||||
}
|
||||
client = QueryClient({("dc_member", ""): (duplicate, duplicate)})
|
||||
caplog.set_level(
|
||||
logging.ERROR,
|
||||
logger="zhixing_server.modules.sector_radar.infrastructure.tushare",
|
||||
)
|
||||
|
||||
with pytest.raises(SourceContractError, match="duplicate business keys"):
|
||||
make_adapter(client).fetch_sector_members(TARGET_DATE, ("BK0001.DC",))
|
||||
|
||||
messages = "\n".join(record.getMessage() for record in caplog.records)
|
||||
assert "api_name=dc_member" in messages
|
||||
assert "partition_key=merged" in messages
|
||||
assert "validation=dc_member returned duplicate business keys" in messages
|
||||
assert "private-payload-marker" not in messages
|
||||
assert len(caplog.records) == 1
|
||||
|
||||
|
||||
def test_dc_member_reloads_by_sector_when_the_all_market_call_hits_limit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setitem(source_module.ROW_LIMITS, "dc_member", 2)
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"dc_member",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "A",
|
||||
},
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000002.SZ",
|
||||
"name": "B",
|
||||
},
|
||||
),
|
||||
(
|
||||
"dc_member",
|
||||
"BK0001.DC",
|
||||
): (
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "A",
|
||||
},
|
||||
),
|
||||
(
|
||||
"dc_member",
|
||||
"BK0002.DC",
|
||||
): (
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0002.DC",
|
||||
"con_code": "600000.SH",
|
||||
"name": "C",
|
||||
},
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_sector_members(
|
||||
TARGET_DATE,
|
||||
("BK0001.DC", "BK0002.DC"),
|
||||
)
|
||||
|
||||
assert [row.stock_code for row in result.rows] == ["000001.SZ", "600000.SH"]
|
||||
assert [snapshot.partition_key for snapshot in result.snapshots] == [
|
||||
"all",
|
||||
"BK0001.DC",
|
||||
"BK0002.DC",
|
||||
]
|
||||
|
||||
|
||||
def test_dc_member_preserves_an_explicit_empty_partition() -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"dc_member",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"trade_date": "20260828",
|
||||
"ts_code": "BK0001.DC",
|
||||
"con_code": "000001.SZ",
|
||||
"name": "A",
|
||||
},
|
||||
),
|
||||
("dc_member", "BK0002.DC"): (),
|
||||
}
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_sector_members(
|
||||
TARGET_DATE,
|
||||
("BK0001.DC", "BK0002.DC"),
|
||||
)
|
||||
|
||||
assert [row.sector_code for row in result.rows] == ["BK0001.DC"]
|
||||
assert [snapshot.partition_key for snapshot in result.snapshots] == [
|
||||
"all",
|
||||
"BK0002.DC",
|
||||
]
|
||||
assert result.snapshots[1].row_count == 0
|
||||
|
||||
|
||||
def test_stock_basic_requests_only_current_listings() -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"stock_basic",
|
||||
"L",
|
||||
): (
|
||||
{
|
||||
"ts_code": "000001.SZ",
|
||||
"symbol": "000001",
|
||||
"name": "L",
|
||||
"market": "主板",
|
||||
"exchange": "SZSE",
|
||||
"list_status": "L",
|
||||
"list_date": "20200101",
|
||||
"delist_date": None,
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_stock_basics()
|
||||
|
||||
assert {row.list_status for row in result.rows} == {"L"}
|
||||
assert [snapshot.partition_key for snapshot in result.snapshots] == ["L"]
|
||||
assert [call[1]["list_status"] for call in client.calls] == ["L"]
|
||||
|
||||
|
||||
def test_stock_basic_rejects_a_non_listed_row_from_the_l_partition() -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
("stock_basic", "L"): (
|
||||
{
|
||||
"ts_code": "000001.SZ",
|
||||
"symbol": "000001",
|
||||
"name": "unexpected",
|
||||
"market": "主板",
|
||||
"exchange": "SZSE",
|
||||
"list_status": "D",
|
||||
"list_date": "20200101",
|
||||
"delist_date": "20260828",
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
with pytest.raises(SourceContractError, match="unexpected list_status"):
|
||||
make_adapter(client).fetch_stock_basics()
|
||||
|
||||
|
||||
def test_suspend_timing_may_be_missing_while_suspend_type_remains_required() -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"suspend_d",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"ts_code": "000001.SZ",
|
||||
"trade_date": "20260828",
|
||||
"suspend_timing": None,
|
||||
"suspend_type": "S",
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_suspensions(TARGET_DATE)
|
||||
|
||||
assert result.rows[0].suspend_timing is None
|
||||
assert result.rows[0].suspend_type == "S"
|
||||
|
||||
missing_type = QueryClient(
|
||||
{
|
||||
(
|
||||
"suspend_d",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"ts_code": "000001.SZ",
|
||||
"trade_date": "20260828",
|
||||
"suspend_timing": None,
|
||||
"suspend_type": None,
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
with pytest.raises(SourceContractError, match="suspend_type must be a non-empty string"):
|
||||
make_adapter(missing_type).fetch_suspensions(TARGET_DATE)
|
||||
|
||||
|
||||
def test_source_snapshot_hash_is_order_stable_and_excludes_token_params() -> None:
|
||||
first = build_source_snapshot(
|
||||
api_name="daily",
|
||||
params={"trade_date": "20260828", "token": "secret"},
|
||||
rows=({"ts_code": "2"}, {"ts_code": "1"}),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
second = build_source_snapshot(
|
||||
api_name="daily",
|
||||
params={"trade_date": "20260828"},
|
||||
rows=({"ts_code": "1"}, {"ts_code": "2"}),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=OBSERVED_AT,
|
||||
)
|
||||
|
||||
assert first.snapshot_id == second.snapshot_id
|
||||
assert "secret" not in repr(first)
|
||||
|
||||
|
||||
def test_source_snapshot_identity_includes_schema_and_limit_metadata() -> None:
|
||||
first = build_source_snapshot(
|
||||
api_name="daily",
|
||||
params={"trade_date": "20260828"},
|
||||
rows=({"ts_code": "000001.SZ"},),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=OBSERVED_AT,
|
||||
returned_fields=("ts_code",),
|
||||
row_limit=1,
|
||||
)
|
||||
changed_schema = build_source_snapshot(
|
||||
api_name="daily",
|
||||
params={"trade_date": "20260828"},
|
||||
rows=({"ts_code": "000001.SZ"},),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=OBSERVED_AT,
|
||||
returned_fields=("name", "ts_code"),
|
||||
row_limit=1,
|
||||
)
|
||||
changed_limit = build_source_snapshot(
|
||||
api_name="daily",
|
||||
params={"trade_date": "20260828"},
|
||||
rows=({"ts_code": "000001.SZ"},),
|
||||
target_trade_date=TARGET_DATE,
|
||||
observed_at=OBSERVED_AT,
|
||||
returned_fields=("ts_code",),
|
||||
row_limit=2,
|
||||
)
|
||||
|
||||
assert first.content_sha256 != changed_schema.content_sha256
|
||||
assert first.snapshot_id != changed_schema.snapshot_id
|
||||
assert first.content_sha256 != changed_limit.content_sha256
|
||||
assert first.snapshot_id != changed_limit.snapshot_id
|
||||
|
||||
|
||||
def test_capability_probe_classifies_errors_without_exposing_provider_text() -> None:
|
||||
client = QueryClient({("daily", ""): RuntimeError("权限不足 private-detail")})
|
||||
|
||||
probe = make_adapter(client).probe(TARGET_DATE)
|
||||
|
||||
by_name = {result.api_name: result for result in probe.interfaces}
|
||||
assert by_name["daily"].status is CapabilityStatus.FORBIDDEN
|
||||
assert "private-detail" not in repr(probe)
|
||||
assert len(probe.interfaces) == 7
|
||||
|
||||
|
||||
def test_sector_index_uses_independent_concept_and_industry_params() -> None:
|
||||
client = QueryClient(
|
||||
{
|
||||
(
|
||||
"dc_index",
|
||||
"",
|
||||
): (
|
||||
{
|
||||
"ts_code": "BK0001.DC",
|
||||
"trade_date": "20260828",
|
||||
"name": "示例",
|
||||
"idx_type": "概念板块",
|
||||
"level": "一级",
|
||||
"pct_change": "1",
|
||||
"leading_code": "000001.SZ",
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
result = make_adapter(client).fetch_sector_indices(TARGET_DATE, SectorType.CONCEPT)
|
||||
|
||||
assert result.rows[0].sector_type is SectorType.CONCEPT
|
||||
assert client.calls[0][1]["idx_type"] == "概念板块"
|
||||
Reference in New Issue
Block a user