Skip to content

Commit fc7594e

Browse files
Give clone() the host it is being built for (#27)
1 parent 0787458 commit fc7594e

3 files changed

Lines changed: 148 additions & 43 deletions

File tree

‎aiohttp_client_middlewares/rate_limit.py‎

Lines changed: 35 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,6 @@
1313
import math
1414
import time
1515
from abc import ABC, abstractmethod
16-
from collections import defaultdict
1716

1817
from aiohttp import ClientHandlerType, ClientRequest, ClientResponse, ClientTimeout
1918

@@ -46,11 +45,15 @@ async def acquire(self) -> float:
4645
"""
4746

4847
@abstractmethod
49-
def clone(self) -> "RateLimiter":
50-
"""Return a fresh limiter with the same configuration.
51-
52-
Per-domain mode clones the configured limiter once per target
53-
host, so state (queued slots, accrued tokens) must not carry over.
48+
def clone(self, host: str, /) -> "RateLimiter":
49+
"""Return a fresh limiter, configured the same, scoped to *host*.
50+
51+
Per-domain mode calls this the first time it meets a host, so state
52+
(queued slots, accrued tokens) must not carry over. Threads racing on
53+
that first contact may each build one and only one is kept, so the
54+
call itself should have no side effects. An algorithm that keeps its
55+
state in-process can ignore *host*; one that keeps it in a shared
56+
backend needs it in the key, or every host draws on one limit.
5457
"""
5558

5659
def release(self) -> None:
@@ -147,8 +150,11 @@ async def acquire(self) -> float:
147150
self._tokens -= 1.0
148151
return max(0.0, -self._tokens) * self._interval
149152

150-
def clone(self) -> "TokenBucket":
151-
"""Return a fresh, full bucket with the same rate and burst."""
153+
def clone(self, host: str, /) -> "TokenBucket":
154+
"""Return a fresh, full bucket with the same rate and burst.
155+
156+
The bucket's state is per-object, so *host* needs no part in it.
157+
"""
152158
return TokenBucket(rate=self._rate, burst=int(self._burst))
153159

154160
def release(self) -> None:
@@ -182,7 +188,7 @@ class RateLimitMiddleware:
182188
:param RateLimiter limiter: The :class:`RateLimiter` to throttle with --
183189
for example ``TokenBucket(rate=5.0, burst=2)``. With
184190
``per_domain=True`` it acts as a template: each target host gets
185-
``limiter.clone()`` the first time that host is seen.
191+
``limiter.clone(host)`` the first time that host is seen.
186192
:param bool per_domain: When ``True``, keep an independent limiter per
187193
target host instead of a single global one. Limiters are keyed on the
188194
URL host only (port and scheme are not distinguished) and are never
@@ -198,31 +204,39 @@ def __init__(
198204
) -> None:
199205
if not isinstance(limiter, RateLimiter):
200206
raise TypeError(f"limiter must be a RateLimiter, got {limiter!r}")
201-
self._per_domain = per_domain
202-
self._global_limiter: RateLimiter | None = None
203-
if per_domain:
204-
self._domain_limiters: dict[str, RateLimiter] = defaultdict(limiter.clone)
205-
else:
206-
self._global_limiter = limiter
207+
# The one limiter in global mode; the template to clone in per-domain
208+
# mode. Whether the per-host dict exists is what says which.
209+
self._limiter = limiter
210+
self._domain_limiters: dict[str, RateLimiter] | None = (
211+
{} if per_domain else None
212+
)
207213

208214
@property
209215
def per_domain(self) -> bool:
210216
"""Whether this middleware keeps one limiter per target host.
211217
212-
Read-only: the limiters are built once in ``__init__``, so flipping
213-
this afterwards could not take effect.
218+
Read-only: which limiter a request gets is decided in ``__init__``,
219+
so flipping this afterwards could not take effect.
214220
"""
215-
return self._per_domain
221+
return self._domain_limiters is not None
216222

217223
def _get_limiter(self, request: ClientRequest) -> RateLimiter:
218-
if self._global_limiter is not None:
219-
return self._global_limiter
224+
limiters = self._domain_limiters
225+
if limiters is None:
226+
return self._limiter
220227
# aiohttp raises InvalidUrlClientError for host-less URLs before
221228
# any middleware runs (on redirects too), so ``host`` is only
222229
# ``None`` in the type; the assert narrows it for mypy.
223230
domain = request.url.host
224231
assert domain is not None
225-
return self._domain_limiters[domain]
232+
limiter = limiters.get(domain)
233+
if limiter is None:
234+
# setdefault, not an assignment: threads racing for a host they
235+
# have not seen before must all leave with the limiter that was
236+
# stored, or each gets a private full budget and the burst
237+
# allowance is briefly multiplied by the number of racers.
238+
limiter = limiters.setdefault(domain, self._limiter.clone(domain))
239+
return limiter
226240

227241
async def __call__(
228242
self,

‎docs/api.rst‎

Lines changed: 14 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -86,8 +86,10 @@ Rate limiting
8686
:class:`RateLimitMiddleware` accepts. Implementations provide async
8787
``acquire()``, which reserves a slot and returns its delay as a non-negative
8888
finite number of seconds -- ``wait()`` takes that on trust -- and
89-
``clone()``, which returns a fresh limiter with the same configuration
90-
(used once per host by per-domain mode)::
89+
``clone(host)``, which returns a fresh limiter with the same configuration
90+
scoped to one host (called when per-domain mode first meets a host; a
91+
racing thread's extra clone is discarded, so it should have no side
92+
effects)::
9193

9294
class RedisLimiter(RateLimiter):
9395
def __init__(self, redis, key, script):
@@ -100,14 +102,16 @@ Rate limiting
100102
ms = await self._redis.evalsha(self._script, 1, self._key)
101103
return ms / 1000
102104

103-
def clone(self):
104-
return RedisLimiter(self._redis, self._key, self._script)
105+
def clone(self, host):
106+
return RedisLimiter(self._redis, f"{self._key}:{host}", self._script)
105107

106-
The sketch keeps one key across clones, and ``clone()`` takes no arguments
107-
so it cannot learn the host; give each host its own limiter rather than
108-
using ``per_domain=True``. It also leaves ``release()`` at the default
109-
no-op, and a cancelled round trip can leave a reservation nobody holds;
110-
an expiry on each reservation covers both.
108+
Putting *host* in the key is what makes ``per_domain=True`` mean a budget
109+
per host for a shared backend; a limiter that keeps its state in-process,
110+
like :class:`TokenBucket`, has nothing to key and can ignore it. Redirects
111+
pick hosts too, so give those keys an expiry of their own rather than let a
112+
shared backend keep one for every host ever seen. The sketch also leaves
113+
``release()`` at the default no-op, and a cancelled round trip can leave a
114+
reservation nobody holds; an expiry on each reservation covers both.
111115

112116
``wait(timeout=None)`` is supplied by the base class. It charges async
113117
acquisition against *timeout* once ``acquire()`` returns -- without bounding
@@ -150,7 +154,7 @@ Rate limiting
150154

151155
:param RateLimiter limiter: The :class:`RateLimiter` to throttle with --
152156
for example ``TokenBucket(rate=5.0, burst=2)``. With ``per_domain=True``
153-
it acts as a template: each target host gets ``limiter.clone()`` the
157+
it acts as a template: each target host gets ``limiter.clone(host)`` the
154158
first time that host is seen.
155159
:param bool per_domain: Keep an independent limiter per target host instead
156160
of a single global one. Limiters are keyed on the URL host only (port

‎tests/test_rate_limit.py‎

Lines changed: 99 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
"""Tests for the rate-limiting middleware."""
22

33
import asyncio
4+
import threading
45
import time
56
from types import SimpleNamespace
67
from unittest import mock
@@ -215,25 +216,111 @@ def test_per_domain_uses_distinct_limiters() -> None:
215216

216217
assert limiter_a is not limiter_b # distinct hosts -> isolated limiters
217218
assert limiter_a is limiter_a_again # same host -> same limiter
219+
assert middleware._domain_limiters is not None
218220
assert len(middleware._domain_limiters) == 2
219221

220222

223+
class _KeyedByHost(RateLimiter):
224+
"""A limiter whose state lives under a key, as a Redis-backed one would.
225+
226+
Every host it is cloned for is recorded on a list the clones share with
227+
their template, so a test can see how often the middleware asks.
228+
"""
229+
230+
def __init__(
231+
self, prefix: str, host: str = "", clones: list[str] | None = None
232+
) -> None:
233+
self._prefix = prefix
234+
self.key = f"{prefix}:{host}" if host else prefix
235+
self.clones = [] if clones is None else clones
236+
237+
async def acquire(self) -> float:
238+
return 0.0
239+
240+
def clone(self, host: str, /) -> "_KeyedByHost":
241+
self.clones.append(host)
242+
return _KeyedByHost(self._prefix, host, self.clones)
243+
244+
245+
def test_clone_is_given_the_host_it_is_built_for() -> None:
246+
"""A shared backend can only scope per host if clone() is told the host.
247+
248+
A host seen again gets the stored limiter back without another clone.
249+
"""
250+
template = _KeyedByHost("rl")
251+
middleware = RateLimitMiddleware(template, per_domain=True)
252+
253+
a = middleware._get_limiter(_fake_request("a.example"))
254+
b = middleware._get_limiter(_fake_request("b.example"))
255+
a_again = middleware._get_limiter(_fake_request("a.example"))
256+
257+
assert isinstance(a, _KeyedByHost) and isinstance(b, _KeyedByHost)
258+
assert (a.key, b.key) == ("rl:a.example", "rl:b.example")
259+
assert a_again is a
260+
assert template.clones == ["a.example", "b.example"]
261+
262+
263+
class _BlocksInClone(RateLimiter):
264+
"""Limiter whose ``clone()`` holds each caller until they have all arrived."""
265+
266+
def __init__(self, barrier: threading.Barrier, host: str = "") -> None:
267+
self._barrier = barrier
268+
self.host = host
269+
270+
async def acquire(self) -> float:
271+
return 0.0
272+
273+
def clone(self, host: str, /) -> "_BlocksInClone":
274+
self._barrier.wait()
275+
return _BlocksInClone(self._barrier, host)
276+
277+
278+
def test_first_contact_hands_every_racer_the_stored_limiter() -> None:
279+
"""Threads meeting a new host at once must all leave with the same limiter.
280+
281+
Each keeping the one it built would hand every racer a private, full
282+
budget, multiplying the burst allowance by the number of racers for that
283+
instant. ``dict.setdefault`` is what makes the winner the one everybody
284+
gets; a plain assignment does not. The barrier sits inside ``clone()``,
285+
so this also pins that the miss path stays lock-free: every racer builds
286+
one and the extra clones are discarded.
287+
"""
288+
racers = 4
289+
barrier = threading.Barrier(racers, timeout=10)
290+
middleware = RateLimitMiddleware(_BlocksInClone(barrier), per_domain=True)
291+
handed_out: list[RateLimiter] = []
292+
293+
def race() -> None:
294+
handed_out.append(middleware._get_limiter(_fake_request("new.example")))
295+
296+
threads = [threading.Thread(target=race) for _ in range(racers)]
297+
for thread in threads:
298+
thread.start()
299+
for thread in threads:
300+
thread.join(timeout=10)
301+
302+
assert len(handed_out) == racers
303+
assert middleware._domain_limiters is not None
304+
stored = middleware._domain_limiters["new.example"]
305+
assert all(limiter is stored for limiter in handed_out)
306+
307+
221308
def test_global_mode_shares_one_limiter() -> None:
222309
"""Without ``per_domain`` every host goes through the one limiter."""
223310
middleware = RateLimitMiddleware(TokenBucket(rate=10.0, burst=1))
224311

225312
limiter_a = middleware._get_limiter(_fake_request("a.example"))
226313
limiter_b = middleware._get_limiter(_fake_request("b.example"))
227314

228-
assert limiter_a is limiter_b is middleware._global_limiter
315+
assert limiter_a is limiter_b is middleware._limiter
229316

230317

231318
@pytest.mark.parametrize("per_domain", (False, True))
232319
def test_per_domain_is_readable_and_read_only(per_domain: bool) -> None:
233320
"""``per_domain`` reports the configured mode and cannot be reassigned.
234321
235-
The limiters are built once in ``__init__``, so a writable attribute
236-
would silently do nothing.
322+
Which limiter a request gets is decided in ``__init__``, so a writable
323+
attribute would silently do nothing.
237324
"""
238325
middleware = RateLimitMiddleware(
239326
TokenBucket(rate=10.0, burst=1), per_domain=per_domain
@@ -276,7 +363,7 @@ async def test_middleware_bails_before_sleeping_when_timeout_known() -> None:
276363
async def handler(req: ClientRequest) -> ClientResponse:
277364
raise AssertionError("a doomed request must never be sent")
278365

279-
bucket = middleware._global_limiter
366+
bucket = middleware._limiter
280367
assert isinstance(bucket, TokenBucket)
281368
assert await bucket.acquire() == 0.0 # drain the burst slot
282369

@@ -296,7 +383,7 @@ async def test_middleware_cancel_during_sleep_releases_slot() -> None:
296383
async def handler(req: ClientRequest) -> ClientResponse:
297384
raise AssertionError("the cancelled request must never be sent")
298385

299-
bucket = middleware._global_limiter
386+
bucket = middleware._limiter
300387
assert isinstance(bucket, TokenBucket)
301388
assert await bucket.acquire() == 0.0 # drain the burst slot
302389

@@ -314,7 +401,7 @@ def test_limiter_injection_is_used_directly() -> None:
314401
"""The caller-provided limiter is the one throttling, not a copy of it."""
315402
bucket = TokenBucket(rate=100.0, burst=1)
316403
middleware = RateLimitMiddleware(bucket)
317-
assert middleware._global_limiter is bucket
404+
assert middleware._limiter is bucket
318405

319406

320407
def test_non_limiter_rejected() -> None:
@@ -330,7 +417,7 @@ async def test_token_bucket_clone_is_fresh(clock: _FakeClock) -> None:
330417
await template.acquire()
331418
assert await template.acquire() > 0.0 # template drained into debt
332419

333-
fresh = template.clone()
420+
fresh = template.clone("example.com")
334421
assert await fresh.acquire() == 0.0 # full burst again
335422
assert await fresh.acquire() == 0.0
336423
assert await fresh.acquire() == pytest.approx(0.1) # same rate as the template
@@ -355,7 +442,7 @@ def __init__(self, delay: float) -> None:
355442
async def acquire(self) -> float:
356443
return self._delay
357444

358-
def clone(self) -> "_FixedDelay":
445+
def clone(self, host: str, /) -> "_FixedDelay":
359446
return _FixedDelay(self._delay)
360447

361448

@@ -405,7 +492,7 @@ async def acquire(self) -> float:
405492
await asyncio.sleep(0)
406493
return 0.0
407494

408-
def clone(self) -> "_AwaitsToReserve":
495+
def clone(self, host: str, /) -> "_AwaitsToReserve":
409496
return _AwaitsToReserve()
410497

411498

@@ -442,7 +529,7 @@ async def acquire(self) -> float:
442529
def release(self) -> None:
443530
self.releases += 1
444531

445-
def clone(self) -> "_ElapsedAcquire":
532+
def clone(self, host: str, /) -> "_ElapsedAcquire":
446533
return _ElapsedAcquire(self._clock, self._spend, self._delay)
447534

448535

@@ -489,7 +576,7 @@ async def acquire(self) -> float:
489576
def release(self) -> None:
490577
self.releases += 1
491578

492-
def clone(self) -> "_CancelledAcquire":
579+
def clone(self, host: str, /) -> "_CancelledAcquire":
493580
return _CancelledAcquire()
494581

495582

@@ -572,7 +659,7 @@ async def acquire(self) -> float:
572659
def release(self) -> None:
573660
self.releases += 1
574661

575-
def clone(self) -> "_RecordsReleases":
662+
def clone(self, host: str, /) -> "_RecordsReleases":
576663
return _RecordsReleases(self._delay, self._error)
577664

578665

0 commit comments

Comments
 (0)