From 8b58e49dbad035e7179f1788e4c163c302036e6f Mon Sep 17 00:00:00 2001 From: dragoncage Date: Fri, 26 Jun 2026 23:25:55 -0700 Subject: [PATCH] fix: throttle invalid api key attempts --- server.py | 51 +++++++++++++++++++++++++++++----------- tests/test_api_auth.py | 53 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 90 insertions(+), 14 deletions(-) diff --git a/server.py b/server.py index 14a045b..f6a2b9a 100644 --- a/server.py +++ b/server.py @@ -61,8 +61,16 @@ 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 '."}, status=401, @@ -70,25 +78,40 @@ async def middleware(self, request: web.Request, handler) -> web.Response: 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: diff --git a/tests/test_api_auth.py b/tests/test_api_auth.py index c09a9a6..a4bcd60 100644 --- a/tests/test_api_auth.py +++ b/tests/test_api_auth.py @@ -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):