mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(cli): serialize session replacement with token renewal
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
08dbab7914
commit
501baea61c
2 changed files with 134 additions and 6 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue