From f261ea506db4c9a7c864c9b04c6265594218608b Mon Sep 17 00:00:00 2001 From: Blake Bertuccelli-Booth <46652+bbertucc@users.noreply.github.com> Date: Mon, 18 May 2026 14:49:57 -0500 Subject: [PATCH] feat(auth): labelled API keys for per-key usage attribution MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit API-key requests were unattributable. The middleware stored request.state.api_key but the request logger only read request.state.identity (the session user), so every API-key call logged as anonymous. The key value is never logged (correct), but nothing derived from it was either — two keys were indistinguishable in every log line, so "how much is this key being used" had no answer. API_KEYS now accepts an optional label per entry: API_KEYS="zach:,daisy:" # labelled API_KEYS="," # bare — still works Each entry is either a bare key or a 'label:key' pair, split on the first colon — the same encoding AUTH_BASIC_USERS uses for 'username:hash'. A bare key gets a derived 'key-' label (first 8 hex of sha256) so even unlabelled keys stay distinct in logs without exposing the key. A key that itself contains a colon must be given an explicit label. Changes: - api_key_auth: _load_api_keys returns a {key: label} dict instead of a set; _is_valid_key becomes _match_key returning the matched label (or None). On a valid request, request.state.api_key_label is set. - logging_middleware: _identity_fields now also surfaces api_key_label, and accumulates fields from both auth paths instead of returning early — a request can carry a session identity, an API key label, or neither. - config: API_KEYS field description documents the label syntax. Backward compatible: an existing flat comma-separated API_KEYS with no colons keeps working unchanged; each key just gets a fingerprint label. Co-Authored-By: Claude Opus 4.7 (1M context) --- src/config.py | 13 ++- src/middleware/api_key_auth.py | 83 +++++++++++++------ src/middleware/logging_middleware.py | 36 +++++--- .../api/test_api_authentication.py | 18 ++-- tests/unit/middleware/test_api_key_auth.py | 77 ++++++++++++++++- .../middleware/test_logging_middleware.py | 58 +++++++++++++ 6 files changed, 238 insertions(+), 47 deletions(-) create mode 100644 tests/unit/middleware/test_logging_middleware.py diff --git a/src/config.py b/src/config.py index b047255e..508f9d04 100644 --- a/src/config.py +++ b/src/config.py @@ -65,7 +65,18 @@ class Settings(BaseSettings): # API Key Authentication Configuration enable_api_key_auth: bool = Field(default=True, description="Enable API key authentication for API endpoints") api_key_header_name: str = Field(default="X-API-Key", description="Header name for API key authentication") - api_keys: SecretStr | None = Field(default=None, description="Comma-separated list of valid API keys") + api_keys: SecretStr | None = Field( + default=None, + description=( + "Comma-separated list of valid API keys. Each entry is either a " + "bare key or a 'label:key' pair (split on the first colon, like " + "AUTH_BASIC_USERS). The label is recorded on every authenticated " + "request log so per-key usage is attributable; the key itself is " + "never logged. Bare keys get a derived 'key-' label. " + "Keys that themselves contain a colon must be given an explicit " + "label to avoid the leading segment being parsed as one." + ), + ) # Viewer Authentication (optional, layered on top of API keys) # AUTH_MODE=none keeps today's behaviour: no login, no cookies, no identity. diff --git a/src/middleware/api_key_auth.py b/src/middleware/api_key_auth.py index f685e301..068bf019 100644 --- a/src/middleware/api_key_auth.py +++ b/src/middleware/api_key_auth.py @@ -1,5 +1,6 @@ """API Key authentication middleware for FastAPI.""" +import hashlib import logging import secrets from collections.abc import Awaitable, Callable @@ -15,6 +16,15 @@ logger = logging.getLogger(__name__) +def _fingerprint(api_key: str) -> str: + """Stable, non-reversible identifier for an unlabelled key. + + Lets two unlabelled keys still be told apart in logs without ever + logging the key itself. + """ + return "key-" + hashlib.sha256(api_key.encode()).hexdigest()[:8] + + class APIKeyAuthMiddleware(BaseHTTPMiddleware): """ Middleware for API key authentication. @@ -31,29 +41,50 @@ def __init__(self, app: Any) -> None: app: FastAPI application instance """ super().__init__(app) - # Cache API keys at initialization to avoid reloading on every request - self._cached_keys: set[str] = self._load_api_keys() + # Cache API keys at initialization to avoid reloading on every request. + # Maps secret key -> human label (for usage attribution in logs). + self._cached_keys: dict[str, str] = self._load_api_keys() - def _load_api_keys(self) -> set[str]: + def _load_api_keys(self) -> dict[str, str]: """ - Load valid API keys from settings. + Load valid API keys from settings into a {key: label} map. + + ``API_KEYS`` is comma-separated. Each entry is either a bare key or a + ``label:key`` pair (split on the first colon, mirroring how + ``AUTH_BASIC_USERS`` encodes ``username:hash``). The label is logged + on every authenticated request so per-key usage is attributable; the + key itself is never logged. A bare key gets a derived + ``key-`` label so even unlabelled keys stay distinct. Returns: - Set of valid API key strings + Mapping of valid API key string -> label. """ if not settings.api_keys: logger.warning("No API keys configured! All authenticated requests will be rejected.") - return set() - - # Parse comma-separated keys from SecretStr - keys_str = settings.api_keys.get_secret_value() - keys = {key.strip() for key in keys_str.split(",") if key.strip()} + return {} + + keys: dict[str, str] = {} + for entry in settings.api_keys.get_secret_value().split(","): + entry = entry.strip() + if not entry: + continue + if ":" in entry: + label, _, key = entry.partition(":") + label = label.strip() + key = key.strip() + else: + key = entry + label = "" + if not key: + logger.warning("API key entry skipped: empty key after parsing") + continue + keys[key] = label or _fingerprint(key) if not keys: logger.warning("API keys configured but empty after parsing!") - return set() + return {} - logger.info(f"Loaded {len(keys)} API key(s) for authentication") + logger.info("Loaded %d API key(s) for authentication", len(keys)) return keys async def dispatch(self, request: Request, call_next: Callable[[Request], Awaitable[Response]]) -> Response: @@ -84,14 +115,18 @@ async def dispatch(self, request: Request, call_next: Callable[[Request], Awaita ) # Validate API key using constant-time comparison - if not self._is_valid_key(api_key): + label = self._match_key(api_key) + if label is None: logger.warning( f"Invalid API key for {request.method} {request.url.path} from {self._get_client_ip(request)}" ) return self._unauthorized_response(detail="Invalid API key") - # API key is valid, add to request state for potential use in handlers + # API key is valid. Stash the key for handlers and the label for + # request logging (LoggingMiddleware reads request.state.api_key_label + # so per-key usage is attributable without ever logging the key). request.state.api_key = api_key + request.state.api_key_label = label # Process request return await call_next(request) @@ -243,28 +278,24 @@ def _is_stream_token_request(self, request: Request) -> bool: return True - def _is_valid_key(self, provided_key: str) -> bool: + def _match_key(self, provided_key: str) -> str | None: """ - Validate API key using constant-time comparison. + Validate an API key and return its label. - Uses secrets.compare_digest() to prevent timing attacks. - Uses cached keys loaded at initialization for optimal performance. + Uses secrets.compare_digest() per key to prevent timing attacks on + the comparison itself. Uses cached keys loaded at initialization. Args: provided_key: API key from request header Returns: - True if key is valid + The matching key's label, or None if no key matches. """ - if not self._cached_keys: - return False - - # Use constant-time comparison to prevent timing attacks - for valid_key in self._cached_keys: + for valid_key, label in self._cached_keys.items(): if secrets.compare_digest(provided_key, valid_key): - return True + return label - return False + return None def _get_client_ip(self, request: Request) -> str: """ diff --git a/src/middleware/logging_middleware.py b/src/middleware/logging_middleware.py index 3e85606b..ac0dc9ff 100644 --- a/src/middleware/logging_middleware.py +++ b/src/middleware/logging_middleware.py @@ -11,24 +11,40 @@ def _identity_fields(request: Request) -> dict[str, str | None]: - """Pull identity attributes off ``request.state`` if SessionAuthMiddleware - populated them. Returns ``{}`` for anonymous requests so log lines from + """Pull caller-identity attributes off ``request.state`` for logging. + + Two independent auth paths can populate identity, and a request may use + either: + + * SessionAuthMiddleware sets ``request.state.identity`` (an ``Identity``) + for viewer logins — surfaced as ``user_sub`` / ``user_email`` / + ``user_provider``. + * APIKeyAuthMiddleware sets ``request.state.api_key_label`` for + programmatic clients — surfaced as ``api_key_label``. The key itself + is never logged. + + Returns ``{}`` for fully anonymous requests so log lines from AUTH_MODE=none deployments stay shape-identical to today's. - The ``isinstance`` rather than truthiness check matters: existing mock + The ``isinstance`` rather than truthiness checks matter: existing mock tests sometimes use a MagicMock for ``request.state`` whose attribute access always returns a truthy child mock. """ from ..auth.base import Identity # lazy to keep middleware import-light + fields: dict[str, str | None] = {} + identity = getattr(request.state, "identity", None) - if not isinstance(identity, Identity): - return {} - return { - "user_sub": identity.sub, - "user_email": identity.email, - "user_provider": identity.provider_id, - } + if isinstance(identity, Identity): + fields["user_sub"] = identity.sub + fields["user_email"] = identity.email + fields["user_provider"] = identity.provider_id + + api_key_label = getattr(request.state, "api_key_label", None) + if isinstance(api_key_label, str): + fields["api_key_label"] = api_key_label + + return fields class LoggingMiddleware(BaseHTTPMiddleware): diff --git a/tests/integration/api/test_api_authentication.py b/tests/integration/api/test_api_authentication.py index d9d8df60..438dca19 100644 --- a/tests/integration/api/test_api_authentication.py +++ b/tests/integration/api/test_api_authentication.py @@ -76,10 +76,10 @@ async def test_valid_api_key_allows_access(enable_api_key_auth): try: # Patch the middleware's validation method - def mock_is_valid_key(self, provided_key: str) -> bool: - return provided_key == "test-key-123" + def mock_match_key(self, provided_key: str) -> str | None: + return "test-label" if provided_key == "test-key-123" else None - with patch.object(APIKeyAuthMiddleware, '_is_valid_key', mock_is_valid_key): + with patch.object(APIKeyAuthMiddleware, '_match_key', mock_match_key): async with AsyncClient( transport=ASGITransport(app=app), base_url="http://test", @@ -150,10 +150,10 @@ async def test_multiple_api_keys_supported(enable_api_key_auth): # Directly patch the method that validates keys - def mock_is_valid_key(self, provided_key: str) -> bool: - return provided_key in {"key-1", "key-2", "key-3"} + def mock_match_key(self, provided_key: str) -> str | None: + return provided_key if provided_key in {"key-1", "key-2", "key-3"} else None - with patch.object(APIKeyAuthMiddleware, '_is_valid_key', mock_is_valid_key): + with patch.object(APIKeyAuthMiddleware, '_match_key', mock_match_key): mock_job_service = AsyncMock() mock_job_service.get_job.return_value = { "job_id": "test", @@ -224,10 +224,10 @@ async def test_api_key_auth_unaffected_by_public_docs(): from src.middleware.api_key_auth import APIKeyAuthMiddleware with patch("src.dependencies.get_job_service") as mock_job_service_dep: - def mock_is_valid_key(self, provided_key: str) -> bool: - return provided_key == "api-key-123" + def mock_match_key(self, provided_key: str) -> str | None: + return "test-label" if provided_key == "api-key-123" else None - with patch.object(APIKeyAuthMiddleware, '_is_valid_key', mock_is_valid_key): + with patch.object(APIKeyAuthMiddleware, '_match_key', mock_match_key): mock_job_service = AsyncMock() mock_job_service.get_job.return_value = { "job_id": "test-id", diff --git a/tests/unit/middleware/test_api_key_auth.py b/tests/unit/middleware/test_api_key_auth.py index 4ae24cff..28cb491d 100644 --- a/tests/unit/middleware/test_api_key_auth.py +++ b/tests/unit/middleware/test_api_key_auth.py @@ -5,7 +5,7 @@ import pytest from fastapi import Request, Response from fastapi.responses import JSONResponse -from src.middleware.api_key_auth import APIKeyAuthMiddleware +from src.middleware.api_key_auth import APIKeyAuthMiddleware, _fingerprint @pytest.fixture @@ -433,3 +433,78 @@ def tracked_load(self): assert call_count == initial_call_count assert call_count == 1 assert call_next.call_count == 5 + + +# ---------- labelled API keys (usage attribution) ---------- + + +def _middleware_with(raw_keys: str) -> APIKeyAuthMiddleware: + """Build a middleware instance whose API_KEYS is `raw_keys`.""" + mock_app = MagicMock() + with patch("src.middleware.api_key_auth.settings") as mock_settings: + mock_settings.api_key_header_name = "X-API-Key" + mock_settings.environment = "production" + mock_settings.api_keys = MagicMock() + mock_settings.api_keys.get_secret_value.return_value = raw_keys + return APIKeyAuthMiddleware(mock_app) + + +@pytest.mark.unit +def test_fingerprint_is_stable_and_does_not_leak_key(): + fp = _fingerprint("super-secret-key") + assert fp == _fingerprint("super-secret-key") # deterministic + assert fp.startswith("key-") + assert "super-secret-key" not in fp + + +@pytest.mark.unit +def test_labelled_keys_parsed_into_key_to_label_map(): + mw = _middleware_with("zach:key-zzz,daisy:key-ddd") + assert mw._cached_keys == {"key-zzz": "zach", "key-ddd": "daisy"} + + +@pytest.mark.unit +def test_bare_key_gets_fingerprint_label(): + mw = _middleware_with("plainkey") + assert mw._cached_keys["plainkey"] == _fingerprint("plainkey") + + +@pytest.mark.unit +def test_mixed_labelled_and_bare_keys(): + mw = _middleware_with("zach:key-a,key-b") + assert mw._cached_keys["key-a"] == "zach" + assert mw._cached_keys["key-b"] == _fingerprint("key-b") + + +@pytest.mark.unit +def test_label_split_on_first_colon_only(): + # A key that itself contains colons works when explicitly labelled: + # everything after the first colon is the key. + mw = _middleware_with("svc:weird:key:with:colons") + assert mw._cached_keys == {"weird:key:with:colons": "svc"} + + +@pytest.mark.unit +@pytest.mark.asyncio +async def test_valid_labelled_key_sets_label_on_request_state(): + mw = _middleware_with("zach:key-zzz,daisy:key-ddd") + request = create_mock_request("/api/v1/documents/submit", {"X-API-Key": "key-ddd"}) + call_next = AsyncMock(return_value=Response(status_code=200)) + + response = await mw.dispatch(request, call_next) + + assert response.status_code == 200 + assert request.state.api_key == "key-ddd" + assert request.state.api_key_label == "daisy" + + +@pytest.mark.unit +@pytest.mark.asyncio +async def test_bare_key_sets_fingerprint_label_on_request_state(): + mw = _middleware_with("plainkey") + request = create_mock_request("/api/v1/documents/submit", {"X-API-Key": "plainkey"}) + call_next = AsyncMock(return_value=Response(status_code=200)) + + await mw.dispatch(request, call_next) + + assert request.state.api_key_label == _fingerprint("plainkey") diff --git a/tests/unit/middleware/test_logging_middleware.py b/tests/unit/middleware/test_logging_middleware.py new file mode 100644 index 00000000..cd6bed1d --- /dev/null +++ b/tests/unit/middleware/test_logging_middleware.py @@ -0,0 +1,58 @@ +"""Unit tests for request logging middleware identity extraction.""" + +from datetime import UTC, datetime +from types import SimpleNamespace + +import pytest +from src.auth.base import Identity +from src.middleware.logging_middleware import _identity_fields + +pytestmark = pytest.mark.unit + + +def _request(**state_attrs): + """Minimal stand-in for a Request: only request.state is accessed.""" + return SimpleNamespace(state=SimpleNamespace(**state_attrs)) + + +def _identity() -> Identity: + now = datetime.now(UTC) + return Identity( + sub="zach", + email="zach@example.com", + name="Zach", + provider_id="basic", + issued_at=now, + expires_at=now, + ) + + +def test_anonymous_request_yields_no_fields(): + assert _identity_fields(_request()) == {} + + +def test_session_identity_surfaces_user_fields(): + fields = _identity_fields(_request(identity=_identity())) + assert fields == { + "user_sub": "zach", + "user_email": "zach@example.com", + "user_provider": "basic", + } + + +def test_api_key_label_surfaces_without_session_identity(): + fields = _identity_fields(_request(api_key_label="daisy")) + assert fields == {"api_key_label": "daisy"} + + +def test_session_and_api_key_both_surface(): + fields = _identity_fields(_request(identity=_identity(), api_key_label="svc")) + assert fields["user_sub"] == "zach" + assert fields["api_key_label"] == "svc" + + +def test_non_identity_value_is_ignored(): + # A truthy non-Identity (e.g. a MagicMock in upstream tests) must not + # be mistaken for a real identity. + fields = _identity_fields(_request(identity="not-an-identity", api_key_label=123)) + assert fields == {}