From 339351c579fe1bfa1963979e452877d0e83e6fb0 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Fri, 25 Apr 2025 21:30:29 -0700 Subject: [PATCH] fix(ui_sso.py): support experimental jwt keys for UI auth w/ SSO (#10326) --- litellm/proxy/management_endpoints/ui_sso.py | 28 ++++++++++++++++++-- litellm/proxy/proxy_server.py | 24 ++++++++--------- 2 files changed, 38 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 4f53c037022..7414488bc8a 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -40,7 +40,7 @@ from litellm.proxy._types import ( TeamMemberAddRequest, UserAPIKeyAuth, ) -from litellm.proxy.auth.auth_checks import get_user_object +from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object from litellm.proxy.auth.auth_utils import _has_user_setup_sso from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -60,7 +60,7 @@ from litellm.proxy.management_endpoints.sso_helper_utils import ( from litellm.proxy.management_endpoints.team_endpoints import new_team, team_member_add from litellm.proxy.management_endpoints.types import CustomOpenID from litellm.proxy.utils import PrismaClient, ProxyLogging -from litellm.secret_managers.main import str_to_bool +from litellm.secret_managers.main import get_secret_bool, str_to_bool from litellm.types.proxy.management_endpoints.ui_sso import * if TYPE_CHECKING: @@ -621,6 +621,30 @@ async def auth_callback(request: Request): # noqa: PLR0915 import jwt + if get_secret_bool("EXPERIMENTAL_UI_LOGIN"): + _user_info: Optional[LiteLLM_UserTable] = None + if ( + user_defined_values is not None + and user_defined_values["user_id"] is not None + ): + _user_info = LiteLLM_UserTable( + user_id=user_defined_values["user_id"], + user_role=user_defined_values["user_role"] or user_role, + models=[], + max_budget=litellm.max_ui_session_budget, + ) + if _user_info is None: + raise HTTPException( + status_code=401, + detail={ + "error": "User Information is required for experimental UI login" + }, + ) + + key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( + _user_info + ) + jwt_token = jwt.encode( # type: ignore { "user_id": user_id, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 569887f0acc..5252cb9a973 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -809,9 +809,9 @@ model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter( dual_cache=user_api_key_cache ) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) -redis_usage_cache: Optional[RedisCache] = ( - None # redis cache used for tracking spend, tpm/rpm limits -) +redis_usage_cache: Optional[ + RedisCache +] = None # redis cache used for tracking spend, tpm/rpm limits user_custom_auth = None user_custom_key_generate = None user_custom_sso = None @@ -1137,9 +1137,9 @@ async def update_cache( # noqa: PLR0915 _id = "team_id:{}".format(team_id) try: # Fetch the existing cost for the given user - existing_spend_obj: Optional[LiteLLM_TeamTable] = ( - await user_api_key_cache.async_get_cache(key=_id) - ) + existing_spend_obj: Optional[ + LiteLLM_TeamTable + ] = await user_api_key_cache.async_get_cache(key=_id) if existing_spend_obj is None: # do nothing if team not in api key cache return @@ -2811,9 +2811,9 @@ async def initialize( # noqa: PLR0915 user_api_base = api_base dynamic_config[user_model]["api_base"] = api_base if api_version: - os.environ["AZURE_API_VERSION"] = ( - api_version # set this for azure - litellm can read this from the env - ) + os.environ[ + "AZURE_API_VERSION" + ] = api_version # set this for azure - litellm can read this from the env if max_tokens: # model-specific param dynamic_config[user_model]["max_tokens"] = max_tokens if temperature: # model-specific param @@ -7791,9 +7791,9 @@ async def get_config_list( hasattr(sub_field_info, "description") and sub_field_info.description is not None ): - nested_fields[idx].field_description = ( - sub_field_info.description - ) + nested_fields[ + idx + ].field_description = sub_field_info.description idx += 1 _stored_in_db = None