fix(cli): renew saved sessions when assigning teams

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-29 16:57:33 +00:00
parent 0b40c71489
commit eef5541e60
4 changed files with 290 additions and 15 deletions

View file

@ -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}

View file

@ -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")

View file

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

View file

@ -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()