diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index cbbd114cc39..4b7f2e7a7f8 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -71,7 +71,12 @@ router = APIRouter() @router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False) -async def serve_login_page(request: Request, source: Optional[str] = None, key: Optional[str] = None, error: Optional[str] = None): +async def serve_login_page( + request: Request, + source: Optional[str] = None, + key: Optional[str] = None, + error: Optional[str] = None, +): """ Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/" @@ -165,7 +170,9 @@ async def serve_login_page(request: Request, source: Optional[str] = None, key: server_root_path = os.getenv("SERVER_ROOT_PATH", "") if server_root_path != "": proxy_base_url += server_root_path - form_action = proxy_base_url + "/sso/key/generate" # CHANGE BACK to /sso/key/generate + form_action = ( + proxy_base_url + "/sso/key/generate" + ) # CHANGE BACK to /sso/key/generate unified_login_html = f""" @@ -391,11 +398,14 @@ async def serve_login_page(request: Request, source: Optional[str] = None, key: """ from fastapi.responses import HTMLResponse + return HTMLResponse(content=unified_login_html, status_code=200) @router.get("/sso/login", tags=["experimental"], include_in_schema=False) -async def sso_login_redirect(request: Request, source: Optional[str] = None, key: Optional[str] = None): +async def sso_login_redirect( + request: Request, source: Optional[str] = None, key: Optional[str] = None +): """ Handles SSO login redirect - this is what the "Login with SSO" button points to """ @@ -424,7 +434,7 @@ async def sso_login_redirect(request: Request, source: Optional[str] = None, key request=request, sso_callback_route="sso/callback", ) - + # Store CLI key in state for OAuth flow cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state( source=source, @@ -437,11 +447,14 @@ async def sso_login_redirect(request: Request, source: Optional[str] = None, key from litellm_enterprise.proxy.auth.custom_sso_handler import ( EnterpriseCustomSSOHandler, ) + return await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( request=request, ) except ImportError: - raise ValueError("Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise.") + raise ValueError( + "Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise." + ) # Check if we should use SSO handler if ( @@ -818,15 +831,16 @@ async def check_and_update_if_proxy_admin_id( async def auth_callback(request: Request, state: Optional[str] = None): # noqa: PLR0915 """Verify login""" verbose_proxy_logger.info(f"Starting SSO callback with state: {state}") - + # Check if this is a CLI login (state starts with our CLI prefix) from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX + if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"): # Extract the key ID from the state key_id = state.split(":", 1)[1] verbose_proxy_logger.info(f"CLI SSO callback detected for key: {key_id}") return await cli_sso_callback(request, key=key_id) - + from litellm.proxy._types import LiteLLM_JWTAuth from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.proxy_server import ( @@ -901,7 +915,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: status_code=401, detail="Result not returned by SSO provider.", ) - + return await SSOAuthenticationHandler.get_redirect_response_from_openid( result=result, request=request, @@ -911,28 +925,26 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: ) - - async def cli_sso_callback(request: Request, key: Optional[str] = None): """CLI SSO callback - generates the key with pre-specified ID""" verbose_proxy_logger.info(f"CLI SSO callback for key: {key}") - + from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, ) from litellm.proxy.proxy_server import prisma_client - - if not key or not key.startswith('sk-'): + + if not key or not key.startswith("sk-"): raise HTTPException( status_code=400, - detail="Invalid key parameter. Must be a valid key ID starting with 'sk-'" + detail="Invalid key parameter. Must be a valid key ID starting with 'sk-'", ) - + if prisma_client is None: raise HTTPException( status_code=500, detail=CommonProxyErrors.db_not_connected_error.value ) - + # Generate a simple key for CLI usage with the pre-specified key ID try: await generate_key_helper_fn( @@ -946,63 +958,57 @@ async def cli_sso_callback(request: Request, key: Optional[str] = None): table_name="key", token=key, # Use the pre-specified key ID ) - + verbose_proxy_logger.info(f"Generated CLI key: {key}") - + # Return success page from fastapi.responses import HTMLResponse from litellm.proxy.common_utils.html_forms.cli_sso_success import ( render_cli_sso_success_page, ) - + html_content = render_cli_sso_success_page() return HTMLResponse(content=html_content, status_code=200) - + except Exception as e: verbose_proxy_logger.error(f"Error generating CLI key: {e}") - raise HTTPException( - status_code=500, - detail=f"Failed to generate key: {str(e)}" - ) + raise HTTPException(status_code=500, detail=f"Failed to generate key: {str(e)}") @router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False) async def cli_poll_key(key_id: str): """CLI polling endpoint - checks if key exists in DB""" from litellm.proxy.proxy_server import prisma_client - - if not key_id.startswith('sk-'): - raise HTTPException( - status_code=400, - detail="Invalid key ID format" - ) - + + if not key_id.startswith("sk-"): + raise HTTPException(status_code=400, detail="Invalid key ID format") + if prisma_client is None: raise HTTPException( status_code=500, detail=CommonProxyErrors.db_not_connected_error.value ) - + try: # Check if key exists in database from litellm.proxy.utils import hash_token + hashed_token = hash_token(key_id) - + key_obj = await prisma_client.db.litellm_verificationtoken.find_unique( where={"token": hashed_token} ) - + if key_obj: verbose_proxy_logger.info(f"CLI key found: {key_id}") return {"status": "ready", "key": key_id} else: return {"status": "pending"} - + except Exception as e: verbose_proxy_logger.error(f"Error polling for CLI key: {e}") raise HTTPException( - status_code=500, - detail=f"Error checking key status: {str(e)}" + status_code=500, detail=f"Error checking key status: {str(e)}" ) @@ -1039,9 +1045,9 @@ async def insert_sso_user( if user_defined_values.get("max_budget") is None: user_defined_values["max_budget"] = litellm.max_internal_user_budget if user_defined_values.get("budget_duration") is None: - user_defined_values["budget_duration"] = ( - litellm.internal_user_budget_duration - ) + user_defined_values[ + "budget_duration" + ] = litellm.internal_user_budget_duration if user_defined_values["user_role"] is None: user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY @@ -1104,6 +1110,7 @@ class SSOAuthenticationHandler: """ Handler for SSO Authentication across all SSO providers """ + @staticmethod async def get_sso_login_redirect( redirect_url: str, @@ -1239,9 +1246,9 @@ class SSOAuthenticationHandler: if state: redirect_params["state"] = state elif "okta" in generic_authorization_endpoint: - redirect_params["state"] = ( - uuid.uuid4().hex - ) # set state param for okta - required + redirect_params[ + "state" + ] = uuid.uuid4().hex # set state param for okta - required return await generic_sso.get_login_redirect(**redirect_params) # type: ignore raise ValueError( "Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso" @@ -1456,7 +1463,6 @@ class SSOAuthenticationHandler: _new_team_request.update(_default_team_params) team_request = NewTeamRequest(**_new_team_request) return team_request - @staticmethod def _get_cli_state(source: Optional[str], key: Optional[str]) -> Optional[str]: @@ -1469,13 +1475,15 @@ class SSOAuthenticationHandler: LITELLM_CLI_SESSION_TOKEN_PREFIX, LITELLM_CLI_SOURCE_IDENTIFIER, ) - return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}" if source == LITELLM_CLI_SOURCE_IDENTIFIER and key else None - - + return ( + f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}" + if source == LITELLM_CLI_SOURCE_IDENTIFIER and key + else None + ) @staticmethod - async def get_redirect_response_from_openid( # noqa: PLR0915 + async def get_redirect_response_from_openid( # noqa: PLR0915 result: Union[OpenID, dict, CustomOpenID], request: Request, received_response: Optional[dict] = None, @@ -1495,14 +1503,18 @@ class SSOAuthenticationHandler: ) from litellm.proxy.utils import get_custom_url, get_prisma_client_or_throw from litellm.types.proxy.ui_sso import ReturnedUITokenObject - prisma_client = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy") + prisma_client = get_prisma_client_or_throw( + "Prisma client is None, connect a database to your proxy" + ) # User is Authe'd in - generate key for the UI to access Proxy verbose_proxy_logger.info(f"SSO callback result: {result}") user_email: Optional[str] = getattr(result, "email", None) - user_id: Optional[str] = getattr(result, "id", None) if result is not None else None + user_id: Optional[str] = ( + getattr(result, "id", None) if result is not None else None + ) if user_email is not None and os.getenv("ALLOWED_EMAIL_DOMAINS") is not None: email_domain = user_email.split("@")[1] @@ -1687,7 +1699,8 @@ class SSOAuthenticationHandler: redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303) redirect_response.set_cookie(key="token", value=jwt_token) return redirect_response - + + class MicrosoftSSOHandler: """ Handles Microsoft SSO callback response and returns a CustomOpenID object @@ -1756,9 +1769,9 @@ class MicrosoftSSOHandler: # if user is trying to get the raw sso response for debugging, return the raw sso response if return_raw_sso_response: - original_msft_result[MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY] = ( - user_team_ids - ) + original_msft_result[ + MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY + ] = user_team_ids return original_msft_result or {} result = MicrosoftSSOHandler.openid_from_response( @@ -1826,9 +1839,9 @@ class MicrosoftSSOHandler: # Fetch user membership from Microsoft Graph API all_group_ids = [] - next_link: Optional[str] = ( - MicrosoftSSOHandler.graph_api_user_groups_endpoint - ) + next_link: Optional[ + str + ] = MicrosoftSSOHandler.graph_api_user_groups_endpoint auth_headers = {"Authorization": f"Bearer {access_token}"} page_count = 0 @@ -2199,16 +2212,16 @@ async def process_login(request: Request): form_data = await request.form() username = form_data.get("username") password = form_data.get("password") - + if not username or not password: return RedirectResponse(url="/sso/key/generate?error=1", status_code=303) - + # Import the actual login function from proxy_server from litellm.proxy.proxy_server import login - + # Call the real login function that handles all the authentication properly return await login(request) - + except Exception as e: verbose_proxy_logger.error(f"Error processing login: {e}") return RedirectResponse(url="/sso/key/generate?error=1", status_code=303)