mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(mcp): bind per-user OAuth credentials to the authenticated LiteLLM caller
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
cbdaa3b153
commit
88b3cfd789
7 changed files with 614 additions and 1 deletions
|
|
@ -54,6 +54,9 @@ from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
|
|||
relative_request_url,
|
||||
revoke_refresh_token,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
|
||||
enforce_oauth_identity_binding,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
TOKEN_NO_CACHE_HEADERS,
|
||||
build_upstream_oauth2_token_request,
|
||||
|
|
@ -1133,11 +1136,26 @@ async def exchange_token_with_server(
|
|||
server_id=resolved_server.server_id,
|
||||
)
|
||||
|
||||
# Bind the exchanged token to the LiteLLM caller BEFORE it is returned, stored, or cached, so a
|
||||
# token minted for a different upstream principal never becomes usable under the caller's user_id.
|
||||
resolved_user_id: Final = (
|
||||
await _extract_user_id_from_request(request)
|
||||
if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None
|
||||
else None
|
||||
)
|
||||
if isinstance(token_response, dict):
|
||||
await enforce_oauth_identity_binding(
|
||||
server=resolved_server,
|
||||
token_response=token_response,
|
||||
litellm_user_id=resolved_user_id,
|
||||
grant_type=grant_type,
|
||||
)
|
||||
|
||||
# Store server-side when the server is configured for per-user OAuth and
|
||||
# the calling client has provided a valid LiteLLM identity.
|
||||
# Errors are non-fatal: the token is still returned to the client.
|
||||
if resolved_server.needs_user_oauth_token:
|
||||
user_id: Final = await _extract_user_id_from_request(request)
|
||||
user_id: Final = resolved_user_id
|
||||
if user_id:
|
||||
try:
|
||||
await _store_per_user_token_server_side(
|
||||
|
|
|
|||
|
|
@ -2174,6 +2174,8 @@ class MCPServerManager:
|
|||
allow_elicitation=bool(server_config.get("allow_elicitation", False)),
|
||||
timeout=server_config.get("timeout", None),
|
||||
max_concurrent_requests=server_config.get("max_concurrent_requests", None),
|
||||
token_validation=server_config.get("token_validation", None),
|
||||
oauth_identity_binding=server_config.get("oauth_identity_binding", None),
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
|
||||
|
|
|
|||
249
litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py
Normal file
249
litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py
Normal file
|
|
@ -0,0 +1,249 @@
|
|||
"""Per-user OAuth identity binding: verify the upstream OIDC principal matches the LiteLLM caller.
|
||||
|
||||
Closes the confused-deputy gap where a browser authenticated upstream as one principal produces a
|
||||
token that the relay stores under a different, LiteLLM-authenticated principal: before the token
|
||||
endpoint returns, stores, or caches an exchanged token for an identity-bound server, the id_token
|
||||
is validated (signature via the pinned issuer's JWKS, issuer, audience, expiry) and its principal
|
||||
claim is compared to the caller's trusted LiteLLM identity. Mismatches fail closed in enforce mode
|
||||
and are logged in audit mode.
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Literal
|
||||
|
||||
import jwt
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer
|
||||
|
||||
_ALLOWED_ID_TOKEN_ALGORITHMS: Final = (
|
||||
"RS256",
|
||||
"RS384",
|
||||
"RS512",
|
||||
"ES256",
|
||||
"ES384",
|
||||
"ES512",
|
||||
"PS256",
|
||||
"PS384",
|
||||
"PS512",
|
||||
)
|
||||
_JWKS_CACHE_TTL_SECONDS: Final = 3600
|
||||
_jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS)
|
||||
|
||||
JwksFetcher = Callable[[MCPOAuthIdentityBinding], Awaitable[list[Mapping[str, object]]]]
|
||||
CallerPrincipalLoader = Callable[[str, MCPOAuthIdentityBinding], Awaitable[str | None]]
|
||||
|
||||
_RejectionCode = Literal["oauth_principal_mismatch", "oauth_identity_binding_failed"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BindingRejection:
|
||||
code: _RejectionCode
|
||||
description: str
|
||||
|
||||
|
||||
async def _fetch_issuer_jwks(binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]:
|
||||
jwks_url: Final[str] = binding.jwks_url or await _discover_jwks_url(binding.issuer)
|
||||
cached: Final = await _jwks_cache.async_get_cache(jwks_url)
|
||||
if isinstance(cached, list):
|
||||
return cached
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response: Final = await client.get(jwks_url)
|
||||
response.raise_for_status()
|
||||
document: Final = response.json()
|
||||
keys: Final = document.get("keys") if isinstance(document, dict) else None
|
||||
if not isinstance(keys, list):
|
||||
raise TypeError(f"JWKS document at {jwks_url} has no 'keys' array")
|
||||
await _jwks_cache.async_set_cache(jwks_url, keys, ttl=_JWKS_CACHE_TTL_SECONDS)
|
||||
return keys
|
||||
|
||||
|
||||
async def _discover_jwks_url(issuer: str) -> str:
|
||||
discovery_url: Final = f"{issuer.rstrip('/')}/.well-known/openid-configuration"
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response: Final = await client.get(discovery_url)
|
||||
response.raise_for_status()
|
||||
metadata: Final = response.json()
|
||||
jwks_uri: Final = metadata.get("jwks_uri") if isinstance(metadata, dict) else None
|
||||
if not isinstance(jwks_uri, str) or not jwks_uri:
|
||||
raise ValueError(f"OIDC discovery at {discovery_url} returned no jwks_uri")
|
||||
return jwks_uri
|
||||
|
||||
|
||||
def _select_signing_key(id_token: str, keys: list[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection":
|
||||
header: Final = jwt.get_unverified_header(id_token)
|
||||
kid: Final = header.get("kid")
|
||||
for key in keys:
|
||||
if kid is None or key.get("kid") == kid:
|
||||
return jwt.PyJWK(dict(key))
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description=f"id_token signing key (kid={kid!r}) not found in the issuer's JWKS",
|
||||
)
|
||||
|
||||
|
||||
def _decode_id_token(
|
||||
id_token: str,
|
||||
binding: MCPOAuthIdentityBinding,
|
||||
signing_key: "jwt.PyJWK",
|
||||
) -> "Mapping[str, object] | _BindingRejection":
|
||||
try:
|
||||
return jwt.decode(
|
||||
id_token,
|
||||
signing_key.key,
|
||||
algorithms=list(_ALLOWED_ID_TOKEN_ALGORITHMS),
|
||||
issuer=binding.issuer,
|
||||
audience=binding.audiences if binding.audiences else None,
|
||||
options={
|
||||
"require": ["iss", "exp"],
|
||||
"verify_aud": bool(binding.audiences),
|
||||
},
|
||||
)
|
||||
except jwt.InvalidTokenError as exc:
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description=f"id_token validation failed: {exc}",
|
||||
)
|
||||
|
||||
|
||||
def _upstream_principal(
|
||||
claims: Mapping[str, object],
|
||||
binding: MCPOAuthIdentityBinding,
|
||||
) -> "str | _BindingRejection":
|
||||
principal: Final = claims.get(binding.principal_claim)
|
||||
if not isinstance(principal, str) or not principal:
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description=f"id_token has no usable '{binding.principal_claim}' claim",
|
||||
)
|
||||
if binding.principal_claim == "email" and binding.require_email_verified and claims.get("email_verified") is not True:
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description="id_token email is not verified (email_verified is not true)",
|
||||
)
|
||||
return principal
|
||||
|
||||
|
||||
async def _load_caller_principal(litellm_user_id: str, binding: MCPOAuthIdentityBinding) -> str | None:
|
||||
if binding.caller_field == "user_id":
|
||||
return litellm_user_id
|
||||
from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
load_active_user_by_id,
|
||||
)
|
||||
|
||||
loaded: Final = await load_active_user_by_id(litellm_user_id)
|
||||
if isinstance(loaded, str):
|
||||
return None
|
||||
return loaded.user_email
|
||||
|
||||
|
||||
def _principals_match(upstream: str, caller: str, binding: MCPOAuthIdentityBinding) -> bool:
|
||||
if binding.principal_claim == "email" or binding.caller_field == "user_email":
|
||||
return upstream.strip().casefold() == caller.strip().casefold()
|
||||
return upstream == caller
|
||||
|
||||
|
||||
async def _evaluate_binding(
|
||||
binding: MCPOAuthIdentityBinding,
|
||||
token_response: Mapping[str, object],
|
||||
litellm_user_id: str | None,
|
||||
grant_type: str,
|
||||
jwks_fetcher: JwksFetcher,
|
||||
caller_principal_loader: CallerPrincipalLoader,
|
||||
) -> _BindingRejection | None:
|
||||
id_token: Final = token_response.get("id_token")
|
||||
if not isinstance(id_token, str) or not id_token:
|
||||
if grant_type == "refresh_token":
|
||||
return None
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description="the upstream token response carries no id_token to bind the credential to a principal",
|
||||
)
|
||||
if not litellm_user_id:
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description="the request carries no resolvable LiteLLM user identity to bind the credential to",
|
||||
)
|
||||
try:
|
||||
keys: Final = await jwks_fetcher(binding)
|
||||
except Exception as exc: # noqa: BLE001 # a JWKS fetch failure must fail closed, not surface as a 500
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description=f"could not fetch the issuer's JWKS: {exc}",
|
||||
)
|
||||
signing_key: Final = _select_signing_key(id_token, keys)
|
||||
if isinstance(signing_key, _BindingRejection):
|
||||
return signing_key
|
||||
claims: Final = _decode_id_token(id_token, binding, signing_key)
|
||||
if isinstance(claims, _BindingRejection):
|
||||
return claims
|
||||
upstream: Final = _upstream_principal(claims, binding)
|
||||
if isinstance(upstream, _BindingRejection):
|
||||
return upstream
|
||||
caller: Final = await caller_principal_loader(litellm_user_id, binding)
|
||||
if not caller:
|
||||
return _BindingRejection(
|
||||
code="oauth_identity_binding_failed",
|
||||
description=f"the LiteLLM user has no '{binding.caller_field}' to compare the upstream principal against",
|
||||
)
|
||||
if not _principals_match(upstream, caller, binding):
|
||||
return _BindingRejection(
|
||||
code="oauth_principal_mismatch",
|
||||
description="The browser account does not match the selected credential owner.",
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
async def enforce_oauth_identity_binding(
|
||||
server: MCPServer,
|
||||
token_response: Mapping[str, object],
|
||||
litellm_user_id: str | None,
|
||||
grant_type: str,
|
||||
jwks_fetcher: JwksFetcher = _fetch_issuer_jwks,
|
||||
caller_principal_loader: CallerPrincipalLoader = _load_caller_principal,
|
||||
) -> None:
|
||||
"""Validate the exchanged token's upstream principal against the LiteLLM caller.
|
||||
|
||||
No-op when the server has no binding or it is disabled. In enforce mode a failure raises 403
|
||||
before the caller returns, stores, or caches the token; in audit mode failures are logged only.
|
||||
A refresh_token grant without an id_token is allowed in both modes: the stored credential keeps
|
||||
the binding established at the original authorization_code exchange.
|
||||
"""
|
||||
binding: Final = server.oauth_identity_binding
|
||||
if binding is None or binding.mode == "disabled":
|
||||
return
|
||||
rejection: Final = await _evaluate_binding(
|
||||
binding=binding,
|
||||
token_response=token_response,
|
||||
litellm_user_id=litellm_user_id,
|
||||
grant_type=grant_type,
|
||||
jwks_fetcher=jwks_fetcher,
|
||||
caller_principal_loader=caller_principal_loader,
|
||||
)
|
||||
if rejection is None:
|
||||
return
|
||||
if binding.mode == "audit":
|
||||
verbose_logger.warning(
|
||||
"oauth_identity_binding audit: server=%s user=%s grant=%s rejected=%s (%s)",
|
||||
server.server_id,
|
||||
litellm_user_id,
|
||||
grant_type,
|
||||
rejection.code,
|
||||
rejection.description,
|
||||
)
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": rejection.code,
|
||||
"error_description": rejection.description,
|
||||
"server_id": server.server_id,
|
||||
"credential_owner": "caller",
|
||||
"credential_stored": False,
|
||||
},
|
||||
)
|
||||
|
|
@ -2135,6 +2135,28 @@ if MCP_AVAILABLE:
|
|||
"""Persist the OAuth2 access token obtained by the calling user."""
|
||||
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
|
||||
await _authorize_and_fetch_mcp_server(prisma_client, user_api_key_dict, server_id)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
|
||||
global_mcp_server_manager as _manager,
|
||||
)
|
||||
|
||||
# This endpoint accepts an opaque token with no upstream identity validation, so it must be
|
||||
# closed for identity-bound servers or it becomes a bypass of the token-relay binding check.
|
||||
registry_server: Final = _manager.get_mcp_server_by_id(server_id)
|
||||
binding: Final = registry_server.oauth_identity_binding if registry_server else None
|
||||
if binding is not None and binding.mode == "enforce":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={
|
||||
"error": "oauth_identity_binding_enforced",
|
||||
"error_description": (
|
||||
"Direct credential storage is disabled for this server: its OAuth identity "
|
||||
"binding is enforced and this endpoint cannot validate the token's principal. "
|
||||
"Complete the OAuth flow through the gateway instead."
|
||||
),
|
||||
"server_id": server_id,
|
||||
"credential_stored": False,
|
||||
},
|
||||
)
|
||||
user_id: Final = user_api_key_dict.user_id or ""
|
||||
if not user_id:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -38,6 +38,26 @@ class MCPOAuthMetadata(BaseModel):
|
|||
usable in memory but must never be persisted as configuration."""
|
||||
|
||||
|
||||
class MCPOAuthIdentityBinding(BaseModel):
|
||||
"""Per-server policy binding stored per-user OAuth credentials to the authenticated LiteLLM caller.
|
||||
|
||||
When enabled for an interactive oauth2 server, the token relay validates the upstream OIDC
|
||||
``id_token`` (signature via the pinned issuer's JWKS, issuer, audience, expiry) and compares its
|
||||
principal claim to the LiteLLM caller's trusted identity before the token is returned, stored,
|
||||
or cached. ``audit`` logs mismatches without changing behavior; ``enforce`` fails closed with
|
||||
403 ``oauth_principal_mismatch`` and disables the direct ``oauth-user-credential`` POST, which
|
||||
would otherwise bypass validation with an arbitrary opaque token.
|
||||
"""
|
||||
|
||||
mode: Literal["disabled", "audit", "enforce"] = "disabled"
|
||||
issuer: str
|
||||
jwks_url: str | None = None
|
||||
audiences: list[str] = []
|
||||
principal_claim: str = "email"
|
||||
caller_field: Literal["user_email", "user_id"] = "user_email"
|
||||
require_email_verified: bool = True
|
||||
|
||||
|
||||
class MCPServer(BaseModel):
|
||||
server_id: str
|
||||
name: str
|
||||
|
|
@ -172,6 +192,7 @@ class MCPServer(BaseModel):
|
|||
# response (supports dot-notation for nested fields, e.g. "team.enterprise_id").
|
||||
# Tokens that fail validation are rejected before storage.
|
||||
token_validation: dict[str, Any] | None = None
|
||||
oauth_identity_binding: MCPOAuthIdentityBinding | None = None
|
||||
# Optional TTL override (seconds) for the Redis per-user token cache, capped
|
||||
# at the token's expires_in minus the expiry buffer so a cached entry never
|
||||
# outlives the token. Defaults to the token's expires_in minus the expiry
|
||||
|
|
|
|||
|
|
@ -0,0 +1,240 @@
|
|||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import jwt
|
||||
import pytest
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
|
||||
enforce_oauth_identity_binding,
|
||||
)
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer
|
||||
|
||||
ISSUER: Final = "https://idp.example.com"
|
||||
AUDIENCE: Final = "litellm-client"
|
||||
KID: Final = "test-key"
|
||||
|
||||
_PRIVATE_KEY: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
_PRIVATE_PEM: Final = _PRIVATE_KEY.private_bytes(
|
||||
encoding=serialization.Encoding.PEM,
|
||||
format=serialization.PrivateFormat.PKCS8,
|
||||
encryption_algorithm=serialization.NoEncryption(),
|
||||
)
|
||||
_PUBLIC_JWK: Final = {
|
||||
**jwt.algorithms.RSAAlgorithm.to_jwk(_PRIVATE_KEY.public_key(), as_dict=True),
|
||||
"kid": KID,
|
||||
"alg": "RS256",
|
||||
"use": "sig",
|
||||
}
|
||||
|
||||
|
||||
def _sign_id_token(claims: Mapping[str, object]) -> str:
|
||||
payload: Final = {
|
||||
"iss": ISSUER,
|
||||
"aud": AUDIENCE,
|
||||
"exp": int(time.time()) + 300,
|
||||
"iat": int(time.time()),
|
||||
**claims,
|
||||
}
|
||||
return jwt.encode(payload, _PRIVATE_PEM, algorithm="RS256", headers={"kid": KID})
|
||||
|
||||
|
||||
async def _jwks_fetcher(_binding: MCPOAuthIdentityBinding) -> list[Mapping[str, object]]:
|
||||
return [_PUBLIC_JWK]
|
||||
|
||||
|
||||
def _caller_loader(email: str | None):
|
||||
async def load(_user_id: str, _binding: MCPOAuthIdentityBinding) -> str | None:
|
||||
return email
|
||||
|
||||
return load
|
||||
|
||||
|
||||
def _server(mode: str = "enforce", **binding_overrides: object) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="srv-1",
|
||||
name="srv-1",
|
||||
url="https://mcp.example.com",
|
||||
transport=MCPTransport.http,
|
||||
oauth_identity_binding=MCPOAuthIdentityBinding(
|
||||
mode=mode,
|
||||
issuer=ISSUER,
|
||||
audiences=[AUDIENCE],
|
||||
**binding_overrides,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_matching_principal_passes():
|
||||
token: Final = _sign_id_token({"email": "Alice@Example.com", "email_verified": True})
|
||||
result: Final = await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at", "id_token": token},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mismatched_principal_rejected():
|
||||
token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True})
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at", "id_token": token},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
assert exc_info.value.detail["error"] == "oauth_principal_mismatch"
|
||||
assert exc_info.value.detail["credential_stored"] is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_id_token_rejected_on_authorization_code():
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at"},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
assert exc_info.value.detail["error"] == "oauth_identity_binding_failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_without_id_token_allowed():
|
||||
result: Final = await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at"},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="refresh_token",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_with_mismatched_id_token_rejected():
|
||||
token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True})
|
||||
with pytest.raises(HTTPException):
|
||||
await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at", "id_token": token},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="refresh_token",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audit_mode_logs_but_does_not_reject():
|
||||
token: Final = _sign_id_token({"email": "mallory@example.com", "email_verified": True})
|
||||
result: Final = await enforce_oauth_identity_binding(
|
||||
server=_server(mode="audit"),
|
||||
token_response={"access_token": "at", "id_token": token},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unverified_email_rejected():
|
||||
token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": False})
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at", "id_token": token},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert exc_info.value.detail["error"] == "oauth_identity_binding_failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrong_issuer_rejected():
|
||||
payload: Final = {
|
||||
"iss": "https://evil.example.com",
|
||||
"aud": AUDIENCE,
|
||||
"exp": int(time.time()) + 300,
|
||||
"email": "alice@example.com",
|
||||
"email_verified": True,
|
||||
}
|
||||
token: Final = jwt.encode(payload, _PRIVATE_PEM, algorithm="RS256", headers={"kid": KID})
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at", "id_token": token},
|
||||
litellm_user_id="user-a",
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert exc_info.value.detail["error"] == "oauth_identity_binding_failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_litellm_identity_rejected():
|
||||
token: Final = _sign_id_token({"email": "alice@example.com", "email_verified": True})
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await enforce_oauth_identity_binding(
|
||||
server=_server(),
|
||||
token_response={"access_token": "at", "id_token": token},
|
||||
litellm_user_id=None,
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader("alice@example.com"),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disabled_binding_is_noop():
|
||||
result: Final = await enforce_oauth_identity_binding(
|
||||
server=_server(mode="disabled"),
|
||||
token_response={"access_token": "at"},
|
||||
litellm_user_id=None,
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader(None),
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_binding_is_noop():
|
||||
server: Final = MCPServer(
|
||||
server_id="srv-2",
|
||||
name="srv-2",
|
||||
url="https://mcp.example.com",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
result: Final = await enforce_oauth_identity_binding(
|
||||
server=server,
|
||||
token_response={"access_token": "at"},
|
||||
litellm_user_id=None,
|
||||
grant_type="authorization_code",
|
||||
jwks_fetcher=_jwks_fetcher,
|
||||
caller_principal_loader=_caller_loader(None),
|
||||
)
|
||||
assert result is None
|
||||
|
|
@ -4638,6 +4638,67 @@ async def test_store_mcp_oauth_user_credential_returns_status():
|
|||
assert result.expires_at == "2099-01-01T00:00:00+00:00"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_mcp_oauth_user_credential_blocked_when_identity_binding_enforced():
|
||||
"""The direct opaque-token POST must be closed for enforce-mode identity-bound servers,
|
||||
otherwise it bypasses the token-relay principal check."""
|
||||
from litellm.proxy._types import MCPOAuthUserCredentialRequest
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer
|
||||
|
||||
if not mgmt_endpoints.MCP_AVAILABLE:
|
||||
pytest.skip("MCP module not installed")
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
store_mcp_oauth_user_credential,
|
||||
)
|
||||
|
||||
server_id = "srv-binding-1"
|
||||
bound_server = MCPServer(
|
||||
server_id=server_id,
|
||||
name=server_id,
|
||||
url="https://mcp.example.com",
|
||||
transport=MCPTransport.http,
|
||||
oauth_identity_binding=MCPOAuthIdentityBinding(mode="enforce", issuer="https://idp.example.com"),
|
||||
)
|
||||
store_mock = AsyncMock(return_value=None)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: mirrors the existing store-credential tests in this file
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=_make_prisma_client(),
|
||||
),
|
||||
patch( # test-quality-ok: mirrors the existing store-credential tests in this file
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)),
|
||||
),
|
||||
patch( # test-quality-ok: mirrors the existing store-credential tests in this file
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=True,
|
||||
),
|
||||
patch.object( # test-quality-ok: registry is a module-level singleton; injecting it would change the endpoint signature
|
||||
manager_module.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
return_value=bound_server,
|
||||
),
|
||||
patch( # test-quality-ok: asserting the DB write is never reached is the point of the test
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.store_user_oauth_credential",
|
||||
new=store_mock,
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await store_mcp_oauth_user_credential(
|
||||
server_id=server_id,
|
||||
payload=MCPOAuthUserCredentialRequest(access_token="opaque-tok", expires_in=3600),
|
||||
user_api_key_dict=_make_user_auth("user-123"),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert exc_info.value.detail["error"] == "oauth_identity_binding_enforced"
|
||||
store_mock.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_mcp_oauth_user_credential_only_deletes_oauth():
|
||||
"""delete_mcp_oauth_user_credential only deletes OAuth2 credentials, not BYOK."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue