diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 5adf54c1627..1756fe9c6f6 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -99,24 +99,24 @@ def determine_role_from_groups( ) -> Optional[LitellmUserRoles]: """ Determine the highest privilege role for a user based on their groups. - + Role hierarchy (highest to lowest): - proxy_admin - proxy_admin_viewer - internal_user - internal_user_viewer - + Args: user_groups: List of group names from the SSO token role_mappings: RoleMappings configuration object - + Returns: The highest privilege role found, or default_role if no matches, or None """ if not role_mappings.roles: # No role mappings configured, return default_role return role_mappings.default_role - + # Role hierarchy (highest to lowest) role_hierarchy = [ LitellmUserRoles.PROXY_ADMIN, @@ -124,20 +124,22 @@ def determine_role_from_groups( LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] - + # Convert user_groups to a set for efficient lookup user_groups_set = set(user_groups) if isinstance(user_groups, list) else set() - + # Find the highest privilege role the user belongs to for role in role_hierarchy: if role in role_mappings.roles: role_groups = role_mappings.roles[role] - if isinstance(role_groups, list) and user_groups_set.intersection(set(role_groups)): + if isinstance(role_groups, list) and user_groups_set.intersection( + set(role_groups) + ): verbose_proxy_logger.debug( f"User groups {user_groups} matched role '{role.value}' via groups: {role_groups}" ) return role - + # No matching groups found, return default_role verbose_proxy_logger.debug( f"User groups {user_groups} did not match any role mappings, using default_role: {role_mappings.default_role}" @@ -326,9 +328,7 @@ def generic_response_convertor( "GENERIC_USER_PROVIDER_ATTRIBUTE", "provider" ) - generic_user_role_attribute_name = os.getenv( - "GENERIC_USER_ROLE_ATTRIBUTE", "role" - ) + generic_user_role_attribute_name = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role") verbose_proxy_logger.debug( f" generic_user_id_attribute_name: {generic_user_id_attribute_name}\n generic_user_email_attribute_name: {generic_user_email_attribute_name}" @@ -345,12 +345,15 @@ def generic_response_convertor( # Determine user role based on role_mappings if available # Only apply role_mappings for GENERIC SSO provider user_role: Optional[LitellmUserRoles] = None - - if role_mappings is not None and role_mappings.provider.lower() in ["generic", "okta"]: + + if role_mappings is not None and role_mappings.provider.lower() in [ + "generic", + "okta", + ]: # Use role_mappings to determine role from groups group_claim = role_mappings.group_claim user_groups_raw: Any = get_nested_value(response, group_claim) - + # Handle different formats: could be a list, string (comma-separated), or single value user_groups: List[str] = [] if isinstance(user_groups_raw, list): @@ -361,7 +364,7 @@ def generic_response_convertor( elif user_groups_raw is not None: # Single value user_groups = [str(user_groups_raw)] - + if user_groups: user_role = determine_role_from_groups(user_groups, role_mappings) verbose_proxy_logger.debug( @@ -373,10 +376,12 @@ def generic_response_convertor( verbose_proxy_logger.debug( f"No groups found in '{group_claim}', using default_role: {role_mappings.default_role}" ) - + # Fallback to existing logic if role_mappings not used if user_role is None: - user_role_from_sso = get_nested_value(response, generic_user_role_attribute_name) + user_role_from_sso = get_nested_value( + response, generic_user_role_attribute_name + ) if user_role_from_sso is not None: role = get_litellm_user_role(user_role_from_sso) if role is not None: @@ -399,7 +404,9 @@ def generic_response_convertor( ) -def _setup_generic_sso_env_vars(generic_client_id: str, redirect_url: str) -> Tuple[str, List[str], str, str, str, bool]: +def _setup_generic_sso_env_vars( + generic_client_id: str, redirect_url: str +) -> Tuple[str, List[str], str, str, str, bool]: """Setup and validate Generic SSO environment variables.""" generic_client_secret = os.getenv("GENERIC_CLIENT_SECRET", None) generic_scope = os.getenv("GENERIC_SCOPE", "openid email profile").split(" ") @@ -492,7 +499,43 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: verbose_proxy_logger.debug( f"Could not load role_mappings from database: {e}. Continuing with existing role logic." ) + + generic_role_mappings = os.getenv("GENERIC_ROLE_MAPPINGS_ROLES", None) + generic_role_mappings_group_claim = os.getenv( + "GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None + ) + generic_role_mappoings_default_role = os.getenv( + "GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None + ) + if generic_role_mappings is not None: + verbose_proxy_logger.debug( + "Found role_mappings for generic provider in environment variables" + ) + import ast + try: + generic_user_role_mappings_data: Dict[ + LitellmUserRoles, List[str] + ] = ast.literal_eval(generic_role_mappings) + if isinstance(generic_user_role_mappings_data, dict): + from litellm.types.proxy.management_endpoints.ui_sso import ( + RoleMappings, + ) + + role_mappings_data = { + "provider": "generic", + "group_claim": generic_role_mappings_group_claim, + "default_role": generic_role_mappoings_default_role, + "roles": generic_user_role_mappings_data, + } + + role_mappings = RoleMappings(**role_mappings_data) + verbose_proxy_logger.debug( + f"Loaded role_mappings from environments for provider '{role_mappings.provider}'." + ) + return role_mappings + except TypeError as e: + verbose_proxy_logger.warning(f"Error decoding role mappings from environment variables: {e}. Continuing with existing role logic.") return role_mappings @@ -529,7 +572,7 @@ async def get_generic_sso_response( # Get role_mappings from SSO settings if available role_mappings = await _setup_role_mappings() - + def response_convertor(response, client): nonlocal received_response # return for user debugging received_response = response @@ -1217,20 +1260,24 @@ async def insert_sso_user( role_mappings_configured = False 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" ) - + # Get SSO config from dedicated table 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) role_mappings_data = sso_settings_dict.get("role_mappings") role_mappings_configured = role_mappings_data is not None + generic_user_role_mappings = os.getenv("GENERIC_USER_ROLE_MAPPINGS", None) + if generic_user_role_mappings is not None: + role_mappings_configured = True + except Exception as e: # If we can't check role_mappings, continue with existing logic verbose_proxy_logger.debug( @@ -1240,7 +1287,10 @@ async def insert_sso_user( # Apply default_internal_user_params if litellm.default_internal_user_params: # If role_mappings is configured and user_role is already set from SSO, preserve it - if role_mappings_configured and user_defined_values.get("user_role") is not None: + if ( + role_mappings_configured + and user_defined_values.get("user_role") is not None + ): # Preserve the SSO-extracted role, but apply other defaults preserved_role = user_defined_values.get("user_role") user_defined_values.update(litellm.default_internal_user_params) # type: ignore diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index ad4f53dac4b..60bdb7d12cb 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -1227,3 +1227,113 @@ class TestProxySettingEndpoints: assert retrieved_role_mappings["provider"] == "google" assert retrieved_role_mappings["group_claim"] == "groups" assert retrieved_role_mappings["default_role"] == LitellmUserRoles.INTERNAL_USER + + def test_setup_role_mappings_custom_logic_with_env_vars(self, monkeypatch): + """Test the _setup_role_mappings function directly with custom role mapping logic from environment variables""" + import asyncio + import os + from litellm.proxy.management_endpoints.ui_sso import _setup_role_mappings + from litellm.proxy._types import LitellmUserRoles + + # Set up environment variables for custom role mappings using valid Python dict format + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['custom-admin-group'], 'internal_user': ['custom-user-group'], 'proxy_admin_viewer': ['custom-viewer-group']}") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "custom-groups") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", "internal_user_viewer") + + # Debug: Print environment variables + print("GENERIC_ROLE_MAPPINGS_ROLES:", os.getenv("GENERIC_ROLE_MAPPINGS_ROLES")) + print("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM:", os.getenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM")) + print("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE:", os.getenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE")) + + # Run the async function + role_mappings = asyncio.run(_setup_role_mappings()) + + # Debug: Print result + print("role_mappings result:", role_mappings) + + # Verify role_mappings is returned correctly from environment variables + assert role_mappings is not None + assert role_mappings.provider == "generic" + assert role_mappings.group_claim == "custom-groups" + assert role_mappings.default_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY + assert role_mappings.roles[LitellmUserRoles.PROXY_ADMIN] == ["custom-admin-group"] + assert role_mappings.roles[LitellmUserRoles.INTERNAL_USER] == ["custom-user-group"] + assert role_mappings.roles[LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY] == ["custom-viewer-group"] + + def test_setup_role_mappings_custom_logic_with_no_config(self, monkeypatch): + """Test the _setup_role_mappings function returns None when no configuration is available""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + from litellm.proxy.management_endpoints.ui_sso import _setup_role_mappings + + # Ensure environment variables are not set + monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_ROLES", raising=False) + monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", raising=False) + monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", raising=False) + + # Mock the prisma client to return None (no database record) + mock_prisma = MagicMock() + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + # Run the async function + role_mappings = asyncio.run(_setup_role_mappings()) + + # Should return None when no configuration is available + assert role_mappings is None + + def test_get_sso_settings_with_env_role_mappings(self, mock_proxy_config, mock_auth, monkeypatch): + import json + from unittest.mock import AsyncMock, MagicMock + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", '{"proxy_admin": ["custom-admin-group"], "internal_user": ["custom-user-group"], "proxy_admin_viewer": ["custom-viewer-group"]}') + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "custom-groups") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", "internal_user_viewer") + + mock_prisma = MagicMock() + mock_db_record = MagicMock() + mock_db_record.sso_settings = { + "google_client_id": "test_google_client_id", + "role_mappings": { + "provider": "google", + "group_claim": "db-groups", + "default_role": "proxy_admin", + "roles": { + "proxy_admin": ["db-admin-group"], + }, + }, + } + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_db_record) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + from litellm.proxy.proxy_server import proxy_config + monkeypatch.setattr( + proxy_config, "_decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables + ) + + response = client.get("/get/sso_settings") + + assert response.status_code == 200 + data = response.json() + + values = data["values"] + assert "role_mappings" in values + assert values["role_mappings"] is not None + + # The database values shoeld override the environment variables + assert values["role_mappings"]["provider"] == "google" + assert values["role_mappings"]["group_claim"] == "db-groups" + assert values["role_mappings"]["default_role"] == LitellmUserRoles.PROXY_ADMIN + assert values["role_mappings"]["roles"][LitellmUserRoles.PROXY_ADMIN] == ["db-admin-group"] + + # Verify that the database was checked but environment variables took priority + mock_prisma.db.litellm_ssoconfig.find_unique.assert_called_once_with( + where={"id": "sso_config"} + ) + + # Verify other SSO settings are still correctly returned + assert values["google_client_id"] == "test_google_client_id" + + # Verify field_schema is still present + assert "field_schema" in data + assert "properties" in data["field_schema"] + assert "role_mappings" in data["field_schema"]["properties"]