mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
refactor(auth): separate JWT identity and OAuth authorization
This commit is contained in:
parent
c7e4160ee6
commit
e035682ed1
5 changed files with 300 additions and 183 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue