fix(mcp): reuse the standard JWT auth builder for OAuth ownership

This commit is contained in:
Joshua Valluru 2026-09-15 16:50:55 -07:00
parent eda98f38d9
commit 88d0371a46
2 changed files with 88 additions and 24 deletions

View file

@ -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__)

View file

@ -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()