feat: Extract roles from id_token for Generic SSO

Co-authored-by: ishaan <ishaan@berri.ai>
This commit is contained in:
Cursor Agent 2026-01-08 09:14:30 +00:00
parent 6c00f6f342
commit bbb3a51c92
2 changed files with 360 additions and 0 deletions

View file

@ -291,6 +291,49 @@ async def google_login(
return HTMLResponse(content=html_form, status_code=200)
def get_roles_from_id_token(id_token: Optional[str]) -> List[str]:
"""
Extract roles from an OIDC id_token JWT.
Many OIDC providers (including Azure AD when used via Generic SSO) include
roles in the 'roles' claim of the id_token rather than in the userinfo endpoint.
Args:
id_token (Optional[str]): The JWT id_token from the SSO provider
Returns:
List[str]: List of role names found in the id_token
"""
if not id_token:
verbose_proxy_logger.debug("No id_token provided for role extraction")
return []
try:
import jwt
# Decode the JWT without signature verification
# (signature is already verified by fastapi_sso)
decoded_token = jwt.decode(id_token, options={"verify_signature": False})
# Check for 'roles' claim (common in Azure AD and other OIDC providers)
roles = decoded_token.get("roles", [])
if roles and isinstance(roles, list):
verbose_proxy_logger.debug(
f"Found {len(roles)} role(s) in id_token: {roles}"
)
return roles
else:
verbose_proxy_logger.debug(
"No roles found in id_token or roles claim is not a list"
)
return []
except Exception as e:
verbose_proxy_logger.error(f"Error extracting roles from id_token: {e}")
return []
def generic_response_convertor(
response,
jwt_handler: JWTHandler,
@ -570,6 +613,29 @@ async def get_generic_sso_response(
access_token_str: Optional[str] = generic_sso.access_token
process_sso_jwt_access_token(access_token_str, sso_jwt_handler, result)
# Extract roles from id_token if available
# This is important for providers like Azure AD that include roles in the id_token
# rather than in the userinfo endpoint response
id_token_str: Optional[str] = getattr(generic_sso, "id_token", None)
if id_token_str and result is not None:
id_token_roles = get_roles_from_id_token(id_token_str)
if id_token_roles:
verbose_proxy_logger.debug(
f"Extracted roles from id_token: {id_token_roles}"
)
# If result doesn't have a user_role set, try to set it from id_token roles
current_user_role = getattr(result, "user_role", None)
if current_user_role is None:
# Check if any id_token role is a valid LitellmUserRoles
for role_str in id_token_roles:
role = get_litellm_user_role(role_str)
if role is not None:
result.user_role = role
verbose_proxy_logger.debug(
f"Set user_role to '{role.value}' from id_token roles"
)
break
except Exception as e:
verbose_proxy_logger.exception(
f"Error verifying and processing generic SSO: {e}. Passed in headers: {additional_generic_sso_headers_dict}"

View file

@ -3386,3 +3386,297 @@ class TestSSOReadinessEndpoint:
)
finally:
app.dependency_overrides.clear()
class TestGetRolesFromIdToken:
"""Test get_roles_from_id_token function for Generic SSO role extraction from id_token"""
def test_extracts_roles_from_valid_id_token(self):
"""
Test that get_roles_from_id_token correctly extracts roles from a valid JWT id_token.
This is important for Azure AD which returns roles in the id_token.
"""
import base64
import json
from litellm.proxy.management_endpoints.ui_sso import get_roles_from_id_token
# Create a mock JWT with roles claim (header.payload.signature)
payload = {
"sub": "user123",
"email": "test@example.com",
"roles": ["proxy_admin", "internal_user"],
}
payload_json = json.dumps(payload)
payload_b64 = base64.urlsafe_b64encode(payload_json.encode()).decode().rstrip("=")
# Minimal JWT structure: header.payload.signature
header = base64.urlsafe_b64encode(json.dumps({"alg": "none"}).encode()).decode().rstrip("=")
mock_id_token = f"{header}.{payload_b64}."
# Act
result = get_roles_from_id_token(mock_id_token)
# Assert
assert result == ["proxy_admin", "internal_user"]
def test_returns_empty_list_when_no_roles_claim(self):
"""
Test that get_roles_from_id_token returns empty list when no roles claim exists.
"""
import base64
import json
from litellm.proxy.management_endpoints.ui_sso import get_roles_from_id_token
# Create a mock JWT without roles claim
payload = {
"sub": "user123",
"email": "test@example.com",
}
payload_json = json.dumps(payload)
payload_b64 = base64.urlsafe_b64encode(payload_json.encode()).decode().rstrip("=")
header = base64.urlsafe_b64encode(json.dumps({"alg": "none"}).encode()).decode().rstrip("=")
mock_id_token = f"{header}.{payload_b64}."
# Act
result = get_roles_from_id_token(mock_id_token)
# Assert
assert result == []
def test_returns_empty_list_when_id_token_is_none(self):
"""
Test that get_roles_from_id_token returns empty list when id_token is None.
"""
from litellm.proxy.management_endpoints.ui_sso import get_roles_from_id_token
# Act
result = get_roles_from_id_token(None)
# Assert
assert result == []
def test_returns_empty_list_when_id_token_is_empty_string(self):
"""
Test that get_roles_from_id_token returns empty list when id_token is empty string.
"""
from litellm.proxy.management_endpoints.ui_sso import get_roles_from_id_token
# Act
result = get_roles_from_id_token("")
# Assert
assert result == []
def test_returns_empty_list_when_roles_is_not_a_list(self):
"""
Test that get_roles_from_id_token returns empty list when roles claim is not a list.
"""
import base64
import json
from litellm.proxy.management_endpoints.ui_sso import get_roles_from_id_token
# Create a mock JWT with roles as a string instead of list
payload = {
"sub": "user123",
"roles": "proxy_admin", # String instead of list
}
payload_json = json.dumps(payload)
payload_b64 = base64.urlsafe_b64encode(payload_json.encode()).decode().rstrip("=")
header = base64.urlsafe_b64encode(json.dumps({"alg": "none"}).encode()).decode().rstrip("=")
mock_id_token = f"{header}.{payload_b64}."
# Act
result = get_roles_from_id_token(mock_id_token)
# Assert
assert result == []
def test_handles_invalid_jwt_gracefully(self):
"""
Test that get_roles_from_id_token handles invalid JWT tokens gracefully.
"""
from litellm.proxy.management_endpoints.ui_sso import get_roles_from_id_token
# Act
result = get_roles_from_id_token("invalid.jwt.token")
# Assert
assert result == []
def test_extracts_single_role(self):
"""
Test that get_roles_from_id_token correctly extracts a single role from id_token.
This mimics Azure AD app roles scenario.
"""
import base64
import json
from litellm.proxy.management_endpoints.ui_sso import get_roles_from_id_token
# Create a mock JWT with single role in list
payload = {
"sub": "user123",
"email": "test@example.com",
"roles": ["proxy_admin"],
}
payload_json = json.dumps(payload)
payload_b64 = base64.urlsafe_b64encode(payload_json.encode()).decode().rstrip("=")
header = base64.urlsafe_b64encode(json.dumps({"alg": "none"}).encode()).decode().rstrip("=")
mock_id_token = f"{header}.{payload_b64}."
# Act
result = get_roles_from_id_token(mock_id_token)
# Assert
assert result == ["proxy_admin"]
class TestGenericSSOIdTokenRoleExtraction:
"""Test that Generic SSO extracts roles from id_token when userinfo doesn't have roles"""
@pytest.mark.asyncio
async def test_get_generic_sso_response_extracts_role_from_id_token(self):
"""
Test that get_generic_sso_response extracts user role from id_token when
the userinfo response doesn't contain the role.
This is the fix for Azure AD users using Generic SSO.
"""
import base64
import json
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
# Create mock id_token with proxy_admin role
payload = {
"sub": "user123",
"email": "test@example.com",
"roles": ["proxy_admin"],
}
payload_json = json.dumps(payload)
payload_b64 = base64.urlsafe_b64encode(payload_json.encode()).decode().rstrip("=")
header = base64.urlsafe_b64encode(json.dumps({"alg": "none"}).encode()).decode().rstrip("=")
mock_id_token = f"{header}.{payload_b64}."
# Create mock result without user_role (simulating userinfo response without roles)
mock_result = CustomOpenID(
id="user123",
email="test@example.com",
display_name="Test User",
team_ids=[],
user_role=None, # No role from userinfo
)
# Mock the generic_sso object
mock_generic_sso = MagicMock()
mock_generic_sso.verify_and_process = AsyncMock(return_value=mock_result)
mock_generic_sso.access_token = "mock_access_token"
mock_generic_sso.id_token = mock_id_token # id_token contains the role
mock_request = MagicMock(spec=Request)
mock_request.query_params = {}
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_ID": "test-client-id",
"GENERIC_CLIENT_SECRET": "test-secret",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/authorize",
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
},
):
with patch(
"fastapi_sso.sso.generic.create_provider"
) as mock_create_provider:
mock_create_provider.return_value = lambda **kwargs: mock_generic_sso
result, received_response = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
generic_client_id="test-client-id",
redirect_url="https://example.com/callback",
)
# Assert that the user_role was extracted from id_token
assert isinstance(result, CustomOpenID)
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
@pytest.mark.asyncio
async def test_get_generic_sso_response_does_not_override_existing_role(self):
"""
Test that get_generic_sso_response does NOT override user_role if already set
from userinfo response.
"""
import base64
import json
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import CustomOpenID
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
# Create mock id_token with proxy_admin role
payload = {
"sub": "user123",
"roles": ["proxy_admin"],
}
payload_json = json.dumps(payload)
payload_b64 = base64.urlsafe_b64encode(payload_json.encode()).decode().rstrip("=")
header = base64.urlsafe_b64encode(json.dumps({"alg": "none"}).encode()).decode().rstrip("=")
mock_id_token = f"{header}.{payload_b64}."
# Create mock result WITH user_role already set (from userinfo or role_mappings)
mock_result = CustomOpenID(
id="user123",
email="test@example.com",
display_name="Test User",
team_ids=[],
user_role=LitellmUserRoles.INTERNAL_USER, # Role already set
)
mock_generic_sso = MagicMock()
mock_generic_sso.verify_and_process = AsyncMock(return_value=mock_result)
mock_generic_sso.access_token = "mock_access_token"
mock_generic_sso.id_token = mock_id_token
mock_request = MagicMock(spec=Request)
mock_request.query_params = {}
mock_jwt_handler = MagicMock(spec=JWTHandler)
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
with patch.dict(
os.environ,
{
"GENERIC_CLIENT_ID": "test-client-id",
"GENERIC_CLIENT_SECRET": "test-secret",
"GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/authorize",
"GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token",
"GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo",
},
):
with patch(
"fastapi_sso.sso.generic.create_provider"
) as mock_create_provider:
mock_create_provider.return_value = lambda **kwargs: mock_generic_sso
result, received_response = await get_generic_sso_response(
request=mock_request,
jwt_handler=mock_jwt_handler,
sso_jwt_handler=None,
generic_client_id="test-client-id",
redirect_url="https://example.com/callback",
)
# Assert that the original role was NOT overridden
assert isinstance(result, CustomOpenID)
assert result.user_role == LitellmUserRoles.INTERNAL_USER