From 501baea61cbbef6c9c902ab535516412e0ee5d8d Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 29 Sep 2026 17:50:41 +0000 Subject: [PATCH] fix(cli): serialize session replacement with token renewal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/client/cli/commands/auth.py | 52 +++++++++-- .../proxy/client/cli/test_auth_commands.py | 88 +++++++++++++++++++ 2 files changed, 134 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index b0bbd28f372..13266e2f72d 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -3,11 +3,13 @@ import sys import time import webbrowser from collections.abc import Callable, Mapping, Sequence +from pathlib import Path from typing import Any, Final, TypeVar from urllib.parse import urlencode import click import requests +from filelock import BaseFileLock, FileLock, Timeout from rich.console import Console from rich.table import Table from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never @@ -239,6 +241,10 @@ def _renewal_reader(vault: SecretVault) -> Callable[[], Mapping[str, object] | N return reload +def _credential_lock() -> BaseFileLock: + return FileLock(str(Path.home() / ".litellm-token.lock"), timeout=30, mode=0o600) + + def get_stored_api_key( expected_base_url: str | None = None, *, @@ -256,6 +262,26 @@ def get_stored_api_key( return None if expected_base_url is not None and token_data.get("base_url") != expected_base_url.rstrip("/"): return None + if is_cli_token_fresh(token_data) or not token_data.get("refresh_token"): + return _key_from_record(token_data, vault) + try: + with _credential_lock(): + return _get_stored_api_key(expected_base_url, vault) + except (OSError, Timeout) as error: + _warn(f"Could not lock the saved login: {error}") + return None + + +def _get_stored_api_key(expected_base_url: str | None, vault: SecretVault) -> str | None: + token_data: Final = load_token(vault=vault) + if token_data is None: + return None + if expected_base_url is not None and token_data.get("base_url") != expected_base_url.rstrip("/"): + return None + return _key_from_record(token_data, vault) + + +def _key_from_record(token_data: Mapping[str, object], vault: SecretVault) -> str | None: return fresh_api_key( token_data, _renewal_saver(vault), @@ -868,6 +894,14 @@ def _replace_stored_token( "it expires on its own." ) raise click.ClickException("The login did not select the requested team; your saved login has not changed") + try: + with _credential_lock(): + return _persist_replacement(record, http, vault) + except (OSError, Timeout) as error: + return CredentialNotSaved(f"Could not lock the saved login: {error}") + + +def _persist_replacement(record: CliTokenData, http: Http, vault: SecretVault) -> SecretSave: previous: Final = load_token(vault=vault) previous_secret: Final = vault.read() stored: Final = save_token(record, vault=vault) @@ -1051,6 +1085,14 @@ def login_to_proxy( def logout(ctx: click.Context): """Logout and clear stored authentication""" vault: Final = context_secret_vault(ctx) + try: + with _credential_lock(): + _logout(vault) + except (OSError, Timeout) as error: + raise click.ClickException(f"Could not lock the saved login: {error}") from error + + +def _logout(vault: SecretVault) -> None: token_data: Final = load_token(vault=vault) revocation: Final = revoke_stored_credential(token_data, requests.Session()) if token_data is not None else None match revocation: @@ -1124,15 +1166,13 @@ def print_token(ctx: click.Context): click.echo(keychain_unreadable_notice(vault), err=True) sys.exit(1) + saved_base_url: Final = token_data.get("base_url") api_key: Final = ( ctx_obj.get("api_key") if issued_for_this_server and ctx_obj.get("api_key_from_token_file") - else fresh_api_key( - token_data, - _renewal_saver(vault), - requests.Session(), - reload=_renewal_reader(vault), - warn=_warn, + else get_stored_api_key( + saved_base_url if isinstance(saved_base_url, str) else None, + vault=vault, ) ) if not api_key: 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 a2fe8e25603..3a23c49689e 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -2,8 +2,10 @@ import json import os import stat import time +from concurrent.futures import ThreadPoolExecutor from dataclasses import replace from pathlib import Path +from threading import Event from unittest.mock import Mock, patch @@ -28,6 +30,7 @@ from litellm.proxy.client.cli import cli from litellm.proxy.client.cli.commands import claude_settings as claude_settings_module from litellm.proxy.client.cli.commands.claude_settings import SettingsFileOwner from litellm.proxy.client.cli.commands.auth import ( + _replace_stored_token, get_stored_api_key, load_token, login, @@ -2328,6 +2331,91 @@ def test_assign_key_handles_keychain_update_without_metadata(isolated_home, secr ) +@pytest.mark.parametrize("reader", ["stored-key", "print-token"]) +def test_partial_replacement_preserves_a_concurrently_refreshed_login(isolated_home, secret_vault_factory, reader): + vault = secret_vault_factory() + save_cli_token(CliTokenRecord(**_pkce_record(team_id="team-a")), vault=vault) + refreshing = Event() + replacement_staged = Event() + rotation_saved = Event() + replace_file = os.replace + http = Mock() + + def refresh_response(*_, **__): + refreshing.set() + replacement_staged.wait(1) + return _FakeHttpResponse(200, {**PKCE_TOKEN_RESPONSE, "team_id": "team-a"}) + + def commit_metadata(source, target): + if json.loads(Path(source).read_text())["team_id"] == "team-b": + replacement_staged.set() + assert rotation_saved.wait(5), "renewal did not persist its rotated credential" + raise OSError("cannot replace team metadata") + replace_file(source, target) + rotation_saved.set() + + def read_key(): + if reader == "stored-key": + return get_stored_api_key(PKCE_BASE_URL, vault=vault) + result = CliRunner().invoke(print_token, obj={"base_url": PKCE_BASE_URL, "secret_vault": vault}) + assert result.exit_code == 0, result.output + return result.output.strip() + + http.post.side_effect = refresh_response + with ( + patch("litellm.proxy.client.cli.commands.auth.requests.Session", return_value=http), + patch("litellm.litellm_core_utils.private_json.os.replace", side_effect=commit_metadata), + ThreadPoolExecutor(max_workers=2) as executor, + ): + renewal = executor.submit(read_key) + assert refreshing.wait(5), "renewal did not reach the proxy" + replacement = executor.submit( + _replace_stored_token, + _pkce_record(key="session-new", refresh_token="llm_srefresh_new", team_id="team-b"), + _FakeSession(), + vault, + "team-b", + ) + assert renewal.result(timeout=10) == "sk-cli-rotated" + assert isinstance(replacement.result(timeout=10), CredentialNotSaved) + + saved = load_token(vault=vault) + assert (saved["key"], saved["refresh_token"], saved["team_id"]) == ( + "sk-cli-rotated", + "llm_srefresh_rotated", + "team-a", + ) + + +def test_reading_a_fresh_login_does_not_need_a_writable_lock(isolated_home, secret_vault_factory, capsys): + vault = secret_vault_factory() + (isolated_home / ".litellm-token.lock").mkdir() + assert get_stored_api_key(PKCE_BASE_URL, vault=vault) is None + save_cli_token(CliTokenRecord(**_pkce_record(expires_at=time.time() + 3600)), vault=vault) + + assert get_stored_api_key(PKCE_BASE_URL, vault=vault) == "sk-cli-old" + assert capsys.readouterr().err == "" + + +def test_lock_failure_leaves_saved_login_and_refresh_token_untouched(isolated_home, secret_vault_factory): + vault = secret_vault_factory() + save_cli_token(CliTokenRecord(**_pkce_record(team_id="team-a")), vault=vault) + before = load_token(vault=vault) + (isolated_home / ".litellm-token.lock").mkdir() + http = _FakeSession() + + assert get_stored_api_key(PKCE_BASE_URL, vault=vault) is None + outcome = _replace_stored_token(_pkce_record(team_id="team-b"), http, vault, "team-b") + result = CliRunner().invoke(logout, obj={"secret_vault": vault}) + + assert isinstance(outcome, CredentialNotSaved) + assert "Could not lock the saved login" in outcome.detail + assert result.exit_code == 1, result.output + assert "Could not lock the saved login" in result.output + assert load_token(vault=vault) == before + assert http.posts == [] + + 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)))