mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
Merge pull request #41314 from BerriAI/litellm_fix_mcp_jwt_oauth_persistence
fix(mcp): authorize JWT OAuth credential persistence
This commit is contained in:
commit
9cd787386e
9 changed files with 1505 additions and 161 deletions
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue