diff --git a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py index 297c68c71..0989460ff 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -2,23 +2,43 @@ from __future__ import annotations +import json +import re import threading from dataclasses import dataclass from datetime import timedelta -from typing import Any, Literal +from typing import Literal, NoReturn, cast import jwt -from jwt import PyJWKClient +from jwt import PyJWKClient, PyJWKClientConnectionError, PyJWKClientError, PyJWKSetError +from jwt.algorithms import AllowedPublicKeys +from jwt.types import Options from loguru import logger from shared.core.config import settings from shared.core.exceptions.domain_exceptions import AuthException +from shared.core.logging import LogEvent JWKS_ENDPOINT_PATH = "/api/auth/jwks" JWKS_CACHE_TTL_SECONDS = 60 * 60 +JWT_KEY_ID_MAX_LENGTH = 64 +JWT_KEY_ID_UNSAFE_PATTERN = re.compile(r"[^A-Za-z0-9._:-]") +JWT_ALGORITHMS: tuple[str, ...] = ("HS256", "RS256", "EdDSA") +JWT_STRUCTURE_ONLY_DECODE_OPTIONS: Options = { + "verify_signature": False, +} READ_ONLY_PERMISSION: Literal["read_only"] = "read_only" FULL_ACCESS_PERMISSION: Literal["full_access"] = "full_access" Permission = Literal["read_only", "full_access"] +JWTFailureReason = Literal[ + "jwt_missing_key_id", + "jwt_unknown_key_id", + "jwt_expired", + "jwt_invalid", + "jwks_unavailable", + "jwks_invalid", +] +VerificationKey = AllowedPublicKeys | str | bytes @dataclass(frozen=True) @@ -27,6 +47,30 @@ class DashboardJWTIdentity: permission: Permission +@dataclass(frozen=True) +class _DashboardJWTHeader: + algorithm: object + key_id: str + + +@dataclass(frozen=True) +class _DashboardJWTTelemetry: + jwt_kid_present: bool + jwt_algorithm: str | None = None + jwt_kid: str | None = None + + def to_log_data(self) -> dict[str, object]: + log_data: dict[str, object] = { + "auth_component": "dashboard_jwt", + "jwt_kid_present": self.jwt_kid_present, + } + if self.jwt_algorithm is not None: + log_data["jwt_algorithm"] = self.jwt_algorithm + if self.jwt_kid is not None: + log_data["jwt_kid"] = self.jwt_kid + return log_data + + class DashboardJWTAuthenticationService: """Validate Dashboard-issued JWTs through the configured JWKS endpoint.""" @@ -40,47 +84,78 @@ def decode_user_id(self, token: str) -> str: def decode_identity(self, token: str) -> DashboardJWTIdentity: """Decode and validate a JWT, returning the user ID and permission.""" - try: - payload = self._decode_payload(token) - user_id = payload.get("id") - if not isinstance(user_id, str) or not user_id: - raise AuthException(user_message="Token missing 'id' claim") - - permission = _normalize_permission(payload.get("permission")) - return DashboardJWTIdentity(user_id=user_id, permission=permission) - except jwt.ExpiredSignatureError: - raise AuthException(user_message="Token has expired") - except jwt.InvalidTokenError as exc: - logger.warning(f"Invalid JWT token: {exc}") - raise AuthException(user_message="Invalid token") - - def _decode_payload(self, token: str) -> dict[str, Any]: - key = self._get_verification_key(token) - payload: dict[str, Any] = jwt.decode( - token, - key, - algorithms=["HS256", "RS256", "EdDSA"], - leeway=timedelta(seconds=30), - options={"verify_aud": False}, + header = _parse_header_or_reject(token) + telemetry_context = _build_telemetry_context(header) + _assert_token_structure_or_reject(token, telemetry_context) + key = self._resolve_verification_key_or_reject( + key_id=header.key_id, + telemetry_context=telemetry_context, ) - return payload + payload = _verify_payload_or_reject( + token=token, + key=key, + telemetry_context=telemetry_context, + ) + return _build_identity_or_reject(payload, telemetry_context) - def _get_verification_key(self, token: str) -> Any: - """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" + def _resolve_verification_key_or_reject( + self, + *, + key_id: str, + telemetry_context: _DashboardJWTTelemetry, + ) -> VerificationKey: + """Resolve a JWT verification key and classify JWKS failures locally.""" try: - jwks_client = self._get_jwks_client() - signing_key = jwks_client.get_signing_key_from_jwt(token) - return signing_key.key - except jwt.PyJWKClientError as exc: - logger.error(f"Failed to fetch JWKS: {exc}") - raise AuthException( - internal_message=( - f"Failed to fetch verification key from JWKS endpoint: {exc}" - ) + key = self._get_verification_key(key_id) + except PyJWKClientConnectionError as error: + _reject_jwks_dependency( + failure_reason="jwks_unavailable", + telemetry_context=telemetry_context, + original_exception=error, + ) + except ( + json.JSONDecodeError, + UnicodeDecodeError, + PyJWKSetError, + jwt.PyJWKError, + ) as error: + _reject_jwks_dependency( + failure_reason="jwks_invalid", + telemetry_context=telemetry_context, + original_exception=error, + ) + except PyJWKClientError as error: + _reject_jwks_dependency( + failure_reason="jwks_invalid", + telemetry_context=telemetry_context, + original_exception=error, + ) + + if key is None: + _reject_client_jwt( + failure_reason="jwt_unknown_key_id", + telemetry_context=telemetry_context, ) - except jwt.PyJWKSetError as exc: - logger.error(f"Invalid JWKS format: {exc}") - raise AuthException(internal_message=f"Invalid JWKS format: {exc}") + + return key + + def _get_verification_key(self, key_id: str) -> VerificationKey | None: + """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" + jwks_client = self._get_jwks_client() + signing_keys = jwks_client.get_signing_keys() + signing_key = jwks_client.match_kid(signing_keys, key_id) + if signing_key is not None: + return cast(VerificationKey, signing_key.key) + + refreshed_signing_keys = jwks_client.get_signing_keys(refresh=True) + refreshed_signing_key = jwks_client.match_kid( + refreshed_signing_keys, + key_id, + ) + if refreshed_signing_key is None: + return None + + return cast(VerificationKey, refreshed_signing_key.key) def _get_jwks_client(self) -> PyJWKClient: """Return a cached JWKS client for Dashboard token verification.""" @@ -96,11 +171,188 @@ def _get_jwks_client(self) -> PyJWKClient: lifespan=JWKS_CACHE_TTL_SECONDS, timeout=30, ) - logger.info(f"Initialized JWKS client with endpoint: {jwks_url}") return self._jwks_client +def _parse_header_or_reject(token: str) -> _DashboardJWTHeader: + try: + unverified_header = cast(dict[str, object], jwt.get_unverified_header(token)) + except jwt.InvalidTokenError: + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=_build_telemetry_context_from_values( + algorithm=None, + key_id=None, + ), + ) + + algorithm = unverified_header.get("alg") + key_id_value = unverified_header.get("kid") + key_id = key_id_value if isinstance(key_id_value, str) else None + if key_id is None or not key_id.strip(): + _reject_client_jwt( + failure_reason="jwt_missing_key_id", + telemetry_context=_build_telemetry_context_from_values( + algorithm=algorithm, + key_id=key_id, + ), + ) + + return _DashboardJWTHeader(algorithm=algorithm, key_id=key_id) + + +def _build_telemetry_context( + header: _DashboardJWTHeader, +) -> _DashboardJWTTelemetry: + return _build_telemetry_context_from_values( + algorithm=header.algorithm, + key_id=header.key_id, + ) + + +def _build_telemetry_context_from_values( + *, + algorithm: object, + key_id: str | None, +) -> _DashboardJWTTelemetry: + is_key_id_present = key_id is not None and bool(key_id.strip()) + jwt_algorithm: str | None = None + if isinstance(algorithm, str) and algorithm in JWT_ALGORITHMS: + jwt_algorithm = algorithm + jwt_kid: str | None = None + if is_key_id_present and key_id is not None: + jwt_kid = _sanitize_key_id(key_id) + return _DashboardJWTTelemetry( + jwt_kid_present=is_key_id_present, + jwt_algorithm=jwt_algorithm, + jwt_kid=jwt_kid, + ) + + +def _assert_token_structure_or_reject( + token: str, + telemetry_context: _DashboardJWTTelemetry, +) -> None: + """Reject structurally invalid JWTs before touching Dashboard JWKS.""" + try: + # This decode only checks token structure; verified claims come from + # _verify_payload_or_reject after the signing key is resolved. + jwt.decode( + token, + options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, + ) + except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=telemetry_context, + ) + + +def _verify_payload_or_reject( + *, + token: str, + key: VerificationKey, + telemetry_context: _DashboardJWTTelemetry, +) -> dict[str, object]: + try: + payload = cast( + dict[str, object], + jwt.decode( + token, + key, + algorithms=list(JWT_ALGORITHMS), + leeway=timedelta(seconds=30), + options={"verify_aud": False}, + ), + ) + except jwt.ExpiredSignatureError: + _reject_client_jwt( + failure_reason="jwt_expired", + telemetry_context=telemetry_context, + ) + except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=telemetry_context, + ) + + return payload + + +def _build_identity_or_reject( + payload: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, +) -> DashboardJWTIdentity: + user_id = payload.get("id") + if not isinstance(user_id, str) or not user_id: + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=telemetry_context, + ) + + permission = _normalize_permission(payload.get("permission")) + return DashboardJWTIdentity(user_id=user_id, permission=permission) + + +def _sanitize_key_id(key_id: str) -> str: + sanitized_key_id = JWT_KEY_ID_UNSAFE_PATTERN.sub("_", key_id) + return sanitized_key_id[:JWT_KEY_ID_MAX_LENGTH] + + +def _reject_client_jwt( + *, + failure_reason: JWTFailureReason, + telemetry_context: _DashboardJWTTelemetry, +) -> NoReturn: + _log_dashboard_jwt_auth_failure( + failure_reason=failure_reason, + telemetry_context=telemetry_context, + is_jwks_dependency_failure=False, + ) + raise AuthException() from None + + +def _reject_jwks_dependency( + *, + failure_reason: JWTFailureReason, + telemetry_context: _DashboardJWTTelemetry, + original_exception: Exception, +) -> NoReturn: + _log_dashboard_jwt_auth_failure( + failure_reason=failure_reason, + telemetry_context=telemetry_context, + is_jwks_dependency_failure=True, + original_exception=original_exception, + ) + raise AuthException() from None + + +def _log_dashboard_jwt_auth_failure( + *, + failure_reason: JWTFailureReason, + telemetry_context: _DashboardJWTTelemetry, + is_jwks_dependency_failure: bool, + original_exception: Exception | None = None, +) -> None: + log_data: dict[str, object] = { + **telemetry_context.to_log_data(), + "failure_reason": failure_reason, + } + message = f"Dashboard JWT authentication failed: {failure_reason}" + if is_jwks_dependency_failure: + logger.bind( + event=LogEvent.EXCEPTION_SYSTEM.value, + **log_data, + ).opt(exception=original_exception).error(message) + return + + logger.bind( + event=LogEvent.EXCEPTION_CLIENT.value, + **log_data, + ).warning(message) + + def _normalize_permission(value: object) -> Permission: if value == READ_ONLY_PERMISSION: return READ_ONLY_PERMISSION diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py new file mode 100644 index 000000000..5d38d3a5a --- /dev/null +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -0,0 +1,652 @@ +from __future__ import annotations + +import base64 +import json +import sys +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Protocol, cast + +import jwt +import pytest +from fastapi import FastAPI, Header +from httpx import ASGITransport, AsyncClient, Response +from loguru import logger +from pytest import MonkeyPatch + +from tests.support.import_environment import ( + configure_import_environment, + ensure_import_paths, +) +from tests.support.dashboard_jwt import ( + create_dashboard_rsa_jwk as _create_rsa_jwk, + create_dashboard_rsa_private_key as _create_rsa_private_key, + create_dashboard_rsa_token as _create_rsa_token, + serve_dashboard_jwks as _serve_jwks, +) + +configure_import_environment() +ensure_import_paths() + + +class _LoguruMessage(Protocol): + @property + def record(self) -> Mapping[str, object]: + raise NotImplementedError + + +@dataclass(frozen=True) +class _CapturedAuthLog: + level: str + event: str + message: str + extra: Mapping[str, object] + exception_type: str | None + exception_message: str | None + + +@dataclass +class _FakeLogfireExceptionHelper: + exception: BaseException + level: str = "error" + is_recording_exception: bool = True + + def no_record_exception(self) -> None: + self.is_recording_exception = False + + +class _AuthLogCapture: + def __init__(self) -> None: + self.records: list[_CapturedAuthLog] = [] + + def capture(self, message: _LoguruMessage) -> None: + record = message.record + extra = cast(Mapping[str, object], record["extra"]) + if extra.get("auth_component") != "dashboard_jwt": + return + + level: object = record["level"] + exception: object | None = record.get("exception") + exception_type: str | None = None + exception_message: str | None = None + if exception is not None: + exception_value: object | None = getattr(exception, "value", None) + exception_type_value: object | None = getattr(exception, "type", None) + if exception_type_value is not None: + exception_type = str(getattr(exception_type_value, "__name__", "")) + if exception_value is not None: + exception_message = str(exception_value) + self.records.append( + _CapturedAuthLog( + level=str(getattr(level, "name", level)), + event=str(extra.get("event", "")), + message=str(record["message"]), + extra=dict(extra), + exception_type=exception_type, + exception_message=exception_message, + ) + ) + + +@contextmanager +def _capture_auth_logs() -> Iterator[_AuthLogCapture]: + log_capture = _AuthLogCapture() + log_sink_id = logger.add(log_capture.capture) + try: + yield log_capture + finally: + logger.remove(log_sink_id) + + +def _prepare_api_app_imports() -> None: + api_root: str = str(Path(__file__).resolve().parents[2]) + if api_root in sys.path: + sys.path.remove(api_root) + sys.path.insert(0, api_root) + + cached_module_names: list[str] = list(sys.modules) + for module_name in cached_module_names: + if module_name == "app" or module_name.startswith("app."): + sys.modules.pop(module_name, None) + + +def _use_dashboard_endpoint( + monkeypatch: MonkeyPatch, + endpoint: str, +) -> None: + from shared.core.config import settings + + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + endpoint, + ) + + +def _create_auth_exception() -> BaseException: + from shared.core.exceptions.domain_exceptions import AuthException + + return AuthException() + + +def _downgrade_logfire_exception(helper: _FakeLogfireExceptionHelper) -> None: + from shared.core.logging import _downgrade_expected_logfire_exception + + _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] + + +def _create_authentication_app() -> FastAPI: + _prepare_api_app_imports() + + from app.core.exception_handlers import setup_exception_handlers + from app.services.auth.dashboard_jwt_authentication_service import ( + DashboardJWTAuthenticationService, + ) + + authentication_service = DashboardJWTAuthenticationService() + app = FastAPI() + + @app.get("/protected") + async def read_protected_resource( + authorization: str = Header(), + ) -> dict[str, str]: + _, _, token = authorization.partition(" ") + identity = authentication_service.decode_identity(token) + return { + "user_id": identity.user_id, + "permission": identity.permission, + } + + setup_exception_handlers(app) + return app + + +async def _request_with_token(token: str) -> Response: + app = _create_authentication_app() + transport = ASGITransport(app=app) + async with AsyncClient(transport=transport, base_url="http://test") as client: + return await client.get( + "/protected", + headers={"Authorization": f"Bearer {token}"}, + ) + + +def _assert_unauthenticated_response( + response: Response, + *, + expected_message: str, + token: str, +) -> None: + assert response.status_code == 401 + response_json = cast(dict[str, object], response.json()) + error = cast(dict[str, object], response_json["error"]) + assert error["code"] == "UNAUTHENTICATED" + assert error["message"] == expected_message + + serialized_response = json.dumps(response_json, default=str) + assert "failure_reason" not in serialized_response + assert "auth_component" not in serialized_response + assert "dashboard_jwt" not in serialized_response + assert "jwt_algorithm" not in serialized_response + assert "jwt_kid" not in serialized_response + assert "payload" not in serialized_response + assert "contract-dashboard-user" not in serialized_response + assert token not in serialized_response + assert f"Bearer {token}" not in serialized_response + token_segments = token.split(".") + if len(token_segments) > 1: + assert token_segments[1] not in serialized_response + + +def _assert_log_excludes_token( + auth_log: _CapturedAuthLog, + *, + token: str, +) -> None: + serialized_log = _serialize_auth_log(auth_log) + assert token not in serialized_log + assert f"Bearer {token}" not in serialized_log + token_segments = token.split(".") + if len(token_segments) > 1: + assert token_segments[1] not in serialized_log + assert "contract-dashboard-user" not in serialized_log + + +def _serialize_auth_log(auth_log: _CapturedAuthLog) -> str: + return json.dumps( + { + "message": auth_log.message, + "extra": auth_log.extra, + "exception_type": auth_log.exception_type, + "exception_message": auth_log.exception_message, + }, + default=str, + ) + + +def _assert_client_auth_log(auth_log: _CapturedAuthLog) -> None: + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.exception_type is None + assert auth_log.exception_message is None + assert "error_category" not in auth_log.extra + + +def _assert_jwks_dependency_auth_log(auth_log: _CapturedAuthLog) -> None: + assert auth_log.level == "ERROR" + assert auth_log.event == "exception.system" + assert auth_log.exception_type is not None + assert auth_log.exception_message is not None + assert "error_category" not in auth_log.extra + + +def _assert_log_excludes_jwks_body( + auth_log: _CapturedAuthLog, + *, + jwks_body: bytes, +) -> None: + raw_body = jwks_body.decode("utf-8", errors="ignore") + if raw_body: + assert raw_body not in _serialize_auth_log(auth_log) + + +def _create_token_without_key_id() -> str: + return jwt.encode( + { + "id": "contract-dashboard-user", + "exp": datetime.now(timezone.utc) + timedelta(minutes=5), + }, + "contract-secret-with-at-least-32-bytes", + algorithm="HS256", + ) + + +def _base64url_encode(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii") + + +def _create_token_with_malformed_json_payload(*, key_id: str) -> str: + header: dict[str, object] = {"alg": "RS256", "kid": key_id, "typ": "JWT"} + header_segment = _base64url_encode(json.dumps(header).encode("utf-8")) + payload_segment = _base64url_encode(b"not-json") + signature_segment = _base64url_encode(b"signature") + return f"{header_segment}.{payload_segment}.{signature_segment}" + + +@pytest.mark.asyncio +async def test_missing_key_id_is_a_client_warning_without_fetching_jwks( + monkeypatch: MonkeyPatch, +) -> None: + token = _create_token_without_key_id() + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_missing_key_id" + assert auth_log.extra["jwt_algorithm"] == "HS256" + assert auth_log.extra["jwt_kid_present"] is False + assert "jwt_kid" not in auth_log.extra + + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( + monkeypatch: MonkeyPatch, +) -> None: + token = "not-a-jwt" + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_kid_present"] is False + assert "jwt_algorithm" not in auth_log.extra + assert "jwt_kid" not in auth_log.extra + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_malformed_jwt_payload_is_a_client_warning_without_fetching_jwks( + monkeypatch: MonkeyPatch, +) -> None: + key_id = "malformed-payload-key" + token = _create_token_with_malformed_json_payload(key_id=key_id) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_raw_response(b"unavailable", status_code=503) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid_present"] is True + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_unknown_key_id_refreshes_once_and_remains_a_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + attacker_key_id = f"unknown\nkey:{'x' * 80}" + token = _create_rsa_token( + signing_key, + key_id=attacker_key_id, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id="known-key")]} + ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 2 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_unknown_key_id" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid_present"] is True + assert auth_log.extra["jwt_kid"] == "unknown_key:" + ("x" * 52) + + _assert_log_excludes_token(auth_log, token=token) + serialized_log = json.dumps(auth_log.extra, default=str) + assert attacker_key_id not in serialized_log + + +@pytest.mark.asyncio +async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + token = _create_rsa_token(signing_key, key_id="unavailable-key") + jwks_body = b"dashboard-jwks-secret-body" + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_raw_response(jwks_body, status_code=503) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_jwks_dependency_auth_log(auth_log) + assert auth_log.exception_type == "PyJWKClientConnectionError" + assert auth_log.extra["failure_reason"] == "jwks_unavailable" + assert auth_log.extra["jwt_kid"] == "unavailable-key" + + _assert_log_excludes_token(auth_log, token=token) + _assert_log_excludes_jwks_body(auth_log, jwks_body=jwks_body) + + +@pytest.mark.parametrize( + "jwks_body", + [ + b'{"keys": []}', + b"not-json", + ], + ids=["empty-key-set", "malformed-json"], +) +@pytest.mark.asyncio +async def test_invalid_jwks_is_a_system_error_with_an_unchanged_response( + monkeypatch: MonkeyPatch, + jwks_body: bytes, +) -> None: + signing_key = _create_rsa_private_key() + token = _create_rsa_token(signing_key, key_id="invalid-jwks-key") + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_raw_response(jwks_body) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_jwks_dependency_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwks_invalid" + assert auth_log.extra["jwt_kid"] == "invalid-jwks-key" + + _assert_log_excludes_token(auth_log, token=token) + _assert_log_excludes_jwks_body(auth_log, jwks_body=jwks_body) + + +@pytest.mark.asyncio +async def test_valid_keyed_jwt_returns_identity_without_auth_rejection_log( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "valid-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + payload_overrides={"permission": "read_only"}, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + assert response.status_code == 200 + assert response.json() == { + "user_id": "contract-dashboard-user", + "permission": "read_only", + } + assert jwks_server.state.request_count == 1 + assert log_capture.records == [] + + +@pytest.mark.asyncio +async def test_expired_jwt_retains_response_and_logs_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "expired-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + payload_overrides={"exp": datetime.now(timezone.utc) - timedelta(minutes=5)}, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_expired" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_invalid_signature_retains_response_and_logs_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + jwks_key = _create_rsa_private_key() + key_id = "invalid-signature-key" + token = _create_rsa_token(signing_key, key_id=key_id) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(jwks_key, key_id=key_id)]} + ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_missing_user_claim_retains_response_and_logs_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "missing-user-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + payload_overrides={"id": None}, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "unsupported-algorithm-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + algorithm="PS256", + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _assert_client_auth_log(auth_log) + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_kid_present"] is True + assert auth_log.extra["jwt_kid"] == key_id + assert "jwt_algorithm" not in auth_log.extra + _assert_log_excludes_token(auth_log, token=token) + + +def test_logfire_exception_callback_downgrades_auth_exceptions_by_status() -> None: + helper = _FakeLogfireExceptionHelper(exception=_create_auth_exception()) + + _downgrade_logfire_exception(helper) + + assert helper.level == "warning" + assert helper.is_recording_exception is False diff --git a/apps/api/tests/contract/test_dashboard_token_permission_contract.py b/apps/api/tests/contract/test_dashboard_token_permission_contract.py index 4bd74edbe..baf1216e6 100644 --- a/apps/api/tests/contract/test_dashboard_token_permission_contract.py +++ b/apps/api/tests/contract/test_dashboard_token_permission_contract.py @@ -1,47 +1,15 @@ from collections.abc import Callable from contextlib import AbstractAsyncContextManager -from datetime import datetime, timedelta, timezone -from typing import Literal, cast +from typing import cast from uuid import uuid4 -import jwt import pytest from httpx import AsyncClient from pytest import MonkeyPatch from shared.testing.contract_runtime import seed_contract_developer from tests.support.contract_database import ContractDatabase - -Permission = Literal["read_only", "full_access"] - - -def _use_dashboard_token( - api_client: AsyncClient, - monkeypatch: MonkeyPatch, - *, - user_id: str, - permission: Permission | None, -) -> None: - jwt_secret = f"contract-jwt-secret-{uuid4().hex[:12]}" - payload: dict[str, object] = { - "id": user_id, - "exp": datetime.now(timezone.utc) + timedelta(minutes=5), - } - if permission is not None: - payload["permission"] = permission - - token = jwt.encode(payload, jwt_secret, algorithm="HS256") - - from app.services.auth.dashboard_jwt_authentication_service import ( - get_dashboard_jwt_authentication_service, - ) - - monkeypatch.setattr( - get_dashboard_jwt_authentication_service(), - "_get_verification_key", - lambda _token: jwt_secret, - ) - api_client.headers.update({"Authorization": f"Bearer {token}"}) +from tests.support.dashboard_jwt import use_dashboard_jwks_token async def _seed_dashboard_user() -> str: @@ -105,19 +73,21 @@ async def test_read_only_dashboard_token_can_read_but_cannot_parse_or_archive( user_id=user_id, namespace="contract-permission", ) - _use_dashboard_token( + + with use_dashboard_jwks_token( api_client, monkeypatch, user_id=user_id, permission="read_only", - ) - - list_jobs_response = await api_client.get("/api/v1/jobs") - get_document_response = await api_client.get(f"/api/v1/documents/{document_id}") - create_job_response = await api_client.post("/api/v1/jobs", json=payload) - archive_document_response = await api_client.post( - f"/api/v1/documents/{document_id}/archive" - ) + ): + list_jobs_response = await api_client.get("/api/v1/jobs") + get_document_response = await api_client.get( + f"/api/v1/documents/{document_id}" + ) + create_job_response = await api_client.post("/api/v1/jobs", json=payload) + archive_document_response = await api_client.post( + f"/api/v1/documents/{document_id}/archive" + ) assert list_jobs_response.status_code == 200 assert get_document_response.status_code == 200 @@ -146,14 +116,13 @@ async def test_dashboard_token_without_permission_claim_keeps_full_access( async with api_client_factory() as api_client: user_id = await _seed_dashboard_user() - _use_dashboard_token( + + with use_dashboard_jwks_token( api_client, monkeypatch, user_id=user_id, - permission=None, - ) - - response = await api_client.post("/api/v1/jobs", json=payload) + ): + response = await api_client.post("/api/v1/jobs", json=payload) assert response.status_code == 200 response_json = cast(dict[str, object], response.json()) diff --git a/apps/api/tests/contract/test_job_creation_contract.py b/apps/api/tests/contract/test_job_creation_contract.py index 5d419de87..a375f5151 100644 --- a/apps/api/tests/contract/test_job_creation_contract.py +++ b/apps/api/tests/contract/test_job_creation_contract.py @@ -1,12 +1,11 @@ from collections.abc import Callable from contextlib import AbstractAsyncContextManager -from datetime import datetime, timedelta, timezone +from datetime import datetime, timezone import json import socket from typing import cast from uuid import uuid4 -import jwt import pytest from httpx import AsyncClient from pytest import MonkeyPatch @@ -15,6 +14,7 @@ from shared.testing.contract_runtime import get_contract_database_url from tests.support.contract_database import ContractDatabase +from tests.support.dashboard_jwt import use_dashboard_jwks_token async def _create_contract_engine() -> AsyncEngine: @@ -695,15 +695,6 @@ async def test_should_reject_authenticated_user_id_missing_from_user_table( monkeypatch: MonkeyPatch, ) -> None: user_id = f"contract-missing-user-{uuid4().hex[:12]}" - jwt_secret = f"contract-jwt-secret-{uuid4().hex[:12]}" - token = jwt.encode( - { - "id": user_id, - "exp": datetime.now(timezone.utc) + timedelta(minutes=5), - }, - jwt_secret, - algorithm="HS256", - ) payload: dict[str, str] = { "namespace": "contract-jobs", "source_type": "file", @@ -712,18 +703,12 @@ async def test_should_reject_authenticated_user_id_missing_from_user_table( } async with api_client_factory() as api_client: - from app.services.auth.dashboard_jwt_authentication_service import ( - get_dashboard_jwt_authentication_service, - ) - - monkeypatch.setattr( - get_dashboard_jwt_authentication_service(), - "_get_verification_key", - lambda _token: jwt_secret, - ) - - api_client.headers.update({"Authorization": f"Bearer {token}"}) - response = await api_client.post("/api/v1/jobs", json=payload) + with use_dashboard_jwks_token( + api_client, + monkeypatch, + user_id=user_id, + ): + response = await api_client.post("/api/v1/jobs", json=payload) assert response.status_code == 401 assert response.headers["x-request-id"] diff --git a/apps/api/tests/support/dashboard_jwt.py b/apps/api/tests/support/dashboard_jwt.py new file mode 100644 index 000000000..70d1b6738 --- /dev/null +++ b/apps/api/tests/support/dashboard_jwt.py @@ -0,0 +1,172 @@ +from __future__ import annotations + +import json +import threading +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass, field +from datetime import datetime, timedelta, timezone +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Literal, cast + +import jwt +from cryptography.hazmat.primitives.asymmetric.rsa import ( + RSAPrivateKey, + generate_private_key, +) +from httpx import AsyncClient +from pytest import MonkeyPatch + +DashboardPermission = Literal["read_only", "full_access"] + + +@dataclass +class JWKSResponseState: + body: bytes = b'{"keys": []}' + status_code: int = 200 + request_count: int = 0 + lock: threading.Lock = field(default_factory=threading.Lock) + + def record_request(self) -> tuple[bytes, int]: + with self.lock: + self.request_count += 1 + return self.body, self.status_code + + def set_json_response( + self, + response: Mapping[str, object], + *, + status_code: int = 200, + ) -> None: + with self.lock: + self.body = json.dumps(response).encode("utf-8") + self.status_code = status_code + + def set_raw_response( + self, + body: bytes, + *, + status_code: int = 200, + ) -> None: + with self.lock: + self.body = body + self.status_code = status_code + + +@dataclass(frozen=True) +class LocalJWKSServer: + endpoint: str + state: JWKSResponseState + + +def create_dashboard_rsa_private_key() -> RSAPrivateKey: + return generate_private_key(public_exponent=65537, key_size=2048) + + +def create_dashboard_rsa_jwk( + private_key: RSAPrivateKey, + *, + key_id: str, +) -> dict[str, object]: + jwk = cast( + dict[str, object], + jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key(), as_dict=True), + ) + return { + **jwk, + "kid": key_id, + "use": "sig", + "alg": "RS256", + } + + +def create_dashboard_rsa_token( + private_key: RSAPrivateKey, + *, + key_id: str, + user_id: str = "contract-dashboard-user", + permission: DashboardPermission | None = None, + expires_at: datetime | None = None, + payload_overrides: Mapping[str, object] | None = None, + algorithm: str = "RS256", +) -> str: + payload: dict[str, object] = { + "id": user_id, + "exp": expires_at or datetime.now(timezone.utc) + timedelta(minutes=5), + } + if permission is not None: + payload["permission"] = permission + payload.update(payload_overrides or {}) + + return jwt.encode( + payload, + private_key, + algorithm=algorithm, + headers={"kid": key_id}, + ) + + +def _create_jwks_handler( + state: JWKSResponseState, +) -> type[BaseHTTPRequestHandler]: + class JWKSRequestHandler(BaseHTTPRequestHandler): + def do_GET(self) -> None: + body, status_code = state.record_request() + self.send_response(status_code) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, format: str, *args: object) -> None: + return + + return JWKSRequestHandler + + +@contextmanager +def serve_dashboard_jwks() -> Iterator[LocalJWKSServer]: + state = JWKSResponseState() + server = ThreadingHTTPServer(("127.0.0.1", 0), _create_jwks_handler(state)) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + host, port = server.server_address + + try: + yield LocalJWKSServer(endpoint=f"http://{host}:{port}", state=state) + finally: + server.shutdown() + server.server_close() + thread.join(timeout=5) + + +@contextmanager +def use_dashboard_jwks_token( + api_client: AsyncClient, + monkeypatch: MonkeyPatch, + *, + user_id: str, + permission: DashboardPermission | None = None, +) -> Iterator[str]: + key_id = f"contract-dashboard-key-{user_id}" + signing_key = create_dashboard_rsa_private_key() + token = create_dashboard_rsa_token( + signing_key, + key_id=key_id, + user_id=user_id, + permission=permission, + ) + + with serve_dashboard_jwks() as jwks_server: + from shared.core.config import settings + + jwks_server.state.set_json_response( + {"keys": [create_dashboard_rsa_jwk(signing_key, key_id=key_id)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + api_client.headers.update({"Authorization": f"Bearer {token}"}) + yield token diff --git a/packages/shared-python/shared/core/exceptions/knowhere_exception.py b/packages/shared-python/shared/core/exceptions/knowhere_exception.py index 9158be63f..424431a0b 100644 --- a/packages/shared-python/shared/core/exceptions/knowhere_exception.py +++ b/packages/shared-python/shared/core/exceptions/knowhere_exception.py @@ -59,13 +59,14 @@ raise KnowhereException(code=ErrorCode.INVALID_ARGUMENT, ...) """ -from typing import Any, Dict, Optional +from typing import Any, Dict, Literal, Optional from shared.core.response.ErrorCode import ErrorCode, ErrorCodeMapper # Default messages for auto-sanitization DEFAULT_5XX_USER_MESSAGE = "An internal system error occurred. Please contact support." DEFAULT_4XX_USER_MESSAGE = "Invalid request. Please check your input." +LogErrorCategory = Literal["client", "system"] class KnowhereException(Exception): @@ -215,13 +216,10 @@ def to_log(self) -> Dict[str, Any]: - details: Additional structured data - original_exception: Wrapped exception info """ - # Determine error category based on HTTP status - error_category = "system" if self.http_status_code >= 500 else "client" - log_data: Dict[str, Any] = { "error_code": self.code.value, "http_status": self.http_status_code, - "error_category": error_category, + "error_category": _get_error_category(self.http_status_code), "exception_class": self.__class__.__name__, "internal_message": self.internal_message, "user_message": self.user_message, @@ -324,3 +322,11 @@ def _reconstruct_knowhere_exception(cls, state): obj = cls.__new__(cls) obj.__setstate__(state) return obj + + +def _get_error_category(http_status_code: int) -> LogErrorCategory: + """Return the stable log category derived from the public HTTP status.""" + if http_status_code >= 500: + return "system" + + return "client"