mirror of
https://github.com/tiennm99/litellm.git
synced 2026-07-20 08:18:16 +00:00
[Feat] AI Gateway Auth - Allow using JWTs for signing in with Proxy CLI (#16756)
* fix auth * get_cli_jwt_auth_token * fix linting * test fixes * docs * test fixes * fix refactor
This commit is contained in:
@@ -9,6 +9,26 @@ Use the litellm cli to authenticate to the LiteLLM Gateway. This is great if you
|
||||
|
||||
## Usage
|
||||
|
||||
### Prerequisites - Start LiteLLM Proxy with Beta Flag
|
||||
|
||||
:::warning[Beta Feature - Required]
|
||||
|
||||
CLI SSO Authentication is currently in beta. You must set this environment variable **when starting up your LiteLLM Proxy**:
|
||||
|
||||
```bash
|
||||
export EXPERIMENTAL_UI_LOGIN="True"
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
Or add it to your proxy startup command:
|
||||
|
||||
```bash
|
||||
EXPERIMENTAL_UI_LOGIN="True" litellm --config config.yaml
|
||||
```
|
||||
|
||||
:::
|
||||
|
||||
### Steps
|
||||
|
||||
1. **Install the CLI**
|
||||
|
||||
@@ -33,6 +53,8 @@ Use the litellm cli to authenticate to the LiteLLM Gateway. This is great if you
|
||||
|
||||
2. **Set up environment variables**
|
||||
|
||||
On your local machine, set the proxy URL:
|
||||
|
||||
```bash
|
||||
export LITELLM_PROXY_URL=http://localhost:4000
|
||||
```
|
||||
|
||||
@@ -67,7 +67,26 @@ For an indepth guide, see [CLI Authentication](./cli_sso).
|
||||
|
||||
:::
|
||||
|
||||
### Prerequisites
|
||||
|
||||
:::warning[Beta Feature - Required Environment Variable]
|
||||
|
||||
CLI SSO Authentication is currently in beta. You must set this environment variable **when starting up your LiteLLM Proxy**:
|
||||
|
||||
```bash
|
||||
export EXPERIMENTAL_UI_LOGIN="True"
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
Or add it to your proxy startup command:
|
||||
|
||||
```bash
|
||||
EXPERIMENTAL_UI_LOGIN="True" litellm --config config.yaml
|
||||
```
|
||||
|
||||
:::
|
||||
|
||||
### Steps
|
||||
|
||||
1. **Set up the proxy URL**
|
||||
|
||||
|
||||
@@ -1044,6 +1044,8 @@ 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"
|
||||
CLI_SSO_SESSION_CACHE_KEY_PREFIX = "cli_sso_session"
|
||||
CLI_JWT_TOKEN_NAME = "cli-jwt-token"
|
||||
|
||||
########################### DB CRON JOB NAMES ###########################
|
||||
DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job"
|
||||
|
||||
@@ -1352,6 +1352,57 @@ class ExperimentalUIJWTToken:
|
||||
|
||||
return encrypt_value_helper(valid_token.model_dump_json(exclude_none=True))
|
||||
|
||||
@staticmethod
|
||||
def get_cli_jwt_auth_token(
|
||||
user_info: LiteLLM_UserTable, team_id: Optional[str] = None
|
||||
) -> str:
|
||||
"""
|
||||
Generate a JWT token for CLI authentication with 24-hour expiration.
|
||||
|
||||
Args:
|
||||
user_info: User information from the database
|
||||
team_id: Team ID for the user (optional, uses user's team if available)
|
||||
|
||||
Returns:
|
||||
Encrypted JWT token string
|
||||
"""
|
||||
from datetime import timedelta
|
||||
|
||||
from litellm.constants import CLI_JWT_TOKEN_NAME
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
if user_info.user_role is None:
|
||||
raise Exception("User role is required for CLI JWT login")
|
||||
|
||||
# Calculate expiration time (24 hours from now - matching old CLI key behavior)
|
||||
expiration_time = get_utc_datetime() + timedelta(hours=24)
|
||||
|
||||
# Format the expiration time as ISO 8601 string
|
||||
expires = expiration_time.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "+00:00"
|
||||
|
||||
# Use provided team_id, or fall back to user's teams if available
|
||||
_team_id = team_id
|
||||
if _team_id is None and hasattr(user_info, "teams") and user_info.teams:
|
||||
# Use first team if user has teams
|
||||
_team_id = user_info.teams[0] if len(user_info.teams) > 0 else None
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token=CLI_JWT_TOKEN_NAME,
|
||||
key_name=CLI_JWT_TOKEN_NAME,
|
||||
key_alias=CLI_JWT_TOKEN_NAME,
|
||||
max_budget=litellm.max_ui_session_budget,
|
||||
expires=expires,
|
||||
user_id=user_info.user_id,
|
||||
team_id=_team_id,
|
||||
models=user_info.models,
|
||||
max_parallel_requests=None,
|
||||
user_role=LitellmUserRoles(user_info.user_role),
|
||||
)
|
||||
|
||||
return encrypt_value_helper(valid_token.model_dump_json(exclude_none=True))
|
||||
|
||||
@staticmethod
|
||||
def get_key_object_from_ui_hash_key(
|
||||
hashed_token: str,
|
||||
|
||||
@@ -97,24 +97,6 @@ def display_teams_table(teams: List[Dict[str, Any]]) -> None:
|
||||
console.print(table)
|
||||
|
||||
|
||||
def get_user_teams(base_url: str, api_key: str, user_id: str) -> List[Dict[str, Any]]:
|
||||
"""Fetch teams for the current user"""
|
||||
from litellm.proxy.client import Client
|
||||
|
||||
client = Client(base_url=base_url, api_key=api_key)
|
||||
try:
|
||||
response = client.teams.list_v2(user_id=user_id)
|
||||
# Extract just the teams array from the paginated response
|
||||
if isinstance(response, dict) and 'teams' in response:
|
||||
return response['teams']
|
||||
else:
|
||||
# Fallback in case the response structure is different
|
||||
return response if isinstance(response, list) else []
|
||||
except Exception as e:
|
||||
click.echo(f"❌ Error fetching teams: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def get_key_input():
|
||||
"""Get a single key input from the user (cross-platform)"""
|
||||
try:
|
||||
@@ -279,57 +261,172 @@ def prompt_team_selection_fallback(teams: List[Dict[str, Any]]) -> Optional[Dict
|
||||
return None
|
||||
|
||||
|
||||
def update_key_with_team(base_url: str, api_key: str, team_id: str) -> bool:
|
||||
"""Update the API key to be associated with the selected team"""
|
||||
from litellm.proxy._types import SpecialModelNames
|
||||
from litellm.proxy.client import Client
|
||||
|
||||
client = Client(base_url=base_url, api_key=api_key)
|
||||
try:
|
||||
client.keys.update(key=api_key, team_id=team_id, models=[SpecialModelNames.all_team_models.value])
|
||||
click.echo(f"✅ Successfully assigned key to team: {team_id}")
|
||||
return True
|
||||
except requests.exceptions.HTTPError as e:
|
||||
# Bubble up the response text for detailed error info
|
||||
error_msg = e.response.text if e.response else str(e)
|
||||
click.echo(f"❌ Error updating key with team: {error_msg}")
|
||||
return False
|
||||
except Exception as e:
|
||||
click.echo(f"❌ Error updating key with team: {e}")
|
||||
return False
|
||||
|
||||
|
||||
# Polling-based authentication - no local server needed
|
||||
|
||||
def _handle_team_assignment(base_url: str, api_key: str, user_id: str) -> None:
|
||||
"""Handle team fetching and assignment for the authenticated user."""
|
||||
click.echo("\n" + "="*60)
|
||||
click.echo("📋 Fetching your teams...")
|
||||
def _poll_for_authentication(
|
||||
base_url: str, key_id: str
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Poll the server for authentication completion and handle team selection.
|
||||
|
||||
teams = get_user_teams(
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
user_id=user_id,
|
||||
)
|
||||
Returns:
|
||||
Dictionary with authentication data if successful, None otherwise
|
||||
"""
|
||||
poll_url = f"{base_url}/sso/cli/poll/{key_id}"
|
||||
timeout = 300 # 5 minute timeout
|
||||
poll_interval = 2 # Poll every 2 seconds
|
||||
|
||||
if teams:
|
||||
# Prompt for team selection (will display teams interactively)
|
||||
selected_team = prompt_team_selection(teams)
|
||||
|
||||
if selected_team:
|
||||
team_id = selected_team.get('team_id')
|
||||
if team_id:
|
||||
click.echo(f"\n🔄 Assigning your key to team: {selected_team.get('team_alias', team_id)}")
|
||||
success = update_key_with_team(base_url, api_key, team_id)
|
||||
if success:
|
||||
click.echo(f"✅ Your CLI key is now associated with team: {selected_team.get('team_alias', team_id)}")
|
||||
click.echo(f"🎯 You can now access models: {', '.join(selected_team.get('models', ['All models']))}")
|
||||
else:
|
||||
click.echo("⚠️ Key assignment failed, but you can still use the CLI")
|
||||
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":
|
||||
# Check if we need team selection first
|
||||
if data.get("requires_team_selection"):
|
||||
# Server returned teams list without JWT - need to select team
|
||||
teams = data.get("teams", [])
|
||||
user_id = data.get("user_id")
|
||||
|
||||
if teams and len(teams) > 1:
|
||||
# User has multiple teams - let them select
|
||||
jwt_with_team = _handle_team_selection_during_polling(
|
||||
base_url=base_url,
|
||||
key_id=key_id,
|
||||
teams=teams
|
||||
)
|
||||
|
||||
# Use the team-specific JWT if selection succeeded
|
||||
if jwt_with_team:
|
||||
return {
|
||||
"api_key": jwt_with_team,
|
||||
"user_id": user_id,
|
||||
"teams": teams,
|
||||
"team_id": None # Set by server in JWT
|
||||
}
|
||||
else:
|
||||
# Selection failed or was skipped - poll again without team_id
|
||||
click.echo("⚠️ Team selection skipped, retrying...")
|
||||
continue
|
||||
else:
|
||||
# Shouldn't happen, but fallback
|
||||
click.echo("⚠️ No teams available, retrying...")
|
||||
continue
|
||||
else:
|
||||
# JWT is ready (single team or team already selected)
|
||||
api_key = data.get("key")
|
||||
user_id = data.get("user_id")
|
||||
teams = data.get("teams", [])
|
||||
team_id = data.get("team_id")
|
||||
|
||||
# Show which team was assigned
|
||||
if team_id and len(teams) == 1:
|
||||
click.echo(f"\n✅ Automatically assigned to team: {team_id}")
|
||||
|
||||
if api_key:
|
||||
return {
|
||||
"api_key": api_key,
|
||||
"user_id": user_id,
|
||||
"teams": teams,
|
||||
"team_id": team_id
|
||||
}
|
||||
elif data.get("status") == "pending":
|
||||
# Still pending
|
||||
if attempt % 10 == 0: # Show progress every 20 seconds
|
||||
click.echo("Still waiting for authentication...")
|
||||
else:
|
||||
click.echo("ℹ️ Continuing without team assignment. You can assign a team later using the CLI.")
|
||||
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)
|
||||
|
||||
# Timeout reached
|
||||
return None
|
||||
|
||||
|
||||
def _handle_team_selection_during_polling(
|
||||
base_url: str, key_id: str, teams: List[str]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Handle team selection and re-poll with selected team_id.
|
||||
|
||||
Args:
|
||||
teams: List of team IDs (strings)
|
||||
|
||||
Returns:
|
||||
The JWT token with the selected team, or None if selection was skipped
|
||||
"""
|
||||
if not teams:
|
||||
click.echo("ℹ️ No teams found. You can create or join teams using the web interface.")
|
||||
return None
|
||||
|
||||
click.echo("\n" + "="*60)
|
||||
click.echo("📋 Select a team for your CLI session...")
|
||||
|
||||
# Display teams as simple list since we only have IDs
|
||||
console = Console()
|
||||
table = Table(title="Available Teams")
|
||||
table.add_column("Index", style="cyan", no_wrap=True)
|
||||
table.add_column("Team ID", style="green")
|
||||
|
||||
for i, team in enumerate(teams):
|
||||
table.add_row(str(i + 1), team)
|
||||
|
||||
console.print(table)
|
||||
|
||||
# Simple selection
|
||||
team_id: Optional[str] = None
|
||||
while True:
|
||||
try:
|
||||
choice = click.prompt(
|
||||
"\nSelect a team by entering the index number (or 'skip' to use first team)",
|
||||
type=str
|
||||
).strip()
|
||||
|
||||
if choice.lower() == 'skip':
|
||||
team_id = teams[0] if teams else None
|
||||
break
|
||||
|
||||
index = int(choice) - 1
|
||||
if 0 <= index < len(teams):
|
||||
team_id = teams[index]
|
||||
break
|
||||
else:
|
||||
click.echo(f"❌ Invalid selection. Please enter a number between 1 and {len(teams)}")
|
||||
except ValueError:
|
||||
click.echo("❌ Invalid input. Please enter a number or 'skip'")
|
||||
except KeyboardInterrupt:
|
||||
click.echo("\n❌ Team selection cancelled.")
|
||||
return None
|
||||
|
||||
if not team_id:
|
||||
click.echo("ℹ️ No team selected.")
|
||||
return None
|
||||
|
||||
click.echo(f"\n🔄 Generating JWT for team: {team_id}")
|
||||
|
||||
# Re-poll with team_id to get JWT with correct team
|
||||
try:
|
||||
poll_url = f"{base_url}/sso/cli/poll/{key_id}?team_id={team_id}"
|
||||
response = requests.get(poll_url, timeout=10)
|
||||
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
if data.get("status") == "ready":
|
||||
jwt_token = data.get("key")
|
||||
if jwt_token:
|
||||
click.echo(f"✅ Successfully generated JWT for team: {team_id}")
|
||||
return jwt_token
|
||||
|
||||
click.echo(f"❌ Failed to get JWT with team. Status: {response.status_code}")
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
click.echo(f"❌ Error getting JWT with team: {e}")
|
||||
return None
|
||||
|
||||
|
||||
@click.command(name="login")
|
||||
@@ -337,7 +434,6 @@ def _handle_team_assignment(base_url: str, api_key: str, user_id: str) -> None:
|
||||
def login(ctx: click.Context):
|
||||
"""Login to LiteLLM proxy using SSO authentication"""
|
||||
from litellm._uuid import uuid
|
||||
|
||||
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
|
||||
from litellm.proxy.client.cli.interface import show_commands
|
||||
|
||||
@@ -365,59 +461,37 @@ def login(ctx: click.Context):
|
||||
# Open browser
|
||||
webbrowser.open(sso_url)
|
||||
|
||||
# Poll for key creation
|
||||
# Poll for authentication completion
|
||||
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
|
||||
auth_result = _poll_for_authentication(base_url=base_url, key_id=key_id)
|
||||
|
||||
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")
|
||||
|
||||
# Handle team assignment
|
||||
_handle_team_assignment(base_url, api_key, data.get("user_id"))
|
||||
|
||||
# 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}")
|
||||
if auth_result:
|
||||
api_key = auth_result["api_key"]
|
||||
user_id = auth_result["user_id"]
|
||||
|
||||
time.sleep(poll_interval)
|
||||
|
||||
click.echo("❌ Authentication timed out. Please try again.")
|
||||
return
|
||||
# Save token data (simplified for CLI - we just need the key)
|
||||
save_token({
|
||||
'key': api_key,
|
||||
'user_id': user_id or 'cli-user',
|
||||
'user_email': 'unknown',
|
||||
'user_role': 'cli',
|
||||
'auth_header_name': 'Authorization',
|
||||
'jwt_token': '',
|
||||
'timestamp': time.time()
|
||||
})
|
||||
|
||||
click.echo("\n✅ Login successful!")
|
||||
click.echo(f"JWT Token: {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
|
||||
else:
|
||||
click.echo("❌ Authentication timed out. Please try again.")
|
||||
return
|
||||
|
||||
except KeyboardInterrupt:
|
||||
click.echo("\n❌ Authentication cancelled by user.")
|
||||
|
||||
@@ -707,71 +707,22 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
|
||||
)
|
||||
|
||||
|
||||
async def _regenerate_cli_key(
|
||||
existing_key: str, new_key: str, user_id: Optional[str] = None
|
||||
) -> None:
|
||||
"""Regenerate an existing CLI key with a new token"""
|
||||
from litellm.proxy._types import RegenerateKeyRequest, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
regenerate_key_fn,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(f"Regenerating existing CLI key: {existing_key}")
|
||||
|
||||
admin_user_dict = UserAPIKeyAuth.get_litellm_cli_user_api_key_auth()
|
||||
|
||||
regenerate_request = RegenerateKeyRequest(
|
||||
key=existing_key,
|
||||
new_key=new_key,
|
||||
duration="24hr",
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
await regenerate_key_fn(
|
||||
key=existing_key, data=regenerate_request, user_api_key_dict=admin_user_dict
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(f"Regenerated CLI key: {new_key}")
|
||||
|
||||
|
||||
async def _create_new_cli_key(
|
||||
key: str,
|
||||
user_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Create a new CLI key"""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
generate_key_helper_fn,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info("Creating new CLI key")
|
||||
|
||||
await generate_key_helper_fn(
|
||||
request_type="key",
|
||||
duration="24hr",
|
||||
key_max_budget=litellm.max_ui_session_budget,
|
||||
aliases={},
|
||||
config={},
|
||||
spend=0,
|
||||
user_id=user_id,
|
||||
table_name="key",
|
||||
token=key,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(f"Created new CLI key: {key}")
|
||||
|
||||
|
||||
async def cli_sso_callback(
|
||||
request: Request,
|
||||
key: Optional[str] = None,
|
||||
existing_key: Optional[str] = None,
|
||||
result: Optional[Union[OpenID, dict]] = None,
|
||||
):
|
||||
"""CLI SSO callback - regenerates existing CLI key or creates new one"""
|
||||
"""CLI SSO callback - stores session info for JWT generation on polling"""
|
||||
verbose_proxy_logger.info(
|
||||
f"CLI SSO callback for key: {key}, existing_key: {existing_key}"
|
||||
)
|
||||
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not key or not key.startswith("sk-"):
|
||||
raise HTTPException(
|
||||
@@ -784,24 +735,60 @@ async def cli_sso_callback(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
if result is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="SSO authentication failed - no result returned from provider",
|
||||
)
|
||||
|
||||
# After None check, cast to non-None type for type checker
|
||||
result_non_none: Union[OpenID, dict] = cast(Union[OpenID, dict], result)
|
||||
|
||||
parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result(
|
||||
result=result
|
||||
result=result_non_none
|
||||
)
|
||||
verbose_proxy_logger.debug(f"parsed_openid_result: {parsed_openid_result}")
|
||||
|
||||
try:
|
||||
if existing_key:
|
||||
await _regenerate_cli_key(
|
||||
existing_key=existing_key,
|
||||
new_key=key,
|
||||
user_id=parsed_openid_result.get("user_id"),
|
||||
)
|
||||
else:
|
||||
await _create_new_cli_key(
|
||||
key=key,
|
||||
user_id=parsed_openid_result.get("user_id"),
|
||||
# Get full user info from DB
|
||||
user_info = await get_user_info_from_db(
|
||||
result=result_non_none,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_email=parsed_openid_result.get("user_email"),
|
||||
user_defined_values=None,
|
||||
alternate_user_id=parsed_openid_result.get("user_id"),
|
||||
)
|
||||
|
||||
if user_info is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail="Failed to retrieve user information from SSO"
|
||||
)
|
||||
|
||||
# Store session info in cache (10 min TTL)
|
||||
from litellm.constants import CLI_SSO_SESSION_CACHE_KEY_PREFIX
|
||||
|
||||
# Get all teams from user_info - CLI will let user select which one
|
||||
teams: List[str] = []
|
||||
if hasattr(user_info, "teams") and user_info.teams:
|
||||
teams = user_info.teams if isinstance(user_info.teams, list) else []
|
||||
|
||||
session_data = {
|
||||
"user_id": user_info.user_id,
|
||||
"user_role": user_info.user_role,
|
||||
"models": user_info.models if hasattr(user_info, "models") else [],
|
||||
"user_email": parsed_openid_result.get("user_email"),
|
||||
"teams": teams,
|
||||
}
|
||||
|
||||
cache_key = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:{key}"
|
||||
user_api_key_cache.set_cache(key=cache_key, value=session_data, ttl=600)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}"
|
||||
)
|
||||
|
||||
# Return success page
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
@@ -813,47 +800,103 @@ async def cli_sso_callback(
|
||||
return HTMLResponse(content=html_content, status_code=200)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error with CLI key: {e}")
|
||||
verbose_proxy_logger.error(f"Error with CLI SSO callback: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Failed to process CLI key: {str(e)}"
|
||||
status_code=500, detail=f"Failed to process CLI SSO: {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._types import LiteLLM_VerificationToken
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
|
||||
"""
|
||||
CLI polling endpoint - retrieves session from cache and generates JWT.
|
||||
|
||||
Flow:
|
||||
1. First poll (no team_id): Returns teams list without generating JWT
|
||||
2. Second poll (with team_id): Generates JWT with selected team and deletes session
|
||||
|
||||
Args:
|
||||
key_id: The session key ID
|
||||
team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams.
|
||||
"""
|
||||
from litellm.constants import CLI_SSO_SESSION_CACHE_KEY_PREFIX
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
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
|
||||
# Look up session in cache
|
||||
cache_key = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:{key_id}"
|
||||
session_data = user_api_key_cache.get_cache(key=cache_key)
|
||||
|
||||
hashed_token = hash_token(key_id)
|
||||
if session_data:
|
||||
user_teams = session_data.get("teams", [])
|
||||
user_id = session_data["user_id"]
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"CLI poll: user={user_id}, team_id={team_id}, user_teams={user_teams}, num_teams={len(user_teams)}"
|
||||
)
|
||||
|
||||
# If no team_id provided and user has teams, return teams list for selection
|
||||
# Don't generate JWT yet - let CLI select a team first
|
||||
if team_id is None and len(user_teams) > 1:
|
||||
verbose_proxy_logger.info(
|
||||
f"Returning teams list for user {user_id} to select from: {user_teams}"
|
||||
)
|
||||
return {
|
||||
"status": "ready",
|
||||
"user_id": user_id,
|
||||
"teams": user_teams,
|
||||
"requires_team_selection": True,
|
||||
}
|
||||
|
||||
# Validate team_id if provided
|
||||
if team_id is not None:
|
||||
if team_id not in user_teams:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"User does not belong to team: {team_id}. Available teams: {user_teams}",
|
||||
)
|
||||
else:
|
||||
# If no team_id provided and user has 0 or 1 team, use first team (or None)
|
||||
team_id = user_teams[0] if len(user_teams) > 0 else None
|
||||
|
||||
key_obj = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed_token}
|
||||
)
|
||||
key_obj = cast(LiteLLM_VerificationToken, key_obj)
|
||||
# Create user object for JWT generation
|
||||
user_info = LiteLLM_UserTable(
|
||||
user_id=user_id,
|
||||
user_role=session_data["user_role"],
|
||||
models=session_data.get("models", []),
|
||||
max_budget=litellm.max_ui_session_budget,
|
||||
)
|
||||
|
||||
if key_obj:
|
||||
verbose_proxy_logger.info(f"CLI key found: {key_id}")
|
||||
return {"status": "ready", "key": key_id, "user_id": key_obj.user_id}
|
||||
# Generate CLI JWT on-demand (24hr expiration)
|
||||
# Pass selected team_id to ensure JWT has correct team
|
||||
jwt_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
||||
user_info=user_info, team_id=team_id
|
||||
)
|
||||
|
||||
# Delete cache entry (single-use)
|
||||
user_api_key_cache.delete_cache(key=cache_key)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"CLI JWT generated for user: {user_id}, team: {team_id}"
|
||||
)
|
||||
return {
|
||||
"status": "ready",
|
||||
"key": jwt_token,
|
||||
"user_id": user_id,
|
||||
"team_id": team_id,
|
||||
"teams": user_teams,
|
||||
}
|
||||
else:
|
||||
return {"status": "pending"}
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error polling for CLI key: {e}")
|
||||
verbose_proxy_logger.error(f"Error polling for CLI JWT: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Error checking key status: {str(e)}"
|
||||
status_code=500, detail=f"Error checking session status: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -188,16 +188,19 @@ class TestLoginCommand:
|
||||
self.runner = CliRunner()
|
||||
|
||||
def test_login_success(self):
|
||||
"""Test successful login flow"""
|
||||
"""Test successful login flow with single team (JWT generated immediately)"""
|
||||
mock_context = Mock()
|
||||
mock_context.obj = {"base_url": "https://test.example.com"}
|
||||
|
||||
# Mock the requests for successful authentication
|
||||
# Mock the requests for successful authentication with single team
|
||||
mock_response = Mock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"status": "ready",
|
||||
"key": "sk-test-api-key-123"
|
||||
"key": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt",
|
||||
"user_id": "test-user-123",
|
||||
"team_id": "team-1",
|
||||
"teams": ["team-1"]
|
||||
}
|
||||
|
||||
with patch('webbrowser.open') as mock_browser, \
|
||||
@@ -210,7 +213,7 @@ class TestLoginCommand:
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert "✅ Login successful!" in result.output
|
||||
assert "API Key: sk-test-api-key-123" in result.output
|
||||
assert "Automatically assigned to team: team-1" in result.output
|
||||
|
||||
# Verify browser was opened with correct URL
|
||||
mock_browser.assert_called_once()
|
||||
@@ -218,11 +221,11 @@ class TestLoginCommand:
|
||||
assert "https://test.example.com/sso/key/generate" in call_args
|
||||
assert "sk-test-uuid-123" in call_args
|
||||
|
||||
# Verify token was saved
|
||||
# Verify JWT 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'
|
||||
assert saved_data['key'] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt"
|
||||
assert saved_data['user_id'] == 'test-user-123'
|
||||
|
||||
# Verify commands were shown
|
||||
mock_show_commands.assert_called_once()
|
||||
@@ -444,98 +447,107 @@ class TestCLIKeyRegenerationFlow:
|
||||
"""Setup for each test"""
|
||||
self.runner = CliRunner()
|
||||
|
||||
def test_login_with_existing_key_regeneration_flow(self):
|
||||
"""Test complete login flow when user has existing key - should regenerate it"""
|
||||
def test_login_with_team_selection_flow(self):
|
||||
"""Test complete login flow when user has multiple teams - should prompt for selection"""
|
||||
mock_context = Mock()
|
||||
mock_context.obj = {"base_url": "https://test.example.com"}
|
||||
|
||||
# Mock existing stored key
|
||||
existing_key = "sk-existing-key-123"
|
||||
|
||||
# Mock successful regeneration response
|
||||
mock_response = Mock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
# Mock first response - requires team selection
|
||||
mock_first_response = Mock()
|
||||
mock_first_response.status_code = 200
|
||||
mock_first_response.json.return_value = {
|
||||
"status": "ready",
|
||||
"key": "sk-regenerated-key-456" # New regenerated key
|
||||
"requires_team_selection": True,
|
||||
"user_id": "test-user-456",
|
||||
"teams": ["team-alpha", "team-beta", "team-gamma"]
|
||||
}
|
||||
|
||||
# Mock second response after team selection - JWT with selected team
|
||||
mock_second_response = Mock()
|
||||
mock_second_response.status_code = 200
|
||||
mock_second_response.json.return_value = {
|
||||
"status": "ready",
|
||||
"key": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.team-beta.jwt",
|
||||
"user_id": "test-user-456",
|
||||
"team_id": "team-beta",
|
||||
"teams": ["team-alpha", "team-beta", "team-gamma"]
|
||||
}
|
||||
|
||||
# Simulate user selecting team #2 (team-beta)
|
||||
with patch('webbrowser.open') as mock_browser, \
|
||||
patch('requests.get', return_value=mock_response) as mock_get, \
|
||||
patch('litellm.proxy.client.cli.commands.auth.get_stored_api_key', return_value=existing_key) as mock_get_stored, \
|
||||
patch('requests.get', side_effect=[mock_first_response, mock_second_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('litellm._uuid.uuid.uuid4', return_value='new-session-uuid-789'):
|
||||
patch('litellm._uuid.uuid.uuid4', return_value='session-uuid-456'), \
|
||||
patch('click.prompt', return_value='2'): # User selects index 2
|
||||
|
||||
result = self.runner.invoke(login, obj=mock_context.obj)
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert "✅ Login successful!" in result.output
|
||||
assert "API Key: sk-regenerated-key-4..." in result.output
|
||||
assert "team-beta" in result.output
|
||||
|
||||
# Verify existing key was retrieved
|
||||
mock_get_stored.assert_called_once()
|
||||
|
||||
# Verify browser was opened with correct URL including existing key
|
||||
# Verify browser was opened
|
||||
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 "source=litellm-cli" in call_args
|
||||
assert "key=sk-new-session-uuid-789" in call_args
|
||||
assert f"existing_key={existing_key}" in call_args
|
||||
|
||||
# Verify polling was done with correct session key
|
||||
mock_get.assert_called()
|
||||
# Check that the polling URL was called (should be the first call)
|
||||
first_call_args = mock_get.call_args_list[0]
|
||||
poll_url = first_call_args[0][0]
|
||||
assert "sk-new-session-uuid-789" in poll_url
|
||||
# Verify two polling requests were made
|
||||
assert mock_get.call_count == 2
|
||||
|
||||
# Verify regenerated key was saved
|
||||
# First poll should be without team_id
|
||||
first_poll_url = mock_get.call_args_list[0][0][0]
|
||||
assert "sk-session-uuid-456" in first_poll_url
|
||||
assert "team_id=" not in first_poll_url
|
||||
|
||||
# Second poll should include team_id=team-beta
|
||||
second_poll_url = mock_get.call_args_list[1][0][0]
|
||||
assert "team_id=team-beta" in second_poll_url
|
||||
|
||||
# Verify JWT was saved
|
||||
mock_save.assert_called_once()
|
||||
saved_data = mock_save.call_args[0][0]
|
||||
assert saved_data['key'] == 'sk-regenerated-key-456'
|
||||
assert saved_data['user_id'] == 'cli-user'
|
||||
assert saved_data['key'] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.team-beta.jwt"
|
||||
assert saved_data['user_id'] == 'test-user-456'
|
||||
|
||||
mock_show_commands.assert_called_once()
|
||||
|
||||
def test_login_without_existing_key_creation_flow(self):
|
||||
"""Test complete login flow when user has no existing key - should create new one"""
|
||||
def test_login_without_teams_flow(self):
|
||||
"""Test complete login flow when user has no teams - JWT generated without team"""
|
||||
mock_context = Mock()
|
||||
mock_context.obj = {"base_url": "https://test.example.com"}
|
||||
|
||||
# Mock no existing key
|
||||
# Mock response with no teams
|
||||
mock_response = Mock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"status": "ready",
|
||||
"key": "sk-new-created-key-789"
|
||||
"key": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.no-team.jwt",
|
||||
"user_id": "test-user-solo",
|
||||
"team_id": None,
|
||||
"teams": []
|
||||
}
|
||||
|
||||
with patch('webbrowser.open') as mock_browser, \
|
||||
patch('requests.get', return_value=mock_response), \
|
||||
patch('litellm.proxy.client.cli.commands.auth.get_stored_api_key', return_value=None) as mock_get_stored, \
|
||||
patch('litellm.proxy.client.cli.commands.auth.save_token') as mock_save, \
|
||||
patch('litellm.proxy.client.cli.interface.show_commands'), \
|
||||
patch('litellm._uuid.uuid.uuid4', return_value='new-session-uuid-999'):
|
||||
patch('litellm._uuid.uuid.uuid4', return_value='session-uuid-solo'):
|
||||
|
||||
result = self.runner.invoke(login, obj=mock_context.obj)
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert "✅ Login successful!" in result.output
|
||||
|
||||
# Verify existing key check was done
|
||||
mock_get_stored.assert_called_once()
|
||||
|
||||
# Verify browser was opened with correct URL WITHOUT existing key
|
||||
# Verify browser was opened
|
||||
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 "source=litellm-cli" in call_args
|
||||
assert "key=sk-new-session-uuid-999" in call_args
|
||||
assert "existing_key=" not in call_args # Should not include existing_key param
|
||||
assert "key=sk-session-uuid-solo" in call_args
|
||||
|
||||
# Verify new key was saved
|
||||
# Verify JWT was saved
|
||||
mock_save.assert_called_once()
|
||||
saved_data = mock_save.call_args[0][0]
|
||||
assert saved_data['key'] == 'sk-new-created-key-789'
|
||||
assert saved_data['key'] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.no-team.jwt"
|
||||
assert saved_data['user_id'] == 'test-user-solo'
|
||||
|
||||
@@ -249,6 +249,10 @@ class TestCallbackManagementEndpoints:
|
||||
# Verify dynamic_params structure
|
||||
assert isinstance(first_config["dynamic_params"], dict)
|
||||
|
||||
|
||||
# Check if at least one callback has detailed parameter configuration
|
||||
has_detailed_params = any(
|
||||
config.get("dynamic_params") and len(config.get("dynamic_params", {})) > 0
|
||||
for config in response_data
|
||||
)
|
||||
assert has_detailed_params, "Expected at least one callback to have detailed parameter configuration"
|
||||
|
||||
|
||||
@@ -1392,22 +1392,41 @@ class TestCLIKeyRegenerationFlow:
|
||||
"""Test the end-to-end CLI key regeneration flow"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_sso_callback_regenerate_existing_key(self):
|
||||
"""Test CLI SSO callback regenerating an existing key"""
|
||||
async def test_cli_sso_callback_stores_session(self):
|
||||
"""Test CLI SSO callback stores session data in cache for JWT generation"""
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
|
||||
|
||||
# Mock request
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
||||
# Test data
|
||||
existing_key = "sk-existing-key-123"
|
||||
new_key = "sk-new-key-456"
|
||||
session_key = "sk-session-456"
|
||||
|
||||
# Mock user info
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
user_id="test-user-123",
|
||||
user_role="internal_user",
|
||||
teams=["team1", "team2"],
|
||||
models=["gpt-4"]
|
||||
)
|
||||
|
||||
# Mock the regenerate helper function
|
||||
# Mock SSO result
|
||||
mock_sso_result = {
|
||||
"user_email": "test@example.com",
|
||||
"user_id": "test-user-123"
|
||||
}
|
||||
|
||||
# Mock cache
|
||||
mock_cache = MagicMock()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso._regenerate_cli_key"
|
||||
) as mock_regenerate, patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
|
||||
return_value=mock_user_info
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", MagicMock()
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
|
||||
), patch(
|
||||
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
|
||||
return_value="<html>Success</html>",
|
||||
@@ -1415,46 +1434,65 @@ class TestCLIKeyRegenerationFlow:
|
||||
|
||||
# Act
|
||||
result = await cli_sso_callback(
|
||||
request=mock_request, key=new_key, existing_key=existing_key
|
||||
request=mock_request, key=session_key, existing_key=None, result=mock_sso_result
|
||||
)
|
||||
|
||||
# Assert
|
||||
mock_regenerate.assert_called_once_with(
|
||||
existing_key=existing_key, new_key=new_key, user_id=None
|
||||
)
|
||||
# Assert - verify session was stored in cache
|
||||
mock_cache.set_cache.assert_called_once()
|
||||
call_args = mock_cache.set_cache.call_args
|
||||
|
||||
# Verify cache key format
|
||||
assert "cli_sso_session:" in call_args.kwargs["key"]
|
||||
assert session_key in call_args.kwargs["key"]
|
||||
|
||||
# Verify session data structure
|
||||
session_data = call_args.kwargs["value"]
|
||||
assert session_data["user_id"] == "test-user-123"
|
||||
assert session_data["user_role"] == "internal_user"
|
||||
assert session_data["teams"] == ["team1", "team2"]
|
||||
assert session_data["models"] == ["gpt-4"]
|
||||
|
||||
# Verify TTL
|
||||
assert call_args.kwargs["ttl"] == 600 # 10 minutes
|
||||
|
||||
assert result.status_code == 200
|
||||
assert "Success" in result.body.decode()
|
||||
# Verify response contains success message (response is HTML)
|
||||
assert result.body is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_sso_callback_create_new_key(self):
|
||||
"""Test CLI SSO callback creating a new key when no existing key provided"""
|
||||
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
|
||||
|
||||
# Mock request
|
||||
mock_request = MagicMock(spec=Request)
|
||||
async def test_cli_poll_key_returns_teams_for_selection(self):
|
||||
"""Test CLI poll endpoint returns teams for user selection when multiple teams exist"""
|
||||
from litellm.proxy.management_endpoints.ui_sso import cli_poll_key
|
||||
|
||||
# Test data
|
||||
new_key = "sk-new-key-789"
|
||||
session_key = "sk-session-789"
|
||||
session_data = {
|
||||
"user_id": "test-user-456",
|
||||
"user_role": "internal_user",
|
||||
"teams": ["team-a", "team-b", "team-c"],
|
||||
"models": ["gpt-4"]
|
||||
}
|
||||
|
||||
# Mock the create helper function
|
||||
# Mock cache
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.get_cache.return_value = session_data
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.ui_sso._create_new_cli_key"
|
||||
) as mock_create, patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", MagicMock()
|
||||
), patch(
|
||||
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
|
||||
return_value="<html>Success</html>",
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
|
||||
):
|
||||
|
||||
# Act
|
||||
result = await cli_sso_callback(
|
||||
request=mock_request, key=new_key, existing_key=None
|
||||
)
|
||||
# Act - First poll without team_id
|
||||
result = await cli_poll_key(key_id=session_key, team_id=None)
|
||||
|
||||
# Assert
|
||||
mock_create.assert_called_once_with(key=new_key, user_id=None)
|
||||
assert result.status_code == 200
|
||||
assert "Success" in result.body.decode()
|
||||
# Assert - should return teams list for selection
|
||||
assert result["status"] == "ready"
|
||||
assert result["requires_team_selection"] is True
|
||||
assert result["user_id"] == "test-user-456"
|
||||
assert result["teams"] == ["team-a", "team-b", "team-c"]
|
||||
assert "key" not in result # JWT should not be generated yet
|
||||
|
||||
# Verify session was NOT deleted
|
||||
mock_cache.delete_cache.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_callback_routes_to_cli_with_existing_key(self):
|
||||
@@ -1543,40 +1581,65 @@ class TestCLIKeyRegenerationFlow:
|
||||
assert "https://test.litellm.ai/sso/callback" == redirect_url
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_sso_callback_regenerate_vs_create_flow(self):
|
||||
"""Test CLI SSO callback calls regenerate_key_fn when existing_key provided, generate_key_helper_fn when not"""
|
||||
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
|
||||
async def test_cli_poll_key_generates_jwt_with_team(self):
|
||||
"""Test CLI poll endpoint generates JWT when team_id is provided"""
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.management_endpoints.ui_sso import cli_poll_key
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
# Test data
|
||||
session_key = "sk-session-999"
|
||||
selected_team = "team-b"
|
||||
session_data = {
|
||||
"user_id": "test-user-789",
|
||||
"user_role": "internal_user",
|
||||
"teams": ["team-a", "team-b", "team-c"],
|
||||
"models": ["gpt-4"],
|
||||
"user_email": "test@example.com"
|
||||
}
|
||||
|
||||
# Mock user info
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
user_id="test-user-789",
|
||||
user_role="internal_user",
|
||||
teams=["team-a", "team-b", "team-c"],
|
||||
models=["gpt-4"]
|
||||
)
|
||||
|
||||
# Mock cache
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.get_cache.return_value = session_data
|
||||
|
||||
mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.token"
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.regenerate_key_fn"
|
||||
) as mock_regenerate, patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn"
|
||||
) as mock_generate, patch(
|
||||
"litellm.proxy._types.UserAPIKeyAuth.get_litellm_cli_user_api_key_auth"
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", mock_cache
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", MagicMock()
|
||||
), patch(
|
||||
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
|
||||
return_value="<html>Success</html>",
|
||||
):
|
||||
"litellm.proxy.proxy_server.prisma_client"
|
||||
) as mock_prisma, patch(
|
||||
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
|
||||
return_value=mock_jwt_token
|
||||
) as mock_get_jwt:
|
||||
|
||||
# Mock the user lookup
|
||||
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user_info)
|
||||
|
||||
# Test regeneration path
|
||||
await cli_sso_callback(
|
||||
mock_request, key="sk-new-123", existing_key="sk-existing-456"
|
||||
)
|
||||
mock_regenerate.assert_called_once()
|
||||
mock_generate.assert_not_called()
|
||||
# Act - Second poll with team_id
|
||||
result = await cli_poll_key(key_id=session_key, team_id=selected_team)
|
||||
|
||||
# Reset mocks
|
||||
mock_regenerate.reset_mock()
|
||||
mock_generate.reset_mock()
|
||||
|
||||
# Test creation path
|
||||
await cli_sso_callback(mock_request, key="sk-new-789", existing_key=None)
|
||||
mock_regenerate.assert_not_called()
|
||||
mock_generate.assert_called_once()
|
||||
# Assert - should return JWT
|
||||
assert result["status"] == "ready"
|
||||
assert result["key"] == mock_jwt_token
|
||||
assert result["user_id"] == "test-user-789"
|
||||
assert result["team_id"] == selected_team
|
||||
assert result["teams"] == ["team-a", "team-b", "team-c"]
|
||||
|
||||
# Verify JWT was generated with correct team
|
||||
mock_get_jwt.assert_called_once()
|
||||
jwt_call_args = mock_get_jwt.call_args
|
||||
assert jwt_call_args.kwargs["team_id"] == selected_team
|
||||
|
||||
# Verify session was deleted after JWT generation
|
||||
mock_cache.delete_cache.assert_called_once()
|
||||
|
||||
|
||||
class TestGetAppRolesFromIdToken:
|
||||
|
||||
Reference in New Issue
Block a user