diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 2345b0263e3..64b3233536f 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -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]: diff --git a/tests/test_litellm/proxy/auth/test_cli_auth.py b/tests/test_litellm/proxy/auth/test_cli_auth.py new file mode 100644 index 00000000000..2faf6436523 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_cli_auth.py @@ -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]