From bcf9e53c27045e49853abb2bafab96feb2ee5989 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 20 Aug 2026 16:20:47 -0700 Subject: [PATCH] feat(sso): source generic OIDC user claims from ID/access token when UserInfo is incomplete (#37696) Some IdPs, ADFS among them, return only `sub` from UserInfo and put the real identity claims in the ID token or the access token. Those users land in the Admin UI with no username, email, groups or teams. Adds an opt-in `GENERIC_INCLUDE_TOKEN_CLAIMS` that merges token claims into the UserInfo response before the existing `GENERIC_USER_*_ATTRIBUTE` mappings run. Precedence is UserInfo, then id_token, then access token, and it applies to both the PKCE and non-PKCE login flows. With the flag unset, behavior is unchanged. Co-authored-by: Yassin Kortam --- litellm/proxy/management_endpoints/ui_sso.py | 61 +++- .../proxy/management_endpoints/test_ui_sso.py | 288 ++++++++++++++++++ 2 files changed, 339 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 1ebcb53fd6b..3c135650de9 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -19,6 +19,7 @@ import secrets from collections.abc import Mapping, Sequence from copy import deepcopy from html import escape +from types import MappingProxyType from typing import ( TYPE_CHECKING, Annotated, @@ -245,6 +246,7 @@ def _team_detail_db(repo: "_HasTeamDetailTable") -> "_PrismaTableActions[_TeamDe _MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str]) +_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object]) def _decode_model_aliases(value: object) -> object: @@ -1002,6 +1004,30 @@ def process_sso_jwt_access_token( return None +def _decode_sso_token_claims(token: str | None) -> Mapping[str, object]: + if not token: + return MappingProxyType({}) + try: + return MappingProxyType( + _SSO_TOKEN_CLAIMS_ADAPTER.validate_python(jwt.decode(token, options={"verify_signature": False})) + ) + except (jwt.exceptions.InvalidTokenError, ValidationError): + verbose_proxy_logger.debug("SSO token is not a decodable JWT, skipping token claims") + return MappingProxyType({}) + + +def _merge_sso_token_claims( + userinfo: Mapping[str, object], + id_token: str | None, + access_token: str | None, +) -> Mapping[str, object]: + sources: Final = (userinfo, _decode_sso_token_claims(id_token), _decode_sso_token_claims(access_token)) + claim_names: Final = frozenset(key for source in sources for key in source) + return MappingProxyType( + {key: next((source[key] for source in sources if source.get(key) is not None), None) for key in claim_names} + ) + + async def _raise_if_sso_exceeds_free_user_limit(premium_user: bool, prisma_client: PrismaClient | None) -> None: """Free tier allows SSO for up to 5 billable users; beyond that requires an Enterprise license.""" if premium_user is True: @@ -1534,12 +1560,34 @@ async def get_generic_sso_response( role_mappings: Final = await _setup_role_mappings() team_mappings: Final = await _setup_team_mappings() + generic_include_token_claims: Final = os.getenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "false").lower() == "true" - def response_convertor(response, client): + def response_convertor(response: Mapping[str, object], httpx_session: object): nonlocal received_response # return for user debugging - received_response = response + response_id_token: Final = response.get("id_token") + response_access_token: Final = response.get("access_token") + id_token: Final = ( + response_id_token if isinstance(response_id_token, str) and response_id_token else generic_sso.id_token + ) + access_token: Final = ( + response_access_token + if isinstance(response_access_token, str) and response_access_token + else generic_sso.access_token + ) + claims: Final = ( + _merge_sso_token_claims( + userinfo=response, + id_token=id_token, + access_token=access_token, + ) + if generic_include_token_claims + else response + ) + received_response = { # mutable-ok: preserve the existing dict return contract + key: value for key, value in claims.items() if key not in _OAUTH_TOKEN_FIELDS + } return generic_response_convertor( - response=response, + response=claims, jwt_handler=jwt_handler, sso_jwt_handler=sso_jwt_handler, role_mappings=role_mappings, @@ -1641,13 +1689,6 @@ async def get_generic_sso_response( # Pass the full response so custom response_convertor implementations # can access all fields (including id_token for claim extraction). result = response_convertor(combined_response, generic_sso) - # Strip bearer credentials from combined_response before storing in - # received_response. received_response may appear in restricted-group - # error messages — bearer tokens (access_token, id_token, refresh_token) - # must not be exposed to callers. - # Assign directly rather than relying on nonlocal mutation so that Pyright - # can track that received_response is non-None from this point on. - received_response = {k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS} sso_assertion = assertion_from_sso_login( combined_response.get("id_token"), combined_response.get("refresh_token") ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 74f89242f7f..66cb07ef2c0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -1606,6 +1606,294 @@ async def test_get_generic_sso_response_with_empty_headers(): assert result == mock_sso_response +@pytest.mark.asyncio +async def test_get_generic_sso_response_includes_token_claims_when_enabled(monkeypatch): + import jwt as pyjwt + + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + from litellm.proxy._types import LitellmUserRoles + + mock_request = MagicMock(spec=Request) + mock_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler.get_all_jwt_team_ids.return_value = ["team-from-userinfo"] + mock_sso_jwt_handler.get_team_ids_from_jwt.return_value = [] + + userinfo = { + "sub": "subject-only", + "groups": ["admins"], + "access_token": "", + } + access_token = pyjwt.encode( + { + "upn": "token-user@example.com", + "email": "token-user@example.com", + "given_name": "Token", + "family_name": "User", + "display_name": "Token User", + }, + "test-secret", + algorithm="HS256", + ) + mock_sso_instance = MagicMock() + mock_sso_instance.access_token = access_token + mock_sso_instance.id_token = None + + def fake_create_provider(*, response_convertor, **_kwargs): + mock_sso_instance.verify_and_process = AsyncMock( + side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object()) + ) + return MagicMock(return_value=mock_sso_instance) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true") + monkeypatch.setenv("GENERIC_USER_ID_ATTRIBUTE", "upn") + monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email") + monkeypatch.setenv("GENERIC_USER_FIRST_NAME_ATTRIBUTE", "given_name") + monkeypatch.setenv("GENERIC_USER_LAST_NAME_ATTRIBUTE", "family_name") + monkeypatch.setenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "display_name") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['admins']}") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "groups") + + with patch("fastapi_sso.sso.base.DiscoveryDocument"): + with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider): + result, received_response, _, _ = await get_generic_sso_response( + request=mock_request, + jwt_handler=mock_jwt_handler, + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=mock_sso_jwt_handler, + ) + + assert isinstance(result, CustomOpenID) + assert result.id == "token-user@example.com" + assert result.email == "token-user@example.com" + assert result.first_name == "Token" + assert result.last_name == "User" + assert result.display_name == "Token User" + assert result.team_ids == ["team-from-userinfo"] + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + assert received_response is not None + assert "access_token" not in received_response + assert "id_token" not in received_response + assert "refresh_token" not in received_response + + +@pytest.mark.asyncio +async def test_get_generic_sso_response_does_not_include_token_claims_when_disabled(monkeypatch): + import jwt as pyjwt + + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + + mock_request = MagicMock(spec=Request) + mock_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler.get_all_jwt_team_ids.return_value = ["team-from-userinfo"] + mock_sso_jwt_handler.get_team_ids_from_jwt.return_value = [] + access_token = pyjwt.encode({"upn": "token-user@example.com"}, "test-secret", algorithm="HS256") + userinfo = {"sub": "subject-only", "groups": ["admins"], "access_token": ""} + mock_sso_instance = MagicMock() + mock_sso_instance.access_token = access_token + mock_sso_instance.id_token = None + + def fake_create_provider(*, response_convertor, **_kwargs): + mock_sso_instance.verify_and_process = AsyncMock( + side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object()) + ) + return MagicMock(return_value=mock_sso_instance) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "false") + monkeypatch.setenv("GENERIC_USER_ID_ATTRIBUTE", "upn") + monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email") + monkeypatch.setenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "display_name") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['admins']}") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "groups") + + with patch("fastapi_sso.sso.base.DiscoveryDocument"): + with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider): + result, received_response, _, _ = await get_generic_sso_response( + request=mock_request, + jwt_handler=mock_jwt_handler, + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=mock_sso_jwt_handler, + ) + + assert isinstance(result, CustomOpenID) + assert result.id is None + assert result.email is None + assert result.display_name is None + assert result.team_ids == ["team-from-userinfo"] + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + assert received_response == {"sub": "subject-only", "groups": ["admins"]} + + +def test_merge_sso_token_claims_precedence_and_invalid_tokens(): + import jwt as pyjwt + + from litellm.proxy.management_endpoints.ui_sso import _merge_sso_token_claims + + id_token = pyjwt.encode( + {"preferred_username": "id-user", "email": "id@example.com", "id_only": "id-value"}, + "test-secret", + algorithm="HS256", + ) + access_token = pyjwt.encode( + {"preferred_username": "access-user", "email": "access@example.com", "access_only": "access-value"}, + "test-secret", + algorithm="HS256", + ) + + merged = _merge_sso_token_claims( + userinfo={"preferred_username": "userinfo-user", "email": None, "userinfo_only": "userinfo-value"}, + id_token=id_token, + access_token=access_token, + ) + + assert merged["preferred_username"] == "userinfo-user" + assert merged["email"] == "id@example.com" + assert merged["id_only"] == "id-value" + assert merged["access_only"] == "access-value" + + userinfo_only = _merge_sso_token_claims( + userinfo={"sub": "userinfo-user", "email": "userinfo@example.com"}, + id_token=pyjwt.encode({}, "test-secret", algorithm="HS256"), + access_token="opaque-access-token", + ) + + assert userinfo_only == {"sub": "userinfo-user", "email": "userinfo@example.com"} + + +@pytest.mark.asyncio +async def test_get_generic_sso_response_pkce_merges_token_claims_and_excludes_credentials(monkeypatch): + """The real PKCE path merges access-token claims and keeps bearer credentials out of received_response. + + Only the PKCE verifier cache and the HTTP transport are injected, so + prepare_token_exchange_parameters, _pkce_token_exchange and the claim merge all run for real. + """ + import jwt as pyjwt + from starlette.requests import Request as StarletteRequest + + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + + access_token = pyjwt.encode( + {"sub": "token-user", "email": "token-user@example.com"}, "test-secret", algorithm="HS256" + ) + request = StarletteRequest( + { + "type": "http", + "method": "GET", + "path": "/sso/callback", + "query_string": b"code=test-code&state=test-state", + "headers": [(b"cookie", b"litellm_oauth_state=test-state")], + } + ) + + pkce_cache = MagicMock(redis_cache=None) + pkce_cache.async_get_cache = AsyncMock(return_value={"code_verifier": "test-code-verifier"}) + pkce_cache.async_delete_cache = AsyncMock() + + token_endpoint_response = MagicMock(status_code=200) + token_endpoint_response.json.return_value = { + "access_token": access_token, + "id_token": "id-token-secret", + "refresh_token": "refresh-token-secret", + } + token_client = MagicMock() + token_client.post = AsyncMock(return_value=token_endpoint_response) + + userinfo_endpoint_response = MagicMock(status_code=200) + userinfo_endpoint_response.json.return_value = {"sub": "userinfo-user"} + userinfo_client = MagicMock() + userinfo_client.get = AsyncMock(return_value=userinfo_endpoint_response) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_CLIENT_USE_PKCE", "true") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true") + monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email") + + with ( + patch("litellm.proxy.proxy_server.redis_usage_cache", None), + patch("litellm.proxy.proxy_server.user_api_key_cache", pkce_cache), + patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client", + side_effect=[token_client, userinfo_client], + ), + ): + result, received_response, _, _ = await get_generic_sso_response( + request=request, + jwt_handler=MagicMock(spec=JWTHandler), + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=None, + ) + + # The real token exchange ran: it forwarded the cached verifier to the token endpoint. + assert token_client.post.await_args.kwargs["data"]["code_verifier"] == "test-code-verifier" + # UserInfo wins for sub; email exists only on the access token, so the merge must supply it. + assert isinstance(result, CustomOpenID) + assert result.email == "token-user@example.com" + assert received_response == {"sub": "userinfo-user", "email": "token-user@example.com"} + pkce_cache.async_delete_cache.assert_awaited_once_with(key="pkce_verifier:test-state") + + +@pytest.mark.asyncio +async def test_get_generic_sso_response_ignores_opaque_and_empty_token_claims(monkeypatch): + import jwt as pyjwt + + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + + mock_request = MagicMock(spec=Request) + mock_jwt_handler = MagicMock(spec=JWTHandler) + userinfo = { + "preferred_username": "userinfo-user", + "email": "userinfo@example.com", + "sub": "User Info", + } + mock_sso_instance = MagicMock() + mock_sso_instance.access_token = "opaque-access-token" + mock_sso_instance.id_token = pyjwt.encode({}, "test-secret", algorithm="HS256") + + def fake_create_provider(*, response_convertor, **_kwargs): + mock_sso_instance.verify_and_process = AsyncMock( + side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object()) + ) + return MagicMock(return_value=mock_sso_instance) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true") + + with patch("fastapi_sso.sso.base.DiscoveryDocument"): + with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider): + result, received_response, _, _ = await get_generic_sso_response( + request=mock_request, + jwt_handler=mock_jwt_handler, + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=None, + ) + + assert isinstance(result, CustomOpenID) + assert result.id == "userinfo-user" + assert result.email == "userinfo@example.com" + assert result.display_name == "User Info" + assert received_response == userinfo + + class TestCLISSOCallbackFunction: """Test the cli_sso_callback function specifically"""