fix(sector-radar): 支持当前上市股票资金流补拉
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user