From eef5541e60f993c80ffbe1dfad6e96761eda2e06 Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 29 Sep 2026 16:57:33 +0000 Subject: [PATCH] fix(cli): renew saved sessions when assigning teams Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/client/cli/commands/_cli_context.py | 3 +- litellm/proxy/client/cli/commands/auth.py | 57 +++-- litellm/proxy/client/cli/commands/teams.py | 23 ++ .../proxy/client/cli/test_auth_commands.py | 222 ++++++++++++++++++ 4 files changed, 290 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/client/cli/commands/_cli_context.py b/litellm/proxy/client/cli/commands/_cli_context.py index 74c29653d16..5ccbb30392c 100644 --- a/litellm/proxy/client/cli/commands/_cli_context.py +++ b/litellm/proxy/client/cli/commands/_cli_context.py @@ -1,7 +1,7 @@ from typing import Final import click -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict class CliContextValues(TypedDict): @@ -9,6 +9,7 @@ class CliContextValues(TypedDict): base_url: ReadOnly[str] api_key: ReadOnly[str | None] + api_key_from_token_file: ReadOnly[NotRequired[bool]] _UNSET_CLI_CONTEXT: Final[CliContextValues] = {"base_url": "", "api_key": None} diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 5c2ce9eb53c..2520c09a535 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -634,7 +634,7 @@ def match_requested_team(teams: Sequence[CliTeam], requested_team: str | None) - def _poll_for_authentication( - base_url: str, key_id: str, poll_secret: str, team: str | None = None + base_url: str, key_id: str, poll_secret: str, team: str | None = None, required_team_id: str | None = None ) -> CliAuthResult | None: """ Poll the server for authentication completion and handle team selection. @@ -655,6 +655,11 @@ def _poll_for_authentication( team_details: Final = data.get("team_details") user_id = data.get("user_id") normalized_teams: Final[list[CliTeam]] = _normalize_teams(teams, team_details) + if ( + required_team_id is not None + and match_requested_team(normalized_teams, required_team_id) != required_team_id + ): + raise click.ClickException("The requested team is not available for this login") if not normalized_teams: click.echo("Warning: No teams available for selection.") return None @@ -674,7 +679,7 @@ def _poll_for_authentication( "api_key": jwt_with_team, "user_id": user_id, "teams": teams, - "team_id": None, # Set by server in JWT + "team_id": match_requested_team(normalized_teams, team), } click.echo("Team selection cancelled or JWT generation failed.") @@ -852,7 +857,17 @@ def _finish_login(base_url: str, api_key: str, config_claude: bool, stored: Secr show_commands() -def _replace_stored_token(record: CliTokenData, http: Http, vault: SecretVault) -> SecretSave: +def _replace_stored_token( + record: CliTokenData, http: Http, vault: SecretVault, required_team_id: str | None = None +) -> SecretSave: + if required_team_id is not None and record.get("team_id") != required_team_id: + refused_revocation: Final = revoke_stored_credential(record, http) + if refused_revocation is not None: + click.echo( + f"Could not revoke the rejected login's refresh token on the proxy ({refused_revocation.reason}); " + "it expires on its own." + ) + raise click.ClickException("The login did not select the requested team; your saved login has not changed") previous: Final = load_token(vault=vault) stored: Final = save_token(record, vault=vault) if previous is None or isinstance(stored, CredentialNotSaved): @@ -866,14 +881,17 @@ def _replace_stored_token(record: CliTokenData, http: Http, vault: SecretVault) return stored -def _pkce_login(base_url: str, config_claude: bool, vault: SecretVault, team: str | None) -> None: +def _pkce_login( + base_url: str, config_claude: bool, vault: SecretVault, team: str | None, required_team_id: str | None = None +) -> bool: http: Final = requests.Session() credential: Final = run_pkce_login(base_url, http, echo=click.echo, team=team) if isinstance(credential, PkceFailure): click.echo(f"Authentication failed: {credential.reason}") - return - stored: Final = _replace_stored_token(pkce_token_record(base_url, credential), http, vault) + return False + stored: Final = _replace_stored_token(pkce_token_record(base_url, credential), http, vault, required_team_id) _finish_login(base_url, credential.access_token, config_claude, stored) + return not isinstance(stored, (CredentialNotSaved, CredentialNotRecorded)) @click.command(name="login") @@ -911,6 +929,12 @@ def _pkce_login(base_url: str, config_claude: bool, vault: SecretVault, team: st @click.pass_context def login(ctx: click.Context, config_claude: bool, pkce: bool, team: str | None) -> None: """Login to LiteLLM proxy using SSO authentication""" + login_to_proxy(ctx, config_claude, pkce, team) + + +def login_to_proxy( + ctx: click.Context, config_claude: bool, pkce: bool, team: str | None, required_team_id: str | None = None +) -> bool: from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER ctx_obj: Final[CliContextObj] = ctx.obj @@ -924,8 +948,7 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool, team: str | None) try: if pkce: - _pkce_login(base_url, config_claude, context_secret_vault(ctx), team) - return + return _pkce_login(base_url, config_claude, context_secret_vault(ctx), team, required_team_id) cli_sso_flow: Final = _start_cli_sso_flow(base_url=base_url) key_id: Final = cli_sso_flow["login_id"] poll_secret: Final = cli_sso_flow["poll_secret"] @@ -955,8 +978,12 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool, team: str | None) # Poll for authentication completion click.echo("Waiting for authentication...") - auth_result: Final = _poll_for_authentication( - base_url=base_url, key_id=key_id, poll_secret=poll_secret, team=team + auth_result: Final = ( + _poll_for_authentication( + base_url=base_url, key_id=key_id, poll_secret=poll_secret, team=team, required_team_id=required_team_id + ) + if required_team_id is not None + else _poll_for_authentication(base_url=base_url, key_id=key_id, poll_secret=poll_secret, team=team) ) if auth_result: @@ -975,31 +1002,33 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool, team: str | None) "auth_header_name": "Authorization", "jwt_token": "", "timestamp": time.time(), + "team_id": auth_result["team_id"], }, requests.Session(), context_secret_vault(ctx), + required_team_id, ) _finish_login(base_url, api_key, config_claude, stored) - return + return not isinstance(stored, (CredentialNotSaved, CredentialNotRecorded)) else: click.echo("Authentication timed out. Please try again.") click.echo( "The proxy never reported the browser sign-in as finished. If you did complete it, " "check the proxy logs for /sso/callback errors and confirm SSO is configured on the proxy." ) - return + return False except KeyboardInterrupt: click.echo("\nAuthentication cancelled by user.") - return + return False except click.ClickException: # Login itself already succeeded; only the post-login step failed, so this # must not be relabelled as an authentication failure by the handler below. raise except Exception as e: click.echo(f"Authentication failed: {e}") - return + return False @click.command(name="logout") diff --git a/litellm/proxy/client/cli/commands/teams.py b/litellm/proxy/client/cli/commands/teams.py index 1a941786f19..fe47254102e 100644 --- a/litellm/proxy/client/cli/commands/teams.py +++ b/litellm/proxy/client/cli/commands/teams.py @@ -12,6 +12,7 @@ from typing_extensions import ReadOnly, TypedDict from litellm.proxy.client import Client from ._cli_context import cli_context_values +from .auth import context_secret_vault, load_token, login_to_proxy class _TeamRow(TypedDict): @@ -141,6 +142,28 @@ def assign_key(ctx: click.Context, team_id: str | None): click.echo("No API key found. Please login first using 'litellm login'") raise click.Abort() + stored_token: Final = load_token(vault=context_secret_vault(ctx)) + if ( + stored_token is not None + and stored_token.get("base_url") == context["base_url"].rstrip("/") + and (context.get("api_key_from_token_file", False) or stored_token.get("key") == api_key) + and (stored_token.get("refresh_token") or not api_key.startswith("sk-")) + ): + if not context.get("api_key_from_token_file", False): + raise click.ClickException("Unset --api-key and LITELLM_PROXY_API_KEY to switch your saved CLI session") + click.echo("Signing in again to select the team for your CLI session") + saved: Final = login_to_proxy( + ctx, + config_claude=False, + pkce=bool(stored_token.get("refresh_token")), + team=team_id, + required_team_id=team_id, + ) + if not saved: + raise click.ClickException("CLI session team assignment did not complete") + click.echo(f"Successfully assigned CLI session to team: {team_id}" if team_id else "CLI session team selected") + return + try: # If no team_id provided, show teams and let user select if not team_id: diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index 543df16e23e..3e922e6c8e7 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -2,6 +2,7 @@ import json import os import stat import time +from dataclasses import replace from pathlib import Path from unittest.mock import Mock, patch @@ -28,11 +29,13 @@ from litellm.proxy.client.cli.commands import claude_settings as claude_settings from litellm.proxy.client.cli.commands.claude_settings import SettingsFileOwner from litellm.proxy.client.cli.commands.auth import ( get_stored_api_key, + load_token, login, logout, print_token, whoami, ) +from litellm.proxy.client.cli.commands.pkce_login import PkceFailure, RevocationUnavailable @pytest.fixture @@ -2129,3 +2132,222 @@ class TestRequestedTeamLoginOption: assert result.exit_code == 0, result.output sso_start.assert_not_called() assert pkce_login.call_args.args[3] == "Beta Team" + + +@pytest.mark.parametrize("pkce,refreshed", [(False, False), (True, False), (True, True)]) +def test_assign_key_replaces_saved_session_and_retains_login_protocol(isolated_home, monkeypatch, pkce, refreshed): + monkeypatch.setenv(DISABLE_KEYRING_ENV_VAR, "1") + record = ( + _pkce_record(key="session-old", expires_at=time.time() + 3600, team_id="team-a") + if pkce + else { + "base_url": PKCE_BASE_URL, + "key": "session-old", + "user_id": "u1", + "user_role": "cli", + "timestamp": time.time(), + } + ) + save_cli_token(CliTokenRecord(**record)) + _FakeSession.instances.clear() + with ( + patch( + "litellm.proxy.client.cli.main.get_stored_api_key", + return_value="session-refreshed" if refreshed else "session-old", + ), + patch("litellm.proxy.client.cli.commands.teams.Client") as client, + patch( + "litellm.proxy.client.cli.commands.auth.run_pkce_login", + return_value=replace(_pkce_credential(), access_token="session-new"), + ) as run_pkce, + patch( + "litellm.proxy.client.cli.commands.auth._start_cli_sso_flow", + return_value={ + "login_id": "login-id", + "poll_secret": "poll-secret", + "user_code": "ABCD-EFGH", + }, + ), + patch( + "litellm.proxy.client.cli.commands.auth._poll_for_authentication", + return_value={ + "api_key": "session-new", + "user_id": "u1", + "teams": ["team-a", "team-b"], + "team_id": "team-b", + }, + ) as poll, + patch("litellm.proxy.client.cli.commands.auth.webbrowser.open"), + patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession), + patch("litellm.proxy.client.cli.interface.show_commands"), + ): + result = CliRunner().invoke(cli, ["--base-url", PKCE_BASE_URL, "teams", "assign-key", "--team-id", "team-b"]) + + assert result.exit_code == 0, result.output + saved = load_token() + assert saved is not None + assert (saved["key"], saved["team_id"], saved["user_id"]) == ("session-new", "team-b", "u1") + assert "Successfully assigned CLI session to team: team-b" in result.output + client.return_value.keys.update.assert_not_called() + if pkce: + assert saved["refresh_token"] == "llm_srefresh_fresh" + assert run_pkce.call_args.kwargs["team"] == "team-b" + poll.assert_not_called() + assert ( + next(session for session in _FakeSession.instances if session.posts).posts[0][1]["token"] + == "llm_srefresh_old" + ) + else: + assert "refresh_token" not in saved + assert poll.call_args.kwargs["team"] == "team-b" + run_pkce.assert_not_called() + + +@pytest.mark.parametrize("outcome", ["denied", "wrong-team", "storage-failure"]) +def test_assign_key_does_not_replace_saved_session_when_login_cannot_be_used(isolated_home, monkeypatch, outcome): + monkeypatch.setenv(DISABLE_KEYRING_ENV_VAR, "1") + save_cli_token(CliTokenRecord(**_pkce_record(key="session-old", expires_at=time.time() + 3600, team_id="team-a"))) + before = load_token() + _FakeSession.instances.clear() + credential = ( + PkceFailure("access_denied") + if outcome == "denied" + else replace( + _pkce_credential(), access_token="session-new", team_id="team-c" if outcome == "wrong-team" else "team-b" + ) + ) + with ( + patch("litellm.proxy.client.cli.commands.teams.Client"), + patch("litellm.proxy.client.cli.commands.auth.run_pkce_login", return_value=credential), + patch( + "litellm.proxy.client.cli.commands.auth.save_token", return_value=CredentialNotSaved("read-only") + ) as save, + patch("litellm.proxy.client.cli.commands.auth.requests.Session", _FakeSession), + patch("litellm.proxy.client.cli.interface.show_commands"), + ): + result = CliRunner().invoke(cli, ["--base-url", PKCE_BASE_URL, "teams", "assign-key", "--team-id", "team-b"]) + + assert result.exit_code != 0, result.output + assert load_token() == before + assert "Successfully assigned" not in result.output + if outcome != "storage-failure": + save.assert_not_called() + posts = [post for session in _FakeSession.instances for post in session.posts] + assert posts == ( + [ + ( + f"{PKCE_BASE_URL}/revoke", + {"token": "llm_srefresh_fresh", "token_type_hint": "refresh_token", "client_id": "llm_dcrc_abc"}, + ) + ] + if outcome == "wrong-team" + else [] + ) + + +def test_assign_key_keeps_explicit_virtual_key_update_with_saved_session(isolated_home, monkeypatch): + monkeypatch.setenv(DISABLE_KEYRING_ENV_VAR, "1") + save_cli_token(CliTokenRecord(**_pkce_record(key="session-old", expires_at=time.time() + 3600))) + before = load_token() + with ( + patch("litellm.proxy.client.cli.commands.teams.Client") as client, + patch("litellm.proxy.client.cli.commands.auth.run_pkce_login") as run_pkce, + ): + client.return_value.teams.list.return_value = [] + result = CliRunner().invoke( + cli, + [ + "--base-url", + PKCE_BASE_URL, + "--api-key", + "sk-virtual-key", + "teams", + "assign-key", + "--team-id", + "team-b", + ], + ) + + assert result.exit_code == 0, result.output + client.return_value.keys.update.assert_called_once_with(key="sk-virtual-key", team_id="team-b") + run_pkce.assert_not_called() + assert load_token() == before + + +def test_assign_key_warns_when_a_rejected_login_cannot_be_revoked(isolated_home, monkeypatch): + monkeypatch.setenv(DISABLE_KEYRING_ENV_VAR, "1") + save_cli_token(CliTokenRecord(**_pkce_record(key="session-old", expires_at=time.time() + 3600))) + before = load_token() + with ( + patch( + "litellm.proxy.client.cli.commands.auth.run_pkce_login", + return_value=replace(_pkce_credential(), team_id="team-c"), + ), + patch( + "litellm.proxy.client.cli.commands.auth.revoke_stored_credential", + return_value=RevocationUnavailable("connection unavailable"), + ) as revoke, + ): + result = CliRunner().invoke(cli, ["--base-url", PKCE_BASE_URL, "teams", "assign-key", "--team-id", "team-b"]) + + assert result.exit_code != 0 + assert load_token() == before + assert "Could not revoke the rejected login's refresh token" in result.output + assert "connection unavailable" in result.output + assert "your saved login has not changed" in result.output + assert revoke.call_args.args[0]["refresh_token"] == "llm_srefresh_fresh" + + +@pytest.mark.parametrize("multiple_teams", [False, True]) +def test_assign_key_refuses_unavailable_team_in_sso_poll(isolated_home, monkeypatch, multiple_teams): + monkeypatch.setenv(DISABLE_KEYRING_ENV_VAR, "1") + save_cli_token(CliTokenRecord(base_url=PKCE_BASE_URL, key="session-old", timestamp=time.time())) + before = load_token() + with ( + patch( + "litellm.proxy.client.cli.commands.auth._start_cli_sso_flow", + return_value={"login_id": "login-id", "poll_secret": "poll-secret", "user_code": "ABCD-EFGH"}, + ), + patch("litellm.proxy.client.cli.commands.auth.webbrowser.open"), + patch("litellm.proxy.client.cli.commands.auth.requests.get") as get, + ): + get.return_value.status_code = 200 + get.return_value.json.return_value = { + "status": "ready", + "requires_team_selection": multiple_teams, + "teams": ["team-c", "team-d"] if multiple_teams else ["team-c"], + "key": "session-wrong-team", + "team_id": "team-c", + "user_id": "u1", + } + result = CliRunner().invoke(cli, ["--base-url", PKCE_BASE_URL, "teams", "assign-key", "--team-id", "team-b"]) + + assert result.exit_code != 0 + assert "requested team" in result.output + assert load_token() == before + assert "Successfully assigned" not in result.output + assert get.call_count == 1 + + +@pytest.mark.parametrize("from_env", [False, True]) +def test_assign_key_refuses_an_explicit_saved_session_override(isolated_home, monkeypatch, from_env): + monkeypatch.setenv(DISABLE_KEYRING_ENV_VAR, "1") + save_cli_token(CliTokenRecord(**_pkce_record(key="session-old", expires_at=time.time() + 3600))) + before = load_token() + with ( + patch("litellm.proxy.client.cli.commands.teams.Client") as client, + patch("litellm.proxy.client.cli.commands.auth.run_pkce_login") as run_pkce, + ): + result = CliRunner().invoke( + cli, + ["--base-url", PKCE_BASE_URL] + + ([] if from_env else ["--api-key", "session-old"]) + + ["teams", "assign-key", "--team-id", "team-b"], + env={"LITELLM_PROXY_API_KEY": "session-old"} if from_env else {}, + ) + + assert result.exit_code != 0 + assert "Unset --api-key and LITELLM_PROXY_API_KEY" in result.output + assert load_token() == before + run_pkce.assert_not_called() + client.return_value.keys.update.assert_not_called()