Adding role mappings to SSOConfig DB

This commit is contained in:
yuneng-jiang
2025-12-16 15:48:26 -08:00
parent 9274860aa2
commit 42d7d757a3
5 changed files with 182 additions and 2 deletions
+1 -1
View File
@@ -594,7 +594,7 @@ def _build_sso_user_update_data(
user_id: Optional[str],
) -> dict:
"""
Build the update data dictionary for SSO user upsert
Build the update data dictionary for SSO user upsert.
Args:
result: The SSO response containing user information
+1
View File
@@ -3578,6 +3578,7 @@ class ProxyConfig:
)
if sso_settings is not None:
# Capitalize all keys in sso_settings dictionary
sso_settings.sso_settings.pop("role_mappings", None)
uppercase_sso_settings = {
key.upper(): value
for key, value in sso_settings.sso_settings.items()
@@ -433,10 +433,21 @@ async def get_sso_settings():
if sso_db_record and sso_db_record.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:
from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings
if isinstance(role_mappings_data, dict):
role_mappings = RoleMappings(**role_mappings_data)
elif isinstance(role_mappings_data, RoleMappings):
role_mappings = role_mappings_data
decrypted_sso_settings_dict = proxy_config._decrypt_and_set_db_env_variables(environment_variables=sso_settings_dict)
# Build SSO config with database values or environment fallback
sso_config = SSOConfig(
google_client_id=decrypted_sso_settings_dict.get("google_client_id", None),
google_client_secret=decrypted_sso_settings_dict.get("google_client_secret", None),
@@ -451,6 +462,7 @@ async def get_sso_settings():
proxy_base_url=decrypted_sso_settings_dict.get("proxy_base_url", None),
user_email=decrypted_sso_settings_dict.get("user_email"),
ui_access_mode=decrypted_sso_settings_dict.get("ui_access_mode"),
role_mappings=role_mappings,
)
# Get the schema for UI display
@@ -1,10 +1,12 @@
from typing import List, Literal, Optional, Union
from typing import Dict, List, Literal, Optional, Union
from pydantic import Field
from typing_extensions import TypedDict
from litellm.types.utils import LiteLLMPydanticObjectBase
from litellm.proxy._types import LitellmUserRoles
class LiteLLM_UpperboundKeyGenerateParams(LiteLLMPydanticObjectBase):
"""
@@ -60,6 +62,30 @@ class AccessControl_UI_AccessMode(LiteLLMPydanticObjectBase):
sso_group_jwt_field: str
class RoleMappings(LiteLLMPydanticObjectBase):
"""
Configuration for mapping SSO groups to LiteLLM roles.
The system will look at the group_claim field in the SSO token to determine
which role to assign the user based on the roles mapping.
"""
provider: str = Field(
description="SSO Provider name (e.g., 'google', 'microsoft', 'generic')"
)
group_claim: str = Field(
description="The field name in the SSO token that contains the groups array (e.g., 'groups', 'roles')"
)
default_role: Optional[LitellmUserRoles] = Field(
default=None,
description="Default role to assign if user's groups don't match any role mappings. Must be a valid LitellmUserRoles value (e.g., 'proxy_admin', 'internal_user', 'proxy_admin_viewer')"
)
roles: Dict[LitellmUserRoles, List[str]] = Field(
default_factory=dict,
description="Mapping of LiteLLM role names to arrays of SSO group names. Example: {'proxy_admin': ['group-1', 'group-2'], 'proxy_admin_viewer': ['group-3']}"
)
class SSOConfig(LiteLLMPydanticObjectBase):
"""
Configuration for SSO environment variables and settings
@@ -127,6 +153,12 @@ class SSOConfig(LiteLLMPydanticObjectBase):
description="Access mode for the UI",
)
# Role Mappings
role_mappings: Optional[RoleMappings] = Field(
default=None,
description="Configuration for mapping SSO groups to LiteLLM roles based on group claims in the SSO token",
)
class DefaultTeamSSOParams(LiteLLMPydanticObjectBase):
"""
@@ -290,6 +290,10 @@ class TestProxySettingEndpoints:
assert "google_client_id" in data["field_schema"]["properties"]
assert "description" in data["field_schema"]["properties"]["google_client_id"]
# Verify role_mappings is present in response (can be None if not set)
assert "role_mappings" in values
assert values["role_mappings"] is None
# Verify find_unique was called with correct parameters
mock_prisma.db.litellm_ssoconfig.find_unique.assert_called_once()
call_args = mock_prisma.db.litellm_ssoconfig.find_unique.call_args
@@ -863,6 +867,10 @@ class TestProxySettingEndpoints:
assert values["google_client_secret"] == "decrypted_google_secret"
assert values["microsoft_client_id"] == "decrypted_microsoft_id"
assert values["proxy_base_url"] == "https://decrypted.example.com"
# Verify role_mappings is present in response (can be None if not set)
assert "role_mappings" in values
assert values["role_mappings"] is None
def test_update_sso_settings_to_database(self, mock_proxy_config, mock_auth, monkeypatch):
"""Test updating SSO settings saves to the dedicated database table"""
@@ -1062,6 +1070,7 @@ class TestProxySettingEndpoints:
assert values.get("google_client_id") is None
assert values.get("google_client_secret") is None
assert values.get("microsoft_client_id") is None
assert values.get("role_mappings") is None
def test_update_sso_settings_no_database_connection(self, mock_proxy_config, mock_auth, monkeypatch):
"""Test updating SSO settings when database is not connected"""
@@ -1088,3 +1097,129 @@ class TestProxySettingEndpoints:
data = response.json()
assert "error" in data["detail"]
assert "Database not connected" in data["detail"]["error"]
def test_get_sso_settings_with_role_mappings(self, mock_proxy_config, mock_auth, monkeypatch):
"""Test getting SSO settings when role_mappings is present in database"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LitellmUserRoles
# Mock the prisma client with database record containing role_mappings
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": "groups",
"default_role": LitellmUserRoles.INTERNAL_USER,
"roles": {
LitellmUserRoles.PROXY_ADMIN: ["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)
# Mock decryption to return the values as-is (role_mappings should not be passed to decryption)
from litellm.proxy.proxy_server import proxy_config
def mock_decrypt(environment_variables):
# role_mappings should not be in environment_variables since it's extracted before decryption
assert "role_mappings" not in environment_variables
return environment_variables
monkeypatch.setattr(
proxy_config, "_decrypt_and_set_db_env_variables", mock_decrypt
)
response = client.get("/get/sso_settings")
assert response.status_code == 200
data = response.json()
# Verify role_mappings is returned correctly
values = data["values"]
assert "role_mappings" in values
assert values["role_mappings"] is not None
assert values["role_mappings"]["provider"] == "google"
assert values["role_mappings"]["group_claim"] == "groups"
assert values["role_mappings"]["default_role"] == LitellmUserRoles.INTERNAL_USER
assert values["role_mappings"]["roles"][LitellmUserRoles.PROXY_ADMIN] == ["admin-group"]
def test_role_mappings_stored_and_retrieved(self, mock_proxy_config, mock_auth, monkeypatch):
"""Test that role_mappings is properly stored and retrieved from SSO settings"""
import json
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LitellmUserRoles
monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key")
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
# Mock the prisma client
mock_prisma = MagicMock()
mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock()
mock_prisma.db.litellm_config = MagicMock()
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.update = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
# Mock encryption to return values as-is
from litellm.proxy.proxy_server import proxy_config
monkeypatch.setattr(proxy_config, "_encrypt_env_variables", lambda environment_variables: environment_variables)
# SSO settings with role_mappings
role_mappings_data = {
"provider": "google",
"group_claim": "groups",
"default_role": LitellmUserRoles.INTERNAL_USER,
"roles": {
LitellmUserRoles.PROXY_ADMIN: ["admin-group"],
LitellmUserRoles.INTERNAL_USER: ["user-group"],
},
}
new_sso_settings = {
"google_client_id": "test_google_id",
"role_mappings": role_mappings_data,
}
response = client.patch("/update/sso_settings", json=new_sso_settings)
assert response.status_code == 200
data = response.json()
assert data["status"] == "success"
assert "role_mappings" in data["settings"]
# Verify role_mappings structure in response
returned_role_mappings = data["settings"]["role_mappings"]
assert returned_role_mappings["provider"] == "google"
assert returned_role_mappings["group_claim"] == "groups"
assert returned_role_mappings["default_role"] == LitellmUserRoles.INTERNAL_USER
assert returned_role_mappings["roles"][LitellmUserRoles.PROXY_ADMIN] == ["admin-group"]
# Verify upsert was called with role_mappings in the data
assert mock_prisma.db.litellm_ssoconfig.upsert.called
call_args = mock_prisma.db.litellm_ssoconfig.upsert.call_args
create_data = call_args.kwargs["data"]["create"]
stored_sso_settings = json.loads(create_data["sso_settings"])
assert "role_mappings" in stored_sso_settings
assert stored_sso_settings["role_mappings"]["provider"] == "google"
# Now test retrieving role_mappings
mock_db_record = MagicMock()
mock_db_record.sso_settings = stored_sso_settings
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_db_record)
monkeypatch.setattr(
proxy_config, "_decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables
)
get_response = client.get("/get/sso_settings")
assert get_response.status_code == 200
get_data = get_response.json()
# Verify role_mappings is returned correctly
assert "role_mappings" in get_data["values"]
retrieved_role_mappings = get_data["values"]["role_mappings"]
assert retrieved_role_mappings is not None
assert retrieved_role_mappings["provider"] == "google"
assert retrieved_role_mappings["group_claim"] == "groups"
assert retrieved_role_mappings["default_role"] == LitellmUserRoles.INTERNAL_USER