mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
team mapping
This commit is contained in:
parent
87c1e6fd68
commit
395ccac651
5 changed files with 120 additions and 7 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue