mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
feat: Extract roles from id_token for Generic SSO
Co-authored-by: ishaan <ishaan@berri.ai>
This commit is contained in:
parent
6c00f6f342
commit
bbb3a51c92
2 changed files with 360 additions and 0 deletions
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue