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"] == "概念板块"