Skip to content

Commit cecc2b0

Browse files
authored
fix schwab market data transient retry
1 parent 8374940 commit cecc2b0

5 files changed

Lines changed: 172 additions & 16 deletions

File tree

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
44

55
[project]
66
name = "quant-platform-kit"
7-
version = "0.7.36"
7+
version = "0.7.37"
88
description = "Shared broker adapters, domain models, execution ports, and notification utilities for QuantStrategyLab strategies."
99
readme = "README.md"
1010
requires-python = ">=3.9"

setup.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33

44
setup(
55
name="quant-platform-kit",
6-
version="0.7.36",
6+
version="0.7.37",
77
description="Shared broker adapters, domain models, execution ports, and notification utilities for QuantStrategyLab strategies.",
88
package_dir={"": "src"},
99
packages=find_packages(where="src"),

src/quant_platform_kit/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
used by older strategy repositories.
55
"""
66

7-
__version__ = "0.7.36"
7+
__version__ = "0.7.37"
88

99
from .common.models import (
1010
ExecutionReport,

src/quant_platform_kit/schwab/market_data.py

Lines changed: 103 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,102 @@
11
from __future__ import annotations
22

3-
from datetime import datetime
4-
from typing import Any
3+
import os
4+
import time
5+
from datetime import datetime, timezone
6+
from email.utils import parsedate_to_datetime
7+
from typing import Any, Callable, Optional
58

69
from quant_platform_kit.common.models import QuoteSnapshot
710

11+
RETRYABLE_STATUS_CODES = {429, 500, 502, 503, 504}
12+
DEFAULT_HTTP_MAX_ATTEMPTS = 4
13+
DEFAULT_HTTP_BACKOFF_SECONDS = 1.0
14+
DEFAULT_HTTP_MAX_BACKOFF_SECONDS = 8.0
15+
16+
17+
def _env_int(name: str, default: int, *, minimum: int, maximum: int) -> int:
18+
raw_value = os.environ.get(name)
19+
if not raw_value:
20+
return default
21+
try:
22+
value = int(raw_value)
23+
except ValueError:
24+
return default
25+
return min(max(value, minimum), maximum)
26+
27+
28+
def _env_float(name: str, default: float, *, minimum: float, maximum: float) -> float:
29+
raw_value = os.environ.get(name)
30+
if not raw_value:
31+
return default
32+
try:
33+
value = float(raw_value)
34+
except ValueError:
35+
return default
36+
return min(max(value, minimum), maximum)
37+
38+
39+
def _header_value(headers: Any, name: str) -> Optional[str]:
40+
if not headers:
41+
return None
42+
if hasattr(headers, "get"):
43+
value = headers.get(name)
44+
if value is None:
45+
value = headers.get(name.lower())
46+
if value is None:
47+
value = headers.get(name.upper())
48+
return str(value).strip() if value is not None else None
49+
return None
50+
51+
52+
def _retry_after_seconds(response: Any, fallback_seconds: float, max_seconds: float) -> float:
53+
raw_value = _header_value(getattr(response, "headers", None), "Retry-After")
54+
if not raw_value:
55+
return min(fallback_seconds, max_seconds)
56+
try:
57+
return min(max(float(raw_value), 0.0), max_seconds)
58+
except ValueError:
59+
pass
60+
61+
try:
62+
retry_at = parsedate_to_datetime(raw_value)
63+
except (TypeError, ValueError):
64+
return min(fallback_seconds, max_seconds)
65+
if retry_at.tzinfo is None:
66+
retry_at = retry_at.replace(tzinfo=timezone.utc)
67+
wait_seconds = (retry_at - datetime.now(timezone.utc)).total_seconds()
68+
return min(max(wait_seconds, 0.0), max_seconds)
69+
70+
71+
def _request_with_retries(request_fn: Callable[[], Any]) -> Any:
72+
max_attempts = _env_int("QPK_SCHWAB_HTTP_MAX_ATTEMPTS", DEFAULT_HTTP_MAX_ATTEMPTS, minimum=1, maximum=8)
73+
backoff_seconds = _env_float(
74+
"QPK_SCHWAB_HTTP_BACKOFF_SECONDS",
75+
DEFAULT_HTTP_BACKOFF_SECONDS,
76+
minimum=0.0,
77+
maximum=30.0,
78+
)
79+
max_backoff_seconds = _env_float(
80+
"QPK_SCHWAB_HTTP_MAX_BACKOFF_SECONDS",
81+
DEFAULT_HTTP_MAX_BACKOFF_SECONDS,
82+
minimum=0.0,
83+
maximum=60.0,
84+
)
85+
86+
response = None
87+
for attempt in range(1, max_attempts + 1):
88+
response = request_fn()
89+
status_code = getattr(response, "status_code", None)
90+
if status_code not in RETRYABLE_STATUS_CODES or attempt >= max_attempts:
91+
return response
92+
93+
fallback_seconds = backoff_seconds * (2 ** (attempt - 1))
94+
wait_seconds = _retry_after_seconds(response, fallback_seconds, max_backoff_seconds)
95+
if wait_seconds > 0:
96+
time.sleep(wait_seconds)
97+
98+
return response
99+
8100

9101
def decode_response_json(response: Any, context: str) -> Any:
10102
if response.status_code not in (200, 201):
@@ -18,12 +110,14 @@ def decode_response_json(response: Any, context: str) -> Any:
18110
def fetch_default_daily_price_history_candles(api_client: Any, symbol: str) -> list[dict[str, Any]]:
19111
from schwab import client
20112

21-
response = api_client.get_price_history(
22-
symbol,
23-
period_type=client.Client.PriceHistory.PeriodType.YEAR,
24-
period=client.Client.PriceHistory.Period.TWO_YEARS,
25-
frequency_type=client.Client.PriceHistory.FrequencyType.DAILY,
26-
frequency=client.Client.PriceHistory.Frequency.DAILY,
113+
response = _request_with_retries(
114+
lambda: api_client.get_price_history(
115+
symbol,
116+
period_type=client.Client.PriceHistory.PeriodType.YEAR,
117+
period=client.Client.PriceHistory.Period.TWO_YEARS,
118+
frequency_type=client.Client.PriceHistory.FrequencyType.DAILY,
119+
frequency=client.Client.PriceHistory.Frequency.DAILY,
120+
)
27121
)
28122
payload = decode_response_json(response, f"{symbol} history")
29123
candles = payload.get("candles")
@@ -33,7 +127,7 @@ def fetch_default_daily_price_history_candles(api_client: Any, symbol: str) -> l
33127

34128

35129
def fetch_quotes(api_client: Any, symbols: list[str] | tuple[str, ...]) -> dict[str, QuoteSnapshot]:
36-
payload = decode_response_json(api_client.get_quotes(symbols), "Quotes")
130+
payload = decode_response_json(_request_with_retries(lambda: api_client.get_quotes(symbols)), "Quotes")
37131
as_of = datetime.utcnow()
38132
snapshots: dict[str, QuoteSnapshot] = {}
39133
for symbol in symbols:

tests/test_schwab_market_data.py

Lines changed: 66 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,13 +5,18 @@
55
import unittest
66
from unittest.mock import patch
77

8-
from quant_platform_kit.schwab.market_data import fetch_default_daily_price_history_candles, fetch_quotes
8+
from quant_platform_kit.schwab.market_data import (
9+
decode_response_json,
10+
fetch_default_daily_price_history_candles,
11+
fetch_quotes,
12+
)
913

1014

1115
class FakeResponse:
12-
def __init__(self, payload, status_code=200):
16+
def __init__(self, payload, status_code=200, headers=None):
1317
self._payload = payload
1418
self.status_code = status_code
19+
self.headers = headers or {}
1520
self.text = str(payload)
1621

1722
def json(self):
@@ -32,7 +37,7 @@ def get_quotes(self, symbols):
3237

3338

3439
class SchwabMarketDataTests(unittest.TestCase):
35-
def test_fetch_default_daily_price_history_candles(self) -> None:
40+
def _install_fake_schwab_module(self):
3641
schwab_module = types.ModuleType("schwab")
3742
client_module = types.ModuleType("schwab.client")
3843
client_module.Client = types.SimpleNamespace(
@@ -43,13 +48,70 @@ def test_fetch_default_daily_price_history_candles(self) -> None:
4348
Frequency=types.SimpleNamespace(DAILY="DAILY"),
4449
)
4550
)
51+
return patch.dict(sys.modules, {"schwab": schwab_module, "schwab.client": client_module})
4652

47-
with patch.dict(sys.modules, {"schwab": schwab_module, "schwab.client": client_module}):
53+
def test_fetch_default_daily_price_history_candles(self) -> None:
54+
with self._install_fake_schwab_module():
4855
candles = fetch_default_daily_price_history_candles(FakeClient(), "QQQ")
4956

5057
self.assertEqual(len(candles), 2)
5158
self.assertEqual(candles[-1]["close"], 11.0)
5259

60+
def test_fetch_default_daily_price_history_retries_rate_limit(self) -> None:
61+
class RateLimitedClient:
62+
def __init__(self):
63+
self.calls = 0
64+
65+
def get_price_history(self, symbol, **_kwargs):
66+
self.calls += 1
67+
if self.calls == 1:
68+
return FakeResponse({"error": "rate limited"}, status_code=429, headers={"Retry-After": "0.25"})
69+
return FakeResponse({"candles": [{"close": 12.0}]})
70+
71+
rate_limited_client = RateLimitedClient()
72+
with self._install_fake_schwab_module(), patch(
73+
"quant_platform_kit.schwab.market_data.time.sleep"
74+
) as sleep_mock:
75+
candles = fetch_default_daily_price_history_candles(rate_limited_client, "SOXL")
76+
77+
self.assertEqual(candles, [{"close": 12.0}])
78+
self.assertEqual(rate_limited_client.calls, 2)
79+
sleep_mock.assert_called_once_with(0.25)
80+
81+
def test_fetch_quotes_retries_transient_server_error(self) -> None:
82+
class FlakyQuoteClient:
83+
def __init__(self):
84+
self.calls = 0
85+
86+
def get_quotes(self, symbols):
87+
self.calls += 1
88+
if self.calls < 3:
89+
return FakeResponse({"error": "unavailable"}, status_code=503)
90+
return FakeClient().get_quotes(symbols)
91+
92+
flaky_client = FlakyQuoteClient()
93+
with patch("quant_platform_kit.schwab.market_data.time.sleep") as sleep_mock:
94+
snapshots = fetch_quotes(flaky_client, ["TQQQ"])
95+
96+
self.assertEqual(snapshots["TQQQ"].last_price, 100.0)
97+
self.assertEqual(flaky_client.calls, 3)
98+
self.assertEqual([call.args[0] for call in sleep_mock.call_args_list], [1.0, 2.0])
99+
100+
def test_retry_exhaustion_keeps_original_error_context(self) -> None:
101+
class AlwaysRateLimitedClient:
102+
def get_price_history(self, symbol, **_kwargs):
103+
return FakeResponse({"error": "rate limited"}, status_code=429)
104+
105+
with self._install_fake_schwab_module(), patch(
106+
"quant_platform_kit.schwab.market_data.time.sleep"
107+
), patch.dict("os.environ", {"QPK_SCHWAB_HTTP_MAX_ATTEMPTS": "2"}):
108+
with self.assertRaisesRegex(RuntimeError, "SOXL history failed: 429"):
109+
fetch_default_daily_price_history_candles(AlwaysRateLimitedClient(), "SOXL")
110+
111+
def test_decode_response_json_still_reports_non_retryable_errors(self) -> None:
112+
with self.assertRaisesRegex(RuntimeError, "Quotes failed: 400"):
113+
decode_response_json(FakeResponse({"error": "bad request"}, status_code=400), "Quotes")
114+
53115
def test_fetch_quotes_returns_snapshots(self) -> None:
54116
snapshots = fetch_quotes(FakeClient(), ["TQQQ", "BOXX"])
55117

0 commit comments

Comments
 (0)