Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 26 additions & 12 deletions freerelay/middleware/rate_limit.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

from __future__ import annotations

import asyncio
import logging
import time
from collections import OrderedDict
Expand Down Expand Up @@ -44,10 +45,15 @@ def __init__(self, rate: float, capacity: int) -> None:
rate: Tokens added per second.
capacity: Maximum tokens (burst capacity).
"""
if rate <= 0:
raise ValueError("rate must be greater than zero")
if capacity <= 0:
raise ValueError("capacity must be greater than zero")

self.rate = rate
self.capacity = capacity
self.tokens = float(capacity)
self.last_refill = time.time()
self.last_refill = time.monotonic()

def consume(self) -> bool:
"""
Expand All @@ -57,7 +63,7 @@ def consume(self) -> bool:
True if a token was available (request allowed).
False if no tokens (request should be rate limited).
"""
now = time.time()
now = time.monotonic()
elapsed = now - self.last_refill
self.tokens = min(self.capacity, self.tokens + elapsed * self.rate)
self.last_refill = now
Expand All @@ -83,12 +89,18 @@ def __init__(
requests_per_minute: int = 60,
burst_capacity: int = 10,
) -> None:
if requests_per_minute < 1:
raise ValueError("requests_per_minute must be at least 1")
if burst_capacity < 1:
raise ValueError("burst_capacity must be at least 1")

super().__init__(app) # type: ignore[arg-type]
self.requests_per_minute = requests_per_minute
self.burst_capacity = burst_capacity
self._buckets: OrderedDict[str, TokenBucket] = OrderedDict()
self._rate = requests_per_minute / 60.0 # tokens per second
self._request_count = 0
self._state_lock = asyncio.Lock()

def _get_bucket(self, client_ip: str) -> TokenBucket:
"""Get or create a token bucket for a client with LRU eviction."""
Expand All @@ -108,7 +120,7 @@ def _get_bucket(self, client_ip: str) -> TokenBucket:

def _cleanup_stale_buckets(self) -> None:
"""Remove buckets that haven't been used in over 5 minutes."""
now = time.time()
now = time.monotonic()
stale_threshold = 300.0 # 5 minutes
stale_keys = [
key
Expand All @@ -132,17 +144,19 @@ async def dispatch(
if any(path.startswith(p) for p in _SKIP_PATHS):
return await call_next(request)

# Periodic cleanup of stale buckets
self._request_count += 1
if self._request_count % _CLEANUP_INTERVAL == 0:
self._cleanup_stale_buckets()
async with self._state_lock:
# Periodic cleanup of stale buckets
self._request_count += 1
if self._request_count % _CLEANUP_INTERVAL == 0:
self._cleanup_stale_buckets()

# Identify by user_id if available (from AuthMiddleware), otherwise IP
user_id = getattr(request.state, "user_id", None)
client_id = user_id or (request.client.host if request.client else "unknown")
bucket = self._get_bucket(client_id)
# Identify by user_id if available (from AuthMiddleware), otherwise IP
user_id = getattr(request.state, "user_id", None)
client_id = user_id or (request.client.host if request.client else "unknown")
bucket = self._get_bucket(client_id)
allowed = bucket.consume()

if not bucket.consume():
if not allowed:
logger.warning("Rate limit exceeded for %s", client_id)
return JSONResponse(
status_code=429,
Expand Down
43 changes: 43 additions & 0 deletions tests/unit/test_rate_limit.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,43 @@
from unittest.mock import patch

import pytest

from freerelay.middleware.rate_limit import TokenBucket


@pytest.mark.parametrize("rate, capacity", [(0, 1), (-1, 1), (1, 0), (1, -1)])
def test_token_bucket_rejects_non_positive_settings(rate, capacity):
with pytest.raises(ValueError):
TokenBucket(rate, capacity)


def test_token_bucket_uses_monotonic_clock_for_refills():
with patch("freerelay.middleware.rate_limit.time.monotonic", side_effect=[10.0, 10.0, 10.0, 11.0]):
bucket = TokenBucket(rate=1, capacity=1)
assert bucket.consume()
assert not bucket.consume()

bucket.tokens = 0
assert bucket.consume()


@pytest.mark.asyncio
async def test_concurrent_requests_share_bucket_atomically():
from unittest.mock import AsyncMock, Mock

from freerelay.middleware.rate_limit import RateLimitMiddleware

middleware = RateLimitMiddleware(Mock(), requests_per_minute=1, burst_capacity=1)
request = Mock()
request.url.path = "/v1/chat/completions"
request.state.user_id = "tenant"
request.client.host = "127.0.0.1"
call_next = AsyncMock()

responses = await __import__("asyncio").gather(
middleware.dispatch(request, call_next),
middleware.dispatch(request, call_next),
)

assert sum(response.status_code == 429 for response in responses) == 1
assert call_next.await_count == 1
Loading