mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
Merge pull request #41485 from BerriAI/litellm_jwt_token_exchange_grant
feat(proxy): add RFC 8693 token exchange for IdP JWTs on the gateway token endpoint
This commit is contained in:
commit
078a60478f
18 changed files with 1253 additions and 60 deletions
|
|
@ -262,7 +262,12 @@ async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | No
|
|||
return loaded if isinstance(loaded, str) else None
|
||||
|
||||
|
||||
async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResolutionFailure":
|
||||
UserRowSource = Literal["cache", "database"]
|
||||
|
||||
|
||||
async def load_active_user_by_id(
|
||||
user_id: str, source: UserRowSource = "cache"
|
||||
) -> "LiteLLM_UserTable | _KeyResolutionFailure":
|
||||
"""Load a live litellm user by id, returning the record when the user is active or a precise
|
||||
failure otherwise. The interactive DCR client authenticates via SSO, so its refresh envelope seals a
|
||||
user subject; renewing it must re-check the user is still live (present and not SCIM-deactivated) so a
|
||||
|
|
@ -273,7 +278,11 @@ async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResol
|
|||
``HTTPException``, a SCIM-deactivated user, and, unlike the key path, a missing user. ``get_user_object``
|
||||
catches every DB failure and re-raises a bare ``ValueError`` (a deleted user and a real outage look
|
||||
identical, the original error surviving only as ``__context__``), so the outage check walks the cause
|
||||
chain, and a missing user falls through to ``no_active_key`` rather than an opaque gateway fault."""
|
||||
chain, and a missing user falls through to ``no_active_key`` rather than an opaque gateway fault.
|
||||
``source="database"`` reads the row from the database, never the cache, so the credential mint refuses
|
||||
a user that a writer deactivated or deleted without evicting the cached row, and it leaves the fresh
|
||||
row in the cache for the requests the credential makes next. Every other caller keeps the cache read,
|
||||
so introspection, which a resource server may call per request, stays off the database."""
|
||||
from litellm.proxy._types import (
|
||||
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
)
|
||||
|
|
@ -296,6 +305,7 @@ async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResol
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=source == "database",
|
||||
)
|
||||
except (ProxyException, HTTPException):
|
||||
return "no_active_key"
|
||||
|
|
|
|||
|
|
@ -59,6 +59,11 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
|
|||
register_aggregate_client,
|
||||
relative_request_url,
|
||||
revoke_refresh_token,
|
||||
supported_grant_types,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.idp_token_exchange import (
|
||||
exchange_idp_subject_token,
|
||||
token_exchange_available,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
|
||||
RefreshOwnershipProven,
|
||||
|
|
@ -1980,6 +1985,9 @@ async def token_endpoint(
|
|||
refresh_token: str | None = Form(None),
|
||||
scope: str | None = Form(None),
|
||||
resource: str | None = Form(None),
|
||||
subject_token: str | None = Form(None),
|
||||
subject_token_type: str | None = Form(None),
|
||||
requested_token_type: str | None = Form(None),
|
||||
mcp_server_name: str | None = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -2010,6 +2018,10 @@ async def token_endpoint(
|
|||
cache=user_api_key_cache,
|
||||
resource=resource,
|
||||
mint_proxy_credential=mint_proxy_credential,
|
||||
subject_token=subject_token,
|
||||
subject_token_type=subject_token_type,
|
||||
requested_token_type=requested_token_type,
|
||||
exchange_subject_token=exchange_idp_subject_token,
|
||||
)
|
||||
|
||||
lookup_name: Final = mcp_server_name or client_id
|
||||
|
|
@ -2131,7 +2143,9 @@ async def introspect_endpoint(token: str = Form(...)) -> Response:
|
|||
async def native_client_auth_discovery(request: Request) -> JSONResponse:
|
||||
"""The versioned contract a native client (``lite login --pkce``, or a CLI in any other
|
||||
language) reads to sign a user in through the browser and obtain a proxy credential."""
|
||||
return JSONResponse(native_client_auth_contract(request), headers=TOKEN_NO_CACHE_HEADERS)
|
||||
return JSONResponse(
|
||||
native_client_auth_contract(request, token_exchange_available()), headers=TOKEN_NO_CACHE_HEADERS
|
||||
)
|
||||
|
||||
|
||||
# Per RFC 6749 §4.1.2.1, an IdP that rejects an OAuth authorization request
|
||||
|
|
@ -2619,7 +2633,7 @@ def _build_aggregate_protected_resource_response(request: Request) -> dict:
|
|||
}
|
||||
|
||||
|
||||
def _build_aggregate_authorization_server_response(request: Request) -> dict:
|
||||
def _build_aggregate_authorization_server_response(request: Request, token_exchange_available: bool) -> dict:
|
||||
"""RFC 8414 metadata for the gateway as the aggregate authorization server.
|
||||
|
||||
The issuer is ``{base}/mcp`` and must stay equal to the value the
|
||||
|
|
@ -2638,7 +2652,7 @@ def _build_aggregate_authorization_server_response(request: Request) -> dict:
|
|||
"registration_endpoint": f"{request_base_url}/register",
|
||||
"response_types_supported": ["code"],
|
||||
"scopes_supported": [],
|
||||
"grant_types_supported": ["authorization_code", "refresh_token"],
|
||||
"grant_types_supported": supported_grant_types(token_exchange_available),
|
||||
"code_challenge_methods_supported": ["S256"],
|
||||
"token_endpoint_auth_methods_supported": ["none", "client_secret_post"],
|
||||
}
|
||||
|
|
@ -2676,7 +2690,7 @@ async def oauth_authorization_server_aggregate(request: Request):
|
|||
per-server row win here instead would serve an issuer of {base} against a resource that
|
||||
advertised {base}/mcp, which fails the RFC 8414 issuer check and breaks the front door.
|
||||
"""
|
||||
return _build_aggregate_authorization_server_response(request)
|
||||
return _build_aggregate_authorization_server_response(request, token_exchange_available())
|
||||
|
||||
|
||||
# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name}
|
||||
|
|
@ -2902,7 +2916,9 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
# advertises that), so this does not affect it. A request without redirect_uris is not
|
||||
# a DCR request, so the legacy single-server-or-dummy fallback is kept for it.
|
||||
if data.get("redirect_uris"):
|
||||
return await register_aggregate_client(request=request, request_body=data)
|
||||
return await register_aggregate_client(
|
||||
request=request, request_body=data, token_exchange_available=token_exchange_available()
|
||||
)
|
||||
resolved: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if resolved:
|
||||
return await register_client_with_server(
|
||||
|
|
|
|||
|
|
@ -51,7 +51,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
|||
from fastapi import HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict, assert_never
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -187,6 +187,52 @@ class MintProxyCredential(Protocol):
|
|||
) -> Awaitable[MintedProxyCredential | ProxyCredentialMintFailure]: ...
|
||||
|
||||
|
||||
TOKEN_EXCHANGE_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:token-exchange"
|
||||
|
||||
|
||||
def supported_grant_types(token_exchange_available: bool) -> tuple[str, ...]:
|
||||
"""The grants ``/token`` can serve on this deployment. The RFC 8693 exchange is listed
|
||||
only where the JWT auth that proves a subject token is on, backed by a database, and
|
||||
licensed, so a client never selects a grant the gateway would then refuse."""
|
||||
if token_exchange_available:
|
||||
return ("authorization_code", "refresh_token", TOKEN_EXCHANGE_GRANT_TYPE)
|
||||
return ("authorization_code", "refresh_token")
|
||||
|
||||
|
||||
"""RFC 8693: a native client that already holds a token from the customer's identity
|
||||
provider trades it for the proxy-API credential without a browser round trip."""
|
||||
|
||||
_IssuedTokenType = Literal["urn:ietf:params:oauth:token-type:access_token"]
|
||||
ACCESS_TOKEN_TOKEN_TYPE: Final[_IssuedTokenType] = "urn:ietf:params:oauth:token-type:access_token"
|
||||
SUBJECT_TOKEN_TYPES: Final = frozenset(
|
||||
{
|
||||
"urn:ietf:params:oauth:token-type:jwt",
|
||||
"urn:ietf:params:oauth:token-type:id_token",
|
||||
ACCESS_TOKEN_TOKEN_TYPE,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class SubjectIdentity(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
user_id: str = Field(min_length=1)
|
||||
team_id: str | None = None
|
||||
|
||||
|
||||
class SubjectTokenRefusal(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
error: Literal["unsupported_grant_type", "invalid_request", "temporarily_unavailable"]
|
||||
description: str = Field(min_length=1)
|
||||
|
||||
|
||||
class ExchangeSubjectToken(Protocol):
|
||||
"""Injected RFC 8693 subject-token verifier ``(subject_token, request)``: proves the
|
||||
IdP token the way the proxy's own JWT auth does and names the litellm user and team it
|
||||
stands for, or says why this gateway will not take it."""
|
||||
|
||||
def __call__(self, subject_token: str, request: Request, /) -> Awaitable[SubjectIdentity | SubjectTokenRefusal]: ...
|
||||
|
||||
|
||||
class ConsentTeam(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
team_id: str = Field(min_length=1)
|
||||
|
|
@ -213,6 +259,12 @@ async def _refuse_proxy_credential(user_id: str, team_id: str | None) -> ProxyCr
|
|||
return "unresolvable"
|
||||
|
||||
|
||||
async def _refuse_subject_token(subject_token: str, request: Request) -> SubjectTokenRefusal:
|
||||
return SubjectTokenRefusal(
|
||||
error="unsupported_grant_type", description="this gateway is not configured to exchange IdP tokens"
|
||||
)
|
||||
|
||||
|
||||
async def _unavailable_vendor_credential(user_id: str, server_id: str) -> VendorCredentialState:
|
||||
return "unavailable"
|
||||
|
||||
|
|
@ -318,7 +370,9 @@ def open_gateway_dcr_client(client_id: str) -> GatewayDcrClient | None:
|
|||
return _open_sealed(client_id, GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient, _CLIENT_RECORD_DEBUG_KEY)
|
||||
|
||||
|
||||
async def register_aggregate_client(request: Request, request_body: Mapping[str, object]) -> Response:
|
||||
async def register_aggregate_client(
|
||||
request: Request, request_body: Mapping[str, object], token_exchange_available: bool
|
||||
) -> Response:
|
||||
"""RFC 7591 dynamic registration against the gateway itself, statelessly.
|
||||
|
||||
Only ``redirect_uris`` is authoritative; every client is registered as a public
|
||||
|
|
@ -382,7 +436,7 @@ async def register_aggregate_client(request: Request, request_body: Mapping[str,
|
|||
"client_id_issued_at": int(now.timestamp()),
|
||||
"redirect_uris": list(raw_uris),
|
||||
"token_endpoint_auth_method": "none",
|
||||
"grant_types": ["authorization_code", "refresh_token"],
|
||||
"grant_types": list(supported_grant_types(token_exchange_available)),
|
||||
"response_types": ["code"],
|
||||
},
|
||||
)
|
||||
|
|
@ -580,7 +634,7 @@ class NativeClientAuthContract(TypedDict):
|
|||
revocation_endpoint_auth_methods_supported: ReadOnly[tuple[str, ...]]
|
||||
|
||||
|
||||
def native_client_auth_contract(request: Request) -> NativeClientAuthContract:
|
||||
def native_client_auth_contract(request: Request, token_exchange_available: bool) -> NativeClientAuthContract:
|
||||
"""The versioned discovery document at ``/.well-known/litellm-cli-auth``: everything a
|
||||
native client (in any language) needs to run the sign-in without reading LiteLLM
|
||||
source. ``resource`` is the exact value to send as the RFC 8707 ``resource`` parameter
|
||||
|
|
@ -595,7 +649,7 @@ def native_client_auth_contract(request: Request) -> NativeClientAuthContract:
|
|||
"revocation_endpoint": f"{base_url}/revoke",
|
||||
"resource": base_url,
|
||||
"response_types_supported": ("code",),
|
||||
"grant_types_supported": ("authorization_code", "refresh_token"),
|
||||
"grant_types_supported": supported_grant_types(token_exchange_available),
|
||||
"code_challenge_methods_supported": ("S256",),
|
||||
"token_endpoint_auth_methods_supported": ("none",),
|
||||
"revocation_endpoint_auth_methods_supported": ("none",),
|
||||
|
|
@ -1033,20 +1087,26 @@ class _ProxyCredentialTokenResponse(TypedDict):
|
|||
refresh_token: ReadOnly[str]
|
||||
user_id: ReadOnly[str]
|
||||
team_id: ReadOnly[str | None]
|
||||
issued_token_type: NotRequired[ReadOnly[_IssuedTokenType]]
|
||||
|
||||
|
||||
def _proxy_credential_response(
|
||||
minted: MintedProxyCredential, principal: SessionPrincipal, keys: SessionSigningKeys, now: datetime
|
||||
minted: MintedProxyCredential,
|
||||
principal: SessionPrincipal,
|
||||
keys: SessionSigningKeys,
|
||||
now: datetime,
|
||||
issued_token_type: _IssuedTokenType | None = None,
|
||||
) -> Response:
|
||||
"""The proxy-API token response: the access token is the very credential ``lite
|
||||
login`` stores (accepted on every proxy route with user and team attribution), and
|
||||
the refresh token is a gateway-sealed rotating token bound to the team the credential
|
||||
was minted for, so a renewal keeps the team the user consented to."""
|
||||
was minted for, so a renewal keeps the team the user consented to. A token exchange
|
||||
also states ``issued_token_type``, which RFC 8693 section 2.2.1 requires."""
|
||||
bound_principal: Final = principal.model_copy(update=MappingProxyType({"team_id": minted.team_id}))
|
||||
refresh: Final = mint_session_refresh_token(bound_principal, keys, now)
|
||||
if not isinstance(refresh, MintedSessionToken):
|
||||
return _oauth_error(500, "server_error", "failed to mint the session credential")
|
||||
body: Final[_ProxyCredentialTokenResponse] = {
|
||||
credential: Final[_ProxyCredentialTokenResponse] = {
|
||||
"access_token": minted.key,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": minted.expires_in,
|
||||
|
|
@ -1054,7 +1114,10 @@ def _proxy_credential_response(
|
|||
"user_id": minted.user_id,
|
||||
"team_id": minted.team_id,
|
||||
}
|
||||
return JSONResponse(status_code=200, content=body, headers=TOKEN_NO_CACHE_HEADERS)
|
||||
if issued_token_type is None:
|
||||
return JSONResponse(status_code=200, content=credential, headers=TOKEN_NO_CACHE_HEADERS)
|
||||
exchanged: Final[_ProxyCredentialTokenResponse] = {**credential, "issued_token_type": issued_token_type}
|
||||
return JSONResponse(status_code=200, content=exchanged, headers=TOKEN_NO_CACHE_HEADERS)
|
||||
|
||||
|
||||
def _reload_failure_response(failure: ReloadUserFailure) -> Response:
|
||||
|
|
@ -1073,6 +1136,16 @@ def _reload_failure_response(failure: ReloadUserFailure) -> Response:
|
|||
assert_never(failure)
|
||||
|
||||
|
||||
def _subject_token_refusal_response(refusal: SubjectTokenRefusal) -> Response:
|
||||
match refusal.error:
|
||||
case "temporarily_unavailable":
|
||||
return _oauth_error(503, refusal.error, refusal.description)
|
||||
case "unsupported_grant_type" | "invalid_request":
|
||||
return _oauth_error(400, refusal.error, refusal.description)
|
||||
case _:
|
||||
assert_never(refusal.error)
|
||||
|
||||
|
||||
def _mint_failure_response(failure: ProxyCredentialMintFailure) -> Response:
|
||||
match failure:
|
||||
case "not_a_member":
|
||||
|
|
@ -1116,11 +1189,16 @@ async def aggregate_token(
|
|||
cache: DualCache,
|
||||
resource: str | None = None,
|
||||
mint_proxy_credential: MintProxyCredential = _refuse_proxy_credential,
|
||||
subject_token: str | None = None,
|
||||
subject_token_type: str | None = None,
|
||||
requested_token_type: str | None = None,
|
||||
exchange_subject_token: ExchangeSubjectToken = _refuse_subject_token,
|
||||
) -> Response:
|
||||
"""The aggregate token verb: authorization_code and refresh_token grants for the
|
||||
identity-only session pair, or for the proxy-API credential when the grant was issued
|
||||
with that audience. Every path re-validates the litellm user live before minting, so a
|
||||
deactivated user cannot obtain or renew a session."""
|
||||
with that audience, and the RFC 8693 token exchange that turns an IdP token straight
|
||||
into the proxy-API credential. Every path re-validates the litellm user live before
|
||||
minting, so a deactivated user cannot obtain or renew a session."""
|
||||
if master_key is None:
|
||||
verbose_logger.error("mcp_gateway_dcr token grant rejected: no master_key configured")
|
||||
return _oauth_error(500, "server_error", "the gateway has no master key configured")
|
||||
|
|
@ -1159,7 +1237,20 @@ async def aggregate_token(
|
|||
now=now,
|
||||
issue=issue,
|
||||
)
|
||||
return _oauth_error(400, "unsupported_grant_type", "grant_type must be authorization_code or refresh_token")
|
||||
if grant_type == TOKEN_EXCHANGE_GRANT_TYPE:
|
||||
return await _token_exchange_grant(
|
||||
subject_token=subject_token,
|
||||
subject_token_type=subject_token_type,
|
||||
requested_token_type=requested_token_type,
|
||||
client_id=client_id,
|
||||
exchange_subject_token=exchange_subject_token,
|
||||
issue=issue,
|
||||
)
|
||||
return _oauth_error(
|
||||
400,
|
||||
"unsupported_grant_type",
|
||||
f"grant_type must be authorization_code, refresh_token, or {TOKEN_EXCHANGE_GRANT_TYPE}",
|
||||
)
|
||||
|
||||
|
||||
class _GrantIssuer:
|
||||
|
|
@ -1211,10 +1302,9 @@ class _GrantIssuer:
|
|||
async def _issue_proxy_credential(
|
||||
self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str
|
||||
) -> Response:
|
||||
if self._resource is not None and not is_proxy_api_resource(self._request, self._resource):
|
||||
return _oauth_error(
|
||||
400, "invalid_target", "resource does not match the proxy API this grant was issued for"
|
||||
)
|
||||
target_refusal: Final = self._proxy_api_target_refusal()
|
||||
if target_refusal is not None:
|
||||
return target_refusal
|
||||
minted: Final = await self._mint_proxy_credential(principal.user_id, principal.team_id)
|
||||
if not isinstance(minted, MintedProxyCredential):
|
||||
return _mint_failure_response(minted)
|
||||
|
|
@ -1223,6 +1313,33 @@ class _GrantIssuer:
|
|||
return refusal
|
||||
return _proxy_credential_response(minted, principal, self._keys, self._now)
|
||||
|
||||
async def exchange(
|
||||
self, subject_token: str, client_id: str, exchange_subject_token: ExchangeSubjectToken
|
||||
) -> Response:
|
||||
"""The RFC 8693 tail: prove the IdP token, then mint. No single-use marker, because
|
||||
the subject token stays a valid proof for as long as the IdP says it is and every
|
||||
exchange mints a fresh credential and refresh token of its own."""
|
||||
target_refusal: Final = self._proxy_api_target_refusal()
|
||||
if target_refusal is not None:
|
||||
return target_refusal
|
||||
identity: Final = await exchange_subject_token(subject_token, self._request)
|
||||
if isinstance(identity, SubjectTokenRefusal):
|
||||
return _subject_token_refusal_response(identity)
|
||||
principal: Final = SessionPrincipal(
|
||||
user_id=identity.user_id, client_id=client_id, audience=PROXY_API_AUDIENCE, team_id=identity.team_id
|
||||
)
|
||||
minted: Final = await self._mint_proxy_credential(principal.user_id, principal.team_id)
|
||||
if not isinstance(minted, MintedProxyCredential):
|
||||
return _mint_failure_response(minted)
|
||||
return _proxy_credential_response(
|
||||
minted, principal, self._keys, self._now, issued_token_type=ACCESS_TOKEN_TOKEN_TYPE
|
||||
)
|
||||
|
||||
def _proxy_api_target_refusal(self) -> Response | None:
|
||||
if self._resource is None or is_proxy_api_resource(self._request, self._resource):
|
||||
return None
|
||||
return _oauth_error(400, "invalid_target", "resource does not match the proxy API this grant was issued for")
|
||||
|
||||
async def _claim_refusal(self, claim_key: str, claim_ttl_seconds: int, replayed: str) -> Response | None:
|
||||
return _claim_refusal(
|
||||
await self._guard.claim(claim_key, claim_ttl_seconds), replayed=_oauth_error(400, "invalid_grant", replayed)
|
||||
|
|
@ -1297,6 +1414,32 @@ async def _refresh_token_grant(
|
|||
)
|
||||
|
||||
|
||||
async def _token_exchange_grant(
|
||||
subject_token: str | None,
|
||||
subject_token_type: str | None,
|
||||
requested_token_type: str | None,
|
||||
client_id: str,
|
||||
exchange_subject_token: ExchangeSubjectToken,
|
||||
issue: _GrantIssuer,
|
||||
) -> Response:
|
||||
"""RFC 8693 token exchange for a registered native client that already holds an IdP
|
||||
token: the gateway proves the token the way its JWT auth does and answers with the
|
||||
proxy-API credential, so a fresh laptop with only an IdP login gets a gateway key
|
||||
without a browser round trip. The client must be registered because the refresh token
|
||||
in the answer is bound to it."""
|
||||
if not is_gateway_dcr_client_id(client_id) or open_gateway_dcr_client(client_id) is None:
|
||||
return _oauth_error(401, "invalid_client", "unknown or malformed client_id")
|
||||
if not subject_token or not subject_token_type:
|
||||
return _oauth_error(400, "invalid_request", "subject_token and subject_token_type are required")
|
||||
if subject_token_type not in SUBJECT_TOKEN_TYPES:
|
||||
return _oauth_error(
|
||||
400, "invalid_request", f"subject_token_type must be one of {', '.join(sorted(SUBJECT_TOKEN_TYPES))}"
|
||||
)
|
||||
if requested_token_type is not None and requested_token_type != ACCESS_TOKEN_TOKEN_TYPE:
|
||||
return _oauth_error(400, "invalid_request", f"requested_token_type must be {ACCESS_TOKEN_TOKEN_TYPE}")
|
||||
return await issue.exchange(subject_token, client_id, exchange_subject_token)
|
||||
|
||||
|
||||
async def revoke_refresh_token(token: str, client_id: str, master_key: str | None, cache: DualCache) -> Response:
|
||||
"""RFC 7009 revocation for the gateway's refresh tokens: burn the presented token's
|
||||
``jti`` so neither the holder nor a thief can rotate it again. Access tokens are
|
||||
|
|
|
|||
217
litellm/proxy/_experimental/mcp_server/idp_token_exchange.py
Normal file
217
litellm/proxy/_experimental/mcp_server/idp_token_exchange.py
Normal file
|
|
@ -0,0 +1,217 @@
|
|||
"""The identity-provider side of the RFC 8693 token exchange on ``POST /token``: a native
|
||||
client that already holds a JWT from the customer's IdP trades it for the same proxy-API
|
||||
credential ``lite login`` stores, proven by the proxy's own JWT auth (signature, claims,
|
||||
and the user and team sync it performs), so no browser round trip is needed."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Literal, Protocol
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import SubjectIdentity, SubjectTokenRefusal
|
||||
from litellm.proxy._types import JWTAuthBuilderResult, ProxyException
|
||||
from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
|
||||
EXCHANGE_ROUTE: Final = "/token"
|
||||
REJECTED_SUBJECT_TOKEN: Final = "subject_token was rejected by the gateway's JWT auth"
|
||||
SUBJECT_TOKEN_CHECK_UNAVAILABLE: Final = (
|
||||
"the gateway could not verify subject_token because its identity provider or database is unavailable; retry"
|
||||
)
|
||||
SUBJECT_TOKEN_CHECK_FAULTED: Final = (
|
||||
"the gateway could not verify subject_token because its database reported a fault that is not a transient "
|
||||
"outage; retrying will not help until the gateway deployment is repaired"
|
||||
)
|
||||
GatewayOutage = Literal["retryable", "faulted"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TokenExchangePrerequisites:
|
||||
"""The deployment-level gates ``user_api_key_auth`` applies before it verifies any JWT
|
||||
bearer, plus the JWT-to-virtual-key mapping it consults first: a gateway that maps
|
||||
tokens authenticates a JWT as its mapped key, with that key's models and budget, or
|
||||
refuses an unmapped one, and the exchange proves the token through ``auth_builder``
|
||||
alone, so it would mint the user's own credential past that policy. Discovery and
|
||||
registration advertise the exchange grant only when every gate holds, and an exchange
|
||||
attempt is refused naming the first one that does not."""
|
||||
|
||||
jwt_auth_enabled: bool
|
||||
has_database: bool
|
||||
licensed: bool
|
||||
maps_jwts_to_virtual_keys: bool
|
||||
|
||||
@property
|
||||
def available(self) -> bool:
|
||||
return self.jwt_auth_enabled and self.has_database and self.licensed and not self.maps_jwts_to_virtual_keys
|
||||
|
||||
def refusal(self) -> SubjectTokenRefusal | None:
|
||||
if not self.jwt_auth_enabled:
|
||||
return SubjectTokenRefusal(
|
||||
error="unsupported_grant_type",
|
||||
description="JWT auth is not enabled on this gateway, so it cannot exchange IdP tokens",
|
||||
)
|
||||
if not self.has_database:
|
||||
return SubjectTokenRefusal(
|
||||
error="unsupported_grant_type",
|
||||
description="this gateway has no database, so it cannot exchange IdP tokens",
|
||||
)
|
||||
if not self.licensed:
|
||||
return SubjectTokenRefusal(
|
||||
error="unsupported_grant_type",
|
||||
description="JWT auth is an enterprise only feature; no license is set",
|
||||
)
|
||||
if self.maps_jwts_to_virtual_keys:
|
||||
return SubjectTokenRefusal(
|
||||
error="unsupported_grant_type",
|
||||
description="this gateway maps IdP tokens to virtual keys, which the exchange does not serve",
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def read_token_exchange_prerequisites() -> TokenExchangePrerequisites:
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # rebound after startup, so read them per call
|
||||
general_settings,
|
||||
jwt_handler,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
)
|
||||
|
||||
return TokenExchangePrerequisites(
|
||||
jwt_auth_enabled=general_settings.get("enable_jwt_auth", False) is True,
|
||||
has_database=prisma_client is not None,
|
||||
licensed=premium_user is True,
|
||||
maps_jwts_to_virtual_keys=_maps_jwts_to_virtual_keys(jwt_handler),
|
||||
)
|
||||
|
||||
|
||||
def _maps_jwts_to_virtual_keys(jwt_handler: JWTHandler) -> bool:
|
||||
if not hasattr(jwt_handler, "litellm_jwtauth"):
|
||||
return False
|
||||
return jwt_handler.litellm_jwtauth.is_virtual_key_mapping_configured()
|
||||
|
||||
|
||||
def token_exchange_available() -> bool:
|
||||
return read_token_exchange_prerequisites().available
|
||||
|
||||
|
||||
class AuthorizeSubjectToken(Protocol):
|
||||
"""Injected JWT authorization ``(subject_token, request_headers)``: the proxy's
|
||||
``JWTAuthManager.auth_builder`` in production, which raises when the token is not
|
||||
acceptable and otherwise names the user and team it resolved."""
|
||||
|
||||
def __call__(
|
||||
self, subject_token: str, request_headers: Mapping[str, str], /
|
||||
) -> Awaitable[JWTAuthBuilderResult]: ...
|
||||
|
||||
|
||||
async def exchange_idp_subject_token(subject_token: str, request: Request) -> SubjectIdentity | SubjectTokenRefusal:
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # rebound after startup, so read them per call
|
||||
general_settings,
|
||||
jwt_handler,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
async def authorize(token: str, request_headers: Mapping[str, str]) -> JWTAuthBuilderResult:
|
||||
return await JWTAuthManager.auth_builder(
|
||||
api_key=token,
|
||||
jwt_handler=jwt_handler,
|
||||
request_data={},
|
||||
general_settings=general_settings,
|
||||
route=EXCHANGE_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=request_headers,
|
||||
request_method="POST",
|
||||
)
|
||||
|
||||
return await identity_from_subject_token(
|
||||
subject_token,
|
||||
request_headers=request.headers,
|
||||
prerequisites=read_token_exchange_prerequisites(),
|
||||
is_jwt=jwt_handler.is_jwt,
|
||||
authorize=authorize,
|
||||
)
|
||||
|
||||
|
||||
async def identity_from_subject_token(
|
||||
subject_token: str,
|
||||
request_headers: Mapping[str, str],
|
||||
prerequisites: TokenExchangePrerequisites,
|
||||
is_jwt: Callable[[str], bool],
|
||||
authorize: AuthorizeSubjectToken,
|
||||
) -> SubjectIdentity | SubjectTokenRefusal:
|
||||
"""Apply the same gates ``user_api_key_auth`` applies to a JWT bearer, then let the
|
||||
proxy's JWT auth prove the token. A rejection comes back as ``invalid_request``, which
|
||||
RFC 8693 section 2.2.2 prescribes for an invalid or unacceptable subject token, and a
|
||||
check the gateway could not complete (the IdP's JWKS unreachable with no cached copy,
|
||||
the auth database down) as ``temporarily_unavailable``, so the client retries instead
|
||||
of treating a valid token as bad, worded by whether retrying can help. The reason stays
|
||||
in the proxy log: this endpoint is public and JWT auth's own wording can name the JWKS
|
||||
URL it fetched or quote the IdP's response."""
|
||||
unmet: Final = prerequisites.refusal()
|
||||
if unmet is not None:
|
||||
return unmet
|
||||
if not is_jwt(subject_token):
|
||||
return SubjectTokenRefusal(error="invalid_request", description="subject_token is not a JWT")
|
||||
try:
|
||||
result: Final = await authorize(subject_token, request_headers)
|
||||
except HTTPException as denied:
|
||||
return _refusal_for(denied, denied.detail)
|
||||
except ProxyException as denied:
|
||||
return _refusal_for(denied, denied.message)
|
||||
except Exception as denied: # noqa: BLE001 # auth_jwt raises a plain Exception on signature and claim failures
|
||||
return _refusal_for(denied, denied)
|
||||
user_id: Final = result["user_id"]
|
||||
if user_id is None:
|
||||
return SubjectTokenRefusal(error="invalid_request", description="subject_token names no user the gateway knows")
|
||||
return SubjectIdentity(user_id=user_id, team_id=result["team_id"])
|
||||
|
||||
|
||||
def _refusal_for(denied: Exception, reason: object) -> SubjectTokenRefusal:
|
||||
outage: Final = _gateway_could_not_verify(denied)
|
||||
if outage is None:
|
||||
verbose_proxy_logger.warning("token exchange refused a subject_token: %s", reason)
|
||||
return SubjectTokenRefusal(error="invalid_request", description=REJECTED_SUBJECT_TOKEN)
|
||||
verbose_proxy_logger.error("token exchange could not verify a subject_token, %s: %s", outage, reason)
|
||||
return SubjectTokenRefusal(error="temporarily_unavailable", description=_check_unavailable_description(outage))
|
||||
|
||||
|
||||
def _check_unavailable_description(outage: GatewayOutage) -> str:
|
||||
match outage:
|
||||
case "retryable":
|
||||
return SUBJECT_TOKEN_CHECK_UNAVAILABLE
|
||||
case "faulted":
|
||||
return SUBJECT_TOKEN_CHECK_FAULTED
|
||||
case _:
|
||||
assert_never(outage)
|
||||
|
||||
|
||||
def _gateway_could_not_verify(denied: Exception) -> GatewayOutage | None:
|
||||
"""A database fault anywhere in the chain (``get_user_object`` wraps prisma failures in a
|
||||
bare ``ValueError``) or a 5xx from JWT auth (the IdP's JWKS unreachable with no cached
|
||||
copy) is the gateway failing, not the token. A fault retrying cannot clear (a missing or
|
||||
version-skewed query engine) is named as such, the way the mint path words it, so the
|
||||
client is not told to wait on a deployment that needs repair."""
|
||||
fault: Final = PrismaDBExceptionHandler.find_database_service_unavailable_error_in_chain(denied)
|
||||
if fault is not None:
|
||||
return "faulted" if PrismaDBExceptionHandler.is_permanent_database_fault(fault) else "retryable"
|
||||
return "retryable" if _is_server_error(denied) else None
|
||||
|
||||
|
||||
def _is_server_error(denied: Exception) -> bool:
|
||||
match denied:
|
||||
case HTTPException(status_code=status_code):
|
||||
return status_code >= 500
|
||||
case ProxyException(code=code):
|
||||
return code.isdigit() and int(code) >= 500
|
||||
case _:
|
||||
return False
|
||||
|
|
@ -16,7 +16,7 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
|
|||
ReloadUserFailure,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, effective_user_role
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
CliSsoTeamDetail,
|
||||
fetch_cli_sso_team_details,
|
||||
|
|
@ -42,7 +42,7 @@ async def mint_proxy_credential(
|
|||
user_id: str, team_id: str | None
|
||||
) -> MintedProxyCredential | ProxyCredentialMintFailure:
|
||||
"""Mint the ``lite login`` credential for a consented grant. Membership is checked
|
||||
live, so a team the user left between consent and redemption (or between refreshes)
|
||||
live against the database row, so a team the user left between consent and redemption (or between refreshes)
|
||||
refuses the grant instead of minting a credential attributed to a team they are no
|
||||
longer on. The team is exactly the one the consent page sealed into the grant; nothing
|
||||
is picked on the user's behalf here, so a refresh can never move the credential, and a
|
||||
|
|
@ -51,12 +51,12 @@ async def mint_proxy_credential(
|
|||
posting the consent form without one. Memberships whose team rows are gone count as no
|
||||
team at all, the way ``lite login`` treats them, so they can never lock a user out. The
|
||||
user row handed to the minter carries no team list, exactly like ``lite login``'s, so
|
||||
the minter's own first-team fallback stays inert."""
|
||||
user: Final = await load_active_user_by_id(user_id)
|
||||
the minter's own first-team fallback stays inert. The credential carries the role the
|
||||
proxy already enforces for the user on every request, so a row with no role (JWT auth's
|
||||
upsert writes none) mints as an internal user instead of being refused."""
|
||||
user: Final = await load_active_user_by_id(user_id, source="database")
|
||||
if isinstance(user, str):
|
||||
return user
|
||||
if user.user_role is None:
|
||||
return "no_active_key"
|
||||
if team_id is not None and team_id not in user.teams:
|
||||
return "not_a_member"
|
||||
details: Final = await _team_details(user.teams) if user.teams else ()
|
||||
|
|
@ -68,7 +68,9 @@ async def mint_proxy_credential(
|
|||
if selected is None:
|
||||
return "not_a_member"
|
||||
key: Final = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
||||
user_info=LiteLLM_UserTable(user_id=user.user_id, user_role=user.user_role, models=user.models),
|
||||
user_info=LiteLLM_UserTable(
|
||||
user_id=user.user_id, user_role=effective_user_role(user.user_role).value, models=user.models
|
||||
),
|
||||
team_id=team_id,
|
||||
team_alias=selected.team_alias,
|
||||
team_models=selected.team_models,
|
||||
|
|
|
|||
|
|
@ -23684,6 +23684,17 @@
|
|||
],
|
||||
"title": "Refresh Token"
|
||||
},
|
||||
"requested_token_type": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Requested Token Type"
|
||||
},
|
||||
"resource": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -23705,6 +23716,28 @@
|
|||
}
|
||||
],
|
||||
"title": "Scope"
|
||||
},
|
||||
"subject_token": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Subject Token"
|
||||
},
|
||||
"subject_token_type": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Subject Token Type"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
|
|
@ -23758,6 +23791,17 @@
|
|||
],
|
||||
"title": "Refresh Token"
|
||||
},
|
||||
"requested_token_type": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Requested Token Type"
|
||||
},
|
||||
"resource": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -23779,6 +23823,28 @@
|
|||
}
|
||||
],
|
||||
"title": "Scope"
|
||||
},
|
||||
"subject_token": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Subject Token"
|
||||
},
|
||||
"subject_token_type": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Subject Token Type"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
|
|
|
|||
|
|
@ -534,6 +534,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/mcp-rest/tools/call",
|
||||
"/v1/mcp/tools",
|
||||
"/introspect",
|
||||
"/token",
|
||||
]
|
||||
|
||||
# MCP server CRUD routes — control-plane. Gated by DISABLE_ADMIN_ENDPOINTS.
|
||||
|
|
|
|||
|
|
@ -1216,21 +1216,19 @@ async def common_checks(
|
|||
return True
|
||||
|
||||
|
||||
def effective_user_role(user_role: str | None) -> LitellmUserRoles:
|
||||
try:
|
||||
return LitellmUserRoles(user_role)
|
||||
except ValueError:
|
||||
return LitellmUserRoles.INTERNAL_USER
|
||||
|
||||
|
||||
def _get_user_role(
|
||||
user_obj: LiteLLM_UserTable | None,
|
||||
) -> LitellmUserRoles | None:
|
||||
if user_obj is None:
|
||||
return None
|
||||
|
||||
_user: Final = user_obj
|
||||
|
||||
_user_role: Final = _user.user_role
|
||||
try:
|
||||
role: Final = LitellmUserRoles(_user_role)
|
||||
except ValueError:
|
||||
return LitellmUserRoles.INTERNAL_USER
|
||||
|
||||
return role
|
||||
return effective_user_role(user_obj.user_role)
|
||||
|
||||
|
||||
def _is_api_route_allowed(
|
||||
|
|
@ -2577,7 +2575,7 @@ async def get_user_object(
|
|||
raise Exception("No db connected")
|
||||
try:
|
||||
db_access_time_key: Final = f"user_id:{user_id}"
|
||||
should_check_db: Final = _should_check_db(
|
||||
should_check_db: Final = bool(check_db_only) or _should_check_db(
|
||||
key=db_access_time_key,
|
||||
last_db_access_time=last_db_access_time,
|
||||
db_cache_expiry=db_cache_expiry,
|
||||
|
|
|
|||
|
|
@ -1867,7 +1867,7 @@ class JWTAuthManager:
|
|||
|
||||
@staticmethod
|
||||
def get_team_id_from_header(
|
||||
request_headers: dict | None,
|
||||
request_headers: Mapping[str, str] | None,
|
||||
allowed_team_ids: set[str],
|
||||
fallback_to_db_teams: bool = False,
|
||||
) -> str | None:
|
||||
|
|
@ -2037,7 +2037,7 @@ class JWTAuthManager:
|
|||
async def _attach_team_from_header_for_admin(
|
||||
admin_result: JWTAuthBuilderResult,
|
||||
route: str,
|
||||
request_headers: dict | None,
|
||||
request_headers: Mapping[str, str] | None,
|
||||
jwt_handler: JWTHandler,
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
@ -2293,7 +2293,7 @@ class JWTAuthManager:
|
|||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
request_headers: dict | None = None,
|
||||
request_headers: Mapping[str, str] | None = None,
|
||||
request_method: str | None = None,
|
||||
) -> JWTAuthBuilderResult:
|
||||
return await JWTAuthManager.authorize_jwt(
|
||||
|
|
@ -2390,7 +2390,7 @@ class JWTAuthManager:
|
|||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
request_headers: dict[str, str] | None = None,
|
||||
request_headers: Mapping[str, str] | None = None,
|
||||
request_method: str | None = None,
|
||||
provisioning: _JWTProvisioning | None = None,
|
||||
) -> JWTAuthBuilderResult:
|
||||
|
|
|
|||
|
|
@ -3335,6 +3335,7 @@ async def team_member_add(
|
|||
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
premium_user,
|
||||
|
|
@ -3429,6 +3430,10 @@ async def team_member_add(
|
|||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
)
|
||||
|
||||
await evict_and_broadcast(
|
||||
cache_keys=tuple(sorted(user.user_id for user in updated_users)),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
await _evict_created_membership_caches(
|
||||
user_ids=(tm.user_id for tm in updated_team_memberships),
|
||||
team_id=data.team_id,
|
||||
|
|
|
|||
|
|
@ -7568,6 +7568,69 @@ async def test_reload_active_user_by_id_permanent_engine_fault_is_faulted(proxy_
|
|||
assert await _reload_active_user_by_id("sso-user-7") == "faulted"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_active_user_by_id_reads_the_row_from_the_database_not_the_cache(proxy_globals):
|
||||
"""JWT auth caches the user it creates before it adds that user to the JWT's team, and adding a
|
||||
member never evicts the cached row, so a credential minted off the cached row refused the very first
|
||||
token exchange as not a member. The database source has to read the row from the database and leave
|
||||
the fresh row in the cache for the requests the credential makes next."""
|
||||
from litellm.proxy._experimental.mcp_server.bridge_token_flow import load_active_user_by_id
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
cache = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key="fresh-jwt-user", value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=[]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=["team-a"])
|
||||
)
|
||||
proxy_globals.user_api_key_cache = cache
|
||||
proxy_globals.prisma_client = prisma
|
||||
|
||||
loaded = await load_active_user_by_id("fresh-jwt-user", source="database")
|
||||
|
||||
assert not isinstance(loaded, str)
|
||||
assert loaded.teams == ["team-a"]
|
||||
cached = await cache.async_get_cache(key="fresh-jwt-user", model_type=LiteLLM_UserTable)
|
||||
assert cached is not None
|
||||
assert cached.teams == ["team-a"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_active_user_by_id_serves_a_cached_row_without_a_database_read(proxy_globals):
|
||||
"""Introspection and refresh revalidation run per call, so the loader's default source is the cache: a
|
||||
cached row answers without a database read, and only a caller that asks for the database row pays for
|
||||
one."""
|
||||
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
|
||||
_reload_active_user_by_id,
|
||||
load_active_user_by_id,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
cache = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key="cached-jwt-user",
|
||||
value=LiteLLM_UserTable(user_id="cached-jwt-user", teams=["team-a"]),
|
||||
model_type=LiteLLM_UserTable,
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_UserTable(user_id="cached-jwt-user", teams=[])
|
||||
)
|
||||
proxy_globals.user_api_key_cache = cache
|
||||
proxy_globals.prisma_client = prisma
|
||||
|
||||
loaded = await load_active_user_by_id("cached-jwt-user")
|
||||
|
||||
assert not isinstance(loaded, str)
|
||||
assert loaded.teams == ["team-a"]
|
||||
assert await _reload_active_user_by_id("cached-jwt-user") is None
|
||||
prisma.db.litellm_usertable.find_unique.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_uses_client_secret_basic_when_configured():
|
||||
"""LIT-4091: a server with token_endpoint_auth_method=client_secret_basic must send the
|
||||
|
|
@ -11048,6 +11111,43 @@ def test_native_client_login_walks_discovery_consent_token_refresh_and_revoke(mo
|
|||
assert stranger.json()["error"] == "invalid_client"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"jwt_auth_enabled, virtual_key_claim_field, exchange_servable",
|
||||
[(True, None, True), (False, None, False), (True, "client_id", False)],
|
||||
ids=["jwt auth on", "jwt auth off", "jwts mapped to virtual keys"],
|
||||
)
|
||||
def test_discovery_advertises_the_exchange_grant_only_where_the_gateway_can_serve_it(
|
||||
monkeypatch, jwt_auth_enabled, virtual_key_claim_field, exchange_servable
|
||||
):
|
||||
"""Every document a native client reads before it picks a grant (the versioned contract, the
|
||||
aggregate authorization-server metadata, and the registration response) lists the RFC 8693
|
||||
exchange exactly when the running proxy can serve it: JWT auth on, a database, a license, and
|
||||
no JWT-to-virtual-key mapping, since the exchange would mint past the mapped key's policy."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
|
||||
client, _session_cookie, _minted = _native_client_app(monkeypatch)
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=DualCache(),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(virtual_key_claim_field=virtual_key_claim_field),
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.jwt_handler", handler)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": jwt_auth_enabled})
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
exchange_grant = ["urn:ietf:params:oauth:grant-type:token-exchange"] if exchange_servable else []
|
||||
expected = ["authorization_code", "refresh_token", *exchange_grant]
|
||||
|
||||
assert client.get("/.well-known/litellm-cli-auth").json()["grant_types_supported"] == expected
|
||||
assert client.get("/.well-known/oauth-authorization-server/mcp").json()["grant_types_supported"] == expected
|
||||
registered = client.post("/register", json={"redirect_uris": ["http://127.0.0.1:51234/callback"]})
|
||||
assert registered.status_code == 201
|
||||
assert registered.json()["grant_types"] == expected
|
||||
|
||||
|
||||
def test_native_client_authorize_without_the_proxy_resource_keeps_the_mcp_flow(monkeypatch):
|
||||
"""A registered client asking for the MCP resource (or no resource) never sees the consent
|
||||
page, so existing MCP clients are untouched by the native-client arm."""
|
||||
|
|
@ -11847,13 +11947,17 @@ async def test_oauth_refresh_revalidates_the_same_active_user_rule(
|
|||
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id
|
||||
|
||||
handler, _ = jwt_oauth_identity
|
||||
handler.user_api_key_cache.set_cache(
|
||||
"jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", metadata={"scim_active": state != "inactive"})
|
||||
)
|
||||
user_id: Final = f"jwt-owner-{state}"
|
||||
row: Final = LiteLLM_UserTable(user_id=user_id, metadata={"scim_active": state != "inactive"})
|
||||
proxy_server.prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=row)
|
||||
if state == "missing_database":
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
expected: Final = None if state == "active" else "no_active_key" if state == "inactive" else "unresolvable"
|
||||
assert await _reload_active_user_by_id("jwt-owner") == expected
|
||||
assert await _reload_active_user_by_id(user_id) == expected
|
||||
if state != "missing_database":
|
||||
cached: Final = handler.user_api_key_cache.get_cache(user_id, model_type=LiteLLM_UserTable)
|
||||
assert cached is not None
|
||||
assert cached.metadata == row.metadata
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -15,13 +15,18 @@ from starlette.requests import Request
|
|||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
|
||||
_AUTH_CODE_DEBUG_KEY,
|
||||
ACCESS_TOKEN_TOKEN_TYPE,
|
||||
CONNECT_FLOW_COOKIE_PREFIX,
|
||||
GATEWAY_AUTH_CODE_PREFIX,
|
||||
GATEWAY_AUTH_CODE_TTL_SECONDS,
|
||||
MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS,
|
||||
MAX_CLIENT_ID_LENGTH,
|
||||
SUBJECT_TOKEN_TYPES,
|
||||
TOKEN_EXCHANGE_GRANT_TYPE,
|
||||
ConsentTeam,
|
||||
MintedProxyCredential,
|
||||
SubjectIdentity,
|
||||
SubjectTokenRefusal,
|
||||
_GatewayAuthCode,
|
||||
_open_sealed,
|
||||
_seal,
|
||||
|
|
@ -90,9 +95,11 @@ def _request(path="/authorize", query="", cookies=None, method="GET"):
|
|||
)
|
||||
|
||||
|
||||
async def _register(redirect_uris) -> dict:
|
||||
async def _register(redirect_uris, token_exchange_available=True) -> dict:
|
||||
response = await register_aggregate_client(
|
||||
request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris}
|
||||
request=_request(path="/register", method="POST"),
|
||||
request_body={"redirect_uris": redirect_uris},
|
||||
token_exchange_available=token_exchange_available,
|
||||
)
|
||||
return json.loads(response.body)
|
||||
|
||||
|
|
@ -105,6 +112,7 @@ async def _reload_user_active(user_id: str):
|
|||
async def test_register_mints_stateless_public_client():
|
||||
body = await _register([REDIRECT_URI])
|
||||
assert body["token_endpoint_auth_method"] == "none"
|
||||
assert body["grant_types"] == ["authorization_code", "refresh_token", TOKEN_EXCHANGE_GRANT_TYPE]
|
||||
assert "client_secret" not in body
|
||||
assert body["redirect_uris"] == [REDIRECT_URI]
|
||||
assert is_gateway_dcr_client_id(body["client_id"])
|
||||
|
|
@ -113,11 +121,18 @@ async def test_register_mints_stateless_public_client():
|
|||
assert record.redirect_uris == (REDIRECT_URI,)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_omits_the_exchange_grant_where_the_gateway_cannot_serve_it():
|
||||
body = await _register([REDIRECT_URI], token_exchange_available=False)
|
||||
assert body["grant_types"] == ["authorization_code", "refresh_token"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("redirect_uris", [VSCODE_REDIRECT_URIS, MAX_LENGTH_REDIRECT_URIS])
|
||||
async def test_register_four_callbacks_preserves_metadata(redirect_uris: tuple[str, ...]) -> None:
|
||||
response: Final = await register_aggregate_client(
|
||||
request=_request(path="/register", method="POST"),
|
||||
token_exchange_available=True,
|
||||
request_body={
|
||||
"client_name": "Visual Studio Code",
|
||||
"client_uri": "https://code.visualstudio.com",
|
||||
|
|
@ -143,6 +158,7 @@ async def test_register_four_callbacks_preserves_metadata(redirect_uris: tuple[s
|
|||
async def test_register_rejects_five_valid_callbacks() -> None:
|
||||
response: Final = await register_aggregate_client(
|
||||
request=_request(path="/register", method="POST"),
|
||||
token_exchange_available=True,
|
||||
request_body={"redirect_uris": [*VSCODE_REDIRECT_URIS, "http://127.0.0.1:33419/"]},
|
||||
)
|
||||
assert response.status_code == 400
|
||||
|
|
@ -156,6 +172,7 @@ async def test_register_rejects_five_valid_callbacks() -> None:
|
|||
async def test_register_four_callbacks_preserves_encoded_size_guard() -> None:
|
||||
response: Final = await register_aggregate_client(
|
||||
request=_request(path="/register", method="POST"),
|
||||
token_exchange_available=True,
|
||||
request_body={"redirect_uris": [f"https://client.example/{index}/".ljust(256, "é") for index in range(4)]},
|
||||
)
|
||||
assert response.status_code == 400
|
||||
|
|
@ -208,6 +225,7 @@ async def test_register_rejects_userinfo_spoofed_origin():
|
|||
response = await register_aggregate_client(
|
||||
request=_request(path="/register", method="POST"),
|
||||
request_body={"redirect_uris": ["https://claude.ai@attacker.example/callback"]},
|
||||
token_exchange_available=True,
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"] == "invalid_redirect_uri"
|
||||
|
|
@ -228,7 +246,9 @@ async def test_register_rejects_userinfo_spoofed_origin():
|
|||
)
|
||||
async def test_register_rejects_bad_redirect_uris(redirect_uris):
|
||||
response = await register_aggregate_client(
|
||||
request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris}
|
||||
request=_request(path="/register", method="POST"),
|
||||
request_body={"redirect_uris": redirect_uris},
|
||||
token_exchange_available=True,
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"] in ("invalid_redirect_uri", "invalid_client_metadata")
|
||||
|
|
@ -1948,7 +1968,7 @@ async def test_revoke_refuses_unknown_clients_and_a_missing_master_key():
|
|||
|
||||
|
||||
def test_native_client_auth_contract_points_every_endpoint_at_this_proxy():
|
||||
assert json.loads(json.dumps(native_client_auth_contract(_request("/.well-known/litellm-cli-auth")))) == {
|
||||
assert json.loads(json.dumps(native_client_auth_contract(_request("/.well-known/litellm-cli-auth"), True))) == {
|
||||
"contract_version": 1,
|
||||
"issuer": "https://llm.example.com",
|
||||
"authorization_endpoint": "https://llm.example.com/authorize",
|
||||
|
|
@ -1957,13 +1977,22 @@ def test_native_client_auth_contract_points_every_endpoint_at_this_proxy():
|
|||
"revocation_endpoint": "https://llm.example.com/revoke",
|
||||
"resource": "https://llm.example.com",
|
||||
"response_types_supported": ["code"],
|
||||
"grant_types_supported": ["authorization_code", "refresh_token"],
|
||||
"grant_types_supported": [
|
||||
"authorization_code",
|
||||
"refresh_token",
|
||||
"urn:ietf:params:oauth:grant-type:token-exchange",
|
||||
],
|
||||
"code_challenge_methods_supported": ["S256"],
|
||||
"token_endpoint_auth_methods_supported": ["none"],
|
||||
"revocation_endpoint_auth_methods_supported": ["none"],
|
||||
}
|
||||
|
||||
|
||||
def test_native_client_auth_contract_omits_the_exchange_grant_where_the_gateway_cannot_serve_it():
|
||||
contract = native_client_auth_contract(_request("/.well-known/litellm-cli-auth"), False)
|
||||
assert list(contract["grant_types_supported"]) == ["authorization_code", "refresh_token"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"resource, expected",
|
||||
[
|
||||
|
|
@ -2148,3 +2177,180 @@ async def test_gateway_owned_resource_stays_scoped_through_consent_and_refresh(a
|
|||
)
|
||||
assert renewed.status_code == 200
|
||||
assert _opened_principal(json.loads(renewed.body)).resource_server_id == "github-id"
|
||||
|
||||
|
||||
JWT_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:jwt"
|
||||
IDP_TOKEN = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1MSJ9.idp-signature"
|
||||
|
||||
|
||||
class _Exchanger:
|
||||
def __init__(self, result=None):
|
||||
self.calls = []
|
||||
self.result = result
|
||||
|
||||
async def __call__(self, subject_token, request):
|
||||
self.calls.append((subject_token, request.url.path))
|
||||
if self.result is not None:
|
||||
return self.result
|
||||
return SubjectIdentity(user_id="u1", team_id="team-b")
|
||||
|
||||
|
||||
async def _exchange_native(client_id, minter, exchanger, cache=None, **overrides):
|
||||
arguments = {
|
||||
"grant_type": TOKEN_EXCHANGE_GRANT_TYPE,
|
||||
"subject_token": IDP_TOKEN,
|
||||
"subject_token_type": JWT_SUBJECT_TOKEN_TYPE,
|
||||
"exchange_subject_token": exchanger,
|
||||
}
|
||||
return await _redeem_native(None, client_id, minter, cache=cache, **{**arguments, **overrides})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_mints_the_proxy_credential_for_the_idp_subject():
|
||||
"""RFC 8693: a registered native client trades the IdP token it already holds for the
|
||||
same credential the consent flow mints, attributed to the user and team the gateway's
|
||||
JWT auth resolved, with a rotating refresh token bound to that team and the client.
|
||||
The exchange can be repeated while the IdP token lives; nothing is burned."""
|
||||
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
minter, exchanger, cache = _Minter(), _Exchanger(), DualCache()
|
||||
response = await _exchange_native(client_id, minter, exchanger, cache=cache)
|
||||
assert response.status_code == 200
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
body = json.loads(response.body)
|
||||
assert exchanger.calls == [(IDP_TOKEN, "/token")]
|
||||
assert minter.calls == [("u1", "team-b")]
|
||||
assert body["issued_token_type"] == ACCESS_TOKEN_TOKEN_TYPE
|
||||
assert body["access_token"] == "sk-cli-u1"
|
||||
assert body["token_type"] == "Bearer"
|
||||
assert body["expires_in"] == 3600
|
||||
assert (body["user_id"], body["team_id"]) == ("u1", "team-b")
|
||||
principal = _opened_refresh(body["refresh_token"], client_id)
|
||||
assert (principal.user_id, principal.client_id, principal.audience, principal.team_id) == (
|
||||
"u1",
|
||||
client_id,
|
||||
"proxy_api",
|
||||
"team-b",
|
||||
)
|
||||
again = await _exchange_native(client_id, minter, exchanger, cache=cache)
|
||||
assert again.status_code == 200
|
||||
assert json.loads(again.body)["refresh_token"] != body["refresh_token"]
|
||||
assert minter.calls == [("u1", "team-b"), ("u1", "team-b")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exchanged_credential_refreshes_and_rotates_like_a_consented_one():
|
||||
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
minter, cache = _Minter(), DualCache()
|
||||
exchanged = json.loads((await _exchange_native(client_id, minter, _Exchanger(), cache=cache)).body)
|
||||
refreshed = await _refresh_native(exchanged["refresh_token"], client_id, minter, cache)
|
||||
assert refreshed.status_code == 200
|
||||
body = json.loads(refreshed.body)
|
||||
assert "issued_token_type" not in body
|
||||
assert (body["access_token"], body["user_id"], body["team_id"]) == ("sk-cli-u1", "u1", "team-b")
|
||||
assert body["refresh_token"] != exchanged["refresh_token"]
|
||||
assert minter.calls == [("u1", "team-b"), ("u1", "team-b")]
|
||||
replay = await _refresh_native(exchanged["refresh_token"], client_id, minter, cache)
|
||||
assert replay.status_code == 400
|
||||
assert json.loads(replay.body)["error"] == "invalid_grant"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_for_a_teamless_subject_mints_a_teamless_credential():
|
||||
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
minter = _Minter()
|
||||
response = await _exchange_native(client_id, minter, _Exchanger(SubjectIdentity(user_id="u2")))
|
||||
assert response.status_code == 200
|
||||
body = json.loads(response.body)
|
||||
assert minter.calls == [("u2", None)]
|
||||
assert (body["user_id"], body["team_id"]) == ("u2", None)
|
||||
assert _opened_refresh(body["refresh_token"], client_id).team_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("subject_token_type", sorted(SUBJECT_TOKEN_TYPES))
|
||||
async def test_token_exchange_accepts_every_advertised_subject_token_type(subject_token_type):
|
||||
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
response = await _exchange_native(client_id, _Minter(), _Exchanger(), subject_token_type=subject_token_type)
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_without_an_idp_exchanger_is_unsupported():
|
||||
"""A gateway that wires no IdP verifier into the endpoint answers the way it always
|
||||
answered an unknown grant, and never reaches the minter."""
|
||||
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
minter = _Minter()
|
||||
response = await _redeem_native(
|
||||
None,
|
||||
client_id,
|
||||
minter,
|
||||
grant_type=TOKEN_EXCHANGE_GRANT_TYPE,
|
||||
subject_token=IDP_TOKEN,
|
||||
subject_token_type=JWT_SUBJECT_TOKEN_TYPE,
|
||||
)
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"] == "unsupported_grant_type"
|
||||
assert minter.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"overrides, status, error",
|
||||
[
|
||||
({"subject_token": None}, 400, "invalid_request"),
|
||||
({"subject_token": ""}, 400, "invalid_request"),
|
||||
({"subject_token_type": None}, 400, "invalid_request"),
|
||||
({"subject_token_type": "urn:ietf:params:oauth:token-type:saml2"}, 400, "invalid_request"),
|
||||
({"requested_token_type": "urn:ietf:params:oauth:token-type:refresh_token"}, 400, "invalid_request"),
|
||||
({"resource": "https://other.example.com"}, 400, "invalid_target"),
|
||||
({"resource": "https://llm.example.com/mcp"}, 400, "invalid_target"),
|
||||
({"client_id": "llm_dcrc_forged"}, 401, "invalid_client"),
|
||||
({"client_id": "not-a-gateway-client"}, 401, "invalid_client"),
|
||||
],
|
||||
)
|
||||
async def test_token_exchange_refuses_a_malformed_request_before_touching_the_idp_token(overrides, status, error):
|
||||
registered = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
minter, exchanger = _Minter(), _Exchanger()
|
||||
response = await _exchange_native(
|
||||
overrides.get("client_id", registered),
|
||||
minter,
|
||||
exchanger,
|
||||
**{name: value for name, value in overrides.items() if name != "client_id"},
|
||||
)
|
||||
assert response.status_code == status
|
||||
assert json.loads(response.body)["error"] == error
|
||||
assert exchanger.calls == []
|
||||
assert minter.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"error, status",
|
||||
[("unsupported_grant_type", 400), ("invalid_request", 400), ("temporarily_unavailable", 503)],
|
||||
)
|
||||
async def test_token_exchange_relays_the_idp_refusal_and_never_mints(error, status):
|
||||
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
minter = _Minter()
|
||||
exchanger = _Exchanger(SubjectTokenRefusal(error=error, description="subject_token was rejected: bad signature"))
|
||||
response = await _exchange_native(client_id, minter, exchanger)
|
||||
assert response.status_code == status
|
||||
body = json.loads(response.body)
|
||||
assert (body["error"], body["error_description"]) == (error, "subject_token was rejected: bad signature")
|
||||
assert minter.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"failure, status, error",
|
||||
[
|
||||
("not_a_member", 400, "invalid_grant"),
|
||||
("team_required", 400, "invalid_grant"),
|
||||
("no_active_key", 400, "invalid_grant"),
|
||||
("unavailable", 503, "temporarily_unavailable"),
|
||||
],
|
||||
)
|
||||
async def test_token_exchange_relays_a_mint_refusal(failure, status, error):
|
||||
client_id = (await _register([LOOPBACK_REDIRECT_URI]))["client_id"]
|
||||
response = await _exchange_native(client_id, _Minter(failure), _Exchanger())
|
||||
assert response.status_code == status
|
||||
assert json.loads(response.body)["error"] == error
|
||||
|
|
|
|||
|
|
@ -0,0 +1,237 @@
|
|||
import logging
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from prisma.engine.errors import BinaryNotFoundError
|
||||
from prisma.errors import DataError
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import SubjectIdentity, SubjectTokenRefusal
|
||||
from litellm.proxy._experimental.mcp_server.idp_token_exchange import (
|
||||
REJECTED_SUBJECT_TOKEN,
|
||||
SUBJECT_TOKEN_CHECK_FAULTED,
|
||||
SUBJECT_TOKEN_CHECK_UNAVAILABLE,
|
||||
TokenExchangePrerequisites,
|
||||
identity_from_subject_token,
|
||||
token_exchange_available,
|
||||
)
|
||||
from litellm.proxy._types import JWTIssuerConfig, LiteLLM_JWTAuth, ProxyException
|
||||
from litellm.proxy.auth.handle_jwt import JWKSUnreachableError, JWTHandler, jwks_unavailable_exception
|
||||
|
||||
IDP_JWT = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1MSJ9.idp-signature"
|
||||
REQUEST_HEADERS = {"x-litellm-team-id": "team-b", "user-agent": "lite/0.1"}
|
||||
EVERY_GATE_HOLDS = {
|
||||
"jwt_auth_enabled": True,
|
||||
"has_database": True,
|
||||
"licensed": True,
|
||||
"maps_jwts_to_virtual_keys": False,
|
||||
}
|
||||
JWKS_URL = "https://idp.example.com/.well-known/jwks.json"
|
||||
JWKS_DOWN = jwks_unavailable_exception(JWKSUnreachableError(f"ConnectError fetching {JWKS_URL} after 3 attempts"))
|
||||
|
||||
|
||||
def _authorized(user_id="u1", team_id="team-b"):
|
||||
return {
|
||||
"is_proxy_admin": False,
|
||||
"team_object": None,
|
||||
"user_object": None,
|
||||
"end_user_object": None,
|
||||
"org_object": None,
|
||||
"token": IDP_JWT,
|
||||
"team_id": team_id,
|
||||
"user_id": user_id,
|
||||
"user_email": None,
|
||||
"end_user_id": None,
|
||||
"org_id": None,
|
||||
"team_membership": None,
|
||||
"jwt_claims": {"sub": user_id},
|
||||
"agent_id": None,
|
||||
}
|
||||
|
||||
|
||||
class _Authorizer:
|
||||
def __init__(self, result=None, raises=None):
|
||||
self.calls = []
|
||||
self.result = result if result is not None else _authorized()
|
||||
self.raises = raises
|
||||
|
||||
async def __call__(self, subject_token, request_headers):
|
||||
self.calls.append((subject_token, dict(request_headers)))
|
||||
if self.raises is not None:
|
||||
raise self.raises
|
||||
return self.result
|
||||
|
||||
|
||||
async def _identity(authorizer, subject_token=IDP_JWT, **unmet):
|
||||
return await identity_from_subject_token(
|
||||
subject_token,
|
||||
request_headers=REQUEST_HEADERS,
|
||||
prerequisites=TokenExchangePrerequisites(**{**EVERY_GATE_HOLDS, **unmet}),
|
||||
is_jwt=JWTHandler.is_jwt,
|
||||
authorize=authorizer,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_jwt_the_proxy_accepts_names_its_user_and_team():
|
||||
"""The subject token goes to the proxy's own JWT auth with the caller's headers (that is
|
||||
where the team header is read), and the identity it resolved is what gets minted."""
|
||||
authorizer = _Authorizer()
|
||||
assert await _identity(authorizer) == SubjectIdentity(user_id="u1", team_id="team-b")
|
||||
assert authorizer.calls == [(IDP_JWT, REQUEST_HEADERS)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_jwt_that_resolves_no_team_names_a_teamless_identity():
|
||||
assert await _identity(_Authorizer(_authorized(team_id=None))) == SubjectIdentity(user_id="u1", team_id=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"unmet, subject_token, error, mentions",
|
||||
[
|
||||
({"jwt_auth_enabled": False}, IDP_JWT, "unsupported_grant_type", "JWT auth is not enabled"),
|
||||
({"has_database": False}, IDP_JWT, "unsupported_grant_type", "no database"),
|
||||
({"licensed": False}, IDP_JWT, "unsupported_grant_type", "enterprise"),
|
||||
({"maps_jwts_to_virtual_keys": True}, IDP_JWT, "unsupported_grant_type", "virtual keys"),
|
||||
({}, "sk-litellm-virtual-key", "invalid_request", "not a JWT"),
|
||||
],
|
||||
)
|
||||
async def test_the_gates_user_api_key_auth_applies_refuse_before_any_verification(
|
||||
unmet, subject_token, error, mentions
|
||||
):
|
||||
authorizer = _Authorizer()
|
||||
refusal = await _identity(authorizer, subject_token=subject_token, **unmet)
|
||||
assert isinstance(refusal, SubjectTokenRefusal)
|
||||
assert refusal.error == error
|
||||
assert mentions in refusal.description
|
||||
assert authorizer.calls == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"unmet",
|
||||
[
|
||||
{},
|
||||
{"jwt_auth_enabled": False},
|
||||
{"has_database": False},
|
||||
{"licensed": False},
|
||||
{"maps_jwts_to_virtual_keys": True},
|
||||
],
|
||||
)
|
||||
def test_the_grant_is_available_exactly_when_every_gate_holds(unmet):
|
||||
prerequisites = TokenExchangePrerequisites(**{**EVERY_GATE_HOLDS, **unmet})
|
||||
assert prerequisites.available is (unmet == {})
|
||||
assert (prerequisites.refusal() is None) is prerequisites.available
|
||||
|
||||
|
||||
MAPPED_ISSUER = JWTIssuerConfig(
|
||||
issuer="https://idp.example.test", audience="litellm-gateway", virtual_key_claim_field="client_id"
|
||||
)
|
||||
|
||||
|
||||
def _running_jwt_handler(litellm_jwtauth):
|
||||
handler = JWTHandler()
|
||||
if litellm_jwtauth is not None:
|
||||
handler.update_environment(prisma_client=None, user_api_key_cache=DualCache(), litellm_jwtauth=litellm_jwtauth)
|
||||
return handler
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"general_settings, prisma_client, premium_user, litellm_jwtauth, expected",
|
||||
[
|
||||
({"enable_jwt_auth": True}, object(), True, LiteLLM_JWTAuth(), True),
|
||||
({"enable_jwt_auth": True}, object(), True, None, True),
|
||||
({}, object(), True, LiteLLM_JWTAuth(), False),
|
||||
({"enable_jwt_auth": True}, None, True, LiteLLM_JWTAuth(), False),
|
||||
({"enable_jwt_auth": True}, object(), False, LiteLLM_JWTAuth(), False),
|
||||
({"enable_jwt_auth": True}, object(), True, LiteLLM_JWTAuth(virtual_key_claim_field="client_id"), False),
|
||||
({"enable_jwt_auth": True}, object(), True, LiteLLM_JWTAuth(issuers=[MAPPED_ISSUER]), False),
|
||||
],
|
||||
)
|
||||
def test_availability_is_read_from_the_running_proxy(
|
||||
monkeypatch, general_settings, prisma_client, premium_user, litellm_jwtauth, expected
|
||||
):
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", premium_user)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.jwt_handler", _running_jwt_handler(litellm_jwtauth))
|
||||
assert token_exchange_available() is expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"raised, reason",
|
||||
[
|
||||
(HTTPException(status_code=403, detail="User not allowed to access this route"), "not allowed"),
|
||||
(ProxyException(message="Token expired", type="auth_error", param="token", code=401), "Token expired"),
|
||||
(Exception("Validation fails: signature verification failed"), "signature verification failed"),
|
||||
(Exception("Invalid JWT Submitted"), "Invalid JWT"),
|
||||
(Exception(f"Failed to fetch keys from {JWKS_URL}: 502 Bad Gateway from the IdP"), JWKS_URL),
|
||||
(ValueError("User doesn't exist in db. 'user_id'=u1. Got error - not found"), "not found"),
|
||||
],
|
||||
)
|
||||
async def test_a_jwt_the_proxy_rejects_is_refused_with_the_reason_kept_in_the_log(raised, reason, caplog):
|
||||
"""The endpoint is public, so the response never quotes JWT auth's wording (it can name
|
||||
the JWKS URL or relay the IdP's reply); the operator reads the reason in the proxy log."""
|
||||
caplog.set_level(logging.WARNING, logger="LiteLLM Proxy")
|
||||
refusal = await _identity(_Authorizer(raises=raised))
|
||||
assert refusal == SubjectTokenRefusal(error="invalid_request", description=REJECTED_SUBJECT_TOKEN)
|
||||
assert reason in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_jwt_that_resolves_no_user_cannot_be_exchanged():
|
||||
refusal = await _identity(_Authorizer(_authorized(user_id=None)))
|
||||
assert refusal == SubjectTokenRefusal(
|
||||
error="invalid_request", description="subject_token names no user the gateway knows"
|
||||
)
|
||||
|
||||
|
||||
def _user_lookup_wrapping_a_database_outage():
|
||||
p1001 = DataError(
|
||||
data={"user_facing_error": {"message": "Can't reach database server at `127.0.0.1`:`5432`", "meta": {}}}
|
||||
)
|
||||
try:
|
||||
raise p1001
|
||||
except DataError as outage:
|
||||
try:
|
||||
raise ValueError(f"User doesn't exist in db. 'user_id'=u1. Got error - {outage}")
|
||||
except ValueError as wrapped:
|
||||
return wrapped
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"raised, reason",
|
||||
[
|
||||
(JWKS_DOWN, JWKS_URL),
|
||||
(HTTPException(status_code=503, detail="the auth database is not reachable"), "not reachable"),
|
||||
(_user_lookup_wrapping_a_database_outage(), "Can't reach database server"),
|
||||
],
|
||||
)
|
||||
async def test_an_idp_or_gateway_outage_is_reported_as_retryable_not_as_a_bad_token(raised, reason, caplog):
|
||||
caplog.set_level(logging.ERROR, logger="LiteLLM Proxy")
|
||||
refusal = await _identity(_Authorizer(raises=raised))
|
||||
assert refusal == SubjectTokenRefusal(error="temporarily_unavailable", description=SUBJECT_TOKEN_CHECK_UNAVAILABLE)
|
||||
assert reason in caplog.text
|
||||
|
||||
|
||||
def _user_lookup_wrapping_a_fault_retrying_cannot_clear():
|
||||
try:
|
||||
raise BinaryNotFoundError("query engine binary not found")
|
||||
except BinaryNotFoundError as fault:
|
||||
try:
|
||||
raise ValueError(f"User doesn't exist in db. 'user_id'=u1. Got error - {fault}")
|
||||
except ValueError as wrapped:
|
||||
return wrapped
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_database_fault_retrying_cannot_clear_is_not_reported_as_a_transient_outage(caplog):
|
||||
"""The status stays 503 (the only OAuth error a client reads as the server's fault, and what
|
||||
the mint path answers to the same fault) but the wording must not tell the client to wait."""
|
||||
caplog.set_level(logging.ERROR, logger="LiteLLM Proxy")
|
||||
refusal = await _identity(_Authorizer(raises=_user_lookup_wrapping_a_fault_retrying_cannot_clear()))
|
||||
assert refusal == SubjectTokenRefusal(error="temporarily_unavailable", description=SUBJECT_TOKEN_CHECK_FAULTED)
|
||||
assert "retrying will not help" in refusal.description
|
||||
assert "faulted: " in caplog.text and "query engine binary not found" in caplog.text
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
"""Tests for minting the ``lite login`` credential from a consented native-client grant."""
|
||||
|
||||
from unittest.mock import ANY, AsyncMock
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -8,7 +8,9 @@ from litellm.constants import CLI_JWT_EXPIRATION_HOURS
|
|||
from litellm.models.user import LiteLLM_UserTable
|
||||
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import ConsentTeam, MintedProxyCredential
|
||||
from litellm.proxy._experimental.mcp_server.proxy_api_credentials import lookup_consent_teams, mint_proxy_credential
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail
|
||||
|
||||
_LOAD_USER = "litellm.proxy._experimental.mcp_server.proxy_api_credentials.load_active_user_by_id"
|
||||
|
|
@ -67,10 +69,21 @@ async def test_mint_passes_user_lookup_failures_through(failure, load_user, fetc
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mint_refuses_a_user_without_a_role(load_user, fetch_teams):
|
||||
load_user.return_value = _user(user_role=None)
|
||||
assert await mint_proxy_credential("u1", None) == "no_active_key"
|
||||
fetch_teams.assert_not_awaited()
|
||||
@pytest.mark.parametrize(
|
||||
"stored_role, minted_role",
|
||||
[
|
||||
(None, LitellmUserRoles.INTERNAL_USER),
|
||||
("made_up_role", LitellmUserRoles.INTERNAL_USER),
|
||||
("proxy_admin", LitellmUserRoles.PROXY_ADMIN),
|
||||
],
|
||||
)
|
||||
async def test_mint_carries_the_role_the_proxy_enforces_for_the_user(load_user, fetch_teams, stored_role, minted_role):
|
||||
"""A user JWT auth upserted has no role in the database, and the proxy already treats
|
||||
such a user as an internal user on every request, so the credential says the same."""
|
||||
load_user.return_value = _user(user_role=stored_role)
|
||||
minted = await mint_proxy_credential("u1", "team-a")
|
||||
assert isinstance(minted, MintedProxyCredential)
|
||||
assert _decoded(minted).user_role == minted_role
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -79,7 +92,7 @@ async def test_mint_refuses_a_teamless_grant_for_a_team_member(load_user, fetch_
|
|||
is refused for a user with teams instead of minting an unscoped credential or drifting
|
||||
onto the first team, on redemption and on every refresh alike."""
|
||||
assert await mint_proxy_credential("u1", None) == "team_required"
|
||||
load_user.assert_awaited_once_with("u1")
|
||||
load_user.assert_awaited_once_with("u1", source="database")
|
||||
fetch_teams.assert_awaited_once_with(ANY, ["team-a", "team-b"])
|
||||
|
||||
|
||||
|
|
@ -114,6 +127,53 @@ async def test_mint_honors_the_consented_team(load_user, fetch_teams):
|
|||
assert decoded.team_model_aliases == {"fast": "gpt-5.4-mini"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mint_reads_the_users_teams_from_the_database_not_a_stale_cached_row(fetch_teams, monkeypatch):
|
||||
"""JWT auth caches the user it creates before it adds that user to the JWT's team, and adding a member
|
||||
never evicts the cached row, so a mint off the cached row refused the very first token exchange as not
|
||||
a member. The mint has to read the database row, whatever the cache holds."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
cache = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key="stale-cache-user", value=_user(user_id="stale-cache-user", teams=[]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_user(user_id="stale-cache-user", teams=["team-a"])
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
|
||||
minted = await mint_proxy_credential("stale-cache-user", "team-a")
|
||||
|
||||
assert isinstance(minted, MintedProxyCredential)
|
||||
assert minted.team_id == "team-a"
|
||||
assert _decoded(minted).team_id == "team-a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mint_refuses_a_user_scim_deactivated_after_the_cache_last_saw_them_active(fetch_teams, monkeypatch):
|
||||
"""SCIM deactivation writes the user row without evicting the cached copy, so a mint off the cache would
|
||||
keep issuing credentials for the management-object TTL. The mint reads the database row, so the
|
||||
deactivated user is refused on the first refresh after the deactivation."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
cache = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key="deactivated-user", value=_user(user_id="deactivated-user", teams=["team-a"]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_user(user_id="deactivated-user", teams=["team-a"], metadata={"scim_active": False})
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
|
||||
assert await mint_proxy_credential("deactivated-user", "team-a") == "no_active_key"
|
||||
fetch_teams.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mint_refuses_a_team_the_user_is_not_on(load_user, fetch_teams):
|
||||
assert await mint_proxy_credential("u1", "team-c") == "not_a_member"
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional
|
||||
|
|
@ -916,6 +917,32 @@ async def test_get_user_object_wraps_db_outage_as_valueerror_preserving_context(
|
|||
assert isinstance(exc_info.value.__context__, ConnectionError)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
|
||||
"""A database-only read is never answered by the per-worker negative memo: a row created after a miss on
|
||||
this worker is returned within db_cache_expiry seconds instead of raising UserNotFoundError, so the token
|
||||
exchange mints for a user JWT auth just accepted."""
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
user_id = "memo-probe-user"
|
||||
monkeypatch.setitem(auth_checks.last_db_access_time, f"user_id:{user_id}", (None, time.time()))
|
||||
db_row = LiteLLM_UserTable(user_id=user_id, user_email=None, user_role="internal_user")
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
|
||||
|
||||
result = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=UserApiKeyCache(),
|
||||
user_id_upsert=False,
|
||||
check_db_only=True,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.user_id == user_id
|
||||
mock_prisma_client.db.litellm_usertable.find_unique.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_object_upsert_includes_user_email():
|
||||
"""Test that user_email is included when creating a new user via get_user_object upsert"""
|
||||
|
|
@ -8740,6 +8767,23 @@ async def test_access_group_model_fallback_uses_the_injected_database(channel: s
|
|||
reader.assert_awaited_once_with(where={"access_group_id": "group-a"})
|
||||
|
||||
|
||||
def test_jwt_team_role_reaches_the_gateway_token_endpoint_by_default():
|
||||
"""The RFC 8693 token exchange authorizes the IdP JWT against ``POST /token`` itself, and JWT
|
||||
auth only binds a team from a multi-team claim when that team may call the route, so the
|
||||
default team allowlist has to cover the gateway's token endpoint or the exchange would mint
|
||||
teamless credentials for every ``team_ids_jwt_field`` deployment."""
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth
|
||||
from litellm.proxy.auth.auth_checks import allowed_routes_check
|
||||
|
||||
assert allowed_routes_check(
|
||||
user_role=LitellmUserRoles.TEAM, user_route="/token", litellm_proxy_roles=LiteLLM_JWTAuth()
|
||||
)
|
||||
assert not allowed_routes_check(
|
||||
user_role=LitellmUserRoles.TEAM,
|
||||
user_route="/token",
|
||||
litellm_proxy_roles=LiteLLM_JWTAuth(team_allowed_routes=[]),
|
||||
)
|
||||
|
||||
def test_route_skips_budget_checks_marks_only_spend_free_routes() -> None:
|
||||
assert route_skips_budget_checks(route="/v1/models") is True
|
||||
assert route_skips_budget_checks(route="/spend/logs") is True
|
||||
|
|
|
|||
|
|
@ -627,6 +627,7 @@ def test_virtual_key_llm_api_routes_denies_spend_logs_v2():
|
|||
"/mcp/tools/call",
|
||||
"/mcp-rest/tools/call",
|
||||
"/mcp/tools/list",
|
||||
"/token",
|
||||
],
|
||||
)
|
||||
def test_mcp_inference_routes_classified_as_llm_api(route):
|
||||
|
|
|
|||
|
|
@ -13316,6 +13316,77 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp
|
|||
assert created_user_id not in mock_audit.call_args.kwargs["existing_user_ids"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_add_evicts_the_new_members_cached_user_row_on_every_worker(monkeypatch):
|
||||
"""Auth admits a team-bound credential off the teams list of the cached user row. The add wrote the
|
||||
new team to the database row only, so a worker still holding the old row refused the member's
|
||||
credential with 403 until the management-object TTL expired. The add now evicts the row here and
|
||||
broadcasts the eviction to the other workers, the way /team/member_delete already does."""
|
||||
from litellm.proxy._types import TeamMemberAddRequest
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_member_add
|
||||
|
||||
team_id = "team-b"
|
||||
user_id = "dev-1"
|
||||
cache = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key=user_id, value=LiteLLM_UserTable(user_id=user_id, teams=["team-a"]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
broadcast = AsyncMock()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", AsyncMock())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id")
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation", broadcast
|
||||
)
|
||||
|
||||
updated_team = MagicMock()
|
||||
updated_team.model_dump.return_value = {
|
||||
"team_id": team_id,
|
||||
"members_with_roles": [{"user_id": user_id, "role": "user"}],
|
||||
}
|
||||
|
||||
async def fake_add_team_members_to_team(**kwargs):
|
||||
return updated_team, [LiteLLM_UserTable(user_id=user_id, teams=["team-a", team_id])], []
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: team_member_add has no injection seam for its prisma-backed helpers
|
||||
"litellm.proxy.management_endpoints.team_endpoints.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=LiteLLM_TeamTable(team_id=team_id, members_with_roles=[]),
|
||||
),
|
||||
patch( # test-quality-ok: team_member_add has no injection seam for its prisma-backed helpers
|
||||
"litellm.proxy.management_endpoints.team_endpoints._validate_team_member_add_permissions",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch( # test-quality-ok: team_member_add has no injection seam for its prisma-backed helpers
|
||||
"litellm.proxy.management_endpoints.team_endpoints._validate_and_populate_member_user_info",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch( # test-quality-ok: team_member_add has no injection seam for its prisma-backed helpers
|
||||
"litellm.proxy.management_endpoints.team_endpoints._resolve_existing_member_user_ids",
|
||||
new_callable=AsyncMock,
|
||||
return_value=frozenset({user_id}),
|
||||
),
|
||||
patch( # test-quality-ok: team_member_add has no injection seam for its prisma-backed helpers
|
||||
"litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team",
|
||||
side_effect=fake_add_team_members_to_team,
|
||||
),
|
||||
patch( # test-quality-ok: team_member_add has no injection seam for its prisma-backed helpers
|
||||
"litellm.proxy.management_endpoints.team_endpoints._create_team_member_add_audit_logs",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
):
|
||||
await team_member_add(
|
||||
data=TeamMemberAddRequest(team_id=team_id, member=Member(user_id=user_id, role="user")),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1"),
|
||||
)
|
||||
|
||||
assert await cache.async_get_cache(key=user_id, model_type=LiteLLM_UserTable) is None
|
||||
broadcast.assert_awaited_once_with(cache_key=user_id)
|
||||
|
||||
|
||||
def test_validate_member_user_id_provisioning_caps_the_ids_it_echoes_back():
|
||||
"""A large member list must not echo every id back in the error body."""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
|
|
|
|||
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -24767,10 +24767,16 @@ export interface components {
|
|||
redirect_uri?: string;
|
||||
/** Refresh Token */
|
||||
refresh_token?: string | null;
|
||||
/** Requested Token Type */
|
||||
requested_token_type?: string | null;
|
||||
/** Resource */
|
||||
resource?: string | null;
|
||||
/** Scope */
|
||||
scope?: string | null;
|
||||
/** Subject Token */
|
||||
subject_token?: string | null;
|
||||
/** Subject Token Type */
|
||||
subject_token_type?: string | null;
|
||||
};
|
||||
/** Body_token_endpoint_token_post */
|
||||
Body_token_endpoint_token_post: {
|
||||
|
|
@ -24788,10 +24794,16 @@ export interface components {
|
|||
redirect_uri?: string;
|
||||
/** Refresh Token */
|
||||
refresh_token?: string | null;
|
||||
/** Requested Token Type */
|
||||
requested_token_type?: string | null;
|
||||
/** Resource */
|
||||
resource?: string | null;
|
||||
/** Scope */
|
||||
scope?: string | null;
|
||||
/** Subject Token */
|
||||
subject_token?: string | null;
|
||||
/** Subject Token Type */
|
||||
subject_token_type?: string | null;
|
||||
};
|
||||
/** Body_upload_logo_upload_logo_post */
|
||||
Body_upload_logo_upload_logo_post: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue