diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 451f110a894..503cd7f1a40 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -77,14 +77,18 @@ router = APIRouter() @router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False) -async def google_login(request: Request, source: Optional[str] = None, key: Optional[str] = None): # noqa: PLR0915 +async def google_login( + request: Request, source: Optional[str] = None, key: Optional[str] = None +): # noqa: PLR0915 """ 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/" Example: """ from litellm.proxy.proxy_server import ( + _license_check, premium_user, + prisma_client, user_custom_ui_sso_sign_in_handler, ) @@ -106,12 +110,23 @@ async def google_login(request: Request, source: Optional[str] = None, key: Opti or generic_client_id is not None ): if premium_user is not True: - raise ProxyException( - message="You must be a LiteLLM Enterprise user to use SSO. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this", - type=ProxyErrorTypes.auth_error, - param="premium_user", - code=status.HTTP_403_FORBIDDEN, - ) + # Check if under 'free SSO user' limit + if prisma_client is not None: + total_users = await prisma_client.db.litellm_usertable.count() + if total_users and total_users > 5: + raise ProxyException( + message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this", + type=ProxyErrorTypes.auth_error, + param="premium_user", + code=status.HTTP_403_FORBIDDEN, + ) + else: + raise ProxyException( + message=CommonProxyErrors.db_not_connected_error.value, + type=ProxyErrorTypes.auth_error, + param="premium_user", + code=status.HTTP_403_FORBIDDEN, + ) ####### Detect DB + MASTER KEY in .env ####### missing_env_vars = show_missing_vars_in_env() @@ -124,7 +139,7 @@ async def google_login(request: Request, source: Optional[str] = None, key: Opti 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, @@ -137,11 +152,14 @@ async def google_login(request: Request, source: Optional[str] = None, key: Opti 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 ( @@ -525,15 +543,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 ( @@ -608,7 +627,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, @@ -618,28 +637,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( @@ -653,63 +670,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)}" ) @@ -811,6 +822,7 @@ class SSOAuthenticationHandler: """ Handler for SSO Authentication across all SSO providers """ + @staticmethod async def get_sso_login_redirect( redirect_url: str, @@ -1163,7 +1175,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]: @@ -1176,13 +1187,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, @@ -1202,14 +1215,18 @@ class SSOAuthenticationHandler: ) from litellm.proxy.utils import 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] @@ -1394,7 +1411,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