Merge pull request #20111 from BerriAI/litellm_sso_map_teams

[Feature] SSO Config Team Mappings
This commit is contained in:
yuneng-jiang 2026-02-02 14:18:25 -08:00 • committed by GitHub
commit f1227ce5a8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 120 additions and 7 deletions

View file

@ -326,6 +326,7 @@ def generic_response_convertor(
jwt_handler: JWTHandler,
sso_jwt_handler: Optional[JWTHandler] = None,
role_mappings: Optional["RoleMappings"] = None,
team_mappings: Optional["TeamMappings"] = None,
) -> CustomOpenID:
generic_user_id_attribute_name = os.getenv(
"GENERIC_USER_ID_ATTRIBUTE", "preferred_username"
@ -359,8 +360,20 @@ def generic_response_convertor(
team_ids = sso_jwt_handler.get_team_ids_from_jwt(cast(dict, response))
all_teams.extend(team_ids)
team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response))
all_teams.extend(team_ids)
if team_mappings is not None and team_mappings.team_ids_jwt_field is not None:
team_ids_from_db_mapping: Optional[List[str]] = get_nested_value(
data=cast(dict, response),
key_path=team_mappings.team_ids_jwt_field,
default=[],
)
if team_ids_from_db_mapping:
all_teams.extend(team_ids_from_db_mapping)
verbose_proxy_logger.debug(
f"Loaded team_ids from DB team_mappings.team_ids_jwt_field='{team_mappings.team_ids_jwt_field}': {team_ids_from_db_mapping}"
)
else:
team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response))
all_teams.extend(team_ids)
# Determine user role based on role_mappings if available
# Only apply role_mappings for GENERIC SSO provider
@ -484,6 +497,43 @@ def _setup_generic_sso_env_vars(
)
async def _setup_team_mappings() -> Optional["TeamMappings"]:
"""Setup team mappings from SSO database settings."""
team_mappings: Optional["TeamMappings"] = None
try:
from litellm.proxy.utils import get_prisma_client_or_throw
prisma_client = get_prisma_client_or_throw(
"Prisma client is None, connect a database to your proxy"
)
sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique(
where={"id": "sso_config"}
)
if sso_db_record and sso_db_record.sso_settings:
sso_settings_dict = dict(sso_db_record.sso_settings)
team_mappings_data = sso_settings_dict.get("team_mappings")
if team_mappings_data:
from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings
if isinstance(team_mappings_data, dict):
team_mappings = TeamMappings(**team_mappings_data)
elif isinstance(team_mappings_data, TeamMappings):
team_mappings = team_mappings_data
if team_mappings and team_mappings.team_ids_jwt_field:
verbose_proxy_logger.debug(
f"Loaded team_mappings with team_ids_jwt_field: '{team_mappings.team_ids_jwt_field}'"
)
except Exception as e:
verbose_proxy_logger.debug(
f"Could not load team_mappings from database: {e}. Continuing with config-based team mapping."
)
return team_mappings
async def _setup_role_mappings() -> Optional["RoleMappings"]:
"""Setup role mappings from SSO database settings."""
role_mappings: Optional["RoleMappings"] = None
@ -494,7 +544,6 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
"Prisma client is None, connect a database to your proxy"
)
# Get SSO config from dedicated table
sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique(
where={"id": "sso_config"}
)
@ -515,7 +564,6 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
f"Loaded role_mappings for provider '{role_mappings.provider}'"
)
except Exception as e:
# If we can't load role_mappings, continue with existing logic
verbose_proxy_logger.debug(
f"Could not load role_mappings from database: {e}. Continuing with existing role logic."
)
@ -590,8 +638,8 @@ async def get_generic_sso_response(
userinfo_endpoint=generic_userinfo_endpoint,
)
# Get role_mappings from SSO settings if available
role_mappings = await _setup_role_mappings()
team_mappings = await _setup_team_mappings()
def response_convertor(response, client):
nonlocal received_response # return for user debugging
@ -601,6 +649,7 @@ async def get_generic_sso_response(
jwt_handler=jwt_handler,
sso_jwt_handler=sso_jwt_handler,
role_mappings=role_mappings,
team_mappings=team_mappings,
)
SSOProvider = create_provider(

View file

@ -3911,8 +3911,9 @@ class ProxyConfig:
where={"id": "sso_config"}
)
if sso_settings is not None:
# Capitalize all keys in sso_settings dictionary
sso_settings.sso_settings.pop("role_mappings", None)
sso_settings.sso_settings.pop("team_mappings", None)
sso_settings.sso_settings.pop("ui_access_mode", None)
uppercase_sso_settings = {
key.upper(): value
for key, value in sso_settings.sso_settings.items()

View file

@ -449,7 +449,6 @@ async def get_sso_settings():
# Load settings from database
sso_settings_dict = dict(sso_db_record.sso_settings)
# Extract role_mappings before removing it (it's a dict, not an env variable)
role_mappings_data = sso_settings_dict.pop("role_mappings", None)
role_mappings = None
if role_mappings_data:
@ -460,6 +459,16 @@ async def get_sso_settings():
elif isinstance(role_mappings_data, RoleMappings):
role_mappings = role_mappings_data
team_mappings_data = sso_settings_dict.pop("team_mappings", None)
team_mappings = None
if team_mappings_data:
from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings
if isinstance(team_mappings_data, dict):
team_mappings = TeamMappings(**team_mappings_data)
elif isinstance(team_mappings_data, TeamMappings):
team_mappings = team_mappings_data
decrypted_sso_settings_dict = proxy_config._decrypt_and_set_db_env_variables(
environment_variables=sso_settings_dict
)
@ -495,6 +504,7 @@ async def get_sso_settings():
user_email=decrypted_sso_settings_dict.get("user_email"),
ui_access_mode=decrypted_sso_settings_dict.get("ui_access_mode"),
role_mappings=role_mappings,
team_mappings=team_mappings,
)
# Get the schema for UI display

View file

@ -86,6 +86,20 @@ class RoleMappings(LiteLLMPydanticObjectBase):
)
class TeamMappings(LiteLLMPydanticObjectBase):
"""
Configuration for mapping SSO JWT fields to team IDs.
This allows configuring team_ids_jwt_field via the database instead of
requiring config file changes and restarts.
"""
team_ids_jwt_field: Optional[str] = Field(
default=None,
description="The field name in the SSO/JWT token that contains the team IDs array (e.g., 'groups', 'teams'). Supports dot notation for nested fields.",
)
class SSOConfig(LiteLLMPydanticObjectBase):
"""
Configuration for SSO environment variables and settings
@ -159,6 +173,12 @@ class SSOConfig(LiteLLMPydanticObjectBase):
description="Configuration for mapping SSO groups to LiteLLM roles based on group claims in the SSO token",
)
# Team Mappings
team_mappings: Optional[TeamMappings] = Field(
default=None,
description="Configuration for mapping SSO JWT fields to team IDs. Takes precedence over config file settings.",
)
class DefaultTeamSSOParams(LiteLLMPydanticObjectBase):
"""

View file

@ -25,12 +25,14 @@ from litellm.proxy.management_endpoints.ui_sso import (
MicrosoftSSOHandler,
SSOAuthenticationHandler,
normalize_email,
_setup_team_mappings,
)
from litellm.types.proxy.management_endpoints.ui_sso import (
DefaultTeamSSOParams,
MicrosoftGraphAPIUserGroupDirectoryObject,
MicrosoftGraphAPIUserGroupResponse,
MicrosoftServicePrincipalTeam,
TeamMappings,
)
@ -3815,3 +3817,34 @@ class TestCustomMicrosoftSSO:
)
assert isinstance(sso, MicrosoftSSO)
@pytest.mark.asyncio
async def test_setup_team_mappings():
"""Test _setup_team_mappings function loads team mappings from database."""
# Arrange
mock_prisma = MagicMock()
mock_sso_config = MagicMock()
mock_sso_config.sso_settings = {
"team_mappings": {
"team_ids_jwt_field": "groups"
}
}
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(
return_value=mock_sso_config
)
with patch(
"litellm.proxy.utils.get_prisma_client_or_throw",
return_value=mock_prisma,
):
# Act
result = await _setup_team_mappings()
# Assert
assert result is not None
assert isinstance(result, TeamMappings)
assert result.team_ids_jwt_field == "groups"
mock_prisma.db.litellm_ssoconfig.find_unique.assert_called_once_with(
where={"id": "sso_config"}
)