mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): authorize per-user OAuth credential writes
This commit is contained in:
parent
ece1de73b4
commit
97211bc356
8 changed files with 369 additions and 116 deletions
|
|
@ -304,24 +304,55 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol
|
|||
assert_never(identity.subject_type)
|
||||
|
||||
|
||||
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."""
|
||||
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."""
|
||||
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._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler # 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
|
||||
)
|
||||
|
||||
token: Final = _litellm_key_from_request(request)
|
||||
if token is not None and JWTHandler.is_jwt(token):
|
||||
return await _extract_jwt_user_id(request, token)
|
||||
resolved: Final = await _resolve_active_litellm_key(request)
|
||||
if not isinstance(resolved, _ResolvedKey):
|
||||
# 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
|
||||
return _active_key_user_id(resolved.key)
|
||||
if write_route is not None and server_id is not None:
|
||||
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,
|
||||
)
|
||||
if not await can_access_mcp_server(auth, server_id, global_mcp_server_manager.get_allowed_mcp_servers):
|
||||
return None
|
||||
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 None
|
||||
return auth.user_id
|
||||
|
||||
|
||||
async def _extract_jwt_user_id(request: Request, token: str) -> str | None:
|
||||
async def _resolve_jwt_auth(
|
||||
request: Request,
|
||||
token: str,
|
||||
write_route: str | None,
|
||||
) -> "UserAPIKeyAuth | 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
|
||||
|
|
@ -353,7 +384,7 @@ async def _extract_jwt_user_id(request: Request, token: str) -> str | None:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if isinstance(mapped, UserAPIKeyAuth):
|
||||
return None if await _key_owner_scim_deactivated(mapped) else _active_key_user_id(mapped)
|
||||
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(
|
||||
|
|
@ -361,19 +392,20 @@ async def _extract_jwt_user_id(request: Request, token: str) -> str | None:
|
|||
jwt_handler=jwt_handler,
|
||||
request_data={},
|
||||
general_settings=general_settings,
|
||||
route=request.url.path,
|
||||
route=write_route or request.url.path,
|
||||
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=True,
|
||||
identity_only=write_route is None,
|
||||
allow_provisioning=False,
|
||||
)
|
||||
resolved_user: Final = identity["user_object"]
|
||||
if resolved_user is not None and isinstance(_active_user_record(resolved_user), str):
|
||||
return None
|
||||
return identity["user_id"]
|
||||
return JWTAuthManager.user_api_key_auth_from_result(identity)
|
||||
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
|
||||
|
|
|
|||
|
|
@ -1218,12 +1218,26 @@ 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.
|
||||
can_store: Final = (
|
||||
await _user_can_reach_mcp_server(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
|
||||
)
|
||||
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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ JWT token must have 'litellm_proxy_admin' in scope.
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import fnmatch
|
||||
import hashlib
|
||||
import os
|
||||
|
|
@ -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,
|
||||
|
|
@ -2268,36 +2269,49 @@ class JWTAuthManager:
|
|||
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.
|
||||
|
||||
Public OAuth endpoints use identity_only to resolve an existing credential owner
|
||||
without authorizing the OAuth route or provisioning users/teams. The returned
|
||||
identity does not grant permission to execute an MCP or model request.
|
||||
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,
|
||||
)
|
||||
|
||||
# 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):
|
||||
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 jwt_handler.get_oidc_userinfo(token=api_key)
|
||||
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 jwt_handler.auth_jwt(token=api_key)
|
||||
jwt_valid_token = await 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):
|
||||
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",
|
||||
)
|
||||
|
||||
# Check RBAC
|
||||
rbac_role: Final = jwt_handler.get_rbac_role(token=jwt_valid_token)
|
||||
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token)
|
||||
if not identity_only:
|
||||
await JWTAuthManager.check_rbac_role(
|
||||
jwt_handler,
|
||||
handler,
|
||||
jwt_valid_token,
|
||||
general_settings,
|
||||
request_data,
|
||||
|
|
@ -2306,30 +2320,30 @@ class JWTAuthManager:
|
|||
)
|
||||
|
||||
# Check Scope Based Access
|
||||
scopes: Final = jwt_handler.get_scopes(token=jwt_valid_token)
|
||||
scopes: Final = handler.get_scopes(token=jwt_valid_token)
|
||||
if (
|
||||
not identity_only
|
||||
and jwt_handler.litellm_jwtauth.enforce_scope_based_access
|
||||
and jwt_handler.litellm_jwtauth.scope_mappings
|
||||
and 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:
|
||||
|
|
@ -2338,12 +2352,12 @@ 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,
|
||||
)
|
||||
|
||||
if identity_only:
|
||||
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,
|
||||
|
|
@ -2352,7 +2366,7 @@ class JWTAuthManager:
|
|||
end_user_id=None,
|
||||
team_id=None,
|
||||
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,
|
||||
|
|
@ -2361,13 +2375,37 @@ class JWTAuthManager:
|
|||
user_id_upsert=False,
|
||||
)
|
||||
except UserNotFoundError:
|
||||
if not jwt_handler.is_admin(scopes=scopes):
|
||||
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 jwt_handler.is_admin(scopes=scopes) else identity_user_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,
|
||||
|
|
@ -2384,7 +2422,7 @@ class JWTAuthManager:
|
|||
|
||||
# Check admin access
|
||||
admin_result: Final = await JWTAuthManager.check_admin_access(
|
||||
jwt_handler,
|
||||
handler,
|
||||
scopes,
|
||||
route,
|
||||
user_id,
|
||||
|
|
@ -2399,7 +2437,7 @@ 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,
|
||||
|
|
@ -2409,8 +2447,8 @@ class JWTAuthManager:
|
|||
|
||||
# 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
|
||||
|
|
@ -2420,9 +2458,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:
|
||||
|
|
@ -2447,7 +2485,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=(handler.litellm_jwtauth.team_id_upsert and not db_team_fallback),
|
||||
)
|
||||
except HTTPException:
|
||||
if not db_team_fallback:
|
||||
|
|
@ -2459,7 +2497,7 @@ 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,
|
||||
|
|
@ -2474,7 +2512,7 @@ class JWTAuthManager:
|
|||
requested_model=request_data.get("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,
|
||||
|
|
@ -2498,7 +2536,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=handler.litellm_jwtauth.team_id_upsert,
|
||||
)
|
||||
|
||||
if team_id and not JWTAuthManager._team_has_passthrough_route_access(
|
||||
|
|
@ -2509,7 +2547,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).
|
||||
(
|
||||
|
|
@ -2525,7 +2563,7 @@ 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,
|
||||
|
|
@ -2538,7 +2576,7 @@ class JWTAuthManager:
|
|||
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_handler=handler,
|
||||
jwt_valid_token=jwt_valid_token,
|
||||
user_object=user_object,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -2556,9 +2594,9 @@ class JWTAuthManager:
|
|||
user_id=user_id,
|
||||
requested_model=request_data.get("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=handler.litellm_jwtauth.team_id_upsert,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -2586,7 +2624,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=handler.litellm_jwtauth.team_id_upsert,
|
||||
)
|
||||
elif db_team_fallback and team_id == header_team_id:
|
||||
JWTAuthManager._validate_header_team_in_db_membership(
|
||||
|
|
@ -2596,7 +2634,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,
|
||||
|
|
@ -2606,10 +2644,11 @@ class JWTAuthManager:
|
|||
)
|
||||
|
||||
## MAP USER TO TEAMS
|
||||
await JWTAuthManager.map_user_to_teams(
|
||||
user_object=user_object,
|
||||
team_object=team_object,
|
||||
)
|
||||
if allow_provisioning:
|
||||
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(
|
||||
|
|
@ -2638,3 +2677,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={
|
||||
|
|
|
|||
|
|
@ -11422,6 +11422,7 @@ def _oauth_identity_jwt(
|
|||
issuer: str = "https://idp.example.test",
|
||||
owner: str | None = "jwt-owner",
|
||||
scope: str = "",
|
||||
claims: dict[str, object] | None = None,
|
||||
) -> str:
|
||||
import jwt
|
||||
|
||||
|
|
@ -11434,6 +11435,7 @@ def _oauth_identity_jwt(
|
|||
"aud": audience,
|
||||
"exp": int(time.time()) + expires_in,
|
||||
"scope": scope,
|
||||
**(claims or {}),
|
||||
},
|
||||
signing_key,
|
||||
algorithm="RS256",
|
||||
|
|
@ -11443,12 +11445,14 @@ def _oauth_identity_jwt(
|
|||
@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,
|
||||
|
|
@ -11461,6 +11465,12 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user(
|
|||
|
||||
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(
|
||||
|
|
@ -11525,7 +11535,12 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user(
|
|||
assert response.status_code == 200
|
||||
assert json.loads(response.body)["access_token"] == "upstream-token"
|
||||
users.create.assert_not_awaited()
|
||||
if not policy_allowed or owner_state in ("inactive", "database_error") or (owner_state == "missing" and not admin):
|
||||
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()
|
||||
|
|
@ -11755,9 +11770,7 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity(
|
|||
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 ""
|
||||
)
|
||||
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")
|
||||
|
|
@ -11823,3 +11836,139 @@ async def test_oauth_refresh_revalidates_the_same_active_user_rule(
|
|||
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 _extract_user_id_from_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 _extract_user_id_from_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)
|
||||
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 _extract_user_id_from_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 _extract_user_id_from_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()
|
||||
|
|
|
|||
|
|
@ -6792,10 +6792,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("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, existing_user: bool, model_allowed: bool
|
||||
monkeypatch: pytest.MonkeyPatch, identity_only: bool, allow_provisioning: bool, existing_user: bool, model_allowed: bool
|
||||
) -> None:
|
||||
from litellm.proxy._types import ScopeMapping
|
||||
from litellm.proxy.auth.auth_checks import UserNotFoundError
|
||||
|
|
@ -6841,6 +6842,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users(
|
|||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
identity_only=identity_only,
|
||||
allow_provisioning=allow_provisioning,
|
||||
)
|
||||
if not identity_only and not model_allowed:
|
||||
with pytest.raises(HTTPException) as denial:
|
||||
|
|
@ -6848,7 +6850,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users(
|
|||
assert denial.value.status_code == 403
|
||||
users.create.assert_not_awaited()
|
||||
return
|
||||
if identity_only and not existing_user:
|
||||
if (identity_only or not allow_provisioning) and not existing_user:
|
||||
with pytest.raises(UserNotFoundError):
|
||||
await pending
|
||||
else:
|
||||
|
|
@ -6856,7 +6858,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 existing_user else 1)
|
||||
assert users.create.await_count == (0 if identity_only or not allow_provisioning or existing_user else 1)
|
||||
|
||||
|
||||
def _entra_agent_registry() -> AgentRegistry:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue