feat: support role_mappings from environment variables (#19498)

* feat: support role_mappings from environment variables

* fix linter
This commit is contained in:
Misha 2026-01-24 06:54:23 +03:00 • committed by GitHub
parent 5c61586e65
commit de538456e3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 183 additions and 23 deletions

View file

@ -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

View file

@ -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"]