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

670 lines
21 KiB
Python
Raw Normal View History

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()
2026-08-31 09:47:09 +08:00
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"] == "概念板块"