fix(sector-radar): 支持当前上市股票资金流补拉

This commit is contained in:
yuxuanhui
2026-08-31 14:26:28 +08:00
parent 1cd7b5cb38
commit 2ffd0163f2
20 changed files with 1077 additions and 93 deletions
@@ -33,7 +33,11 @@ from ..domain.models import (
StockDailyFact,
StockFactStatus,
)
from ..domain.normalize import normalize_memberships, normalize_stock_facts
from ..domain.normalize import (
is_current_listed_stock,
normalize_memberships,
normalize_stock_facts,
)
from ..domain.persistence import (
DailyAggregateRecord,
MembershipRecord,
@@ -461,18 +465,20 @@ class BuildSectorRadar:
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 completed group or fetch and checkpoint it immediately."""
"""Replay a compatible completed group or fetch and checkpoint it immediately."""
try:
snapshots = reusable.get(source_group)
if snapshots is None:
result = fetch()
else:
result = SourceResult(
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(
@@ -540,6 +546,22 @@ class BuildSectorRadar:
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,
@@ -558,23 +580,16 @@ class BuildSectorRadar:
publication_id,
PublicationSourceGroup.MONEYFLOW_DC,
reusable,
lambda: self.source.fetch_moneyflow_dc(target),
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}
),
)
memberships = normalize_memberships(indices, members)
candidate_codes = tuple(
sorted(
{
item.stock_code
for item in memberships
if item.status is MembershipStatus.AVAILABLE and item.stock_code is not None
}
)
)
stock_facts = normalize_stock_facts(
target_trade_date=target,
candidate_codes=candidate_codes,
candidate_codes=member_codes,
stock_basics=stock_basics,
suspensions=suspensions,
daily=daily,
@@ -117,7 +117,7 @@ def normalize_stock_facts(
Args:
target_trade_date: Date whose point-in-time lifecycle is evaluated.
candidate_codes: Union of stocks in that date's sector memberships.
stock_basics: All explicit Tushare listing-status partitions.
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.
@@ -165,7 +165,7 @@ def normalize_stock_facts(
turnover_yuan = None
net_amount_yuan = None
if basic is None or not _is_lifecycle_candidate(basic, target_trade_date):
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
@@ -197,7 +197,23 @@ def normalize_stock_facts(
return tuple(records)
def _is_lifecycle_candidate(stock: StockBasicRow, target: date) -> bool:
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")):
@@ -205,9 +221,7 @@ def _is_lifecycle_candidate(stock: StockBasicRow, target: date) -> bool:
market = stock.market or ""
if "北交" in market or "B股" in market.upper():
return False
if stock.list_date is None or stock.list_date > target:
return False
return stock.delist_date is None or target <= stock.delist_date
return stock.list_date is not None and stock.list_date <= target
def _is_suspend_event(value: str) -> bool:
@@ -42,7 +42,11 @@ class SectorRadarSource(Protocol):
def fetch_daily(self, trade_date: date) -> SourceResult[DailyRow]: ...
def fetch_moneyflow_dc(self, trade_date: date) -> SourceResult[MoneyflowDcRow]: ...
def fetch_moneyflow_dc(
self,
trade_date: date,
candidate_codes: Sequence[str],
) -> SourceResult[MoneyflowDcRow]: ...
def probe(self, trade_date: date) -> CapabilityProbeResult: ...
@@ -5,12 +5,14 @@ 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
@@ -84,6 +86,7 @@ _SECTOR_TYPE_PARAM = {
SectorType.CONCEPT: "概念板块",
SectorType.INDUSTRY: "行业板块",
}
_MONEYFLOW_WORKERS = 2
class TushareSectorRadarAdapter:
@@ -104,12 +107,11 @@ class TushareSectorRadarAdapter:
"""Create an adapter around one already-authenticated SDK client."""
self._client = client
self._sleep_fn = sleep_fn
self._request_interval_seconds = max(0.0, request_interval_seconds)
self._now_fn = now_fn
self._coordinator = request_coordinator or RequestCoordinator(
max_retries=max_retries,
backoff_seconds=backoff_seconds,
request_interval_seconds=request_interval_seconds,
cooldown_seconds=cooldown_seconds,
wait_fn=sleep_fn,
sleep_fn=sleep_fn,
@@ -265,24 +267,19 @@ class TushareSectorRadarAdapter:
)
def fetch_stock_basics(self) -> SourceResult[StockBasicRow]:
"""Fetch every documented listing status instead of relying on the L default."""
"""Fetch the build-time current ``L`` listings in one explicit partition."""
snapshots: list[SourceSnapshot] = []
rows: list[StockBasicRow] = []
for status in ("L", "D", "P", "G", "UN"):
snapshot = self._fetch_snapshot(
"stock_basic",
{"exchange": "", "list_status": status},
target_trade_date=None,
partition_key=status,
)
snapshots.append(snapshot)
parsed = tuple(StockBasicRow.from_mapping(row) for row in snapshot.rows)
if any(row.list_status != status for row in parsed):
raise SourceContractError("stock_basic returned an unexpected list_status")
rows.extend(parsed)
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(tuple(snapshots), tuple(sorted(rows, key=lambda row: row.ts_code)))
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."""
@@ -315,19 +312,121 @@ class TushareSectorRadarAdapter:
self._require_unique(rows, key=lambda row: row.ts_code, api_name="daily")
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
def fetch_moneyflow_dc(self, trade_date: date) -> SourceResult[MoneyflowDcRow]:
"""Fetch a full-market DC moneyflow snapshot in its documented source unit."""
def fetch_moneyflow_dc(
self,
trade_date: date,
candidate_codes: Sequence[str],
) -> SourceResult[MoneyflowDcRow]:
"""Fetch full-market moneyflow and refill uncovered current candidates."""
snapshot = self._fetch_snapshot(
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",
)
self._reject_limit(snapshot)
rows = tuple(MoneyflowDcRow.from_mapping(row) for row in snapshot.rows)
self._require_target_date(rows, trade_date, "moneyflow_dc")
self._require_unique(rows, key=lambda row: row.ts_code, api_name="moneyflow_dc")
return SourceResult((snapshot,), tuple(sorted(rows, key=lambda row: row.ts_code)))
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."""
@@ -360,7 +459,7 @@ class TushareSectorRadarAdapter:
("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)),
("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))
@@ -385,7 +484,6 @@ class TushareSectorRadarAdapter:
return method(fields=fields, **params)
result = self._coordinator.call(api_name, request)
self._sleep_fn(self._request_interval_seconds)
try:
columns = getattr(result, "columns", None)
returned_fields = (
@@ -33,11 +33,12 @@ class TushareSourceError(RuntimeError):
class RequestCoordinator:
"""Coordinate retries and shared rate-limit cooling for one provider client.
"""Coordinate retries, rate-limit cooling, and optional request start spacing.
Normal requests are not serialized. Only a classified provider limit creates
a shared cooldown. Injectable time functions keep long cooldowns deterministic
in tests without coupling the coordinator to any business bounded context.
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__(
@@ -45,6 +46,7 @@ class RequestCoordinator:
*,
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,
@@ -56,6 +58,7 @@ class RequestCoordinator:
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
@@ -63,6 +66,7 @@ class RequestCoordinator:
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
@@ -77,7 +81,7 @@ class RequestCoordinator:
last_error: BaseException | None = None
for attempt in range(self.max_retries + 1):
self._wait_for_cooldown(method_name)
self._wait_for_request_start(method_name)
try:
result = request()
except Exception as exc:
@@ -124,17 +128,29 @@ class RequestCoordinator:
return self.call(method_name, operation)
def _wait_for_cooldown(self, method_name: str) -> None:
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:
delay = self._cooldown_until - self.clock()
if delay <= 0:
return
logger.info(
"provider_rate_limit_wait method=%s wait_seconds=%.1f",
method_name,
delay,
)
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:
@@ -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:
@@ -1,5 +1,6 @@
import logging
from collections.abc import Sequence
from dataclasses import replace
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
@@ -51,6 +52,7 @@ class FakeRadarSource:
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, ...]
@@ -239,8 +241,13 @@ class FakeRadarSource:
)
return self._result("daily", trade_date, rows)
def fetch_moneyflow_dc(self, trade_date: date) -> SourceResult[MoneyflowDcRow]:
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(
@@ -288,6 +295,28 @@ def test_successful_build_is_idempotent_and_failed_retry_preserves_last_good() -
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:
@@ -386,6 +415,85 @@ def test_unknown_membership_is_persisted_as_partial_and_retried_independently()
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)
@@ -76,7 +76,9 @@ def test_cli_main_returns_summary_exit_code_and_json(
@staticmethod
def from_token(token: str, **kwargs: object) -> object:
assert token == "secret-token"
assert kwargs
assert kwargs["max_retries"] == 3
assert kwargs["backoff_seconds"] == 1.0
assert kwargs["request_interval_seconds"] == 0.2
return object()
class FakeBuild:
@@ -1,4 +1,5 @@
import logging
import threading
from collections.abc import Mapping
from datetime import UTC, date, datetime
from decimal import Decimal
@@ -24,11 +25,13 @@ 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:
self.calls.append((api_name, kwargs))
partition = str(kwargs.get("ts_code") or kwargs.get("list_status") or "")
response = self.responses.get((api_name, partition), ())
with self._lock:
self.calls.append((api_name, kwargs))
response = self.responses.get((api_name, partition), ())
if isinstance(response, BaseException):
raise response
return response
@@ -44,6 +47,22 @@ def make_adapter(client: object) -> TushareSectorRadarAdapter:
)
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(
{
@@ -98,7 +117,7 @@ def test_daily_and_moneyflow_keep_source_units_and_distinguish_missing_from_zero
adapter = make_adapter(client)
daily = adapter.fetch_daily(TARGET_DATE)
moneyflow = adapter.fetch_moneyflow_dc(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")
@@ -109,6 +128,168 @@ def test_daily_and_moneyflow_keep_source_units_and_distinguish_missing_from_zero
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(
{
@@ -305,32 +486,54 @@ def test_dc_member_preserves_an_explicit_empty_partition() -> None:
assert result.snapshots[1].row_count == 0
def test_stock_basic_explicitly_requests_all_lifecycle_statuses() -> None:
responses = {
(
"stock_basic",
status,
): (
{
"ts_code": f"00000{index}.SZ",
"symbol": f"00000{index}",
"name": status,
"market": None if status == "D" else "主板",
"exchange": "SZSE",
"list_status": status,
"list_date": "20200101",
"delist_date": None,
},
)
for index, status in enumerate(("L", "D", "P", "G", "UN"), start=1)
}
client = QueryClient(responses)
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", "D", "P", "G", "UN"}
assert next(row for row in result.rows if row.list_status == "D").market is None
assert [call[1]["list_status"] for call in client.calls] == ["L", "D", "P", "G", "UN"]
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: