From 92e182b898a557849f374ba66af708999626806d Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 12:29:03 -0700 Subject: [PATCH 1/9] fix(mcp): persist OAuth credentials for validated JWT users --- .../mcp_server/bridge_token_flow.py | 56 ++++ .../mcp_server/discoverable_endpoints.py | 5 +- .../mcp_server/test_discoverable_endpoints.py | 277 +++++++++++++++++- 3 files changed, 335 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 35a30127e27..8dcc49c2fd8 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -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`` diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index bafe33d0a6b..94b6348b0f0 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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, ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 9ea870d3210..af4c9eea770 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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" From eda98f38d940ecd065aee71d4b0b34ba1be98a06 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 12:41:12 -0700 Subject: [PATCH 2/9] fix(mcp): preserve canonical JWT owner lookup without cached identity --- .../mcp_server/bridge_token_flow.py | 10 ++++-- .../mcp_server/test_discoverable_endpoints.py | 36 +++++++++++++++++++ 2 files changed, 43 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 8dcc49c2fd8..3817935bf71 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -198,7 +198,9 @@ async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | No return loaded if isinstance(loaded, str) else None -async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResolutionFailure": +async def load_active_user_by_id( + user_id: str, *, sso_user_id: str | None = None, user_email: str | None = None +) -> "LiteLLM_UserTable | _KeyResolutionFailure": """Load a live litellm user by id, returning the record when the user is active or a precise failure otherwise. The interactive DCR client authenticates via SSO, so its refresh envelope seals a user subject; renewing it must re-check the user is still live (present and not SCIM-deactivated) so a @@ -232,6 +234,8 @@ async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResol prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, user_id_upsert=False, + sso_user_id=sso_user_id, + user_email=user_email, ) except (ProxyException, HTTPException): return "no_active_key" @@ -352,7 +356,7 @@ async def _extract_jwt_user_id(token: str) -> str | None: 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) + user_id, user_email, 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 @@ -361,7 +365,7 @@ async def _extract_jwt_user_id(token: str) -> str | None: ) if not owner_id or valid_email is False: return None - owner: Final = await load_active_user_by_id(owner_id) + owner: Final = await load_active_user_by_id(owner_id, sso_user_id=owner_id, user_email=user_email) 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__) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index af4c9eea770..ce985064eb0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11428,6 +11428,7 @@ def _oauth_identity_jwt( { "sub": "not-the-configured-user-id", "identity": {"user_id": owner}, + "email": "owner@example.test", "iss": issuer, "aud": audience, "exp": int(time.time()) + expires_in, @@ -11649,3 +11650,38 @@ async def test_oauth_jwt_uses_rbac_user_object_id(jwt_oauth_identity: tuple["JWT ] request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) assert await _extract_user_id_from_request(request) == "jwt-owner" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("identity", ["sso", "email"]) +@pytest.mark.parametrize("inactive", [False, True]) +async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + identity: str, + inactive: bool, +) -> None: + 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 + external_id: Final = f"external-{identity}-{inactive}" + handler.litellm_jwtauth.user_email_jwt_field = "email" + owner: Final = LiteLLM_UserTable( + user_id="canonical-oauth-owner", + user_email="owner@example.test", + metadata={"scim_active": not inactive}, + organization_memberships=[], + ) + database: Final = MagicMock() + table: Final = database.db.litellm_usertable + table.find_unique = AsyncMock(side_effect=[None, owner if identity == "sso" else None]) + table.find_first = AsyncMock(return_value=owner) + table.update = AsyncMock(return_value=owner) + monkeypatch.setattr(proxy_server, "prisma_client", database) + request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key, owner=external_id)}"}) + assert await _extract_user_id_from_request(request) == (None if inactive else "canonical-oauth-owner") + assert table.find_unique.await_count == 2 + if identity == "email": + table.find_first.assert_awaited_once() From 88d0371a4670f57ca3e4c3824cb6846fe3fdacd7 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 16:50:55 -0700 Subject: [PATCH 3/9] fix(mcp): reuse the standard JWT auth builder for OAuth ownership --- .../mcp_server/bridge_token_flow.py | 38 ++++++---- .../mcp_server/test_discoverable_endpoints.py | 74 ++++++++++++++++--- 2 files changed, 88 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 3817935bf71..2f3e50b803f 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -314,15 +314,15 @@ async def _extract_user_id_from_request(request: Request) -> str | None: token: Final = _litellm_key_from_request(request) if token is not None and JWTHandler.is_jwt(token): - return await _extract_jwt_user_id(token) + return await _extract_jwt_user_id(request, 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 +async def _extract_jwt_user_id(request: Request, token: str) -> str | None: + from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.handle_jwt import JWTAuthManager # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle _resolve_jwt_to_virtual_key, # pyright: ignore[reportPrivateUsage] # reuse admission mapping policy without provisioning a new key @@ -339,11 +339,11 @@ async def _extract_jwt_user_id(token: str) -> str | None: 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(): + 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 mapped: Final = await _resolve_jwt_to_virtual_key( jwt_claims=claims, jwt_handler=jwt_handler, @@ -356,16 +356,24 @@ async def _extract_jwt_user_id(token: str) -> str | None: return None if await _key_owner_scim_deactivated(mapped) else _active_key_user_id(mapped) if mapped is not None: return None - user_id, user_email, 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 + identity: Final = await JWTAuthManager.auth_builder( + api_key=token, + jwt_handler=jwt_handler, + request_data={}, + general_settings=general_settings, + route=request.url.path, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + request_headers=dict(request.headers), + request_method=request.method, ) - if not owner_id or valid_email is False: + owner_id: Final = identity["user_id"] + if not owner_id: return None - owner: Final = await load_active_user_by_id(owner_id, sso_user_id=owner_id, user_email=user_email) + # Admin JWTs can return before auth_builder loads the canonical database user. + owner: Final = await load_active_user_by_id(owner_id, sso_user_id=owner_id, user_email=identity["user_email"]) 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__) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index ce985064eb0..d1ab8809861 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11421,6 +11421,7 @@ def _oauth_identity_jwt( audience: str = "litellm-proxy", issuer: str = "https://idp.example.test", owner: str | None = "jwt-owner", + scope: str = "", ) -> str: import jwt @@ -11432,6 +11433,7 @@ def _oauth_identity_jwt( "iss": issuer, "aud": audience, "exp": int(time.time()) + expires_in, + "scope": scope, }, signing_key, algorithm="RS256", @@ -11440,9 +11442,11 @@ def _oauth_identity_jwt( @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Authorization", "x-litellm-api-key"]) +@pytest.mark.parametrize("policy_allowed", [False, True]) async def test_oauth_exchange_stores_token_for_validated_jwt_user( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], header: str, + policy_allowed: bool, monkeypatch: pytest.MonkeyPatch, ) -> None: import httpx @@ -11451,7 +11455,8 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( from litellm.proxy._types import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer - _, signing_key = jwt_oauth_identity + handler, signing_key = jwt_oauth_identity + handler.litellm_jwtauth.enforce_team_based_model_access = not policy_allowed bearer: Final = _oauth_identity_jwt(signing_key) request: Final = _token_request({header: f"Bearer {bearer}"}) server: Final = MCPServer( @@ -11501,6 +11506,10 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( code_verifier=None, ) assert response.status_code == 200 + assert json.loads(response.body)["access_token"] == "upstream-token" + if not policy_allowed: + table.upsert.assert_not_awaited() + return 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}} @@ -11525,6 +11534,8 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( "scim_inactive", "custom_validate", "missing_database", + "denied_route", + "required_team", ], ) async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( @@ -11537,6 +11548,7 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( 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 + from litellm.proxy._types import LitellmUserRoles, RoleBasedPermissions handler, signing_key = jwt_oauth_identity key: Final = ( @@ -11561,6 +11573,18 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( ) if rejection == "custom_validate": handler.litellm_jwtauth.custom_validate = lambda claims: False + if rejection == "denied_route": + handler.litellm_jwtauth.enforce_rbac = True + monkeypatch.setattr( + proxy_server, + "general_settings", + { + "enable_jwt_auth": True, + "role_permissions": [RoleBasedPermissions(role=LitellmUserRoles.INTERNAL_USER, routes=["/models"])], + }, + ) + if rejection == "required_team": + handler.litellm_jwtauth.enforce_team_based_model_access = True assert await _extract_user_id_from_request(_token_request({"Authorization": f"Bearer {bearer}"})) is None @@ -11586,7 +11610,9 @@ async def test_oauth_jwt_cannot_override_explicit_litellm_key( @pytest.mark.asyncio -@pytest.mark.parametrize("mapping", ["active", "blocked", "inactive_owner", "fallback", "pending", "reject"]) +@pytest.mark.parametrize( + "mapping", ["active", "blocked", "inactive_owner", "fallback", "pending", "reject", "custom_reject"] +) async def test_oauth_jwt_uses_configured_virtual_key_owner( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], mapping: str, @@ -11598,6 +11624,8 @@ async def test_oauth_jwt_uses_configured_virtual_key_owner( handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.virtual_key_claim_field = "sub" + if mapping == "custom_reject": + handler.litellm_jwtauth.custom_validate = lambda claims: False handler.litellm_jwtauth.unregistered_jwt_client_behavior = ( UnregisteredJWTClientBehavior.AUTO_REGISTER if mapping == "pending" @@ -11637,9 +11665,15 @@ async def test_oauth_jwt_respects_custom_validation_and_email_policy( @pytest.mark.asyncio -async def test_oauth_jwt_uses_rbac_user_object_id(jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"]) -> None: +@pytest.mark.parametrize("route_allowed", [False, True]) +async def test_oauth_jwt_uses_rbac_user_object_id( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + route_allowed: bool, +) -> None: + from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request - from litellm.proxy._types import LitellmUserRoles, RoleMapping + from litellm.proxy._types import LitellmUserRoles, RoleBasedPermissions, RoleMapping handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.user_id_jwt_field = "sub" @@ -11648,26 +11682,43 @@ async def test_oauth_jwt_uses_rbac_user_object_id(jwt_oauth_identity: tuple["JWT handler.litellm_jwtauth.role_mappings = [ RoleMapping(role="litellm-proxy", internal_role=LitellmUserRoles.INTERNAL_USER) ] + handler.litellm_jwtauth.enforce_rbac = True + monkeypatch.setattr( + proxy_server, + "general_settings", + { + "enable_jwt_auth": True, + "role_permissions": [ + RoleBasedPermissions( + role=LitellmUserRoles.INTERNAL_USER, + routes=["/token"] if route_allowed else ["/models"], + ) + ], + }, + ) request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) - assert await _extract_user_id_from_request(request) == "jwt-owner" + assert await _extract_user_id_from_request(request) == ("jwt-owner" if route_allowed else None) @pytest.mark.asyncio @pytest.mark.parametrize("identity", ["sso", "email"]) @pytest.mark.parametrize("inactive", [False, True]) +@pytest.mark.parametrize("admin", [False, True]) async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch, identity: str, inactive: bool, + admin: bool, ) -> None: 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 - external_id: Final = f"external-{identity}-{inactive}" + external_id: Final = f"external-{identity}-{inactive}-{admin}" handler.litellm_jwtauth.user_email_jwt_field = "email" + handler.litellm_jwtauth.admin_allowed_routes = ["/token"] owner: Final = LiteLLM_UserTable( user_id="canonical-oauth-owner", user_email="owner@example.test", @@ -11676,12 +11727,17 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( ) database: Final = MagicMock() table: Final = database.db.litellm_usertable - table.find_unique = AsyncMock(side_effect=[None, owner if identity == "sso" else None]) + table.find_unique = AsyncMock(side_effect=[None, owner if identity == "sso" else None, owner]) table.find_first = AsyncMock(return_value=owner) table.update = AsyncMock(return_value=owner) monkeypatch.setattr(proxy_server, "prisma_client", database) - request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key, owner=external_id)}"}) + bearer: Final = _oauth_identity_jwt( + signing_key, owner=external_id, scope="litellm_proxy_admin" if admin else "" + ) + request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) assert await _extract_user_id_from_request(request) == (None if inactive else "canonical-oauth-owner") - assert table.find_unique.await_count == 2 + assert table.find_unique.await_count == (2 if admin else 3) + if not admin: + assert table.find_unique.call_args.kwargs["where"] == {"user_id": "canonical-oauth-owner"} if identity == "email": table.find_first.assert_awaited_once() From 53318796fd27d7e59bb86850f163836dda7a45e8 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 17:13:10 -0700 Subject: [PATCH 4/9] fix(mcp): separate JWT identity lookup from request authorization --- .../mcp_server/bridge_token_flow.py | 22 +-- litellm/proxy/auth/handle_jwt.py | 71 ++++++++-- .../mcp_server/test_discoverable_endpoints.py | 126 +++++++++++++----- .../proxy/auth/test_handle_jwt.py | 71 +++++++++- 4 files changed, 237 insertions(+), 53 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 2f3e50b803f..6d11a4d2607 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -198,9 +198,7 @@ async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | No return loaded if isinstance(loaded, str) else None -async def load_active_user_by_id( - user_id: str, *, sso_user_id: str | None = None, user_email: str | None = None -) -> "LiteLLM_UserTable | _KeyResolutionFailure": +async def load_active_user_by_id(user_id: str) -> "LiteLLM_UserTable | _KeyResolutionFailure": """Load a live litellm user by id, returning the record when the user is active or a precise failure otherwise. The interactive DCR client authenticates via SSO, so its refresh envelope seals a user subject; renewing it must re-check the user is still live (present and not SCIM-deactivated) so a @@ -234,8 +232,6 @@ async def load_active_user_by_id( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, user_id_upsert=False, - sso_user_id=sso_user_id, - user_email=user_email, ) except (ProxyException, HTTPException): return "no_active_key" @@ -247,6 +243,10 @@ async def load_active_user_by_id( return "no_active_key" if user_object is None: return "no_active_key" + return _active_user_record(user_object) + + +def _active_user_record(user_object: "LiteLLM_UserTable") -> "LiteLLM_UserTable | Literal['no_active_key']": if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False: return "no_active_key" return user_object @@ -336,7 +336,7 @@ async def _extract_jwt_user_id(request: Request, token: str) -> str | None: user_api_key_cache, ) - if general_settings.get("enable_jwt_auth") is not True or premium_user is not True: + if general_settings.get("enable_jwt_auth") is not True or premium_user is not True or prisma_client is None: return None try: if jwt_handler.litellm_jwtauth.is_virtual_key_mapping_configured(): @@ -368,13 +368,13 @@ async def _extract_jwt_user_id(request: Request, token: str) -> str | None: proxy_logging_obj=proxy_logging_obj, request_headers=dict(request.headers), request_method=request.method, + identity_only=True, ) - owner_id: Final = identity["user_id"] - if not owner_id: + resolved_user: Final = identity["user_object"] + if resolved_user is None: return None - # Admin JWTs can return before auth_builder loads the canonical database user. - owner: Final = await load_active_user_by_id(owner_id, sso_user_id=owner_id, user_email=identity["user_email"]) - return None if isinstance(owner, str) else owner.user_id + owner: Final = _active_user_record(resolved_user) + return None if isinstance(owner, str) else identity["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 diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 4304542fc83..678769b375b 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1674,6 +1674,7 @@ class JWTAuthManager: proxy_logging_obj: ProxyLogging, route: str, org_alias: str | None = None, + user_id_upsert: bool | None = None, ) -> tuple[ LiteLLM_UserTable | None, LiteLLM_OrganizationTable | None, @@ -1737,7 +1738,11 @@ class JWTAuthManager: user_id=user_id, user_email=user_email, sso_user_id=user_id, - upsert=jwt_handler.is_upsert_user_id(valid_user_email=valid_user_email), + upsert=( + jwt_handler.is_upsert_user_id(valid_user_email=valid_user_email) + if user_id_upsert is None + else user_id_upsert + ), ), team_id=team_id, ) @@ -2209,8 +2214,14 @@ class JWTAuthManager: proxy_logging_obj: ProxyLogging, request_headers: dict | None = None, request_method: str | None = None, + identity_only: bool = False, ) -> JWTAuthBuilderResult: - """Main authentication and authorization builder""" + """Build JWT authentication and authorization context. + + Public OAuth endpoints use identity_only to resolve an existing credential owner + without authorizing the OAuth route or provisioning users/teams. The returned + identity does not grant permission to execute an MCP or model request. + """ # Check if OIDC UserInfo endpoint is enabled, but fall back to standard # JWT auth if the token itself is a well-formed JWT (3-part structure). if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not jwt_handler.is_jwt(token=api_key): @@ -2231,18 +2242,23 @@ class JWTAuthManager: # Check RBAC rbac_role: Final = jwt_handler.get_rbac_role(token=jwt_valid_token) - await JWTAuthManager.check_rbac_role( - jwt_handler, - jwt_valid_token, - general_settings, - request_data, - route, - rbac_role, - ) + if not identity_only: + await JWTAuthManager.check_rbac_role( + jwt_handler, + jwt_valid_token, + general_settings, + request_data, + route, + rbac_role, + ) # Check Scope Based Access scopes: Final = jwt_handler.get_scopes(token=jwt_valid_token) - if jwt_handler.litellm_jwtauth.enforce_scope_based_access and jwt_handler.litellm_jwtauth.scope_mappings: + if ( + not identity_only + and jwt_handler.litellm_jwtauth.enforce_scope_based_access + and jwt_handler.litellm_jwtauth.scope_mappings + ): JWTAuthManager.check_scope_based_access( scope_mappings=jwt_handler.litellm_jwtauth.scope_mappings, scopes=scopes, @@ -2268,6 +2284,39 @@ class JWTAuthManager: elif rbac_role == LitellmUserRoles.INTERNAL_USER: user_id = object_id + if identity_only: + identity_user, _, _, _, identity_user_id = await JWTAuthManager.get_objects( + user_id=user_id, + user_email=user_email, + org_id=None, + end_user_id=None, + team_id=None, + valid_user_email=valid_user_email, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + route=route, + user_id_upsert=False, + ) + return JWTAuthBuilderResult( + is_proxy_admin=False, + # Admin admission uses the claim ID; other callers use the canonical DB ID. + user_id=user_id if jwt_handler.is_admin(scopes=scopes) else identity_user_id, + user_email=identity_user.user_email if identity_user is not None else user_email, + user_object=identity_user, + team_id=None, + team_object=None, + org_id=None, + org_object=None, + end_user_id=None, + end_user_object=None, + team_membership=None, + token=api_key, + jwt_claims=jwt_valid_token, + ) + # Check admin access admin_result: Final = await JWTAuthManager.check_admin_access( jwt_handler, scopes, route, user_id, org_id, api_key, jwt_valid_token, user_email=user_email diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index d1ab8809861..bf557667892 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7127,12 +7127,12 @@ async def test_build_oauth_protected_resource_response_obo_end_to_end(): global_mcp_server_manager.registry.clear() -def _token_request(headers): +def _token_request(headers, path="/token"): """A real Starlette request with case-insensitive headers (matches production).""" from starlette.requests import Request raw = [(k.lower().encode(), v.encode()) for k, v in headers.items()] - return Request({"type": "http", "method": "POST", "path": "/token", "headers": raw, "query_string": b""}) + return Request({"type": "http", "method": "POST", "path": path, "headers": raw, "query_string": b""}) @pytest.fixture @@ -11443,10 +11443,12 @@ def _oauth_identity_jwt( @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Authorization", "x-litellm-api-key"]) @pytest.mark.parametrize("policy_allowed", [False, True]) +@pytest.mark.parametrize("admin", [False, True]) async def test_oauth_exchange_stores_token_for_validated_jwt_user( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], header: str, policy_allowed: bool, + admin: bool, monkeypatch: pytest.MonkeyPatch, ) -> None: import httpx @@ -11456,9 +11458,9 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( from litellm.types.mcp_server.mcp_server_manager import MCPServer handler, signing_key = jwt_oauth_identity - handler.litellm_jwtauth.enforce_team_based_model_access = not policy_allowed - bearer: Final = _oauth_identity_jwt(signing_key) - request: Final = _token_request({header: f"Bearer {bearer}"}) + handler.litellm_jwtauth.custom_validate = lambda claims: policy_allowed + bearer: Final = _oauth_identity_jwt(signing_key, scope="litellm_proxy_admin" if admin else "") + request: Final = _token_request({header: f"Bearer {bearer}"}, path="/jwt-oauth-server/token") server: Final = MCPServer( server_id="jwt-oauth-server", name="jwt-oauth-server", @@ -11534,8 +11536,6 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( "scim_inactive", "custom_validate", "missing_database", - "denied_route", - "required_team", ], ) async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( @@ -11548,7 +11548,6 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( 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 - from litellm.proxy._types import LitellmUserRoles, RoleBasedPermissions handler, signing_key = jwt_oauth_identity key: Final = ( @@ -11573,18 +11572,6 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( ) if rejection == "custom_validate": handler.litellm_jwtauth.custom_validate = lambda claims: False - if rejection == "denied_route": - handler.litellm_jwtauth.enforce_rbac = True - monkeypatch.setattr( - proxy_server, - "general_settings", - { - "enable_jwt_auth": True, - "role_permissions": [RoleBasedPermissions(role=LitellmUserRoles.INTERNAL_USER, routes=["/models"])], - }, - ) - if rejection == "required_team": - handler.litellm_jwtauth.enforce_team_based_model_access = True assert await _extract_user_id_from_request(_token_request({"Authorization": f"Bearer {bearer}"})) is None @@ -11666,7 +11653,7 @@ async def test_oauth_jwt_respects_custom_validation_and_email_policy( @pytest.mark.asyncio @pytest.mark.parametrize("route_allowed", [False, True]) -async def test_oauth_jwt_uses_rbac_user_object_id( +async def test_oauth_jwt_identity_preserves_separate_mcp_route_authorization( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch, route_allowed: bool, @@ -11674,6 +11661,7 @@ async def test_oauth_jwt_uses_rbac_user_object_id( from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request from litellm.proxy._types import LitellmUserRoles, RoleBasedPermissions, RoleMapping + from litellm.proxy.auth.handle_jwt import JWTAuthManager handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.user_id_jwt_field = "sub" @@ -11691,13 +11679,32 @@ async def test_oauth_jwt_uses_rbac_user_object_id( "role_permissions": [ RoleBasedPermissions( role=LitellmUserRoles.INTERNAL_USER, - routes=["/token"] if route_allowed else ["/models"], + routes=["mcp_routes"] if route_allowed else ["/models"], ) ], }, ) - request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) - assert await _extract_user_id_from_request(request) == ("jwt-owner" if route_allowed else None) + bearer: Final = _oauth_identity_jwt(signing_key) + request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/example/token") + assert await _extract_user_id_from_request(request) == "jwt-owner" + admission: Final = JWTAuthManager.auth_builder( + api_key=bearer, + jwt_handler=handler, + request_data={}, + general_settings=proxy_server.general_settings, + route="/mcp/example", + prisma_client=proxy_server.prisma_client, + user_api_key_cache=handler.user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_server.proxy_logging_obj, + request_method="POST", + ) + if route_allowed: + assert (await admission)["user_id"] == "jwt-owner" + else: + with pytest.raises(HTTPException) as denial: + await admission + assert denial.value.status_code == 403 @pytest.mark.asyncio @@ -11714,11 +11721,12 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( 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 + from litellm.proxy.auth.handle_jwt import JWTAuthManager handler, signing_key = jwt_oauth_identity external_id: Final = f"external-{identity}-{inactive}-{admin}" handler.litellm_jwtauth.user_email_jwt_field = "email" - handler.litellm_jwtauth.admin_allowed_routes = ["/token"] + handler.litellm_jwtauth.admin_allowed_routes = ["mcp_routes"] owner: Final = LiteLLM_UserTable( user_id="canonical-oauth-owner", user_email="owner@example.test", @@ -11727,7 +11735,7 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( ) database: Final = MagicMock() table: Final = database.db.litellm_usertable - table.find_unique = AsyncMock(side_effect=[None, owner if identity == "sso" else None, owner]) + table.find_unique = AsyncMock(side_effect=[None, owner if identity == "sso" else None]) table.find_first = AsyncMock(return_value=owner) table.update = AsyncMock(return_value=owner) monkeypatch.setattr(proxy_server, "prisma_client", database) @@ -11735,9 +11743,67 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( signing_key, owner=external_id, scope="litellm_proxy_admin" if admin else "" ) request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) - assert await _extract_user_id_from_request(request) == (None if inactive else "canonical-oauth-owner") - assert table.find_unique.await_count == (2 if admin else 3) - if not admin: - assert table.find_unique.call_args.kwargs["where"] == {"user_id": "canonical-oauth-owner"} + stored_owner: Final = await _extract_user_id_from_request(request) + assert stored_owner == (None if inactive else external_id if admin else "canonical-oauth-owner") + assert table.find_unique.await_count == 2 if identity == "email": table.find_first.assert_awaited_once() + if not inactive: + admission: Final = await JWTAuthManager.auth_builder( + api_key=bearer, + jwt_handler=handler, + request_data={}, + general_settings=proxy_server.general_settings, + route="/mcp/example", + prisma_client=database, + user_api_key_cache=handler.user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_server.proxy_logging_obj, + ) + assert stored_owner == admission["user_id"] + + +@pytest.mark.asyncio +async def test_oauth_jwt_identity_does_not_provision_or_synchronize_teams( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], +) -> None: + 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 + handler.litellm_jwtauth.enforce_team_based_model_access = True + handler.litellm_jwtauth.team_id_default = "new-team" + handler.litellm_jwtauth.team_id_upsert = True + handler.litellm_jwtauth.sync_user_role_and_teams = True + owner: Final = LiteLLM_UserTable(user_id="jwt-owner", teams=["existing-team"]) + handler.user_api_key_cache.set_cache("jwt-owner", owner) + request: Final = _token_request( + {"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}, path="/example/token" + ) + assert await _extract_user_id_from_request(request) == "jwt-owner" + assert owner.teams == ["existing-team"] + proxy_server.prisma_client.db.litellm_teamtable.find_unique.assert_not_called() + proxy_server.prisma_client.db.litellm_teamtable.upsert.assert_not_called() + proxy_server.prisma_client.db.litellm_usertable.update.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("state", ["active", "inactive", "missing_database"]) +async def test_oauth_refresh_revalidates_the_same_active_user_rule( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + state: str, +) -> None: + from litellm.models.user import LiteLLM_UserTable + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id + + handler, _ = jwt_oauth_identity + handler.user_api_key_cache.set_cache( + "jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", metadata={"scim_active": state != "inactive"}) + ) + if state == "missing_database": + monkeypatch.setattr(proxy_server, "prisma_client", None) + expected: Final = None if state == "active" else "no_active_key" if state == "inactive" else "unresolvable" + assert await _reload_active_user_by_id("jwt-owner") == expected diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 94226b5404d..9e15a442a18 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -2,7 +2,7 @@ import asyncio import re import time from collections.abc import Mapping, Sequence -from typing import Optional +from typing import Final, Optional from unittest.mock import AsyncMock, MagicMock, patch from fastapi import HTTPException @@ -6786,3 +6786,72 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla } assert mock_patch.call_args.kwargs["teams_ids_to_add_user_to"] == [] assert user.teams == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("identity_only", [False, True]) +@pytest.mark.parametrize("existing_user", [False, True]) +@pytest.mark.parametrize("model_allowed", [False, True]) +async def test_auth_builder_identity_lookup_does_not_provision_users( + monkeypatch: pytest.MonkeyPatch, identity_only: bool, existing_user: bool, model_allowed: bool +) -> None: + from litellm.proxy._types import ScopeMapping + from litellm.proxy.auth.auth_checks import UserNotFoundError + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + private_key, jwk = _get_rsa_key_and_jwk("identity-mode") + cache: Final = UserApiKeyCache() + cache.set_cache("litellm_jwt_auth_keys_https://identity.example/jwks", [jwk]) + user_id: Final = f"identity-mode-{identity_only}-{existing_user}-{model_allowed}" + user: Final = LiteLLM_UserTable(user_id=user_id, organization_memberships=[]) + if existing_user: + cache.set_cache(user_id, user) + database: Final = MagicMock() + users: Final = database.db.litellm_usertable + users.find_unique = AsyncMock(return_value=None) + users.find_first = AsyncMock(return_value=None) + users.create = AsyncMock(return_value=user) + handler: Final = JWTHandler() + handler.update_environment( + prisma_client=database, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + user_id_jwt_field="sub", + user_id_upsert=True, + enforce_scope_based_access=True, + scope_mappings=[ScopeMapping(scope="allowed", models=["allowed-model"])], + ), + ) + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", "https://identity.example/jwks") + monkeypatch.setenv("JWT_ISSUER", "https://identity.example") + monkeypatch.setenv("JWT_AUDIENCE", "gateway") + token: Final = _encode_rsa_jwt( + private_key, "https://identity.example", "gateway", "identity-mode", {"sub": user_id, "scope": "allowed"} + ) + pending: Final = JWTAuthManager.auth_builder( + api_key=token, + jwt_handler=handler, + request_data={"model": "allowed-model" if model_allowed else "forbidden-model"}, + general_settings={}, + route="/example/token" if identity_only else "/mcp/example", + prisma_client=database, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + identity_only=identity_only, + ) + if not identity_only and not model_allowed: + with pytest.raises(HTTPException) as denial: + await pending + assert denial.value.status_code == 403 + users.create.assert_not_awaited() + return + if identity_only and not existing_user: + with pytest.raises(UserNotFoundError): + await pending + else: + result: Final = await pending + assert result["user_id"] == user_id + assert result["user_object"] is not None + assert result["user_object"].user_id == user_id + assert users.create.await_count == (0 if identity_only or existing_user else 1) From 61e3b5ddae1fdbbf12f6fa087b8bf0ee3e318d12 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 17:31:07 -0700 Subject: [PATCH 5/9] fix(mcp): persist OAuth credentials for rowless JWT admins --- .../mcp_server/bridge_token_flow.py | 5 ++- litellm/proxy/auth/handle_jwt.py | 36 +++++++++++-------- .../mcp_server/test_discoverable_endpoints.py | 18 +++++++++- 3 files changed, 40 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 6d11a4d2607..b693cc046c8 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -371,10 +371,9 @@ async def _extract_jwt_user_id(request: Request, token: str) -> str | None: identity_only=True, ) resolved_user: Final = identity["user_object"] - if resolved_user is None: + if resolved_user is not None and isinstance(_active_user_record(resolved_user), str): return None - owner: Final = _active_user_record(resolved_user) - return None if isinstance(owner, str) else identity["user_id"] + return identity["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 diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index d5fcc578167..2d8cfb614ce 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -62,6 +62,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.user_repository import UserRepository from litellm.types.agents import AgentResponse +from litellm.types.proxy.auth.auth_checks import UserNotFoundError from .auth_checks import ( _allowed_routes_check, @@ -2343,21 +2344,26 @@ class JWTAuthManager: ) if identity_only: - identity_user, _, _, _, identity_user_id = await JWTAuthManager.get_objects( - user_id=user_id, - user_email=user_email, - org_id=None, - end_user_id=None, - team_id=None, - valid_user_email=valid_user_email, - jwt_handler=jwt_handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - route=route, - user_id_upsert=False, - ) + try: + identity_user, _, _, _, identity_user_id = await JWTAuthManager.get_objects( + user_id=user_id, + user_email=user_email, + org_id=None, + end_user_id=None, + team_id=None, + valid_user_email=valid_user_email, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + route=route, + user_id_upsert=False, + ) + except UserNotFoundError: + if not jwt_handler.is_admin(scopes=scopes): + raise + identity_user, identity_user_id = None, user_id return JWTAuthBuilderResult( is_proxy_admin=False, # Admin admission uses the claim ID; other callers use the canonical DB ID. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index bf557667892..f6dc7932e2f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11444,11 +11444,13 @@ def _oauth_identity_jwt( @pytest.mark.parametrize("header", ["Authorization", "x-litellm-api-key"]) @pytest.mark.parametrize("policy_allowed", [False, True]) @pytest.mark.parametrize("admin", [False, True]) +@pytest.mark.parametrize("owner_state", ["active", "missing", "inactive", "database_error"]) async def test_oauth_exchange_stores_token_for_validated_jwt_user( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], header: str, policy_allowed: bool, admin: bool, + owner_state: str, monkeypatch: pytest.MonkeyPatch, ) -> None: import httpx @@ -11475,6 +11477,7 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( 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.models.user import LiteLLM_UserTable from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.types.llms.custom_http import httpxSpecialProvider @@ -11485,6 +11488,18 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( return httpx.Response(200, json={"access_token": "upstream-token", "token_type": "Bearer"}) database: Final = MagicMock() + users: Final = database.db.litellm_usertable + users.find_unique = AsyncMock(return_value=None) + users.find_first = AsyncMock(return_value=None) + users.create = AsyncMock() + if owner_state in ("missing", "database_error"): + handler.user_api_key_cache.delete_cache("jwt-owner") + if owner_state == "database_error": + users.find_unique.side_effect = RuntimeError("database unavailable") + if owner_state == "inactive": + handler.user_api_key_cache.set_cache( + "jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", metadata={"scim_active": False}) + ) table: Final = database.db.litellm_mcpusercredentials table.find_unique = AsyncMock(return_value=None) table.upsert = AsyncMock() @@ -11509,7 +11524,8 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( ) assert response.status_code == 200 assert json.loads(response.body)["access_token"] == "upstream-token" - if not policy_allowed: + users.create.assert_not_awaited() + if not policy_allowed or owner_state in ("inactive", "database_error") or (owner_state == "missing" and not admin): table.upsert.assert_not_awaited() return table.upsert.assert_awaited_once() From 97211bc356d47d39214311e7cea432f858d28445 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 19:00:47 -0700 Subject: [PATCH 6/9] fix(mcp): authorize per-user OAuth credential writes --- .../mcp_server/bridge_token_flow.py | 64 +++++-- .../mcp_server/discoverable_endpoints.py | 24 ++- .../mcp_server/ui_session_utils.py | 13 ++ litellm/proxy/auth/handle_jwt.py | 172 +++++++++++++----- litellm/proxy/auth/user_api_key_auth.py | 37 +--- .../mcp_management_endpoints.py | 10 +- .../mcp_server/test_discoverable_endpoints.py | 157 +++++++++++++++- .../proxy/auth/test_handle_jwt.py | 8 +- 8 files changed, 369 insertions(+), 116 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index b693cc046c8..962e39d7dd6 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -304,24 +304,55 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol assert_never(identity.subject_type) -async def _extract_user_id_from_request(request: Request) -> str | None: - """The litellm ``user_id`` for the token request, so a per-user token is stored under the same - identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome - (including a transient DB outage) collapses to ``None`` here and the caller simply skips the store; - the bridge mint, which must status those outcomes differently, consumes - :func:`_resolve_active_litellm_key` directly.""" +async def _extract_user_id_from_request(request: Request, server_id: str | None = None) -> str | None: + """Resolve identity for binding, or authorize the credential-write action for a target server.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # registry imports auth helpers + global_mcp_server_manager, + ) + from litellm.proxy._experimental.mcp_server.ui_session_utils import ( + can_access_mcp_server, # noqa: PLC0415 # proxy import cycle + ) + from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle + from litellm.proxy.auth.route_checks import RouteChecks # noqa: PLC0415 # proxy import cycle + from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle + _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action + ) token: Final = _litellm_key_from_request(request) - if token is not None and JWTHandler.is_jwt(token): - return await _extract_jwt_user_id(request, token) - resolved: Final = await _resolve_active_litellm_key(request) - if not isinstance(resolved, _ResolvedKey): + # The OAuth relay is public; the optional server-side write is the same protected action + # as the direct credential endpoint. Authorize that action without rewriting the Request. + write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential" if server_id is not None else None + resolved: Final = ( + await _resolve_jwt_auth(request, token, write_route) + if token is not None and JWTHandler.is_jwt(token) + else await _resolve_active_litellm_key(request) + ) + auth: Final = resolved.key if isinstance(resolved, _ResolvedKey) else resolved + if not isinstance(auth, UserAPIKeyAuth) or not _active_key_user_id(auth): return None - return _active_key_user_id(resolved.key) + if write_route is not None and server_id is not None: + try: + RouteChecks.is_virtual_key_allowed_to_call_route(route=write_route, valid_token=auth, request=request) + await _run_centralized_common_checks( + user_api_key_auth_obj=auth, + request=request, + request_data={}, + route=write_route, + ) + if not await can_access_mcp_server(auth, server_id, global_mcp_server_manager.get_allowed_mcp_servers): + return None + except Exception as exc: # noqa: BLE001 # authorization failure must never write credentials + verbose_logger.debug("OAuth credential write not authorized (%s)", type(exc).__name__) + return None + return auth.user_id -async def _extract_jwt_user_id(request: Request, token: str) -> str | None: +async def _resolve_jwt_auth( + request: Request, + token: str, + write_route: str | None, +) -> "UserAPIKeyAuth | None": from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.handle_jwt import JWTAuthManager # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle @@ -353,7 +384,7 @@ async def _extract_jwt_user_id(request: Request, token: str) -> str | None: proxy_logging_obj=proxy_logging_obj, ) if isinstance(mapped, UserAPIKeyAuth): - return None if await _key_owner_scim_deactivated(mapped) else _active_key_user_id(mapped) + return None if await _key_owner_scim_deactivated(mapped) or not _active_key_user_id(mapped) else mapped if mapped is not None: return None identity: Final = await JWTAuthManager.auth_builder( @@ -361,19 +392,20 @@ async def _extract_jwt_user_id(request: Request, token: str) -> str | None: jwt_handler=jwt_handler, request_data={}, general_settings=general_settings, - route=request.url.path, + route=write_route or request.url.path, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, proxy_logging_obj=proxy_logging_obj, request_headers=dict(request.headers), request_method=request.method, - identity_only=True, + identity_only=write_route is None, + allow_provisioning=False, ) resolved_user: Final = identity["user_object"] if resolved_user is not None and isinstance(_active_user_record(resolved_user), str): return None - return identity["user_id"] + return JWTAuthManager.user_api_key_auth_from_result(identity) except Exception as exc: # noqa: BLE001 # public OAuth exchange stays available; unvalidated identities never write credentials verbose_logger.debug("OAuth JWT identity could not be validated (%s)", type(exc).__name__) return None diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 94b6348b0f0..56968745ea9 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1218,12 +1218,26 @@ async def exchange_token_with_server( user_id: Final = resolved_user_id if user_id: try: - await _store_per_user_token_server_side( - server=resolved_server, - user_id=user_id, - token_response=token_response, - identity_binding_proof=binding_proof, + # Identity binding above must retain the verified caller even when a write is + # denied. Authorize persistence separately, immediately before its side effect. + can_store: Final = ( + await _user_can_reach_mcp_server(user_id, resolved_server.server_id) + if bridge_identity is not None + else await _extract_user_id_from_request(request, resolved_server.server_id) == user_id ) + if can_store: + await _store_per_user_token_server_side( + server=resolved_server, + user_id=user_id, + token_response=token_response, + identity_binding_proof=binding_proof, + ) + else: + verbose_logger.warning( + "OAuth credential storage not authorized for user=%s server=%s", + user_id, + resolved_server.server_id, + ) except Exception as exc: verbose_logger.warning( "exchange_token_with_server: server-side storage failed for user=%s server=%s: %s", diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index 188bfce1484..107a4818de1 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -2,6 +2,7 @@ from __future__ import annotations +from collections.abc import Awaitable, Callable from typing import Final from fastapi import HTTPException @@ -137,3 +138,15 @@ async def build_effective_auth_contexts( if admitted_context is None: return team_contexts return [*team_contexts, admitted_context] + + +async def can_access_mcp_server( + user_api_key_auth: UserAPIKeyAuth, + server_id: str, + allowed_servers: Callable[[UserAPIKeyAuth], Awaitable[list[str]]], +) -> bool: + """Resolve server access through the same credential contexts as MCP management.""" + for context in await build_effective_auth_contexts(user_api_key_auth): + if server_id in await allowed_servers(context): + return True + return False diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 2d8cfb614ce..f2bdbdd9341 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -9,6 +9,7 @@ JWT token must have 'litellm_proxy_admin' in scope. from __future__ import annotations import asyncio +import copy import fnmatch import hashlib import os @@ -54,7 +55,7 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_checks import can_team_access_model from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.auth.team_grants import team_model_aliases +from litellm.proxy.auth.team_grants import team_grants, team_model_aliases from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, get_management_object_ttl, @@ -2268,36 +2269,49 @@ class JWTAuthManager: request_headers: dict | None = None, request_method: str | None = None, identity_only: bool = False, + allow_provisioning: bool = True, ) -> JWTAuthBuilderResult: """Build JWT authentication and authorization context. - Public OAuth endpoints use identity_only to resolve an existing credential owner - without authorizing the OAuth route or provisioning users/teams. The returned - identity does not grant permission to execute an MCP or model request. + identity_only resolves the caller for OAuth identity binding and grants no permission. + Credential writes use full authorization with allow_provisioning=False: resolve the + existing policy context without creating users/teams or synchronizing membership. + A private handler configuration keeps that restriction out of concurrent normal requests. """ + handler: Final = jwt_handler if allow_provisioning else copy.copy(jwt_handler) + if not allow_provisioning: + handler.update_environment( + prisma_client=jwt_handler.prisma_client, + user_api_key_cache=jwt_handler.user_api_key_cache, + litellm_jwtauth=jwt_handler.litellm_jwtauth.model_copy( + update={"user_id_upsert": False, "team_id_upsert": False, "sync_user_role_and_teams": False} + ), + leeway=jwt_handler.leeway, + ) + # Check if OIDC UserInfo endpoint is enabled, but fall back to standard # JWT auth if the token itself is a well-formed JWT (3-part structure). - if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not jwt_handler.is_jwt(token=api_key): + if handler.litellm_jwtauth.oidc_userinfo_enabled and not handler.is_jwt(token=api_key): verbose_proxy_logger.debug("OIDC UserInfo is enabled. Fetching user info from UserInfo endpoint.") # Use the access token to fetch user info from OIDC UserInfo endpoint - jwt_valid_token: dict = await jwt_handler.get_oidc_userinfo(token=api_key) + jwt_valid_token: dict = await handler.get_oidc_userinfo(token=api_key) else: # Default behavior: decode and validate the JWT token - jwt_valid_token = await jwt_handler.auth_jwt(token=api_key) + jwt_valid_token = await handler.auth_jwt(token=api_key) # Check custom validate - if jwt_handler.litellm_jwtauth.custom_validate: - if not jwt_handler.litellm_jwtauth.custom_validate(jwt_valid_token): + if handler.litellm_jwtauth.custom_validate: + if not handler.litellm_jwtauth.custom_validate(jwt_valid_token): raise HTTPException( status_code=403, detail="Invalid JWT token", ) # Check RBAC - rbac_role: Final = jwt_handler.get_rbac_role(token=jwt_valid_token) + rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) if not identity_only: await JWTAuthManager.check_rbac_role( - jwt_handler, + handler, jwt_valid_token, general_settings, request_data, @@ -2306,30 +2320,30 @@ class JWTAuthManager: ) # Check Scope Based Access - scopes: Final = jwt_handler.get_scopes(token=jwt_valid_token) + scopes: Final = handler.get_scopes(token=jwt_valid_token) if ( not identity_only - and jwt_handler.litellm_jwtauth.enforce_scope_based_access - and jwt_handler.litellm_jwtauth.scope_mappings + and handler.litellm_jwtauth.enforce_scope_based_access + and handler.litellm_jwtauth.scope_mappings ): JWTAuthManager.check_scope_based_access( - scope_mappings=jwt_handler.litellm_jwtauth.scope_mappings, + scope_mappings=handler.litellm_jwtauth.scope_mappings, scopes=scopes, request_data=request_data, general_settings=general_settings, ) - object_id = jwt_handler.get_object_id(token=jwt_valid_token, default_value=None) + object_id = handler.get_object_id(token=jwt_valid_token, default_value=None) # Get basic user info - user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(jwt_handler, jwt_valid_token) + user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token) # Get IDs - org_id: Final = jwt_handler.get_org_id(token=jwt_valid_token, default_value=None) - end_user_id: Final = jwt_handler.get_end_user_id(token=jwt_valid_token, default_value=None) + org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None) + end_user_id: Final = handler.get_end_user_id(token=jwt_valid_token, default_value=None) team_id: str | None = None team_object: LiteLLM_TeamTable | None = None - object_id = jwt_handler.get_object_id(token=jwt_valid_token, default_value=None) + object_id = handler.get_object_id(token=jwt_valid_token, default_value=None) if rbac_role and object_id: if rbac_role == LitellmUserRoles.TEAM: @@ -2338,12 +2352,12 @@ class JWTAuthManager: user_id = object_id agent_id: Final = JWTAuthManager.resolve_agent_id( - jwt_handler=jwt_handler, + jwt_handler=handler, jwt_valid_token=jwt_valid_token, - agent_registry=jwt_handler.agent_lookup, + agent_registry=handler.agent_lookup, ) - if identity_only: + if identity_only or (not allow_provisioning and handler.is_admin(scopes=scopes)): try: identity_user, _, _, _, identity_user_id = await JWTAuthManager.get_objects( user_id=user_id, @@ -2352,7 +2366,7 @@ class JWTAuthManager: end_user_id=None, team_id=None, valid_user_email=valid_user_email, - jwt_handler=jwt_handler, + jwt_handler=handler, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -2361,13 +2375,37 @@ class JWTAuthManager: user_id_upsert=False, ) except UserNotFoundError: - if not jwt_handler.is_admin(scopes=scopes): + if not handler.is_admin(scopes=scopes): raise identity_user, identity_user_id = None, user_id + if not identity_only: + admin: Final = await JWTAuthManager.check_admin_access( + handler, + scopes, + route, + user_id, + org_id, + api_key, + jwt_valid_token, + user_email=user_email, + agent_id=agent_id, + ) + if admin is not None: + await JWTAuthManager._attach_team_from_header_for_admin( + admin_result=admin, + route=route, + request_headers=request_headers, + jwt_handler=handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + return {**admin, "user_object": identity_user} return JWTAuthBuilderResult( is_proxy_admin=False, # Admin admission uses the claim ID; other callers use the canonical DB ID. - user_id=user_id if jwt_handler.is_admin(scopes=scopes) else identity_user_id, + user_id=user_id if handler.is_admin(scopes=scopes) else identity_user_id, user_email=identity_user.user_email if identity_user is not None else user_email, user_object=identity_user, team_id=None, @@ -2384,7 +2422,7 @@ class JWTAuthManager: # Check admin access admin_result: Final = await JWTAuthManager.check_admin_access( - jwt_handler, + handler, scopes, route, user_id, @@ -2399,7 +2437,7 @@ class JWTAuthManager: admin_result=admin_result, route=route, request_headers=request_headers, - jwt_handler=jwt_handler, + jwt_handler=handler, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -2409,8 +2447,8 @@ class JWTAuthManager: # Get team with model access ## Check if team_id is specified via x-litellm-team-id header - all_team_ids: Final = JWTAuthManager.get_all_team_ids(jwt_handler, jwt_valid_token) - specific_team_id: Final = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None) + all_team_ids: Final = JWTAuthManager.get_all_team_ids(handler, jwt_valid_token) + specific_team_id: Final = handler.get_team_id(token=jwt_valid_token, default_value=None) # The DB fallback only applies when the token carries no team identity at # all. `get_all_jwt_team_ids` ignores `team_id_default` so a configured @@ -2420,9 +2458,9 @@ class JWTAuthManager: # the RBAC team-role path (which already set `team_id`); otherwise a # provisional x-litellm-team-id header could override an RBAC-asserted team. db_team_fallback: Final = ( - jwt_handler.litellm_jwtauth.fallback_to_db_teams - and not jwt_handler.get_all_jwt_team_ids(token=jwt_valid_token) - and not jwt_handler.get_team_alias(token=jwt_valid_token, default_value=None) + handler.litellm_jwtauth.fallback_to_db_teams + and not handler.get_all_jwt_team_ids(token=jwt_valid_token) + and not handler.get_team_alias(token=jwt_valid_token, default_value=None) and team_id is None ) if specific_team_id and not db_team_fallback: @@ -2447,7 +2485,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=(jwt_handler.litellm_jwtauth.team_id_upsert and not db_team_fallback), + team_id_upsert=(handler.litellm_jwtauth.team_id_upsert and not db_team_fallback), ) except HTTPException: if not db_team_fallback: @@ -2459,7 +2497,7 @@ class JWTAuthManager: team_id, team_object, ) = await JWTAuthManager.find_and_validate_specific_team_id( - jwt_handler, + handler, jwt_valid_token, prisma_client, user_api_key_cache, @@ -2474,7 +2512,7 @@ class JWTAuthManager: requested_model=request_data.get("model"), route=route, request_method=request_method, - jwt_handler=jwt_handler, + jwt_handler=handler, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -2498,7 +2536,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=handler.litellm_jwtauth.team_id_upsert, ) if team_id and not JWTAuthManager._team_has_passthrough_route_access( @@ -2509,7 +2547,7 @@ class JWTAuthManager: JWTAuthManager._raise_team_passthrough_route_denial(route=route) # Extract alias fields for resolution (if configured) - org_alias: Final = jwt_handler.get_org_alias(token=jwt_valid_token, default_value=None) + org_alias: Final = handler.get_org_alias(token=jwt_valid_token, default_value=None) # get_objects returns effective_user_id for downstream spend attribution (GH #26789). ( @@ -2525,7 +2563,7 @@ class JWTAuthManager: end_user_id=end_user_id, team_id=team_id, valid_user_email=valid_user_email, - jwt_handler=jwt_handler, + jwt_handler=handler, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -2538,7 +2576,7 @@ class JWTAuthManager: resolved_org_id: Final = org_object.organization_id if org_object else org_id await JWTAuthManager.sync_user_role_and_teams( - jwt_handler=jwt_handler, + jwt_handler=handler, jwt_valid_token=jwt_valid_token, user_object=user_object, prisma_client=prisma_client, @@ -2556,9 +2594,9 @@ class JWTAuthManager: user_id=user_id, requested_model=request_data.get("model"), route=route, - jwt_handler=jwt_handler, - enforce_team_based_model_access=jwt_handler.litellm_jwtauth.enforce_team_based_model_access, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + jwt_handler=handler, + enforce_team_based_model_access=handler.litellm_jwtauth.enforce_team_based_model_access, + team_id_upsert=handler.litellm_jwtauth.team_id_upsert, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -2586,7 +2624,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=handler.litellm_jwtauth.team_id_upsert, ) elif db_team_fallback and team_id == header_team_id: JWTAuthManager._validate_header_team_in_db_membership( @@ -2596,7 +2634,7 @@ class JWTAuthManager: if not JWTAuthManager._is_team_route_allowed( route=route, request_method=request_method, - jwt_handler=jwt_handler, + jwt_handler=handler, ): raise HTTPException( status_code=403, @@ -2606,10 +2644,11 @@ class JWTAuthManager: ) ## MAP USER TO TEAMS - await JWTAuthManager.map_user_to_teams( - user_object=user_object, - team_object=team_object, - ) + if allow_provisioning: + await JWTAuthManager.map_user_to_teams( + user_object=user_object, + team_object=team_object, + ) # Validate that a valid rbac id is returned for spend tracking JWTAuthManager.validate_object_id( @@ -2638,3 +2677,38 @@ class JWTAuthManager: jwt_claims=jwt_valid_token, agent_id=agent_id, ) + + @staticmethod + def user_api_key_auth_from_result( + result: JWTAuthBuilderResult, + parent_otel_span: Span | None = None, + ) -> UserAPIKeyAuth: + """Keep JWT identity and permission attribution identical across consumers.""" + user: Final = result["user_object"] + admin: Final = result["is_proxy_admin"] + return UserAPIKeyAuth( + api_key=None, + user_role=( + LitellmUserRoles.PROXY_ADMIN + if admin + else LitellmUserRoles(user.user_role) + if user is not None and user.user_role is not None + else LitellmUserRoles.INTERNAL_USER + ), + user_id=result["user_id"], + user_email=result["user_email"], + team_id=result["team_id"], + org_id=result["org_id"], + end_user_id=result["end_user_id"], + parent_otel_span=parent_otel_span, + jwt_claims=result["jwt_claims"], + agent_id=result.get("agent_id"), + user_tpm_limit=user.tpm_limit if user is not None and not admin else None, + user_rpm_limit=user.rpm_limit if user is not None and not admin else None, + user_model_max_budget=user.model_max_budget if user is not None and not admin else None, + **team_grants( + team_object=result["team_object"], + team_membership=result.get("team_membership"), + user_id=result["user_id"], + ), + ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index c5297ac83dc..ba267114bac 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1669,13 +1669,11 @@ async def _user_api_key_auth_builder( is_proxy_admin: Final = result["is_proxy_admin"] team_id: Final = result["team_id"] - team_object: Final = result["team_object"] user_id: Final = result["user_id"] user_email: Final = result["user_email"] user_object: Final = result["user_object"] end_user_id = result["end_user_id"] org_id: Final = result["org_id"] - team_membership: Final[LiteLLM_TeamMembership | None] = result.get("team_membership", None) jwt_claims = result.get("jwt_claims", None) agent_id: Final[str | None] = result.get("agent_id") @@ -1693,40 +1691,9 @@ async def _user_api_key_auth_builder( value=_JWT_PROXY_ADMIN_SENTINEL, ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl, ) - return UserAPIKeyAuth( - api_key=None, - user_role=LitellmUserRoles.PROXY_ADMIN, - user_id=user_id, - user_email=user_email, - team_id=team_id, - org_id=org_id, - end_user_id=end_user_id, - parent_otel_span=parent_otel_span, - jwt_claims=jwt_claims, - agent_id=agent_id, - **team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id), - ) + return JWTAuthManager.user_api_key_auth_from_result(result, parent_otel_span) - valid_token = UserAPIKeyAuth( - api_key=None, - team_id=team_id, - user_role=( - LitellmUserRoles(user_object.user_role) - if user_object is not None and user_object.user_role is not None - else LitellmUserRoles.INTERNAL_USER - ), - user_id=user_id, - user_email=user_email, - org_id=org_id, - parent_otel_span=parent_otel_span, - end_user_id=end_user_id, - user_tpm_limit=(user_object.tpm_limit if user_object is not None else None), - user_rpm_limit=(user_object.rpm_limit if user_object is not None else None), - user_model_max_budget=(user_object.model_max_budget if user_object is not None else None), - jwt_claims=jwt_claims, - agent_id=agent_id, - **team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id), - ) + valid_token = JWTAuthManager.user_api_key_auth_from_result(result, parent_otel_span) # AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key. # JWT policy (RBAC, scope, custom_validate, email-domain) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 4c97bbaf5de..918a55bb9ce 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -170,6 +170,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.ui_session_utils import ( admitted_user_context, build_effective_auth_contexts, + can_access_mcp_server, is_ui_session_credential, ) from litellm.proxy._types import ( @@ -2483,10 +2484,11 @@ if MCP_AVAILABLE: ) return server - allowed_server_ids: Final[set[str]] = set() - for auth_context in await build_effective_auth_contexts(user_api_key_dict): - allowed_server_ids.update(await global_mcp_server_manager.get_allowed_mcp_servers(auth_context)) - if server is None or server.server_id not in allowed_server_ids: + if server is None or not await can_access_mcp_server( + user_api_key_dict, + server.server_id, + global_mcp_server_manager.get_allowed_mcp_servers, + ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f6dc7932e2f..8c179eea4cb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11422,6 +11422,7 @@ def _oauth_identity_jwt( issuer: str = "https://idp.example.test", owner: str | None = "jwt-owner", scope: str = "", + claims: dict[str, object] | None = None, ) -> str: import jwt @@ -11434,6 +11435,7 @@ def _oauth_identity_jwt( "aud": audience, "exp": int(time.time()) + expires_in, "scope": scope, + **(claims or {}), }, signing_key, algorithm="RS256", @@ -11443,12 +11445,14 @@ def _oauth_identity_jwt( @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Authorization", "x-litellm-api-key"]) @pytest.mark.parametrize("policy_allowed", [False, True]) +@pytest.mark.parametrize("server_allowed", [False, True]) @pytest.mark.parametrize("admin", [False, True]) @pytest.mark.parametrize("owner_state", ["active", "missing", "inactive", "database_error"]) async def test_oauth_exchange_stores_token_for_validated_jwt_user( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], header: str, policy_allowed: bool, + server_allowed: bool, admin: bool, owner_state: str, monkeypatch: pytest.MonkeyPatch, @@ -11461,6 +11465,12 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( handler, signing_key = jwt_oauth_identity handler.litellm_jwtauth.custom_validate = lambda claims: policy_allowed + from litellm.proxy._experimental.mcp_server import mcp_server_manager + + manager: Final = MagicMock() + manager.get_allowed_mcp_servers = AsyncMock(return_value=["jwt-oauth-server"] if server_allowed else []) + manager.invalidate_user_oauth_token_cache = AsyncMock() + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) bearer: Final = _oauth_identity_jwt(signing_key, scope="litellm_proxy_admin" if admin else "") request: Final = _token_request({header: f"Bearer {bearer}"}, path="/jwt-oauth-server/token") server: Final = MCPServer( @@ -11525,7 +11535,12 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( assert response.status_code == 200 assert json.loads(response.body)["access_token"] == "upstream-token" users.create.assert_not_awaited() - if not policy_allowed or owner_state in ("inactive", "database_error") or (owner_state == "missing" and not admin): + if ( + not server_allowed + or not policy_allowed + or owner_state in ("inactive", "database_error") + or (owner_state == "missing" and not admin) + ): table.upsert.assert_not_awaited() return table.upsert.assert_awaited_once() @@ -11755,9 +11770,7 @@ async def test_oauth_jwt_resolves_canonical_owner_without_cached_identity( table.find_first = AsyncMock(return_value=owner) table.update = AsyncMock(return_value=owner) monkeypatch.setattr(proxy_server, "prisma_client", database) - bearer: Final = _oauth_identity_jwt( - signing_key, owner=external_id, scope="litellm_proxy_admin" if admin else "" - ) + bearer: Final = _oauth_identity_jwt(signing_key, owner=external_id, scope="litellm_proxy_admin" if admin else "") request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) stored_owner: Final = await _extract_user_id_from_request(request) assert stored_owner == (None if inactive else external_id if admin else "canonical-oauth-owner") @@ -11823,3 +11836,139 @@ async def test_oauth_refresh_revalidates_the_same_active_user_rule( monkeypatch.setattr(proxy_server, "prisma_client", None) expected: Final = None if state == "active" else "no_active_key" if state == "inactive" else "unresolvable" assert await _reload_active_user_by_id("jwt-owner") == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mapped", [False, True]) +@pytest.mark.parametrize("state", ["allowed", "route_denied", "server_denied", "blocked", "expired", "lookup_error", "cancelled"]) +async def test_oauth_credential_write_keeps_virtual_key_permissions( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + mapped: bool, + state: str, +) -> None: + import asyncio + + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._types import UserAPIKeyAuth, hash_token + from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key + + handler, signing_key = jwt_oauth_identity + key: Final = "sk-oauth-permission-test" + hashed: Final = hash_token(key) + credential: Final = UserAPIKeyAuth( + token=hashed, + user_id="jwt-owner", + blocked=state == "blocked", + expires=datetime.now(timezone.utc) - timedelta(seconds=60) if state == "expired" else None, + allowed_routes=["openai_routes"] if state == "route_denied" else ["mcp_routes"], + agent_id="agent-scope", + org_id="org-scope", + end_user_id="end-user-scope", + ) + handler.user_api_key_cache.set_cache(hashed, credential) + if mapped: + handler.litellm_jwtauth.virtual_key_claim_field = "sub" + handler.user_api_key_cache.set_cache(jwt_key_mapping_cache_key("sub", "not-the-configured-user-id"), hashed) + manager: Final = MagicMock() + manager.get_allowed_mcp_servers = AsyncMock( + return_value=[] if state == "server_denied" else ["server-a"], + side_effect=(asyncio.CancelledError() if state == "cancelled" else RuntimeError("permission lookup unavailable") if state == "lookup_error" else None), + ) + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + bearer: Final = _oauth_identity_jwt(signing_key) if mapped else key + request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/server-a/token") + if state == "cancelled": + with pytest.raises(asyncio.CancelledError): + await _extract_user_id_from_request(request, "server-a") + manager.get_allowed_mcp_servers.assert_awaited_once() + return + assert await _extract_user_id_from_request(request, "server-a") == ("jwt-owner" if state == "allowed" else None) + if state in ("allowed", "server_denied", "lookup_error"): + manager.get_allowed_mcp_servers.assert_awaited_once() + writer: Final = manager.get_allowed_mcp_servers.call_args.args[0] + assert (writer.user_id, writer.token, writer.org_id, writer.agent_id, writer.end_user_id) == ( + "jwt-owner", + hashed, + "org-scope", + "agent-scope", + "end-user-scope", + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("server_id", ["team-a-server", "team-b-server"]) +async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_roster( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + server_id: str, +) -> None: + from litellm.models.user import LiteLLM_UserTable + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._types import LiteLLM_TeamTable, Member + + handler, signing_key = jwt_oauth_identity + handler.litellm_jwtauth.team_id_jwt_field = "team" + handler.litellm_jwtauth.team_id_upsert = True + handler.litellm_jwtauth.user_id_upsert = True + handler.litellm_jwtauth.sync_user_role_and_teams = True + handler.user_api_key_cache.set_cache("jwt-owner", LiteLLM_UserTable(user_id="jwt-owner", teams=["a", "b"])) + handler.user_api_key_cache.set_cache( + "team_id:a", + LiteLLM_TeamTable(team_id="a", models=[], members_with_roles=[Member(user_id="jwt-owner", role="user")]), + ) + manager: Final = MagicMock() + manager.get_allowed_mcp_servers = AsyncMock(return_value=["team-a-server"]) + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + bearer: Final = _oauth_identity_jwt(signing_key, claims={"team": "a"}) + request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path=f"/{server_id}/token") + assert await _extract_user_id_from_request(request, server_id) == ( + "jwt-owner" if server_id == "team-a-server" else None + ) + manager.get_allowed_mcp_servers.assert_awaited_once() + writer: Final = manager.get_allowed_mcp_servers.call_args.args[0] + assert writer.team_id == "a" + assert not writer.mcp_admitted_user_subject + proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called() + proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() + proxy_server.prisma_client.db.litellm_usertable.update.assert_not_called() + assert handler.litellm_jwtauth.user_id_upsert and handler.litellm_jwtauth.team_id_upsert + assert handler.litellm_jwtauth.sync_user_role_and_teams + + +@pytest.mark.asyncio +async def test_oauth_write_denial_does_not_erase_identity_binding( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy._experimental.mcp_server import discoverable_endpoints, mcp_server_manager + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + _, signing_key = jwt_oauth_identity + monkeypatch.setenv("LITELLM_SALT_KEY", "oauth-identity-binding-test-salt") + manager: Final = MagicMock() + manager.get_allowed_mcp_servers = AsyncMock(return_value=[]) + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + server: Final = MCPServer( + server_id="bound-server", name="bound-server", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client", + token_url="https://upstream.example.test/token", + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", issuer="https://upstream.example.test", audiences=["client"], + ), + ) + request: Final = _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}) + code: Final = discoverable_endpoints.seal_bridge_authorization_code( + "upstream-code", "another-owner", server.server_id, "bound-nonce", + ) + with pytest.raises(HTTPException) as denied: + await discoverable_endpoints.exchange_token_with_server( + request=request, mcp_server=server, grant_type="authorization_code", code=code, + redirect_uri="http://localhost/callback", client_id="client", client_secret=None, code_verifier="verifier", + ) + assert denied.value.status_code == 403 + assert denied.value.detail == {"error": "oauth_principal_mismatch"} + manager.get_allowed_mcp_servers.assert_not_awaited() diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 4d86e395358..babbf88dc29 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -6792,10 +6792,11 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla @pytest.mark.asyncio @pytest.mark.parametrize("identity_only", [False, True]) +@pytest.mark.parametrize("allow_provisioning", [False, True]) @pytest.mark.parametrize("existing_user", [False, True]) @pytest.mark.parametrize("model_allowed", [False, True]) async def test_auth_builder_identity_lookup_does_not_provision_users( - monkeypatch: pytest.MonkeyPatch, identity_only: bool, existing_user: bool, model_allowed: bool + monkeypatch: pytest.MonkeyPatch, identity_only: bool, allow_provisioning: bool, existing_user: bool, model_allowed: bool ) -> None: from litellm.proxy._types import ScopeMapping from litellm.proxy.auth.auth_checks import UserNotFoundError @@ -6841,6 +6842,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( parent_otel_span=None, proxy_logging_obj=MagicMock(), identity_only=identity_only, + allow_provisioning=allow_provisioning, ) if not identity_only and not model_allowed: with pytest.raises(HTTPException) as denial: @@ -6848,7 +6850,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( assert denial.value.status_code == 403 users.create.assert_not_awaited() return - if identity_only and not existing_user: + if (identity_only or not allow_provisioning) and not existing_user: with pytest.raises(UserNotFoundError): await pending else: @@ -6856,7 +6858,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( assert result["user_id"] == user_id assert result["user_object"] is not None assert result["user_object"].user_id == user_id - assert users.create.await_count == (0 if identity_only or existing_user else 1) + assert users.create.await_count == (0 if identity_only or not allow_provisioning or existing_user else 1) def _entra_agent_registry() -> AgentRegistry: From c7e4160ee6ce40493e5e3d896a30ac1668a55c70 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 19:15:34 -0700 Subject: [PATCH 7/9] fix(mcp): enforce OAuth write policy across signed callbacks --- .../mcp_server/bridge_token_flow.py | 54 ++++---- .../mcp_server/discoverable_endpoints.py | 58 +++++---- .../test_user_api_key_auth.py | 1 + .../mcp_server/test_discoverable_endpoints.py | 118 +++++++++++++++++- 4 files changed, 181 insertions(+), 50 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 962e39d7dd6..5fe0929773c 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -306,18 +306,8 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol async def _extract_user_id_from_request(request: Request, server_id: str | None = None) -> str | None: """Resolve identity for binding, or authorize the credential-write action for a target server.""" - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # registry imports auth helpers - global_mcp_server_manager, - ) - from litellm.proxy._experimental.mcp_server.ui_session_utils import ( - can_access_mcp_server, # noqa: PLC0415 # proxy import cycle - ) from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle - from litellm.proxy.auth.route_checks import RouteChecks # noqa: PLC0415 # proxy import cycle - from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle - _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action - ) token: Final = _litellm_key_from_request(request) # The OAuth relay is public; the optional server-side write is the same protected action @@ -331,23 +321,39 @@ async def _extract_user_id_from_request(request: Request, server_id: str | None auth: Final = resolved.key if isinstance(resolved, _ResolvedKey) else resolved if not isinstance(auth, UserAPIKeyAuth) or not _active_key_user_id(auth): return None - if write_route is not None and server_id is not None: - try: - RouteChecks.is_virtual_key_allowed_to_call_route(route=write_route, valid_token=auth, request=request) - await _run_centralized_common_checks( - user_api_key_auth_obj=auth, - request=request, - request_data={}, - route=write_route, - ) - if not await can_access_mcp_server(auth, server_id, global_mcp_server_manager.get_allowed_mcp_servers): - return None - except Exception as exc: # noqa: BLE001 # authorization failure must never write credentials - verbose_logger.debug("OAuth credential write not authorized (%s)", type(exc).__name__) - return None + if server_id is not None and not await can_store_oauth_credential(request, auth, server_id): + return None return auth.user_id +async def can_store_oauth_credential(request: Request, auth: "UserAPIKeyAuth", server_id: str) -> bool: + """Apply the same write policy to request credentials and verified signed-callback users.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # registry imports auth helpers + global_mcp_server_manager, + ) + from litellm.proxy._experimental.mcp_server.ui_session_utils import ( + can_access_mcp_server, # noqa: PLC0415 # proxy import cycle + ) + from litellm.proxy.auth.route_checks import RouteChecks # noqa: PLC0415 # proxy import cycle + from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle + _run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action + ) + + write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential" + try: + RouteChecks.is_virtual_key_allowed_to_call_route(route=write_route, valid_token=auth, request=request) + await _run_centralized_common_checks( + user_api_key_auth_obj=auth, + request=request, + request_data={}, + route=write_route, + ) + return await can_access_mcp_server(auth, server_id, global_mcp_server_manager.get_allowed_mcp_servers) + except Exception as exc: # noqa: BLE001 # authorization failure must never write credentials + verbose_logger.debug("OAuth credential write not authorized (%s)", type(exc).__name__) + return False + + async def _resolve_jwt_auth( request: Request, token: str, diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 56968745ea9..9ba67f966a2 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -29,9 +29,11 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _BridgeRefreshReady, _extract_user_id_from_request, _finish_bridge_mint, + _litellm_key_from_request, # pyright: ignore[reportPrivateUsage] # shared credential precedence for authorization issuance _prepare_bridge_mint, _prepare_bridge_refresh, _reload_active_user_by_id, + can_store_oauth_credential, ) from litellm.proxy._experimental.mcp_server.faults import ( CallerRejected, @@ -836,16 +838,29 @@ async def _user_can_reach_mcp_server(user_id: str, server_id: str) -> bool: return server_id in await global_mcp_server_manager.get_allowed_mcp_servers(admitted) -async def _bridge_authorize_access_denial( - litellm_user_id: str, +async def _resolve_oauth_authorization_user( + request: Request, mcp_server: MCPServer, redirect_uri: str, state: str, -) -> RedirectResponse | None: - """The denial redirect for a signed-in user who cannot reach the target server, or None to proceed.""" - if await _user_can_reach_mcp_server(litellm_user_id, mcp_server.server_id): - return None - return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) + enforce_binding: bool, +) -> str | RedirectResponse: + """Resolve the authorization subject without replacing denied credentials with cookie grants.""" + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # proxy import cycle + _user_id_from_session_cookie, + ) + + request_user_id: Final = ( + await _extract_user_id_from_request(request, mcp_server.server_id) if enforce_binding else None + ) + if enforce_binding and request_user_id is None and _litellm_key_from_request(request): + return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) + user_id: Final = request_user_id or _user_id_from_session_cookie(request) + if user_id is None: + return _redirect_to_litellm_login(request) + if not await _user_can_reach_mcp_server(user_id, mcp_server.server_id): + return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) + return user_id async def authorize_with_server( @@ -911,23 +926,12 @@ async def authorize_with_server( # Seal the authenticated caller into state so the token exchange cannot select another credential owner. litellm_user_id: str | None = None if enforce_binding or (resolved_server.is_dcr_bridge and resolved_server.is_oauth_delegate): - from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # inline import avoids a module-load circular import - _user_id_from_session_cookie, + subject: Final = await _resolve_oauth_authorization_user( + request, resolved_server, redirect_uri, state, enforce_binding ) - - litellm_user_id = ( - await _extract_user_id_from_request(request) if enforce_binding else None - ) or _user_id_from_session_cookie(request) - if litellm_user_id is None: - return _redirect_to_litellm_login(request) - denial: Final = await _bridge_authorize_access_denial( - litellm_user_id=litellm_user_id, - mcp_server=resolved_server, - redirect_uri=redirect_uri, - state=state, - ) - if denial is not None: - return denial + if isinstance(subject, RedirectResponse): + return subject + litellm_user_id = subject oauth_nonce: Final = secrets.token_urlsafe(32) if enforce_binding else None encoded_state: Final = encode_state_with_base_url( @@ -1220,8 +1224,14 @@ async def exchange_token_with_server( try: # Identity binding above must retain the verified caller even when a write is # denied. Authorize persistence separately, immediately before its side effect. + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + # A sealed code delegates a verified user for this authorized server. Raw + # request credentials retain their own JWT/key restrictions during resolution. can_store: Final = ( - await _user_can_reach_mcp_server(user_id, resolved_server.server_id) + await can_store_oauth_credential( + request, await MCPRequestHandler.reload_admitted_user(user_id), resolved_server.server_id + ) if bridge_identity is not None else await _extract_user_id_from_request(request, resolved_server.server_id) == user_id ) diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index 0cdf3500d50..a8fce58c60b 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -1069,6 +1069,7 @@ async def test_jwt_non_admin_team_route_access(monkeypatch): mock_jwt_response = { "is_proxy_admin": False, + "jwt_claims": {}, "team_id": None, "team_object": None, "user_id": None, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 8c179eea4cb..1739ac6d743 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11171,8 +11171,8 @@ async def test_identity_bound_authorization_carries_nonce_and_caller_through_cal "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", new=AsyncMock(return_value="alice")), patch( # test-quality-ok: isolate user access lookup while testing nonce and caller preservation - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._bridge_authorize_access_denial", - new=AsyncMock(return_value=None)), + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._user_can_reach_mcp_server", + new=AsyncMock(return_value=True)), ): authorized = await authorize_with_server( request, server, "client", "http://127.0.0.1:6274/callback", state="client-state", @@ -11972,3 +11972,117 @@ async def test_oauth_write_denial_does_not_erase_identity_binding( assert denied.value.status_code == 403 assert denied.value.detail == {"error": "oauth_principal_mismatch"} manager.get_allowed_mcp_servers.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("admin_only", [False, True]) +async def test_signed_oauth_callback_honors_credential_write_policy( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + admin_only: bool, +) -> None: + import httpx + 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._experimental.mcp_server import discoverable_endpoints, mcp_server_manager + from litellm.proxy._types import MCPTransport + from litellm.types.llms.custom_http import httpxSpecialProvider + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server: Final = MCPServer( + server_id="signed-server", name="signed-server", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client", + token_url="https://upstream.example.test/token", + ) + monkeypatch.setattr(proxy_server, "general_settings", { + "enable_jwt_auth": True, + "admin_only_routes": [f"/v1/mcp/server/{server.server_id}/oauth-user-credential"] if admin_only else [], + }) + monkeypatch.setenv("LITELLM_SALT_KEY", "signed-oauth-test-salt") + manager: Final = MagicMock() + manager.get_allowed_mcp_servers = AsyncMock(return_value=[server.server_id]) + manager.invalidate_user_oauth_token_cache = AsyncMock() + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + table: Final = proxy_server.prisma_client.db.litellm_mcpusercredentials + table.find_unique = AsyncMock(return_value=None) + table.upsert = AsyncMock() + clients: Final = LLMClientCache() + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", clients) + + def upstream_response(outbound: httpx.Request) -> httpx.Response: + assert outbound.url == server.token_url + assert b"code=upstream-code" in outbound.content + return httpx.Response(200, json={"access_token": "upstream-token", "token_type": "Bearer"}) + + 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=_token_request({}, path="/signed-server/token"), mcp_server=server, + grant_type="authorization_code", + code=discoverable_endpoints.seal_bridge_authorization_code("upstream-code", "jwt-owner", server.server_id), + redirect_uri="http://localhost/callback", client_id="client", client_secret=None, code_verifier=None, + ) + assert response.status_code == 200 + assert json.loads(response.body)["access_token"] == "upstream-token" + if admin_only: + table.upsert.assert_not_awaited() + else: + table.upsert.assert_awaited_once() + assert table.upsert.call_args.kwargs["where"]["user_id_server_id"] == { + "user_id": "jwt-owner", "server_id": server.server_id, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("allowed", [False, True]) +async def test_identity_bound_authorize_preserves_presented_jwt_permissions( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + allowed: bool, +) -> None: + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server import byok_oauth_endpoints, discoverable_endpoints, mcp_server_manager + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + _, signing_key = jwt_oauth_identity + monkeypatch.setenv("LITELLM_SALT_KEY", "authorize-policy-test-salt") + server: Final = MCPServer( + server_id="bound-server", name="bound-server", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client", + authorization_url="https://upstream.example.test/authorize", token_url="https://upstream.example.test/token", + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", issuer="https://upstream.example.test", audiences=["client"], + ), + ) + manager: Final = MagicMock() + # The full user roster permits the server; the presented JWT may have narrower access. + manager.get_allowed_mcp_servers = AsyncMock( + side_effect=lambda auth: [server.server_id] if allowed or auth.mcp_admitted_user_subject else [], + ) + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + monkeypatch.setattr( # test-quality-ok: session-cookie decoder is the separate authentication boundary; a valid cookie must not override a denied explicit credential + byok_oauth_endpoints, "_user_id_from_session_cookie", lambda request: "jwt-owner", + ) + response: Final = await discoverable_endpoints.authorize_with_server( + request=_token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}), + mcp_server=server, client_id="client", redirect_uri="http://127.0.0.1:6274/callback", + state="client-state", code_challenge="pkce-challenge", code_challenge_method="S256", + ) + redirect: Final = urlparse(response.headers["location"]) + query: Final = parse_qs(redirect.query) + if allowed: + assert redirect.hostname == "upstream.example.test" + assert query["nonce"] and response.headers.get("set-cookie") + else: + assert redirect.hostname == "127.0.0.1" + assert query["error"] == ["access_denied"] + assert query["state"] == ["client-state"] + assert "set-cookie" not in response.headers From e035682ed17295b9c5f0a363e266e890de31061c Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 22:15:57 -0700 Subject: [PATCH 8/9] refactor(auth): separate JWT identity and OAuth authorization --- .../mcp_server/bridge_token_flow.py | 68 +++-- .../mcp_server/discoverable_endpoints.py | 5 +- litellm/proxy/auth/handle_jwt.py | 276 ++++++++++-------- .../mcp_server/test_discoverable_endpoints.py | 34 ++- .../proxy/auth/test_handle_jwt.py | 100 +++++-- 5 files changed, 300 insertions(+), 183 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index 5fe0929773c..f001ee87dbd 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -24,6 +24,7 @@ if TYPE_CHECKING: UpstreamTokenGrant, ) from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.handle_jwt import JWTIdentity def _litellm_key_from_request(request: Request) -> str | None: @@ -304,26 +305,37 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol assert_never(identity.subject_type) -async def _extract_user_id_from_request(request: Request, server_id: str | None = None) -> str | None: - """Resolve identity for binding, or authorize the credential-write action for a target server.""" +async def _extract_user_id_from_request(request: Request) -> str | None: + """Resolve the caller for identity binding without granting credential-write permission.""" + from litellm.proxy.auth.handle_jwt import JWTIdentity # noqa: PLC0415 # proxy import cycle + + resolved: Final = await _resolve_request_auth(request) + if isinstance(resolved, JWTIdentity): + return resolved.user_id + return _active_key_user_id(resolved) if resolved is not None else None + + +async def authorize_oauth_credential_request(request: Request, server_id: str) -> str | None: from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle + + resolved: Final = await _resolve_request_auth(request, f"/v1/mcp/server/{server_id}/oauth-user-credential") + if not isinstance(resolved, UserAPIKeyAuth) or not _active_key_user_id(resolved): + return None + if not await can_store_oauth_credential(request, resolved, server_id): + return None + return resolved.user_id + + +async def _resolve_request_auth( + request: Request, write_route: str | None = None +) -> "UserAPIKeyAuth | JWTIdentity | None": from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle token: Final = _litellm_key_from_request(request) - # The OAuth relay is public; the optional server-side write is the same protected action - # as the direct credential endpoint. Authorize that action without rewriting the Request. - write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential" if server_id is not None else None - resolved: Final = ( - await _resolve_jwt_auth(request, token, write_route) - if token is not None and JWTHandler.is_jwt(token) - else await _resolve_active_litellm_key(request) - ) - auth: Final = resolved.key if isinstance(resolved, _ResolvedKey) else resolved - if not isinstance(auth, UserAPIKeyAuth) or not _active_key_user_id(auth): - return None - if server_id is not None and not await can_store_oauth_credential(request, auth, server_id): - return None - return auth.user_id + if token is not None and JWTHandler.is_jwt(token): + return await _resolve_jwt_auth(request, token, write_route) + resolved: Final = await _resolve_active_litellm_key(request) + return resolved.key if isinstance(resolved, _ResolvedKey) else None async def can_store_oauth_credential(request: Request, auth: "UserAPIKeyAuth", server_id: str) -> bool: @@ -358,7 +370,7 @@ async def _resolve_jwt_auth( request: Request, token: str, write_route: str | None, -) -> "UserAPIKeyAuth | None": +) -> "UserAPIKeyAuth | JWTIdentity | None": from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.handle_jwt import JWTAuthManager # noqa: PLC0415 # proxy import cycle from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle @@ -393,25 +405,35 @@ async def _resolve_jwt_auth( return None if await _key_owner_scim_deactivated(mapped) or not _active_key_user_id(mapped) else mapped if mapped is not None: return None - identity: Final = await JWTAuthManager.auth_builder( + if write_route is None: + identity: Final = await JWTAuthManager.resolve_identity( + api_key=token, + 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 identity.user_object is not None and isinstance(_active_user_record(identity.user_object), str): + return None + return identity + authorized: Final = await JWTAuthManager.authorize_jwt( api_key=token, jwt_handler=jwt_handler, request_data={}, general_settings=general_settings, - route=write_route or request.url.path, + route=write_route, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=None, proxy_logging_obj=proxy_logging_obj, request_headers=dict(request.headers), request_method=request.method, - identity_only=write_route is None, - allow_provisioning=False, ) - resolved_user: Final = identity["user_object"] + resolved_user: Final = authorized["user_object"] if resolved_user is not None and isinstance(_active_user_record(resolved_user), str): return None - return JWTAuthManager.user_api_key_auth_from_result(identity) + return JWTAuthManager.user_api_key_auth_from_result(authorized) 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 diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 9ba67f966a2..7e7189c8a6b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _prepare_bridge_mint, _prepare_bridge_refresh, _reload_active_user_by_id, + authorize_oauth_credential_request, can_store_oauth_credential, ) from litellm.proxy._experimental.mcp_server.faults import ( @@ -851,7 +852,7 @@ async def _resolve_oauth_authorization_user( ) request_user_id: Final = ( - await _extract_user_id_from_request(request, mcp_server.server_id) if enforce_binding else None + await authorize_oauth_credential_request(request, mcp_server.server_id) if enforce_binding else None ) if enforce_binding and request_user_id is None and _litellm_key_from_request(request): return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) @@ -1233,7 +1234,7 @@ async def exchange_token_with_server( request, await MCPRequestHandler.reload_admitted_user(user_id), resolved_server.server_id ) if bridge_identity is not None - else await _extract_user_id_from_request(request, resolved_server.server_id) == user_id + else await authorize_oauth_credential_request(request, resolved_server.server_id) == user_id ) if can_store: await _store_per_user_token_server_side( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index f2bdbdd9341..4fe44eb1dc8 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -9,13 +9,13 @@ JWT token must have 'litellm_proxy_admin' in scope. from __future__ import annotations import asyncio -import copy import fnmatch import hashlib import os import re import time from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast import httpx @@ -130,6 +130,19 @@ class _UserInfoResponse(Protocol): def json(self) -> dict[str, object]: ... +@dataclass(frozen=True, slots=True) +class JWTIdentity: + user_id: str | None + user_object: LiteLLM_UserTable | None + agent_id: str | None + + +@dataclass(frozen=True, slots=True) +class _JWTProvisioning: + user_id_upsert: bool + team_id_upsert: bool + + class AgentLookup(Protocol): """The registered-agent lookups a JWT agent claim is matched against.""" @@ -1473,6 +1486,7 @@ class JWTAuthManager: user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, + team_id_upsert: bool | None = None, ) -> tuple[str | None, LiteLLM_TeamTable | None]: """Find and validate specific team ID from team_id_jwt_field or team_alias_jwt_field""" individual_team_id = jwt_handler.get_team_id(token=jwt_valid_token, default_value=None) @@ -1500,7 +1514,9 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert + if team_id_upsert is None + else team_id_upsert, ) return individual_team_id, team_object except HTTPException as e: @@ -2017,6 +2033,7 @@ class JWTAuthManager: user_api_key_cache: UserApiKeyCache, parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, + team_id_upsert: bool | None = None, ) -> None: """Attach team context from x-litellm-team-id to an admin result. @@ -2034,7 +2051,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert if team_id_upsert is None else team_id_upsert, ) except Exception as e: # Fall back to pre-PR admin behavior: honor the admin's @@ -2268,64 +2285,119 @@ class JWTAuthManager: proxy_logging_obj: ProxyLogging, request_headers: dict | None = None, request_method: str | None = None, - identity_only: bool = False, - allow_provisioning: bool = True, ) -> JWTAuthBuilderResult: - """Build JWT authentication and authorization context. + return await JWTAuthManager.authorize_jwt( + api_key=api_key, + jwt_handler=jwt_handler, + request_data=request_data, + general_settings=general_settings, + route=route, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + request_headers=request_headers, + request_method=request_method, + provisioning=_JWTProvisioning( + user_id_upsert=jwt_handler.litellm_jwtauth.user_id_upsert, + team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + ), + ) - identity_only resolves the caller for OAuth identity binding and grants no permission. - Credential writes use full authorization with allow_provisioning=False: resolve the - existing policy context without creating users/teams or synchronizing membership. - A private handler configuration keeps that restriction out of concurrent normal requests. - """ - handler: Final = jwt_handler if allow_provisioning else copy.copy(jwt_handler) - if not allow_provisioning: - handler.update_environment( - prisma_client=jwt_handler.prisma_client, - user_api_key_cache=jwt_handler.user_api_key_cache, - litellm_jwtauth=jwt_handler.litellm_jwtauth.model_copy( - update={"user_id_upsert": False, "team_id_upsert": False, "sync_user_role_and_teams": False} - ), - leeway=jwt_handler.leeway, + @staticmethod + async def authenticate_jwt(api_key: str, jwt_handler: JWTHandler) -> dict[str, object]: + claims: Final = ( + await jwt_handler.get_oidc_userinfo(token=api_key) + if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not jwt_handler.is_jwt(token=api_key) + else await jwt_handler.auth_jwt(token=api_key) + ) + validate: Final = jwt_handler.litellm_jwtauth.custom_validate + if validate is not None and not validate(claims): + raise HTTPException(status_code=403, detail="Invalid JWT token") + return claims + + @staticmethod + async def resolve_identity( + api_key: str, + jwt_handler: JWTHandler, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, + ) -> JWTIdentity: + claims: Final = await JWTAuthManager.authenticate_jwt(api_key, jwt_handler) + return await JWTAuthManager._resolve_claim_identity( + claims, jwt_handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj + ) + + @staticmethod + async def _resolve_claim_identity( + claims: dict[str, object], + jwt_handler: JWTHandler, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, + ) -> JWTIdentity: + claim_user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(jwt_handler, claims) + user_id: Final = ( + jwt_handler.get_object_id(token=claims, default_value=None) or claim_user_id + if jwt_handler.get_rbac_role(token=claims) == LitellmUserRoles.INTERNAL_USER + else claim_user_id + ) + agent_id: Final = JWTAuthManager.resolve_agent_id(jwt_handler, claims, jwt_handler.agent_lookup) + is_admin: Final = jwt_handler.is_admin(scopes=jwt_handler.get_scopes(token=claims)) + try: + user, _, _, _, canonical_id = await JWTAuthManager.get_objects( + user_id=user_id, + user_email=user_email, + org_id=None, + end_user_id=None, + team_id=None, + valid_user_email=valid_user_email, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + route="", + user_id_upsert=False, ) + except UserNotFoundError: + if not is_admin: + raise + return JWTIdentity(user_id=user_id, user_object=None, agent_id=agent_id) + return JWTIdentity(user_id=user_id if is_admin else canonical_id, user_object=user, agent_id=agent_id) - # Check if OIDC UserInfo endpoint is enabled, but fall back to standard - # JWT auth if the token itself is a well-formed JWT (3-part structure). - if handler.litellm_jwtauth.oidc_userinfo_enabled and not handler.is_jwt(token=api_key): - verbose_proxy_logger.debug("OIDC UserInfo is enabled. Fetching user info from UserInfo endpoint.") - # Use the access token to fetch user info from OIDC UserInfo endpoint - jwt_valid_token: dict = await handler.get_oidc_userinfo(token=api_key) - else: - # Default behavior: decode and validate the JWT token - jwt_valid_token = await handler.auth_jwt(token=api_key) - - # Check custom validate - if handler.litellm_jwtauth.custom_validate: - if not handler.litellm_jwtauth.custom_validate(jwt_valid_token): - raise HTTPException( - status_code=403, - detail="Invalid JWT token", - ) + @staticmethod + async def authorize_jwt( + api_key: str, + jwt_handler: JWTHandler, + request_data: dict[str, object], + general_settings: dict[str, object], + route: str, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, + request_headers: dict[str, str] | None = None, + request_method: str | None = None, + provisioning: _JWTProvisioning | None = None, + ) -> JWTAuthBuilderResult: + """Resolve and authorize JWT context; only normal admission supplies provisioning.""" + handler: Final = jwt_handler + jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler) + team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False + model: Final = request_data.get("model") + requested_model: Final = model if isinstance(model, str) else None # Check RBAC rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) - if not identity_only: - await JWTAuthManager.check_rbac_role( - handler, - jwt_valid_token, - general_settings, - request_data, - route, - rbac_role, - ) + await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role) # Check Scope Based Access scopes: Final = handler.get_scopes(token=jwt_valid_token) - if ( - not identity_only - and handler.litellm_jwtauth.enforce_scope_based_access - and handler.litellm_jwtauth.scope_mappings - ): + if handler.litellm_jwtauth.enforce_scope_based_access and handler.litellm_jwtauth.scope_mappings: JWTAuthManager.check_scope_based_access( scope_mappings=handler.litellm_jwtauth.scope_mappings, scopes=scopes, @@ -2357,69 +2429,6 @@ class JWTAuthManager: agent_registry=handler.agent_lookup, ) - if identity_only or (not allow_provisioning and handler.is_admin(scopes=scopes)): - try: - identity_user, _, _, _, identity_user_id = await JWTAuthManager.get_objects( - user_id=user_id, - user_email=user_email, - org_id=None, - end_user_id=None, - team_id=None, - valid_user_email=valid_user_email, - jwt_handler=handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - route=route, - user_id_upsert=False, - ) - except UserNotFoundError: - if not handler.is_admin(scopes=scopes): - raise - identity_user, identity_user_id = None, user_id - if not identity_only: - admin: Final = await JWTAuthManager.check_admin_access( - handler, - scopes, - route, - user_id, - org_id, - api_key, - jwt_valid_token, - user_email=user_email, - agent_id=agent_id, - ) - if admin is not None: - await JWTAuthManager._attach_team_from_header_for_admin( - admin_result=admin, - route=route, - request_headers=request_headers, - jwt_handler=handler, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - return {**admin, "user_object": identity_user} - return JWTAuthBuilderResult( - is_proxy_admin=False, - # Admin admission uses the claim ID; other callers use the canonical DB ID. - user_id=user_id if handler.is_admin(scopes=scopes) else identity_user_id, - user_email=identity_user.user_email if identity_user is not None else user_email, - user_object=identity_user, - team_id=None, - team_object=None, - org_id=None, - org_object=None, - end_user_id=None, - end_user_object=None, - team_membership=None, - token=api_key, - jwt_claims=jwt_valid_token, - agent_id=agent_id, - ) - # Check admin access admin_result: Final = await JWTAuthManager.check_admin_access( handler, @@ -2442,7 +2451,13 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, + team_id_upsert=team_id_upsert, ) + if provisioning is None: + identity: Final = await JWTAuthManager._resolve_claim_identity( + jwt_valid_token, handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj + ) + return {**admin_result, "user_object": identity.user_object} return admin_result # Get team with model access @@ -2485,7 +2500,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=(handler.litellm_jwtauth.team_id_upsert and not db_team_fallback), + team_id_upsert=(team_id_upsert and not db_team_fallback), ) except HTTPException: if not db_team_fallback: @@ -2503,13 +2518,14 @@ class JWTAuthManager: user_api_key_cache, parent_otel_span, proxy_logging_obj, + team_id_upsert=team_id_upsert, ) if not team_object and not team_id: ## CHECK USER GROUP ACCESS team_id, team_object = await JWTAuthManager.find_team_with_model_access( team_ids=all_team_ids, - requested_model=request_data.get("model"), + requested_model=requested_model, route=route, request_method=request_method, jwt_handler=handler, @@ -2536,7 +2552,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=team_id_upsert, ) if team_id and not JWTAuthManager._team_has_passthrough_route_access( @@ -2570,18 +2586,20 @@ class JWTAuthManager: proxy_logging_obj=proxy_logging_obj, route=route, org_alias=org_alias, + user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False, ) # Derive org_id from org_object if resolved by alias resolved_org_id: Final = org_object.organization_id if org_object else org_id - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler=handler, - jwt_valid_token=jwt_valid_token, - user_object=user_object, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - ) + if provisioning is not None: + await JWTAuthManager.sync_user_role_and_teams( + jwt_handler=handler, + jwt_valid_token=jwt_valid_token, + user_object=user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) # If JWT did not resolve team_id, attempt a team fallback. if team_id is None and db_team_fallback: @@ -2592,11 +2610,11 @@ class JWTAuthManager: ) = await JWTAuthManager._resolve_db_team_fallback( user_object=user_object, user_id=user_id, - requested_model=request_data.get("model"), + requested_model=requested_model, route=route, jwt_handler=handler, enforce_team_based_model_access=handler.litellm_jwtauth.enforce_team_based_model_access, - team_id_upsert=handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=team_id_upsert, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, @@ -2624,7 +2642,7 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=handler.litellm_jwtauth.team_id_upsert, + team_id_upsert=team_id_upsert, ) elif db_team_fallback and team_id == header_team_id: JWTAuthManager._validate_header_team_in_db_membership( @@ -2644,7 +2662,7 @@ class JWTAuthManager: ) ## MAP USER TO TEAMS - if allow_provisioning: + if provisioning is not None: await JWTAuthManager.map_user_to_teams( user_object=user_object, team_object=team_object, @@ -2654,7 +2672,7 @@ class JWTAuthManager: JWTAuthManager.validate_object_id( user_id=user_id, team_id=team_id, - enforce_rbac=general_settings.get("enforce_rbac", False), + enforce_rbac=bool(general_settings.get("enforce_rbac", False)), is_proxy_admin=False, ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 1739ac6d743..8d300c7c508 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -6980,6 +6980,11 @@ async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool: new_callable=AsyncMock, return_value="admin-user", ), + patch( # test-quality-ok: this control tests persistence by auth mode; write-policy behavior is covered separately + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.authorize_oauth_credential_request", + new_callable=AsyncMock, + return_value="admin-user", + ), patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._store_per_user_token_server_side", new_callable=AsyncMock, @@ -11168,7 +11173,7 @@ async def test_identity_bound_authorization_carries_nonce_and_caller_through_cal "path": "/authorize", "query_string": b"", "headers": []}) with ( patch( # test-quality-ok: isolate authenticated request resolution from the real encrypted OAuth round trip - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.authorize_oauth_credential_request", new=AsyncMock(return_value="alice")), patch( # test-quality-ok: isolate user access lookup while testing nonce and caller preservation "litellm.proxy._experimental.mcp_server.discoverable_endpoints._user_can_reach_mcp_server", @@ -11569,17 +11574,24 @@ async def test_oauth_exchange_stores_token_for_validated_jwt_user( "missing_database", ], ) +@pytest.mark.parametrize("credential_write", [False, True]) async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch, rejection: str, + credential_write: bool, ) -> 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 + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( + _extract_user_id_from_request, authorize_oauth_credential_request, + ) + allowed_servers: Final = AsyncMock(return_value=["server-a"]) + monkeypatch.setattr(mcp_server_manager.global_mcp_server_manager, "get_allowed_mcp_servers", allowed_servers) handler, signing_key = jwt_oauth_identity key: Final = ( rsa.generate_private_key(public_exponent=65537, key_size=2048) if rejection == "signature" else signing_key @@ -11603,7 +11615,13 @@ async def test_oauth_jwt_identity_rejects_untrusted_or_inactive_owner( ) 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 + request: Final = _token_request({"Authorization": f"Bearer {bearer}"}) + result: Final = ( + await authorize_oauth_credential_request(request, "server-a") + if credential_write else await _extract_user_id_from_request(request) + ) + assert result is None + allowed_servers.assert_not_awaited() @pytest.mark.asyncio @@ -11850,7 +11868,7 @@ async def test_oauth_credential_write_keeps_virtual_key_permissions( import asyncio from litellm.proxy._experimental.mcp_server import mcp_server_manager - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import authorize_oauth_credential_request from litellm.proxy._types import UserAPIKeyAuth, hash_token from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key @@ -11881,10 +11899,10 @@ async def test_oauth_credential_write_keeps_virtual_key_permissions( request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path="/server-a/token") if state == "cancelled": with pytest.raises(asyncio.CancelledError): - await _extract_user_id_from_request(request, "server-a") + await authorize_oauth_credential_request(request, "server-a") manager.get_allowed_mcp_servers.assert_awaited_once() return - assert await _extract_user_id_from_request(request, "server-a") == ("jwt-owner" if state == "allowed" else None) + assert await authorize_oauth_credential_request(request, "server-a") == ("jwt-owner" if state == "allowed" else None) if state in ("allowed", "server_denied", "lookup_error"): manager.get_allowed_mcp_servers.assert_awaited_once() writer: Final = manager.get_allowed_mcp_servers.call_args.args[0] @@ -11907,7 +11925,7 @@ async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_ros from litellm.models.user import LiteLLM_UserTable from litellm.proxy import proxy_server from litellm.proxy._experimental.mcp_server import mcp_server_manager - from litellm.proxy._experimental.mcp_server.bridge_token_flow import _extract_user_id_from_request + from litellm.proxy._experimental.mcp_server.bridge_token_flow import authorize_oauth_credential_request from litellm.proxy._types import LiteLLM_TeamTable, Member handler, signing_key = jwt_oauth_identity @@ -11925,7 +11943,7 @@ async def test_oauth_writer_preserves_claimed_team_instead_of_expanding_user_ros monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) bearer: Final = _oauth_identity_jwt(signing_key, claims={"team": "a"}) request: Final = _token_request({"Authorization": f"Bearer {bearer}"}, path=f"/{server_id}/token") - assert await _extract_user_id_from_request(request, server_id) == ( + assert await authorize_oauth_credential_request(request, server_id) == ( "jwt-owner" if server_id == "team-a-server" else None ) manager.get_allowed_mcp_servers.assert_awaited_once() diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index babbf88dc29..6fbbb37fbc3 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -6791,12 +6791,11 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla @pytest.mark.asyncio -@pytest.mark.parametrize("identity_only", [False, True]) -@pytest.mark.parametrize("allow_provisioning", [False, True]) +@pytest.mark.parametrize("operation", ["identity", "authorize", "admit"]) @pytest.mark.parametrize("existing_user", [False, True]) @pytest.mark.parametrize("model_allowed", [False, True]) -async def test_auth_builder_identity_lookup_does_not_provision_users( - monkeypatch: pytest.MonkeyPatch, identity_only: bool, allow_provisioning: bool, existing_user: bool, model_allowed: bool +async def test_jwt_identity_and_authorization_keep_provisioning_in_admission( + monkeypatch: pytest.MonkeyPatch, operation: str, existing_user: bool, model_allowed: bool ) -> None: from litellm.proxy._types import ScopeMapping from litellm.proxy.auth.auth_checks import UserNotFoundError @@ -6805,7 +6804,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( private_key, jwk = _get_rsa_key_and_jwk("identity-mode") cache: Final = UserApiKeyCache() cache.set_cache("litellm_jwt_auth_keys_https://identity.example/jwks", [jwk]) - user_id: Final = f"identity-mode-{identity_only}-{existing_user}-{model_allowed}" + user_id: Final = f"identity-mode-{operation}-{existing_user}-{model_allowed}" user: Final = LiteLLM_UserTable(user_id=user_id, organization_memberships=[]) if existing_user: cache.set_cache(user_id, user) @@ -6831,26 +6830,38 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( token: Final = _encode_rsa_jwt( private_key, "https://identity.example", "gateway", "identity-mode", {"sub": user_id, "scope": "allowed"} ) - pending: Final = JWTAuthManager.auth_builder( - api_key=token, - jwt_handler=handler, + common: Final = { + "api_key": token, + "jwt_handler": handler, + "prisma_client": database, + "user_api_key_cache": cache, + "parent_otel_span": None, + "proxy_logging_obj": MagicMock(), + } + if operation == "identity": + if not existing_user: + with pytest.raises(UserNotFoundError): + await JWTAuthManager.resolve_identity(**common) + else: + identity: Final = await JWTAuthManager.resolve_identity(**common) + assert identity.user_id == user_id + assert identity.user_object is not None and identity.user_object.user_id == user_id + users.create.assert_not_awaited() + return + authorize: Final = JWTAuthManager.auth_builder if operation == "admit" else JWTAuthManager.authorize_jwt + pending: Final = authorize( + **common, request_data={"model": "allowed-model" if model_allowed else "forbidden-model"}, general_settings={}, - route="/example/token" if identity_only else "/mcp/example", - prisma_client=database, - user_api_key_cache=cache, - parent_otel_span=None, - proxy_logging_obj=MagicMock(), - identity_only=identity_only, - allow_provisioning=allow_provisioning, + route="/mcp/example", ) - if not identity_only and not model_allowed: + if not model_allowed: with pytest.raises(HTTPException) as denial: await pending assert denial.value.status_code == 403 users.create.assert_not_awaited() return - if (identity_only or not allow_provisioning) and not existing_user: + if operation == "authorize" and not existing_user: with pytest.raises(UserNotFoundError): await pending else: @@ -6858,7 +6869,7 @@ async def test_auth_builder_identity_lookup_does_not_provision_users( assert result["user_id"] == user_id assert result["user_object"] is not None assert result["user_object"].user_id == user_id - assert users.create.await_count == (0 if identity_only or not allow_provisioning or existing_user else 1) + assert users.create.await_count == (0 if operation == "authorize" or existing_user else 1) def _entra_agent_registry() -> AgentRegistry: @@ -6997,6 +7008,14 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a ) jwt_handler.bind_agent_lookup(_entra_agent_registry()) + if identity_only: + identity = await JWTAuthManager.resolve_identity( + api_key=token, jwt_handler=jwt_handler, prisma_client=None, + user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + ) + assert identity.agent_id == "canonical-agent-id" + return + result = await JWTAuthManager.auth_builder( api_key=token, jwt_handler=jwt_handler, @@ -7007,10 +7026,9 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, - identity_only=identity_only, ) - assert result["is_proxy_admin"] is (is_admin_token and not identity_only) + assert result["is_proxy_admin"] is is_admin_token assert result["agent_id"] == "canonical-agent-id" @@ -7025,6 +7043,14 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch ) jwt_handler.bind_agent_lookup(_entra_agent_registry()) + if identity_only: + with pytest.raises(HTTPException) as denial: + await JWTAuthManager.resolve_identity( + api_key=token, jwt_handler=jwt_handler, prisma_client=None, + user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + ) + assert denial.value.status_code == 403 + return with pytest.raises(HTTPException) as exc_info: await JWTAuthManager.auth_builder( api_key=token, @@ -7036,7 +7062,39 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, - identity_only=identity_only, ) assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("admission", [False, True]) +async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatch, admission: bool): + from litellm.proxy.management_endpoints import team_endpoints + + handler, token = _entra_signed_app_token( + monkeypatch, azp="canonical-agent-id", scope=LiteLLM_JWTAuth().admin_jwt_scope, + ) + handler.bind_agent_lookup(_entra_agent_registry()) + handler.litellm_jwtauth.team_id_upsert = True + handler.litellm_jwtauth.admin_allowed_routes = ["openai_routes"] + database = MagicMock() + database.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + create_team = AsyncMock(return_value=LiteLLM_TeamTable(team_id="new-team").model_dump()) + monkeypatch.setattr(team_endpoints, "new_team", create_team) + resolve = JWTAuthManager.auth_builder if admission else JWTAuthManager.authorize_jwt + + result = await resolve( + api_key=token, jwt_handler=handler, request_data={}, general_settings={}, + route="/chat/completions", prisma_client=database, + user_api_key_cache=handler.user_api_key_cache, parent_otel_span=None, + proxy_logging_obj=MagicMock(), request_headers={"x-litellm-team-id": "new-team"}, + ) + + assert result["is_proxy_admin"] is True + if admission: + create_team.assert_awaited_once() + assert result["team_id"] == "new-team" + else: + create_team.assert_not_awaited() + assert result["team_id"] is None From ee676d59f259dbb1283b741b7856aaf8559ee3cf Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 23:08:14 -0700 Subject: [PATCH 9/9] fix(mcp): preserve browser OAuth for unrelated bearer tokens --- .../mcp_server/bridge_token_flow.py | 63 ++++++++ .../mcp_server/discoverable_endpoints.py | 7 +- .../mcp_server/test_discoverable_endpoints.py | 147 +++++++++++++++++- 3 files changed, 207 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py index f001ee87dbd..2b13baa624b 100644 --- a/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py +++ b/litellm/proxy/_experimental/mcp_server/bridge_token_flow.py @@ -1,6 +1,8 @@ """Bridge token flow: litellm identity resolution and the DCR-bridge oauth_delegate mint/refresh pipeline.""" import math +import os +import secrets from dataclasses import dataclass from datetime import datetime, timezone from typing import TYPE_CHECKING, Final, Literal @@ -12,6 +14,9 @@ from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + _V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator +) from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: @@ -49,6 +54,64 @@ def _litellm_key_from_request(request: Request) -> str | None: return None +async def oauth_authorization_uses_gateway_credential(request: Request) -> bool: + """Classify credentials for browser authorize; candidates still require full authorization.""" + from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # startup owns the active auth configuration + jwt_handler, + master_key, + user_custom_auth, + ) + + if "x-litellm-api-key" in request.headers: + return True + token: Final = _litellm_key_from_request(request) + if token is None: + return "authorization" in request.headers + if token.startswith("sk-") or (master_key and secrets.compare_digest(token.encode(), master_key.encode())): + return True + if user_custom_auth is not None or jwt_handler.litellm_jwtauth.oidc_userinfo_enabled: + return True + if not JWTHandler.is_jwt(token): + return await _opaque_bearer_is_gateway_credential(token) + claims: Final = JWTHandler.get_unverified_claims(token) + issuer: Final = claims.get("iss") if claims is not None else None + global_issuer: Final = os.getenv("JWT_ISSUER") + # An unscoped global validator can accept issuers absent from the configured issuer list. + if not isinstance(issuer, str) or not issuer or not global_issuer: + return True + return issuer == global_issuer or any( + issuer == configured.issuer for configured in jwt_handler.litellm_jwtauth.issuers or () + ) + + +async def _opaque_bearer_is_gateway_credential(token: str) -> bool: + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( + is_envelope, # noqa: PLC0415 # envelope imports bridge types + is_refresh_envelope, + ) + from litellm.proxy._types import hash_token # noqa: PLC0415 # proxy import cycle + from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken # noqa: PLC0415 # proxy import cycle + from litellm.proxy.auth.resolvers.exceptions import KeyNotFoundError # noqa: PLC0415 # proxy import cycle + from litellm.proxy.auth.resolvers.store import IdentityStore # noqa: PLC0415 # proxy import cycle + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # startup owns the identity store dependencies + prisma_client, + user_api_key_cache, + ) + + if is_envelope(token) or is_refresh_envelope(token) or token.startswith(_V2_GCM_PREFIX): + return True + try: + if ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) is not None: + return True + await IdentityStore(prisma_client, user_api_key_cache).resolve(hashed_token=hash_token(token)) + except KeyNotFoundError: + return False + except Exception as exc: # noqa: BLE001 # an identity lookup fault must not permit cookie fallback + verbose_logger.debug("OAuth bearer ownership could not be checked (%s)", type(exc).__name__) + return True + + def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool: """``True`` when the presented key is neither blocked nor past its expiry. diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 7e7189c8a6b..ffb27d5f92e 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -29,12 +29,12 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _BridgeRefreshReady, _extract_user_id_from_request, _finish_bridge_mint, - _litellm_key_from_request, # pyright: ignore[reportPrivateUsage] # shared credential precedence for authorization issuance _prepare_bridge_mint, _prepare_bridge_refresh, _reload_active_user_by_id, authorize_oauth_credential_request, can_store_oauth_credential, + oauth_authorization_uses_gateway_credential, ) from litellm.proxy._experimental.mcp_server.faults import ( CallerRejected, @@ -851,10 +851,11 @@ async def _resolve_oauth_authorization_user( _user_id_from_session_cookie, ) + use_gateway_credential: Final = enforce_binding and await oauth_authorization_uses_gateway_credential(request) request_user_id: Final = ( - await authorize_oauth_credential_request(request, mcp_server.server_id) if enforce_binding else None + await authorize_oauth_credential_request(request, mcp_server.server_id) if use_gateway_credential else None ) - if enforce_binding and request_user_id is None and _litellm_key_from_request(request): + if use_gateway_credential and request_user_id is None: return _bridge_access_denied_redirect(redirect_uri, state, mcp_server) user_id: Final = request_user_id or _user_id_from_session_cookie(request) if user_id is None: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 8d300c7c508..aa45b2f6793 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11170,7 +11170,7 @@ async def test_identity_bound_authorization_carries_nonce_and_caller_through_cal ), ) request = Request({"type": "http", "scheme": "https", "server": ("proxy.example.com", 443), - "path": "/authorize", "query_string": b"", "headers": []}) + "path": "/authorize", "query_string": b"", "headers": [(b"authorization", b"Bearer sk-alice")]}) with ( patch( # test-quality-ok: isolate authenticated request resolution from the real encrypted OAuth round trip "litellm.proxy._experimental.mcp_server.discoverable_endpoints.authorize_oauth_credential_request", @@ -12059,18 +12059,52 @@ async def test_signed_oauth_callback_honors_credential_write_policy( @pytest.mark.asyncio @pytest.mark.parametrize("allowed", [False, True]) +@pytest.mark.parametrize("credential", [ + "jwt", "key", "expired_jwt", "wrong_audience", "bad_signature", "malformed_jwt", "missing_issuer", + "foreign_explicit", "blank_explicit", "unknown_key", "blocked_key", "expired_key", "opaque_record", + "opaque_outage", "opaque_oidc", "opaque_custom", "foreign_unscoped", "foreign_configured", "encrypted", "invalid_encrypted", "envelope", "master", +]) async def test_identity_bound_authorize_preserves_presented_jwt_permissions( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch, allowed: bool, + credential: str, ) -> None: + import jwt + from datetime import datetime, timedelta, timezone from urllib.parse import parse_qs, urlparse - from litellm.proxy._experimental.mcp_server import byok_oauth_endpoints, discoverable_endpoints, mcp_server_manager + from litellm.models.user import LiteLLM_UserTable + from litellm.proxy import proxy_server + from litellm.proxy._types import JWTIssuerConfig, UserAPIKeyAuth, hash_token + from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints, mcp_server_manager from litellm.proxy._types import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer - _, signing_key = jwt_oauth_identity + handler, signing_key = jwt_oauth_identity + master: Final = "browser-session-test-signing-key-123456789" + monkeypatch.setattr(proxy_server, "master_key", master) + monkeypatch.setattr(proxy_server, "user_custom_auth", (lambda: None) if credential == "opaque_custom" else None) + handler.litellm_jwtauth.oidc_userinfo_enabled = credential == "opaque_oidc" + if credential == "foreign_unscoped": + monkeypatch.delenv("JWT_ISSUER") + if credential == "foreign_configured": + handler.litellm_jwtauth.issuers = [JWTIssuerConfig( + issuer="https://unrelated.example.test", jwks_url="https://idp.example.test/jwks", + audience="litellm-proxy", user_id_jwt_field="identity.user_id", + )] + proxy_server.prisma_client.get_data = AsyncMock( + return_value=None, side_effect=RuntimeError("database unavailable") if credential == "opaque_outage" else None, + ) + handler.user_api_key_cache.set_cache("cookie-owner", LiteLLM_UserTable(user_id="cookie-owner")) + key: Final = "opaque-record" if credential == "opaque_record" else "sk-browser-gateway-key" + if credential in ("key", "blocked_key", "expired_key", "opaque_record"): + handler.user_api_key_cache.set_cache(hash_token(key), UserAPIKeyAuth( + token=hash_token(key), user_id="jwt-owner", blocked=credential in ("blocked_key", "opaque_record"), + expires=datetime.now(timezone.utc) - timedelta(seconds=60) if credential == "expired_key" else None, + )) monkeypatch.setenv("LITELLM_SALT_KEY", "authorize-policy-test-salt") server: Final = MCPServer( server_id="bound-server", name="bound-server", transport=MCPTransport.http, @@ -12086,21 +12120,120 @@ async def test_identity_bound_authorize_preserves_presented_jwt_permissions( side_effect=lambda auth: [server.server_id] if allowed or auth.mcp_admitted_user_subject else [], ) monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) - monkeypatch.setattr( # test-quality-ok: session-cookie decoder is the separate authentication boundary; a valid cookie must not override a denied explicit credential - byok_oauth_endpoints, "_user_id_from_session_cookie", lambda request: "jwt-owner", + bearer: Final = ( + key if credential in ("key", "blocked_key", "expired_key", "opaque_record", "unknown_key") + else "opaque-bearer" if credential in ("opaque_outage", "opaque_oidc", "opaque_custom") + else "not.a.jwt" if credential == "malformed_jwt" + else "llm_env_invalid" if credential == "envelope" + else "v2:gcm:invalid" if credential == "invalid_encrypted" + else master if credential == "master" + else ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( + LiteLLM_UserTable(user_id="jwt-owner", user_role="internal_user"), + ) if credential == "encrypted" + else jwt.encode({"iss": "https://idp.example.test"}, "wrong-signing-key-at-least-32-bytes", algorithm="HS256") + if credential == "bad_signature" + else jwt.encode({"sub": "jwt-owner"}, signing_key, algorithm="RS256") if credential == "missing_issuer" + else _oauth_identity_jwt( + signing_key, + expires_in=-60 if credential == "expired_jwt" else 300, + audience="another-service" if credential == "wrong_audience" else "litellm-proxy", + issuer="https://unrelated.example.test" if credential.startswith("foreign_") or credential == "blank_explicit" else "https://idp.example.test", + ) + ) + cookie: Final = jwt.encode( + {"user_id": "cookie-owner", "login_method": "sso", "exp": int(time.time()) + 300}, master, algorithm="HS256", ) response: Final = await discoverable_endpoints.authorize_with_server( - request=_token_request({"Authorization": f"Bearer {_oauth_identity_jwt(signing_key)}"}), + request=_token_request({ + "Authorization": f"Bearer {bearer}", "Cookie": f"token={cookie}", + **({"x-litellm-api-key": bearer} if credential == "foreign_explicit" else {}), + **({"x-litellm-api-key": ""} if credential == "blank_explicit" else {}), + }), mcp_server=server, client_id="client", redirect_uri="http://127.0.0.1:6274/callback", state="client-state", code_challenge="pkce-challenge", code_challenge_method="S256", ) redirect: Final = urlparse(response.headers["location"]) query: Final = parse_qs(redirect.query) - if allowed: + if allowed and credential in ("jwt", "key", "foreign_unscoped", "foreign_configured"): assert redirect.hostname == "upstream.example.test" assert query["nonce"] and response.headers.get("set-cookie") + assert all(call.args[0].user_id == "jwt-owner" for call in manager.get_allowed_mcp_servers.await_args_list) else: assert redirect.hostname == "127.0.0.1" assert query["error"] == ["access_denied"] assert query["state"] == ["client-state"] assert "set-cookie" not in response.headers + + proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called() + proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() + proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("credential", ["none", "opaque", "foreign_jwt"]) +@pytest.mark.parametrize("cookie_state", ["allowed", "server_denied", "expired", "missing"]) +async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + credential: str, + cookie_state: str, +) -> None: + import jwt + from urllib.parse import parse_qs, urlparse + + from litellm.models.user import LiteLLM_UserTable + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import discoverable_endpoints, mcp_server_manager + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer + + handler, signing_key = jwt_oauth_identity + master: Final = "browser-session-test-signing-key-123456789" + monkeypatch.setattr(proxy_server, "master_key", master) + monkeypatch.setattr(proxy_server, "user_custom_auth", None) + monkeypatch.setenv("LITELLM_SALT_KEY", "authorize-policy-test-salt") + handler.user_api_key_cache.set_cache("cookie-owner", LiteLLM_UserTable(user_id="cookie-owner")) + proxy_server.prisma_client.get_data = AsyncMock(return_value=None) + server: Final = MCPServer( + server_id="bound-server", name="bound-server", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client", + authorization_url="https://upstream.example.test/authorize", token_url="https://upstream.example.test/token", + oauth_identity_binding=MCPOAuthIdentityBinding( + mode="enforce", issuer="https://upstream.example.test", audiences=["client"], + ), + ) + manager: Final = MagicMock() + manager.get_allowed_mcp_servers = AsyncMock(return_value=[] if cookie_state == "server_denied" else [server.server_id]) + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + bearer: Final = ( + _oauth_identity_jwt(signing_key, issuer="https://unrelated.example.test") + if credential == "foreign_jwt" else "unrelated-upstream-bearer" + ) + cookie: Final = jwt.encode( + {"user_id": "cookie-owner", "login_method": "sso", "exp": int(time.time()) + (-60 if cookie_state == "expired" else 300)}, + master, algorithm="HS256", + ) + response: Final = await discoverable_endpoints.authorize_with_server( + request=_token_request({ + **({"Authorization": f"Bearer {bearer}"} if credential != "none" else {}), + **({"Cookie": f"token={cookie}"} if cookie_state != "missing" else {}), + }), + mcp_server=server, client_id="client", redirect_uri="http://127.0.0.1:6274/callback", + state="client-state", code_challenge="pkce-challenge", code_challenge_method="S256", + ) + redirect: Final = urlparse(response.headers["location"]) + query: Final = parse_qs(redirect.query) + if cookie_state == "allowed": + assert redirect.hostname == "upstream.example.test" + assert query["nonce"] and response.headers.get("set-cookie") + manager.get_allowed_mcp_servers.assert_awaited_once() + assert manager.get_allowed_mcp_servers.call_args.args[0].user_id == "cookie-owner" + elif cookie_state == "server_denied": + assert query["error"] == ["access_denied"] + assert query["state"] == ["client-state"] + else: + assert redirect.path == "/sso/key/generate" + manager.get_allowed_mcp_servers.assert_not_awaited() + proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called() + proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() + proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()