From af90aa5e140aea3e982774f6b01b48bd999bf132 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 02:09:55 +0800 Subject: [PATCH 1/8] fix: classify dashboard jwt telemetry --- .../dashboard_jwt_authentication_service.py | 191 +++++- ...t_dashboard_jwt_authentication_contract.py | 582 ++++++++++++++++++ ...est_dashboard_token_permission_contract.py | 65 +- .../contract/test_job_creation_contract.py | 31 +- apps/api/tests/support/dashboard_jwt.py | 172 ++++++ .../core/exceptions/domain_exceptions.py | 7 +- .../core/exceptions/knowhere_exception.py | 26 +- packages/shared-python/shared/core/logging.py | 2 +- 8 files changed, 962 insertions(+), 114 deletions(-) create mode 100644 apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py create mode 100644 apps/api/tests/support/dashboard_jwt.py 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..52339508a 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,37 @@ 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, cast import jwt -from jwt import PyJWKClient -from loguru import logger +from jwt import PyJWKClient, PyJWKClientConnectionError, PyJWKClientError, PyJWKSetError +from jwt.algorithms import AllowedPublicKeys from shared.core.config import settings from shared.core.exceptions.domain_exceptions import AuthException 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") 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) @@ -41,46 +55,115 @@ 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) + unverified_header = cast(dict[str, object], jwt.get_unverified_header(token)) + except jwt.InvalidTokenError: + raise _create_auth_exception( + user_message="Invalid token", + failure_reason="jwt_invalid", + exception_context=_build_exception_context( + algorithm=None, + key_id=None, + ), + ) from 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 + exception_context = _build_exception_context( + algorithm=algorithm, + key_id=key_id, + ) + + if key_id is None or not key_id.strip(): + raise _create_auth_exception( + failure_reason="jwt_missing_key_id", + exception_context=exception_context, + ) + + try: + key = self._get_verification_key(key_id) + if key is None: + raise _create_auth_exception( + failure_reason="jwt_unknown_key_id", + exception_context=exception_context, + ) + + payload = self._decode_payload(token, key) user_id = payload.get("id") if not isinstance(user_id, str) or not user_id: - raise AuthException(user_message="Token missing 'id' claim") + raise _create_auth_exception( + user_message="Token missing 'id' claim", + failure_reason="jwt_invalid", + exception_context=exception_context, + ) 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}, + raise _create_auth_exception( + user_message="Token has expired", + failure_reason="jwt_expired", + exception_context=exception_context, + ) from None + except PyJWKClientConnectionError: + raise _create_auth_exception( + failure_reason="jwks_unavailable", + error_category="system", + exception_context=exception_context, + ) from None + except (json.JSONDecodeError, UnicodeDecodeError, PyJWKSetError, jwt.PyJWKError): + raise _create_auth_exception( + failure_reason="jwks_invalid", + error_category="system", + exception_context=exception_context, + ) from None + except PyJWKClientError: + raise _create_auth_exception( + failure_reason="jwks_invalid", + error_category="system", + exception_context=exception_context, + ) from None + except jwt.InvalidTokenError: + raise _create_auth_exception( + user_message="Invalid token", + failure_reason="jwt_invalid", + exception_context=exception_context, + ) from None + + def _decode_payload( + self, + token: str, + key: VerificationKey, + ) -> dict[str, object]: + payload = cast( + dict[str, object], + jwt.decode( + token, + key, + algorithms=list(JWT_ALGORITHMS), + leeway=timedelta(seconds=30), + options={"verify_aud": False}, + ), ) return payload - def _get_verification_key(self, token: str) -> Any: + def _get_verification_key(self, key_id: str) -> VerificationKey | None: """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" - 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}" - ) - ) - except jwt.PyJWKSetError as exc: - logger.error(f"Invalid JWKS format: {exc}") - raise AuthException(internal_message=f"Invalid JWKS format: {exc}") + 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 +179,51 @@ 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 _build_exception_context( + *, + algorithm: object, + key_id: str | None, +) -> dict[str, object]: + is_key_id_present = key_id is not None and bool(key_id.strip()) + context: dict[str, object] = { + "auth_component": "dashboard_jwt", + "jwt_kid_present": is_key_id_present, + } + if isinstance(algorithm, str) and algorithm in JWT_ALGORITHMS: + context["jwt_algorithm"] = algorithm + if is_key_id_present and key_id is not None: + context["jwt_kid"] = _sanitize_key_id(key_id) + return context + + +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 _create_auth_exception( + *, + failure_reason: JWTFailureReason, + user_message: str = "Authentication required", + error_category: Literal["client", "system"] | None = None, + exception_context: dict[str, object] | None = None, +) -> AuthException: + context: dict[str, object] = { + **(exception_context or {}), + "failure_reason": failure_reason, + } + return AuthException( + user_message=user_message, + internal_message=f"Dashboard JWT authentication failed: {failure_reason}", + error_category=error_category, + exception_context=context, + ) + + 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..869ddbd9a --- /dev/null +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -0,0 +1,582 @@ +from __future__ import annotations + +import json +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING, 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 app.core.exception_handlers import setup_exception_handlers +from app.services.auth.dashboard_jwt_authentication_service import ( + DashboardJWTAuthenticationService, +) +from shared.core.config import settings +from shared.core.exceptions.domain_exceptions import AuthException +from shared.core.logging import _downgrade_expected_logfire_exception +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, +) + +if TYPE_CHECKING: + from logfire.types import ExceptionCallbackHelper + + +class _LoguruMessage(Protocol): + @property + def record(self) -> Mapping[str, object]: ... + + +@dataclass(frozen=True) +class _CapturedAuthLog: + level: str + event: str + message: str + extra: Mapping[str, object] + + +@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 = record["level"] + self.records.append( + _CapturedAuthLog( + level=str(getattr(level, "name", level)), + event=str(extra.get("event", "")), + message=str(record["message"]), + extra=dict(extra), + ) + ) + + +@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 _create_authentication_app() -> FastAPI: + 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 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 = json.dumps(auth_log.extra, default=str) + 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 _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", + ) + + +@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: + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + 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 auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + 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: + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Invalid token", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + 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_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")]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + 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 auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + 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") + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_raw_response(b"unavailable", status_code=503) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + 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 auth_log.level == "ERROR" + assert auth_log.event == "exception.system" + assert auth_log.extra["error_category"] == "system" + assert auth_log.extra["failure_reason"] == "jwks_unavailable" + assert auth_log.extra["jwt_kid"] == "unavailable-key" + + _assert_log_excludes_token(auth_log, token=token) + + +@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) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + 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 auth_log.level == "ERROR" + assert auth_log.event == "exception.system" + assert auth_log.extra["error_category"] == "system" + 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) + + +@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)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + 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)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Token has expired", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["error_category"] == "client" + 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)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Invalid token", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["error_category"] == "client" + 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)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Token missing 'id' claim", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["error_category"] == "client" + 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)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Invalid token", + token=token, + ) + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + 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_keeps_system_category_auth_errors_recorded() -> None: + helper = _FakeLogfireExceptionHelper( + exception=AuthException( + error_category="system", + exception_context={ + "auth_component": "dashboard_jwt", + "failure_reason": "jwks_unavailable", + }, + ) + ) + + _downgrade_expected_logfire_exception( + cast("ExceptionCallbackHelper", helper), + ) + + assert helper.level == "error" + assert helper.is_recording_exception is True + + +def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: + helper = _FakeLogfireExceptionHelper(exception=AuthException()) + + _downgrade_expected_logfire_exception( + cast("ExceptionCallbackHelper", 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/domain_exceptions.py b/packages/shared-python/shared/core/exceptions/domain_exceptions.py index 57ddd79a1..6b8b518be 100644 --- a/packages/shared-python/shared/core/exceptions/domain_exceptions.py +++ b/packages/shared-python/shared/core/exceptions/domain_exceptions.py @@ -42,9 +42,10 @@ # User sees: "An internal system error occurred. Please contact support." """ +from collections.abc import Mapping from typing import Any, Dict, List, Optional, TypedDict -from shared.core.exceptions.knowhere_exception import KnowhereException +from shared.core.exceptions.knowhere_exception import ErrorCategory, KnowhereException from shared.core.response.ErrorCode import ErrorCode, SubCode # ============================================================================ @@ -112,12 +113,16 @@ def __init__( self, user_message: str = "Authentication required", internal_message: Optional[str] = None, + error_category: ErrorCategory | None = None, + exception_context: Mapping[str, object] | None = None, ): super().__init__( code=ErrorCode.UNAUTHENTICATED, internal_message=internal_message or user_message, user_message=user_message, details={}, # Empty for security + error_category=error_category, + exception_context=exception_context, ) diff --git a/packages/shared-python/shared/core/exceptions/knowhere_exception.py b/packages/shared-python/shared/core/exceptions/knowhere_exception.py index 9158be63f..22c91fbc3 100644 --- a/packages/shared-python/shared/core/exceptions/knowhere_exception.py +++ b/packages/shared-python/shared/core/exceptions/knowhere_exception.py @@ -59,13 +59,15 @@ raise KnowhereException(code=ErrorCode.INVALID_ARGUMENT, ...) """ -from typing import Any, Dict, Optional +from collections.abc import Mapping +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." +ErrorCategory = Literal["client", "system"] class KnowhereException(Exception): @@ -121,6 +123,8 @@ def __init__( details: Optional[Dict[str, Any]] = None, http_status_code: Optional[int] = None, original_exception: Optional[Exception] = None, + error_category: ErrorCategory | None = None, + exception_context: Mapping[str, object] | None = None, ): """ Initialize a KnowhereException. @@ -134,6 +138,10 @@ def __init__( details: Optional structured data to include in response (must be safe). http_status_code: Override HTTP status (auto-derived from code if None). original_exception: The underlying exception being wrapped (for logging). + error_category: Optional telemetry category override. Defaults to the + category derived from the HTTP status. + exception_context: Internal-only structured telemetry fields. These are + included in logs and never returned to clients. """ super().__init__(internal_message) self.code = code @@ -143,6 +151,11 @@ def __init__( http_status_code or ErrorCodeMapper.get_http_status_from_error_code(code) ) self.original_exception = original_exception + default_error_category: ErrorCategory = ( + "system" if self.http_status_code >= 500 else "client" + ) + self.error_category: ErrorCategory = error_category or default_error_category + self.exception_context: Dict[str, object] = dict(exception_context or {}) # ======================================================================= # SECURITY: Auto-sanitize user_message based on HTTP status @@ -215,13 +228,11 @@ 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] = { + **self.exception_context, "error_code": self.code.value, "http_status": self.http_status_code, - "error_category": error_category, + "error_category": self.error_category, "exception_class": self.__class__.__name__, "internal_message": self.internal_message, "user_message": self.user_message, @@ -275,8 +286,9 @@ def logging(self, **extra_context): } # Log at appropriate level with appropriate event - if self.http_status_code >= 500: - # 5xx: ERROR level with stacktrace + if self.error_category == "system": + # System-category errors use ERROR level with a stacktrace even when + # their public HTTP status intentionally remains a 4xx response. logger.bind(event=LogEvent.EXCEPTION_SYSTEM.value, **log_data).opt( exception=self ).error(self.internal_message) diff --git a/packages/shared-python/shared/core/logging.py b/packages/shared-python/shared/core/logging.py index bb246f18e..60b815666 100644 --- a/packages/shared-python/shared/core/logging.py +++ b/packages/shared-python/shared/core/logging.py @@ -84,7 +84,7 @@ def _is_expected_client_exception(exception: BaseException) -> bool: from shared.core.exceptions.knowhere_exception import KnowhereException if isinstance(exception, KnowhereException): - return 400 <= exception.http_status_code < 500 + return exception.error_category == "client" try: from fastapi import HTTPException as FastAPIHTTPException From abd83fbde7777d6f59b24a194a734f5f222368f7 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 02:26:37 +0800 Subject: [PATCH 2/8] test: address jwt telemetry code scanning comments --- ...est_dashboard_jwt_authentication_contract.py | 17 +++++------------ 1 file changed, 5 insertions(+), 12 deletions(-) diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index 869ddbd9a..3351f5a31 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -5,7 +5,7 @@ from contextlib import contextmanager from dataclasses import dataclass from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Protocol, cast +from typing import Protocol, cast import jwt import pytest @@ -28,13 +28,10 @@ serve_dashboard_jwks as _serve_jwks, ) -if TYPE_CHECKING: - from logfire.types import ExceptionCallbackHelper - - class _LoguruMessage(Protocol): @property - def record(self) -> Mapping[str, object]: ... + def record(self) -> Mapping[str, object]: + raise NotImplementedError @dataclass(frozen=True) @@ -563,9 +560,7 @@ def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() ) ) - _downgrade_expected_logfire_exception( - cast("ExceptionCallbackHelper", helper), - ) + _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] assert helper.level == "error" assert helper.is_recording_exception is True @@ -574,9 +569,7 @@ def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: helper = _FakeLogfireExceptionHelper(exception=AuthException()) - _downgrade_expected_logfire_exception( - cast("ExceptionCallbackHelper", helper), - ) + _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] assert helper.level == "warning" assert helper.is_recording_exception is False From 8e555022725f42a303674a9e13ecc100b3bfec23 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 10:11:22 +0800 Subject: [PATCH 3/8] test: isolate dashboard jwt contract imports --- ...t_dashboard_jwt_authentication_contract.py | 135 ++++++++++-------- 1 file changed, 74 insertions(+), 61 deletions(-) diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index 3351f5a31..f9de2cf68 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -1,11 +1,13 @@ from __future__ import annotations 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 typing import Protocol, cast +from pathlib import Path +from typing import Literal, Protocol, cast import jwt import pytest @@ -14,13 +16,10 @@ from loguru import logger from pytest import MonkeyPatch -from app.core.exception_handlers import setup_exception_handlers -from app.services.auth.dashboard_jwt_authentication_service import ( - DashboardJWTAuthenticationService, +from tests.support.import_environment import ( + configure_import_environment, + ensure_import_paths, ) -from shared.core.config import settings -from shared.core.exceptions.domain_exceptions import AuthException -from shared.core.logging import _downgrade_expected_logfire_exception from tests.support.dashboard_jwt import ( create_dashboard_rsa_jwk as _create_rsa_jwk, create_dashboard_rsa_private_key as _create_rsa_private_key, @@ -28,6 +27,10 @@ serve_dashboard_jwks as _serve_jwks, ) +configure_import_environment() +ensure_import_paths() + + class _LoguruMessage(Protocol): @property def record(self) -> Mapping[str, object]: @@ -83,7 +86,57 @@ def _capture_auth_logs() -> Iterator[_AuthLogCapture]: logger.remove(log_sink_id) +def _prepare_api_app_imports() -> None: + api_root = str(Path(__file__).resolve().parents[2]) + if api_root in sys.path: + sys.path.remove(api_root) + sys.path.insert(0, api_root) + + for module_name in list(sys.modules): + 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( + *, + error_category: Literal["client", "system"] | None = None, + exception_context: Mapping[str, object] | None = None, +) -> BaseException: + from shared.core.exceptions.domain_exceptions import AuthException + + return AuthException( + error_category=error_category, + exception_context=exception_context, + ) + + +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() @@ -170,11 +223,7 @@ async def test_missing_key_id_is_a_client_warning_without_fetching_jwks( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -204,11 +253,7 @@ async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -245,11 +290,7 @@ async def test_unknown_key_id_refreshes_once_and_remains_a_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id="known-key")]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -283,11 +324,7 @@ async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: jwks_server.state.set_raw_response(b"unavailable", status_code=503) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -327,11 +364,7 @@ async def test_invalid_jwks_is_a_system_error_with_an_unchanged_response( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: jwks_server.state.set_raw_response(jwks_body) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -369,11 +402,7 @@ async def test_valid_keyed_jwt_returns_identity_without_auth_rejection_log( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) assert response.status_code == 200 @@ -402,11 +431,7 @@ async def test_expired_jwt_retains_response_and_logs_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -441,11 +466,7 @@ async def test_invalid_signature_retains_response_and_logs_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(jwks_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -483,11 +504,7 @@ async def test_missing_user_claim_retains_response_and_logs_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -525,11 +542,7 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -551,7 +564,7 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() -> None: helper = _FakeLogfireExceptionHelper( - exception=AuthException( + exception=_create_auth_exception( error_category="system", exception_context={ "auth_component": "dashboard_jwt", @@ -560,16 +573,16 @@ def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() ) ) - _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] + _downgrade_logfire_exception(helper) assert helper.level == "error" assert helper.is_recording_exception is True def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: - helper = _FakeLogfireExceptionHelper(exception=AuthException()) + helper = _FakeLogfireExceptionHelper(exception=_create_auth_exception()) - _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] + _downgrade_logfire_exception(helper) assert helper.level == "warning" assert helper.is_recording_exception is False From 9645eb5552efd95f2338b708792c7b356b56e245 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 10:21:02 +0800 Subject: [PATCH 4/8] test: type dashboard jwt import isolation helper --- .../contract/test_dashboard_jwt_authentication_contract.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index f9de2cf68..f38320e05 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -87,12 +87,13 @@ def _capture_auth_logs() -> Iterator[_AuthLogCapture]: def _prepare_api_app_imports() -> None: - api_root = str(Path(__file__).resolve().parents[2]) + 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) - for module_name in list(sys.modules): + 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) From 1cd8a011a5a8794e5e0947eebbdaf1bb5f4f35ce Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 14:46:42 +0800 Subject: [PATCH 5/8] fix: reject malformed jwt payloads before jwks lookup --- .../dashboard_jwt_authentication_service.py | 32 ++++++++++++++ ...t_dashboard_jwt_authentication_contract.py | 44 +++++++++++++++++++ 2 files changed, 76 insertions(+) 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 52339508a..518634e7f 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -12,6 +12,7 @@ import jwt from jwt import PyJWKClient, PyJWKClientConnectionError, PyJWKClientError, PyJWKSetError from jwt.algorithms import AllowedPublicKeys +from jwt.types import Options from shared.core.config import settings from shared.core.exceptions.domain_exceptions import AuthException @@ -21,6 +22,16 @@ 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, + "verify_exp": False, + "verify_nbf": False, + "verify_iat": False, + "verify_aud": False, + "verify_iss": False, + "verify_sub": False, + "verify_jti": False, +} READ_ONLY_PERMISSION: Literal["read_only"] = "read_only" FULL_ACCESS_PERMISSION: Literal["full_access"] = "full_access" Permission = Literal["read_only", "full_access"] @@ -81,6 +92,7 @@ def decode_identity(self, token: str) -> DashboardJWTIdentity: ) try: + self._reject_malformed_token_before_jwks_lookup(token, exception_context) key = self._get_verification_key(key_id) if key is None: raise _create_auth_exception( @@ -147,6 +159,26 @@ def _decode_payload( ) return payload + def _reject_malformed_token_before_jwks_lookup( + self, + token: str, + exception_context: dict[str, object], + ) -> None: + """Reject structurally invalid JWTs before touching Dashboard JWKS.""" + try: + # This decode only checks token structure; verified claims come from + # _decode_payload after the signing key is resolved. + jwt.decode( + token, + options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, + ) + except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + raise _create_auth_exception( + user_message="Invalid token", + failure_reason="jwt_invalid", + exception_context=exception_context, + ) from None + 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() diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index f38320e05..4d0bc5097 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -1,5 +1,6 @@ from __future__ import annotations +import base64 import json import sys from collections.abc import Iterator, Mapping @@ -216,6 +217,18 @@ def _create_token_without_key_id() -> str: ) +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, @@ -275,6 +288,37 @@ async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( _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="Invalid token", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + 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, From e7c246358246a243111c4d17c1d0d7169a02c9bd Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 16:25:25 +0800 Subject: [PATCH 6/8] fix: keep dashboard jwt telemetry local --- .../dashboard_jwt_authentication_service.py | 148 +++++++++++------- ...t_dashboard_jwt_authentication_contract.py | 143 +++++++++-------- .../core/exceptions/domain_exceptions.py | 7 +- .../core/exceptions/knowhere_exception.py | 30 ++-- packages/shared-python/shared/core/logging.py | 2 +- 5 files changed, 187 insertions(+), 143 deletions(-) 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 518634e7f..300109e3d 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -7,15 +7,17 @@ import threading from dataclasses import dataclass from datetime import timedelta -from typing import Literal, cast +from typing import Literal, NoReturn, cast import jwt 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 @@ -68,79 +70,80 @@ def decode_identity(self, token: str) -> DashboardJWTIdentity: try: unverified_header = cast(dict[str, object], jwt.get_unverified_header(token)) except jwt.InvalidTokenError: - raise _create_auth_exception( - user_message="Invalid token", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=_build_exception_context( + telemetry_context=_build_telemetry_context( algorithm=None, key_id=None, ), - ) from 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 - exception_context = _build_exception_context( + telemetry_context = _build_telemetry_context( algorithm=algorithm, key_id=key_id, ) if key_id is None or not key_id.strip(): - raise _create_auth_exception( + _reject_client_jwt( failure_reason="jwt_missing_key_id", - exception_context=exception_context, + telemetry_context=telemetry_context, ) try: - self._reject_malformed_token_before_jwks_lookup(token, exception_context) + self._reject_malformed_token_before_jwks_lookup(token, telemetry_context) key = self._get_verification_key(key_id) if key is None: - raise _create_auth_exception( + _reject_client_jwt( failure_reason="jwt_unknown_key_id", - exception_context=exception_context, + telemetry_context=telemetry_context, ) payload = self._decode_payload(token, key) user_id = payload.get("id") if not isinstance(user_id, str) or not user_id: - raise _create_auth_exception( - user_message="Token missing 'id' claim", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=exception_context, + telemetry_context=telemetry_context, ) permission = _normalize_permission(payload.get("permission")) return DashboardJWTIdentity(user_id=user_id, permission=permission) except jwt.ExpiredSignatureError: - raise _create_auth_exception( - user_message="Token has expired", + _reject_client_jwt( failure_reason="jwt_expired", - exception_context=exception_context, - ) from None - except PyJWKClientConnectionError: - raise _create_auth_exception( + telemetry_context=telemetry_context, + ) + except PyJWKClientConnectionError as error: + _reject_jwks_dependency( failure_reason="jwks_unavailable", - error_category="system", - exception_context=exception_context, - ) from None - except (json.JSONDecodeError, UnicodeDecodeError, PyJWKSetError, jwt.PyJWKError): - raise _create_auth_exception( + telemetry_context=telemetry_context, + original_exception=error, + ) + except ( + json.JSONDecodeError, + UnicodeDecodeError, + PyJWKSetError, + jwt.PyJWKError, + ) as error: + _reject_jwks_dependency( failure_reason="jwks_invalid", - error_category="system", - exception_context=exception_context, - ) from None - except PyJWKClientError: - raise _create_auth_exception( + telemetry_context=telemetry_context, + original_exception=error, + ) + except PyJWKClientError as error: + _reject_jwks_dependency( failure_reason="jwks_invalid", - error_category="system", - exception_context=exception_context, - ) from None + telemetry_context=telemetry_context, + original_exception=error, + ) except jwt.InvalidTokenError: - raise _create_auth_exception( - user_message="Invalid token", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=exception_context, - ) from None + telemetry_context=telemetry_context, + ) def _decode_payload( self, @@ -162,7 +165,7 @@ def _decode_payload( def _reject_malformed_token_before_jwks_lookup( self, token: str, - exception_context: dict[str, object], + telemetry_context: dict[str, object], ) -> None: """Reject structurally invalid JWTs before touching Dashboard JWKS.""" try: @@ -173,11 +176,10 @@ def _reject_malformed_token_before_jwks_lookup( options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, ) except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): - raise _create_auth_exception( - user_message="Invalid token", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=exception_context, - ) from None + telemetry_context=telemetry_context, + ) def _get_verification_key(self, key_id: str) -> VerificationKey | None: """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" @@ -215,7 +217,7 @@ def _get_jwks_client(self) -> PyJWKClient: return self._jwks_client -def _build_exception_context( +def _build_telemetry_context( *, algorithm: object, key_id: str | None, @@ -237,23 +239,57 @@ def _sanitize_key_id(key_id: str) -> str: return sanitized_key_id[:JWT_KEY_ID_MAX_LENGTH] -def _create_auth_exception( +def _reject_client_jwt( *, failure_reason: JWTFailureReason, - user_message: str = "Authentication required", - error_category: Literal["client", "system"] | None = None, - exception_context: dict[str, object] | None = None, -) -> AuthException: - context: dict[str, object] = { - **(exception_context or {}), + telemetry_context: dict[str, object], +) -> 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: dict[str, object], + 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: dict[str, object], + is_jwks_dependency_failure: bool, + original_exception: Exception | None = None, +) -> None: + log_data: dict[str, object] = { + **telemetry_context, "failure_reason": failure_reason, } - return AuthException( - user_message=user_message, - internal_message=f"Dashboard JWT authentication failed: {failure_reason}", - error_category=error_category, - exception_context=context, - ) + 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: diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index 4d0bc5097..5d38d3a5a 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -8,7 +8,7 @@ from dataclasses import dataclass from datetime import datetime, timedelta, timezone from pathlib import Path -from typing import Literal, Protocol, cast +from typing import Protocol, cast import jwt import pytest @@ -44,6 +44,8 @@ class _CapturedAuthLog: event: str message: str extra: Mapping[str, object] + exception_type: str | None + exception_message: str | None @dataclass @@ -66,13 +68,25 @@ def capture(self, message: _LoguruMessage) -> None: if extra.get("auth_component") != "dashboard_jwt": return - level = record["level"] + 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, ) ) @@ -112,17 +126,10 @@ def _use_dashboard_endpoint( ) -def _create_auth_exception( - *, - error_category: Literal["client", "system"] | None = None, - exception_context: Mapping[str, object] | None = None, -) -> BaseException: +def _create_auth_exception() -> BaseException: from shared.core.exceptions.domain_exceptions import AuthException - return AuthException( - error_category=error_category, - exception_context=exception_context, - ) + return AuthException() def _downgrade_logfire_exception(helper: _FakeLogfireExceptionHelper) -> None: @@ -185,6 +192,8 @@ def _assert_unauthenticated_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(".") @@ -197,7 +206,7 @@ def _assert_log_excludes_token( *, token: str, ) -> None: - serialized_log = json.dumps(auth_log.extra, default=str) + 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(".") @@ -206,6 +215,44 @@ def _assert_log_excludes_token( 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( { @@ -249,8 +296,7 @@ async def test_missing_key_id_is_a_client_warning_without_fetching_jwks( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _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 @@ -272,15 +318,14 @@ async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( _assert_unauthenticated_response( response, - expected_message="Invalid token", + 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 auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _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 @@ -303,15 +348,14 @@ async def test_malformed_jwt_payload_is_a_client_warning_without_fetching_jwks( _assert_unauthenticated_response( response, - expected_message="Invalid token", + 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 auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _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 @@ -347,8 +391,7 @@ async def test_unknown_key_id_refreshes_once_and_remains_a_client_warning( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _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 @@ -365,10 +408,11 @@ async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( ) -> 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(b"unavailable", status_code=503) + jwks_server.state.set_raw_response(jwks_body, status_code=503) _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) @@ -381,13 +425,13 @@ async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "ERROR" - assert auth_log.event == "exception.system" - assert auth_log.extra["error_category"] == "system" + _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( @@ -421,13 +465,12 @@ async def test_invalid_jwks_is_a_system_error_with_an_unchanged_response( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "ERROR" - assert auth_log.event == "exception.system" - assert auth_log.extra["error_category"] == "system" + _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 @@ -481,16 +524,14 @@ async def test_expired_jwt_retains_response_and_logs_client_warning( _assert_unauthenticated_response( response, - expected_message="Token has expired", + 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 auth_log.level == "WARNING" - assert auth_log.event == "exception.client" - assert auth_log.extra["error_category"] == "client" + _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 @@ -516,16 +557,14 @@ async def test_invalid_signature_retains_response_and_logs_client_warning( _assert_unauthenticated_response( response, - expected_message="Invalid token", + 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 auth_log.level == "WARNING" - assert auth_log.event == "exception.client" - assert auth_log.extra["error_category"] == "client" + _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 @@ -554,16 +593,14 @@ async def test_missing_user_claim_retains_response_and_logs_client_warning( _assert_unauthenticated_response( response, - expected_message="Token missing 'id' claim", + 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 auth_log.level == "WARNING" - assert auth_log.event == "exception.client" - assert auth_log.extra["error_category"] == "client" + _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 @@ -592,14 +629,13 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( _assert_unauthenticated_response( response, - expected_message="Invalid token", + expected_message="Authentication required", token=token, ) assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _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 @@ -607,24 +643,7 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( _assert_log_excludes_token(auth_log, token=token) -def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() -> None: - helper = _FakeLogfireExceptionHelper( - exception=_create_auth_exception( - error_category="system", - exception_context={ - "auth_component": "dashboard_jwt", - "failure_reason": "jwks_unavailable", - }, - ) - ) - - _downgrade_logfire_exception(helper) - - assert helper.level == "error" - assert helper.is_recording_exception is True - - -def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: +def test_logfire_exception_callback_downgrades_auth_exceptions_by_status() -> None: helper = _FakeLogfireExceptionHelper(exception=_create_auth_exception()) _downgrade_logfire_exception(helper) diff --git a/packages/shared-python/shared/core/exceptions/domain_exceptions.py b/packages/shared-python/shared/core/exceptions/domain_exceptions.py index 6b8b518be..57ddd79a1 100644 --- a/packages/shared-python/shared/core/exceptions/domain_exceptions.py +++ b/packages/shared-python/shared/core/exceptions/domain_exceptions.py @@ -42,10 +42,9 @@ # User sees: "An internal system error occurred. Please contact support." """ -from collections.abc import Mapping from typing import Any, Dict, List, Optional, TypedDict -from shared.core.exceptions.knowhere_exception import ErrorCategory, KnowhereException +from shared.core.exceptions.knowhere_exception import KnowhereException from shared.core.response.ErrorCode import ErrorCode, SubCode # ============================================================================ @@ -113,16 +112,12 @@ def __init__( self, user_message: str = "Authentication required", internal_message: Optional[str] = None, - error_category: ErrorCategory | None = None, - exception_context: Mapping[str, object] | None = None, ): super().__init__( code=ErrorCode.UNAUTHENTICATED, internal_message=internal_message or user_message, user_message=user_message, details={}, # Empty for security - error_category=error_category, - exception_context=exception_context, ) diff --git a/packages/shared-python/shared/core/exceptions/knowhere_exception.py b/packages/shared-python/shared/core/exceptions/knowhere_exception.py index 22c91fbc3..424431a0b 100644 --- a/packages/shared-python/shared/core/exceptions/knowhere_exception.py +++ b/packages/shared-python/shared/core/exceptions/knowhere_exception.py @@ -59,7 +59,6 @@ raise KnowhereException(code=ErrorCode.INVALID_ARGUMENT, ...) """ -from collections.abc import Mapping from typing import Any, Dict, Literal, Optional from shared.core.response.ErrorCode import ErrorCode, ErrorCodeMapper @@ -67,7 +66,7 @@ # 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." -ErrorCategory = Literal["client", "system"] +LogErrorCategory = Literal["client", "system"] class KnowhereException(Exception): @@ -123,8 +122,6 @@ def __init__( details: Optional[Dict[str, Any]] = None, http_status_code: Optional[int] = None, original_exception: Optional[Exception] = None, - error_category: ErrorCategory | None = None, - exception_context: Mapping[str, object] | None = None, ): """ Initialize a KnowhereException. @@ -138,10 +135,6 @@ def __init__( details: Optional structured data to include in response (must be safe). http_status_code: Override HTTP status (auto-derived from code if None). original_exception: The underlying exception being wrapped (for logging). - error_category: Optional telemetry category override. Defaults to the - category derived from the HTTP status. - exception_context: Internal-only structured telemetry fields. These are - included in logs and never returned to clients. """ super().__init__(internal_message) self.code = code @@ -151,11 +144,6 @@ def __init__( http_status_code or ErrorCodeMapper.get_http_status_from_error_code(code) ) self.original_exception = original_exception - default_error_category: ErrorCategory = ( - "system" if self.http_status_code >= 500 else "client" - ) - self.error_category: ErrorCategory = error_category or default_error_category - self.exception_context: Dict[str, object] = dict(exception_context or {}) # ======================================================================= # SECURITY: Auto-sanitize user_message based on HTTP status @@ -229,10 +217,9 @@ def to_log(self) -> Dict[str, Any]: - original_exception: Wrapped exception info """ log_data: Dict[str, Any] = { - **self.exception_context, "error_code": self.code.value, "http_status": self.http_status_code, - "error_category": self.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, @@ -286,9 +273,8 @@ def logging(self, **extra_context): } # Log at appropriate level with appropriate event - if self.error_category == "system": - # System-category errors use ERROR level with a stacktrace even when - # their public HTTP status intentionally remains a 4xx response. + if self.http_status_code >= 500: + # 5xx: ERROR level with stacktrace logger.bind(event=LogEvent.EXCEPTION_SYSTEM.value, **log_data).opt( exception=self ).error(self.internal_message) @@ -336,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" diff --git a/packages/shared-python/shared/core/logging.py b/packages/shared-python/shared/core/logging.py index 60b815666..bb246f18e 100644 --- a/packages/shared-python/shared/core/logging.py +++ b/packages/shared-python/shared/core/logging.py @@ -84,7 +84,7 @@ def _is_expected_client_exception(exception: BaseException) -> bool: from shared.core.exceptions.knowhere_exception import KnowhereException if isinstance(exception, KnowhereException): - return exception.error_category == "client" + return 400 <= exception.http_status_code < 500 try: from fastapi import HTTPException as FastAPIHTTPException From fe6b88023b57f56ade23bb1c2ba679c649936751 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 16:59:47 +0800 Subject: [PATCH 7/8] refactor: simplify jwt structure decode options --- .../services/auth/dashboard_jwt_authentication_service.py | 7 ------- 1 file changed, 7 deletions(-) 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 300109e3d..79f9ec989 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -26,13 +26,6 @@ JWT_ALGORITHMS: tuple[str, ...] = ("HS256", "RS256", "EdDSA") JWT_STRUCTURE_ONLY_DECODE_OPTIONS: Options = { "verify_signature": False, - "verify_exp": False, - "verify_nbf": False, - "verify_iat": False, - "verify_aud": False, - "verify_iss": False, - "verify_sub": False, - "verify_jti": False, } READ_ONLY_PERMISSION: Literal["read_only"] = "read_only" FULL_ACCESS_PERMISSION: Literal["full_access"] = "full_access" From d26b6372ee3aefd08f99b625346927d8c7b33ca3 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 17:58:03 +0800 Subject: [PATCH 8/8] refactor: clarify dashboard jwt auth pipeline --- .../dashboard_jwt_authentication_service.py | 256 +++++++++++------- 1 file changed, 162 insertions(+), 94 deletions(-) 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 79f9ec989..0989460ff 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -47,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.""" @@ -60,55 +84,29 @@ 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: - 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( - 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 - telemetry_context = _build_telemetry_context( - algorithm=algorithm, - key_id=key_id, + 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, ) + payload = _verify_payload_or_reject( + token=token, + key=key, + telemetry_context=telemetry_context, + ) + return _build_identity_or_reject(payload, telemetry_context) - if key_id is None or not key_id.strip(): - _reject_client_jwt( - failure_reason="jwt_missing_key_id", - telemetry_context=telemetry_context, - ) - + 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: - self._reject_malformed_token_before_jwks_lookup(token, telemetry_context) key = self._get_verification_key(key_id) - if key is None: - _reject_client_jwt( - failure_reason="jwt_unknown_key_id", - telemetry_context=telemetry_context, - ) - - payload = self._decode_payload(token, key) - 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) - except jwt.ExpiredSignatureError: - _reject_client_jwt( - failure_reason="jwt_expired", - telemetry_context=telemetry_context, - ) except PyJWKClientConnectionError as error: _reject_jwks_dependency( failure_reason="jwks_unavailable", @@ -132,48 +130,15 @@ def decode_identity(self, token: str) -> DashboardJWTIdentity: telemetry_context=telemetry_context, original_exception=error, ) - except jwt.InvalidTokenError: - _reject_client_jwt( - failure_reason="jwt_invalid", - telemetry_context=telemetry_context, - ) - - def _decode_payload( - self, - token: str, - key: VerificationKey, - ) -> dict[str, object]: - payload = cast( - dict[str, object], - jwt.decode( - token, - key, - algorithms=list(JWT_ALGORITHMS), - leeway=timedelta(seconds=30), - options={"verify_aud": False}, - ), - ) - return payload - def _reject_malformed_token_before_jwks_lookup( - self, - token: str, - telemetry_context: dict[str, object], - ) -> None: - """Reject structurally invalid JWTs before touching Dashboard JWKS.""" - try: - # This decode only checks token structure; verified claims come from - # _decode_payload after the signing key is resolved. - jwt.decode( - token, - options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, - ) - except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + if key is None: _reject_client_jwt( - failure_reason="jwt_invalid", + failure_reason="jwt_unknown_key_id", telemetry_context=telemetry_context, ) + 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() @@ -210,21 +175,124 @@ def _get_jwks_client(self) -> PyJWKClient: 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, -) -> dict[str, object]: +) -> _DashboardJWTTelemetry: is_key_id_present = key_id is not None and bool(key_id.strip()) - context: dict[str, object] = { - "auth_component": "dashboard_jwt", - "jwt_kid_present": is_key_id_present, - } + jwt_algorithm: str | None = None if isinstance(algorithm, str) and algorithm in JWT_ALGORITHMS: - context["jwt_algorithm"] = algorithm + jwt_algorithm = algorithm + jwt_kid: str | None = None if is_key_id_present and key_id is not None: - context["jwt_kid"] = _sanitize_key_id(key_id) - return context + 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: @@ -235,7 +303,7 @@ def _sanitize_key_id(key_id: str) -> str: def _reject_client_jwt( *, failure_reason: JWTFailureReason, - telemetry_context: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, ) -> NoReturn: _log_dashboard_jwt_auth_failure( failure_reason=failure_reason, @@ -248,7 +316,7 @@ def _reject_client_jwt( def _reject_jwks_dependency( *, failure_reason: JWTFailureReason, - telemetry_context: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, original_exception: Exception, ) -> NoReturn: _log_dashboard_jwt_auth_failure( @@ -263,12 +331,12 @@ def _reject_jwks_dependency( def _log_dashboard_jwt_auth_failure( *, failure_reason: JWTFailureReason, - telemetry_context: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, is_jwks_dependency_failure: bool, original_exception: Exception | None = None, ) -> None: log_data: dict[str, object] = { - **telemetry_context, + **telemetry_context.to_log_data(), "failure_reason": failure_reason, } message = f"Dashboard JWT authentication failed: {failure_reason}"