From e035682ed17295b9c5f0a363e266e890de31061c Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 22:15:57 -0700 Subject: [PATCH] refactor(auth): separate JWT identity and OAuth authorization --- .../mcp_server/bridge_token_flow.py | 68 +++-- .../mcp_server/discoverable_endpoints.py | 5 +- litellm/proxy/auth/handle_jwt.py | 276 ++++++++++-------- .../mcp_server/test_discoverable_endpoints.py | 34 ++- .../proxy/auth/test_handle_jwt.py | 100 +++++-- 5 files changed, 300 insertions(+), 183 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 5fe0929773c..f001ee87dbd 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -24,6 +24,7 @@ if TYPE_CHECKING: UpstreamTokenGrant, ) from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.handle_jwt import JWTIdentity def _litellm_key_from_request(request: Request) -> str | None: @@ -304,26 +305,37 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol assert_never(identity.subject_type) -async def _extract_user_id_from_request(request: Request, server_id: str | None = None) -> str | None: - """Resolve identity for binding, or authorize the credential-write action for a target server.""" +async def _extract_user_id_from_request(request: Request) -> str | None: + """Resolve the caller for identity binding without granting credential-write permission.""" + from litellm.proxy.auth.handle_jwt import JWTIdentity # noqa: PLC0415 # proxy import cycle + + resolved: Final = await _resolve_request_auth(request) + if isinstance(resolved, JWTIdentity): + return resolved.user_id + return _active_key_user_id(resolved) if resolved is not None else None + + +async def authorize_oauth_credential_request(request: Request, server_id: str) -> str | None: from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle + + resolved: Final = await _resolve_request_auth(request, f"/v1/mcp/server/{server_id}/oauth-user-credential") + if not isinstance(resolved, UserAPIKeyAuth) or not _active_key_user_id(resolved): + return None + if not await can_store_oauth_credential(request, resolved, server_id): + return None + return resolved.user_id + + +async def _resolve_request_auth( + request: Request, write_route: str | None = None +) -> "UserAPIKeyAuth | JWTIdentity | None": from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle token: Final = _litellm_key_from_request(request) - # The OAuth relay is public; the optional server-side write is the same protected action - # as the direct credential endpoint. Authorize that action without rewriting the Request. - write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential" if server_id is not None else None - resolved: Final = ( - await _resolve_jwt_auth(request, token, write_route) - if token is not None and JWTHandler.is_jwt(token) - else await _resolve_active_litellm_key(request) - ) - auth: Final = resolved.key if isinstance(resolved, _ResolvedKey) else resolved - if not isinstance(auth, UserAPIKeyAuth) or not _active_key_user_id(auth): - return None - if server_id is not None and not await can_store_oauth_credential(request, auth, server_id): - return None - return auth.user_id + if token is not None and JWTHandler.is_jwt(token): + return await _resolve_jwt_auth(request, token, write_route) + resolved: Final = await _resolve_active_litellm_key(request) + return resolved.key if isinstance(resolved, _ResolvedKey) else None async def can_store_oauth_credential(request: Request, auth: "UserAPIKeyAuth", server_id: str) -> bool: @@ -358,7 +370,7 @@ async def _resolve_jwt_auth( request: Request, token: str, write_route: str | None, -) -> "UserAPIKeyAuth | None": +) -> "UserAPIKeyAuth | JWTIdentity | None": from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.handle_jwt import JWTAuthManager # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle @@ -393,25 +405,35 @@ async def _resolve_jwt_auth( return None if await _key_owner_scim_deactivated(mapped) or not _active_key_user_id(mapped) else mapped if mapped is not None: return None - identity: Final = await JWTAuthManager.auth_builder( + if write_route is None: + identity: Final = await JWTAuthManager.resolve_identity( + api_key=token, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + if identity.user_object is not None and isinstance(_active_user_record(identity.user_object), str): + return None + return identity + authorized: Final = await JWTAuthManager.authorize_jwt( api_key=token, jwt_handler=jwt_handler, request_data={}, general_settings=general_settings, - route=write_route or request.url.path, + route=write_route, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, proxy_logging_obj=proxy_logging_obj, request_headers=dict(request.headers), request_method=request.method, - identity_only=write_route is None, - allow_provisioning=False, ) - resolved_user: Final = identity["user_object"] + resolved_user: Final = authorized["user_object"] if resolved_user is not None and isinstance(_active_user_record(resolved_user), str): return None - return JWTAuthManager.user_api_key_auth_from_result(identity) + return JWTAuthManager.user_api_key_auth_from_result(authorized) except Exception as exc: # noqa: BLE001 # public OAuth exchange stays available; unvalidated identities never write credentials verbose_logger.debug("OAuth JWT identity could not be validated (%s)", type(exc).__name__) return None diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 9ba67f966a2..7e7189c8a6b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _prepare_bridge_mint, _prepare_bridge_refresh, _reload_active_user_by_id, + authorize_oauth_credential_request, can_store_oauth_credential, ) from litellm.proxy._experimental.mcp_server.faults import ( @@ -851,7 +852,7 @@ async def _resolve_oauth_authorization_user( ) request_user_id: Final = ( - await _extract_user_id_from_request(request, mcp_server.server_id) if enforce_binding else None + await authorize_oauth_credential_request(request, mcp_server.server_id) if enforce_binding else None ) if enforce_binding and request_user_id is None and _litellm_key_from_request(request): return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) @@ -1233,7 +1234,7 @@ async def exchange_token_with_server( request, await MCPRequestHandler.reload_admitted_user(user_id), resolved_server.server_id ) if bridge_identity is not None - else await _extract_user_id_from_request(request, resolved_server.server_id) == user_id + else await authorize_oauth_credential_request(request, resolved_server.server_id) == user_id ) if can_store: await _store_per_user_token_server_side( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index f2bdbdd9341..4fe44eb1dc8 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -9,13 +9,13 @@ JWT token must have 'litellm_proxy_admin' in scope. from __future__ import annotations import asyncio -import copy import fnmatch import hashlib import os import re import time from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast import httpx @@ -130,6 +130,19 @@ class _UserInfoResponse(Protocol): def json(self) -> dict[str, object]: ... +@dataclass(frozen=True, slots=True) +class JWTIdentity: + user_id: str | None + user_object: LiteLLM_UserTable | None + agent_id: str | None + + +@dataclass(frozen=True, slots=True) +class _JWTProvisioning: + user_id_upsert: bool + team_id_upsert: bool + + class AgentLookup(Protocol): """The registered-agent lookups a JWT agent claim is matched against.""" @@ -1473,6 +1486,7 @@ class JWTAuthManager: user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, + team_id_upsert: bool | None = None, ) -> tuple[str | None, LiteLLM_TeamTable | None]: """Find and validate specific team ID from team_id_jwt_field or team_alias_jwt_field""" individual_team_id = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None) @@ -1500,7 +1514,9 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert + if team_id_upsert is None + else team_id_upsert, ) return individual_team_id, team_object except HTTPException as e: @@ -2017,6 +2033,7 @@ class JWTAuthManager: user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, + team_id_upsert: bool | None = None, ) -> None: """Attach team context from x-litellm-team-id to an admin result. @@ -2034,7 +2051,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert if team_id_upsert is None else team_id_upsert, ) except Exception as e: # Fall back to pre-PR admin behavior: honor the admin's @@ -2268,64 +2285,119 @@ class JWTAuthManager: proxy_logging_obj: ProxyLogging, request_headers: dict | None = None, request_method: str | None = None, - identity_only: bool = False, - allow_provisioning: bool = True, ) -> JWTAuthBuilderResult: - """Build JWT authentication and authorization context. + return await JWTAuthManager.authorize_jwt( + api_key=api_key, + jwt_handler=jwt_handler, + request_data=request_data, + general_settings=general_settings, + route=route, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + request_headers=request_headers, + request_method=request_method, + provisioning=_JWTProvisioning( + user_id_upsert=jwt_handler.litellm_jwtauth.user_id_upsert, + team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + ), + ) - identity_only resolves the caller for OAuth identity binding and grants no permission. - Credential writes use full authorization with allow_provisioning=False: resolve the - existing policy context without creating users/teams or synchronizing membership. - A private handler configuration keeps that restriction out of concurrent normal requests. - """ - handler: Final = jwt_handler if allow_provisioning else copy.copy(jwt_handler) - if not allow_provisioning: - handler.update_environment( - prisma_client=jwt_handler.prisma_client, - user_api_key_cache=jwt_handler.user_api_key_cache, - litellm_jwtauth=jwt_handler.litellm_jwtauth.model_copy( - update={"user_id_upsert": False, "team_id_upsert": False, "sync_user_role_and_teams": False} - ), - leeway=jwt_handler.leeway, + @staticmethod + async def authenticate_jwt(api_key: str, jwt_handler: JWTHandler) -> dict[str, object]: + claims: Final = ( + await jwt_handler.get_oidc_userinfo(token=api_key) + if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not jwt_handler.is_jwt(token=api_key) + else await jwt_handler.auth_jwt(token=api_key) + ) + validate: Final = jwt_handler.litellm_jwtauth.custom_validate + if validate is not None and not validate(claims): + raise HTTPException(status_code=403, detail="Invalid JWT token") + return claims + + @staticmethod + async def resolve_identity( + api_key: str, + jwt_handler: JWTHandler, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, + ) -> JWTIdentity: + claims: Final = await JWTAuthManager.authenticate_jwt(api_key, jwt_handler) + return await JWTAuthManager._resolve_claim_identity( + claims, jwt_handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj + ) + + @staticmethod + async def _resolve_claim_identity( + claims: dict[str, object], + jwt_handler: JWTHandler, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, + ) -> JWTIdentity: + claim_user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(jwt_handler, claims) + user_id: Final = ( + jwt_handler.get_object_id(token=claims, default_value=None) or claim_user_id + if jwt_handler.get_rbac_role(token=claims) == LitellmUserRoles.INTERNAL_USER + else claim_user_id + ) + agent_id: Final = JWTAuthManager.resolve_agent_id(jwt_handler, claims, jwt_handler.agent_lookup) + is_admin: Final = jwt_handler.is_admin(scopes=jwt_handler.get_scopes(token=claims)) + try: + user, _, _, _, canonical_id = await JWTAuthManager.get_objects( + user_id=user_id, + user_email=user_email, + org_id=None, + end_user_id=None, + team_id=None, + valid_user_email=valid_user_email, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + route="", + user_id_upsert=False, ) + except UserNotFoundError: + if not is_admin: + raise + return JWTIdentity(user_id=user_id, user_object=None, agent_id=agent_id) + return JWTIdentity(user_id=user_id if is_admin else canonical_id, user_object=user, agent_id=agent_id) - # Check if OIDC UserInfo endpoint is enabled, but fall back to standard - # JWT auth if the token itself is a well-formed JWT (3-part structure). - if handler.litellm_jwtauth.oidc_userinfo_enabled and not handler.is_jwt(token=api_key): - verbose_proxy_logger.debug("OIDC UserInfo is enabled. Fetching user info from UserInfo endpoint.") - # Use the access token to fetch user info from OIDC UserInfo endpoint - jwt_valid_token: dict = await handler.get_oidc_userinfo(token=api_key) - else: - # Default behavior: decode and validate the JWT token - jwt_valid_token = await handler.auth_jwt(token=api_key) - - # Check custom validate - if handler.litellm_jwtauth.custom_validate: - if not handler.litellm_jwtauth.custom_validate(jwt_valid_token): - raise HTTPException( - status_code=403, - detail="Invalid JWT token", - ) + @staticmethod + async def authorize_jwt( + api_key: str, + jwt_handler: JWTHandler, + request_data: dict[str, object], + general_settings: dict[str, object], + route: str, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, + request_headers: dict[str, str] | None = None, + request_method: str | None = None, + provisioning: _JWTProvisioning | None = None, + ) -> JWTAuthBuilderResult: + """Resolve and authorize JWT context; only normal admission supplies provisioning.""" + handler: Final = jwt_handler + jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler) + team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False + model: Final = request_data.get("model") + requested_model: Final = model if isinstance(model, str) else None # Check RBAC rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) - if not identity_only: - await JWTAuthManager.check_rbac_role( - handler, - jwt_valid_token, - general_settings, - request_data, - route, - rbac_role, - ) + await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role) # Check Scope Based Access scopes: Final = handler.get_scopes(token=jwt_valid_token) - if ( - not identity_only - and handler.litellm_jwtauth.enforce_scope_based_access - and handler.litellm_jwtauth.scope_mappings - ): + if handler.litellm_jwtauth.enforce_scope_based_access and handler.litellm_jwtauth.scope_mappings: JWTAuthManager.check_scope_based_access( scope_mappings=handler.litellm_jwtauth.scope_mappings, scopes=scopes, @@ -2357,69 +2429,6 @@ class JWTAuthManager: agent_registry=handler.agent_lookup, ) - if identity_only or (not allow_provisioning and handler.is_admin(scopes=scopes)): - try: - identity_user, _, _, _, identity_user_id = await JWTAuthManager.get_objects( - user_id=user_id, - user_email=user_email, - org_id=None, - end_user_id=None, - team_id=None, - valid_user_email=valid_user_email, - jwt_handler=handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - route=route, - user_id_upsert=False, - ) - except UserNotFoundError: - if not handler.is_admin(scopes=scopes): - raise - identity_user, identity_user_id = None, user_id - if not identity_only: - admin: Final = await JWTAuthManager.check_admin_access( - handler, - scopes, - route, - user_id, - org_id, - api_key, - jwt_valid_token, - user_email=user_email, - agent_id=agent_id, - ) - if admin is not None: - await JWTAuthManager._attach_team_from_header_for_admin( - admin_result=admin, - route=route, - request_headers=request_headers, - jwt_handler=handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - return {**admin, "user_object": identity_user} - return JWTAuthBuilderResult( - is_proxy_admin=False, - # Admin admission uses the claim ID; other callers use the canonical DB ID. - user_id=user_id if handler.is_admin(scopes=scopes) else identity_user_id, - user_email=identity_user.user_email if identity_user is not None else user_email, - user_object=identity_user, - team_id=None, - team_object=None, - org_id=None, - org_object=None, - end_user_id=None, - end_user_object=None, - team_membership=None, - token=api_key, - jwt_claims=jwt_valid_token, - agent_id=agent_id, - ) - # Check admin access admin_result: Final = await JWTAuthManager.check_admin_access( handler, @@ -2442,7 +2451,13 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, + team_id_upsert=team_id_upsert, ) + if provisioning is None: + identity: Final = await JWTAuthManager._resolve_claim_identity( + jwt_valid_token, handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj + ) + return {**admin_result, "user_object": identity.user_object} return admin_result # Get team with model access @@ -2485,7 +2500,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=(handler.litellm_jwtauth.team_id_upsert and not db_team_fallback), + team_id_upsert=(team_id_upsert and not db_team_fallback), ) except HTTPException: if not db_team_fallback: @@ -2503,13 +2518,14 @@ class JWTAuthManager: user_api_key_cache, parent_otel_span, proxy_logging_obj, + team_id_upsert=team_id_upsert, ) if not team_object and not team_id: ## CHECK USER GROUP ACCESS team_id, team_object = await JWTAuthManager.find_team_with_model_access( team_ids=all_team_ids, - requested_model=request_data.get("model"), + requested_model=requested_model, route=route, request_method=request_method, jwt_handler=handler, @@ -2536,7 +2552,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=team_id_upsert, ) if team_id and not JWTAuthManager._team_has_passthrough_route_access( @@ -2570,18 +2586,20 @@ class JWTAuthManager: proxy_logging_obj=proxy_logging_obj, route=route, org_alias=org_alias, + user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False, ) # Derive org_id from org_object if resolved by alias resolved_org_id: Final = org_object.organization_id if org_object else org_id - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler=handler, - jwt_valid_token=jwt_valid_token, - user_object=user_object, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - ) + if provisioning is not None: + await JWTAuthManager.sync_user_role_and_teams( + jwt_handler=handler, + jwt_valid_token=jwt_valid_token, + user_object=user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) # If JWT did not resolve team_id, attempt a team fallback. if team_id is None and db_team_fallback: @@ -2592,11 +2610,11 @@ class JWTAuthManager: ) = await JWTAuthManager._resolve_db_team_fallback( user_object=user_object, user_id=user_id, - requested_model=request_data.get("model"), + requested_model=requested_model, route=route, jwt_handler=handler, enforce_team_based_model_access=handler.litellm_jwtauth.enforce_team_based_model_access, - team_id_upsert=handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=team_id_upsert, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -2624,7 +2642,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=team_id_upsert, ) elif db_team_fallback and team_id == header_team_id: JWTAuthManager._validate_header_team_in_db_membership( @@ -2644,7 +2662,7 @@ class JWTAuthManager: ) ## MAP USER TO TEAMS - if allow_provisioning: + if provisioning is not None: await JWTAuthManager.map_user_to_teams( user_object=user_object, team_object=team_object, @@ -2654,7 +2672,7 @@ class JWTAuthManager: JWTAuthManager.validate_object_id( user_id=user_id, team_id=team_id, - enforce_rbac=general_settings.get("enforce_rbac", False), + enforce_rbac=bool(general_settings.get("enforce_rbac", False)), is_proxy_admin=False, ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 1739ac6d743..8d300c7c508 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -6980,6 +6980,11 @@ async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool: new_callable=AsyncMock, return_value="admin-user", ), + patch( # test-quality-ok: this control tests persistence by auth mode; write-policy behavior is covered separately + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.authorize_oauth_credential_request", + new_callable=AsyncMock, + return_value="admin-user", + ), patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._store_per_user_token_server_side", new_callable=AsyncMock, @@ -11168,7 +11173,7 @@ async def test_identity_bound_authorization_carries_nonce_and_caller_through_cal "path": "/authorize", "query_string": b"", "headers": []}) with ( patch( # test-quality-ok: isolate authenticated request resolution from the real encrypted OAuth round trip - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.authorize_oauth_credential_request", new=AsyncMock(return_value="alice")), patch( # test-quality-ok: isolate user access lookup while testing nonce and caller preservation "litellm.proxy._experimental.mcp_server.discoverable_endpoints._user_can_reach_mcp_server", @@ -11569,17 +11574,24 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( "missing_database", ], ) +@pytest.mark.parametrize("credential_write", [False, True]) async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch, rejection: str, + credential_write: bool, ) -> None: from cryptography.hazmat.primitives.asymmetric import rsa from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( + _extract_user_id_from_request, authorize_oauth_credential_request, + ) + allowed_servers: Final = AsyncMock(return_value=["server-a"]) + monkeypatch.setattr(mcp_server_manager.global_mcp_server_manager, "get_allowed_mcp_servers", allowed_servers) handler, signing_key = jwt_oauth_identity key: Final = ( rsa.generate_private_key(public_exponent=65537, key_size=2048) if rejection == "signature" else signing_key @@ -11603,7 +11615,13 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( ) if rejection == "custom_validate": handler.litellm_jwtauth.custom_validate = lambda claims: False - assert await _extract_user_id_from_request(_token_request({"Authorization": f"Bearer {bearer}"})) is None + request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) + result: Final = ( + await authorize_oauth_credential_request(request, "server-a") + if credential_write else await _extract_user_id_from_request(request) + ) + assert result is None + allowed_servers.assert_not_awaited() @pytest.mark.asyncio @@ -11850,7 +11868,7 @@ async def test_oauth_credential_write_keeps_virtual_key_permissions( import asyncio from litellm.proxy._experimental.mcp_server import mcp_server_manager - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import authorize_oauth_credential_request from litellm.proxy._types import UserAPIKeyAuth, hash_token from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key @@ -11881,10 +11899,10 @@ async def test_oauth_credential_write_keeps_virtual_key_permissions( request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/server-a/token") if state == "cancelled": with pytest.raises(asyncio.CancelledError): - await _extract_user_id_from_request(request, "server-a") + await authorize_oauth_credential_request(request, "server-a") manager.get_allowed_mcp_servers.assert_awaited_once() return - assert await _extract_user_id_from_request(request, "server-a") == ("jwt-owner" if state == "allowed" else None) + assert await authorize_oauth_credential_request(request, "server-a") == ("jwt-owner" if state == "allowed" else None) if state in ("allowed", "server_denied", "lookup_error"): manager.get_allowed_mcp_servers.assert_awaited_once() writer: Final = manager.get_allowed_mcp_servers.call_args.args[0] @@ -11907,7 +11925,7 @@ async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_ros from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server import mcp_server_manager - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import authorize_oauth_credential_request from litellm.proxy._types import LiteLLM_TeamTable, Member handler, signing_key = jwt_oauth_identity @@ -11925,7 +11943,7 @@ async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_ros monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) bearer: Final = _oauth_identity_jwt(signing_key, claims={"team": "a"}) request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path=f"/{server_id}/token") - assert await _extract_user_id_from_request(request, server_id) == ( + assert await authorize_oauth_credential_request(request, server_id) == ( "jwt-owner" if server_id == "team-a-server" else None ) manager.get_allowed_mcp_servers.assert_awaited_once() diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index babbf88dc29..6fbbb37fbc3 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -6791,12 +6791,11 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla @pytest.mark.asyncio -@pytest.mark.parametrize("identity_only", [False, True]) -@pytest.mark.parametrize("allow_provisioning", [False, True]) +@pytest.mark.parametrize("operation", ["identity", "authorize", "admit"]) @pytest.mark.parametrize("existing_user", [False, True]) @pytest.mark.parametrize("model_allowed", [False, True]) -async def test_auth_builder_identity_lookup_does_not_provision_users( - monkeypatch: pytest.MonkeyPatch, identity_only: bool, allow_provisioning: bool, existing_user: bool, model_allowed: bool +async def test_jwt_identity_and_authorization_keep_provisioning_in_admission( + monkeypatch: pytest.MonkeyPatch, operation: str, existing_user: bool, model_allowed: bool ) -> None: from litellm.proxy._types import ScopeMapping from litellm.proxy.auth.auth_checks import UserNotFoundError @@ -6805,7 +6804,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( private_key, jwk = _get_rsa_key_and_jwk("identity-mode") cache: Final = UserApiKeyCache() cache.set_cache("litellm_jwt_auth_keys_https://identity.example/jwks", [jwk]) - user_id: Final = f"identity-mode-{identity_only}-{existing_user}-{model_allowed}" + user_id: Final = f"identity-mode-{operation}-{existing_user}-{model_allowed}" user: Final = LiteLLM_UserTable(user_id=user_id, organization_memberships=[]) if existing_user: cache.set_cache(user_id, user) @@ -6831,26 +6830,38 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( token: Final = _encode_rsa_jwt( private_key, "https://identity.example", "gateway", "identity-mode", {"sub": user_id, "scope": "allowed"} ) - pending: Final = JWTAuthManager.auth_builder( - api_key=token, - jwt_handler=handler, + common: Final = { + "api_key": token, + "jwt_handler": handler, + "prisma_client": database, + "user_api_key_cache": cache, + "parent_otel_span": None, + "proxy_logging_obj": MagicMock(), + } + if operation == "identity": + if not existing_user: + with pytest.raises(UserNotFoundError): + await JWTAuthManager.resolve_identity(**common) + else: + identity: Final = await JWTAuthManager.resolve_identity(**common) + assert identity.user_id == user_id + assert identity.user_object is not None and identity.user_object.user_id == user_id + users.create.assert_not_awaited() + return + authorize: Final = JWTAuthManager.auth_builder if operation == "admit" else JWTAuthManager.authorize_jwt + pending: Final = authorize( + **common, request_data={"model": "allowed-model" if model_allowed else "forbidden-model"}, general_settings={}, - route="/example/token" if identity_only else "/mcp/example", - prisma_client=database, - user_api_key_cache=cache, - parent_otel_span=None, - proxy_logging_obj=MagicMock(), - identity_only=identity_only, - allow_provisioning=allow_provisioning, + route="/mcp/example", ) - if not identity_only and not model_allowed: + if not model_allowed: with pytest.raises(HTTPException) as denial: await pending assert denial.value.status_code == 403 users.create.assert_not_awaited() return - if (identity_only or not allow_provisioning) and not existing_user: + if operation == "authorize" and not existing_user: with pytest.raises(UserNotFoundError): await pending else: @@ -6858,7 +6869,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( assert result["user_id"] == user_id assert result["user_object"] is not None assert result["user_object"].user_id == user_id - assert users.create.await_count == (0 if identity_only or not allow_provisioning or existing_user else 1) + assert users.create.await_count == (0 if operation == "authorize" or existing_user else 1) def _entra_agent_registry() -> AgentRegistry: @@ -6997,6 +7008,14 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a ) jwt_handler.bind_agent_lookup(_entra_agent_registry()) + if identity_only: + identity = await JWTAuthManager.resolve_identity( + api_key=token, jwt_handler=jwt_handler, prisma_client=None, + user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + ) + assert identity.agent_id == "canonical-agent-id" + return + result = await JWTAuthManager.auth_builder( api_key=token, jwt_handler=jwt_handler, @@ -7007,10 +7026,9 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, - identity_only=identity_only, ) - assert result["is_proxy_admin"] is (is_admin_token and not identity_only) + assert result["is_proxy_admin"] is is_admin_token assert result["agent_id"] == "canonical-agent-id" @@ -7025,6 +7043,14 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch ) jwt_handler.bind_agent_lookup(_entra_agent_registry()) + if identity_only: + with pytest.raises(HTTPException) as denial: + await JWTAuthManager.resolve_identity( + api_key=token, jwt_handler=jwt_handler, prisma_client=None, + user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + ) + assert denial.value.status_code == 403 + return with pytest.raises(HTTPException) as exc_info: await JWTAuthManager.auth_builder( api_key=token, @@ -7036,7 +7062,39 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, - identity_only=identity_only, ) assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("admission", [False, True]) +async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatch, admission: bool): + from litellm.proxy.management_endpoints import team_endpoints + + handler, token = _entra_signed_app_token( + monkeypatch, azp="canonical-agent-id", scope=LiteLLM_JWTAuth().admin_jwt_scope, + ) + handler.bind_agent_lookup(_entra_agent_registry()) + handler.litellm_jwtauth.team_id_upsert = True + handler.litellm_jwtauth.admin_allowed_routes = ["openai_routes"] + database = MagicMock() + database.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + create_team = AsyncMock(return_value=LiteLLM_TeamTable(team_id="new-team").model_dump()) + monkeypatch.setattr(team_endpoints, "new_team", create_team) + resolve = JWTAuthManager.auth_builder if admission else JWTAuthManager.authorize_jwt + + result = await resolve( + api_key=token, jwt_handler=handler, request_data={}, general_settings={}, + route="/chat/completions", prisma_client=database, + user_api_key_cache=handler.user_api_key_cache, parent_otel_span=None, + proxy_logging_obj=MagicMock(), request_headers={"x-litellm-team-id": "new-team"}, + ) + + assert result["is_proxy_admin"] is True + if admission: + create_team.assert_awaited_once() + assert result["team_id"] == "new-team" + else: + create_team.assert_not_awaited() + assert result["team_id"] is None