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:
boarder7395 2026-01-28 11:52:52 -05:00 • committed by GitHub
parent 3ab1b9f543
commit 8e4f06583a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 326 additions and 110 deletions

View file

@ -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]:

View 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]