Skip to content

Commit 1bed3fb

Browse files
Add a client-side rate limiting middleware
Add RateLimitMiddleware, a token-bucket client middleware that throttles outgoing requests to a configurable rate and burst, with optional per-domain buckets and numeric Retry-After handling on HTTP 429. Promoted from the example in aiohttp#11969 (which was moved here): drop the demo and module-level logging config, keep a strong reference to the scheduler task so it cannot be garbage collected mid-run, and add tests and documentation.
1 parent 5bf6380 commit 1bed3fb

7 files changed

Lines changed: 329 additions & 4 deletions

File tree

‎CHANGES/12.feature.rst‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
Added :class:`~aiohttp_client_middlewares.RateLimitMiddleware`, a client-side
2+
token-bucket rate limiter with optional per-domain buckets and ``Retry-After``
3+
handling -- by :user:`rodrigobnogueira`.
Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,12 @@
11
"""Client middlewares for :mod:`aiohttp`.
22
33
This package is the canonical home for reusable aiohttp *client* middlewares,
4-
starting with HTTP Digest authentication.
4+
starting with HTTP Digest authentication and client-side rate limiting.
55
"""
66

77
from .digest_auth import DigestAuthMiddleware
8+
from .rate_limit import RateLimitMiddleware
89

910
__version__ = "0.1.0"
1011

11-
__all__ = ("DigestAuthMiddleware",)
12+
__all__ = ("DigestAuthMiddleware", "RateLimitMiddleware")
Lines changed: 146 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,146 @@
1+
"""Client-side rate-limiting middleware for aiohttp.
2+
3+
This middleware throttles outgoing requests using a token-bucket algorithm.
4+
It is *not* server-side rate limiting -- it limits how fast the client sends
5+
requests so it does not overwhelm upstream servers or exceed API quotas.
6+
7+
Features:
8+
- Configurable rate and burst size
9+
- Optional per-domain buckets
10+
- Automatic ``Retry-After`` header handling
11+
"""
12+
13+
import asyncio
14+
import logging
15+
import time
16+
from collections import defaultdict, deque
17+
from http import HTTPStatus
18+
19+
from aiohttp import ClientHandlerType, ClientRequest, ClientResponse
20+
21+
_LOGGER = logging.getLogger(__name__)
22+
23+
24+
class TokenBucket:
25+
"""FIFO token-bucket using an ``asyncio.Event`` queue.
26+
27+
Each caller appends its own event to a FIFO queue and waits. A single
28+
``_schedule`` coroutine services the queue front-to-back, sleeping until
29+
each slot's send time arrives and then unblocking the corresponding
30+
caller. This guarantees strict FIFO ordering even under high concurrency.
31+
"""
32+
33+
def __init__(self, rate: float, burst: int) -> None:
34+
self._interval = 1.0 / rate
35+
self._burst = burst
36+
# Start *burst* intervals in the past so the first ``burst`` acquires
37+
# are instant.
38+
self._next_send = time.monotonic() - burst * self._interval
39+
self._waiters: deque[asyncio.Event] = deque()
40+
self._scheduling = False
41+
# Keep a strong reference to the scheduler task: the event loop only
42+
# holds a weak reference to it, so an otherwise-unreferenced task can
43+
# be garbage collected mid-run.
44+
self._scheduler_task: asyncio.Task[None] | None = None
45+
46+
async def acquire(self) -> None:
47+
"""Reserve the next send slot and wait until it arrives."""
48+
event = asyncio.Event()
49+
self._waiters.append(event)
50+
self._ensure_scheduling()
51+
await event.wait()
52+
53+
def _ensure_scheduling(self) -> None:
54+
"""Start the scheduler loop if it is not already running."""
55+
if not self._scheduling:
56+
self._scheduling = True
57+
self._scheduler_task = asyncio.ensure_future(self._schedule())
58+
59+
async def _schedule(self) -> None:
60+
"""Service waiters in FIFO order, one slot at a time."""
61+
while self._waiters:
62+
now = time.monotonic()
63+
# Cap drift so idle periods never accumulate more than *burst*
64+
# free slots.
65+
self._next_send = max(self._next_send, now - self._burst * self._interval)
66+
self._next_send += self._interval
67+
delay = self._next_send - now
68+
if delay > 0:
69+
await asyncio.sleep(delay)
70+
self._waiters.popleft().set()
71+
self._scheduling = False
72+
73+
74+
class RateLimitMiddleware:
75+
"""Client middleware that throttles requests with a token bucket.
76+
77+
The middleware delays each outgoing request until the bucket grants it a
78+
slot, so the client never sends faster than ``rate`` requests per second
79+
(allowing short bursts of up to ``burst`` requests).
80+
81+
:param float rate: Sustained request rate, in requests per second.
82+
:param int burst: Number of requests allowed to go out back-to-back before
83+
throttling kicks in.
84+
:param bool per_domain: When ``True``, keep an independent bucket per
85+
target host instead of a single global bucket.
86+
:param bool respect_retry_after: When ``True``, sleep for the duration of a
87+
numeric ``Retry-After`` header on an HTTP 429 response before returning
88+
it to the caller.
89+
"""
90+
91+
rate: float
92+
burst: int
93+
per_domain: bool
94+
respect_retry_after: bool
95+
96+
def __init__(
97+
self,
98+
rate: float = 10.0,
99+
burst: int = 10,
100+
per_domain: bool = False,
101+
respect_retry_after: bool = True,
102+
) -> None:
103+
self.rate = rate
104+
self.burst = burst
105+
self.per_domain = per_domain
106+
self.respect_retry_after = respect_retry_after
107+
self._global_bucket = TokenBucket(rate, burst)
108+
self._domain_buckets: dict[str, TokenBucket] = defaultdict(
109+
lambda: TokenBucket(rate, burst)
110+
)
111+
112+
def _get_bucket(self, request: ClientRequest) -> TokenBucket:
113+
if self.per_domain:
114+
domain = request.url.host or "unknown"
115+
return self._domain_buckets[domain]
116+
return self._global_bucket
117+
118+
async def _handle_retry_after(self, response: ClientResponse) -> None:
119+
if response.status != HTTPStatus.TOO_MANY_REQUESTS:
120+
return
121+
retry_after = response.headers.get("Retry-After")
122+
if retry_after:
123+
try:
124+
wait_seconds = float(retry_after)
125+
_LOGGER.info("Server requested Retry-After: %ss", wait_seconds)
126+
await asyncio.sleep(wait_seconds)
127+
except ValueError:
128+
_LOGGER.debug(
129+
"Retry-After is not a number (likely HTTP-date): %s", retry_after
130+
)
131+
132+
async def __call__(
133+
self,
134+
request: ClientRequest,
135+
handler: ClientHandlerType,
136+
) -> ClientResponse:
137+
"""Run the request through the rate limiter."""
138+
bucket = self._get_bucket(request)
139+
await bucket.acquire()
140+
141+
response = await handler(request)
142+
143+
if self.respect_retry_after:
144+
await self._handle_retry_after(response)
145+
146+
return response

‎docs/api.rst‎

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -75,3 +75,38 @@ Digest authentication
7575
# The middleware automatically handles the digest auth handshake.
7676
async with session.get("http://protected.example.com") as resp:
7777
assert resp.status == 200
78+
79+
80+
Rate limiting
81+
-------------
82+
83+
.. class:: RateLimitMiddleware(rate=10.0, burst=10, per_domain=False, respect_retry_after=True)
84+
85+
Client middleware that throttles outgoing requests with a token bucket.
86+
87+
:param float rate: Sustained request rate, in requests per second.
88+
:param int burst: Number of requests allowed to go out back-to-back before
89+
throttling kicks in.
90+
:param bool per_domain: Keep an independent bucket per target host instead
91+
of a single global bucket.
92+
:param bool respect_retry_after: Sleep for the duration of a numeric
93+
``Retry-After`` header on an HTTP 429 response before returning it to the
94+
caller.
95+
96+
The middleware delays each request until the bucket grants it a slot, so the
97+
client never sends faster than ``rate`` requests per second while still
98+
allowing short bursts of up to ``burst`` requests. Slots are served in strict
99+
FIFO order.
100+
101+
**Usage**
102+
103+
::
104+
105+
from aiohttp import ClientSession
106+
from aiohttp_client_middlewares import RateLimitMiddleware
107+
108+
# At most 5 requests/second, bursting up to 2.
109+
rate_limit = RateLimitMiddleware(rate=5.0, burst=2)
110+
async with ClientSession(middlewares=(rate_limit,)) as session:
111+
async with session.get("http://example.com") as resp:
112+
assert resp.status == 200

‎docs/code/rate_limit.py‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,19 @@
1+
"""Quickstart example for the rate-limiting middleware."""
2+
3+
import asyncio
4+
5+
from aiohttp import ClientSession
6+
7+
from aiohttp_client_middlewares import RateLimitMiddleware
8+
9+
10+
async def main() -> None:
11+
# Throttle to at most 5 requests per second, allowing bursts of up to 2.
12+
rate_limit = RateLimitMiddleware(rate=5.0, burst=2)
13+
async with ClientSession(middlewares=(rate_limit,)) as session:
14+
for _ in range(10):
15+
async with session.get("https://httpbin.org/get") as resp:
16+
print("Status:", resp.status)
17+
18+
19+
asyncio.run(main())

‎docs/index.rst‎

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,8 @@ This package collects ready-to-use middlewares for
88

99
- :class:`~aiohttp_client_middlewares.DigestAuthMiddleware` -- HTTP Digest
1010
authentication.
11+
- :class:`~aiohttp_client_middlewares.RateLimitMiddleware` -- client-side
12+
token-bucket rate limiting.
1113

1214

1315
Installation
@@ -21,11 +23,15 @@ Installation
2123
Quickstart
2224
----------
2325

24-
Attach a middleware to a session through the ``middlewares`` argument and
25-
let it handle authentication for every request:
26+
Attach one or more middlewares to a session through the ``middlewares``
27+
argument. For HTTP Digest authentication:
2628

2729
.. literalinclude:: code/digest_auth.py
2830

31+
For client-side rate limiting:
32+
33+
.. literalinclude:: code/rate_limit.py
34+
2935

3036
Contents
3137
--------

‎tests/test_rate_limit.py‎

Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
"""Tests for the rate-limiting middleware."""
2+
3+
import asyncio
4+
import time
5+
6+
from aiohttp import web
7+
from pytest_aiohttp import AiohttpClient
8+
9+
from aiohttp_client_middlewares.rate_limit import RateLimitMiddleware, TokenBucket
10+
11+
12+
async def _ok_handler(request: web.Request) -> web.Response:
13+
return web.Response(text="OK")
14+
15+
16+
def _make_app() -> web.Application:
17+
app = web.Application()
18+
app.router.add_get("/api", _ok_handler)
19+
return app
20+
21+
22+
async def test_token_bucket_allows_burst() -> None:
23+
"""Tokens up to burst size should be available immediately."""
24+
bucket = TokenBucket(rate=10.0, burst=3)
25+
start = time.monotonic()
26+
for _ in range(3):
27+
await bucket.acquire()
28+
elapsed = time.monotonic() - start
29+
# All three should be near-instant (within the burst allowance).
30+
assert elapsed < 0.05
31+
32+
33+
async def test_token_bucket_refills_after_idle() -> None:
34+
"""After draining, idle time should replenish burst slots."""
35+
bucket = TokenBucket(rate=100.0, burst=1)
36+
await bucket.acquire()
37+
await asyncio.sleep(0.05)
38+
start = time.monotonic()
39+
await bucket.acquire()
40+
elapsed = time.monotonic() - start
41+
# Should be near-instant because idle time refilled the slot.
42+
assert elapsed < 0.05
43+
44+
45+
async def test_token_bucket_fifo_ordering() -> None:
46+
"""Concurrent acquires should be served in FIFO order."""
47+
bucket = TokenBucket(rate=100.0, burst=1)
48+
order: list[int] = []
49+
50+
async def numbered_acquire(n: int) -> None:
51+
await bucket.acquire()
52+
order.append(n)
53+
54+
tasks = [asyncio.create_task(numbered_acquire(i)) for i in range(3)]
55+
await asyncio.gather(*tasks)
56+
assert order == [0, 1, 2]
57+
58+
59+
async def test_rate_limit_middleware_throttles(aiohttp_client: AiohttpClient) -> None:
60+
"""Global middleware should throttle requests beyond burst."""
61+
middleware = RateLimitMiddleware(rate=50.0, burst=2)
62+
client = await aiohttp_client(_make_app(), middlewares=(middleware,))
63+
64+
start = time.monotonic()
65+
for _ in range(4):
66+
resp = await client.get("/api")
67+
assert resp.status == 200
68+
elapsed = time.monotonic() - start
69+
70+
# 2 burst + 2 throttled at 50/s ~= 0.04s minimum wait. The upper bound
71+
# catches hangs or accidental double-sleeps while staying generous for CI.
72+
assert 0.02 <= elapsed < 0.5
73+
74+
75+
async def test_rate_limit_middleware_per_domain(aiohttp_client: AiohttpClient) -> None:
76+
"""Per-domain buckets should still throttle requests to the same host."""
77+
middleware = RateLimitMiddleware(rate=100.0, burst=1, per_domain=True)
78+
client = await aiohttp_client(_make_app(), middlewares=(middleware,))
79+
80+
start = time.monotonic()
81+
# Same host, so the two requests share a bucket and the second one waits.
82+
resp1 = await client.get("/api")
83+
resp2 = await client.get("/api")
84+
elapsed = time.monotonic() - start
85+
86+
assert resp1.status == 200
87+
assert resp2.status == 200
88+
assert 0.005 <= elapsed < 0.5
89+
90+
91+
async def test_rate_limit_middleware_respects_retry_after(
92+
aiohttp_client: AiohttpClient,
93+
) -> None:
94+
"""The middleware should sleep on a 429 with a numeric ``Retry-After``."""
95+
call_count = 0
96+
97+
async def rate_limited_handler(request: web.Request) -> web.Response:
98+
nonlocal call_count
99+
call_count += 1
100+
if call_count <= 1:
101+
return web.Response(status=429, headers={"Retry-After": "0.1"})
102+
return web.Response(text="OK")
103+
104+
app = web.Application()
105+
app.router.add_get("/api", rate_limited_handler)
106+
107+
middleware = RateLimitMiddleware(rate=100.0, burst=10, respect_retry_after=True)
108+
client = await aiohttp_client(app, middlewares=(middleware,))
109+
110+
start = time.monotonic()
111+
resp = await client.get("/api")
112+
elapsed = time.monotonic() - start
113+
114+
assert resp.status == 429
115+
assert 0.08 <= elapsed < 0.5

0 commit comments

Comments
 (0)