mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
WIP waiting for okta
This commit is contained in:
parent
42d7d757a3
commit
58330f852d
1 changed files with 127 additions and 7 deletions
|
|
@ -85,6 +85,58 @@ else:
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
def determine_role_from_groups(
|
||||
user_groups: List[str],
|
||||
role_mappings: "RoleMappings",
|
||||
) -> 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,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
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)):
|
||||
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}"
|
||||
)
|
||||
return role_mappings.default_role
|
||||
|
||||
|
||||
def process_sso_jwt_access_token(
|
||||
access_token_str: Optional[str],
|
||||
sso_jwt_handler: Optional[JWTHandler],
|
||||
|
|
@ -243,6 +295,7 @@ def generic_response_convertor(
|
|||
response,
|
||||
jwt_handler: JWTHandler,
|
||||
sso_jwt_handler: Optional[JWTHandler] = None,
|
||||
role_mappings: Optional["RoleMappings"] = None,
|
||||
) -> CustomOpenID:
|
||||
generic_user_id_attribute_name = os.getenv(
|
||||
"GENERIC_USER_ID_ATTRIBUTE", "preferred_username"
|
||||
|
|
@ -281,16 +334,48 @@ def generic_response_convertor(
|
|||
team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response))
|
||||
all_teams.extend(team_ids)
|
||||
|
||||
# Extract user role from SSO response
|
||||
user_role_from_sso = get_nested_value(response, generic_user_role_attribute_name)
|
||||
# Determine user role based on role_mappings if available
|
||||
# Only apply role_mappings for GENERIC SSO provider
|
||||
user_role: Optional[LitellmUserRoles] = None
|
||||
if user_role_from_sso is not None:
|
||||
role = get_litellm_user_role(user_role_from_sso)
|
||||
if role is not None:
|
||||
user_role = role
|
||||
|
||||
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 = 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):
|
||||
user_groups = [str(g) for g in user_groups_raw]
|
||||
elif isinstance(user_groups_raw, str):
|
||||
# Handle comma-separated string
|
||||
user_groups = [g.strip() for g in user_groups_raw.split(",") if g.strip()]
|
||||
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(
|
||||
f"Found valid LitellmUserRoles '{role.value}' from SSO attribute '{generic_user_role_attribute_name}'"
|
||||
f"Determined role '{user_role.value if user_role else None}' from groups '{user_groups}' using role_mappings"
|
||||
)
|
||||
else:
|
||||
# No groups found, use default_role
|
||||
user_role = role_mappings.default_role
|
||||
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)
|
||||
if user_role_from_sso is not None:
|
||||
role = get_litellm_user_role(user_role_from_sso)
|
||||
if role is not None:
|
||||
user_role = role
|
||||
verbose_proxy_logger.debug(
|
||||
f"Found valid LitellmUserRoles '{role.value}' from SSO attribute '{generic_user_role_attribute_name}'"
|
||||
)
|
||||
|
||||
return CustomOpenID(
|
||||
id=get_nested_value(response, generic_user_id_attribute_name),
|
||||
|
|
@ -369,6 +454,40 @@ async def get_generic_sso_response(
|
|||
userinfo_endpoint=generic_userinfo_endpoint,
|
||||
)
|
||||
|
||||
# Get role_mappings from SSO settings if available
|
||||
role_mappings: Optional["RoleMappings"] = 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"
|
||||
)
|
||||
|
||||
# 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")
|
||||
|
||||
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
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
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."
|
||||
)
|
||||
|
||||
def response_convertor(response, client):
|
||||
nonlocal received_response # return for user debugging
|
||||
received_response = response
|
||||
|
|
@ -376,6 +495,7 @@ async def get_generic_sso_response(
|
|||
response=response,
|
||||
jwt_handler=jwt_handler,
|
||||
sso_jwt_handler=sso_jwt_handler,
|
||||
role_mappings=role_mappings,
|
||||
)
|
||||
|
||||
SSOProvider = create_provider(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue