/
/
/
1"""Tests for the throttle_with_retries decorator and ThrottlerManager."""
2
3from __future__ import annotations
4
5import logging
6from collections.abc import Generator, Sequence
7from unittest.mock import AsyncMock, patch
8
9import pytest
10from music_assistant_models.errors import (
11 RateLimited,
12 ResourceTemporarilyUnavailable,
13 RetriesExhausted,
14)
15
16from music_assistant.helpers.throttle_retry import (
17 ThrottlerManager,
18 parse_retry_after,
19 throttle_with_retries,
20)
21
22
23class FakeProvider:
24 """
25 Minimal provider stub for testing the decorator.
26
27 The decorator requires `self.throttler` and `self.logger`.
28 """
29
30 throttler = ThrottlerManager(rate_limit=100, period=0.01, retry_attempts=5, initial_backoff=4)
31
32 def __init__(self) -> None:
33 """Initialize."""
34 self.logger = logging.getLogger("test.fake_provider")
35 self.call_count = 0
36 self._side_effects: list[Exception | str] = []
37
38 def set_side_effects(self, effects: Sequence[Exception | str]) -> None:
39 """
40 Configure what happens on each call.
41
42 :param effects: List of exceptions to raise, or "ok" to return successfully.
43 """
44 self._side_effects = list(effects)
45 self.call_count = 0
46
47 @throttle_with_retries
48 async def api_call(self, value: str) -> str:
49 """Simulate an API call."""
50 self.call_count += 1
51 if self._side_effects:
52 effect = self._side_effects.pop(0)
53 if isinstance(effect, Exception):
54 raise effect
55 return value
56
57
58@pytest.fixture
59def provider() -> FakeProvider:
60 """Create a FakeProvider with fast throttler for tests."""
61 return FakeProvider()
62
63
64@pytest.fixture
65def mock_sleep() -> Generator[AsyncMock]:
66 """Patch asyncio.sleep to capture sleep times without actually sleeping."""
67 with patch(
68 "music_assistant.helpers.throttle_retry.asyncio.sleep", new_callable=AsyncMock
69 ) as mock:
70 yield mock
71
72
73class TestBasicBehavior:
74 """Basic success/failure behavior."""
75
76 async def test_successful_call(self, provider: FakeProvider) -> None:
77 """Successful call passes through cleanly."""
78 result = await provider.api_call("hello")
79 assert result == "hello"
80 assert provider.call_count == 1
81
82 async def test_retries_exhausted(self, provider: FakeProvider, mock_sleep: AsyncMock) -> None:
83 """Exhausting all retries raises RetriesExhausted."""
84 provider.set_side_effects([ResourceTemporarilyUnavailable("fail")] * 5)
85 with pytest.raises(RetriesExhausted):
86 await provider.api_call("test")
87 assert provider.call_count == 5
88
89 async def test_recovery_after_failures(
90 self, provider: FakeProvider, mock_sleep: AsyncMock
91 ) -> None:
92 """Succeeds after transient failures."""
93 provider.set_side_effects(
94 [
95 ResourceTemporarilyUnavailable("fail"),
96 ResourceTemporarilyUnavailable("fail"),
97 "ok",
98 ]
99 )
100 result = await provider.api_call("recovered")
101 assert result == "recovered"
102 assert provider.call_count == 3
103
104
105class TestServerProvidedBackoff:
106 """When the server names a recovery time (e.g. 503 Retry-After), respect it."""
107
108 async def test_server_backoff_respected(
109 self, provider: FakeProvider, mock_sleep: AsyncMock
110 ) -> None:
111 """Server-provided backoff should be honored without doubling."""
112 provider.set_side_effects(
113 [
114 ResourceTemporarilyUnavailable("unavailable", backoff_time=2),
115 ResourceTemporarilyUnavailable("unavailable", backoff_time=2),
116 ResourceTemporarilyUnavailable("unavailable", backoff_time=2),
117 "ok",
118 ]
119 )
120 result = await provider.api_call("ok")
121 assert result == "ok"
122 assert provider.call_count == 4
123
124 # All three retries sleep ~2s (never less, up to +10% citizen jitter), not 2â4â8
125 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
126 for t in sleep_times:
127 assert 2.0 <= t <= 2.2
128
129 async def test_varying_server_backoff(
130 self, provider: FakeProvider, mock_sleep: AsyncMock
131 ) -> None:
132 """Each retry should use the server's current value."""
133 provider.set_side_effects(
134 [
135 ResourceTemporarilyUnavailable("unavailable", backoff_time=3),
136 ResourceTemporarilyUnavailable("unavailable", backoff_time=5),
137 ResourceTemporarilyUnavailable("unavailable", backoff_time=1),
138 "ok",
139 ]
140 )
141 result = await provider.api_call("ok")
142 assert result == "ok"
143
144 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
145 assert 3.0 <= sleep_times[0] <= 3.3
146 assert 5.0 <= sleep_times[1] <= 5.5
147 assert 1.0 <= sleep_times[2] <= 1.1
148
149 async def test_negative_backoff_falls_back_to_exponential(
150 self, provider: FakeProvider, mock_sleep: AsyncMock
151 ) -> None:
152 """A malformed negative Retry-After must not yield a non-positive sleep."""
153 provider.set_side_effects(
154 [ResourceTemporarilyUnavailable("bad header", backoff_time=-5), "ok"]
155 )
156 await provider.api_call("ok")
157
158 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
159 assert 3.0 <= sleep_times[0] <= 5.0
160
161
162class TestRateLimited:
163 """When rate-limited (429), Retry-After is a floor and we escalate above it."""
164
165 async def test_floor_is_never_undercut(
166 self, provider: FakeProvider, mock_sleep: AsyncMock
167 ) -> None:
168 """A large Retry-After dominates until exponential backoff catches up."""
169 provider.set_side_effects([RateLimited("rate limited", backoff_time=50)] * 3 + ["ok"])
170 await provider.api_call("ok")
171
172 # exp_backoff starts at 4 and doubles; 50 stays the floor for these retries
173 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
174 for t in sleep_times:
175 assert 50.0 <= t <= 55.0
176
177 async def test_escalates_above_floor(
178 self, provider: FakeProvider, mock_sleep: AsyncMock
179 ) -> None:
180 """A small Retry-After is honored as a floor while exponential backoff grows."""
181 provider.set_side_effects([RateLimited("rate limited", backoff_time=2)] * 4 + ["ok"])
182 await provider.api_call("ok")
183
184 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
185 # exp from initial=4 dominates the 2s floor and doubles each retry (up to +10%)
186 assert 4.0 <= sleep_times[0] <= 4.4
187 assert 8.0 <= sleep_times[1] <= 8.8
188 assert 16.0 <= sleep_times[2] <= 17.6
189
190 async def test_absurd_retry_after_capped(
191 self, provider: FakeProvider, mock_sleep: AsyncMock
192 ) -> None:
193 """A hostile Retry-After is clamped to MAX_RETRY_AFTER (1 hour)."""
194 provider.set_side_effects([RateLimited("rate limited", backoff_time=999999), "ok"])
195 await provider.api_call("ok")
196
197 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
198 assert 3600.0 <= sleep_times[0] <= 3960.0
199
200
201class TestExponentialBackoffWithJitter:
202 """When no server backoff is provided, use exponential backoff with jitter."""
203
204 async def test_exponential_backoff_increases(
205 self, provider: FakeProvider, mock_sleep: AsyncMock
206 ) -> None:
207 """Backoff should roughly double each retry (with jitter)."""
208 provider.set_side_effects(
209 [
210 ResourceTemporarilyUnavailable("error"),
211 ResourceTemporarilyUnavailable("error"),
212 ResourceTemporarilyUnavailable("error"),
213 ResourceTemporarilyUnavailable("error"),
214 "ok",
215 ]
216 )
217 result = await provider.api_call("ok")
218 assert result == "ok"
219
220 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
221 assert len(sleep_times) == 4
222
223 # With initial_backoff=4 and jitter ±25%:
224 # Attempt 0: base=4, range [3.0, 5.0]
225 # Attempt 1: base=8, range [6.0, 10.0]
226 # Attempt 2: base=16, range [12.0, 20.0]
227 # Attempt 3: base=32, range [24.0, 40.0]
228 assert 3.0 <= sleep_times[0] <= 5.0
229 assert 6.0 <= sleep_times[1] <= 10.0
230 assert 12.0 <= sleep_times[2] <= 20.0
231 assert 24.0 <= sleep_times[3] <= 40.0
232
233 async def test_backoff_capped_at_max(self, mock_sleep: AsyncMock) -> None:
234 """Exponential backoff should not exceed MAX_BACKOFF (120s)."""
235 provider = FakeProvider()
236 # Override with a very high initial_backoff
237 provider.throttler = ThrottlerManager(
238 rate_limit=100, period=0.01, retry_attempts=5, initial_backoff=100
239 )
240 provider.set_side_effects(
241 [
242 ResourceTemporarilyUnavailable("error"),
243 ResourceTemporarilyUnavailable("error"),
244 ResourceTemporarilyUnavailable("error"),
245 ResourceTemporarilyUnavailable("error"),
246 "ok",
247 ]
248 )
249 await provider.api_call("ok")
250
251 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
252 # Jitter is applied before capping, so no value should exceed MAX_BACKOFF (120)
253 for t in sleep_times:
254 assert t <= 120.0
255
256
257class TestMixedBackoff:
258 """Test switching between server-provided and exponential backoff."""
259
260 async def test_server_then_exponential(
261 self, provider: FakeProvider, mock_sleep: AsyncMock
262 ) -> None:
263 """After server-provided backoff, exponential continues from its own counter."""
264 provider.set_side_effects(
265 [
266 # First: server says wait 2s
267 ResourceTemporarilyUnavailable("rate limited", backoff_time=2),
268 # Then: no server guidance, fall back to exponential
269 ResourceTemporarilyUnavailable("error"),
270 ResourceTemporarilyUnavailable("error"),
271 "ok",
272 ]
273 )
274 result = await provider.api_call("ok")
275 assert result == "ok"
276
277 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
278 # Retry 1: server says 2 (honored, up to +10% jitter)
279 assert 2.0 <= sleep_times[0] <= 2.2
280 # Retry 2: exponential starts at initial=4 (unchanged by server retry), jitter ±25%
281 assert 3.0 <= sleep_times[1] <= 5.0
282 # Retry 3: doubled to 8, jitter ±25%
283 assert 6.0 <= sleep_times[2] <= 10.0
284
285 async def test_exponential_then_server(
286 self, provider: FakeProvider, mock_sleep: AsyncMock
287 ) -> None:
288 """Server-provided backoff overrides even after exponential was growing."""
289 provider.set_side_effects(
290 [
291 ResourceTemporarilyUnavailable("error"), # no server guidance
292 ResourceTemporarilyUnavailable("error"), # no server guidance
293 ResourceTemporarilyUnavailable("unavailable", backoff_time=1), # server says 1
294 "ok",
295 ]
296 )
297 result = await provider.api_call("ok")
298 assert result == "ok"
299
300 sleep_times = [call.args[0] for call in mock_sleep.call_args_list]
301 # Retries 1-2: exponential (4 jittered, 8 jittered)
302 assert 3.0 <= sleep_times[0] <= 5.0
303 assert 6.0 <= sleep_times[1] <= 10.0
304 # Retry 3: server says 1, honored (up to +10% jitter)
305 assert 1.0 <= sleep_times[2] <= 1.1
306
307
308class TestParseRetryAfter:
309 """Tests for RFC 9110 Retry-After header parsing."""
310
311 def test_none_returns_zero(self) -> None:
312 """Missing header returns 0."""
313 assert parse_retry_after(None) == 0
314
315 def test_integer_string(self) -> None:
316 """Standard delay-seconds format."""
317 assert parse_retry_after("120") == 120
318 assert parse_retry_after("0") == 0
319 assert parse_retry_after("1") == 1
320
321 def test_negative_clamped_to_zero(self) -> None:
322 """Negative values (non-conforming) are clamped to 0."""
323 assert parse_retry_after("-5") == 0
324
325 def test_http_date(self) -> None:
326 """RFC 9110 HTTP-date format returns seconds until that time."""
327 import datetime # noqa: PLC0415
328
329 # Create a date 60 seconds in the future
330 future = datetime.datetime.now(tz=datetime.UTC) + datetime.timedelta(seconds=60)
331 date_str = future.strftime("%a, %d %b %Y %H:%M:%S GMT")
332 result = parse_retry_after(date_str)
333 # Allow ±2 seconds tolerance for test execution time
334 assert 58 <= result <= 62
335
336 def test_http_date_in_past(self) -> None:
337 """HTTP-date in the past returns 0 (clamped)."""
338 assert parse_retry_after("Mon, 01 Jan 2024 00:00:00 GMT") == 0
339
340 def test_garbage_returns_zero(self) -> None:
341 """Unparsable values return 0."""
342 assert parse_retry_after("not-a-number") == 0
343 assert parse_retry_after("") == 0
344 assert parse_retry_after("1.5") == 0 # floats are not valid per RFC 9110
345