mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
lint fix
This commit is contained in:
parent
e277e24ff3
commit
1a78984c8b
1 changed files with 74 additions and 61 deletions
|
|
@ -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"""
|
||||
<!DOCTYPE html>
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue