mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat: support role_mappings from environment variables (#19498)
* feat: support role_mappings from environment variables * fix linter
This commit is contained in:
parent
5c61586e65
commit
de538456e3
2 changed files with 183 additions and 23 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue