Skip to content
Draft
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
51 changes: 37 additions & 14 deletions server.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,34 +61,57 @@ async def middleware(self, request: web.Request, handler) -> web.Response:
if not self.enabled:
return await handler(request)

now = time.monotonic()
self._prune_request_log(now)

auth_failure_bucket = self._auth_failure_bucket(request)
if self._bucket_exhausted(auth_failure_bucket):
return self._rate_limit_response()

auth_header = request.headers.get("Authorization", "")
if not auth_header.startswith("Bearer "):
self._record_rate_attempt(auth_failure_bucket, now)
return web.json_response(
{"error": "Missing or invalid Authorization header. Use 'Bearer <api_key>'."},
status=401,
)

key = auth_header[7:]
if not self._is_valid_key(key):
self._record_rate_attempt(auth_failure_bucket, now)
return web.json_response({"error": "Invalid API key."}, status=403)

# Rate limiting
now = time.monotonic()
log = self._request_log[key]
# Prune old entries
cutoff = now - self._rate_window
self._request_log[key] = [t for t in log if t > cutoff]
log = self._request_log[key]

if len(log) >= self._rate_limit:
return web.json_response(
{"error": "Rate limit exceeded.", "retry_after": self._rate_window},
status=429,
)
log.append(now)
if self._bucket_exhausted(key):
return self._rate_limit_response()
self._record_rate_attempt(key, now)

return await handler(request)

def _prune_request_log(self, now: float) -> None:
cutoff = now - self._rate_window
for bucket, log in list(self._request_log.items()):
current_log = [t for t in log if t > cutoff]
if current_log:
self._request_log[bucket] = current_log
else:
del self._request_log[bucket]

def _bucket_exhausted(self, bucket: str) -> bool:
return len(self._request_log.get(bucket, [])) >= self._rate_limit

def _record_rate_attempt(self, bucket: str, now: float) -> None:
self._request_log[bucket].append(now)

def _auth_failure_bucket(self, request: web.Request) -> str:
peer = request.remote or "unknown"
return f"auth-failure:{peer}"

def _rate_limit_response(self) -> web.Response:
return web.json_response(
{"error": "Rate limit exceeded.", "retry_after": self._rate_window},
status=429,
)

def _is_valid_key(self, candidate_key: str) -> bool:
valid = False
for stored_key in self._keys:
Expand Down
53 changes: 53 additions & 0 deletions tests/test_api_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,6 +154,59 @@ async def test_rate_limit_exceeded(self, aiohttp_client):
data = await resp.json()
assert "Rate limit" in data["error"]

async def test_invalid_bearer_attempts_are_rate_limited_without_storing_raw_tokens(self, aiohttp_client):
auth = APIKeyAuth(api_keys=["key"], rate_limit=2, rate_window=60)
client = await aiohttp_client(_make_app(auth))

for token in ("wrong-0", "wrong-1"):
resp = await client.get("/test", headers={"Authorization": f"Bearer {token}"})
assert resp.status == 403

resp = await client.get("/test", headers={"Authorization": "Bearer wrong-2"})
assert resp.status == 429
for token in ("wrong-0", "wrong-1", "wrong-2"):
assert token not in auth._request_log
assert all(token not in bucket for bucket in auth._request_log)

async def test_invalid_bearer_rate_limit_blocks_validation_after_threshold(self, aiohttp_client):
auth = APIKeyAuth(api_keys=["key"], rate_limit=2, rate_window=60)
client = await aiohttp_client(_make_app(auth))

for token in ("wrong-0", "wrong-1"):
resp = await client.get("/test", headers={"Authorization": f"Bearer {token}"})
assert resp.status == 403

resp = await client.get("/test", headers={"Authorization": "Bearer key"})
assert resp.status == 429

async def test_missing_and_malformed_auth_attempts_are_rate_limited_without_storing_raw_header(self, aiohttp_client):
auth = APIKeyAuth(api_keys=["key"], rate_limit=2, rate_window=60)
client = await aiohttp_client(_make_app(auth))

resp = await client.get("/test")
assert resp.status == 401

resp = await client.get("/test", headers={"Authorization": "Basic wrong-0"})
assert resp.status == 401

resp = await client.get("/test")
assert resp.status == 429
for private_value in ("Basic wrong-0", "wrong-0"):
assert private_value not in auth._request_log
assert all(private_value not in bucket for bucket in auth._request_log)
assert "" not in auth._request_log

async def test_stale_auth_failure_buckets_are_pruned_on_later_requests(self, aiohttp_client, monkeypatch):
auth = APIKeyAuth(api_keys=["key"], rate_limit=2, rate_window=60)
auth._request_log["auth-failure:old-peer"] = [10.0]
monkeypatch.setattr(server.time, "monotonic", lambda: 100.0)
client = await aiohttp_client(_make_app(auth))

resp = await client.get("/test", headers={"Authorization": "Bearer key"})

assert resp.status == 200
assert "auth-failure:old-peer" not in auth._request_log


class TestServerPrivacyLogging:
async def test_webhook_failure_log_redacts_url_and_exception_details(self, monkeypatch, caplog):
Expand Down