mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Fix team cli auth flow (#19666)
* Cleanup code for user cli auth, and make sure not to prompt user for team multiple times while polling * Adding tests * Cleanup normalize teams some more
This commit is contained in:
parent
3ab1b9f543
commit
8e4f06583a
2 changed files with 326 additions and 110 deletions
|
|
@ -276,6 +276,71 @@ def prompt_team_selection_fallback(
|
|||
|
||||
|
||||
# Polling-based authentication - no local server needed
|
||||
def _poll_for_ready_data(
|
||||
url: str,
|
||||
*,
|
||||
total_timeout: int = 300,
|
||||
poll_interval: int = 2,
|
||||
request_timeout: int = 10,
|
||||
pending_message: Optional[str] = None,
|
||||
pending_log_every: int = 10,
|
||||
other_status_message: Optional[str] = None,
|
||||
other_status_log_every: int = 10,
|
||||
http_error_log_every: int = 10,
|
||||
connection_error_log_every: int = 10,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
for attempt in range(total_timeout // poll_interval):
|
||||
try:
|
||||
response = requests.get(url, timeout=request_timeout)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
status = data.get("status")
|
||||
if status == "ready":
|
||||
return data
|
||||
if status == "pending":
|
||||
if (
|
||||
pending_message
|
||||
and pending_log_every > 0
|
||||
and attempt % pending_log_every == 0
|
||||
):
|
||||
click.echo(pending_message)
|
||||
elif (
|
||||
other_status_message
|
||||
and other_status_log_every > 0
|
||||
and attempt % other_status_log_every == 0
|
||||
):
|
||||
click.echo(other_status_message)
|
||||
elif http_error_log_every > 0 and attempt % http_error_log_every == 0:
|
||||
click.echo(f"Polling error: HTTP {response.status_code}")
|
||||
except requests.RequestException as e:
|
||||
if (
|
||||
connection_error_log_every > 0
|
||||
and attempt % connection_error_log_every == 0
|
||||
):
|
||||
click.echo(f"Connection error (will retry): {e}")
|
||||
time.sleep(poll_interval)
|
||||
return None
|
||||
|
||||
|
||||
def _normalize_teams(teams, team_details):
|
||||
"""If team_details are a
|
||||
|
||||
Args:
|
||||
teams (_type_): _description_
|
||||
team_details (_type_): _description_
|
||||
|
||||
Returns:
|
||||
_type_: _description_
|
||||
"""
|
||||
if isinstance(team_details, list) and team_details:
|
||||
return [
|
||||
{"team_id": i.get("team_id") or i.get("id"), "team_alias": i.get("team_alias")}
|
||||
for i in team_details
|
||||
if isinstance(i, dict) and (i.get("team_id") or i.get("id"))
|
||||
]
|
||||
if isinstance(teams, list):
|
||||
return [{"team_id": str(t), "team_alias": None} for t in teams]
|
||||
return []
|
||||
|
||||
|
||||
def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
|
||||
|
|
@ -286,106 +351,58 @@ def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
|
|||
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
|
||||
data = _poll_for_ready_data(
|
||||
poll_url,
|
||||
pending_message="Still waiting for authentication...",
|
||||
)
|
||||
if not data:
|
||||
return None
|
||||
if data.get("requires_team_selection"):
|
||||
teams = data.get("teams", [])
|
||||
team_details = data.get("team_details")
|
||||
user_id = data.get("user_id")
|
||||
normalized_teams: List[Dict[str, Any]] = _normalize_teams(teams, team_details)
|
||||
if not normalized_teams:
|
||||
click.echo("⚠️ No teams available for selection.")
|
||||
return None
|
||||
|
||||
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.
|
||||
# Newer servers may also return "team_details" containing
|
||||
# objects with both team_id and team_alias. We prefer those
|
||||
# for display, but continue to support the legacy list of
|
||||
# team IDs for backwards compatibility.
|
||||
teams = data.get("teams", [])
|
||||
team_details = data.get("team_details")
|
||||
user_id = data.get("user_id")
|
||||
# User has multiple teams - let them select
|
||||
jwt_with_team = _handle_team_selection_during_polling(
|
||||
base_url=base_url,
|
||||
key_id=key_id,
|
||||
teams=normalized_teams,
|
||||
)
|
||||
|
||||
# Build a normalized list of team objects that always have
|
||||
# "team_id" and optionally "team_alias".
|
||||
normalized_teams: List[Dict[str, Any]] = []
|
||||
if isinstance(team_details, list) and team_details:
|
||||
for item in team_details:
|
||||
if isinstance(item, dict):
|
||||
team_id = item.get("team_id") or item.get("id")
|
||||
if team_id is None:
|
||||
continue
|
||||
normalized_teams.append(
|
||||
{
|
||||
"team_id": team_id,
|
||||
"team_alias": item.get("team_alias"),
|
||||
}
|
||||
)
|
||||
elif isinstance(teams, list):
|
||||
for t in teams:
|
||||
normalized_teams.append(
|
||||
{
|
||||
"team_id": str(t),
|
||||
"team_alias": None,
|
||||
}
|
||||
)
|
||||
# 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
|
||||
}
|
||||
|
||||
if normalized_teams and len(normalized_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=normalized_teams,
|
||||
)
|
||||
click.echo("❌ Team selection cancelled or JWT generation failed.")
|
||||
return None
|
||||
|
||||
# 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")
|
||||
# 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}")
|
||||
# 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(f"Polling error: HTTP {response.status_code}")
|
||||
if api_key:
|
||||
return {
|
||||
"api_key": api_key,
|
||||
"user_id": user_id,
|
||||
"teams": teams,
|
||||
"team_id": team_id,
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
|
||||
|
|
@ -418,25 +435,21 @@ def _handle_team_selection_during_polling(
|
|||
|
||||
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}")
|
||||
poll_url = f"{base_url}/sso/cli/poll/{key_id}?team_id={team_id}"
|
||||
data = _poll_for_ready_data(
|
||||
poll_url,
|
||||
pending_message="Still waiting for team authentication...",
|
||||
other_status_message="Waiting for team authentication to complete...",
|
||||
http_error_log_every=10,
|
||||
)
|
||||
if not data:
|
||||
return None
|
||||
jwt_token = data.get("key")
|
||||
if jwt_token:
|
||||
click.echo(f"✅ Successfully generated JWT for team: {team_id}")
|
||||
return jwt_token
|
||||
|
||||
except Exception as e:
|
||||
click.echo(f"❌ Error getting JWT with team: {e}")
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _render_and_prompt_for_team_selection(teams: List[Dict[str, Any]]) -> Optional[str]:
|
||||
|
|
|
|||
203
tests/test_litellm/proxy/auth/test_cli_auth.py
Normal file
203
tests/test_litellm/proxy/auth/test_cli_auth.py
Normal file
|
|
@ -0,0 +1,203 @@
|
|||
"""
|
||||
Tests for litellm/proxy/client/cli/commands/auth.py
|
||||
|
||||
This module tests the auth commands and their associated functionality.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
from unittest.mock import AsyncMock, patch, Mock, call
|
||||
from litellm.proxy.client.cli.commands.auth import _normalize_teams, _poll_for_ready_data, _poll_for_authentication
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normalize_teams_teams_only():
|
||||
"""Test normalize teams helper function"""
|
||||
teams = ["1", "2", "3"]
|
||||
team_details = []
|
||||
result = _normalize_teams(teams, team_details)
|
||||
assert result == [{"team_id": "1", "team_alias": None}, {"team_id": "2", "team_alias": None}, {"team_id": "3", "team_alias": None}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normalize_teams_with_details_no_aliases():
|
||||
"""Test normalize teams helper function"""
|
||||
teams = ["4", "5", "6"]
|
||||
team_details = [{"team_id": "1"}, {"team_id": "2"}, {"team_id": "3"}]
|
||||
result = _normalize_teams(teams, team_details)
|
||||
assert result == [{"team_id": "1", "team_alias": None}, {"team_id": "2", "team_alias": None}, {"team_id": "3", "team_alias": None}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normalize_teams_with_details_with_aliases():
|
||||
"""Test normalize teams helper function"""
|
||||
teams = ["4", "5", "6"]
|
||||
team_details = [{"team_id": "1", "team_alias": "A"}, {"team_id": "2", "team_alias": "B"}, {"team_id": "3", "team_alias": "C"}]
|
||||
result = _normalize_teams(teams, team_details)
|
||||
assert result == [{"team_id": "1", "team_alias": "A"}, {"team_id": "2", "team_alias": "B"}, {"team_id": "3", "team_alias": "C"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth.requests.get", side_effect=[Mock(status_code=404)])
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
@patch("litellm.proxy.client.cli.commands.auth.time.sleep")
|
||||
async def test_poll_for_ready_404(sleep_mock, click_mock, request_mock):
|
||||
"""Test poll_for_ready function"""
|
||||
actual = _poll_for_ready_data("https://litellm.com", poll_interval=1, total_timeout=1, request_timeout=42)
|
||||
assert actual is None
|
||||
click_mock.assert_called_once_with("Polling error: HTTP 404")
|
||||
request_mock.assert_called_once_with("https://litellm.com", timeout=42)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth.requests.get", side_effect=[Mock(status_code=200, json=Mock(return_value={"status": "ready","json": "data"}))])
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
@patch("litellm.proxy.client.cli.commands.auth.time.sleep")
|
||||
async def test_poll_for_ready_200_ready(sleep_mock, click_mock, request_mock):
|
||||
"""Test poll_for_ready function"""
|
||||
actual = _poll_for_ready_data("https://litellm.com", poll_interval=1, total_timeout=1, request_timeout=42)
|
||||
assert actual == {"status": "ready", "json": "data"}
|
||||
click_mock.assert_not_called()
|
||||
request_mock.assert_called_once_with("https://litellm.com", timeout=42)
|
||||
sleep_mock.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth.requests.get", side_effect=[Mock(status_code=200, json=Mock(return_value={"status": "pending","json": "data"})), Mock(status_code=200, json=Mock(return_value={"status": "ready","json": "data"}))])
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
@patch("litellm.proxy.client.cli.commands.auth.time.sleep")
|
||||
async def test_poll_for_ready_single_pending(sleep_mock, click_mock, request_mock):
|
||||
"""Test poll_for_ready function"""
|
||||
actual = _poll_for_ready_data("https://litellm.com", poll_interval=1, total_timeout=2, request_timeout=42)
|
||||
assert actual == {"status": "ready", "json": "data"}
|
||||
click_mock.assert_not_called()
|
||||
request_mock.assert_has_calls([
|
||||
call("https://litellm.com", timeout=42),
|
||||
call("https://litellm.com", timeout=42)
|
||||
])
|
||||
sleep_mock.assert_called_once_with(1)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth.requests.get", side_effect=[Mock(status_code=200, json=Mock(return_value={"status": "pending","json": "data"})), Mock(status_code=200, json=Mock(return_value={"status": "pending","json": "data"}))])
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
@patch("litellm.proxy.client.cli.commands.auth.time.sleep")
|
||||
async def test_poll_for_ready_pending(sleep_mock, click_mock, request_mock):
|
||||
"""Test poll_for_ready function"""
|
||||
actual = _poll_for_ready_data("https://litellm.com", poll_interval=1, total_timeout=2, request_timeout=42, pending_message="Pending message", pending_log_every=1)
|
||||
assert actual is None
|
||||
click_mock.assert_has_calls([
|
||||
call("Pending message"),
|
||||
call("Pending message")
|
||||
])
|
||||
request_mock.assert_has_calls([
|
||||
call("https://litellm.com", timeout=42),
|
||||
call("https://litellm.com", timeout=42)
|
||||
])
|
||||
sleep_mock.assert_has_calls([
|
||||
call(1),
|
||||
call(1)
|
||||
])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth.requests.get", side_effect=[requests.RequestException("ERROR"),
|
||||
requests.RequestException("ERROR")])
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
@patch("litellm.proxy.client.cli.commands.auth.time.sleep")
|
||||
async def test_poll_for_ready_connection_failure(sleep_mock, click_mock, request_mock):
|
||||
"""Test poll_for_ready function"""
|
||||
actual = _poll_for_ready_data("https://litellm.com", poll_interval=1, total_timeout=2, request_timeout=42)
|
||||
assert actual is None
|
||||
click_mock.assert_called_once_with("Connection error (will retry): ERROR")
|
||||
request_mock.assert_has_calls([
|
||||
call("https://litellm.com", timeout=42),
|
||||
])
|
||||
sleep_mock.assert_has_calls([
|
||||
call(1),
|
||||
call(1)
|
||||
])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth._handle_team_selection_during_polling")
|
||||
@patch("litellm.proxy.client.cli.commands.auth._poll_for_ready_data", return_value=None)
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
async def test_poll_for_authentication_no_data(click_mock, poll_mock, handle_mock):
|
||||
"""Test poll_for_authentication function"""
|
||||
actual = _poll_for_authentication("https://litellm.com", "key-123")
|
||||
assert actual is None
|
||||
poll_mock.assert_called_once_with(
|
||||
"https://litellm.com/sso/cli/poll/key-123",
|
||||
pending_message="Still waiting for authentication...",
|
||||
)
|
||||
handle_mock.assert_not_called()
|
||||
click_mock.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth._handle_team_selection_during_polling")
|
||||
@patch("litellm.proxy.client.cli.commands.auth._poll_for_ready_data", return_value={"requires_team_selection": True, "teams": [], "team_details": []})
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
async def test_poll_for_authentication_no_teams(click_mock, poll_mock, handle_mock):
|
||||
"""Test poll_for_authentication function"""
|
||||
actual = _poll_for_authentication("https://litellm.com", "key-123")
|
||||
assert actual is None
|
||||
poll_mock.assert_called_once_with(
|
||||
"https://litellm.com/sso/cli/poll/key-123",
|
||||
pending_message="Still waiting for authentication...",
|
||||
)
|
||||
handle_mock.assert_not_called()
|
||||
click_mock.assert_called_once()
|
||||
assert "No teams available for selection." in click_mock.call_args[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth._handle_team_selection_during_polling", return_value="jwt-123")
|
||||
@patch("litellm.proxy.client.cli.commands.auth._poll_for_ready_data", return_value={"requires_team_selection": True, "teams": [1, 2], "user_id": "user-123"})
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
async def test_poll_for_authentication_team_selection_success(click_mock, poll_mock, handle_mock):
|
||||
"""Test poll_for_authentication function"""
|
||||
actual = _poll_for_authentication("https://litellm.com", "key-123")
|
||||
assert actual == {"api_key": "jwt-123", "user_id": "user-123", "teams": [1, 2], "team_id": None}
|
||||
poll_mock.assert_called_once_with(
|
||||
"https://litellm.com/sso/cli/poll/key-123",
|
||||
pending_message="Still waiting for authentication...",
|
||||
)
|
||||
handle_mock.assert_called_once_with(
|
||||
base_url="https://litellm.com",
|
||||
key_id="key-123",
|
||||
teams=[{"team_id": "1", "team_alias": None}, {"team_id": "2", "team_alias": None}],
|
||||
)
|
||||
click_mock.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth._handle_team_selection_during_polling", return_value=None)
|
||||
@patch("litellm.proxy.client.cli.commands.auth._poll_for_ready_data", return_value={"requires_team_selection": True, "teams": ["team-1"], "user_id": "user-123"})
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
async def test_poll_for_authentication_team_selection_cancelled(click_mock, poll_mock, handle_mock):
|
||||
"""Test poll_for_authentication function"""
|
||||
actual = _poll_for_authentication("https://litellm.com", "key-123")
|
||||
assert actual is None
|
||||
poll_mock.assert_called_once_with(
|
||||
"https://litellm.com/sso/cli/poll/key-123",
|
||||
pending_message="Still waiting for authentication...",
|
||||
)
|
||||
handle_mock.assert_called_once_with(
|
||||
base_url="https://litellm.com",
|
||||
key_id="key-123",
|
||||
teams=[{"team_id": "team-1", "team_alias": None}],
|
||||
)
|
||||
click_mock.assert_called_once()
|
||||
assert "Team selection cancelled" in click_mock.call_args[0][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.client.cli.commands.auth._handle_team_selection_during_polling")
|
||||
@patch("litellm.proxy.client.cli.commands.auth._poll_for_ready_data", return_value={"key": "jwt-456", "user_id": "user-456", "teams": ["team-1"], "team_id": "team-1"})
|
||||
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
|
||||
async def test_poll_for_authentication_auto_assigned_team(click_mock, poll_mock, handle_mock):
|
||||
"""Test poll_for_authentication function"""
|
||||
actual = _poll_for_authentication("https://litellm.com", "key-123")
|
||||
assert actual == {"api_key": "jwt-456", "user_id": "user-456", "teams": ["team-1"], "team_id": "team-1"}
|
||||
poll_mock.assert_called_once_with(
|
||||
"https://litellm.com/sso/cli/poll/key-123",
|
||||
pending_message="Still waiting for authentication...",
|
||||
)
|
||||
handle_mock.assert_not_called()
|
||||
click_mock.assert_called_once()
|
||||
assert "Automatically assigned to team: team-1" in click_mock.call_args[0][0]
|
||||
Loading…
Add table
Reference in a new issue