fix(sso): pass decoded JWT access token to role mapping during SSO login

During SSO login, bearer tokens are stripped from the OAuth response
before role mapping runs. Custom role claims encoded inside the JWT
access token are lost, so map_jwt_role_to_litellm_role() returns None
and the user falls back to internal_user_viewer.

process_sso_jwt_access_token() now returns the decoded JWT payload, and
a new _sync_user_role_from_jwt_role_map() receives it so
jwt_litellm_role_map works correctly during SSO login.
This commit is contained in:
Ryan Crabbe
2026-03-27 13:50:30 -07:00
parent 25feae9f0f
commit e24819afef
2 changed files with 246 additions and 9 deletions
+87 -7
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,
@@ -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()