refactor(auth): separate JWT identity and OAuth authorization

This commit is contained in:
Joshua Valluru 2026-09-15 22:15:57 -07:00
parent c7e4160ee6
commit e035682ed1
5 changed files with 300 additions and 183 deletions

View file

@ -24,6 +24,7 @@ if TYPE_CHECKING:
UpstreamTokenGrant,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.handle_jwt import JWTIdentity
def _litellm_key_from_request(request: Request) -> str | None:
@ -304,26 +305,37 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol
assert_never(identity.subject_type)
async def _extract_user_id_from_request(request: Request, server_id: str | None = None) -> str | None:
"""Resolve identity for binding, or authorize the credential-write action for a target server."""
async def _extract_user_id_from_request(request: Request) -> str | None:
"""Resolve the caller for identity binding without granting credential-write permission."""
from litellm.proxy.auth.handle_jwt import JWTIdentity # noqa: PLC0415 # proxy import cycle
resolved: Final = await _resolve_request_auth(request)
if isinstance(resolved, JWTIdentity):
return resolved.user_id
return _active_key_user_id(resolved) if resolved is not None else None
async def authorize_oauth_credential_request(request: Request, server_id: str) -> str | None:
from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle
resolved: Final = await _resolve_request_auth(request, f"/v1/mcp/server/{server_id}/oauth-user-credential")
if not isinstance(resolved, UserAPIKeyAuth) or not _active_key_user_id(resolved):
return None
if not await can_store_oauth_credential(request, resolved, server_id):
return None
return resolved.user_id
async def _resolve_request_auth(
request: Request, write_route: str | None = None
) -> "UserAPIKeyAuth | JWTIdentity | None":
from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle
token: Final = _litellm_key_from_request(request)
# The OAuth relay is public; the optional server-side write is the same protected action
# as the direct credential endpoint. Authorize that action without rewriting the Request.
write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential" if server_id is not None else None
resolved: Final = (
await _resolve_jwt_auth(request, token, write_route)
if token is not None and JWTHandler.is_jwt(token)
else await _resolve_active_litellm_key(request)
)
auth: Final = resolved.key if isinstance(resolved, _ResolvedKey) else resolved
if not isinstance(auth, UserAPIKeyAuth) or not _active_key_user_id(auth):
return None
if server_id is not None and not await can_store_oauth_credential(request, auth, server_id):
return None
return auth.user_id
if token is not None and JWTHandler.is_jwt(token):
return await _resolve_jwt_auth(request, token, write_route)
resolved: Final = await _resolve_active_litellm_key(request)
return resolved.key if isinstance(resolved, _ResolvedKey) else None
async def can_store_oauth_credential(request: Request, auth: "UserAPIKeyAuth", server_id: str) -> bool:
@ -358,7 +370,7 @@ async def _resolve_jwt_auth(
request: Request,
token: str,
write_route: str | None,
) -> "UserAPIKeyAuth | None":
) -> "UserAPIKeyAuth | JWTIdentity | None":
from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.handle_jwt import JWTAuthManager # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle
@ -393,25 +405,35 @@ async def _resolve_jwt_auth(
return None if await _key_owner_scim_deactivated(mapped) or not _active_key_user_id(mapped) else mapped
if mapped is not None:
return None
identity: Final = await JWTAuthManager.auth_builder(
if write_route is None:
identity: Final = await JWTAuthManager.resolve_identity(
api_key=token,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
)
if identity.user_object is not None and isinstance(_active_user_record(identity.user_object), str):
return None
return identity
authorized: Final = await JWTAuthManager.authorize_jwt(
api_key=token,
jwt_handler=jwt_handler,
request_data={},
general_settings=general_settings,
route=write_route or request.url.path,
route=write_route,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
request_headers=dict(request.headers),
request_method=request.method,
identity_only=write_route is None,
allow_provisioning=False,
)
resolved_user: Final = identity["user_object"]
resolved_user: Final = authorized["user_object"]
if resolved_user is not None and isinstance(_active_user_record(resolved_user), str):
return None
return JWTAuthManager.user_api_key_auth_from_result(identity)
return JWTAuthManager.user_api_key_auth_from_result(authorized)
except Exception as exc: # noqa: BLE001 # public OAuth exchange stays available; unvalidated identities never write credentials
verbose_logger.debug("OAuth JWT identity could not be validated (%s)", type(exc).__name__)
return None

View file

@ -33,6 +33,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_prepare_bridge_mint,
_prepare_bridge_refresh,
_reload_active_user_by_id,
authorize_oauth_credential_request,
can_store_oauth_credential,
)
from litellm.proxy._experimental.mcp_server.faults import (
@ -851,7 +852,7 @@ async def _resolve_oauth_authorization_user(
)
request_user_id: Final = (
await _extract_user_id_from_request(request, mcp_server.server_id) if enforce_binding else None
await authorize_oauth_credential_request(request, mcp_server.server_id) if enforce_binding else None
)
if enforce_binding and request_user_id is None and _litellm_key_from_request(request):
return _bridge_access_denied_redirect(redirect_uri, state, mcp_server)
@ -1233,7 +1234,7 @@ async def exchange_token_with_server(
request, await MCPRequestHandler.reload_admitted_user(user_id), resolved_server.server_id
)
if bridge_identity is not None
else await _extract_user_id_from_request(request, resolved_server.server_id) == user_id
else await authorize_oauth_credential_request(request, resolved_server.server_id) == user_id
)
if can_store:
await _store_per_user_token_server_side(

View file

@ -9,13 +9,13 @@ JWT token must have 'litellm_proxy_admin' in scope.
from __future__ import annotations
import asyncio
import copy
import fnmatch
import hashlib
import os
import re
import time
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast
import httpx
@ -130,6 +130,19 @@ class _UserInfoResponse(Protocol):
def json(self) -> dict[str, object]: ...
@dataclass(frozen=True, slots=True)
class JWTIdentity:
user_id: str | None
user_object: LiteLLM_UserTable | None
agent_id: str | None
@dataclass(frozen=True, slots=True)
class _JWTProvisioning:
user_id_upsert: bool
team_id_upsert: bool
class AgentLookup(Protocol):
"""The registered-agent lookups a JWT agent claim is matched against."""
@ -1473,6 +1486,7 @@ class JWTAuthManager:
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging,
team_id_upsert: bool | None = None,
) -> tuple[str | None, LiteLLM_TeamTable | None]:
"""Find and validate specific team ID from team_id_jwt_field or team_alias_jwt_field"""
individual_team_id = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None)
@ -1500,7 +1514,9 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert
if team_id_upsert is None
else team_id_upsert,
)
return individual_team_id, team_object
except HTTPException as e:
@ -2017,6 +2033,7 @@ class JWTAuthManager:
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging,
team_id_upsert: bool | None = None,
) -> None:
"""Attach team context from x-litellm-team-id to an admin result.
@ -2034,7 +2051,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert if team_id_upsert is None else team_id_upsert,
)
except Exception as e:
# Fall back to pre-PR admin behavior: honor the admin's
@ -2268,64 +2285,119 @@ class JWTAuthManager:
proxy_logging_obj: ProxyLogging,
request_headers: dict | None = None,
request_method: str | None = None,
identity_only: bool = False,
allow_provisioning: bool = True,
) -> JWTAuthBuilderResult:
"""Build JWT authentication and authorization context.
return await JWTAuthManager.authorize_jwt(
api_key=api_key,
jwt_handler=jwt_handler,
request_data=request_data,
general_settings=general_settings,
route=route,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
request_headers=request_headers,
request_method=request_method,
provisioning=_JWTProvisioning(
user_id_upsert=jwt_handler.litellm_jwtauth.user_id_upsert,
team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert,
),
)
identity_only resolves the caller for OAuth identity binding and grants no permission.
Credential writes use full authorization with allow_provisioning=False: resolve the
existing policy context without creating users/teams or synchronizing membership.
A private handler configuration keeps that restriction out of concurrent normal requests.
"""
handler: Final = jwt_handler if allow_provisioning else copy.copy(jwt_handler)
if not allow_provisioning:
handler.update_environment(
prisma_client=jwt_handler.prisma_client,
user_api_key_cache=jwt_handler.user_api_key_cache,
litellm_jwtauth=jwt_handler.litellm_jwtauth.model_copy(
update={"user_id_upsert": False, "team_id_upsert": False, "sync_user_role_and_teams": False}
),
leeway=jwt_handler.leeway,
@staticmethod
async def authenticate_jwt(api_key: str, jwt_handler: JWTHandler) -> dict[str, object]:
claims: Final = (
await jwt_handler.get_oidc_userinfo(token=api_key)
if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not jwt_handler.is_jwt(token=api_key)
else await jwt_handler.auth_jwt(token=api_key)
)
validate: Final = jwt_handler.litellm_jwtauth.custom_validate
if validate is not None and not validate(claims):
raise HTTPException(status_code=403, detail="Invalid JWT token")
return claims
@staticmethod
async def resolve_identity(
api_key: str,
jwt_handler: JWTHandler,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging,
) -> JWTIdentity:
claims: Final = await JWTAuthManager.authenticate_jwt(api_key, jwt_handler)
return await JWTAuthManager._resolve_claim_identity(
claims, jwt_handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
)
@staticmethod
async def _resolve_claim_identity(
claims: dict[str, object],
jwt_handler: JWTHandler,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging,
) -> JWTIdentity:
claim_user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(jwt_handler, claims)
user_id: Final = (
jwt_handler.get_object_id(token=claims, default_value=None) or claim_user_id
if jwt_handler.get_rbac_role(token=claims) == LitellmUserRoles.INTERNAL_USER
else claim_user_id
)
agent_id: Final = JWTAuthManager.resolve_agent_id(jwt_handler, claims, jwt_handler.agent_lookup)
is_admin: Final = jwt_handler.is_admin(scopes=jwt_handler.get_scopes(token=claims))
try:
user, _, _, _, canonical_id = await JWTAuthManager.get_objects(
user_id=user_id,
user_email=user_email,
org_id=None,
end_user_id=None,
team_id=None,
valid_user_email=valid_user_email,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route="",
user_id_upsert=False,
)
except UserNotFoundError:
if not is_admin:
raise
return JWTIdentity(user_id=user_id, user_object=None, agent_id=agent_id)
return JWTIdentity(user_id=user_id if is_admin else canonical_id, user_object=user, agent_id=agent_id)
# Check if OIDC UserInfo endpoint is enabled, but fall back to standard
# JWT auth if the token itself is a well-formed JWT (3-part structure).
if handler.litellm_jwtauth.oidc_userinfo_enabled and not handler.is_jwt(token=api_key):
verbose_proxy_logger.debug("OIDC UserInfo is enabled. Fetching user info from UserInfo endpoint.")
# Use the access token to fetch user info from OIDC UserInfo endpoint
jwt_valid_token: dict = await handler.get_oidc_userinfo(token=api_key)
else:
# Default behavior: decode and validate the JWT token
jwt_valid_token = await handler.auth_jwt(token=api_key)
# Check custom validate
if handler.litellm_jwtauth.custom_validate:
if not handler.litellm_jwtauth.custom_validate(jwt_valid_token):
raise HTTPException(
status_code=403,
detail="Invalid JWT token",
)
@staticmethod
async def authorize_jwt(
api_key: str,
jwt_handler: JWTHandler,
request_data: dict[str, object],
general_settings: dict[str, object],
route: str,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging,
request_headers: dict[str, str] | None = None,
request_method: str | None = None,
provisioning: _JWTProvisioning | None = None,
) -> JWTAuthBuilderResult:
"""Resolve and authorize JWT context; only normal admission supplies provisioning."""
handler: Final = jwt_handler
jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler)
team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False
model: Final = request_data.get("model")
requested_model: Final = model if isinstance(model, str) else None
# Check RBAC
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token)
if not identity_only:
await JWTAuthManager.check_rbac_role(
handler,
jwt_valid_token,
general_settings,
request_data,
route,
rbac_role,
)
await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role)
# Check Scope Based Access
scopes: Final = handler.get_scopes(token=jwt_valid_token)
if (
not identity_only
and handler.litellm_jwtauth.enforce_scope_based_access
and handler.litellm_jwtauth.scope_mappings
):
if handler.litellm_jwtauth.enforce_scope_based_access and handler.litellm_jwtauth.scope_mappings:
JWTAuthManager.check_scope_based_access(
scope_mappings=handler.litellm_jwtauth.scope_mappings,
scopes=scopes,
@ -2357,69 +2429,6 @@ class JWTAuthManager:
agent_registry=handler.agent_lookup,
)
if identity_only or (not allow_provisioning and handler.is_admin(scopes=scopes)):
try:
identity_user, _, _, _, identity_user_id = await JWTAuthManager.get_objects(
user_id=user_id,
user_email=user_email,
org_id=None,
end_user_id=None,
team_id=None,
valid_user_email=valid_user_email,
jwt_handler=handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
user_id_upsert=False,
)
except UserNotFoundError:
if not handler.is_admin(scopes=scopes):
raise
identity_user, identity_user_id = None, user_id
if not identity_only:
admin: Final = await JWTAuthManager.check_admin_access(
handler,
scopes,
route,
user_id,
org_id,
api_key,
jwt_valid_token,
user_email=user_email,
agent_id=agent_id,
)
if admin is not None:
await JWTAuthManager._attach_team_from_header_for_admin(
admin_result=admin,
route=route,
request_headers=request_headers,
jwt_handler=handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
return {**admin, "user_object": identity_user}
return JWTAuthBuilderResult(
is_proxy_admin=False,
# Admin admission uses the claim ID; other callers use the canonical DB ID.
user_id=user_id if handler.is_admin(scopes=scopes) else identity_user_id,
user_email=identity_user.user_email if identity_user is not None else user_email,
user_object=identity_user,
team_id=None,
team_object=None,
org_id=None,
org_object=None,
end_user_id=None,
end_user_object=None,
team_membership=None,
token=api_key,
jwt_claims=jwt_valid_token,
agent_id=agent_id,
)
# Check admin access
admin_result: Final = await JWTAuthManager.check_admin_access(
handler,
@ -2442,7 +2451,13 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=team_id_upsert,
)
if provisioning is None:
identity: Final = await JWTAuthManager._resolve_claim_identity(
jwt_valid_token, handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
)
return {**admin_result, "user_object": identity.user_object}
return admin_result
# Get team with model access
@ -2485,7 +2500,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=(handler.litellm_jwtauth.team_id_upsert and not db_team_fallback),
team_id_upsert=(team_id_upsert and not db_team_fallback),
)
except HTTPException:
if not db_team_fallback:
@ -2503,13 +2518,14 @@ class JWTAuthManager:
user_api_key_cache,
parent_otel_span,
proxy_logging_obj,
team_id_upsert=team_id_upsert,
)
if not team_object and not team_id:
## CHECK USER GROUP ACCESS
team_id, team_object = await JWTAuthManager.find_team_with_model_access(
team_ids=all_team_ids,
requested_model=request_data.get("model"),
requested_model=requested_model,
route=route,
request_method=request_method,
jwt_handler=handler,
@ -2536,7 +2552,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=handler.litellm_jwtauth.team_id_upsert,
team_id_upsert=team_id_upsert,
)
if team_id and not JWTAuthManager._team_has_passthrough_route_access(
@ -2570,18 +2586,20 @@ class JWTAuthManager:
proxy_logging_obj=proxy_logging_obj,
route=route,
org_alias=org_alias,
user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False,
)
# Derive org_id from org_object if resolved by alias
resolved_org_id: Final = org_object.organization_id if org_object else org_id
await JWTAuthManager.sync_user_role_and_teams(
jwt_handler=handler,
jwt_valid_token=jwt_valid_token,
user_object=user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
if provisioning is not None:
await JWTAuthManager.sync_user_role_and_teams(
jwt_handler=handler,
jwt_valid_token=jwt_valid_token,
user_object=user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
# If JWT did not resolve team_id, attempt a team fallback.
if team_id is None and db_team_fallback:
@ -2592,11 +2610,11 @@ class JWTAuthManager:
) = await JWTAuthManager._resolve_db_team_fallback(
user_object=user_object,
user_id=user_id,
requested_model=request_data.get("model"),
requested_model=requested_model,
route=route,
jwt_handler=handler,
enforce_team_based_model_access=handler.litellm_jwtauth.enforce_team_based_model_access,
team_id_upsert=handler.litellm_jwtauth.team_id_upsert,
team_id_upsert=team_id_upsert,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
@ -2624,7 +2642,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
team_id_upsert=handler.litellm_jwtauth.team_id_upsert,
team_id_upsert=team_id_upsert,
)
elif db_team_fallback and team_id == header_team_id:
JWTAuthManager._validate_header_team_in_db_membership(
@ -2644,7 +2662,7 @@ class JWTAuthManager:
)
## MAP USER TO TEAMS
if allow_provisioning:
if provisioning is not None:
await JWTAuthManager.map_user_to_teams(
user_object=user_object,
team_object=team_object,
@ -2654,7 +2672,7 @@ class JWTAuthManager:
JWTAuthManager.validate_object_id(
user_id=user_id,
team_id=team_id,
enforce_rbac=general_settings.get("enforce_rbac", False),
enforce_rbac=bool(general_settings.get("enforce_rbac", False)),
is_proxy_admin=False,
)

View file

@ -6980,6 +6980,11 @@ async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool:
new_callable=AsyncMock,
return_value="admin-user",
),
patch( # test-quality-ok: this control tests persistence by auth mode; write-policy behavior is covered separately
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.authorize_oauth_credential_request",
new_callable=AsyncMock,
return_value="admin-user",
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._store_per_user_token_server_side",
new_callable=AsyncMock,
@ -11168,7 +11173,7 @@ async def test_identity_bound_authorization_carries_nonce_and_caller_through_cal
"path": "/authorize", "query_string": b"", "headers": []})
with (
patch( # test-quality-ok: isolate authenticated request resolution from the real encrypted OAuth round trip
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request",
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.authorize_oauth_credential_request",
new=AsyncMock(return_value="alice")),
patch( # test-quality-ok: isolate user access lookup while testing nonce and caller preservation
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._user_can_reach_mcp_server",
@ -11569,17 +11574,24 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user(
"missing_database",
],
)
@pytest.mark.parametrize("credential_write", [False, True])
async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
rejection: str,
credential_write: bool,
) -> None:
from cryptography.hazmat.primitives.asymmetric import rsa
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request, authorize_oauth_credential_request,
)
allowed_servers: Final = AsyncMock(return_value=["server-a"])
monkeypatch.setattr(mcp_server_manager.global_mcp_server_manager, "get_allowed_mcp_servers", allowed_servers)
handler, signing_key = jwt_oauth_identity
key: Final = (
rsa.generate_private_key(public_exponent=65537, key_size=2048) if rejection == "signature" else signing_key
@ -11603,7 +11615,13 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner(
)
if rejection == "custom_validate":
handler.litellm_jwtauth.custom_validate = lambda claims: False
assert await _extract_user_id_from_request(_token_request({"Authorization": f"Bearer {bearer}"})) is None
request: Final = _token_request({"Authorization": f"Bearer {bearer}"})
result: Final = (
await authorize_oauth_credential_request(request, "server-a")
if credential_write else await _extract_user_id_from_request(request)
)
assert result is None
allowed_servers.assert_not_awaited()
@pytest.mark.asyncio
@ -11850,7 +11868,7 @@ async def test_oauth_credential_write_keeps_virtual_key_permissions(
import asyncio
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._experimental.mcp_server.bridge_token_flow import authorize_oauth_credential_request
from litellm.proxy._types import UserAPIKeyAuth, hash_token
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
@ -11881,10 +11899,10 @@ async def test_oauth_credential_write_keeps_virtual_key_permissions(
request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/server-a/token")
if state == "cancelled":
with pytest.raises(asyncio.CancelledError):
await _extract_user_id_from_request(request, "server-a")
await authorize_oauth_credential_request(request, "server-a")
manager.get_allowed_mcp_servers.assert_awaited_once()
return
assert await _extract_user_id_from_request(request, "server-a") == ("jwt-owner" if state == "allowed" else None)
assert await authorize_oauth_credential_request(request, "server-a") == ("jwt-owner" if state == "allowed" else None)
if state in ("allowed", "server_denied", "lookup_error"):
manager.get_allowed_mcp_servers.assert_awaited_once()
writer: Final = manager.get_allowed_mcp_servers.call_args.args[0]
@ -11907,7 +11925,7 @@ async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_ros
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._experimental.mcp_server.bridge_token_flow import authorize_oauth_credential_request
from litellm.proxy._types import LiteLLM_TeamTable, Member
handler, signing_key = jwt_oauth_identity
@ -11925,7 +11943,7 @@ async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_ros
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
bearer: Final = _oauth_identity_jwt(signing_key, claims={"team": "a"})
request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path=f"/{server_id}/token")
assert await _extract_user_id_from_request(request, server_id) == (
assert await authorize_oauth_credential_request(request, server_id) == (
"jwt-owner" if server_id == "team-a-server" else None
)
manager.get_allowed_mcp_servers.assert_awaited_once()

View file

@ -6791,12 +6791,11 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla
@pytest.mark.asyncio
@pytest.mark.parametrize("identity_only", [False, True])
@pytest.mark.parametrize("allow_provisioning", [False, True])
@pytest.mark.parametrize("operation", ["identity", "authorize", "admit"])
@pytest.mark.parametrize("existing_user", [False, True])
@pytest.mark.parametrize("model_allowed", [False, True])
async def test_auth_builder_identity_lookup_does_not_provision_users(
monkeypatch: pytest.MonkeyPatch, identity_only: bool, allow_provisioning: bool, existing_user: bool, model_allowed: bool
async def test_jwt_identity_and_authorization_keep_provisioning_in_admission(
monkeypatch: pytest.MonkeyPatch, operation: str, existing_user: bool, model_allowed: bool
) -> None:
from litellm.proxy._types import ScopeMapping
from litellm.proxy.auth.auth_checks import UserNotFoundError
@ -6805,7 +6804,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users(
private_key, jwk = _get_rsa_key_and_jwk("identity-mode")
cache: Final = UserApiKeyCache()
cache.set_cache("litellm_jwt_auth_keys_https://identity.example/jwks", [jwk])
user_id: Final = f"identity-mode-{identity_only}-{existing_user}-{model_allowed}"
user_id: Final = f"identity-mode-{operation}-{existing_user}-{model_allowed}"
user: Final = LiteLLM_UserTable(user_id=user_id, organization_memberships=[])
if existing_user:
cache.set_cache(user_id, user)
@ -6831,26 +6830,38 @@ async def test_auth_builder_identity_lookup_does_not_provision_users(
token: Final = _encode_rsa_jwt(
private_key, "https://identity.example", "gateway", "identity-mode", {"sub": user_id, "scope": "allowed"}
)
pending: Final = JWTAuthManager.auth_builder(
api_key=token,
jwt_handler=handler,
common: Final = {
"api_key": token,
"jwt_handler": handler,
"prisma_client": database,
"user_api_key_cache": cache,
"parent_otel_span": None,
"proxy_logging_obj": MagicMock(),
}
if operation == "identity":
if not existing_user:
with pytest.raises(UserNotFoundError):
await JWTAuthManager.resolve_identity(**common)
else:
identity: Final = await JWTAuthManager.resolve_identity(**common)
assert identity.user_id == user_id
assert identity.user_object is not None and identity.user_object.user_id == user_id
users.create.assert_not_awaited()
return
authorize: Final = JWTAuthManager.auth_builder if operation == "admit" else JWTAuthManager.authorize_jwt
pending: Final = authorize(
**common,
request_data={"model": "allowed-model" if model_allowed else "forbidden-model"},
general_settings={},
route="/example/token" if identity_only else "/mcp/example",
prisma_client=database,
user_api_key_cache=cache,
parent_otel_span=None,
proxy_logging_obj=MagicMock(),
identity_only=identity_only,
allow_provisioning=allow_provisioning,
route="/mcp/example",
)
if not identity_only and not model_allowed:
if not model_allowed:
with pytest.raises(HTTPException) as denial:
await pending
assert denial.value.status_code == 403
users.create.assert_not_awaited()
return
if (identity_only or not allow_provisioning) and not existing_user:
if operation == "authorize" and not existing_user:
with pytest.raises(UserNotFoundError):
await pending
else:
@ -6858,7 +6869,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users(
assert result["user_id"] == user_id
assert result["user_object"] is not None
assert result["user_object"].user_id == user_id
assert users.create.await_count == (0 if identity_only or not allow_provisioning or existing_user else 1)
assert users.create.await_count == (0 if operation == "authorize" or existing_user else 1)
def _entra_agent_registry() -> AgentRegistry:
@ -6997,6 +7008,14 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a
)
jwt_handler.bind_agent_lookup(_entra_agent_registry())
if identity_only:
identity = await JWTAuthManager.resolve_identity(
api_key=token, jwt_handler=jwt_handler, prisma_client=None,
user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None,
)
assert identity.agent_id == "canonical-agent-id"
return
result = await JWTAuthManager.auth_builder(
api_key=token,
jwt_handler=jwt_handler,
@ -7007,10 +7026,9 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a
user_api_key_cache=None,
parent_otel_span=None,
proxy_logging_obj=None,
identity_only=identity_only,
)
assert result["is_proxy_admin"] is (is_admin_token and not identity_only)
assert result["is_proxy_admin"] is is_admin_token
assert result["agent_id"] == "canonical-agent-id"
@ -7025,6 +7043,14 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch
)
jwt_handler.bind_agent_lookup(_entra_agent_registry())
if identity_only:
with pytest.raises(HTTPException) as denial:
await JWTAuthManager.resolve_identity(
api_key=token, jwt_handler=jwt_handler, prisma_client=None,
user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None,
)
assert denial.value.status_code == 403
return
with pytest.raises(HTTPException) as exc_info:
await JWTAuthManager.auth_builder(
api_key=token,
@ -7036,7 +7062,39 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch
user_api_key_cache=None,
parent_otel_span=None,
proxy_logging_obj=None,
identity_only=identity_only,
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize("admission", [False, True])
async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatch, admission: bool):
from litellm.proxy.management_endpoints import team_endpoints
handler, token = _entra_signed_app_token(
monkeypatch, azp="canonical-agent-id", scope=LiteLLM_JWTAuth().admin_jwt_scope,
)
handler.bind_agent_lookup(_entra_agent_registry())
handler.litellm_jwtauth.team_id_upsert = True
handler.litellm_jwtauth.admin_allowed_routes = ["openai_routes"]
database = MagicMock()
database.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
create_team = AsyncMock(return_value=LiteLLM_TeamTable(team_id="new-team").model_dump())
monkeypatch.setattr(team_endpoints, "new_team", create_team)
resolve = JWTAuthManager.auth_builder if admission else JWTAuthManager.authorize_jwt
result = await resolve(
api_key=token, jwt_handler=handler, request_data={}, general_settings={},
route="/chat/completions", prisma_client=database,
user_api_key_cache=handler.user_api_key_cache, parent_otel_span=None,
proxy_logging_obj=MagicMock(), request_headers={"x-litellm-team-id": "new-team"},
)
assert result["is_proxy_admin"] is True
if admission:
create_team.assert_awaited_once()
assert result["team_id"] == "new-team"
else:
create_team.assert_not_awaited()
assert result["team_id"] is None