diff --git a/litellm/litellm_core_utils/cli_credential_lock.py b/litellm/litellm_core_utils/cli_credential_lock.py new file mode 100644 index 00000000000..f38be474868 --- /dev/null +++ b/litellm/litellm_core_utils/cli_credential_lock.py @@ -0,0 +1,74 @@ +import ctypes +import errno +import os +import sys +import time +from collections.abc import Generator +from contextlib import contextmanager +from ctypes import wintypes +from hashlib import sha256 +from pathlib import Path +from typing import TYPE_CHECKING, Final, cast + +from filelock import Timeout + +if TYPE_CHECKING or sys.platform != "win32": + import fcntl + + +@contextmanager +def credential_lock(home: Path, timeout: float = 30) -> Generator[None, None, None]: + """Serialize credential changes without creating or modifying a lock file.""" + if sys.platform == "win32": + with _windows_mutex(home, timeout): + yield + return + fd: Final = os.open(home, os.O_RDONLY | os.O_DIRECTORY) + try: + deadline: Final = time.monotonic() + timeout + while True: + try: + fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB) + break + except OSError as error: + if error.errno not in (errno.EACCES, errno.EAGAIN, errno.EWOULDBLOCK): + raise + if time.monotonic() >= deadline: + raise Timeout(str(home)) from None + time.sleep(0.05) + try: + yield + finally: + fcntl.flock(fd, fcntl.LOCK_UN) + finally: + os.close(fd) + + +@contextmanager +def _windows_mutex(home: Path, timeout: float) -> Generator[None, None, None]: + kernel: Final = ctypes.WinDLL("kernel32", use_last_error=True) + create: Final = ctypes.WINFUNCTYPE( + wintypes.HANDLE, ctypes.c_void_p, wintypes.BOOL, wintypes.LPCWSTR, use_last_error=True + )(("CreateMutexW", kernel)) + wait: Final = ctypes.WINFUNCTYPE(wintypes.DWORD, wintypes.HANDLE, wintypes.DWORD, use_last_error=True)( + ("WaitForSingleObject", kernel) + ) + release: Final = ctypes.WINFUNCTYPE(wintypes.BOOL, wintypes.HANDLE, use_last_error=True)(("ReleaseMutex", kernel)) + close: Final = ctypes.WINFUNCTYPE(wintypes.BOOL, wintypes.HANDLE, use_last_error=True)(("CloseHandle", kernel)) + identity: Final = sha256(os.path.normcase(str(home.resolve())).encode()).hexdigest() + handle: Final = cast(int | None, create(None, False, f"Global\\litellm-cli-{identity}")) + if handle is None: + raise ctypes.WinError(ctypes.get_last_error()) + try: + result: Final = cast(int, wait(handle, max(0, int(timeout * 1000)))) + if result == 0x102: + raise Timeout(str(home)) + if result not in (0, 0x80): + raise ctypes.WinError(ctypes.get_last_error()) + try: + yield + finally: + if not release(handle) and sys.exc_info()[0] is None: + raise ctypes.WinError(ctypes.get_last_error()) + finally: + close(handle) diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 95af9088c38..0bba43cd75a 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -3,18 +3,20 @@ import sys import time import webbrowser from collections.abc import Callable, Mapping, Sequence +from contextlib import AbstractContextManager 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 filelock import Timeout from rich.console import Console from rich.table import Table from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never from litellm.constants import CLI_JWT_EXPIRATION_HOURS +from litellm.litellm_core_utils.cli_credential_lock import credential_lock from litellm.litellm_core_utils.cli_keyring import ( DISABLE_KEYRING_ENV_VAR, SYSTEM_KEYRING, @@ -241,8 +243,8 @@ 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 _credential_lock() -> AbstractContextManager[None]: + return credential_lock(Path.home()) def get_stored_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 3a23c49689e..1c11c5fc673 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -1,6 +1,7 @@ import json import os import stat +import sys import time from concurrent.futures import ThreadPoolExecutor from dataclasses import replace @@ -11,8 +12,10 @@ from unittest.mock import Mock, patch import pytest from click.testing import CliRunner +from filelock import Timeout from litellm.constants import CLI_JWT_EXPIRATION_HOURS +from litellm.litellm_core_utils.cli_credential_lock import _windows_mutex, credential_lock from litellm.litellm_core_utils.cli_keyring import ( DISABLE_KEYRING_ENV_VAR, KeyringDisabled, @@ -39,6 +42,7 @@ from litellm.proxy.client.cli.commands.auth import ( whoami, ) from litellm.proxy.client.cli.commands.pkce_login import PkceFailure, RevocationUnavailable +from tests.test_litellm_rust.support.child_interpreter import run_child_interpreter @pytest.fixture @@ -2387,13 +2391,39 @@ def test_partial_replacement_preserves_a_concurrently_refreshed_login(isolated_h ) +@pytest.mark.parametrize("operation", ["logout", "renewal", "replacement"]) +def test_saved_login_works_with_read_only_home(isolated_home, secret_vault_factory, operation): + vault = secret_vault_factory() + save_cli_token(CliTokenRecord(**_pkce_record(team_id="team-a")), vault=vault) + http = _FakeSession() + isolated_home.chmod(0o500) + try: + with patch("litellm.proxy.client.cli.commands.auth.requests.Session") as session: + session.return_value.post.return_value = _FakeHttpResponse(200, PKCE_TOKEN_RESPONSE) + if operation == "logout": + result = CliRunner().invoke(logout, obj={"secret_vault": vault}) + assert result.exit_code == 0, result.output + assert load_token(vault=vault) is None + elif operation == "renewal": + assert get_stored_api_key(PKCE_BASE_URL, vault=vault) == "sk-cli-rotated" + assert load_token(vault=vault)["refresh_token"] == "llm_srefresh_rotated" + else: + outcome = _replace_stored_token(_pkce_record(team_id="team-b"), http, vault, "team-b") + assert isinstance(outcome, SecretStored) + assert load_token(vault=vault)["team_id"] == "team-b" + finally: + isolated_home.chmod(0o700) + + 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 + failing_boundary = "ctypes.WinDLL" if sys.platform == "win32" else "os.open" + with patch(failing_boundary, side_effect=PermissionError("lock unavailable")): + 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" + with patch(failing_boundary, side_effect=PermissionError("lock unavailable")): + assert get_stored_api_key(PKCE_BASE_URL, vault=vault) == "sk-cli-old" assert capsys.readouterr().err == "" @@ -2401,12 +2431,13 @@ def test_lock_failure_leaves_saved_login_and_refresh_token_untouched(isolated_ho 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}) + failing_boundary = "ctypes.WinDLL" if sys.platform == "win32" else "os.open" + with patch(failing_boundary, side_effect=PermissionError("lock unavailable")): + 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 @@ -2416,6 +2447,101 @@ def test_lock_failure_leaves_saved_login_and_refresh_token_untouched(isolated_ho assert http.posts == [] +@pytest.mark.parametrize("read_only", ["home", "config"]) +def test_credential_lock_serializes_processes_without_writable_directories(isolated_home, read_only): + config = isolated_home / ".litellm" + config.mkdir() + directory = isolated_home if read_only == "home" else config + directory.chmod(0o500) + child = """ +from pathlib import Path +from filelock import Timeout +from litellm.litellm_core_utils.cli_credential_lock import credential_lock +try: + with credential_lock(Path.home(), timeout=0): + print('acquired') +except Timeout: + print('blocked') +""" + env = {**os.environ, "TMPDIR": str(isolated_home / "different-temporary-directory")} + + def interrupted_writer(): + with credential_lock(isolated_home): + blocked = run_child_interpreter(child, env=env, timeout=60) + assert blocked.returncode == 0, blocked.stderr + assert blocked.stdout.strip() == "blocked" + raise RuntimeError("release on error") + + try: + with pytest.raises(RuntimeError, match="release on error"): + interrupted_writer() + released = run_child_interpreter(child, env=env, timeout=60) + assert released.returncode == 0, released.stderr + assert released.stdout.strip() == "acquired" + finally: + directory.chmod(0o700) + + +@pytest.fixture +def windows_mutex_api(): + api = { + "CreateMutexW": Mock(return_value=123), + "WaitForSingleObject": Mock(return_value=0), + "ReleaseMutex": Mock(return_value=True), + "CloseHandle": Mock(return_value=True), + } + with ( + patch("ctypes.WinDLL", create=True), + patch("ctypes.WINFUNCTYPE", return_value=lambda binding: api[binding[0]], create=True), + patch("ctypes.get_last_error", return_value=5, create=True), + patch("ctypes.WinError", side_effect=lambda code=0: OSError(code, "Win32 error"), create=True), + ): + yield api + + +@pytest.mark.parametrize("wait_result", [0, 0x80]) +@pytest.mark.parametrize("release_result", [True, False]) +def test_windows_mutex_preserves_body_errors_and_closes_handle( + tmp_path, windows_mutex_api, wait_result, release_result +): + windows_mutex_api["WaitForSingleObject"].return_value = wait_result + windows_mutex_api["ReleaseMutex"].return_value = release_result + + with pytest.raises(ValueError, match="body failed"), _windows_mutex(tmp_path, timeout=0): + raise ValueError("body failed") + + windows_mutex_api["ReleaseMutex"].assert_called_once_with(123) + windows_mutex_api["CloseHandle"].assert_called_once_with(123) + + +@pytest.mark.parametrize( + ("handle", "wait_result", "error_type", "error_code"), + [(None, 0, OSError, 5), (123, 0xFFFFFFFF, OSError, 5), (123, 0x102, Timeout, None)], +) +def test_windows_mutex_acquisition_failure_preserves_error_and_closes_handle( + tmp_path, windows_mutex_api, handle, wait_result, error_type, error_code +): + windows_mutex_api["CreateMutexW"].return_value = handle + windows_mutex_api["WaitForSingleObject"].return_value = wait_result + + with pytest.raises(error_type) as raised, _windows_mutex(tmp_path, timeout=0): + pytest.fail("failed acquisition entered the critical section") + + assert raised.value.errno == error_code + windows_mutex_api["ReleaseMutex"].assert_not_called() + assert windows_mutex_api["CloseHandle"].call_count == (0 if handle is None else 1) + + +def test_windows_mutex_release_failure_preserves_error_and_closes_handle(tmp_path, windows_mutex_api): + windows_mutex_api["ReleaseMutex"].return_value = False + + with pytest.raises(OSError, match="Win32 error") as raised, _windows_mutex(tmp_path, timeout=0): + pass + + assert raised.value.errno == 5 + windows_mutex_api["CloseHandle"].assert_called_once_with(123) + + 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)))