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"""