Files
zhixing-system/zhixing-server/tests/unit/sector_radar/test_tushare_source.py
T

670 lines
21 KiB
Python

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