From a6527e501044df1732bd9bce8a03654ae49b73b7 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 1 Jul 2025 18:11:19 -0700 Subject: [PATCH] [Feat] Add litellm-proxy cli login for starting to use litellm proxy (#12216) * add handlers for auth commands * add login, logout, whoami * refactor auth * add CLI Authentication Flow * add SSO sign in constants * add itellm-session-token * fixes for managing state with cli * use locally stored context for cli session * add litellm banner + interactive shell * update main.py * update auth to show commands * fix ui sso render * add TestCLISSOCallbackFunction * update banner.py * remove file * fix cli sso success * TestTokenUtilities * fix code qa * fix execute_command * fix cli_sso_callback * fix import * Authentication using CLI --- docs/my-website/docs/proxy/cli_sso.md | 56 +++ docs/my-website/docs/proxy/management_cli.md | 36 ++ docs/my-website/sidebars.js | 1 + litellm/constants.py | 4 + litellm/proxy/client/README.md | 95 +++- litellm/proxy/client/cli/banner.py | 15 + litellm/proxy/client/cli/commands/auth.py | 168 +++++++ litellm/proxy/client/cli/interface.py | 206 +++++++++ litellm/proxy/client/cli/main.py | 36 +- .../html_forms/cli_sso_success.py | 208 +++++++++ litellm/proxy/management_endpoints/ui_sso.py | 138 +++++- .../proxy/client/cli/test_auth_commands.py | 431 ++++++++++++++++++ .../proxy/management_endpoints/test_ui_sso.py | 273 +++++++++++ 13 files changed, 1648 insertions(+), 19 deletions(-) create mode 100644 docs/my-website/docs/proxy/cli_sso.md create mode 100644 litellm/proxy/client/cli/banner.py create mode 100644 litellm/proxy/client/cli/commands/auth.py create mode 100644 litellm/proxy/client/cli/interface.py create mode 100644 litellm/proxy/common_utils/html_forms/cli_sso_success.py create mode 100644 tests/test_litellm/proxy/client/cli/test_auth_commands.py diff --git a/docs/my-website/docs/proxy/cli_sso.md b/docs/my-website/docs/proxy/cli_sso.md new file mode 100644 index 00000000000..f7669d6a25c --- /dev/null +++ b/docs/my-website/docs/proxy/cli_sso.md @@ -0,0 +1,56 @@ +# CLI Authentication + +Use the litellm cli to authenticate to the LiteLLM Gateway. This is great if you're trying to give a large number of developers self-serve access to the LiteLLM Gateway. + + +## Demo + + + +## Usage + + +1. **Install the CLI** + + If you have [uv](https://github.com/astral-sh/uv) installed, you can try this: + + ```shell + uv tool install 'litellm[proxy]' + ``` + + If that works, you'll see something like this: + + ```shell + ... + Installed 2 executables: litellm, litellm-proxy + ``` + + and now you can use the tool by just typing `litellm-proxy` in your terminal: + + ```shell + litellm-proxy + ``` + +2. **Set up environment variables** + + ```bash + export LITELLM_PROXY_URL=http://localhost:4000 + ``` + + *(Replace with your actual proxy URL)* + +3. **Login** + + ```shell + litellm-proxy login + ``` + + This will open a browser window to authenticate. If you have connected LiteLLM Proxy to your SSO provider, you should be able to login with your SSO credentials. Once logged in, you can use the CLI to make requests to the LiteLLM Gateway. + +4. **Make a test request to view models** + + ```shell + litellm-proxy models list + ``` + + This will list all the models available to you. \ No newline at end of file diff --git a/docs/my-website/docs/proxy/management_cli.md b/docs/my-website/docs/proxy/management_cli.md index 6593b88ba4f..9ecc2ae8a34 100644 --- a/docs/my-website/docs/proxy/management_cli.md +++ b/docs/my-website/docs/proxy/management_cli.md @@ -57,6 +57,42 @@ and more, as well as making chat and HTTP requests to the proxy server. - If you see an error, check your environment variables and proxy server status. +## Authentication using CLI + +You can use the CLI to authenticate to the LiteLLM Gateway. This is great if you're trying to give a large number of developers self-serve access to the LiteLLM Gateway. + +:::info + +For an indepth guide, see [CLI Authentication](./cli_sso). + +::: + + + +1. **Set up the proxy URL** + + ```bash + export LITELLM_PROXY_URL=http://localhost:4000 + ``` + + *(Replace with your actual proxy URL)* + +2. **Login** + + ```bash + litellm-proxy login + ``` + + This will open a browser window to authenticate. If you have connected LiteLLM Proxy to your SSO provider, you can login with your SSO credentials. Once logged in, you can use the CLI to make requests to the LiteLLM Gateway. + +3. **Test your authentication** + + ```bash + litellm-proxy models list + ``` + + This will list all the models available to you. + ## Main Commands ### Models Management diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 1e1322a8c01..6e3fd05eea9 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -142,6 +142,7 @@ const sidebars = { "proxy/token_auth", "proxy/service_accounts", "proxy/access_control", + "proxy/cli_sso", "proxy/custom_auth", "proxy/ip_address", "proxy/email", diff --git a/litellm/constants.py b/litellm/constants.py index f0b38d21a12..4ee28adff93 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -758,6 +758,10 @@ HEALTH_CHECK_TIMEOUT_SECONDS = int( UI_SESSION_TOKEN_TEAM_ID = "litellm-dashboard" LITELLM_PROXY_ADMIN_NAME = "default_user_id" +########################### CLI SSO AUTHENTICATION CONSTANTS ########################### +LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli" +LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token" + ########################### DB CRON JOB NAMES ########################### DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job" PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics" diff --git a/litellm/proxy/client/README.md b/litellm/proxy/client/README.md index 06340e6dc17..8e8c5d2b9db 100644 --- a/litellm/proxy/client/README.md +++ b/litellm/proxy/client/README.md @@ -298,4 +298,97 @@ Contributions are welcome! Please check out our [contributing guidelines](../../ ## License -This project is licensed under the MIT License - see the [LICENSE](../../LICENSE) file for details. \ No newline at end of file +This project is licensed under the MIT License - see the [LICENSE](../../LICENSE) file for details. + +## CLI Authentication Flow + +The LiteLLM CLI supports SSO authentication through a polling-based approach that works with any OAuth-compatible SSO provider. + +### How CLI Authentication Works + +```mermaid +sequenceDiagram + participant CLI as CLI + participant Browser as Browser + participant Proxy as LiteLLM Proxy + participant SSO as SSO Provider + + CLI->>CLI: Generate key ID (sk-uuid) + CLI->>Browser: Open /sso/key/generate?source=litellm-cli&key=sk-uuid + + Browser->>Proxy: GET /sso/key/generate?source=litellm-cli&key=sk-uuid + Proxy->>Proxy: Set cli_state = litellm-session-token:sk-uuid + Proxy->>SSO: Redirect with state=litellm-session-token:sk-uuid + + SSO->>Browser: Show login page + Browser->>SSO: User authenticates + SSO->>Proxy: Redirect to /sso/callback?state=litellm-session-token:sk-uuid + + Proxy->>Proxy: Check if state starts with "litellm-session-token:" + Proxy->>Proxy: Generate API key with ID=sk-uuid + Proxy->>Browser: Show success page + + CLI->>Proxy: Poll /sso/cli/poll/sk-uuid + Proxy->>CLI: Return {"status": "ready", "key": "sk-uuid"} + CLI->>CLI: Save key to ~/.litellm/token.json +``` + +### Authentication Commands + +The CLI provides three authentication commands: + +- **`litellm-proxy login`** - Start SSO authentication flow +- **`litellm-proxy logout`** - Clear stored authentication token +- **`litellm-proxy whoami`** - Show current authentication status + +### Authentication Flow Steps + +1. **Generate Session ID**: CLI generates a unique key ID (`sk-{uuid}`) +2. **Open Browser**: CLI opens browser to `/sso/key/generate` with CLI source and key parameters +3. **SSO Redirect**: Proxy sets the formatted state (`litellm-session-token:sk-uuid`) as OAuth state parameter and redirects to SSO provider +4. **User Authentication**: User completes SSO authentication in browser +5. **Callback Processing**: SSO provider redirects back to proxy with state parameter +6. **Key Generation**: Proxy detects CLI login (state starts with "litellm-session-token:") and generates API key with pre-specified ID +7. **Polling**: CLI polls `/sso/cli/poll/{key_id}` endpoint until key is ready +8. **Token Storage**: CLI saves the authentication token to `~/.litellm/token.json` + +### Benefits of This Approach + +- **No Local Server**: No need to run a local callback server +- **Standard OAuth**: Uses OAuth 2.0 state parameter correctly +- **Remote Compatible**: Works with remote proxy servers +- **Secure**: Uses UUID session identifiers +- **Simple Setup**: No additional OAuth redirect URL configuration needed + +### Token Storage + +Authentication tokens are stored in `~/.litellm/token.json` with restricted file permissions (600). The stored token includes: + +```json +{ + "key": "sk-...", + "user_id": "cli-user", + "user_email": "user@example.com", + "user_role": "cli", + "auth_header_name": "Authorization", + "timestamp": 1234567890 +} +``` + +### Usage + +Once authenticated, the CLI will automatically use the stored token for all requests. You no longer need to specify `--api-key` for subsequent commands. + +```bash +# Login +litellm-proxy login + +# Use CLI without specifying API key +litellm-proxy models list + +# Check authentication status +litellm-proxy whoami + +# Logout +litellm-proxy logout +``` \ No newline at end of file diff --git a/litellm/proxy/client/cli/banner.py b/litellm/proxy/client/cli/banner.py new file mode 100644 index 00000000000..58bba9cac2e --- /dev/null +++ b/litellm/proxy/client/cli/banner.py @@ -0,0 +1,15 @@ +# third party imports +import click + +# LiteLLM ASCII banner +LITELLM_BANNER = """ ██╗ ██╗████████╗███████╗██╗ ██╗ ███╗ ███╗ + ██║ ██║╚══██╔══╝██╔════╝██║ ██║ ████╗ ████║ + ██║ ██║ ██║ █████╗ ██║ ██║ ██╔████╔██║ + ██║ ██║ ██║ ██╔══╝ ██║ ██║ ██║╚██╔╝██║ + ███████╗██║ ██║ ███████╗███████╗███████╗██║ ╚═╝ ██║ + ╚══════╝╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝ ╚═╝""" + + +def show_banner(): + """Display the LiteLLM CLI banner.""" + click.echo(f"\n{LITELLM_BANNER}\n") \ No newline at end of file diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py new file mode 100644 index 00000000000..7d89d39ed70 --- /dev/null +++ b/litellm/proxy/client/cli/commands/auth.py @@ -0,0 +1,168 @@ +import json +import os +import time +import webbrowser +from pathlib import Path +from typing import Any, Dict, Optional + +import click + + +# Token storage utilities +def get_token_file_path() -> str: + """Get the path to store the authentication token""" + home_dir = Path.home() + config_dir = home_dir / ".litellm" + config_dir.mkdir(exist_ok=True) + return str(config_dir / "token.json") + +def save_token(token_data: Dict[str, Any]) -> None: + """Save token data to file""" + token_file = get_token_file_path() + with open(token_file, 'w') as f: + json.dump(token_data, f, indent=2) + # Set file permissions to be readable only by owner + os.chmod(token_file, 0o600) + +def load_token() -> Optional[Dict[str, Any]]: + """Load token data from file""" + token_file = get_token_file_path() + if not os.path.exists(token_file): + return None + + try: + with open(token_file, 'r') as f: + return json.load(f) + except (json.JSONDecodeError, IOError): + return None + +def clear_token() -> None: + """Clear stored token""" + token_file = get_token_file_path() + if os.path.exists(token_file): + os.remove(token_file) + +def get_stored_api_key() -> Optional[str]: + """Get the stored API key from token file""" + token_data = load_token() + if token_data and 'key' in token_data: + return token_data['key'] + return None + +# Polling-based authentication - no local server needed + +@click.command(name="login") +@click.pass_context +def login(ctx: click.Context): + """Login to LiteLLM proxy using SSO authentication""" + import uuid + + import requests + + from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER + from litellm.proxy.client.cli.interface import show_commands + + base_url = ctx.obj["base_url"] + + # Generate unique key ID for this login session + key_id = f"sk-{str(uuid.uuid4())}" + + try: + # Construct SSO login URL with CLI source and pre-generated key + sso_url = f"{base_url}/sso/key/generate?source={LITELLM_CLI_SOURCE_IDENTIFIER}&key={key_id}" + + click.echo(f"Opening browser to: {sso_url}") + click.echo("Please complete the SSO authentication in your browser...") + click.echo(f"Session ID: {key_id}") + + # Open browser + webbrowser.open(sso_url) + + # Poll for key creation + click.echo("Waiting for authentication...") + + poll_url = f"{base_url}/sso/cli/poll/{key_id}" + timeout = 300 # 5 minute timeout + poll_interval = 2 # Poll every 2 seconds + + for attempt in range(timeout // poll_interval): + try: + response = requests.get(poll_url, timeout=10) + if response.status_code == 200: + data = response.json() + if data.get("status") == "ready": + # Key is ready - save it + api_key = data.get("key") + if api_key: + # Save token data (simplified for CLI - we just need the key) + save_token({ + 'key': api_key, + 'user_id': 'cli-user', + 'user_email': 'unknown', + 'user_role': 'cli', + 'auth_header_name': 'Authorization', + 'jwt_token': '', + 'timestamp': time.time() + }) + + click.echo("✅ Login successful!") + click.echo(f"API Key: {api_key[:20]}...") + click.echo("You can now use the CLI without specifying --api-key") + + # Show available commands after successful login + click.echo("\n" + "="*60) + show_commands() + return + elif response.status_code == 200: + # Still pending + if attempt % 10 == 0: # Show progress every 20 seconds + click.echo("Still waiting for authentication...") + else: + click.echo(f"Polling error: HTTP {response.status_code}") + + except requests.RequestException as e: + if attempt % 10 == 0: + click.echo(f"Connection error (will retry): {e}") + + time.sleep(poll_interval) + + click.echo("❌ Authentication timed out. Please try again.") + return + + except KeyboardInterrupt: + click.echo("\n❌ Authentication cancelled by user.") + return + except Exception as e: + click.echo(f"❌ Authentication failed: {e}") + return + +@click.command(name="logout") +def logout(): + """Logout and clear stored authentication""" + clear_token() + click.echo("✅ Logged out successfully. Authentication token cleared.") + +@click.command(name="whoami") +def whoami(): + """Show current authentication status""" + token_data = load_token() + + if not token_data: + click.echo("❌ Not authenticated. Run 'litellm-proxy login' to authenticate.") + return + + click.echo("✅ Authenticated") + click.echo(f"User Email: {token_data.get('user_email', 'Unknown')}") + click.echo(f"User ID: {token_data.get('user_id', 'Unknown')}") + click.echo(f"User Role: {token_data.get('user_role', 'Unknown')}") + + # Check if token is still valid (basic timestamp check) + timestamp = token_data.get('timestamp', 0) + age_hours = (time.time() - timestamp) / 3600 + click.echo(f"Token age: {age_hours:.1f} hours") + + if age_hours > 24: + click.echo("⚠️ Warning: Token is more than 24 hours old and may have expired.") + +# Export individual commands instead of grouping them +# login, logout, and whoami will be added as top-level commands \ No newline at end of file diff --git a/litellm/proxy/client/cli/interface.py b/litellm/proxy/client/cli/interface.py new file mode 100644 index 00000000000..48f21e39297 --- /dev/null +++ b/litellm/proxy/client/cli/interface.py @@ -0,0 +1,206 @@ +# stdlib imports +import os +import sys +from typing import TYPE_CHECKING + +# third party imports +import click + +from litellm._logging import verbose_logger + +if TYPE_CHECKING: + pass + + +def styled_prompt(): + """Create a styled blue box prompt for user input.""" + + # Get terminal height to ensure we have enough space + try: + terminal_height = os.get_terminal_size().lines + # Ensure we have at least 5 lines of space (for the box + some buffer) + if terminal_height < 10: + # If terminal is too small, just add some newlines to push content up + click.echo("\n" * 3) + except Exception as e: + # Fallback if we can't get terminal size + verbose_logger.debug(f"Error getting terminal size: {e}") + click.echo("\n" * 3) + + # Unicode box drawing characters + top_left = "┌" + top_right = "┐" + bottom_left = "└" + bottom_right = "┘" + horizontal = "─" + vertical = "│" + + # Create the box with increased width + width = 80 + top_line = top_left + horizontal * (width - 2) + top_right + bottom_line = bottom_left + horizontal * (width - 2) + bottom_right + + # Create styled elements + left_border = click.style(vertical, fg="blue", bold=True) + right_border = click.style(vertical, fg="blue", bold=True) + prompt_text = click.style("> ", fg="cyan", bold=True) + + # Display the complete box structure first to reserve space + click.echo(click.style(top_line, fg="blue", bold=True)) + + # Create empty space in the box for input + empty_space = " " * (width - 4) + click.echo(f"{left_border} {empty_space} {right_border}") + + # Display bottom border to complete the box + click.echo(click.style(bottom_line, fg="blue", bold=True)) + + # Now move cursor up to the input line and get input + click.echo("\033[2A", nl=False) # Move cursor up 2 lines + click.echo(f"\r{left_border} {prompt_text}", nl=False) # Position at start of input line + + try: + # Get user input + user_input = input().strip() + + # Move cursor down to after the box + click.echo("\033[1B") # Move cursor down 1 line + click.echo("") # Add some space after + + except (KeyboardInterrupt, EOFError): + # Move cursor down and add space + click.echo("\033[1B") + click.echo("") + raise + + return user_input + + +def show_commands(): + """Display available commands.""" + commands = [ + ("login", "Authenticate with the LiteLLM proxy server"), + ("logout", "Clear stored authentication"), + ("whoami", "Show current authentication status"), + ("models", "Manage and view model configurations"), + ("credentials", "Manage API credentials"), + ("chat", "Interactive chat with models"), + ("http", "Make HTTP requests to the proxy"), + ("keys", "Manage API keys"), + ("users", "Manage users"), + ("version", "Show version information"), + ("help", "Show this help message"), + ("quit", "Exit the interactive session"), + ] + + click.echo("Available commands:") + for cmd, description in commands: + click.echo(f" {cmd:<20} {description}") + click.echo() + + +def setup_shell(ctx: click.Context): + """Set up the interactive shell with banner and initial info.""" + from .banner import show_banner + + show_banner() + + # Show server connection info + base_url = ctx.obj.get("base_url") + click.secho(f"Connected to LiteLLM server: {base_url}\n", fg="green") + + show_commands() + + +def handle_special_commands(user_input: str) -> bool: + """Handle special commands like exit, help, clear. Returns True if command was handled.""" + if user_input.lower() in ["exit", "quit"]: + click.echo("Goodbye!") + return True + elif user_input.lower() == "help": + click.echo("") # Add space before help + show_commands() + return True + elif user_input.lower() == "clear": + click.clear() + from .banner import show_banner + show_banner() + show_commands() + return True + + return False + + +def execute_command(user_input: str, ctx: click.Context): + """Parse and execute a command.""" + # Parse command and arguments + parts = user_input.split() + command = parts[0] + args = parts[1:] if len(parts) > 1 else [] + + # Import cli here to avoid circular import + from . import main + cli = main.cli + + # Check if command exists + if command not in cli.commands: + click.echo(f"Unknown command: {command}") + click.echo("Type 'help' to see available commands.") + return + + # Execute the command + try: + # Create a new argument list for click to parse + sys.argv = ["litellm-proxy"] + [command] + args + + # Get the command object and invoke it + cmd = cli.commands[command] + + # Create a new context for the subcommand + with ctx.scope(): + cmd.main( + args, + parent=ctx, + standalone_mode=False + ) + + except click.ClickException as e: + e.show() + except click.Abort: + click.echo("Command aborted.") + except SystemExit: + # Prevent the interactive shell from exiting on command errors + pass + except Exception as e: + click.echo(f"Error executing command: {e}") + + +def interactive_shell(ctx: click.Context): + """Run the interactive shell.""" + setup_shell(ctx) + + while True: + try: + # Add some space before the input box to ensure it's positioned well + click.echo("\n") # Extra spacing + + # Show styled prompt + user_input = styled_prompt() + + if not user_input: + continue + + # Handle special commands + if handle_special_commands(user_input): + if user_input.lower() in ["exit", "quit"]: + break + continue + + # Execute regular commands + execute_command(user_input, ctx) + + except (KeyboardInterrupt, EOFError): + click.echo("\nGoodbye!") + break + except Exception as e: + click.echo(f"Error: {e}") \ No newline at end of file diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index 9ecd2f1e19c..fb4a37c3a17 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -1,20 +1,23 @@ # stdlib imports -import sys from typing import Optional # third party imports import click -# local imports -from .commands.models import models -from .commands.credentials import credentials -from .commands.chat import chat -from .commands.http import http -from .commands.keys import keys -from .commands.users import users from litellm._version import version as litellm_version from litellm.proxy.client.health import HealthManagementClient +from .commands.auth import get_stored_api_key, login, logout, whoami +from .commands.chat import chat +from .commands.credentials import credentials +from .commands.http import http +from .commands.keys import keys + +# local imports +from .commands.models import models +from .commands.users import users +from .interface import interactive_shell + def print_version(base_url: str, api_key: Optional[str]): """Print CLI and server version info.""" @@ -32,7 +35,7 @@ def print_version(base_url: str, api_key: Optional[str]): click.echo(f"Could not retrieve server version: {e}") -@click.group() +@click.group(invoke_without_command=True) @click.option( "--version", "-v", is_flag=True, is_eager=True, expose_value=False, help="Show the LiteLLM Proxy CLI and server version and exit.", @@ -61,10 +64,17 @@ def print_version(base_url: str, api_key: Optional[str]): def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None: """LiteLLM Proxy CLI - Manage your LiteLLM proxy server""" ctx.ensure_object(dict) - if sys.stderr.isatty(): - click.secho(f"Accessing LiteLLM server: {base_url} ...\n", fg="yellow", err=True) + + # If no API key provided via flag or environment variable, try to load from saved token + if api_key is None: + api_key = get_stored_api_key() + ctx.obj["base_url"] = base_url ctx.obj["api_key"] = api_key + + # If no subcommand was invoked, start interactive mode + if ctx.invoked_subcommand is None: + interactive_shell(ctx) @cli.command() @@ -74,6 +84,10 @@ def version(ctx: click.Context): print_version(ctx.obj.get("base_url"), ctx.obj.get("api_key")) +# Add authentication commands as top-level commands +cli.add_command(login) +cli.add_command(logout) +cli.add_command(whoami) # Add the models command group cli.add_command(models) # Add the credentials command group diff --git a/litellm/proxy/common_utils/html_forms/cli_sso_success.py b/litellm/proxy/common_utils/html_forms/cli_sso_success.py new file mode 100644 index 00000000000..ba9036f578e --- /dev/null +++ b/litellm/proxy/common_utils/html_forms/cli_sso_success.py @@ -0,0 +1,208 @@ + +from litellm.proxy.client.cli.banner import LITELLM_BANNER + + +def render_cli_sso_success_page() -> str: + """ + Renders the CLI SSO authentication success page with minimal styling + + Returns: + str: HTML content for the success page + """ + + html_content = f""" + + + + + CLI Authentication Successful - LiteLLM + + + + +
+
+ +
+ + + +

Authentication Successful!

+

Your CLI authentication is complete.

+ +
+
+ + + + + CLI Authentication Complete +
+

Your LiteLLM CLI has been successfully authenticated and is ready to use.

+
+ +
+
+ + + + + + Next Steps +
+

Return to your terminal - the CLI will automatically detect the successful authentication.

+

You can now use LiteLLM CLI commands with your authenticated session.

+
+ +
This window will close in 3 seconds...
+
+ + + + + """ + return html_content \ No newline at end of file diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 77fdb64b0e9..383543275a7 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -72,7 +72,7 @@ router = APIRouter() @router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False) -async def google_login(request: Request): # 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/" @@ -111,11 +111,17 @@ async def google_login(request: Request): # noqa: PLR0915 return missing_env_vars ui_username = os.getenv("UI_USERNAME") - # get url from request + # get url from request - always use regular callback, but set state for CLI redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( 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, + key=key, + ) # Check if we should use SSO handler if ( @@ -132,6 +138,7 @@ async def google_login(request: Request): # noqa: PLR0915 microsoft_client_id=microsoft_client_id, google_client_id=google_client_id, generic_client_id=generic_client_id, + state=cli_state, ) elif ui_username is not None: # No Google, Microsoft SSO @@ -495,9 +502,18 @@ async def check_and_update_if_proxy_admin_id( @router.get("/sso/callback", tags=["experimental"], include_in_schema=False) -async def auth_callback(request: Request): # noqa: PLR0915 +async def auth_callback(request: Request, state: Optional[str] = None): # noqa: PLR0915 """Verify login""" - verbose_proxy_logger.info("Starting SSO callback") + 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.management_endpoints.key_management_endpoints import ( @@ -705,7 +721,7 @@ async def auth_callback(request: Request): # noqa: PLR0915 f"user_role: {user_role}; ui_access_mode: {ui_access_mode}" ) ## CHECK IF ROLE ALLOWED TO USE PROXY ## - is_admin_only_access = check_is_admin_only_access(ui_access_mode) + is_admin_only_access = check_is_admin_only_access(ui_access_mode or {}) if is_admin_only_access: has_access = has_admin_ui_access(user_role) if not has_access: @@ -774,6 +790,99 @@ async def auth_callback(request: Request): # noqa: PLR0915 return redirect_response +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-'): + raise HTTPException( + status_code=400, + 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( + request_type="key", + duration="24hr", + key_max_budget=litellm.max_ui_session_budget, + aliases={}, + config={}, + spend=0, + team_id="litellm-cli", + 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)}" + ) + + +@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 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)}" + ) + + async def insert_sso_user( result_openid: Optional[Union[OpenID, dict]], user_defined_values: Optional[SSOUserDefinedValues] = None, @@ -879,6 +988,7 @@ class SSOAuthenticationHandler: google_client_id: Optional[str] = None, microsoft_client_id: Optional[str] = None, generic_client_id: Optional[str] = None, + state: Optional[str] = None, ) -> Optional[RedirectResponse]: """ Step 1. Call Get Login Redirect for the SSO provider. Send the redirect response to `redirect_url` @@ -913,7 +1023,7 @@ class SSOAuthenticationHandler: f"In /google-login/key/generate, \nGOOGLE_REDIRECT_URI: {redirect_url}\nGOOGLE_CLIENT_ID: {google_client_id}" ) with google_sso: - return await google_sso.get_login_redirect() + return await google_sso.get_login_redirect(state=state) # Microsoft SSO Auth elif microsoft_client_id is not None: from fastapi_sso.sso.microsoft import MicrosoftSSO @@ -935,7 +1045,7 @@ class SSOAuthenticationHandler: allow_insecure_http=True, ) with microsoft_sso: - return await microsoft_sso.get_login_redirect() + return await microsoft_sso.get_login_redirect(state=state) elif generic_client_id is not None: from fastapi_sso.sso.base import DiscoveryDocument from fastapi_sso.sso.generic import create_provider @@ -1224,6 +1334,20 @@ 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]: + """ + Checks the request 'source' if a cli state token was passed in + + This is used to authenticate through the CLI login flow + """ + from litellm.constants import ( + 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 class MicrosoftSSOHandler: diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py new file mode 100644 index 00000000000..33e5ba29bf0 --- /dev/null +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -0,0 +1,431 @@ +import json +import os +import tempfile +import time +from pathlib import Path +from unittest.mock import MagicMock, Mock, mock_open, patch + +import pytest +from click.testing import CliRunner + +from litellm.proxy.client.cli.commands.auth import ( + clear_token, + get_stored_api_key, + get_token_file_path, + load_token, + login, + logout, + save_token, + whoami, +) + + +class TestTokenUtilities: + """Test token file utility functions""" + + def test_get_token_file_path(self): + """Test getting token file path""" + with patch('pathlib.Path.home') as mock_home, \ + patch('pathlib.Path.mkdir') as mock_mkdir: + mock_home.return_value = Path('/home/user') + + result = get_token_file_path() + + assert result == '/home/user/.litellm/token.json' + mock_mkdir.assert_called_once_with(exist_ok=True) + + def test_get_token_file_path_creates_directory(self): + """Test that get_token_file_path creates the config directory""" + with patch('pathlib.Path.home') as mock_home, \ + patch('pathlib.Path.mkdir') as mock_mkdir: + mock_home.return_value = Path('/home/user') + + get_token_file_path() + + mock_mkdir.assert_called_once_with(exist_ok=True) + + def test_save_token(self): + """Test saving token data to file""" + token_data = { + 'key': 'test-key', + 'user_id': 'test-user', + 'timestamp': 1234567890 + } + + with patch('builtins.open', mock_open()) as mock_file, \ + patch('litellm.proxy.client.cli.commands.auth.get_token_file_path') as mock_path, \ + patch('os.chmod') as mock_chmod: + + mock_path.return_value = '/test/path/token.json' + + save_token(token_data) + + mock_file.assert_called_once_with('/test/path/token.json', 'w') + mock_file().write.assert_called() + mock_chmod.assert_called_once_with('/test/path/token.json', 0o600) + + # Verify JSON content was written correctly + written_content = ''.join(call[0][0] for call in mock_file().write.call_args_list) + parsed_content = json.loads(written_content) + assert parsed_content == token_data + + def test_load_token_success(self): + """Test loading token data from file successfully""" + token_data = { + 'key': 'test-key', + 'user_id': 'test-user', + 'timestamp': 1234567890 + } + + with patch('builtins.open', mock_open(read_data=json.dumps(token_data))), \ + patch('litellm.proxy.client.cli.commands.auth.get_token_file_path') as mock_path, \ + patch('os.path.exists', return_value=True): + + mock_path.return_value = '/test/path/token.json' + + result = load_token() + + assert result == token_data + + def test_load_token_file_not_exists(self): + """Test loading token when file doesn't exist""" + with patch('litellm.proxy.client.cli.commands.auth.get_token_file_path') as mock_path, \ + patch('os.path.exists', return_value=False): + + mock_path.return_value = '/test/path/token.json' + + result = load_token() + + assert result is None + + def test_load_token_json_decode_error(self): + """Test loading token with invalid JSON""" + with patch('builtins.open', mock_open(read_data='invalid json')), \ + patch('litellm.proxy.client.cli.commands.auth.get_token_file_path') as mock_path, \ + patch('os.path.exists', return_value=True): + + mock_path.return_value = '/test/path/token.json' + + result = load_token() + + assert result is None + + def test_load_token_io_error(self): + """Test loading token with IO error""" + with patch('builtins.open', side_effect=IOError("Permission denied")), \ + patch('litellm.proxy.client.cli.commands.auth.get_token_file_path') as mock_path, \ + patch('os.path.exists', return_value=True): + + mock_path.return_value = '/test/path/token.json' + + result = load_token() + + assert result is None + + def test_clear_token_file_exists(self): + """Test clearing token when file exists""" + with patch('litellm.proxy.client.cli.commands.auth.get_token_file_path') as mock_path, \ + patch('os.path.exists', return_value=True), \ + patch('os.remove') as mock_remove: + + mock_path.return_value = '/test/path/token.json' + + clear_token() + + mock_remove.assert_called_once_with('/test/path/token.json') + + def test_clear_token_file_not_exists(self): + """Test clearing token when file doesn't exist""" + with patch('litellm.proxy.client.cli.commands.auth.get_token_file_path') as mock_path, \ + patch('os.path.exists', return_value=False), \ + patch('os.remove') as mock_remove: + + mock_path.return_value = '/test/path/token.json' + + clear_token() + + mock_remove.assert_not_called() + + def test_get_stored_api_key_success(self): + """Test getting stored API key successfully""" + token_data = { + 'key': 'test-api-key-123', + 'user_id': 'test-user' + } + + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=token_data): + result = get_stored_api_key() + assert result == 'test-api-key-123' + + def test_get_stored_api_key_no_token(self): + """Test getting stored API key when no token exists""" + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=None): + result = get_stored_api_key() + assert result is None + + def test_get_stored_api_key_no_key_field(self): + """Test getting stored API key when token has no key field""" + token_data = { + 'user_id': 'test-user' + } + + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=token_data): + result = get_stored_api_key() + assert result is None + + +class TestLoginCommand: + """Test login CLI command""" + + def setup_method(self): + """Setup for each test""" + self.runner = CliRunner() + + def test_login_success(self): + """Test successful login flow""" + mock_context = Mock() + mock_context.obj = {"base_url": "https://test.example.com"} + + # Mock the requests for successful authentication + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "status": "ready", + "key": "sk-test-api-key-123" + } + + with patch('webbrowser.open') as mock_browser, \ + patch('requests.get', return_value=mock_response) as mock_get, \ + patch('litellm.proxy.client.cli.commands.auth.save_token') as mock_save, \ + patch('litellm.proxy.client.cli.interface.show_commands') as mock_show_commands, \ + patch('uuid.uuid4', return_value='test-uuid-123'): + + result = self.runner.invoke(login, obj=mock_context.obj) + + assert result.exit_code == 0 + assert "✅ Login successful!" in result.output + assert "API Key: sk-test-api-key-123" in result.output + + # Verify browser was opened with correct URL + mock_browser.assert_called_once() + call_args = mock_browser.call_args[0][0] + assert "https://test.example.com/sso/key/generate" in call_args + assert "sk-test-uuid-123" in call_args + + # Verify token was saved + mock_save.assert_called_once() + saved_data = mock_save.call_args[0][0] + assert saved_data['key'] == 'sk-test-api-key-123' + assert saved_data['user_id'] == 'cli-user' + + # Verify commands were shown + mock_show_commands.assert_called_once() + + def test_login_timeout(self): + """Test login timeout scenario""" + mock_context = Mock() + mock_context.obj = {"base_url": "https://test.example.com"} + + # Mock response that never returns ready status + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = {"status": "pending"} + + with patch('webbrowser.open'), \ + patch('requests.get', return_value=mock_response), \ + patch('time.sleep') as mock_sleep, \ + patch('uuid.uuid4', return_value='test-uuid-123'): + + # Mock time.sleep to avoid actual delays in tests + result = self.runner.invoke(login, obj=mock_context.obj) + + assert result.exit_code == 0 + assert "❌ Authentication timed out" in result.output + + def test_login_http_error(self): + """Test login with HTTP error""" + mock_context = Mock() + mock_context.obj = {"base_url": "https://test.example.com"} + + # Mock response with HTTP error + mock_response = Mock() + mock_response.status_code = 500 + + with patch('webbrowser.open'), \ + patch('requests.get', return_value=mock_response), \ + patch('time.sleep'), \ + patch('uuid.uuid4', return_value='test-uuid-123'): + + result = self.runner.invoke(login, obj=mock_context.obj) + + assert result.exit_code == 0 + assert "❌ Authentication timed out" in result.output + + def test_login_request_exception(self): + """Test login with request exception""" + import requests + mock_context = Mock() + mock_context.obj = {"base_url": "https://test.example.com"} + + with patch('webbrowser.open'), \ + patch('requests.get', side_effect=requests.RequestException("Connection failed")), \ + patch('time.sleep'), \ + patch('uuid.uuid4', return_value='test-uuid-123'): + + result = self.runner.invoke(login, obj=mock_context.obj) + + assert result.exit_code == 0 + assert "❌ Authentication timed out" in result.output + + def test_login_keyboard_interrupt(self): + """Test login cancelled by user""" + mock_context = Mock() + mock_context.obj = {"base_url": "https://test.example.com"} + + with patch('webbrowser.open'), \ + patch('requests.get', side_effect=KeyboardInterrupt), \ + patch('uuid.uuid4', return_value='test-uuid-123'): + + result = self.runner.invoke(login, obj=mock_context.obj) + + assert result.exit_code == 0 + assert "❌ Authentication cancelled by user" in result.output + + def test_login_no_api_key_in_response(self): + """Test login when response doesn't contain API key""" + mock_context = Mock() + mock_context.obj = {"base_url": "https://test.example.com"} + + # Mock response without API key + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "status": "ready" + # Missing 'key' field + } + + with patch('webbrowser.open'), \ + patch('requests.get', return_value=mock_response), \ + patch('time.sleep'), \ + patch('uuid.uuid4', return_value='test-uuid-123'): + + result = self.runner.invoke(login, obj=mock_context.obj) + + assert result.exit_code == 0 + assert "❌ Authentication timed out" in result.output + + def test_login_general_exception(self): + """Test login with general exception (not requests exception)""" + mock_context = Mock() + mock_context.obj = {"base_url": "https://test.example.com"} + + with patch('webbrowser.open'), \ + patch('requests.get', side_effect=ValueError("Invalid value")), \ + patch('uuid.uuid4', return_value='test-uuid-123'): + + result = self.runner.invoke(login, obj=mock_context.obj) + + assert result.exit_code == 0 + assert "❌ Authentication failed: Invalid value" in result.output + + +class TestLogoutCommand: + """Test logout CLI command""" + + def setup_method(self): + """Setup for each test""" + self.runner = CliRunner() + + def test_logout_success(self): + """Test successful logout""" + with patch('litellm.proxy.client.cli.commands.auth.clear_token') as mock_clear: + result = self.runner.invoke(logout) + + assert result.exit_code == 0 + assert "✅ Logged out successfully" in result.output + mock_clear.assert_called_once() + + +class TestWhoamiCommand: + """Test whoami CLI command""" + + def setup_method(self): + """Setup for each test""" + self.runner = CliRunner() + + def test_whoami_authenticated(self): + """Test whoami when user is authenticated""" + token_data = { + 'user_email': 'test@example.com', + 'user_id': 'test-user-123', + 'user_role': 'admin', + 'timestamp': time.time() - 3600 # 1 hour ago + } + + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=token_data): + result = self.runner.invoke(whoami) + + assert result.exit_code == 0 + assert "✅ Authenticated" in result.output + assert "test@example.com" in result.output + assert "test-user-123" in result.output + assert "admin" in result.output + assert "Token age: 1.0 hours" in result.output + + def test_whoami_not_authenticated(self): + """Test whoami when user is not authenticated""" + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=None): + result = self.runner.invoke(whoami) + + assert result.exit_code == 0 + assert "❌ Not authenticated" in result.output + assert "Run 'litellm-proxy login'" in result.output + + def test_whoami_old_token(self): + """Test whoami with old token showing warning""" + token_data = { + 'user_email': 'test@example.com', + 'user_id': 'test-user-123', + 'user_role': 'admin', + 'timestamp': time.time() - (25 * 3600) # 25 hours ago + } + + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=token_data): + result = self.runner.invoke(whoami) + + assert result.exit_code == 0 + assert "✅ Authenticated" in result.output + assert "⚠️ Warning: Token is more than 24 hours old" in result.output + + def test_whoami_missing_fields(self): + """Test whoami with token missing some fields""" + token_data = { + 'timestamp': time.time() - 3600 + # Missing user_email, user_id, user_role + } + + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=token_data): + result = self.runner.invoke(whoami) + + assert result.exit_code == 0 + assert "✅ Authenticated" in result.output + assert "Unknown" in result.output # Should show "Unknown" for missing fields + + def test_whoami_no_timestamp(self): + """Test whoami with token missing timestamp""" + token_data = { + 'user_email': 'test@example.com', + 'user_id': 'test-user-123', + 'user_role': 'admin' + # Missing timestamp + } + + with patch('litellm.proxy.client.cli.commands.auth.load_token', return_value=token_data), \ + patch('time.time', return_value=1000): + + result = self.runner.invoke(whoami) + + assert result.exit_code == 0 + assert "✅ Authenticated" in result.output + # Should calculate age based on timestamp=0 + assert "Token age:" in result.output diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 60199b335a5..583e953dca7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -772,3 +772,276 @@ async def test_get_generic_sso_response_with_empty_headers(): ) assert result == mock_sso_response + + +class TestCLISSOCallbackFunction: + """Test the cli_sso_callback function specifically""" + + def test_cli_sso_callback_validation_invalid_key(self): + """Test CLI SSO callback input validation for invalid key format""" + # Test the validation logic without hitting the database + invalid_keys = [ + None, + "", + "invalid-key", + "not-sk-key", + "sk", # too short + ] + + for invalid_key in invalid_keys: + # This should fail validation before any database operations + # We can test this by checking if the key starts with 'sk-' + if not invalid_key or not invalid_key.startswith('sk-'): + # This would trigger the validation error + assert True # Validation works as expected + + +class TestCLIPollingFunction: + """Test the cli_poll_key function specifically""" + + def test_cli_poll_key_validation_invalid_format(self): + """Test CLI polling key format validation""" + # Test key format validation logic + invalid_keys = [ + "invalid-key", + "not-sk-key", + "", + "sk", # too short + ] + + for invalid_key in invalid_keys: + # Validation logic: key must start with 'sk-' + if not invalid_key.startswith('sk-'): + # This would trigger the validation error in the actual function + assert True # Validation works as expected + + +class TestAuthCallbackRouting: + """Test the auth_callback function routing logic""" + + def test_cli_state_detection_and_routing(self): + """Test that CLI states are properly detected and would route to CLI callback""" + from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX + + # Test CLI state detection logic + cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123" + + # This mimics the logic in auth_callback + if cli_state and cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"): + # Extract the key ID from the state + key_id = cli_state.split(":", 1)[1] + assert key_id == "sk-test123" + else: + assert False, "CLI state should have been detected" + + def test_non_cli_state_routing(self): + """Test that non-CLI states don't trigger CLI routing""" + from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX + + non_cli_states = [ + "regular_oauth_state", + "some_random_string", + None, + "", + "not_session_token:something" + ] + + for state in non_cli_states: + # This mimics the routing logic in auth_callback + should_route_to_cli = state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") + assert not should_route_to_cli, f"State '{state}' should not route to CLI" + + +class TestGoogleLoginCLIIntegration: + """Test the google_login function with CLI parameters""" + + def test_google_login_cli_state_generation(self): + """Test that google_login generates CLI state when CLI parameters are provided""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Test the CLI state generation logic used in google_login + source = "litellm-cli" + key = "sk-test123" + + cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key) + + assert cli_state is not None + assert cli_state.startswith("litellm-session-token:") + assert "sk-test123" in cli_state + + def test_google_login_no_cli_state_when_missing_params(self): + """Test that google_login doesn't generate CLI state when CLI parameters are missing""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Test various parameter combinations that shouldn't generate CLI state + test_cases = [ + (None, None), + ("litellm-cli", None), + (None, "sk-test123"), + ("wrong-source", "sk-test123"), + ] + + for source, key in test_cases: + cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key) + assert cli_state is None, f"CLI state should not be generated for source='{source}', key='{key}'" + + +class TestSSOHandlerIntegration: + """Test SSOAuthenticationHandler methods""" + + def test_should_use_sso_handler(self): + """Test the SSO handler detection logic""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Test that SSO handler is used when client IDs are provided + assert SSOAuthenticationHandler.should_use_sso_handler(google_client_id="test") is True + assert SSOAuthenticationHandler.should_use_sso_handler(microsoft_client_id="test") is True + assert SSOAuthenticationHandler.should_use_sso_handler(generic_client_id="test") is True + + # Test that SSO handler is not used when no client IDs are provided + assert SSOAuthenticationHandler.should_use_sso_handler() is False + assert SSOAuthenticationHandler.should_use_sso_handler(None, None, None) is False + + def test_get_redirect_url_for_sso(self): + """Test the redirect URL generation for SSO""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Mock request object + mock_request = MagicMock() + mock_request.base_url = "https://test.litellm.ai/" + + # Test redirect URL generation + redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( + request=mock_request, + sso_callback_route="sso/callback" + ) + + assert redirect_url.startswith("https://test.litellm.ai") + assert "sso/callback" in redirect_url + + +class TestUISSO_FunctionsExistence: + """Test that all the new functions exist and are importable""" + + def test_cli_sso_callback_exists(self): + """Test that cli_sso_callback function exists""" + from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback + assert callable(cli_sso_callback) + + def test_cli_poll_key_exists(self): + """Test that cli_poll_key function exists""" + from litellm.proxy.management_endpoints.ui_sso import cli_poll_key + assert callable(cli_poll_key) + + def test_auth_callback_exists(self): + """Test that auth_callback function exists""" + from litellm.proxy.management_endpoints.ui_sso import auth_callback + assert callable(auth_callback) + + def test_google_login_exists(self): + """Test that google_login function exists""" + from litellm.proxy.management_endpoints.ui_sso import google_login + assert callable(google_login) + + def test_sso_authentication_handler_exists(self): + """Test that SSOAuthenticationHandler class exists with new methods""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + # Check that the class exists + assert SSOAuthenticationHandler is not None + + # Check that the new _get_cli_state method exists + assert hasattr(SSOAuthenticationHandler, '_get_cli_state') + assert callable(SSOAuthenticationHandler._get_cli_state) + + +class TestSSOStateHandling: + """Test the SSO state handling for CLI authentication""" + + def test_get_cli_state_valid(self): + """Test generating CLI state with valid parameters""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + state = SSOAuthenticationHandler._get_cli_state(source="litellm-cli", key="sk-test123") + + assert state is not None + assert state.startswith("litellm-session-token:") + assert "sk-test123" in state + + def test_get_cli_state_invalid_source(self): + """Test generating CLI state with invalid source""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + state = SSOAuthenticationHandler._get_cli_state(source="invalid_source", key="sk-test123") + + assert state is None + + def test_get_cli_state_no_key(self): + """Test generating CLI state without key""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + state = SSOAuthenticationHandler._get_cli_state(source="litellm-cli", key=None) + + assert state is None + + def test_get_cli_state_no_source(self): + """Test generating CLI state without source""" + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + state = SSOAuthenticationHandler._get_cli_state(source=None, key="sk-test123") + + assert state is None + + +class TestStateRouting: + """Test state parameter routing logic""" + + def test_cli_state_detection(self): + """Test detection of CLI state parameters""" + from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX + + # Test CLI state format + cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123" + assert cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") + + # Test extraction of key from state + key_id = cli_state.split(":", 1)[1] + assert key_id == "sk-test123" + + def test_non_cli_state_detection(self): + """Test detection of non-CLI state parameters""" + from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX + + # Test various non-CLI states + test_states = [ + "regular_oauth_state", + "some_random_string", + None, + "", + "not_session_token:something" + ] + + for state in test_states: + if state: + assert not state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:") + else: + assert state != f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:" + + +class TestHTMLIntegration: + """Test HTML rendering integration with CLI flow""" + + def test_html_render_utils_import(self): + """Test that HTML render utils can be imported correctly""" + from litellm.proxy.common_utils.html_forms.cli_sso_success import ( + render_cli_sso_success_page, + ) + + # Test that function exists and is callable + assert callable(render_cli_sso_success_page) + + # Test that it returns expected type + html = render_cli_sso_success_page() + + assert isinstance(html, str) + assert len(html) > 0