Merge pull request #41314 from BerriAI/litellm_fix_mcp_jwt_oauth_persistence

fix(mcp): authorize JWT OAuth credential persistence
This commit is contained in:
joshua-berri 2026-09-16 06:47:40 -07:00 committed by GitHub
commit 9cd787386e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 1505 additions and 161 deletions

View file

@ -1,6 +1,8 @@
"""Bridge token flow: litellm identity resolution and the DCR-bridge oauth_delegate mint/refresh pipeline."""
import math
import os
import secrets
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Final, Literal
@ -12,6 +14,9 @@ from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
_V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
@ -24,6 +29,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:
@ -48,6 +54,64 @@ def _litellm_key_from_request(request: Request) -> str | None:
return None
async def oauth_authorization_uses_gateway_credential(request: Request) -> bool:
"""Classify credentials for browser authorize; candidates still require full authorization."""
from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # startup owns the active auth configuration
jwt_handler,
master_key,
user_custom_auth,
)
if "x-litellm-api-key" in request.headers:
return True
token: Final = _litellm_key_from_request(request)
if token is None:
return "authorization" in request.headers
if token.startswith("sk-") or (master_key and secrets.compare_digest(token.encode(), master_key.encode())):
return True
if user_custom_auth is not None or jwt_handler.litellm_jwtauth.oidc_userinfo_enabled:
return True
if not JWTHandler.is_jwt(token):
return await _opaque_bearer_is_gateway_credential(token)
claims: Final = JWTHandler.get_unverified_claims(token)
issuer: Final = claims.get("iss") if claims is not None else None
global_issuer: Final = os.getenv("JWT_ISSUER")
# An unscoped global validator can accept issuers absent from the configured issuer list.
if not isinstance(issuer, str) or not issuer or not global_issuer:
return True
return issuer == global_issuer or any(
issuer == configured.issuer for configured in jwt_handler.litellm_jwtauth.issuers or ()
)
async def _opaque_bearer_is_gateway_credential(token: str) -> bool:
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
is_envelope, # noqa: PLC0415 # envelope imports bridge types
is_refresh_envelope,
)
from litellm.proxy._types import hash_token # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.resolvers.exceptions import KeyNotFoundError # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.resolvers.store import IdentityStore # noqa: PLC0415 # proxy import cycle
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # startup owns the identity store dependencies
prisma_client,
user_api_key_cache,
)
if is_envelope(token) or is_refresh_envelope(token) or token.startswith(_V2_GCM_PREFIX):
return True
try:
if ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) is not None:
return True
await IdentityStore(prisma_client, user_api_key_cache).resolve(hashed_token=hash_token(token))
except KeyNotFoundError:
return False
except Exception as exc: # noqa: BLE001 # an identity lookup fault must not permit cookie fallback
verbose_logger.debug("OAuth bearer ownership could not be checked (%s)", type(exc).__name__)
return True
def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool:
"""``True`` when the presented key is neither blocked nor past its expiry.
@ -243,6 +307,10 @@ async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResol
return "no_active_key"
if user_object is None:
return "no_active_key"
return _active_user_record(user_object)
def _active_user_record(user_object: "LiteLLM_UserTable") -> "LiteLLM_UserTable | Literal['no_active_key']":
if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False:
return "no_active_key"
return user_object
@ -301,15 +369,137 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol
async def _extract_user_id_from_request(request: Request) -> str | None:
"""The litellm ``user_id`` for the token request, so a per-user token is stored under the same
identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome
(including a transient DB outage) collapses to ``None`` here and the caller simply skips the store;
the bridge mint, which must status those outcomes differently, consumes
:func:`_resolve_active_litellm_key` directly."""
resolved: Final = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
"""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)
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:
"""Apply the same write policy to request credentials and verified signed-callback users."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # registry imports auth helpers
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
can_access_mcp_server, # noqa: PLC0415 # proxy import cycle
)
from litellm.proxy.auth.route_checks import RouteChecks # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle
_run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action
)
write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential"
try:
RouteChecks.is_virtual_key_allowed_to_call_route(route=write_route, valid_token=auth, request=request)
await _run_centralized_common_checks(
user_api_key_auth_obj=auth,
request=request,
request_data={},
route=write_route,
)
return await can_access_mcp_server(auth, server_id, global_mcp_server_manager.get_allowed_mcp_servers)
except Exception as exc: # noqa: BLE001 # authorization failure must never write credentials
verbose_logger.debug("OAuth credential write not authorized (%s)", type(exc).__name__)
return False
async def _resolve_jwt_auth(
request: Request,
token: str,
write_route: str | 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
_resolve_jwt_to_virtual_key, # pyright: ignore[reportPrivateUsage] # reuse admission mapping policy without provisioning a new key
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # proxy globals initialized at startup
general_settings,
jwt_handler,
premium_user,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if general_settings.get("enable_jwt_auth") is not True or premium_user is not True or prisma_client is None:
return None
try:
if jwt_handler.litellm_jwtauth.is_virtual_key_mapping_configured():
claims: Final = await jwt_handler.auth_jwt(token=token)
validate: Final = jwt_handler.litellm_jwtauth.custom_validate
if validate is not None and not validate(claims):
return None
mapped: Final = await _resolve_jwt_to_virtual_key(
jwt_claims=claims,
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 isinstance(mapped, UserAPIKeyAuth):
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
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,
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,
)
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(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
return _active_key_user_id(resolved.key)
_UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"]

View file

@ -32,6 +32,9 @@ 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,
oauth_authorization_uses_gateway_credential,
)
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
@ -836,16 +839,30 @@ async def _user_can_reach_mcp_server(user_id: str, server_id: str) -> bool:
return server_id in await global_mcp_server_manager.get_allowed_mcp_servers(admitted)
async def _bridge_authorize_access_denial(
litellm_user_id: str,
async def _resolve_oauth_authorization_user(
request: Request,
mcp_server: MCPServer,
redirect_uri: str,
state: str,
) -> RedirectResponse | None:
"""The denial redirect for a signed-in user who cannot reach the target server, or None to proceed."""
if await _user_can_reach_mcp_server(litellm_user_id, mcp_server.server_id):
return None
return _bridge_access_denied_redirect(redirect_uri, state, mcp_server)
enforce_binding: bool,
) -> str | RedirectResponse:
"""Resolve the authorization subject without replacing denied credentials with cookie grants."""
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # proxy import cycle
_user_id_from_session_cookie,
)
use_gateway_credential: Final = enforce_binding and await oauth_authorization_uses_gateway_credential(request)
request_user_id: Final = (
await authorize_oauth_credential_request(request, mcp_server.server_id) if use_gateway_credential else None
)
if use_gateway_credential and request_user_id is None:
return _bridge_access_denied_redirect(redirect_uri, state, mcp_server)
user_id: Final = request_user_id or _user_id_from_session_cookie(request)
if user_id is None:
return _redirect_to_litellm_login(request)
if not await _user_can_reach_mcp_server(user_id, mcp_server.server_id):
return _bridge_access_denied_redirect(redirect_uri, state, mcp_server)
return user_id
async def authorize_with_server(
@ -911,23 +928,12 @@ async def authorize_with_server(
# Seal the authenticated caller into state so the token exchange cannot select another credential owner.
litellm_user_id: str | None = None
if enforce_binding or (resolved_server.is_dcr_bridge and resolved_server.is_oauth_delegate):
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # inline import avoids a module-load circular import
_user_id_from_session_cookie,
subject: Final = await _resolve_oauth_authorization_user(
request, resolved_server, redirect_uri, state, enforce_binding
)
litellm_user_id = (
await _extract_user_id_from_request(request) if enforce_binding else None
) or _user_id_from_session_cookie(request)
if litellm_user_id is None:
return _redirect_to_litellm_login(request)
denial: Final = await _bridge_authorize_access_denial(
litellm_user_id=litellm_user_id,
mcp_server=resolved_server,
redirect_uri=redirect_uri,
state=state,
)
if denial is not None:
return denial
if isinstance(subject, RedirectResponse):
return subject
litellm_user_id = subject
oauth_nonce: Final = secrets.token_urlsafe(32) if enforce_binding else None
encoded_state: Final = encode_state_with_base_url(
@ -1218,12 +1224,32 @@ async def exchange_token_with_server(
user_id: Final = resolved_user_id
if user_id:
try:
await _store_per_user_token_server_side(
server=resolved_server,
user_id=user_id,
token_response=token_response,
identity_binding_proof=binding_proof,
# Identity binding above must retain the verified caller even when a write is
# denied. Authorize persistence separately, immediately before its side effect.
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
# A sealed code delegates a verified user for this authorized server. Raw
# request credentials retain their own JWT/key restrictions during resolution.
can_store: Final = (
await can_store_oauth_credential(
request, await MCPRequestHandler.reload_admitted_user(user_id), resolved_server.server_id
)
if bridge_identity is not None
else await authorize_oauth_credential_request(request, resolved_server.server_id) == user_id
)
if can_store:
await _store_per_user_token_server_side(
server=resolved_server,
user_id=user_id,
token_response=token_response,
identity_binding_proof=binding_proof,
)
else:
verbose_logger.warning(
"OAuth credential storage not authorized for user=%s server=%s",
user_id,
resolved_server.server_id,
)
except Exception as exc:
verbose_logger.warning(
"exchange_token_with_server: server-side storage failed for user=%s server=%s: %s",
@ -1236,8 +1262,9 @@ async def exchange_token_with_server(
"exchange_token_with_server: could not resolve a LiteLLM user_id for the request, "
"so the per-user token for server=%s was NOT stored. The authorization_code egress "
"requires the stored token, so the client will be challenged with 401 on reconnect. "
"Ensure the request carries a valid LiteLLM key (x-litellm-api-key or Authorization), "
"or store it via POST /mcp/server/{id}/oauth-user-credential.",
"Ensure the request carries a valid LiteLLM key or enabled JWT identity "
"(x-litellm-api-key or Authorization), "
"or store it via POST /v1/mcp/server/{id}/oauth-user-credential.",
resolved_server.server_id,
)

View file

@ -2,6 +2,7 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import Final
from fastapi import HTTPException
@ -137,3 +138,15 @@ async def build_effective_auth_contexts(
if admitted_context is None:
return team_contexts
return [*team_contexts, admitted_context]
async def can_access_mcp_server(
user_api_key_auth: UserAPIKeyAuth,
server_id: str,
allowed_servers: Callable[[UserAPIKeyAuth], Awaitable[list[str]]],
) -> bool:
"""Resolve server access through the same credential contexts as MCP management."""
for context in await build_effective_auth_contexts(user_api_key_auth):
if server_id in await allowed_servers(context):
return True
return False

View file

@ -15,6 +15,7 @@ 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
@ -54,7 +55,7 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_model_aliases
from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
get_management_object_ttl,
@ -62,6 +63,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.user_repository import UserRepository
from litellm.types.agents import AgentResponse
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
from .auth_checks import (
_allowed_routes_check,
@ -128,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."""
@ -1471,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)
@ -1498,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:
@ -1726,6 +1744,7 @@ class JWTAuthManager:
proxy_logging_obj: ProxyLogging,
route: str,
org_alias: str | None = None,
user_id_upsert: bool | None = None,
) -> tuple[
LiteLLM_UserTable | None,
LiteLLM_OrganizationTable | None,
@ -1789,7 +1808,11 @@ class JWTAuthManager:
user_id=user_id,
user_email=user_email,
sso_user_id=user_id,
upsert=jwt_handler.is_upsert_user_id(valid_user_email=valid_user_email),
upsert=(
jwt_handler.is_upsert_user_id(valid_user_email=valid_user_email)
if user_id_upsert is None
else user_id_upsert
),
),
team_id=team_id,
)
@ -2010,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.
@ -2027,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
@ -2262,57 +2286,136 @@ class JWTAuthManager:
request_headers: dict | None = None,
request_method: str | None = None,
) -> JWTAuthBuilderResult:
"""Main authentication and authorization builder"""
# 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 jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not jwt_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 jwt_handler.get_oidc_userinfo(token=api_key)
else:
# Default behavior: decode and validate the JWT token
jwt_valid_token = await jwt_handler.auth_jwt(token=api_key)
# Check custom validate
if jwt_handler.litellm_jwtauth.custom_validate:
if not jwt_handler.litellm_jwtauth.custom_validate(jwt_valid_token):
raise HTTPException(
status_code=403,
detail="Invalid JWT token",
)
# Check RBAC
rbac_role: Final = jwt_handler.get_rbac_role(token=jwt_valid_token)
await JWTAuthManager.check_rbac_role(
jwt_handler,
jwt_valid_token,
general_settings,
request_data,
route,
rbac_role,
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,
),
)
@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)
@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)
await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role)
# Check Scope Based Access
scopes: Final = jwt_handler.get_scopes(token=jwt_valid_token)
if jwt_handler.litellm_jwtauth.enforce_scope_based_access and jwt_handler.litellm_jwtauth.scope_mappings:
scopes: Final = handler.get_scopes(token=jwt_valid_token)
if handler.litellm_jwtauth.enforce_scope_based_access and handler.litellm_jwtauth.scope_mappings:
JWTAuthManager.check_scope_based_access(
scope_mappings=jwt_handler.litellm_jwtauth.scope_mappings,
scope_mappings=handler.litellm_jwtauth.scope_mappings,
scopes=scopes,
request_data=request_data,
general_settings=general_settings,
)
object_id = jwt_handler.get_object_id(token=jwt_valid_token, default_value=None)
object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
# Get basic user info
user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(jwt_handler, jwt_valid_token)
user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token)
# Get IDs
org_id: Final = jwt_handler.get_org_id(token=jwt_valid_token, default_value=None)
end_user_id: Final = jwt_handler.get_end_user_id(token=jwt_valid_token, default_value=None)
org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None)
end_user_id: Final = handler.get_end_user_id(token=jwt_valid_token, default_value=None)
team_id: str | None = None
team_object: LiteLLM_TeamTable | None = None
object_id = jwt_handler.get_object_id(token=jwt_valid_token, default_value=None)
object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
if rbac_role and object_id:
if rbac_role == LitellmUserRoles.TEAM:
@ -2321,14 +2424,14 @@ class JWTAuthManager:
user_id = object_id
agent_id: Final = JWTAuthManager.resolve_agent_id(
jwt_handler=jwt_handler,
jwt_handler=handler,
jwt_valid_token=jwt_valid_token,
agent_registry=jwt_handler.agent_lookup,
agent_registry=handler.agent_lookup,
)
# Check admin access
admin_result: Final = await JWTAuthManager.check_admin_access(
jwt_handler,
handler,
scopes,
route,
user_id,
@ -2343,18 +2446,24 @@ class JWTAuthManager:
admin_result=admin_result,
route=route,
request_headers=request_headers,
jwt_handler=jwt_handler,
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,
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
## Check if team_id is specified via x-litellm-team-id header
all_team_ids: Final = JWTAuthManager.get_all_team_ids(jwt_handler, jwt_valid_token)
specific_team_id: Final = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None)
all_team_ids: Final = JWTAuthManager.get_all_team_ids(handler, jwt_valid_token)
specific_team_id: Final = handler.get_team_id(token=jwt_valid_token, default_value=None)
# The DB fallback only applies when the token carries no team identity at
# all. `get_all_jwt_team_ids` ignores `team_id_default` so a configured
@ -2364,9 +2473,9 @@ class JWTAuthManager:
# the RBAC team-role path (which already set `team_id`); otherwise a
# provisional x-litellm-team-id header could override an RBAC-asserted team.
db_team_fallback: Final = (
jwt_handler.litellm_jwtauth.fallback_to_db_teams
and not jwt_handler.get_all_jwt_team_ids(token=jwt_valid_token)
and not jwt_handler.get_team_alias(token=jwt_valid_token, default_value=None)
handler.litellm_jwtauth.fallback_to_db_teams
and not handler.get_all_jwt_team_ids(token=jwt_valid_token)
and not handler.get_team_alias(token=jwt_valid_token, default_value=None)
and team_id is None
)
if specific_team_id and not db_team_fallback:
@ -2391,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=(jwt_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:
@ -2403,22 +2512,23 @@ class JWTAuthManager:
team_id,
team_object,
) = await JWTAuthManager.find_and_validate_specific_team_id(
jwt_handler,
handler,
jwt_valid_token,
prisma_client,
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=jwt_handler,
jwt_handler=handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
@ -2442,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=jwt_handler.litellm_jwtauth.team_id_upsert,
team_id_upsert=team_id_upsert,
)
if team_id and not JWTAuthManager._team_has_passthrough_route_access(
@ -2453,7 +2563,7 @@ class JWTAuthManager:
JWTAuthManager._raise_team_passthrough_route_denial(route=route)
# Extract alias fields for resolution (if configured)
org_alias: Final = jwt_handler.get_org_alias(token=jwt_valid_token, default_value=None)
org_alias: Final = handler.get_org_alias(token=jwt_valid_token, default_value=None)
# get_objects returns effective_user_id for downstream spend attribution (GH #26789).
(
@ -2469,25 +2579,27 @@ class JWTAuthManager:
end_user_id=end_user_id,
team_id=team_id,
valid_user_email=valid_user_email,
jwt_handler=jwt_handler,
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,
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=jwt_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:
@ -2498,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=jwt_handler,
enforce_team_based_model_access=jwt_handler.litellm_jwtauth.enforce_team_based_model_access,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
jwt_handler=handler,
enforce_team_based_model_access=handler.litellm_jwtauth.enforce_team_based_model_access,
team_id_upsert=team_id_upsert,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
@ -2530,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=jwt_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(
@ -2540,7 +2652,7 @@ class JWTAuthManager:
if not JWTAuthManager._is_team_route_allowed(
route=route,
request_method=request_method,
jwt_handler=jwt_handler,
jwt_handler=handler,
):
raise HTTPException(
status_code=403,
@ -2550,16 +2662,17 @@ class JWTAuthManager:
)
## MAP USER TO TEAMS
await JWTAuthManager.map_user_to_teams(
user_object=user_object,
team_object=team_object,
)
if provisioning is not None:
await JWTAuthManager.map_user_to_teams(
user_object=user_object,
team_object=team_object,
)
# Validate that a valid rbac id is returned for spend tracking
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,
)
@ -2582,3 +2695,38 @@ class JWTAuthManager:
jwt_claims=jwt_valid_token,
agent_id=agent_id,
)
@staticmethod
def user_api_key_auth_from_result(
result: JWTAuthBuilderResult,
parent_otel_span: Span | None = None,
) -> UserAPIKeyAuth:
"""Keep JWT identity and permission attribution identical across consumers."""
user: Final = result["user_object"]
admin: Final = result["is_proxy_admin"]
return UserAPIKeyAuth(
api_key=None,
user_role=(
LitellmUserRoles.PROXY_ADMIN
if admin
else LitellmUserRoles(user.user_role)
if user is not None and user.user_role is not None
else LitellmUserRoles.INTERNAL_USER
),
user_id=result["user_id"],
user_email=result["user_email"],
team_id=result["team_id"],
org_id=result["org_id"],
end_user_id=result["end_user_id"],
parent_otel_span=parent_otel_span,
jwt_claims=result["jwt_claims"],
agent_id=result.get("agent_id"),
user_tpm_limit=user.tpm_limit if user is not None and not admin else None,
user_rpm_limit=user.rpm_limit if user is not None and not admin else None,
user_model_max_budget=user.model_max_budget if user is not None and not admin else None,
**team_grants(
team_object=result["team_object"],
team_membership=result.get("team_membership"),
user_id=result["user_id"],
),
)

View file

@ -1669,13 +1669,11 @@ async def _user_api_key_auth_builder(
is_proxy_admin: Final = result["is_proxy_admin"]
team_id: Final = result["team_id"]
team_object: Final = result["team_object"]
user_id: Final = result["user_id"]
user_email: Final = result["user_email"]
user_object: Final = result["user_object"]
end_user_id = result["end_user_id"]
org_id: Final = result["org_id"]
team_membership: Final[LiteLLM_TeamMembership | None] = result.get("team_membership", None)
jwt_claims = result.get("jwt_claims", None)
agent_id: Final[str | None] = result.get("agent_id")
@ -1693,40 +1691,9 @@ async def _user_api_key_auth_builder(
value=_JWT_PROXY_ADMIN_SENTINEL,
ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl,
)
return UserAPIKeyAuth(
api_key=None,
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id=user_id,
user_email=user_email,
team_id=team_id,
org_id=org_id,
end_user_id=end_user_id,
parent_otel_span=parent_otel_span,
jwt_claims=jwt_claims,
agent_id=agent_id,
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
return JWTAuthManager.user_api_key_auth_from_result(result, parent_otel_span)
valid_token = UserAPIKeyAuth(
api_key=None,
team_id=team_id,
user_role=(
LitellmUserRoles(user_object.user_role)
if user_object is not None and user_object.user_role is not None
else LitellmUserRoles.INTERNAL_USER
),
user_id=user_id,
user_email=user_email,
org_id=org_id,
parent_otel_span=parent_otel_span,
end_user_id=end_user_id,
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
jwt_claims=jwt_claims,
agent_id=agent_id,
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
)
valid_token = JWTAuthManager.user_api_key_auth_from_result(result, parent_otel_span)
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
# JWT policy (RBAC, scope, custom_validate, email-domain)

View file

@ -170,6 +170,7 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
admitted_user_context,
build_effective_auth_contexts,
can_access_mcp_server,
is_ui_session_credential,
)
from litellm.proxy._types import (
@ -2483,10 +2484,11 @@ if MCP_AVAILABLE:
)
return server
allowed_server_ids: Final[set[str]] = set()
for auth_context in await build_effective_auth_contexts(user_api_key_dict):
allowed_server_ids.update(await global_mcp_server_manager.get_allowed_mcp_servers(auth_context))
if server is None or server.server_id not in allowed_server_ids:
if server is None or not await can_access_mcp_server(
user_api_key_dict,
server.server_id,
global_mcp_server_manager.get_allowed_mcp_servers,
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={

View file

@ -1069,6 +1069,7 @@ async def test_jwt_non_admin_team_route_access(monkeypatch):
mock_jwt_response = {
"is_proxy_admin": False,
"jwt_claims": {},
"team_id": None,
"team_object": None,
"user_id": None,

View file

@ -5,7 +5,7 @@ import json
import time
from base64 import urlsafe_b64encode
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -15,6 +15,9 @@ from litellm.types.mcp import MCPAuth
if TYPE_CHECKING:
import httpx
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -6977,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,
@ -7124,12 +7132,12 @@ async def test_build_oauth_protected_resource_response_obo_end_to_end():
global_mcp_server_manager.registry.clear()
def _token_request(headers):
def _token_request(headers, path="/token"):
"""A real Starlette request with case-insensitive headers (matches production)."""
from starlette.requests import Request
raw = [(k.lower().encode(), v.encode()) for k, v in headers.items()]
return Request({"type": "http", "method": "POST", "path": "/token", "headers": raw, "query_string": b""})
return Request({"type": "http", "method": "POST", "path": path, "headers": raw, "query_string": b""})
@pytest.fixture
@ -11162,14 +11170,14 @@ async def test_identity_bound_authorization_carries_nonce_and_caller_through_cal
),
)
request = Request({"type": "http", "scheme": "https", "server": ("proxy.example.com", 443),
"path": "/authorize", "query_string": b"", "headers": []})
"path": "/authorize", "query_string": b"", "headers": [(b"authorization", b"Bearer sk-alice")]})
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._bridge_authorize_access_denial",
new=AsyncMock(return_value=None)),
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._user_can_reach_mcp_server",
new=AsyncMock(return_value=True)),
):
authorized = await authorize_with_server(
request, server, "client", "http://127.0.0.1:6274/callback", state="client-state",
@ -11374,3 +11382,858 @@ with TestClient(app) as client:
assert responses[path]["status"] == 200, responses[path]
assert responses[path]["body"]["issuer"] == f"http://testserver/gateway/{path}"
assert responses["example/mcp"]["body"]["token_endpoint"] == "http://testserver/gateway/example/token"
@pytest.fixture
def jwt_oauth_identity(monkeypatch: pytest.MonkeyPatch) -> tuple["JWTHandler", "RSAPrivateKey"]:
import jwt
from cryptography.hazmat.primitives.asymmetric import rsa
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_JWTAuth
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
signing_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
cache: Final = UserApiKeyCache()
cache.set_cache(
"litellm_jwt_auth_keys_https://idp.example.test/jwks",
[json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(signing_key.public_key()))],
)
cache.set_cache("jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", user_email="owner@example.test"))
handler: Final = JWTHandler()
handler.update_environment(
prisma_client=None,
user_api_key_cache=cache,
litellm_jwtauth=LiteLLM_JWTAuth(user_id_jwt_field="identity.user_id"),
)
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", "https://idp.example.test/jwks")
monkeypatch.setenv("JWT_ISSUER", "https://idp.example.test")
monkeypatch.setenv("JWT_AUDIENCE", "litellm-proxy")
monkeypatch.setattr(proxy_server, "jwt_handler", handler)
monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": True})
monkeypatch.setattr(proxy_server, "premium_user", True)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
return handler, signing_key
def _oauth_identity_jwt(
signing_key: "RSAPrivateKey",
*,
expires_in: int = 300,
audience: str = "litellm-proxy",
issuer: str = "https://idp.example.test",
owner: str | None = "jwt-owner",
scope: str = "",
claims: dict[str, object] | None = None,
) -> str:
import jwt
return jwt.encode(
{
"sub": "not-the-configured-user-id",
"identity": {"user_id": owner},
"email": "owner@example.test",
"iss": issuer,
"aud": audience,
"exp": int(time.time()) + expires_in,
"scope": scope,
**(claims or {}),
},
signing_key,
algorithm="RS256",
)
@pytest.mark.asyncio
@pytest.mark.parametrize("header", ["Authorization", "x-litellm-api-key"])
@pytest.mark.parametrize("policy_allowed", [False, True])
@pytest.mark.parametrize("server_allowed", [False, True])
@pytest.mark.parametrize("admin", [False, True])
@pytest.mark.parametrize("owner_state", ["active", "missing", "inactive", "database_error"])
async def test_oauth_exchange_stores_token_for_validated_jwt_user(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
header: str,
policy_allowed: bool,
server_allowed: bool,
admin: bool,
owner_state: str,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import httpx
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.custom_validate = lambda claims: policy_allowed
from litellm.proxy._experimental.mcp_server import mcp_server_manager
manager: Final = MagicMock()
manager.get_allowed_mcp_servers = AsyncMock(return_value=["jwt-oauth-server"] if server_allowed else [])
manager.invalidate_user_oauth_token_cache = AsyncMock()
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
bearer: Final = _oauth_identity_jwt(signing_key, scope="litellm_proxy_admin" if admin else "")
request: Final = _token_request({header: f"Bearer {bearer}"}, path="/jwt-oauth-server/token")
server: Final = MCPServer(
server_id="jwt-oauth-server",
name="jwt-oauth-server",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
authorization_url="https://upstream.example.test/authorize",
token_url="https://upstream.example.test/token",
client_id="registered-client",
)
import litellm
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy import proxy_server
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.types.llms.custom_http import httpxSpecialProvider
def upstream_response(outbound: httpx.Request) -> httpx.Response:
assert outbound.url == server.token_url
assert bearer not in str(outbound.headers)
assert bearer.encode() not in outbound.content
return httpx.Response(200, json={"access_token": "upstream-token", "token_type": "Bearer"})
database: Final = MagicMock()
users: Final = database.db.litellm_usertable
users.find_unique = AsyncMock(return_value=None)
users.find_first = AsyncMock(return_value=None)
users.create = AsyncMock()
if owner_state in ("missing", "database_error"):
handler.user_api_key_cache.delete_cache("jwt-owner")
if owner_state == "database_error":
users.find_unique.side_effect = RuntimeError("database unavailable")
if owner_state == "inactive":
handler.user_api_key_cache.set_cache(
"jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", metadata={"scim_active": False})
)
table: Final = database.db.litellm_mcpusercredentials
table.find_unique = AsyncMock(return_value=None)
table.upsert = AsyncMock()
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setenv("LITELLM_SALT_KEY", "oauth-jwt-test-encryption-key")
clients: Final = LLMClientCache()
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", clients)
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream_response)) as transport:
upstream: Final = AsyncHTTPHandler()
await upstream.client.aclose()
upstream.client = transport
clients.set_cache("async_httpx_client" + httpxSpecialProvider.Oauth2Check, upstream)
response: Final = await discoverable_endpoints.exchange_token_with_server(
request=request,
mcp_server=server,
grant_type="authorization_code",
code="upstream-code",
redirect_uri="http://localhost/callback",
client_id="registered-client",
client_secret=None,
code_verifier=None,
)
assert response.status_code == 200
assert json.loads(response.body)["access_token"] == "upstream-token"
users.create.assert_not_awaited()
if (
not server_allowed
or not policy_allowed
or owner_state in ("inactive", "database_error")
or (owner_state == "missing" and not admin)
):
table.upsert.assert_not_awaited()
return
table.upsert.assert_awaited_once()
stored: Final = table.upsert.call_args.kwargs
assert stored["where"] == {"user_id_server_id": {"user_id": "jwt-owner", "server_id": server.server_id}}
credential: Final = stored["data"]["create"]["credential_b64"]
assert "upstream-token" not in credential
decoded: Final = decrypt_value_helper(credential, key="mcp_user_credential")
assert json.loads(decoded)["access_token"] == "upstream-token"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"rejection",
[
"expired",
"audience",
"issuer",
"signature",
"missing_user",
"unknown_user",
"disabled",
"not_premium",
"scim_inactive",
"custom_validate",
"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 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
)
bearer: Final = _oauth_identity_jwt(
key,
expires_in=-60 if rejection == "expired" else 300,
audience="upstream-only" if rejection == "audience" else "litellm-proxy",
issuer="https://untrusted.example.test" if rejection == "issuer" else "https://idp.example.test",
owner=None if rejection == "missing_user" else "unknown" if rejection == "unknown_user" else "jwt-owner",
)
if rejection == "disabled":
monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": False})
if rejection == "not_premium":
monkeypatch.setattr(proxy_server, "premium_user", False)
if rejection == "missing_database":
monkeypatch.setattr(proxy_server, "prisma_client", None)
if rejection == "scim_inactive":
handler.user_api_key_cache.set_cache(
"jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", metadata={"scim_active": False})
)
if rejection == "custom_validate":
handler.litellm_jwtauth.custom_validate = lambda claims: False
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
@pytest.mark.parametrize("blocked", [False, True])
async def test_oauth_jwt_cannot_override_explicit_litellm_key(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
blocked: bool,
) -> None:
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._types import UserAPIKeyAuth, hash_token
handler, signing_key = jwt_oauth_identity
key: Final = "sk-explicit-key"
handler.user_api_key_cache.set_cache(hash_token(key), UserAPIKeyAuth(user_id="key-owner", blocked=blocked))
request: Final = _token_request(
{
"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}",
"x-litellm-api-key": key,
}
)
assert await _extract_user_id_from_request(request) == (None if blocked else "key-owner")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"mapping", ["active", "blocked", "inactive_owner", "fallback", "pending", "reject", "custom_reject"]
)
async def test_oauth_jwt_uses_configured_virtual_key_owner(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
mapping: str,
) -> None:
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._types import UserAPIKeyAuth, UnregisteredJWTClientBehavior, hash_token
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.virtual_key_claim_field = "sub"
if mapping == "custom_reject":
handler.litellm_jwtauth.custom_validate = lambda claims: False
handler.litellm_jwtauth.unregistered_jwt_client_behavior = (
UnregisteredJWTClientBehavior.AUTO_REGISTER
if mapping == "pending"
else UnregisteredJWTClientBehavior.REJECT
if mapping == "reject"
else UnregisteredJWTClientBehavior.FALLBACK_TEAM_MAPPING
)
key_hash: Final = hash_token("sk-mapped-oauth-owner")
handler.user_api_key_cache.set_cache(
jwt_key_mapping_cache_key("sub", "not-the-configured-user-id"),
"__NO_MAPPING__" if mapping in ("fallback", "pending", "reject") else key_hash,
)
handler.user_api_key_cache.set_cache(
key_hash, UserAPIKeyAuth(token=key_hash, user_id="mapped-owner", blocked=mapping == "blocked")
)
handler.user_api_key_cache.set_cache(
"mapped-owner", LiteLLM_UserTable(user_id="mapped-owner", metadata={"scim_active": mapping != "inactive_owner"})
)
request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"})
expected: Final = "jwt-owner" if mapping == "fallback" else "mapped-owner" if mapping == "active" else None
assert await _extract_user_id_from_request(request) == expected
@pytest.mark.asyncio
@pytest.mark.parametrize("allowed_domain", [None, "allowed.example.test"])
async def test_oauth_jwt_respects_custom_validation_and_email_policy(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
allowed_domain: str | None,
) -> None:
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.custom_validate = lambda claims: True
handler.litellm_jwtauth.user_allowed_email_domain = allowed_domain
request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"})
assert await _extract_user_id_from_request(request) == (None if allowed_domain else "jwt-owner")
@pytest.mark.asyncio
@pytest.mark.parametrize("route_allowed", [False, True])
async def test_oauth_jwt_identity_preserves_separate_mcp_route_authorization(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
route_allowed: bool,
) -> None:
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._types import LitellmUserRoles, RoleBasedPermissions, RoleMapping
from litellm.proxy.auth.handle_jwt import JWTAuthManager
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.user_id_jwt_field = "sub"
handler.litellm_jwtauth.roles_jwt_field = "aud"
handler.litellm_jwtauth.object_id_jwt_field = "identity.user_id"
handler.litellm_jwtauth.role_mappings = [
RoleMapping(role="litellm-proxy", internal_role=LitellmUserRoles.INTERNAL_USER)
]
handler.litellm_jwtauth.enforce_rbac = True
monkeypatch.setattr(
proxy_server,
"general_settings",
{
"enable_jwt_auth": True,
"role_permissions": [
RoleBasedPermissions(
role=LitellmUserRoles.INTERNAL_USER,
routes=["mcp_routes"] if route_allowed else ["/models"],
)
],
},
)
bearer: Final = _oauth_identity_jwt(signing_key)
request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/example/token")
assert await _extract_user_id_from_request(request) == "jwt-owner"
admission: Final = JWTAuthManager.auth_builder(
api_key=bearer,
jwt_handler=handler,
request_data={},
general_settings=proxy_server.general_settings,
route="/mcp/example",
prisma_client=proxy_server.prisma_client,
user_api_key_cache=handler.user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_server.proxy_logging_obj,
request_method="POST",
)
if route_allowed:
assert (await admission)["user_id"] == "jwt-owner"
else:
with pytest.raises(HTTPException) as denial:
await admission
assert denial.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize("identity", ["sso", "email"])
@pytest.mark.parametrize("inactive", [False, True])
@pytest.mark.parametrize("admin", [False, True])
async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
identity: str,
inactive: bool,
admin: bool,
) -> None:
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.auth.handle_jwt import JWTAuthManager
handler, signing_key = jwt_oauth_identity
external_id: Final = f"external-{identity}-{inactive}-{admin}"
handler.litellm_jwtauth.user_email_jwt_field = "email"
handler.litellm_jwtauth.admin_allowed_routes = ["mcp_routes"]
owner: Final = LiteLLM_UserTable(
user_id="canonical-oauth-owner",
user_email="owner@example.test",
metadata={"scim_active": not inactive},
organization_memberships=[],
)
database: Final = MagicMock()
table: Final = database.db.litellm_usertable
table.find_unique = AsyncMock(side_effect=[None, owner if identity == "sso" else None])
table.find_first = AsyncMock(return_value=owner)
table.update = AsyncMock(return_value=owner)
monkeypatch.setattr(proxy_server, "prisma_client", database)
bearer: Final = _oauth_identity_jwt(signing_key, owner=external_id, scope="litellm_proxy_admin" if admin else "")
request: Final = _token_request({"Authorization": f"Bearer {bearer}"})
stored_owner: Final = await _extract_user_id_from_request(request)
assert stored_owner == (None if inactive else external_id if admin else "canonical-oauth-owner")
assert table.find_unique.await_count == 2
if identity == "email":
table.find_first.assert_awaited_once()
if not inactive:
admission: Final = await JWTAuthManager.auth_builder(
api_key=bearer,
jwt_handler=handler,
request_data={},
general_settings=proxy_server.general_settings,
route="/mcp/example",
prisma_client=database,
user_api_key_cache=handler.user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_server.proxy_logging_obj,
)
assert stored_owner == admission["user_id"]
@pytest.mark.asyncio
async def test_oauth_jwt_identity_does_not_provision_or_synchronize_teams(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
) -> None:
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
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.enforce_team_based_model_access = True
handler.litellm_jwtauth.team_id_default = "new-team"
handler.litellm_jwtauth.team_id_upsert = True
handler.litellm_jwtauth.sync_user_role_and_teams = True
owner: Final = LiteLLM_UserTable(user_id="jwt-owner", teams=["existing-team"])
handler.user_api_key_cache.set_cache("jwt-owner", owner)
request: Final = _token_request(
{"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}, path="/example/token"
)
assert await _extract_user_id_from_request(request) == "jwt-owner"
assert owner.teams == ["existing-team"]
proxy_server.prisma_client.db.litellm_teamtable.find_unique.assert_not_called()
proxy_server.prisma_client.db.litellm_teamtable.upsert.assert_not_called()
proxy_server.prisma_client.db.litellm_usertable.update.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("state", ["active", "inactive", "missing_database"])
async def test_oauth_refresh_revalidates_the_same_active_user_rule(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
state: str,
) -> None:
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id
handler, _ = jwt_oauth_identity
handler.user_api_key_cache.set_cache(
"jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", metadata={"scim_active": state != "inactive"})
)
if state == "missing_database":
monkeypatch.setattr(proxy_server, "prisma_client", None)
expected: Final = None if state == "active" else "no_active_key" if state == "inactive" else "unresolvable"
assert await _reload_active_user_by_id("jwt-owner") == expected
@pytest.mark.asyncio
@pytest.mark.parametrize("mapped", [False, True])
@pytest.mark.parametrize("state", ["allowed", "route_denied", "server_denied", "blocked", "expired", "lookup_error", "cancelled"])
async def test_oauth_credential_write_keeps_virtual_key_permissions(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
mapped: bool,
state: str,
) -> None:
import asyncio
from litellm.proxy._experimental.mcp_server import mcp_server_manager
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
handler, signing_key = jwt_oauth_identity
key: Final = "sk-oauth-permission-test"
hashed: Final = hash_token(key)
credential: Final = UserAPIKeyAuth(
token=hashed,
user_id="jwt-owner",
blocked=state == "blocked",
expires=datetime.now(timezone.utc) - timedelta(seconds=60) if state == "expired" else None,
allowed_routes=["openai_routes"] if state == "route_denied" else ["mcp_routes"],
agent_id="agent-scope",
org_id="org-scope",
end_user_id="end-user-scope",
)
handler.user_api_key_cache.set_cache(hashed, credential)
if mapped:
handler.litellm_jwtauth.virtual_key_claim_field = "sub"
handler.user_api_key_cache.set_cache(jwt_key_mapping_cache_key("sub", "not-the-configured-user-id"), hashed)
manager: Final = MagicMock()
manager.get_allowed_mcp_servers = AsyncMock(
return_value=[] if state == "server_denied" else ["server-a"],
side_effect=(asyncio.CancelledError() if state == "cancelled" else RuntimeError("permission lookup unavailable") if state == "lookup_error" else None),
)
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
bearer: Final = _oauth_identity_jwt(signing_key) if mapped else key
request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/server-a/token")
if state == "cancelled":
with pytest.raises(asyncio.CancelledError):
await authorize_oauth_credential_request(request, "server-a")
manager.get_allowed_mcp_servers.assert_awaited_once()
return
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]
assert (writer.user_id, writer.token, writer.org_id, writer.agent_id, writer.end_user_id) == (
"jwt-owner",
hashed,
"org-scope",
"agent-scope",
"end-user-scope",
)
@pytest.mark.asyncio
@pytest.mark.parametrize("server_id", ["team-a-server", "team-b-server"])
async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_roster(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
server_id: str,
) -> None:
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 authorize_oauth_credential_request
from litellm.proxy._types import LiteLLM_TeamTable, Member
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.team_id_jwt_field = "team"
handler.litellm_jwtauth.team_id_upsert = True
handler.litellm_jwtauth.user_id_upsert = True
handler.litellm_jwtauth.sync_user_role_and_teams = True
handler.user_api_key_cache.set_cache("jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", teams=["a", "b"]))
handler.user_api_key_cache.set_cache(
"team_id:a",
LiteLLM_TeamTable(team_id="a", models=[], members_with_roles=[Member(user_id="jwt-owner", role="user")]),
)
manager: Final = MagicMock()
manager.get_allowed_mcp_servers = AsyncMock(return_value=["team-a-server"])
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 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()
writer: Final = manager.get_allowed_mcp_servers.call_args.args[0]
assert writer.team_id == "a"
assert not writer.mcp_admitted_user_subject
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
proxy_server.prisma_client.db.litellm_usertable.update.assert_not_called()
assert handler.litellm_jwtauth.user_id_upsert and handler.litellm_jwtauth.team_id_upsert
assert handler.litellm_jwtauth.sync_user_role_and_teams
@pytest.mark.asyncio
async def test_oauth_write_denial_does_not_erase_identity_binding(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.proxy._experimental.mcp_server import discoverable_endpoints, mcp_server_manager
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer
_, signing_key = jwt_oauth_identity
monkeypatch.setenv("LITELLM_SALT_KEY", "oauth-identity-binding-test-salt")
manager: Final = MagicMock()
manager.get_allowed_mcp_servers = AsyncMock(return_value=[])
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
server: Final = MCPServer(
server_id="bound-server", name="bound-server", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client",
token_url="https://upstream.example.test/token",
oauth_identity_binding=MCPOAuthIdentityBinding(
mode="enforce", issuer="https://upstream.example.test", audiences=["client"],
),
)
request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"})
code: Final = discoverable_endpoints.seal_bridge_authorization_code(
"upstream-code", "another-owner", server.server_id, "bound-nonce",
)
with pytest.raises(HTTPException) as denied:
await discoverable_endpoints.exchange_token_with_server(
request=request, mcp_server=server, grant_type="authorization_code", code=code,
redirect_uri="http://localhost/callback", client_id="client", client_secret=None, code_verifier="verifier",
)
assert denied.value.status_code == 403
assert denied.value.detail == {"error": "oauth_principal_mismatch"}
manager.get_allowed_mcp_servers.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("admin_only", [False, True])
async def test_signed_oauth_callback_honors_credential_write_policy(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
admin_only: bool,
) -> None:
import httpx
import litellm
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import discoverable_endpoints, mcp_server_manager
from litellm.proxy._types import MCPTransport
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server: Final = MCPServer(
server_id="signed-server", name="signed-server", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client",
token_url="https://upstream.example.test/token",
)
monkeypatch.setattr(proxy_server, "general_settings", {
"enable_jwt_auth": True,
"admin_only_routes": [f"/v1/mcp/server/{server.server_id}/oauth-user-credential"] if admin_only else [],
})
monkeypatch.setenv("LITELLM_SALT_KEY", "signed-oauth-test-salt")
manager: Final = MagicMock()
manager.get_allowed_mcp_servers = AsyncMock(return_value=[server.server_id])
manager.invalidate_user_oauth_token_cache = AsyncMock()
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
table: Final = proxy_server.prisma_client.db.litellm_mcpusercredentials
table.find_unique = AsyncMock(return_value=None)
table.upsert = AsyncMock()
clients: Final = LLMClientCache()
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", clients)
def upstream_response(outbound: httpx.Request) -> httpx.Response:
assert outbound.url == server.token_url
assert b"code=upstream-code" in outbound.content
return httpx.Response(200, json={"access_token": "upstream-token", "token_type": "Bearer"})
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream_response)) as transport:
upstream: Final = AsyncHTTPHandler()
await upstream.client.aclose()
upstream.client = transport
clients.set_cache("async_httpx_client" + httpxSpecialProvider.Oauth2Check, upstream)
response: Final = await discoverable_endpoints.exchange_token_with_server(
request=_token_request({}, path="/signed-server/token"), mcp_server=server,
grant_type="authorization_code",
code=discoverable_endpoints.seal_bridge_authorization_code("upstream-code", "jwt-owner", server.server_id),
redirect_uri="http://localhost/callback", client_id="client", client_secret=None, code_verifier=None,
)
assert response.status_code == 200
assert json.loads(response.body)["access_token"] == "upstream-token"
if admin_only:
table.upsert.assert_not_awaited()
else:
table.upsert.assert_awaited_once()
assert table.upsert.call_args.kwargs["where"]["user_id_server_id"] == {
"user_id": "jwt-owner", "server_id": server.server_id,
}
@pytest.mark.asyncio
@pytest.mark.parametrize("allowed", [False, True])
@pytest.mark.parametrize("credential", [
"jwt", "key", "expired_jwt", "wrong_audience", "bad_signature", "malformed_jwt", "missing_issuer",
"foreign_explicit", "blank_explicit", "unknown_key", "blocked_key", "expired_key", "opaque_record",
"opaque_outage", "opaque_oidc", "opaque_custom", "foreign_unscoped", "foreign_configured", "encrypted", "invalid_encrypted", "envelope", "master",
])
async def test_identity_bound_authorize_preserves_presented_jwt_permissions(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
allowed: bool,
credential: str,
) -> None:
import jwt
from datetime import datetime, timedelta, timezone
from urllib.parse import parse_qs, urlparse
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._types import JWTIssuerConfig, UserAPIKeyAuth, hash_token
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
from litellm.proxy._experimental.mcp_server import discoverable_endpoints, mcp_server_manager
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer
handler, signing_key = jwt_oauth_identity
master: Final = "browser-session-test-signing-key-123456789"
monkeypatch.setattr(proxy_server, "master_key", master)
monkeypatch.setattr(proxy_server, "user_custom_auth", (lambda: None) if credential == "opaque_custom" else None)
handler.litellm_jwtauth.oidc_userinfo_enabled = credential == "opaque_oidc"
if credential == "foreign_unscoped":
monkeypatch.delenv("JWT_ISSUER")
if credential == "foreign_configured":
handler.litellm_jwtauth.issuers = [JWTIssuerConfig(
issuer="https://unrelated.example.test", jwks_url="https://idp.example.test/jwks",
audience="litellm-proxy", user_id_jwt_field="identity.user_id",
)]
proxy_server.prisma_client.get_data = AsyncMock(
return_value=None, side_effect=RuntimeError("database unavailable") if credential == "opaque_outage" else None,
)
handler.user_api_key_cache.set_cache("cookie-owner", LiteLLM_UserTable(user_id="cookie-owner"))
key: Final = "opaque-record" if credential == "opaque_record" else "sk-browser-gateway-key"
if credential in ("key", "blocked_key", "expired_key", "opaque_record"):
handler.user_api_key_cache.set_cache(hash_token(key), UserAPIKeyAuth(
token=hash_token(key), user_id="jwt-owner", blocked=credential in ("blocked_key", "opaque_record"),
expires=datetime.now(timezone.utc) - timedelta(seconds=60) if credential == "expired_key" else None,
))
monkeypatch.setenv("LITELLM_SALT_KEY", "authorize-policy-test-salt")
server: Final = MCPServer(
server_id="bound-server", name="bound-server", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client",
authorization_url="https://upstream.example.test/authorize", token_url="https://upstream.example.test/token",
oauth_identity_binding=MCPOAuthIdentityBinding(
mode="enforce", issuer="https://upstream.example.test", audiences=["client"],
),
)
manager: Final = MagicMock()
# The full user roster permits the server; the presented JWT may have narrower access.
manager.get_allowed_mcp_servers = AsyncMock(
side_effect=lambda auth: [server.server_id] if allowed or auth.mcp_admitted_user_subject else [],
)
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
bearer: Final = (
key if credential in ("key", "blocked_key", "expired_key", "opaque_record", "unknown_key")
else "opaque-bearer" if credential in ("opaque_outage", "opaque_oidc", "opaque_custom")
else "not.a.jwt" if credential == "malformed_jwt"
else "llm_env_invalid" if credential == "envelope"
else "v2:gcm:invalid" if credential == "invalid_encrypted"
else master if credential == "master"
else ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(
LiteLLM_UserTable(user_id="jwt-owner", user_role="internal_user"),
) if credential == "encrypted"
else jwt.encode({"iss": "https://idp.example.test"}, "wrong-signing-key-at-least-32-bytes", algorithm="HS256")
if credential == "bad_signature"
else jwt.encode({"sub": "jwt-owner"}, signing_key, algorithm="RS256") if credential == "missing_issuer"
else _oauth_identity_jwt(
signing_key,
expires_in=-60 if credential == "expired_jwt" else 300,
audience="another-service" if credential == "wrong_audience" else "litellm-proxy",
issuer="https://unrelated.example.test" if credential.startswith("foreign_") or credential == "blank_explicit" else "https://idp.example.test",
)
)
cookie: Final = jwt.encode(
{"user_id": "cookie-owner", "login_method": "sso", "exp": int(time.time()) + 300}, master, algorithm="HS256",
)
response: Final = await discoverable_endpoints.authorize_with_server(
request=_token_request({
"Authorization": f"Bearer {bearer}", "Cookie": f"token={cookie}",
**({"x-litellm-api-key": bearer} if credential == "foreign_explicit" else {}),
**({"x-litellm-api-key": ""} if credential == "blank_explicit" else {}),
}),
mcp_server=server, client_id="client", redirect_uri="http://127.0.0.1:6274/callback",
state="client-state", code_challenge="pkce-challenge", code_challenge_method="S256",
)
redirect: Final = urlparse(response.headers["location"])
query: Final = parse_qs(redirect.query)
if allowed and credential in ("jwt", "key", "foreign_unscoped", "foreign_configured"):
assert redirect.hostname == "upstream.example.test"
assert query["nonce"] and response.headers.get("set-cookie")
assert all(call.args[0].user_id == "jwt-owner" for call in manager.get_allowed_mcp_servers.await_args_list)
else:
assert redirect.hostname == "127.0.0.1"
assert query["error"] == ["access_denied"]
assert query["state"] == ["client-state"]
assert "set-cookie" not in response.headers
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("credential", ["none", "opaque", "foreign_jwt"])
@pytest.mark.parametrize("cookie_state", ["allowed", "server_denied", "expired", "missing"])
async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
credential: str,
cookie_state: str,
) -> None:
import jwt
from urllib.parse import parse_qs, urlparse
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import discoverable_endpoints, mcp_server_manager
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer
handler, signing_key = jwt_oauth_identity
master: Final = "browser-session-test-signing-key-123456789"
monkeypatch.setattr(proxy_server, "master_key", master)
monkeypatch.setattr(proxy_server, "user_custom_auth", None)
monkeypatch.setenv("LITELLM_SALT_KEY", "authorize-policy-test-salt")
handler.user_api_key_cache.set_cache("cookie-owner", LiteLLM_UserTable(user_id="cookie-owner"))
proxy_server.prisma_client.get_data = AsyncMock(return_value=None)
server: Final = MCPServer(
server_id="bound-server", name="bound-server", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client",
authorization_url="https://upstream.example.test/authorize", token_url="https://upstream.example.test/token",
oauth_identity_binding=MCPOAuthIdentityBinding(
mode="enforce", issuer="https://upstream.example.test", audiences=["client"],
),
)
manager: Final = MagicMock()
manager.get_allowed_mcp_servers = AsyncMock(return_value=[] if cookie_state == "server_denied" else [server.server_id])
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
bearer: Final = (
_oauth_identity_jwt(signing_key, issuer="https://unrelated.example.test")
if credential == "foreign_jwt" else "unrelated-upstream-bearer"
)
cookie: Final = jwt.encode(
{"user_id": "cookie-owner", "login_method": "sso", "exp": int(time.time()) + (-60 if cookie_state == "expired" else 300)},
master, algorithm="HS256",
)
response: Final = await discoverable_endpoints.authorize_with_server(
request=_token_request({
**({"Authorization": f"Bearer {bearer}"} if credential != "none" else {}),
**({"Cookie": f"token={cookie}"} if cookie_state != "missing" else {}),
}),
mcp_server=server, client_id="client", redirect_uri="http://127.0.0.1:6274/callback",
state="client-state", code_challenge="pkce-challenge", code_challenge_method="S256",
)
redirect: Final = urlparse(response.headers["location"])
query: Final = parse_qs(redirect.query)
if cookie_state == "allowed":
assert redirect.hostname == "upstream.example.test"
assert query["nonce"] and response.headers.get("set-cookie")
manager.get_allowed_mcp_servers.assert_awaited_once()
assert manager.get_allowed_mcp_servers.call_args.args[0].user_id == "cookie-owner"
elif cookie_state == "server_denied":
assert query["error"] == ["access_denied"]
assert query["state"] == ["client-state"]
else:
assert redirect.path == "/sso/key/generate"
manager.get_allowed_mcp_servers.assert_not_awaited()
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()

View file

@ -2,7 +2,7 @@ import asyncio
import re
import time
from collections.abc import Mapping, Sequence
from typing import Optional
from typing import Final, Optional
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import HTTPException
@ -6790,6 +6790,88 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla
assert user.teams == []
@pytest.mark.asyncio
@pytest.mark.parametrize("operation", ["identity", "authorize", "admit"])
@pytest.mark.parametrize("existing_user", [False, True])
@pytest.mark.parametrize("model_allowed", [False, True])
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
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
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-{operation}-{existing_user}-{model_allowed}"
user: Final = LiteLLM_UserTable(user_id=user_id, organization_memberships=[])
if existing_user:
cache.set_cache(user_id, user)
database: Final = MagicMock()
users: Final = database.db.litellm_usertable
users.find_unique = AsyncMock(return_value=None)
users.find_first = AsyncMock(return_value=None)
users.create = AsyncMock(return_value=user)
handler: Final = JWTHandler()
handler.update_environment(
prisma_client=database,
user_api_key_cache=cache,
litellm_jwtauth=LiteLLM_JWTAuth(
user_id_jwt_field="sub",
user_id_upsert=True,
enforce_scope_based_access=True,
scope_mappings=[ScopeMapping(scope="allowed", models=["allowed-model"])],
),
)
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", "https://identity.example/jwks")
monkeypatch.setenv("JWT_ISSUER", "https://identity.example")
monkeypatch.setenv("JWT_AUDIENCE", "gateway")
token: Final = _encode_rsa_jwt(
private_key, "https://identity.example", "gateway", "identity-mode", {"sub": user_id, "scope": "allowed"}
)
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="/mcp/example",
)
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 operation == "authorize" and not existing_user:
with pytest.raises(UserNotFoundError):
await pending
else:
result: Final = await pending
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 operation == "authorize" or existing_user else 1)
def _entra_agent_registry() -> AgentRegistry:
registry = AgentRegistry()
registry.register_agent(
@ -6916,7 +6998,8 @@ def _entra_signed_app_token(monkeypatch, azp: str, scope: str) -> tuple[JWTHandl
@pytest.mark.asyncio
@pytest.mark.parametrize("is_admin_token", [False, True], ids=["standard_jwt", "proxy_admin_jwt"])
async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_admin_token: bool):
@pytest.mark.parametrize("identity_only", [False, True])
async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_admin_token: bool, identity_only: bool):
"""auth_builder carries the resolved agent id into JWTAuthBuilderResult on both the admin and standard paths."""
jwt_handler, token = _entra_signed_app_token(
monkeypatch,
@ -6925,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,
@ -6942,7 +7033,8 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a
@pytest.mark.asyncio
async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_check(monkeypatch):
@pytest.mark.parametrize("identity_only", [False, True])
async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_check(monkeypatch, identity_only: bool):
"""An unknown agent claim is rejected even when the token would otherwise be a proxy admin."""
jwt_handler, token = _entra_signed_app_token(
monkeypatch,
@ -6951,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,
@ -6965,3 +7065,36 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch
)
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