import threading from datetime import date import pytest import tushare as ts # pyright: ignore[reportMissingTypeStubs] from zhixing_server.modules.market_data.domain.models import SyncWindow from zhixing_server.modules.market_data.infrastructure.tushare import ( RequestCoordinator, TushareAdapter, ) 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" 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] 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: 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