2026-08-31 14:26:28 +08:00
|
|
|
import threading
|
2026-08-06 22:45:18 +08:00
|
|
|
from datetime import date
|
|
|
|
|
|
|
|
|
|
import pytest
|
|
|
|
|
import tushare as ts # pyright: ignore[reportMissingTypeStubs]
|
|
|
|
|
|
|
|
|
|
from zhixing_server.modules.market_data.domain.models import SyncWindow
|
2026-08-11 11:20:13 +08:00
|
|
|
from zhixing_server.modules.market_data.infrastructure.tushare import (
|
|
|
|
|
RequestCoordinator,
|
|
|
|
|
TushareAdapter,
|
|
|
|
|
)
|
2026-08-06 22:45:18 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
def test_from_token_reuses_api_client_for_pro_bar(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
created_client = object()
|
|
|
|
|
calls: list[dict[str, object]] = []
|
|
|
|
|
|
|
|
|
|
def fake_pro_api(token: str) -> object:
|
|
|
|
|
assert token
|
|
|
|
|
return created_client
|
|
|
|
|
|
|
|
|
|
def fake_pro_bar(**kwargs: object) -> list[dict[str, object]]:
|
|
|
|
|
calls.append(kwargs)
|
|
|
|
|
return [
|
|
|
|
|
{
|
|
|
|
|
"ts_code": "000001.SZ",
|
|
|
|
|
"trade_date": "20240102",
|
|
|
|
|
"close": "10",
|
|
|
|
|
}
|
|
|
|
|
]
|
|
|
|
|
|
|
|
|
|
monkeypatch.setattr(ts, "pro_api", fake_pro_api)
|
|
|
|
|
monkeypatch.setattr(ts, "pro_bar", fake_pro_bar)
|
|
|
|
|
|
|
|
|
|
adapter = TushareAdapter.from_token(
|
|
|
|
|
"test-token",
|
|
|
|
|
request_interval_seconds=0,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
bars = adapter.fetch_bars(
|
|
|
|
|
"000001.SZ",
|
|
|
|
|
SyncWindow(start=date(2024, 1, 2), end=date(2024, 1, 2)),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
assert bars[0].ts_code == "000001.SZ"
|
|
|
|
|
assert calls
|
|
|
|
|
assert calls[0]["api"] is created_client
|
|
|
|
|
assert calls[0]["adj"] == "qfq"
|
2026-08-11 11:20:13 +08:00
|
|
|
assert calls[0]["retry_count"] == 1
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def test_rate_limit_cooldown_is_shared_by_following_requests() -> None:
|
|
|
|
|
current = [0.0]
|
|
|
|
|
waits: list[float] = []
|
|
|
|
|
calls: list[float] = []
|
|
|
|
|
|
|
|
|
|
def clock() -> float:
|
|
|
|
|
return current[0]
|
|
|
|
|
|
|
|
|
|
def wait(seconds: float) -> None:
|
|
|
|
|
waits.append(seconds)
|
|
|
|
|
current[0] += seconds
|
|
|
|
|
|
|
|
|
|
coordinator = RequestCoordinator(
|
|
|
|
|
max_retries=0,
|
|
|
|
|
cooldown_seconds=(60, 120, 180),
|
|
|
|
|
clock=clock,
|
|
|
|
|
wait_fn=wait,
|
|
|
|
|
sleep_fn=wait,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
def rate_limited() -> object:
|
|
|
|
|
calls.append(current[0])
|
|
|
|
|
raise RuntimeError("HTTP 429: too many requests")
|
|
|
|
|
|
|
|
|
|
with pytest.raises(RuntimeError):
|
|
|
|
|
coordinator.call("daily", rate_limited)
|
|
|
|
|
|
|
|
|
|
assert calls == [0]
|
|
|
|
|
assert coordinator.cooldown_until == 60
|
|
|
|
|
|
|
|
|
|
result = coordinator.call("adj_factor", lambda: calls.append(current[0]) or "ok")
|
|
|
|
|
|
|
|
|
|
assert result == "ok"
|
|
|
|
|
assert calls == [0, 60]
|
|
|
|
|
assert waits == [60]
|
|
|
|
|
|
|
|
|
|
|
2026-08-31 14:26:28 +08:00
|
|
|
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]
|
|
|
|
|
|
|
|
|
|
|
2026-08-11 11:20:13 +08:00
|
|
|
def test_pro_bar_qfq_calls_are_bound_to_the_shared_coordinator(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
|
|
|
|
class Client:
|
|
|
|
|
def __init__(self) -> None:
|
|
|
|
|
self.calls: list[str] = []
|
|
|
|
|
|
|
|
|
|
def daily(self, **kwargs: object) -> object:
|
|
|
|
|
self.calls.append("daily")
|
|
|
|
|
return object()
|
|
|
|
|
|
|
|
|
|
def adj_factor(self, **kwargs: object) -> object:
|
|
|
|
|
self.calls.append("adj_factor")
|
|
|
|
|
return object()
|
|
|
|
|
|
|
|
|
|
client = Client()
|
|
|
|
|
pro_bar_calls: list[dict[str, object]] = []
|
|
|
|
|
|
|
|
|
|
def fake_pro_api(token: str) -> object:
|
|
|
|
|
return client
|
|
|
|
|
|
|
|
|
|
def fake_pro_bar(**kwargs: object) -> list[dict[str, object]]:
|
|
|
|
|
pro_bar_calls.append(kwargs)
|
|
|
|
|
api = kwargs["api"]
|
|
|
|
|
assert api is client
|
|
|
|
|
api.daily(ts_code="000001.SZ") # type: ignore[union-attr]
|
|
|
|
|
api.adj_factor(ts_code="000001.SZ") # type: ignore[union-attr]
|
|
|
|
|
return [{"ts_code": "000001.SZ", "trade_date": "20240102", "close": "10"}]
|
|
|
|
|
|
|
|
|
|
monkeypatch.setattr(ts, "pro_api", fake_pro_api)
|
|
|
|
|
monkeypatch.setattr(ts, "pro_bar", fake_pro_bar)
|
|
|
|
|
|
|
|
|
|
adapter = TushareAdapter.from_token("test-token", request_interval_seconds=0)
|
|
|
|
|
adapter.fetch_bars(
|
|
|
|
|
"000001.SZ",
|
|
|
|
|
SyncWindow(start=date(2024, 1, 2), end=date(2024, 1, 2)),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
assert client.calls == ["daily", "adj_factor"]
|
|
|
|
|
assert pro_bar_calls[0]["retry_count"] == 1
|