Merge pull request #24701 from BerriAI/litellm_fix-jwt-role-mappings

fix(sso): pass decoded JWT access token to role mapping during SSO login
This commit is contained in:
ryan-crabbe-berri 2026-03-27 18:06:15 -07:00 committed by GitHub
commit a533de0b08
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 246 additions and 9 deletions

View file

@ -204,7 +204,7 @@ def process_sso_jwt_access_token(
sso_jwt_handler: Optional[JWTHandler],
result: Union[OpenID, dict, None],
role_mappings: Optional["RoleMappings"] = None,
) -> None:
) -> Optional[dict]:
"""
Process SSO JWT access token and extract team IDs and user role if available.
@ -218,6 +218,12 @@ def process_sso_jwt_access_token(
sso_jwt_handler: SSO-specific JWT handler for team ID extraction
result: The SSO result object to update with team IDs and role
role_mappings: Optional role mappings configuration for group-based role determination
Returns:
The decoded access token payload dict, or None if decoding failed or
inputs were missing. Callers can pass this to _sync_user_role_from_jwt_role_map
so it has access to custom role claims (e.g. custom_roles) that are
encoded inside the JWT but stripped from received_response.
"""
if access_token_str and result:
import jwt
@ -230,7 +236,7 @@ def process_sso_jwt_access_token(
verbose_proxy_logger.debug(
"Access token is not a valid JWT (possibly an opaque token), skipping JWT-based extraction"
)
return
return None
# Extract team IDs from access token if sso_jwt_handler is available
if sso_jwt_handler:
@ -306,6 +312,10 @@ def process_sso_jwt_access_token(
f"Set user_role='{user_role}' from JWT access token"
)
return access_token_payload
return None
@router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False)
async def google_login(
@ -817,7 +827,7 @@ async def get_generic_sso_response(
], # sso specific jwt handler - used for restricted sso group access control
generic_client_id: str,
redirect_url: str,
) -> Tuple[Union[OpenID, dict], Optional[dict]]: # return received response
) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload)
# make generic sso provider
from fastapi_sso.sso.base import DiscoveryDocument
from fastapi_sso.sso.generic import create_provider
@ -872,6 +882,7 @@ async def get_generic_sso_response(
code_verifier: Optional[
str
] = None # assigned inside try; initialized for type tracking
access_token_payload: Optional[dict] = None # decoded JWT access token claims
try:
token_exchange_params = (
@ -958,7 +969,7 @@ async def get_generic_sso_response(
)
access_token_str = generic_sso.access_token
process_sso_jwt_access_token(
access_token_payload = process_sso_jwt_access_token(
access_token_str, sso_jwt_handler, result, role_mappings=role_mappings
)
# Delete the single-use PKCE verifier only after all downstream processing
@ -976,7 +987,7 @@ async def get_generic_sso_response(
additional_generic_sso_headers_dict,
)
verbose_proxy_logger.debug("generic result: %s", result)
return result or {}, received_response
return result or {}, received_response, access_token_payload
async def create_team_member_add_task(team_id, user_info):
@ -1176,6 +1187,56 @@ def _build_sso_user_update_data(
return update_data
async def _sync_user_role_from_jwt_role_map(
jwt_handler: Optional[JWTHandler],
received_response: Optional[dict],
user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]],
prisma_client: PrismaClient,
user_api_key_cache: DualCache,
user_defined_values: Optional[SSOUserDefinedValues],
) -> None:
"""
Apply jwt_litellm_role_map during SSO login.
When jwt_litellm_role_map is configured with sync_user_role_and_teams=True,
this ensures SSO users get the same role mapping as API/JWT users. Without
this, the SSO path falls back to INTERNAL_USER_VIEW_ONLY for roles that
don't directly match LitellmUserRoles enum values.
"""
if jwt_handler is None or received_response is None:
return
if not jwt_handler.litellm_jwtauth.sync_user_role_and_teams:
return
if not jwt_handler.litellm_jwtauth.jwt_litellm_role_map:
return
mapped_role = jwt_handler.map_jwt_role_to_litellm_role(received_response)
if mapped_role is None:
return
verbose_proxy_logger.info(
f"SSO jwt_litellm_role_map matched role: {mapped_role.value}"
)
# Update user_defined_values so downstream code uses the mapped role
if user_defined_values is not None:
user_defined_values["user_role"] = mapped_role.value
# Update existing DB record if role differs
if user_info is not None and user_info.user_role != mapped_role.value:
await prisma_client.db.litellm_usertable.update(
where={"user_id": user_info.user_id},
data={"user_role": mapped_role.value},
)
user_info.user_role = mapped_role.value
await user_api_key_cache.async_set_cache(
key=user_info.user_id,
value=user_info.model_dump()
if hasattr(user_info, "model_dump")
else dict(user_info),
)
def apply_user_info_values_to_sso_user_defined_values(
user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]],
user_defined_values: Optional[SSOUserDefinedValues],
@ -1279,6 +1340,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
google_client_id = os.getenv("GOOGLE_CLIENT_ID", None)
generic_client_id = os.getenv("GENERIC_CLIENT_ID", None)
received_response: Optional[dict] = None
access_token_payload: Optional[dict] = None
# get url from request
if master_key is None:
raise ProxyException(
@ -1307,7 +1369,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
)
elif generic_client_id is not None:
result, received_response = await get_generic_sso_response(
result, received_response, access_token_payload = await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,
generic_client_id=generic_client_id,
@ -1345,6 +1407,8 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
received_response=received_response,
generic_client_id=generic_client_id,
ui_access_mode=ui_access_mode,
access_token_payload=access_token_payload,
jwt_handler=jwt_handler,
return_to=cp_return_to,
)
@ -2417,6 +2481,8 @@ class SSOAuthenticationHandler:
received_response: Optional[dict] = None,
generic_client_id: Optional[str] = None,
ui_access_mode: Optional[Dict] = None,
access_token_payload: Optional[dict] = None,
jwt_handler: Optional[JWTHandler] = None,
return_to: Optional[str] = None,
) -> RedirectResponse:
import jwt
@ -2498,6 +2564,20 @@ class SSOAuthenticationHandler:
alternate_user_id=user_id,
)
# Sync user role from JWT claims via jwt_litellm_role_map (if configured).
# This ensures SSO users get the same role mapping as API/JWT users.
# Use the decoded access_token_payload (not received_response) because
# custom role claims (e.g. custom_roles) are encoded inside the JWT
# access token, which is stripped from received_response.
await _sync_user_role_from_jwt_role_map(
jwt_handler=jwt_handler,
received_response=access_token_payload or received_response,
user_info=user_info,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_defined_values=user_defined_values,
)
user_defined_values = apply_user_info_values_to_sso_user_defined_values(
user_info=user_info, user_defined_values=user_defined_values
)
@ -3703,7 +3783,7 @@ async def debug_sso_callback(request: Request):
)
elif generic_client_id is not None:
result, _ = await get_generic_sso_response(
result, _, _ = await get_generic_sso_response(
request=request,
jwt_handler=jwt_handler,
generic_client_id=generic_client_id,

View file

@ -24,6 +24,7 @@ from litellm.proxy.management_endpoints.ui_sso import (
MicrosoftSSOHandler,
SSOAuthenticationHandler,
_setup_team_mappings,
_sync_user_role_from_jwt_role_map,
determine_role_from_groups,
normalize_email,
process_sso_jwt_access_token,
@ -1321,7 +1322,7 @@ async def test_get_generic_sso_response_with_additional_headers():
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
result, received_response = await get_generic_sso_response(
result, received_response, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
@ -1383,7 +1384,7 @@ async def test_get_generic_sso_response_with_empty_headers():
"fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class
):
# Act
result, received_response = await get_generic_sso_response(
result, received_response, _ = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
generic_client_id=generic_client_id,
@ -5254,3 +5255,159 @@ class TestValidateReturnTo:
)
SSOAuthenticationHandler._validate_return_to("https://cp.example.com:3000/ui")
class TestSyncUserRoleFromJwtRoleMap:
"""Tests for _sync_user_role_from_jwt_role_map."""
@staticmethod
def _make_jwt_handler():
from litellm.caching.caching import DualCache
from litellm.proxy._types import (
JWTLiteLLMRoleMap,
LiteLLM_JWTAuth,
LitellmUserRoles,
)
handler = JWTHandler()
handler.update_environment(
prisma_client=None,
user_api_key_cache=DualCache(),
litellm_jwtauth=LiteLLM_JWTAuth(
roles_jwt_field="custom_roles",
user_id_upsert=True,
sync_user_role_and_teams=True,
jwt_litellm_role_map=[
JWTLiteLLMRoleMap(
jwt_role="my-admin",
litellm_role=LitellmUserRoles.PROXY_ADMIN,
),
JWTLiteLLMRoleMap(
jwt_role="my-viewer",
litellm_role=LitellmUserRoles.INTERNAL_USER,
),
],
),
)
return handler
@staticmethod
def _make_sso_values(user_role=None):
from litellm.proxy._types import SSOUserDefinedValues
user_id = "testuser@example.com"
return SSOUserDefinedValues(
models=[],
user_id=user_id,
user_email=user_id,
user_role=user_role,
max_budget=None,
budget_duration=None,
)
@pytest.mark.asyncio
async def test_stripped_response_has_no_roles(self):
"""Bug repro: stripped received_response lacks role claims."""
from litellm.caching.caching import DualCache
handler = self._make_jwt_handler()
sso_values = self._make_sso_values()
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={"token_type": "Bearer", "expires_in": 3600},
user_info=None,
prisma_client=AsyncMock(),
user_api_key_cache=DualCache(),
user_defined_values=sso_values,
)
assert sso_values["user_role"] is None
@pytest.mark.asyncio
async def test_decoded_access_token_maps_role(self):
"""Decoded JWT payload with role claims maps correctly."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import LitellmUserRoles
handler = self._make_jwt_handler()
sso_values = self._make_sso_values()
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={"sub": "testuser@example.com", "custom_roles": ["my-admin"]},
user_info=None,
prisma_client=AsyncMock(),
user_api_key_cache=DualCache(),
user_defined_values=sso_values,
)
assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
@pytest.mark.asyncio
async def test_existing_user_role_updated_in_db_and_cache(self):
"""Existing user with stale role gets updated in DB and cache."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import LitellmUserRoles
handler = self._make_jwt_handler()
cache = DualCache()
prisma = AsyncMock()
prisma.db.litellm_usertable.update = AsyncMock()
user_id = "testuser@example.com"
existing_user = LiteLLM_UserTable(
user_id=user_id,
user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
)
await cache.async_set_cache(key=user_id, value=existing_user.model_dump(), ttl=60)
sso_values = self._make_sso_values(
user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
)
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={"sub": user_id, "custom_roles": ["my-admin"]},
user_info=existing_user,
prisma_client=prisma,
user_api_key_cache=cache,
user_defined_values=sso_values,
)
prisma.db.litellm_usertable.update.assert_called_once_with(
where={"user_id": user_id},
data={"user_role": LitellmUserRoles.PROXY_ADMIN.value},
)
assert existing_user.user_role == LitellmUserRoles.PROXY_ADMIN.value
assert sso_values["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
@pytest.mark.asyncio
async def test_same_role_no_db_write(self):
"""No DB update when the mapped role matches the existing role."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import LitellmUserRoles
handler = self._make_jwt_handler()
prisma = AsyncMock()
prisma.db.litellm_usertable.update = AsyncMock()
existing_user = LiteLLM_UserTable(
user_id="testuser@example.com",
user_role=LitellmUserRoles.PROXY_ADMIN.value,
)
sso_values = self._make_sso_values(
user_role=LitellmUserRoles.PROXY_ADMIN.value,
)
await _sync_user_role_from_jwt_role_map(
jwt_handler=handler,
received_response={"sub": "testuser@example.com", "custom_roles": ["my-admin"]},
user_info=existing_user,
prisma_client=prisma,
user_api_key_cache=DualCache(),
user_defined_values=sso_values,
)
prisma.db.litellm_usertable.update.assert_not_called()