11"""Tests for the rate-limiting middleware."""
22
33import asyncio
4+ import threading
45import time
56from types import SimpleNamespace
67from 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+
221308def 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 ))
232319def 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
320407def 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