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 <yassin@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-08-20 16:20:47 -07:00 • committed by GitHub
parent 7b20828c72
commit bcf9e53c27
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 339 additions and 10 deletions

View file

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

View file

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