mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: correct EntraID app roles JWT claim extraction
- Fix get_app_roles_from_id_token to use 'app_roles' claim instead of 'roles' - Add comprehensive unit tests for EntraID app roles functionality - Prevent regressions in Microsoft EntraID SSO authentication Resolves issue where EntraID app roles were not being extracted correctly from JWT tokens, causing authentication failures for users with assigned app roles in Microsoft EntraID.
This commit is contained in:
parent
f1fbae8409
commit
6b166e046d
2 changed files with 230 additions and 15 deletions
|
|
@ -875,9 +875,9 @@ async def insert_sso_user(
|
|||
if user_defined_values.get("max_budget") is None:
|
||||
user_defined_values["max_budget"] = litellm.max_internal_user_budget
|
||||
if user_defined_values.get("budget_duration") is None:
|
||||
user_defined_values["budget_duration"] = (
|
||||
litellm.internal_user_budget_duration
|
||||
)
|
||||
user_defined_values[
|
||||
"budget_duration"
|
||||
] = litellm.internal_user_budget_duration
|
||||
|
||||
if user_defined_values["user_role"] is None:
|
||||
user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
|
||||
|
|
@ -1118,9 +1118,9 @@ class SSOAuthenticationHandler:
|
|||
generic_authorization_endpoint
|
||||
and "okta" in generic_authorization_endpoint
|
||||
):
|
||||
redirect_params["state"] = (
|
||||
uuid.uuid4().hex
|
||||
) # set state param for okta - required
|
||||
redirect_params[
|
||||
"state"
|
||||
] = uuid.uuid4().hex # set state param for okta - required
|
||||
|
||||
return redirect_params
|
||||
|
||||
|
|
@ -1699,9 +1699,9 @@ class MicrosoftSSOHandler:
|
|||
|
||||
# if user is trying to get the raw sso response for debugging, return the raw sso response
|
||||
if return_raw_sso_response:
|
||||
original_msft_result[MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY] = (
|
||||
user_team_ids
|
||||
)
|
||||
original_msft_result[
|
||||
MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY
|
||||
] = user_team_ids
|
||||
original_msft_result["app_roles"] = app_roles
|
||||
return original_msft_result or {}
|
||||
|
||||
|
|
@ -1739,7 +1739,7 @@ class MicrosoftSSOHandler:
|
|||
Extract app roles from the Microsoft Entra ID (Azure AD) id_token JWT.
|
||||
|
||||
App roles are assigned in the Azure AD Enterprise Application and appear
|
||||
in the 'roles' claim of the id_token.
|
||||
in the 'app_roles' claim of the id_token.
|
||||
|
||||
Args:
|
||||
id_token (Optional[str]): The JWT id_token from Microsoft SSO
|
||||
|
|
@ -1758,8 +1758,8 @@ class MicrosoftSSOHandler:
|
|||
# (signature is already verified by fastapi_sso)
|
||||
decoded_token = jwt.decode(id_token, options={"verify_signature": False})
|
||||
|
||||
# Extract roles claim from the token
|
||||
roles = decoded_token.get("roles", [])
|
||||
# Extract app_roles claim from the token
|
||||
roles = decoded_token.get("app_roles", [])
|
||||
|
||||
if roles and isinstance(roles, list):
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -1817,9 +1817,9 @@ class MicrosoftSSOHandler:
|
|||
|
||||
# Fetch user membership from Microsoft Graph API
|
||||
all_group_ids = []
|
||||
next_link: Optional[str] = (
|
||||
MicrosoftSSOHandler.graph_api_user_groups_endpoint
|
||||
)
|
||||
next_link: Optional[
|
||||
str
|
||||
] = MicrosoftSSOHandler.graph_api_user_groups_endpoint
|
||||
auth_headers = {"Authorization": f"Bearer {access_token}"}
|
||||
page_count = 0
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,215 @@
|
|||
"""
|
||||
Unit tests for EntraID app roles JWT claim extraction.
|
||||
|
||||
This module tests the get_app_roles_from_id_token method to ensure it correctly
|
||||
extracts app roles from Microsoft EntraID JWT tokens and prevents regressions.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
import jwt
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import MicrosoftSSOHandler
|
||||
|
||||
|
||||
class TestEntraIDAppRoles:
|
||||
"""Test EntraID app roles extraction from JWT tokens"""
|
||||
|
||||
@pytest.fixture
|
||||
def sample_jwt_token(self):
|
||||
"""Create a sample JWT token with app_roles claim"""
|
||||
payload = {
|
||||
"sub": "user123",
|
||||
"email": "user@company.com",
|
||||
"app_roles": ["proxy_admin"],
|
||||
"aud": "litellm-app",
|
||||
"iss": "https://login.microsoftonline.com/tenant-id/v2.0",
|
||||
"exp": 9999999999,
|
||||
}
|
||||
return jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
@pytest.fixture
|
||||
def sample_jwt_token_single_role(self):
|
||||
"""Create a sample JWT token with single app role"""
|
||||
payload = {
|
||||
"sub": "user456",
|
||||
"email": "admin@company.com",
|
||||
"app_roles": ["proxy_admin_viewer"],
|
||||
"aud": "litellm-app",
|
||||
"iss": "https://login.microsoftonline.com/tenant-id/v2.0",
|
||||
"exp": 9999999999,
|
||||
}
|
||||
return jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
@pytest.fixture
|
||||
def sample_jwt_token_no_roles(self):
|
||||
"""Create a sample JWT token without app_roles claim"""
|
||||
payload = {
|
||||
"sub": "user789",
|
||||
"email": "user@company.com",
|
||||
"aud": "litellm-app",
|
||||
"iss": "https://login.microsoftonline.com/tenant-id/v2.0",
|
||||
"exp": 9999999999,
|
||||
}
|
||||
return jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
@pytest.fixture
|
||||
def sample_jwt_token_empty_roles(self):
|
||||
"""Create a sample JWT token with empty app_roles array"""
|
||||
payload = {
|
||||
"sub": "user000",
|
||||
"email": "user@company.com",
|
||||
"app_roles": [],
|
||||
"aud": "litellm-app",
|
||||
"iss": "https://login.microsoftonline.com/tenant-id/v2.0",
|
||||
"exp": 9999999999,
|
||||
}
|
||||
return jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
def test_get_app_roles_from_id_token_single_role(
|
||||
self, sample_jwt_token_single_role
|
||||
):
|
||||
"""Test extracting single app role from JWT token"""
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(
|
||||
sample_jwt_token_single_role
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result == ["proxy_admin_viewer"]
|
||||
assert len(result) == 1
|
||||
|
||||
def test_get_app_roles_from_id_token_no_roles_claim(
|
||||
self, sample_jwt_token_no_roles
|
||||
):
|
||||
"""Test handling JWT token without app_roles claim"""
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(
|
||||
sample_jwt_token_no_roles
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_empty_roles(
|
||||
self, sample_jwt_token_empty_roles
|
||||
):
|
||||
"""Test handling JWT token with empty app_roles array"""
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(
|
||||
sample_jwt_token_empty_roles
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_none_input(self):
|
||||
"""Test handling None input"""
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(None)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_empty_string(self):
|
||||
"""Test handling empty string input"""
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token("")
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_invalid_jwt(self):
|
||||
"""Test handling invalid JWT token"""
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token("invalid.jwt.token")
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_malformed_roles(self):
|
||||
"""Test handling JWT with malformed app_roles (not a list)"""
|
||||
# Arrange
|
||||
payload = {
|
||||
"sub": "user123",
|
||||
"app_roles": "not_a_list", # Should be a list
|
||||
"exp": 9999999999,
|
||||
}
|
||||
malformed_token = jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(malformed_token)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_jwt_decode_exception(self):
|
||||
"""Test handling JWT decode exceptions gracefully"""
|
||||
# Arrange
|
||||
invalid_token = "completely.invalid.token"
|
||||
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(invalid_token)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_import_error(self):
|
||||
"""Test handling import error for jwt library"""
|
||||
# Arrange
|
||||
with patch(
|
||||
"builtins.__import__", side_effect=ImportError("No module named 'jwt'")
|
||||
):
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token("any.token")
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
assert len(result) == 0
|
||||
|
||||
def test_get_app_roles_from_id_token_uses_correct_claim_name(
|
||||
self, sample_jwt_token
|
||||
):
|
||||
"""Test that the method uses 'app_roles' claim, not 'roles' claim"""
|
||||
# This test ensures we don't regress to the old bug where 'roles' was used
|
||||
|
||||
# Arrange - Create a token with both claims to verify correct one is used
|
||||
payload = {
|
||||
"sub": "user123",
|
||||
"roles": ["old_roles_claim"], # This should be ignored
|
||||
"app_roles": ["proxy_admin"], # This should be used
|
||||
"exp": 9999999999,
|
||||
}
|
||||
token_with_both_claims = jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(token_with_both_claims)
|
||||
|
||||
# Assert
|
||||
assert result == ["proxy_admin"] # Should use app_roles, not roles
|
||||
assert "old_roles_claim" not in result
|
||||
|
||||
def test_get_app_roles_from_id_token_case_sensitivity(self):
|
||||
"""Test that app roles are extracted as-is (case sensitive)"""
|
||||
# Arrange
|
||||
payload = {
|
||||
"sub": "user123",
|
||||
"app_roles": ["PROXY_ADMIN", "Internal_User"], # Mixed case
|
||||
"exp": 9999999999,
|
||||
}
|
||||
mixed_case_token = jwt.encode(payload, "secret", algorithm="HS256")
|
||||
|
||||
# Act
|
||||
result = MicrosoftSSOHandler.get_app_roles_from_id_token(mixed_case_token)
|
||||
|
||||
# Assert
|
||||
assert result == ["PROXY_ADMIN", "Internal_User"]
|
||||
assert "PROXY_ADMIN" in result
|
||||
assert "Internal_User" in result
|
||||
Loading…
Add table
Reference in a new issue