fix(mcp): persist OAuth credentials for validated JWT users

This commit is contained in:
Joshua Valluru 2026-09-15 12:29:03 -07:00
parent d3929287fe
commit 92e182b898
3 changed files with 335 additions and 3 deletions

View file

@ -306,12 +306,68 @@ async def _extract_user_id_from_request(request: Request) -> str | None:
(including a transient DB outage) collapses to ``None`` here and the caller simply skips the store;
the bridge mint, which must status those outcomes differently, consumes
:func:`_resolve_active_litellm_key` directly."""
from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle
token: Final = _litellm_key_from_request(request)
if token is not None and JWTHandler.is_jwt(token):
return await _extract_jwt_user_id(token)
resolved: Final = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
return None
return _active_key_user_id(resolved.key)
async def _extract_jwt_user_id(token: str) -> str | None:
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.handle_jwt import JWTAuthManager # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle
_resolve_jwt_to_virtual_key, # pyright: ignore[reportPrivateUsage] # reuse admission mapping policy without provisioning a new key
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # proxy globals initialized at startup
general_settings,
jwt_handler,
premium_user,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if general_settings.get("enable_jwt_auth") is not True or premium_user is not True:
return None
try:
claims: Final = await jwt_handler.auth_jwt(token=token)
validate: Final = jwt_handler.litellm_jwtauth.custom_validate
if validate is not None and not validate(claims):
return None
if jwt_handler.litellm_jwtauth.is_virtual_key_mapping_configured():
mapped: Final = await _resolve_jwt_to_virtual_key(
jwt_claims=claims,
jwt_handler=jwt_handler,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
)
if isinstance(mapped, UserAPIKeyAuth):
return None if await _key_owner_scim_deactivated(mapped) else _active_key_user_id(mapped)
if mapped is not None:
return None
user_id, _, valid_email = await JWTAuthManager.get_user_info(jwt_handler, claims)
object_id: Final = jwt_handler.get_object_id(token=claims, default_value=None)
owner_id: Final = (
object_id
if jwt_handler.get_rbac_role(token=claims) == LitellmUserRoles.INTERNAL_USER and object_id
else user_id
)
if not owner_id or valid_email is False:
return None
owner: Final = await load_active_user_by_id(owner_id)
return None if isinstance(owner, str) else owner.user_id
except Exception as exc: # noqa: BLE001 # public OAuth exchange stays available; unvalidated identities never write credentials
verbose_logger.debug("OAuth JWT identity could not be validated (%s)", type(exc).__name__)
return None
_UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"]
"""Why an upstream token response cannot back a bridge envelope:
- ``no_access_token``: the response carries no usable ``access_token``

View file

@ -1236,8 +1236,9 @@ async def exchange_token_with_server(
"exchange_token_with_server: could not resolve a LiteLLM user_id for the request, "
"so the per-user token for server=%s was NOT stored. The authorization_code egress "
"requires the stored token, so the client will be challenged with 401 on reconnect. "
"Ensure the request carries a valid LiteLLM key (x-litellm-api-key or Authorization), "
"or store it via POST /mcp/server/{id}/oauth-user-credential.",
"Ensure the request carries a valid LiteLLM key or enabled JWT identity "
"(x-litellm-api-key or Authorization), "
"or store it via POST /v1/mcp/server/{id}/oauth-user-credential.",
resolved_server.server_id,
)

View file

@ -5,7 +5,7 @@ import json
import time
from base64 import urlsafe_b64encode
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -15,6 +15,9 @@ from litellm.types.mcp import MCPAuth
if TYPE_CHECKING:
import httpx
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -11374,3 +11377,275 @@ with TestClient(app) as client:
assert responses[path]["status"] == 200, responses[path]
assert responses[path]["body"]["issuer"] == f"http://testserver/gateway/{path}"
assert responses["example/mcp"]["body"]["token_endpoint"] == "http://testserver/gateway/example/token"
@pytest.fixture
def jwt_oauth_identity(monkeypatch: pytest.MonkeyPatch) -> tuple["JWTHandler", "RSAPrivateKey"]:
import jwt
from cryptography.hazmat.primitives.asymmetric import rsa
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_JWTAuth
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
signing_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
cache: Final = UserApiKeyCache()
cache.set_cache(
"litellm_jwt_auth_keys_https://idp.example.test/jwks",
[json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(signing_key.public_key()))],
)
cache.set_cache("jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", user_email="owner@example.test"))
handler: Final = JWTHandler()
handler.update_environment(
prisma_client=None,
user_api_key_cache=cache,
litellm_jwtauth=LiteLLM_JWTAuth(user_id_jwt_field="identity.user_id"),
)
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", "https://idp.example.test/jwks")
monkeypatch.setenv("JWT_ISSUER", "https://idp.example.test")
monkeypatch.setenv("JWT_AUDIENCE", "litellm-proxy")
monkeypatch.setattr(proxy_server, "jwt_handler", handler)
monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": True})
monkeypatch.setattr(proxy_server, "premium_user", True)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
return handler, signing_key
def _oauth_identity_jwt(
signing_key: "RSAPrivateKey",
*,
expires_in: int = 300,
audience: str = "litellm-proxy",
issuer: str = "https://idp.example.test",
owner: str | None = "jwt-owner",
) -> str:
import jwt
return jwt.encode(
{
"sub": "not-the-configured-user-id",
"identity": {"user_id": owner},
"iss": issuer,
"aud": audience,
"exp": int(time.time()) + expires_in,
},
signing_key,
algorithm="RS256",
)
@pytest.mark.asyncio
@pytest.mark.parametrize("header", ["Authorization", "x-litellm-api-key"])
async def test_oauth_exchange_stores_token_for_validated_jwt_user(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
header: str,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import httpx
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_, signing_key = jwt_oauth_identity
bearer: Final = _oauth_identity_jwt(signing_key)
request: Final = _token_request({header: f"Bearer {bearer}"})
server: Final = MCPServer(
server_id="jwt-oauth-server",
name="jwt-oauth-server",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
authorization_url="https://upstream.example.test/authorize",
token_url="https://upstream.example.test/token",
client_id="registered-client",
)
import litellm
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.types.llms.custom_http import httpxSpecialProvider
def upstream_response(outbound: httpx.Request) -> httpx.Response:
assert outbound.url == server.token_url
assert bearer not in str(outbound.headers)
assert bearer.encode() not in outbound.content
return httpx.Response(200, json={"access_token": "upstream-token", "token_type": "Bearer"})
database: Final = MagicMock()
table: Final = database.db.litellm_mcpusercredentials
table.find_unique = AsyncMock(return_value=None)
table.upsert = AsyncMock()
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setenv("LITELLM_SALT_KEY", "oauth-jwt-test-encryption-key")
clients: Final = LLMClientCache()
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", clients)
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream_response)) as transport:
upstream: Final = AsyncHTTPHandler()
await upstream.client.aclose()
upstream.client = transport
clients.set_cache("async_httpx_client" + httpxSpecialProvider.Oauth2Check, upstream)
response: Final = await discoverable_endpoints.exchange_token_with_server(
request=request,
mcp_server=server,
grant_type="authorization_code",
code="upstream-code",
redirect_uri="http://localhost/callback",
client_id="registered-client",
client_secret=None,
code_verifier=None,
)
assert response.status_code == 200
table.upsert.assert_awaited_once()
stored: Final = table.upsert.call_args.kwargs
assert stored["where"] == {"user_id_server_id": {"user_id": "jwt-owner", "server_id": server.server_id}}
credential: Final = stored["data"]["create"]["credential_b64"]
assert "upstream-token" not in credential
decoded: Final = decrypt_value_helper(credential, key="mcp_user_credential")
assert json.loads(decoded)["access_token"] == "upstream-token"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"rejection",
[
"expired",
"audience",
"issuer",
"signature",
"missing_user",
"unknown_user",
"disabled",
"not_premium",
"scim_inactive",
"custom_validate",
"missing_database",
],
)
async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
rejection: str,
) -> None:
from cryptography.hazmat.primitives.asymmetric import rsa
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
handler, signing_key = jwt_oauth_identity
key: Final = (
rsa.generate_private_key(public_exponent=65537, key_size=2048) if rejection == "signature" else signing_key
)
bearer: Final = _oauth_identity_jwt(
key,
expires_in=-60 if rejection == "expired" else 300,
audience="upstream-only" if rejection == "audience" else "litellm-proxy",
issuer="https://untrusted.example.test" if rejection == "issuer" else "https://idp.example.test",
owner=None if rejection == "missing_user" else "unknown" if rejection == "unknown_user" else "jwt-owner",
)
if rejection == "disabled":
monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": False})
if rejection == "not_premium":
monkeypatch.setattr(proxy_server, "premium_user", False)
if rejection == "missing_database":
monkeypatch.setattr(proxy_server, "prisma_client", None)
if rejection == "scim_inactive":
handler.user_api_key_cache.set_cache(
"jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", metadata={"scim_active": False})
)
if rejection == "custom_validate":
handler.litellm_jwtauth.custom_validate = lambda claims: False
assert await _extract_user_id_from_request(_token_request({"Authorization": f"Bearer {bearer}"})) is None
@pytest.mark.asyncio
@pytest.mark.parametrize("blocked", [False, True])
async def test_oauth_jwt_cannot_override_explicit_litellm_key(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
blocked: bool,
) -> None:
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._types import UserAPIKeyAuth, hash_token
handler, signing_key = jwt_oauth_identity
key: Final = "sk-explicit-key"
handler.user_api_key_cache.set_cache(hash_token(key), UserAPIKeyAuth(user_id="key-owner", blocked=blocked))
request: Final = _token_request(
{
"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}",
"x-litellm-api-key": key,
}
)
assert await _extract_user_id_from_request(request) == (None if blocked else "key-owner")
@pytest.mark.asyncio
@pytest.mark.parametrize("mapping", ["active", "blocked", "inactive_owner", "fallback", "pending", "reject"])
async def test_oauth_jwt_uses_configured_virtual_key_owner(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
mapping: str,
) -> None:
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._types import UserAPIKeyAuth, UnregisteredJWTClientBehavior, hash_token
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.virtual_key_claim_field = "sub"
handler.litellm_jwtauth.unregistered_jwt_client_behavior = (
UnregisteredJWTClientBehavior.AUTO_REGISTER
if mapping == "pending"
else UnregisteredJWTClientBehavior.REJECT
if mapping == "reject"
else UnregisteredJWTClientBehavior.FALLBACK_TEAM_MAPPING
)
key_hash: Final = hash_token("sk-mapped-oauth-owner")
handler.user_api_key_cache.set_cache(
jwt_key_mapping_cache_key("sub", "not-the-configured-user-id"),
"__NO_MAPPING__" if mapping in ("fallback", "pending", "reject") else key_hash,
)
handler.user_api_key_cache.set_cache(
key_hash, UserAPIKeyAuth(token=key_hash, user_id="mapped-owner", blocked=mapping == "blocked")
)
handler.user_api_key_cache.set_cache(
"mapped-owner", LiteLLM_UserTable(user_id="mapped-owner", metadata={"scim_active": mapping != "inactive_owner"})
)
request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"})
expected: Final = "jwt-owner" if mapping == "fallback" else "mapped-owner" if mapping == "active" else None
assert await _extract_user_id_from_request(request) == expected
@pytest.mark.asyncio
@pytest.mark.parametrize("allowed_domain", [None, "allowed.example.test"])
async def test_oauth_jwt_respects_custom_validation_and_email_policy(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
allowed_domain: str | None,
) -> None:
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.custom_validate = lambda claims: True
handler.litellm_jwtauth.user_allowed_email_domain = allowed_domain
request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"})
assert await _extract_user_id_from_request(request) == (None if allowed_domain else "jwt-owner")
@pytest.mark.asyncio
async def test_oauth_jwt_uses_rbac_user_object_id(jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"]) -> None:
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request
from litellm.proxy._types import LitellmUserRoles, RoleMapping
handler, signing_key = jwt_oauth_identity
handler.litellm_jwtauth.user_id_jwt_field = "sub"
handler.litellm_jwtauth.roles_jwt_field = "aud"
handler.litellm_jwtauth.object_id_jwt_field = "identity.user_id"
handler.litellm_jwtauth.role_mappings = [
RoleMapping(role="litellm-proxy", internal_role=LitellmUserRoles.INTERNAL_USER)
]
request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"})
assert await _extract_user_id_from_request(request) == "jwt-owner"