Skip to content
Merged
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
13 changes: 12 additions & 1 deletion src/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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-<fingerprint>' 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.
Expand Down
83 changes: 57 additions & 26 deletions src/middleware/api_key_auth.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""API Key authentication middleware for FastAPI."""

import hashlib
import logging
import secrets
from collections.abc import Awaitable, Callable
Expand All @@ -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.
Expand All @@ -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-<fingerprint>`` 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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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:
"""
Expand Down
36 changes: 26 additions & 10 deletions src/middleware/logging_middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
18 changes: 9 additions & 9 deletions tests/integration/api/test_api_authentication.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
77 changes: 76 additions & 1 deletion tests/unit/middleware/test_api_key_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Loading
Loading