From 395ccac6515a7d4135d0dce25f9b48402dbc48ed Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 30 Jan 2026 21:06:03 -0800 Subject: [PATCH] team mapping --- litellm/proxy/management_endpoints/ui_sso.py | 59 +++++++++++++++++-- litellm/proxy/proxy_server.py | 3 +- .../proxy_setting_endpoints.py | 12 +++- .../proxy/management_endpoints/ui_sso.py | 20 +++++++ .../proxy/management_endpoints/test_ui_sso.py | 33 +++++++++++ 5 files changed, 120 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 4048b3731c1..2d248dc81f3 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 078ce0edf27..f991ee4c07f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4007,8 +4007,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() diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 4a0268eeede..30ec0766dbf 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -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 diff --git a/litellm/types/proxy/management_endpoints/ui_sso.py b/litellm/types/proxy/management_endpoints/ui_sso.py index c9d998f6a92..6743c4a5b9b 100644 --- a/litellm/types/proxy/management_endpoints/ui_sso.py +++ b/litellm/types/proxy/management_endpoints/ui_sso.py @@ -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): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 5e9078ea876..41096503a2e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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"} + )