This commit is contained in:
tin-berri 2026-10-04 23:11:56 +08:00 • committed by GitHub
commit d6f55149fe
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
29 changed files with 1394 additions and 64 deletions

View file

@ -1974,7 +1974,10 @@ async def authorize(
response_type=response_type,
session_user_id=_session_cookie_user_id(request),
lookup_consent_teams=lookup_consent_teams,
scope=scope,
)
if scope == "proxy:admin":
raise HTTPException(400, "proxy:admin requires the gateway URL as its resource")
return aggregate_authorize(
request=request,
client_id=client_id,
@ -2159,7 +2162,7 @@ async def authorize_complete(
async def revoke_endpoint(request: Request, token: str = Form(...), client_id: str = Form(...)) -> Response:
"""RFC 7009 revocation for the gateway's refresh tokens (``lite logout``): 200 for a known
client whatever the token's state, 503 when the shared single-use record cannot be written;
access tokens expire on their own."""
native/MCP access tokens expire normally; delegated access is revoked immediately."""
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
master_key,
user_api_key_cache,
@ -2648,6 +2651,25 @@ def _jwt_auth_issuers() -> list:
return issuers
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/oauth/api")
def oauth_authorization_server_api(request: Request) -> dict[str, str | tuple[str, ...]]:
base: Final = get_request_base_url(request)
return {
"issuer": f"{base}/oauth/api",
"authorization_endpoint": f"{base}/authorize",
"token_endpoint": f"{base}/token",
"registration_endpoint": f"{base}/register",
"revocation_endpoint": f"{base}/revoke",
"introspection_endpoint": f"{base}/introspect",
"response_types_supported": ("code",),
"grant_types_supported": ("authorization_code", "refresh_token"),
"scopes_supported": ("proxy:admin",),
"code_challenge_methods_supported": ("S256",),
"token_endpoint_auth_methods_supported": ("none",),
"revocation_endpoint_auth_methods_supported": ("none",),
}
@router.get("/.well-known/oauth-protected-resource")
def oauth_protected_resource_root(request: Request) -> dict[str, str | tuple[str, ...]]:
request_base_url: Final = get_request_base_url(request)

View file

@ -26,8 +26,8 @@ this track) and then walks the flow implemented here:
:mod:`.outbound_credentials.session_token`, re-validating that the litellm user is
still active first; the ``refresh_token`` grant rotates the pair the same way.
Nothing here stores state server-side except the single-use code guard (a TTL cache
entry). Every sealed value is authenticated encryption over the proxy salt/master key
Server-side state consists of single-use markers and delegated application token families.
Every sealed value is authenticated encryption over the proxy salt/master key
family, opened totally (bad input maps to an OAuth error, never a raise), and every
identity is a stable reference re-validated live at mint, refresh, and (in the admission
PR) tool-call time. Upstream server credentials never appear anywhere in this flow; they
@ -73,6 +73,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credent
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
SESSION_ISSUER,
SESSION_REFRESH_TTL_SECONDS,
DelegatedGrant,
MintedSessionToken,
OpenedSessionToken,
SessionAudience,
@ -85,6 +86,14 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token i
open_session_refresh_token,
open_session_token,
)
from litellm.proxy.auth.delegated_oauth import (
delegated_identity,
delegated_user,
delegation_active,
is_delegated_callback,
issue_delegated_tokens,
revoke_delegation,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
@ -305,6 +314,7 @@ class _ConnectFlow(BaseModel):
exp: int
resource_server_id: str | None = None
audience: SessionAudience | None = None
delegation: DelegatedGrant | None = None
class _GatewayAuthCode(BaseModel):
@ -324,6 +334,7 @@ class _GatewayAuthCode(BaseModel):
resource_server_id: str | None = None
audience: SessionAudience | None = None
team_id: str | None = None
delegation: DelegatedGrant | None = None
def is_gateway_dcr_client_id(client_id: str | None) -> bool:
@ -566,24 +577,30 @@ async def native_client_authorize(
response_type: str | None,
session_user_id: str | None,
lookup_consent_teams: LookupConsentTeams,
scope: str | None = None,
) -> Response:
"""The authorize verb for a native client that named the proxy API itself as its
RFC 8707 ``resource``: the same client, redirect, PKCE, and sign-in checks as the
aggregate verb plus a loopback-only redirect (the credential this grant mints is the
user's personal proxy key, which belongs on their own machine and never behind a hosted
callback), then the consent page rendered right here (no connect-page interlude, since
there is no per-server vaulting to do) with the flow sealed into the per-flow cookie
and its handle carried only in the form, never in a URL."""
"""Authorize the proxy API through existing sign-in and consent. Native clients use
loopback callbacks; explicit admin delegation requires an approved hosted callback."""
rejected: Final = _rejected_authorize_request(
client_id, redirect_uri, state, code_challenge, code_challenge_method, response_type
)
if rejected is not None:
return rejected
if not is_loopback_redirect_host(urlparse(redirect_uri)):
delegated: Final = scope == "proxy:admin"
if scope not in (None, "", "proxy:admin"):
return _oauth_error(400, "invalid_scope", "unsupported proxy API scope")
if delegated and not is_delegated_callback(redirect_uri):
return _oauth_error(400, "invalid_request", "the application callback is not approved by this gateway")
if not delegated and not is_loopback_redirect_host(urlparse(redirect_uri)):
return _oauth_error(400, "invalid_request", "a proxy-API grant may only redirect to a loopback address")
base_url: Final = get_request_base_url(request)
if session_user_id is None:
return _login_redirect(base_url, request)
if delegated:
try:
await delegated_user(session_user_id)
except HTTPException as exc:
return _oauth_error(exc.status_code, "access_denied", str(exc.detail))
teams: Final = await lookup_consent_teams(session_user_id)
if not isinstance(teams, tuple):
return _consent_lookup_failure_response(teams)
@ -596,6 +613,16 @@ async def native_client_authorize(
code_challenge=code_challenge or "",
resource_server_id=None,
audience=PROXY_API_AUDIENCE,
delegation=(
DelegatedGrant(
grant_id=secrets.token_urlsafe(24),
resource=canonicalize_url_identity(base_url),
redirect_uri=redirect_uri,
expires_at=int(datetime.now(timezone.utc).timestamp()) + SESSION_REFRESH_TTL_SECONDS,
)
if delegated
else None
),
)
page: Final = render_native_client_consent_page(
client_origin=_origin_only(redirect_uri),
@ -603,6 +630,7 @@ async def native_client_authorize(
teams=tuple((team.team_id, team.team_alias or team.team_id) for team in teams),
flow_handle=handle,
complete_url=f"{base_url}/authorize/complete",
delegated=delegated,
)
response: Final = HTMLResponse(page, headers=_CONSENT_PAGE_HEADERS)
_set_flow_cookie(response, request, handle, flow)
@ -708,6 +736,7 @@ def _new_connect_flow(
code_challenge: str,
resource_server_id: str | None,
audience: SessionAudience | None,
delegation: DelegatedGrant | None = None,
) -> _ConnectFlow:
now: Final = datetime.now(timezone.utc)
return _ConnectFlow(
@ -720,6 +749,7 @@ def _new_connect_flow(
exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS,
resource_server_id=resource_server_id,
audience=audience,
delegation=delegation,
)
@ -879,6 +909,8 @@ async def complete_connect_flow(
opened: Final = _open_flow_for(request, flow_handle, session_user_id, now)
if isinstance(opened, Response):
return opened
if opened.delegation is not None and decision is None:
return _oauth_error(400, "invalid_request", "an explicit approval or denial is required")
if decision != "deny":
described: Final = await _describe_opened_flow(opened, lookup_vendor_credential, lookup_server_reachability)
if isinstance(described, Response):
@ -930,6 +962,7 @@ def _approved_flow_response(flow: _ConnectFlow, delivery: str | None, team_id: s
resource_server_id=flow.resource_server_id,
audience=flow.audience,
team_id=(team_id or None) if flow.audience == PROXY_API_AUDIENCE else None,
delegation=flow.delegation,
),
)
callback_url: Final = _append_query_params(flow.redirect_uri, (("code", code), *_state_param(flow)))
@ -1261,8 +1294,8 @@ async def aggregate_token(
class _GrantIssuer:
"""The tail every grant shares once its own proof (code + PKCE, or a refresh token)
has checked out: revalidate the user live, claim the single-use marker, mint. The
claim comes AFTER revalidation and minting so a transient DB 503 never burns a
still-valid code or refresh token, and fails closed when it cannot be recorded."""
Native and MCP grants claim after revalidation. Delegated grants atomically create
a revocable token family as their single-use admission."""
def __init__(
self,
@ -1285,6 +1318,11 @@ class _GrantIssuer:
async def __call__(
self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str
) -> Response:
if principal.delegation is not None:
target_refusal: Final = self._delegated_target_refusal(principal)
if target_refusal is not None:
return target_refusal
return await issue_delegated_tokens(principal, self._keys, self._now)
match principal.audience:
case None:
return await self._issue_session_pair(principal, claim_key, claim_ttl_seconds, replayed)
@ -1293,6 +1331,20 @@ class _GrantIssuer:
case _:
assert_never(principal.audience)
def _delegated_target_refusal(self, principal: SessionPrincipal) -> Response | None:
if principal.delegation is None or (
principal.audience != PROXY_API_AUDIENCE
or principal.delegation.resource != canonicalize_url_identity(get_request_base_url(self._request))
):
return _oauth_error(400, "invalid_target", "the grant belongs to a different gateway")
return self._proxy_api_target_refusal()
async def refresh_delegation(self, principal: SessionPrincipal, jti: str) -> Response:
refusal: Final = self._delegated_target_refusal(principal)
if refusal is not None:
return refusal
return await issue_delegated_tokens(principal, self._keys, self._now, previous_refresh_jti=jti)
async def _issue_session_pair(
self, principal: SessionPrincipal, claim_key: str, claim_ttl_seconds: int, replayed: str
) -> Response:
@ -1385,6 +1437,7 @@ async def _authorization_code_grant(
resource_server_id=parsed.resource_server_id,
audience=parsed.audience,
team_id=parsed.team_id,
delegation=parsed.delegation,
),
claim_key=f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}",
claim_ttl_seconds=parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS,
@ -1408,6 +1461,8 @@ async def _refresh_token_grant(
return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client")
if _resource_conflicts_with_scope(request, resource, opened.principal.resource_server_id):
return _oauth_error(400, "invalid_target", "resource does not match the scope this token was issued for")
if opened.principal.delegation is not None:
return await issue.refresh_delegation(opened.principal, opened.jti)
# Refresh-token rotation (OAuth 2.0 Security BCP section 4.13): the presented refresh token is
# single-use, so a captured or replayed refresh token cannot mint a second pair after the
# legitimate holder rotated.
@ -1446,14 +1501,8 @@ async def _token_exchange_grant(
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
stateless and expire on their own (the proxy-API credential within
``CLI_JWT_EXPIRATION_HOURS``), so per RFC 7009 section 2.2 an unrecognized or already
dead token still answers 200; only an unknown client is refused. A live token whose
burn could not be recorded in the shared backend answers 503 (section 2.2.1), so the
client knows the token still stands and retries instead of reporting a logout that
never happened."""
"""RFC 7009: revoke a delegated family or burn a native/MCP refresh token's JTI.
Invalid tokens return 200; unavailable shared storage returns 503 for retry."""
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 master_key is None:
@ -1464,6 +1513,18 @@ async def revoke_refresh_token(token: str, client_id: str, master_key: str | Non
verbose_logger.error("mcp_gateway_dcr revoke rejected: %s", keys.detail)
return _oauth_error(500, "server_error", "the gateway session signing configuration is invalid")
now: Final = datetime.now(timezone.utc)
candidate: Final = (
open_session_token(token, keys, now, for_revocation=True)
if is_session_token(token)
else open_session_refresh_token(token, keys, now, for_revocation=True)
)
if isinstance(candidate, OpenedSessionToken) and candidate.principal.delegation is not None:
if candidate.principal.client_id == client_id:
try:
await revoke_delegation(candidate.principal)
except HTTPException as exc:
return _oauth_error(exc.status_code, "temporarily_unavailable", str(exc.detail))
return Response(content="{}", media_type="application/json", headers=TOKEN_NO_CACHE_HEADERS)
opened: Final = open_session_refresh_bearer(token, keys, now, expected_client_id=client_id)
if isinstance(opened, SessionRefreshOpened):
burned: Final = await _SingleUseGuard(cache).claim(
@ -1490,6 +1551,8 @@ def _active_introspection_response(opened: OpenedSessionToken) -> Response:
("team_id", principal.team_id),
("resource_server_id", principal.resource_server_id),
("audience", principal.audience),
("scope", principal.delegation.scope if principal.delegation is not None else None),
("aud", principal.delegation.resource if principal.delegation is not None else None),
)
if value is not None
}
@ -1497,7 +1560,7 @@ def _active_introspection_response(opened: OpenedSessionToken) -> Response:
status_code=200,
content={
"active": True,
"iss": SESSION_ISSUER,
"iss": f"{principal.delegation.resource}/oauth/api" if principal.delegation is not None else SESSION_ISSUER,
"sub": principal.user_id,
"client_id": principal.client_id,
"jti": opened.jti,
@ -1540,6 +1603,16 @@ async def introspect_gateway_token(
return _inactive_introspection_response()
if not isinstance(opened, OpenedSessionToken):
return _inactive_introspection_response()
if opened.principal.delegation is not None:
try:
if not await delegation_active(opened.principal, opened.jti if opened.kind == "session_refresh" else None):
return _inactive_introspection_response()
await delegated_identity(opened.principal)
except HTTPException as exc:
if exc.status_code >= 500:
return _oauth_error(503, "temporarily_unavailable", str(exc.detail))
return _inactive_introspection_response()
return _active_introspection_response(opened)
if opened.kind == "session_refresh":
peeked: Final = await _SingleUseGuard(cache).peek(f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}")
if peeked == "unavailable":

View file

@ -233,7 +233,7 @@ def resolve_session_bearer(
if not is_session_token(candidate):
return NotSessionBearer()
opened: Final = open_session_token(candidate, keys, now)
if isinstance(opened, OpenedSessionToken):
if isinstance(opened, OpenedSessionToken) and opened.principal.audience is None:
return SessionBearerAdmitted(principal=opened.principal)
return SessionBearerInvalid(expired=isinstance(opened, SessionExpired))

View file

@ -1,4 +1,4 @@
"""Identity-only session tokens for the gateway-level (aggregate ``/mcp``) DCR front door.
"""Identity-only session tokens for gateway OAuth clients.
A DCR client that signs in through LiteLLM SSO holds ONE bearer that carries ONLY a
litellm identity; unlike the :mod:`.envelope` bridge bearer it seals no upstream
@ -6,7 +6,7 @@ credential, because the custody model vaults every upstream token server-side in
``LiteLLM_MCPUserCredentials`` and egress resolves them by user at call time. The token
is therefore a stable REFERENCE, not an authorization: admission reloads the live user
record and policy on every request, so deactivating the user (or their team) kills
outstanding sessions immediately without a revocation store.
outstanding sessions immediately. Delegated API sessions additionally require a live grant family.
Wire shape: ``llm_session_`` (access) / ``llm_srefresh_`` (refresh) + a JWT signed with
the injected key material: HS256 under the default master-key-derived secret (the same
@ -90,11 +90,17 @@ on open, so a signature-valid token of one kind cannot be replayed as the other
wire prefix is swapped (the prefix is not part of the signed payload; this claim is)."""
SessionAudience = Literal["proxy_api"]
"""The non-MCP audience a session REFRESH token can be minted for. ``None`` (the default and
the only value ever on an MCP wire) means the aggregate MCP gateway; ``"proxy_api"`` means the
refresh grant re-mints the proxy-API CLI credential instead of an MCP session pair. The audience
is read only from the signed claims, never from the request, so a token of one audience can
never be redeemed as the other."""
"""``None`` identifies MCP sessions. ``proxy_api`` identifies native CLI refresh grants
or delegated API token pairs. Admission and renewal use only this signed audience."""
class DelegatedGrant(BaseModel):
model_config = ConfigDict(frozen=True, strict=True, extra="forbid")
grant_id: str = Field(min_length=1)
resource: str = Field(min_length=1)
redirect_uri: str = Field(min_length=1)
expires_at: int
scope: Literal["proxy:admin"] = "proxy:admin"
class SessionPrincipal(BaseModel):
@ -119,6 +125,7 @@ class SessionPrincipal(BaseModel):
resource_server_id: str | None = None
audience: SessionAudience | None = None
team_id: str | None = None
delegation: DelegatedGrant | None = None
class SessionKeys(BaseModel):
@ -302,6 +309,7 @@ class _SessionClaims(BaseModel):
resource_server_id: str | None = None
audience: SessionAudience | None = None
team_id: str | None = None
delegation: DelegatedGrant | None = None
def is_session_token(candidate: str) -> bool:
@ -361,19 +369,30 @@ def open_session_token(
candidate: str,
keys: SessionSigningKeys,
now: datetime,
*,
for_revocation: bool = False,
) -> OpenedSessionToken | SessionTokenOpenError:
"""Validate a session ACCESS ``candidate`` and recover the principal.
Never raises for bad input: every invalid, expired, tampered, or wrong-kind candidate
maps to a distinct ``SessionTokenOpenError`` variant.
"""
return _open(candidate, prefix=SESSION_TOKEN_PREFIX, expected_kind="session", keys=keys, now=now)
return _open(
candidate,
prefix=SESSION_TOKEN_PREFIX,
expected_kind="session",
keys=keys,
now=now,
for_revocation=for_revocation,
)
def open_session_refresh_token(
candidate: str,
keys: SessionSigningKeys,
now: datetime,
*,
for_revocation: bool = False,
) -> OpenedSessionToken | SessionTokenOpenError:
"""Validate a session REFRESH ``candidate`` and recover the principal.
@ -381,7 +400,14 @@ def open_session_refresh_token(
``kind="session_refresh"`` claim is required, so an access token re-prefixed as a
refresh one is rejected as ``SessionMalformed``.
"""
return _open(candidate, prefix=SESSION_REFRESH_PREFIX, expected_kind="session_refresh", keys=keys, now=now)
return _open(
candidate,
prefix=SESSION_REFRESH_PREFIX,
expected_kind="session_refresh",
keys=keys,
now=now,
for_revocation=for_revocation,
)
def _mint(
@ -394,10 +420,15 @@ def _mint(
) -> MintedSessionToken | SessionTokenTooLarge:
"""Sign the claims for either token kind and enforce the size cap. Shared by both mints
so the JWT shape, issuer, and size guard cannot drift between access and refresh."""
bounded_expiry: Final = (
min(expires_at, datetime.fromtimestamp(principal.delegation.expires_at, tz=now.tzinfo))
if principal.delegation is not None
else expires_at
)
claims: Final = _SessionClaims(
iss=SESSION_ISSUER,
iat=int(now.timestamp()),
exp=int(expires_at.timestamp()),
exp=int(bounded_expiry.timestamp()),
jti=secrets.token_urlsafe(16),
kind=kind,
user_id=principal.user_id,
@ -405,12 +436,13 @@ def _mint(
resource_server_id=principal.resource_server_id,
audience=principal.audience,
team_id=principal.team_id,
delegation=principal.delegation,
)
token: Final = prefix + _sign_claims(claims, keys)
size_bytes: Final = len(token.encode("utf-8"))
if size_bytes > MAX_SESSION_TOKEN_BYTES:
return SessionTokenTooLarge(size_bytes=size_bytes, max_bytes=MAX_SESSION_TOKEN_BYTES)
return MintedSessionToken(token=SecretStr(token), expires_at=expires_at)
return MintedSessionToken(token=SecretStr(token), expires_at=bounded_expiry)
def _sign_claims(claims: _SessionClaims, keys: SessionSigningKeys) -> str:
@ -434,6 +466,7 @@ def _open(
expected_kind: SessionTokenKind,
keys: SessionSigningKeys,
now: datetime,
for_revocation: bool = False,
) -> OpenedSessionToken | SessionTokenOpenError:
"""Prefix-route, size-bound, signature-verify, kind-check, and expiry-check an
attacker-controlled candidate, shared by both openers so the security gate is identical
@ -452,7 +485,11 @@ def _open(
return claims
if claims.kind != expected_kind:
return SessionMalformed()
if now.timestamp() >= claims.exp:
if claims.delegation is not None and (claims.audience != "proxy_api" or claims.resource_server_id is not None):
return SessionMalformed()
if claims.delegation is not None and now.timestamp() >= claims.delegation.expires_at:
return SessionExpired()
if now.timestamp() >= claims.exp and not (for_revocation and claims.delegation is not None):
return SessionExpired()
return OpenedSessionToken(
principal=SessionPrincipal(
@ -461,6 +498,7 @@ def _open(
resource_server_id=claims.resource_server_id,
audience=claims.audience,
team_id=claims.team_id,
delegation=claims.delegation,
),
jti=claims.jti,
kind=claims.kind,

View file

@ -32653,6 +32653,41 @@
]
}
},
"/.well-known/oauth-authorization-server/oauth/api": {
"get": {
"operationId": "oauth_authorization_server_api__well_known_oauth_authorization_server_oauth_api_get",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"additionalProperties": {
"anyOf": [
{
"type": "string"
},
{
"items": {
"type": "string"
},
"type": "array"
}
]
},
"title": "Response Oauth Authorization Server Api Well Known Oauth Authorization Server Oauth Api Get",
"type": "object"
}
}
},
"description": "Successful Response"
}
},
"summary": "Oauth Authorization Server Api",
"tags": [
"mcp_byok_oauth"
]
}
},
"/.well-known/oauth-authorization-server/{mcp_server_name}": {
"get": {
"description": "OAuth authorization server discovery endpoint.\n\nSupports both legacy pattern (/{server_name}) and root endpoint.",
@ -35077,6 +35112,41 @@
]
}
},
"/.well-known/oauth-authorization-server/oauth/api": {
"get": {
"operationId": "oauth_authorization_server_api__well_known_oauth_authorization_server_oauth_api_get_2",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"additionalProperties": {
"anyOf": [
{
"type": "string"
},
{
"items": {
"type": "string"
},
"type": "array"
}
]
},
"title": "Response Oauth Authorization Server Api Well Known Oauth Authorization Server Oauth Api Get",
"type": "object"
}
}
},
"description": "Successful Response"
}
},
"summary": "Oauth Authorization Server Api",
"tags": [
"mcp_discoverable"
]
}
},
"/.well-known/oauth-authorization-server/{mcp_server_name}": {
"get": {
"description": "OAuth authorization server discovery endpoint.\n\nSupports both legacy pattern (/{server_name}) and root endpoint.",
@ -36045,7 +36115,7 @@
},
"/revoke": {
"post": {
"description": "RFC 7009 revocation for the gateway's refresh tokens (``lite logout``): 200 for a known\nclient whatever the token's state, 503 when the shared single-use record cannot be written;\naccess tokens expire on their own.",
"description": "RFC 7009 revocation for the gateway's refresh tokens (``lite logout``): 200 for a known\nclient whatever the token's state, 503 when the shared single-use record cannot be written;\nnative/MCP access tokens expire normally; delegated access is revoked immediately.",
"operationId": "revoke_endpoint_revoke_post",
"requestBody": {
"content": {

View file

@ -970,6 +970,7 @@ async def common_checks(
)
skip_all_budget_checks: Final = skip_budget_checks or route_skips_budget_checks(route=route)
fresh_policy: Final = valid_token is not None and valid_token.requires_fresh_policy
membership_user_id: Final = (
valid_token.user_id if valid_token is not None and (bool(_model) or not skip_all_budget_checks) else None
@ -982,6 +983,7 @@ async def common_checks(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=fresh_policy,
)
if team_object is not None and membership_user_id is not None
else None
@ -1014,6 +1016,7 @@ async def common_checks(
llm_router=llm_router,
team_model_aliases=(valid_token.team_model_aliases if valid_token else None),
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
check_db_only=fresh_policy,
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
@ -2358,9 +2361,13 @@ async def _fetch_team_membership_from_db(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
*,
use_writer: bool = False,
) -> LiteLLM_TeamMembership | None:
_ = parent_otel_span, proxy_logging_obj
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client), "team_membership").find_unique(
response: Final = await _dictable_table(
TeamMembershipRepository(prisma_client, use_writer=use_writer), "team_membership"
).find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
include={"litellm_budget_table": True},
)
@ -2414,12 +2421,20 @@ async def get_team_membership(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
*,
check_db_only: bool = False,
) -> Optional["LiteLLM_TeamMembership"]:
"""
Returns team membership object if user is member of team.
Do a isolated check for team membership vs. doing a combined key + team + user + team-membership check, as key might come in frequently for different users/teams. Larger call will slowdown query time. This way we get to cache the constant (key/team/user info) and only update based on the changing value (team membership).
"""
if check_db_only:
if prisma_client is None:
raise HTTPException(status_code=503, detail="The gateway database is unavailable")
return await _fetch_team_membership_from_db(
user_id, team_id, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj, use_writer=True
)
if user_id is None or team_id is None:
return None
@ -3205,7 +3220,7 @@ async def _get_team_object_from_user_api_key_cache(
use_writer: bool = False,
) -> LiteLLM_TeamTableCachedObj:
db_access_time_key: Final = key
should_check_db: Final = _should_check_db(
should_check_db: Final = use_writer or _should_check_db(
key=db_access_time_key,
last_db_access_time=last_db_access_time,
db_cache_expiry=db_cache_expiry,
@ -4145,6 +4160,8 @@ async def get_org_object(
parent_otel_span: Span | None = None,
proxy_logging_obj: ProxyLogging | None = None,
include_budget_table: bool = False,
*,
check_db_only: bool = False,
) -> LiteLLM_OrganizationTable | None:
"""
- Check if org id in proxy Org Table
@ -4169,10 +4186,10 @@ async def get_org_object(
if include_budget_table:
cache_key = f"org_id:{org_id}:with_budget"
# check if in cache
deserialized_org: Final = await user_api_key_cache.async_get_cache(
key=cache_key,
model_type=LiteLLM_OrganizationTable,
deserialized_org: Final = (
None
if check_db_only
else await user_api_key_cache.async_get_cache(key=cache_key, model_type=LiteLLM_OrganizationTable)
)
if deserialized_org is not None:
return deserialized_org
@ -4239,6 +4256,7 @@ async def get_org_object_for_request(
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
proxy_logging_obj: ProxyLogging | None,
check_db_only: bool = False,
) -> LiteLLM_OrganizationTable | None:
try:
org: Final = await get_org_object(
@ -4248,10 +4266,13 @@ async def get_org_object_for_request(
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
include_budget_table=True,
check_db_only=check_db_only,
)
except OrganizationNotFoundError:
return None
except Exception as e:
if check_db_only:
raise
if not PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(e):
verbose_proxy_logger.debug("org lookup failed, continuing without org limits", exc_info=True)
return None
@ -4338,6 +4359,7 @@ async def _get_models_from_access_groups(
prisma_client: DatabaseClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,
check_db_only: bool = False,
) -> list[str]:
"""
Collect model names from unified access groups.
@ -4349,6 +4371,7 @@ async def _get_models_from_access_groups(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=check_db_only,
)
@ -5162,6 +5185,7 @@ async def can_team_access_model(
team_model_aliases: dict[str, str] | None = None,
key_model_aliases: Mapping[str, str] | None = None,
prisma_client: DatabaseClient | None = None,
check_db_only: bool = False,
) -> Literal[True]:
"""
Returns True if the team can access a specific model.
@ -5186,6 +5210,7 @@ async def can_team_access_model(
models_from_groups: Final = await _get_models_from_access_groups(
access_group_ids=team_access_group_ids,
prisma_client=prisma_client,
check_db_only=check_db_only,
)
if models_from_groups:
return _can_object_call_model(
@ -6379,8 +6404,11 @@ async def _organization_max_budget_check(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
include_budget_table=True,
check_db_only=valid_token.requires_fresh_policy,
)
except Exception:
if valid_token.requires_fresh_policy:
raise
# If organization lookup fails, skip the check
return

View file

@ -155,6 +155,7 @@ class UserAPIKeyAuthExceptionHandler:
if (
PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
and not (resolved_identity is not None and resolved_identity.requires_fresh_policy)
and PrismaDBExceptionHandler.is_database_connection_error(e)
):
# log this as a DB failure on prometheus

View file

@ -34,6 +34,7 @@ from litellm.proxy._types import *
from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_metadata
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER,
)
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, Deployment, server_owned_wif_fields_present
from litellm.types.router import reject_server_owned_wif_params as _reject_server_owned_wif_params
@ -2065,6 +2066,10 @@ def request_dispatched_to_provider_pass_through(request: Request) -> bool:
return "endpoint" in request.path_params
def request_dispatched_to_marked_provider_pass_through(request: Request) -> bool:
return getattr(request.scope.get("endpoint"), LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER, False) is True
def get_model_from_request(
request_data: dict,
route: str,

View file

@ -0,0 +1,259 @@
from __future__ import annotations
import os
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import Final, Literal
from urllib.parse import urlparse
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, Response
from pydantic import BaseModel, Field, TypeAdapter
from litellm.proxy._experimental.mcp_server.bridge_token_flow import load_active_user_by_id
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
canonical_resource_uri,
get_request_base_url,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
SessionSigningConfigError,
active_session_signing_keys,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
MintedSessionToken,
OpenedSessionToken,
SessionPrincipal,
SessionSigningKeys,
mint_session_refresh_token,
mint_session_token,
open_session_refresh_token,
open_session_token,
)
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles, UserAPIKeyAuth, hash_token
from litellm.proxy.auth.auth_checks import get_object_permission, get_team_membership, get_team_object
from litellm.proxy.auth.resolvers.grants import user_models
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_grants
_FAMILY_SCRIPT: Final = """
local current = redis.call('GET', KEYS[1])
local revoked = '!revoked'
if ARGV[1] == 'revoke' then
redis.call('SET', KEYS[1], revoked, 'EX', ARGV[4]); return 1
end
if ARGV[1] == 'active' then
return current and current ~= revoked and (ARGV[2] == '' or current == ARGV[2]) and 1 or 0
end
if ARGV[1] == 'create' then
if current then return 0 end
elseif not current or current == revoked then
return 0
elseif current ~= ARGV[2] then
redis.call('SET', KEYS[1], revoked, 'EX', ARGV[4]); return 0
end
redis.call('SET', KEYS[1], ARGV[3], 'EX', ARGV[4])
return 1
"""
class _AuthenticationPolicy(BaseModel):
general_settings: Mapping[str, object] = Field(default_factory=dict)
user_custom_auth: object = None
class _UserBudget(BaseModel):
model_max_budget: Mapping[str, object] | None = None
def is_delegated_callback(redirect_uri: str) -> bool:
try:
parsed: Final = urlparse(redirect_uri)
except ValueError:
return False
return (
redirect_uri.isprintable()
and parsed.scheme == "https"
and bool(parsed.hostname)
and not parsed.username
and not parsed.password
and not parsed.fragment
and redirect_uri
in tuple(value.strip() for value in os.getenv("LITELLM_OAUTH_ADMIN_REDIRECT_URIS", "").split(","))
)
async def _family_transition(
principal: SessionPrincipal,
operation: Literal["create", "rotate", "active", "revoke"],
previous: str = "",
replacement: str = "",
) -> bool:
from litellm.proxy.proxy_server import redis_usage_cache
delegation: Final = principal.delegation
if delegation is None or principal.audience != "proxy_api":
return False
ttl: Final = delegation.expires_at - int(datetime.now(timezone.utc).timestamp())
if ttl <= 0:
return False
if redis_usage_cache is None:
raise HTTPException(503, "Delegated OAuth requires shared Redis coordination")
key: Final = "oauth:delegation:" + delegation.grant_id
try:
result: Final = TypeAdapter(int).validate_python(
await redis_usage_cache.async_register_script(_FAMILY_SCRIPT)(
keys=(key,), args=(operation, previous, replacement, ttl)
),
strict=True,
)
except Exception as exc:
raise HTTPException(503, "Delegated OAuth coordination is unavailable") from exc
return result == 1
async def delegation_active(principal: SessionPrincipal, refresh_jti: str | None = None) -> bool:
return (
principal.delegation is not None
and is_delegated_callback(principal.delegation.redirect_uri)
and await _family_transition(principal, "active", refresh_jti or "")
)
async def revoke_delegation(principal: SessionPrincipal) -> None:
await _family_transition(principal, "revoke")
async def delegated_user(user_id: str) -> LiteLLM_UserTable:
from litellm.proxy import proxy_server
policy: Final = _AuthenticationPolicy.model_validate(vars(proxy_server))
if policy.user_custom_auth is not None or any(
policy.general_settings.get(flag) is True for flag in ("enable_oauth2_auth", "enable_oauth2_proxy_auth")
):
raise HTTPException(503, "Delegated OAuth requires the gateway database authentication policy")
user: Final = await load_active_user_by_id(user_id, source="database")
if isinstance(user, str):
raise HTTPException(503 if user in ("unavailable", "faulted", "unresolvable") else 401, "User is unavailable")
if user.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(403, "Delegated admin access requires an active proxy admin")
return user
async def delegated_identity(principal: SessionPrincipal) -> UserAPIKeyAuth:
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
try:
user: Final = await delegated_user(principal.user_id)
if (principal.team_id is None and user.teams) or (
principal.team_id is not None and principal.team_id not in user.teams
):
raise HTTPException(403, "The selected team is no longer available; authorize again")
team: Final = (
await get_team_object(principal.team_id, prisma_client, user_api_key_cache, check_db_only=True)
if principal.team_id is not None
else None
)
if team is not None and (
team.blocked or not any(member.user_id == user.user_id for member in team.members_with_roles)
):
raise HTTPException(403, "The selected team is blocked or its membership has been removed")
membership: Final = (
await get_team_membership(user.user_id, team.team_id, prisma_client, user_api_key_cache, check_db_only=True)
if team is not None
else None
)
permission: Final = (
await get_object_permission(
user.object_permission_id, prisma_client, user_api_key_cache, check_db_only=True
)
if user.object_permission_id is not None
else None
)
identity: Final = UserAPIKeyAuth.model_validate(
{
**team_grants(team, membership, user.user_id),
"user_id": user.user_id,
"user_role": user.user_role,
"user_email": user.user_email,
"team_id": principal.team_id,
"org_id": team.organization_id if team is not None else user.organization_id,
"models": () if team is not None else user_models(user),
"user_tpm_limit": user.tpm_limit,
"user_rpm_limit": user.rpm_limit,
"user_max_budget": user.max_budget,
"user_model_max_budget": _UserBudget.model_validate(user.model_dump()).model_max_budget,
"user_spend": user.spend,
"object_permission": permission,
"object_permission_id": user.object_permission_id,
}
)
identity.requires_fresh_policy = True
return identity
except HTTPException:
raise
except Exception as exc:
raise HTTPException(503, "Delegated user policy is unavailable") from exc
async def issue_delegated_tokens(
principal: SessionPrincipal,
keys: SessionSigningKeys,
now: datetime,
previous_refresh_jti: str | None = None,
) -> Response:
try:
if principal.delegation is None or not is_delegated_callback(principal.delegation.redirect_uri):
raise HTTPException(400, "The delegated callback is not approved")
await delegated_identity(principal)
access: Final = mint_session_token(principal, keys, now)
refresh: Final = mint_session_refresh_token(principal, keys, now)
if not isinstance(access, MintedSessionToken) or not isinstance(refresh, MintedSessionToken):
raise HTTPException(500, "Could not issue delegated tokens")
opened: Final = open_session_refresh_token(refresh.token.get_secret_value(), keys, now)
if not isinstance(opened, OpenedSessionToken):
raise HTTPException(500, "Could not issue delegated tokens")
if not await _family_transition(
principal, "create" if previous_refresh_jti is None else "rotate", previous_refresh_jti or "", opened.jti
):
raise HTTPException(400, "The grant is expired, revoked, or reused; authorize again")
return JSONResponse(
{
"access_token": access.token.get_secret_value(),
"token_type": "Bearer",
"expires_in": int((access.expires_at - now).total_seconds()),
"refresh_token": refresh.token.get_secret_value(),
"scope": principal.delegation.scope,
},
headers=TOKEN_NO_CACHE_HEADERS,
)
except HTTPException as exc:
return JSONResponse(
{
"error": "temporarily_unavailable" if exc.status_code >= 500 else "invalid_grant",
"error_description": exc.detail,
},
status_code=503 if exc.status_code >= 500 else 400,
headers=TOKEN_NO_CACHE_HEADERS,
)
async def authenticate_delegated_request(request: Request, token: str, route: str) -> UserAPIKeyAuth:
from litellm.proxy.proxy_server import master_key
keys: Final = active_session_signing_keys(master_key) if master_key else None
if keys is None or isinstance(keys, SessionSigningConfigError):
raise HTTPException(503, "The gateway token signing configuration is unavailable")
opened: Final = open_session_token(token, keys, datetime.now(timezone.utc))
if not isinstance(opened, OpenedSessionToken) or opened.principal.delegation is None:
raise HTTPException(401, "Invalid delegated access token")
if opened.principal.delegation.resource != canonical_resource_uri(get_request_base_url(request)):
raise HTTPException(401, "The delegated token belongs to a different gateway")
if not RouteChecks.is_delegated_admin_route(route, request):
raise HTTPException(403, "This route is outside the approved application scope")
if not await delegation_active(opened.principal):
raise HTTPException(401, "The application grant has expired or been revoked")
identity: Final = await delegated_identity(opened.principal)
identity.token = hash_token(token)
identity.key_alias = "oauth-admin"
return identity

View file

@ -15,6 +15,12 @@ from litellm.proxy._types import (
)
from .auth_checks_organization import _user_is_org_admin
from .auth_utils import (
get_request_route_template,
request_dispatched_to_marked_provider_pass_through,
request_dispatched_to_pass_through_endpoint,
request_dispatched_to_provider_pass_through,
)
# Management write routes denied to PROXY_ADMIN_VIEW_ONLY. Adding a new write
# endpoint to a management router REQUIRES adding it here too — the surrounding
@ -68,6 +74,49 @@ _AUTH_ENFORCED_PASS_THROUGH_ROUTE_GROUPS: Final = frozenset(("openai_routes", "l
class RouteChecks:
@staticmethod
def is_delegated_admin_route(route: str, request: Request) -> bool:
template: Final = get_request_route_template(request) or route
excluded: Final = (
*LiteLLMRoutes.master_key_only_routes.value,
*LiteLLMRoutes.mcp_routes.value,
"/jwt/*",
"/config/*",
"/get/config/*",
"/sso/*",
"/session/*",
"/user/auth",
"/user/password/*",
"/callbacks/*",
"/team/{team_id:path}/callback",
"/team/{team_id:path}/callback/{callback_name}",
)
allowed: Final = (
*LiteLLMRoutes.management_routes.value,
*LiteLLMRoutes.self_managed_routes.value,
*LiteLLMRoutes.org_admin_only_routes.value,
*LiteLLMRoutes.openai_routes.value,
*LiteLLMRoutes.anthropic_routes.value,
*LiteLLMRoutes.google_routes.value,
*LiteLLMRoutes.admin_viewer_routes.value,
*LiteLLMRoutes.global_spend_tracking_routes.value,
"/budget/new",
"/budget/update",
"/budget/delete",
"/budget/info",
"/organization/new",
"/organization/update",
"/organization/list",
)
return (
route.isprintable()
and not request_dispatched_to_pass_through_endpoint(request)
and not request_dispatched_to_provider_pass_through(request)
and not request_dispatched_to_marked_provider_pass_through(request)
and not RouteChecks.check_route_access(template, excluded)
and RouteChecks.check_route_access(template, allowed)
)
@staticmethod
def should_call_route(
route: str,

View file

@ -669,7 +669,7 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
"path": ws_scope.get("path", ""),
"state": ws_scope.setdefault("state", {}),
}
for key in ("root_path", "app_root_path"):
for key in ("root_path", "app_root_path", "endpoint", "path_params", "route"):
if key in ws_scope:
synthetic_scope[key] = ws_scope[key]
request: Final = Request(scope=synthetic_scope)
@ -1676,7 +1676,15 @@ async def _user_api_key_auth_builder(
if general_settings.get("enable_oauth2_proxy_auth", False) is True:
return await handle_oauth2_proxy_request(request=request)
if general_settings.get("enable_jwt_auth", False) is True:
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
is_session_bearer_shaped,
)
from litellm.proxy.auth.delegated_oauth import authenticate_delegated_request
if api_key is not None and is_session_bearer_shaped(api_key):
valid_token = await authenticate_delegated_request(request, api_key, route)
if valid_token is None and general_settings.get("enable_jwt_auth", False) is True:
is_jwt = jwt_handler.is_jwt(token=api_key)
verbose_proxy_logger.debug("is_jwt: %s", is_jwt)
if is_jwt:
@ -1908,7 +1916,7 @@ async def _user_api_key_auth_builder(
#### ELSE ####
## CHECK PASS-THROUGH ENDPOINTS ##
if not custom_auth_api_key:
if not custom_auth_api_key and not (valid_token is not None and valid_token.requires_fresh_policy):
response = await check_api_key_for_custom_headers_or_pass_through_endpoints(
request=request,
route=route,
@ -2060,6 +2068,7 @@ async def _user_api_key_auth_builder(
valid_token is not None
and isinstance(valid_token, UserAPIKeyAuth)
and valid_token.user_role == LitellmUserRoles.PROXY_ADMIN
and not valid_token.requires_fresh_policy
):
if valid_token.expires is not None:
current_time = datetime.now(timezone.utc)
@ -2095,6 +2104,7 @@ async def _user_api_key_auth_builder(
and isinstance(valid_token, UserAPIKeyAuth)
and valid_token.team_id is not None
and valid_token.team_id != UI_TEAM_ID
and not valid_token.requires_fresh_policy
):
## UPDATE TEAM VALUES BASED ON CACHED TEAM OBJECT - allows `/team/update` values to work for cached token
try:
@ -2326,8 +2336,11 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=valid_token.requires_fresh_policy,
)
except Exception as e:
if valid_token.requires_fresh_policy:
raise
verbose_logger.debug(
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - %s",
e,
@ -2372,11 +2385,12 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
if prisma_client is not None and _user_id is not None and _team_id is not None:
_cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id)
team_member_info = await user_api_key_cache.async_get_cache(
key=_cache_key,
model_type=LiteLLM_TeamMembership,
team_member_info = (
await get_team_membership(_user_id, _team_id, prisma_client, user_api_key_cache, check_db_only=True)
if valid_token.requires_fresh_policy
else await user_api_key_cache.async_get_cache(key=_cache_key, model_type=LiteLLM_TeamMembership)
)
if team_member_info is None:
if team_member_info is None and not valid_token.requires_fresh_policy:
# read from DB
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
where={
@ -2560,8 +2574,11 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=valid_token.requires_fresh_policy,
)
except HTTPException:
if valid_token.requires_fresh_policy:
raise
token_team_models: Final = _token_team_models(valid_token)
_team_obj = LiteLLM_TeamTableCachedObj(
team_id=valid_token.team_id,
@ -2669,11 +2686,12 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
route=route,
start_time=start_time,
)
virtual_key_auth_obj.via_virtual_key = True
virtual_key_auth_obj.requires_fresh_policy = valid_token.requires_fresh_policy
virtual_key_auth_obj.via_virtual_key = not valid_token.requires_fresh_policy
return virtual_key_auth_obj
async def _safe_fetch(label: str, awaitable):
async def _safe_fetch(label: str, awaitable, *, fail_closed: bool = False):
"""Run an awaitable and return its result. Re-raises authentication /
authorization failures (HTTPException, ProxyException,
BudgetExceededError) so they propagate to the caller.
@ -2693,6 +2711,8 @@ async def _safe_fetch(label: str, awaitable):
)
raise
except Exception as e:
if fail_closed:
raise
verbose_proxy_logger.debug(
"centralized auth: %s fetch swallowed (%s: %s)",
label,
@ -2782,6 +2802,7 @@ async def _inherit_org_identity(
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=user_api_key_auth_obj.requires_fresh_policy,
)
if org_object is None:
return
@ -2908,7 +2929,9 @@ async def _run_centralized_common_checks(
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=user_api_key_auth_obj.requires_fresh_policy,
),
fail_closed=user_api_key_auth_obj.requires_fresh_policy,
)
)
else:
@ -2925,7 +2948,9 @@ async def _run_centralized_common_checks(
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=user_api_key_auth_obj.requires_fresh_policy,
),
fail_closed=user_api_key_auth_obj.requires_fresh_policy,
)
)
else:
@ -2941,6 +2966,7 @@ async def _run_centralized_common_checks(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
),
fail_closed=user_api_key_auth_obj.requires_fresh_policy,
)
)
else:
@ -2960,6 +2986,7 @@ async def _run_centralized_common_checks(
token_end_user_max_budget=user_api_key_auth_obj.end_user_max_budget,
key_end_user_budget_id=key_end_user_budget_id,
),
fail_closed=user_api_key_auth_obj.requires_fresh_policy,
)
)
else:
@ -2975,6 +3002,7 @@ async def _run_centralized_common_checks(
token=user_api_key_auth_obj.token or "",
proxy_logging_obj=proxy_logging_obj,
),
fail_closed=user_api_key_auth_obj.requires_fresh_policy,
)
)
@ -3004,6 +3032,8 @@ async def _run_centralized_common_checks(
end_user_result,
global_spend_result,
):
if user_api_key_auth_obj.requires_fresh_policy and isinstance(r, BaseException):
raise r
if isinstance(r, (ProxyException, litellm.BudgetExceededError)):
raise r
@ -3058,7 +3088,10 @@ async def _run_centralized_common_checks(
# caller. The token is the source of truth for these paths — force
# the admin user_object whenever the token says PROXY_ADMIN, even
# if a DB row was fetched.
if user_api_key_auth_obj.user_role == LitellmUserRoles.PROXY_ADMIN:
if (
user_api_key_auth_obj.user_role == LitellmUserRoles.PROXY_ADMIN
and not user_api_key_auth_obj.requires_fresh_policy
):
user_object = LiteLLM_UserTable(
user_id=user_api_key_auth_obj.user_id or litellm_proxy_admin_name,
user_role=LitellmUserRoles.PROXY_ADMIN,

View file

@ -12,18 +12,28 @@ def render_native_client_consent_page(
teams: Sequence[tuple[str, str]],
flow_handle: str,
complete_url: str,
delegated: bool = False,
) -> str:
"""The consent page a native client's sign-in lands on: who is signed in, which
loopback client asked, which team the credential is attributed to, and an explicit
Approve or Deny that POSTs back to ``complete_url``. Every value is client- or
user-influenced and HTML-escaped; the flow handle travels only in the form body."""
title: Final = "Authorize application access" if delegated else "Authorize CLI access"
client: Final = "An application" if delegated else "A command-line client"
permission: Final = (
"Approving allows this application to manage keys, users, teams and budgets with your current admin "
"permissions until you disconnect or the authorization expires. Only approve an application you trust."
if delegated
else f"Approving issues it a personal credential that expires within {CLI_JWT_EXPIRATION_HOURS} hours. "
"<code>lite logout</code> stops it from being renewed. Only approve if you started this sign-in yourself."
)
return f"""<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<meta name="referrer" content="no-referrer">
<title>Authorize CLI access - LiteLLM</title>
<title>{title} - LiteLLM</title>
<style>
body {{
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, Oxygen, Ubuntu, Cantarell, sans-serif;
@ -57,9 +67,9 @@ button {{ flex: 1; padding: 10px; border-radius: 6px; font-size: 15px; cursor: p
</head>
<body>
<div class="container">
<h1>Authorize CLI access</h1>
<p>A command-line client at <code>{escape(client_origin)}</code> wants to call LiteLLM as <strong>{escape(user_id)}</strong>.</p>
<p>Approving issues it a personal credential that expires within {CLI_JWT_EXPIRATION_HOURS} hours. <code>lite logout</code> stops it from being renewed. Only approve if you started this sign-in yourself.</p>
<h1>{title}</h1>
<p>{client} at <code>{escape(client_origin)}</code> wants to call LiteLLM as <strong>{escape(user_id)}</strong>.</p>
<p>{permission}</p>
<form method="post" action="{escape(complete_url)}">
<input type="hidden" name="flow" value="{escape(flow_handle)}">
{_team_field(teams)}

View file

@ -1,6 +1,15 @@
from typing import Final
from fastapi import Request
from fastapi import APIRouter, Request
from fastapi.routing import APIRoute, APIWebSocketRoute
from litellm.types.passthrough_endpoints.pass_through_endpoints import LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER
def mark_provider_pass_through_routes(router: APIRouter) -> None:
for route in router.routes:
if isinstance(route, (APIRoute, APIWebSocketRoute)):
setattr(route.endpoint, LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER, True)
def get_litellm_virtual_key(request: Request) -> str:

View file

@ -89,7 +89,7 @@ from litellm.proxy.common_utils.resource_ownership import is_proxy_admin
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
)
from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key
from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key, mark_provider_pass_through_routes
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
HttpPassThroughEndpointHelpers,
create_pass_through_route,
@ -4191,3 +4191,6 @@ async def watsonx_proxy_route(
fastapi_response,
user_api_key_dict,
)
mark_provider_pass_through_routes(router)

View file

@ -8,6 +8,7 @@ from fastapi import APIRouter, Depends, Request, Response
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.pass_through_endpoints.common_utils import mark_provider_pass_through_routes
router: Final = APIRouter()
@ -42,3 +43,6 @@ async def openai_passthrough_route(
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
)
mark_provider_pass_through_routes(router)

View file

@ -881,6 +881,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolParam,
HttpxBinaryResponseContent,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
ModelGroupInfoProxy,
@ -12910,6 +12911,8 @@ async def vertex_ai_live_passthrough_endpoint(
)
setattr(vertex_ai_live_passthrough_endpoint, LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER, True)
######################################################################
# /v1/realtime Endpoints

View file

@ -23,6 +23,7 @@ from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata
from litellm.proxy.pass_through_endpoints.common_utils import mark_provider_pass_through_routes
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
create_pass_through_route,
)
@ -225,3 +226,6 @@ async def langfuse_proxy_route(
)
return received_value
mark_provider_pass_through_routes(router)

View file

@ -21,6 +21,7 @@ LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY: Final = "litellm_pass_thro
# custom path that collides with a built-in route never suppresses model-access checks:
# on a collision FastAPI dispatches the built-in handler, which does not carry this flag.
LITELLM_PASS_THROUGH_ENDPOINT_MARKER: Final = "__litellm_pass_through_endpoint__"
LITELLM_PROVIDER_PASS_THROUGH_ENDPOINT_MARKER: Final = "__litellm_provider_pass_through_endpoint__"
class EndpointType(str, Enum):

View file

@ -0,0 +1,211 @@
import base64
import hashlib
import re
import secrets
import time
from pathlib import Path
from typing import Final
from urllib.parse import parse_qs, urlsplit
import httpx
import jwt
import yaml
from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
RESOURCE: Final = "https://gateway.integration.example"
CALLBACK: Final = "https://app.integration.example/callback"
def _authorize(gateway: Gateway, user: str) -> dict[str, str]:
registration: Final = gateway.client.post("/register", json={"redirect_uris": [CALLBACK]})
assert registration.status_code == 201, registration.text
client: Final = string_value(JSON_OBJECT.validate_json(registration.content)["client_id"])
verifier: Final = secrets.token_urlsafe(48)
challenge: Final = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
cookies: Final = {
"token": jwt.encode(
{"user_id": user, "login_method": "username_password", "exp": int(time.time()) + 600},
gateway.key,
algorithm="HS256",
)
}
consent: Final = gateway.client.get(
"/authorize",
params={
"client_id": client,
"redirect_uri": CALLBACK,
"response_type": "code",
"state": user,
"code_challenge": challenge,
"code_challenge_method": "S256",
"resource": RESOURCE,
"scope": "proxy:admin",
},
cookies=cookies,
)
assert consent.status_code == 200, consent.text
flow: Final = re.search(r'name="flow" value="([^"]+)"', consent.text)
assert flow is not None, consent.text
approved: Final = gateway.client.post(
"/authorize/complete",
data={"flow": flow[1], "decision": "approve"},
cookies={**cookies, **dict(consent.cookies)},
)
assert approved.status_code == 303, approved.text
callback: Final = parse_qs(urlsplit(approved.headers["location"]).query)
assert callback["state"] == [user]
return {
"grant_type": "authorization_code",
"client_id": client,
"code": callback["code"][0],
"redirect_uri": CALLBACK,
"code_verifier": verifier,
"resource": RESOURCE,
}
def _refresh(gateway: Gateway, client: str, token: str) -> httpx.Response:
return gateway.client.post(
"/token",
data={
"grant_type": "refresh_token",
"client_id": client,
"refresh_token": token,
"resource": RESOURCE,
},
)
def test_delegated_grants_share_rotation_revocation_and_live_role_checks(gateway: Gateway, tmp_path: Path) -> None:
overrides: Final = {"PROXY_BASE_URL": RESOURCE, "LITELLM_OAUTH_ADMIN_REDIRECT_URIS": CALLBACK}
with (
owned_proxy(gateway, tmp_path, overrides) as first,
owned_proxy(gateway, tmp_path, overrides) as second,
first.scenario() as scenario,
):
admin: Final = scenario.user(user_role="proxy_admin")
for outcome in ("replay", "revoke", "demote"):
grant: Final = _authorize(first, admin)
client: Final = grant["client_id"]
if outcome == "replay":
first.post("/user/update", {"user_id": admin, "user_role": "internal_user"})
assert second.client.post("/token", data=grant).status_code == 400
first.post("/user/update", {"user_id": admin, "user_role": "proxy_admin"})
issued: Final = first.client.post("/token", data=grant)
assert issued.status_code == 200, issued.text
original: Final = JSON_OBJECT.validate_json(issued.content)
access: Final = string_value(original["access_token"])
refresh: Final = string_value(original["refresh_token"])
created: Final = first.post("/key/generate", {"user_id": admin}, key=access)
key: Final = string_value(created["key"])
scenario.cleanups.callback(scenario.delete_key, key)
read: Final = second.request("GET", "/key/info", key=access, params={"key": key})
assert read.status_code == 200, read.text
assert object_value(JSON_OBJECT.validate_json(read.content)["info"])["user_id"] == admin
rotated: Final = _refresh(second, client, refresh)
assert rotated.status_code == 200, rotated.text
pair: Final = JSON_OBJECT.validate_json(rotated.content)
next_access: Final = string_value(pair["access_token"])
next_refresh: Final = string_value(pair["refresh_token"])
assert next_refresh != refresh
assert second.client.post("/token", data=grant).status_code == 400
warmed: Final = first.request("GET", "/user/info", key=next_access, params={"user_id": admin})
assert warmed.status_code == 200, warmed.text
if outcome == "replay":
assert _refresh(first, client, refresh).status_code == 400
elif outcome == "revoke":
revoked: Final = first.client.post("/revoke", data={"client_id": client, "token": next_refresh})
assert revoked.status_code == 200, revoked.text
else:
first.post("/user/update", {"user_id": admin, "user_role": "internal_user"})
for worker in (first, second):
assert worker.client.post("/token", data=grant).status_code == 400
for bearer in (access, next_access):
denied: Final = worker.request("GET", "/user/info", key=bearer, params={"user_id": admin})
assert denied.status_code in (401, 403), denied.text
renewal: Final = _refresh(worker, client, next_refresh)
assert renewal.status_code == 400, renewal.text
def test_delegated_inference_keeps_user_models_budgets_and_attribution(gateway: Gateway, tmp_path: Path) -> None:
config: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
path: Final = tmp_path / "delegated-model-policy.yaml"
path.write_text(
yaml.safe_dump(
{
**config,
"general_settings": {
**object_value(config["general_settings"]),
"enable_jwt_auth": True,
"pass_through_endpoints": [
{
"path": "/v1/chat/completions",
"target": gateway.upstream_url,
"methods": ["GET"],
"auth": True,
"include_subpath": True,
},
{"path": "/laya/v1/systemone", "target": gateway.upstream_url, "auth": True},
],
},
"litellm_settings": {
**object_value(config["litellm_settings"]),
"overwrite_user_with_key_hash": True,
},
}
)
)
overrides: Final = {"PROXY_BASE_URL": RESOURCE, "LITELLM_OAUTH_ADMIN_REDIRECT_URIS": CALLBACK}
with (
owned_proxy(gateway, tmp_path, overrides, config=path) as candidate,
candidate.scenario() as scenario,
httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as upstream,
):
allowed: Final = scenario.model(model="openai/gpt-6.1-sol", input_cost_per_token=0.001, num_retries=0)
blocked: Final = scenario.model(model="openai/gpt-6.1-sol", input_cost_per_token=0.001, num_retries=0)
admin: Final = scenario.user(user_role="proxy_admin", models=[allowed], max_budget=1)
issued: Final = candidate.client.post("/token", data=_authorize(candidate, admin))
assert issued.status_code == 200, issued.text
token: Final = string_value(JSON_OBJECT.validate_json(issued.content)["access_token"])
assert candidate.request("GET", "/v1/chat/completions/health").status_code == 200
for route in ("/v1/chat/completions/health", "/v1/chat/completions", "/laya/v1/systemone"):
assert candidate.request("GET", route, key=token).status_code == 403
assert candidate.request("POST", "/laya/v1/systemone", {}, key=token).status_code == 403
body: Final = {
"model": allowed,
"messages": [{"role": "user", "content": "policy check"}],
"user": admin,
"cache": {"no-cache": True},
}
for model, fallbacks, route, status in (
(allowed, [], "/v1/chat/completions", 200),
(blocked, [], "/v1/chat/completions", 403),
(allowed, [blocked], "/v1/chat/completions", 403),
(allowed, [allowed], "/v1/chat/completions", 200),
(allowed, [], f"/openai/deployments/{allowed}/chat/completions", 200),
(blocked, [], f"/openai/deployments/{blocked}/chat/completions", 403),
):
upstream.get("/__observations").raise_for_status()
response: Final = candidate.request(
"POST", route, {**body, "model": model, "fallbacks": fallbacks}, key=token
)
assert response.status_code == status, response.text
requests: Final = object_value(upstream.get("/__observations").json())["requests"]
assert isinstance(requests, list) and len(requests) == int(status == 200), requests
if requests:
assert object_value(object_value(requests[0])["body"])["user"] == admin
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_UserTable" WHERE user_id=%s', (admin,)),
lambda rows: bool(rows) and float(str(rows[0]["spend"])) > 0,
seconds=70,
)
for budget in (
{"model_max_budget": {allowed: {"max_budget": 0, "budget_duration": "1d"}}},
{"model_max_budget": {}, "max_budget": 0},
):
candidate.post("/user/update", {"user_id": admin, **budget})
denied: Final = candidate.request("POST", "/v1/chat/completions", body, key=token)
assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text
assert upstream.get("/__observations").json()["requests"] == []

View file

@ -1,6 +1,7 @@
"""Tests for the session-token KDF and the edge/token-endpoint resolvers."""
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from cryptography.hazmat.primitives import serialization
@ -127,6 +128,13 @@ def test_resolve_rejects_refresh_token_at_the_edge():
assert result.expired is False
def test_mcp_refuses_a_proxy_api_access_token() -> None:
principal: Final = SessionPrincipal(user_id="admin", client_id="app", audience="proxy_api")
minted: Final = mint_session_token(principal, KEYS, NOW)
assert isinstance(minted, MintedSessionToken)
assert isinstance(resolve_session_bearer(minted.token.get_secret_value(), KEYS, NOW), SessionBearerInvalid)
def test_resolve_wrong_master_key_fails_closed():
other_keys = session_keys_from_master_key("sk-rotated-master-key")
result = resolve_session_bearer(f"Bearer {_access_token()}", other_keys, NOW)

View file

@ -1,6 +1,7 @@
"""Tests for the identity-only gateway session token (mint/open, hostile-input totality)."""
from datetime import datetime, timedelta, timezone
from typing import Final
import jwt
import pytest
@ -17,6 +18,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token i
SESSION_TOKEN_PREFIX,
SESSION_TTL_SECONDS,
AsymmetricSessionKeys,
DelegatedGrant,
MintedSessionToken,
NotASessionToken,
OpenedSessionToken,
@ -113,6 +115,55 @@ def test_refresh_round_trip_recovers_principal_and_caps_ttl():
assert opened.principal == PRINCIPAL
@pytest.mark.parametrize("refresh", (False, True))
def test_delegated_tokens_preserve_consent_and_cannot_outlive_it(refresh: bool) -> None:
deadline: Final = NOW + timedelta(seconds=30)
principal: Final = SessionPrincipal(
user_id="user-123", client_id="llm_client_abc", audience="proxy_api", team_id="team-b",
delegation=DelegatedGrant(
grant_id="grant-1", resource="https://gateway.example", redirect_uri="https://app.example/callback",
expires_at=int(deadline.timestamp()),
),
)
mint: Final = mint_session_refresh_token if refresh else mint_session_token
opener: Final = open_session_refresh_token if refresh else open_session_token
minted: Final = mint(principal, KEYS, NOW)
assert isinstance(minted, MintedSessionToken)
assert minted.expires_at == deadline
opened: Final = opener(minted.token.get_secret_value(), KEYS, NOW)
assert isinstance(opened, OpenedSessionToken)
assert opened.principal == principal
assert isinstance(opener(minted.token.get_secret_value(), KEYS, deadline), SessionExpired)
def test_revocation_accepts_expired_delegated_access_only_until_consent_expires() -> None:
deadline: Final = NOW + timedelta(seconds=SESSION_TTL_SECONDS * 2)
delegation: Final = DelegatedGrant(
grant_id="grant-1", resource="https://gateway.example", redirect_uri="https://app.example/callback",
expires_at=int(deadline.timestamp()),
)
principal: Final = SessionPrincipal(user_id="u1", client_id="c1", audience="proxy_api", delegation=delegation)
minted: Final = mint_session_token(principal, KEYS, NOW)
assert isinstance(minted, MintedSessionToken)
token: Final = minted.token.get_secret_value()
expired: Final = NOW + timedelta(seconds=SESSION_TTL_SECONDS)
assert isinstance(open_session_token(token, KEYS, expired), SessionExpired)
assert isinstance(open_session_token(token, KEYS, expired, for_revocation=True), OpenedSessionToken)
assert isinstance(open_session_token(token, KEYS, deadline, for_revocation=True), SessionExpired)
assert isinstance(open_session_token(_corrupt_signature(token), KEYS, expired, for_revocation=True), SessionBadSignature)
assert isinstance(open_session_token(_mint_access(), KEYS, expired, for_revocation=True), SessionExpired)
@pytest.mark.parametrize("audience,server", ((None, None), ("proxy_api", "mcp-server")))
def test_delegation_cannot_be_reused_as_an_mcp_grant(audience: str | None, server: str | None) -> None:
claims: Final = _valid_claims(
audience=audience, resource_server_id=server,
delegation={"grant_id": "g", "resource": "https://gateway.example", "redirect_uri": "https://app.example/cb",
"expires_at": int(NOW.timestamp()) + 600, "scope": "proxy:admin"},
)
assert isinstance(open_session_token(_sign_claims(claims), KEYS, NOW), SessionMalformed)
def test_access_token_reprefixed_as_refresh_is_rejected_by_signed_kind():
body = _mint_access().removeprefix(SESSION_TOKEN_PREFIX)
swapped = SESSION_REFRESH_PREFIX + body

View file

@ -9386,6 +9386,23 @@ def test_aggregate_wellknown_routes_serve_gateway_metadata():
assert asm.json()["authorization_endpoint"] == "http://testserver/authorize/mcp-session"
assert "none" in asm.json()["token_endpoint_auth_methods_supported"]
api: Final = client.get("/.well-known/oauth-authorization-server/oauth/api")
assert api.status_code == 200
assert api.json() == {
"issuer": "http://testserver/oauth/api",
"authorization_endpoint": "http://testserver/authorize",
"token_endpoint": "http://testserver/token",
"registration_endpoint": "http://testserver/register",
"revocation_endpoint": "http://testserver/revoke",
"introspection_endpoint": "http://testserver/introspect",
"response_types_supported": ["code"],
"grant_types_supported": ["authorization_code", "refresh_token"],
"scopes_supported": ["proxy:admin"],
"code_challenge_methods_supported": ["S256"],
"token_endpoint_auth_methods_supported": ["none"],
"revocation_endpoint_auth_methods_supported": ["none"],
}
def test_as_aggregate_route_reserves_mcp_for_the_aggregate():
"""The single-segment /.well-known/oauth-authorization-server/mcp is reserved for the

View file

@ -7,9 +7,11 @@ from base64 import urlsafe_b64encode
from datetime import datetime, timedelta, timezone
from http.cookies import SimpleCookie
from typing import Final
from unittest.mock import AsyncMock, MagicMock
from urllib.parse import parse_qs, urlparse
import pytest
from fastapi import HTTPException
from starlette.requests import Request
from litellm.caching.caching import DualCache
@ -53,10 +55,15 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credent
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
SESSION_ISSUER,
SESSION_REFRESH_PREFIX,
OpenedSessionToken,
SessionPrincipal,
mint_session_refresh_token,
mint_session_token,
open_session_token,
)
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.auth.delegated_oauth import authenticate_delegated_request
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
MASTER_KEY = "sk-gateway-dcr-flow-tests"
REDIRECT_URI = "https://claude.ai/api/mcp/auth_callback"
@ -1373,9 +1380,30 @@ async def test_resource_resolution_is_identity_not_ip_filtered_access():
LOOPBACK_REDIRECT_URI = "http://127.0.0.1:51234/callback"
PROXY_API_RESOURCE = "https://llm.example.com"
HOSTED_CALLBACK: Final = "https://app.example/callback"
CONSENT_TEAMS = (ConsentTeam(team_id="team-a", team_alias="Team A"), ConsentTeam(team_id="team-b"))
@pytest.fixture
def delegated_database(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
from litellm.proxy import proxy_server
database: Final = MagicMock()
read: Final = AsyncMock(return_value=LiteLLM_UserTable(user_id="u1", user_role="proxy_admin"))
database.writer_db.litellm_usertable.find_unique = read
redis: Final = MagicMock()
redis.async_register_script.return_value = AsyncMock(return_value=1)
redis.async_increment = DualCache().async_increment_cache
monkeypatch.setenv("LITELLM_OAUTH_ADMIN_REDIRECT_URIS", HOSTED_CALLBACK)
monkeypatch.setattr(proxy_server, "master_key", MASTER_KEY)
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "redis_usage_cache", redis)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(proxy_server, "user_custom_auth", None)
return read
class _Minter:
def __init__(self, result=None):
self.calls = []
@ -1475,6 +1503,88 @@ def _opened_refresh(refresh_token, client_id):
return opened.principal
@pytest.mark.asyncio
async def test_hosted_consent_and_code_keep_pkce_client_resource_and_delegation_bound(
delegated_database: AsyncMock,
) -> None:
client_id: Final = (await _register([HOSTED_CALLBACK]))["client_id"]
other_client: Final = (await _register([HOSTED_CALLBACK]))["client_id"]
cache: Final = DualCache()
consent: Final = await _native_authorize(
client_id, redirect_uri=HOSTED_CALLBACK, scope="proxy:admin", lookup=_ConsentTeams(()),
)
assert consent.status_code == 200
omitted: Final = await _complete_consent(consent, cache=cache)
assert omitted.status_code == 400
approved: Final = await _complete_consent(consent, cache=cache, decision="approve")
assert approved.status_code == 303
code: Final = _code_from(approved)
wire: Final = _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, _AUTH_CODE_DEBUG_KEY)
assert wire["delegation"]["scope"] == "proxy:admin"
assert wire["delegation"]["resource"] == PROXY_API_RESOURCE
assert wire["delegation"]["redirect_uri"] == HOSTED_CALLBACK
for presented_client, verifier, resource, callback in (
(other_client, CODE_VERIFIER, PROXY_API_RESOURCE, HOSTED_CALLBACK),
(client_id, "x" * 43, PROXY_API_RESOURCE, HOSTED_CALLBACK),
(client_id, CODE_VERIFIER, "https://other.example", HOSTED_CALLBACK),
(client_id, CODE_VERIFIER, PROXY_API_RESOURCE, "https://other.example/callback"),
):
refused: Final = await _redeem(
code, presented_client, cache=cache, redirect_uri=callback, code_verifier=verifier, resource=resource,
)
assert refused.status_code == 400
issued: Final = await _redeem(
code, client_id, cache=cache, redirect_uri=HOSTED_CALLBACK, resource=PROXY_API_RESOURCE,
)
assert issued.status_code == 200
body: Final = json.loads(issued.body)
opened: Final = open_session_token(
body["access_token"], session_keys_from_master_key(MASTER_KEY), datetime.now(timezone.utc),
)
assert isinstance(opened, OpenedSessionToken)
assert opened.principal.delegation is not None
assert opened.principal.delegation.model_dump() == wire["delegation"]
assert _opened_refresh(body["refresh_token"], client_id) == opened.principal
renewed: Final = await _refresh_native(body["refresh_token"], client_id, None, cache)
assert json.loads(renewed.body)["scope"] == "proxy:admin"
identity: Final = await authenticate_delegated_request(_request("/user/info"), body["access_token"], "/user/info")
assert (identity.user_id, identity.requires_fresh_policy) == ("u1", True)
status, claims = await _introspect(body["access_token"])
assert (status, claims["active"], claims["sub"]) == (200, True, "u1")
with pytest.raises(HTTPException, match=r"403:.*scope"):
await authenticate_delegated_request(_request("/mcp"), body["access_token"], "/mcp")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"scope,callback",
((None, HOSTED_CALLBACK), ("proxy:admin", "https://other.example/cb"), ("admin", HOSTED_CALLBACK)),
)
async def test_hosted_authorize_requires_scope_and_exact_approved_callback(
delegated_database: AsyncMock, scope: str | None, callback: str,
) -> None:
client_id: Final = (await _register([callback]))["client_id"]
response: Final = await _native_authorize(client_id, redirect_uri=callback, scope=scope)
assert response.status_code == 400
assert "set-cookie" not in response.headers
delegated_database.assert_not_awaited()
@pytest.mark.asyncio
async def test_hosted_authorize_reuses_gateway_sso_and_denies_nonadmins(delegated_database: AsyncMock) -> None:
client_id: Final = (await _register([HOSTED_CALLBACK]))["client_id"]
login: Final = await _native_authorize(
client_id, redirect_uri=HOSTED_CALLBACK, scope="proxy:admin", session_user_id=None,
)
assert login.status_code == 303
assert login.headers["location"].startswith(PROXY_API_RESOURCE + "/sso/key/generate?return_to=")
delegated_database.assert_not_awaited()
delegated_database.return_value = LiteLLM_UserTable(user_id="u1", user_role="internal_user")
denied: Final = await _native_authorize(client_id, redirect_uri=HOSTED_CALLBACK, scope="proxy:admin")
assert denied.status_code == 403
assert "set-cookie" not in denied.headers
@pytest.mark.asyncio
async def test_native_authorize_renders_consent_page_and_sets_flow_cookie():
"""A native client (RFC 8707 resource = the proxy itself) gets the server-rendered consent

View file

@ -1,6 +1,7 @@
import asyncio
import json
from unittest.mock import AsyncMock, MagicMock, patch
from typing import Final
import httpx
import pytest
@ -45,6 +46,19 @@ class _EngineHttp500:
status = 500
@pytest.mark.asyncio
async def test_authoritative_policy_cannot_fall_back_during_database_outage() -> None:
identity: Final = UserAPIKeyAuth(user_id="admin")
identity.requires_fresh_policy = True
request: Final = Request({"type": "http", "method": "POST", "path": "/key/generate", "headers": []})
with patch("litellm.proxy.proxy_server.general_settings", {"allow_requests_on_db_unavailable": True}):
with pytest.raises(ProxyException) as error:
await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
httpx.ConnectError("Database unavailable"), request, {}, "/key/generate", None, "", identity
)
assert error.value.code == "503"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"db_error",

View file

@ -0,0 +1,165 @@
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from fastapi import HTTPException, Request
import litellm
from litellm.models.access_group import LiteLLM_AccessGroupTable
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import SessionPrincipal
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_OrganizationTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_UserTable,
Member,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import common_checks
from litellm.proxy.auth.delegated_oauth import delegated_identity
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import ProxyLogging
@pytest.fixture
def database(monkeypatch: pytest.MonkeyPatch) -> MagicMock:
db: Final = MagicMock()
db.writer_db.litellm_usertable.find_unique = AsyncMock(
return_value=LiteLLM_UserTable(user_id="admin", user_role="proxy_admin", teams=[])
)
db.db = db.writer_db
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", db)
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache())
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_auth", None)
return db
@pytest.mark.asyncio
@pytest.mark.parametrize("role,active,status", [("internal_user", True, 403), ("proxy_admin", False, 401)])
async def test_delegated_identity_rechecks_current_admin(
database: MagicMock, role: str, active: bool, status: int
) -> None:
principal: Final = SessionPrincipal(user_id="admin", client_id="app", audience="proxy_api")
assert (await delegated_identity(principal)).user_role == "proxy_admin"
database.writer_db.litellm_usertable.find_unique.return_value = LiteLLM_UserTable(
user_id="admin", user_role=role, metadata={"scim_active": active}
)
with pytest.raises(HTTPException) as error:
await delegated_identity(principal)
assert error.value.status_code == status
@pytest.mark.asyncio
async def test_delegated_identity_uses_current_team_limits_and_roster(database: MagicMock) -> None:
database.writer_db.litellm_usertable.find_unique.return_value = LiteLLM_UserTable(
user_id="admin", user_role="proxy_admin", teams=["team"], rpm_limit=13
)
database.writer_db.litellm_teamtable.find_unique = AsyncMock(
return_value=LiteLLM_TeamTable(
team_id="team",
models=["allowed-model"],
rpm_limit=10,
members_with_roles=[Member(user_id="admin", role="user")],
)
)
database.writer_db.litellm_teammembership.find_unique = AsyncMock(
return_value=LiteLLM_TeamMembership(
user_id="admin", team_id="team", litellm_budget_table=LiteLLM_BudgetTable(rpm_limit=3)
)
)
principal: Final = SessionPrincipal(user_id="admin", client_id="app", audience="proxy_api", team_id="team")
first: Final = await delegated_identity(principal)
assert (first.user_rpm_limit, first.team_rpm_limit, first.team_member_rpm_limit) == (13, 10, 3)
assert first.team_models == ["allowed-model"]
database.writer_db.litellm_teammembership.find_unique.return_value = LiteLLM_TeamMembership(
user_id="admin", team_id="team", litellm_budget_table=LiteLLM_BudgetTable(rpm_limit=1)
)
assert (await delegated_identity(principal)).team_member_rpm_limit == 1
database.writer_db.litellm_teamtable.find_unique.return_value = LiteLLM_TeamTable(team_id="team")
with pytest.raises(HTTPException) as error:
await delegated_identity(principal)
assert error.value.status_code == 403
@pytest.mark.asyncio
async def test_delegated_identity_fails_closed_when_database_is_unavailable(database: MagicMock) -> None:
database.writer_db.litellm_usertable.find_unique.side_effect = httpx.ConnectError("Database unavailable")
with pytest.raises(HTTPException) as error:
await delegated_identity(SessionPrincipal(user_id="admin", client_id="app"))
assert error.value.status_code == 503
async def _common_checks(token: UserAPIKeyAuth, team: LiteLLM_TeamTable | None = None) -> bool:
logging: Final = MagicMock(spec=ProxyLogging)
logging.budget_alerts = AsyncMock()
return await common_checks(
request_body={"model": "gpt-4"},
team_object=team,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/v1/chat/completions",
llm_router=None,
proxy_logging_obj=logging,
valid_token=token,
request=Request(
{"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": [], "query_string": b""}
),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("fresh", [False, True])
async def test_common_checks_reread_organization_budget_for_delegated_identity(
database: MagicMock, fresh: bool
) -> None:
from litellm.proxy import proxy_server
cached: Final = LiteLLM_OrganizationTable(
organization_id="org",
budget_id="budget",
created_by="admin",
updated_by="admin",
litellm_budget_table=LiteLLM_BudgetTable(max_budget=100.0),
)
await proxy_server.user_api_key_cache.async_set_cache(key="org_id:org:with_budget", value=cached)
database.db.litellm_organizationtable.find_unique = AsyncMock(
return_value=cached.model_copy(update={"litellm_budget_table": LiteLLM_BudgetTable(max_budget=0.0)})
)
token: Final = UserAPIKeyAuth(token="token", user_id="admin", org_id="org")
token.requires_fresh_policy = fresh
if not fresh:
assert await _common_checks(token) is True
return
with pytest.raises(litellm.BudgetExceededError):
await _common_checks(token)
@pytest.mark.asyncio
@pytest.mark.parametrize("fresh", [False, True])
async def test_common_checks_reread_team_access_groups_for_delegated_identity(database: MagicMock, fresh: bool) -> None:
from litellm.proxy import proxy_server
cached: Final = LiteLLM_AccessGroupTable(
access_group_id="group", access_group_name="group", access_model_names=["gpt-4"]
)
await proxy_server.user_api_key_cache.async_set_cache(key="access_group_id:group", value=cached)
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(
return_value=cached.model_copy(update={"access_model_names": []})
)
database.writer_db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
team: Final = LiteLLM_TeamTable(team_id="team", models=["other-model"], access_group_ids=["group"])
token: Final = UserAPIKeyAuth(token="token", user_id="admin", team_id="team")
token.requires_fresh_policy = fresh
if not fresh:
assert await _common_checks(token, team) is True
return
with pytest.raises(ProxyException) as error:
await _common_checks(token, team)
assert error.value.type == ProxyErrorTypes.team_model_access_denied

View file

@ -5,6 +5,7 @@ from unittest.mock import MagicMock, patch
import pytest
from fastapi import HTTPException, Request
from starlette.routing import Match
from litellm.proxy._types import (
LiteLLM_OrganizationMembershipTable,
@ -16,6 +17,71 @@ from litellm.proxy._types import (
from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin
from litellm.proxy.auth.route_checks import RouteChecks
@pytest.mark.parametrize(
"route,allowed",
[
("/key/generate", True),
("/user/info", True),
("/team/update", True),
("/team/member_add", True),
("/budget/new", True),
("/budget/update", True),
("/budget/info", True),
("/budget/delete", True),
("/organization/new", True),
("/global/spend/logs", True),
("/v1/chat/completions", True),
("/openai/deployments/demo/chat/completions", True),
("/openai/deployments/org/demo/chat/completions", True),
("/laya/v1/systemone", False),
("/deepgram/listen", False),
("/mcp", False),
("/mcp-rest/tools/call", False),
("/v1/mcp/server", False),
("/jwt/key/mapping/new", False),
("/user/auth", False),
("/user/password/change", False),
("/session/logout", False),
("/team/team-id/callback", False),
("/config/yaml", False),
("/global/spend/reset", False),
("/custom/backend", False),
("/key/generate\n", False),
],
)
def test_delegated_admin_scope_limits_api_access(route: str, allowed: bool) -> None:
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import router
from litellm.proxy.proxy_server import app
scope: Final = {"type": "websocket" if route == "/deepgram/listen" else "http", "method": "POST", "path": route}
matched: Final = next((r for r in (*app.routes, *router.routes) if r.matches(scope)[0] is Match.FULL), None)
request: Final = Request({**scope, **(matched.matches(scope)[1] if matched else {}), "type": "http"})
inference_paths: Final = [*LiteLLMRoutes.openai_routes.value, "/laya/v1/systemone", "/deepgram/listen"]
with patch.object(LiteLLMRoutes.openai_routes, "_value_", inference_paths):
assert RouteChecks.is_delegated_admin_route(route, request) is allowed
def test_marked_provider_pass_through_without_endpoint_param_keeps_model_alias_dispatch() -> None:
from litellm.proxy.auth.auth_utils import (
request_dispatched_to_marked_provider_pass_through,
request_dispatched_to_provider_pass_through,
)
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import laya_proxy_route
request: Final = Request(
{
"type": "http",
"method": "POST",
"path": "/laya/v1/systemone",
"endpoint": laya_proxy_route,
"path_params": {},
}
)
assert request_dispatched_to_marked_provider_pass_through(request)
assert not request_dispatched_to_provider_pass_through(request)
DAILY_ACTIVITY_ROUTE_PAIRS: Final[tuple[tuple[str, str], ...]] = (
("/user/daily/activity", "/user/daily/activity/aggregated"),
("/user/daily/activity", "/user/daily/activity/aggregated/keys"),

View file

@ -373,6 +373,7 @@ async def test_team_member_budget_check_blocks_regenerated_key_after_old_key_exh
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
check_db_only=False,
)
assert "Budget has been exceeded" in str(exc_info.value)
assert "test-user-1" in str(exc_info.value)

View file

@ -5,7 +5,7 @@
import litellm.proxy
import litellm.proxy.proxy_server
from typing import Dict, List, Optional
from typing import Dict, Final, List, Optional
from unittest.mock import MagicMock, patch, AsyncMock
import pytest
@ -23,6 +23,32 @@ from fastapi import WebSocket, HTTPException, status
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles
@pytest.mark.asyncio
@pytest.mark.parametrize("credential", ("llm_session_custom-token", "llm_srefresh_custom-token"))
async def test_custom_auth_owns_credentials_with_delegated_prefixes(
monkeypatch: pytest.MonkeyPatch, credential: str
) -> None:
from typing import Final
from fastapi import Request as HttpRequest
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
async def custom_auth(request: HttpRequest, api_key: str) -> UserAPIKeyAuth:
assert api_key == credential
return UserAPIKeyAuth(user_id="custom-user")
monkeypatch.setattr(litellm.proxy.proxy_server, "user_custom_auth", custom_auth)
monkeypatch.setattr(litellm.proxy.proxy_server, "general_settings", {})
monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", False, raising=False)
request: Final = HttpRequest({"type": "http", "method": "GET", "path": "/key/info", "headers": []})
identity: Final = await _user_api_key_auth_builder(
request, f"Bearer {credential}", "", None, None, None, {}
)
assert identity.user_id == "custom-user"
assert identity.authenticated_by_custom_auth is True
class Request:
def __init__(self, client_ip: Optional[str] = None, headers: Optional[dict] = None):
self.client = MagicMock()
@ -883,31 +909,41 @@ async def test_user_api_key_auth_websocket():
@pytest.mark.asyncio
async def test_user_api_key_auth_websocket_carries_asgi_path():
async def test_user_api_key_auth_websocket_carries_asgi_path() -> None:
"""
The synthetic Request must carry the ASGI scope's ``path`` so
``get_request_route`` returns the real WebSocket path, not a value
reconstructed from the (Host-poisonable) ``websocket.url``.
"""
from litellm.proxy.auth.auth_utils import request_dispatched_to_provider_pass_through
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import openai_websocket_proxy_route, router
provider_route: Final = next(route for route in router.routes if route.endpoint is openai_websocket_proxy_route)
mock_websocket = MagicMock(spec=WebSocket)
mock_websocket.query_params = {"model": "some_model"}
mock_websocket.headers = {"authorization": "Bearer some_api_key"}
mock_websocket.scope = {
"type": "websocket",
"path": "/v1/realtime",
"path": "/openai/v1/responses",
"root_path": "",
"headers": [(b"authorization", b"Bearer some_api_key")],
"endpoint": provider_route.endpoint,
"path_params": {"endpoint": "v1/responses"},
"route": provider_route,
}
mock_websocket.url = URL(url="/v1/realtime")
mock_websocket.url = URL(url="/openai/v1/responses")
with patch("litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True) as mock_user_api_key_auth:
await user_api_key_auth_websocket(mock_websocket)
request_arg = mock_user_api_key_auth.call_args.kwargs["request"]
assert request_arg.scope.get("path") == "/v1/realtime"
assert request_arg.scope.get("path") == "/openai/v1/responses"
assert request_arg.scope.get("root_path") == ""
assert request_arg.scope["endpoint"] is provider_route.endpoint
assert request_arg.scope["path_params"] == {"endpoint": "v1/responses"}
assert request_arg.scope["route"] is provider_route
assert request_dispatched_to_provider_pass_through(request_arg)
@pytest.mark.parametrize("enforce_rbac", [True, False])

View file

@ -178,6 +178,23 @@ export interface paths {
patch?: never;
trace?: never;
};
"/.well-known/oauth-authorization-server/oauth/api": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
/** Oauth Authorization Server Api */
get: operations["oauth_authorization_server_api__well_known_oauth_authorization_server_oauth_api_get"];
put?: never;
post?: never;
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/.well-known/oauth-authorization-server/{mcp_server_name}": {
parameters: {
query?: never;
@ -14307,7 +14324,7 @@ export interface paths {
* Revoke Endpoint
* @description RFC 7009 revocation for the gateway's refresh tokens (``lite logout``): 200 for a known
* client whatever the token's state, 503 when the shared single-use record cannot be written;
* access tokens expire on their own.
* native/MCP access tokens expire normally; delegated access is revoked immediately.
*/
post: operations["revoke_endpoint_revoke_post"];
delete?: never;
@ -51149,6 +51166,28 @@ export interface operations {
};
};
};
oauth_authorization_server_api__well_known_oauth_authorization_server_oauth_api_get: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": {
[key: string]: string | string[];
};
};
};
};
};
oauth_authorization_server_mcp__well_known_oauth_authorization_server__mcp_server_name__get: {
parameters: {
query?: never;