mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
0b40c71489
commit
eef5541e60
4 changed files with 290 additions and 15 deletions
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue