diff --git a/litellm/litellm_core_utils/asyncify.py b/litellm/litellm_core_utils/asyncify.py index b58e707b8f8..17cfb3ef84f 100644 --- a/litellm/litellm_core_utils/asyncify.py +++ b/litellm/litellm_core_utils/asyncify.py @@ -1,5 +1,6 @@ import asyncio import functools +import threading from collections.abc import Awaitable, Callable from typing import Final @@ -68,6 +69,18 @@ def asyncify( return wrapper +def is_event_loop_running() -> bool: + try: + asyncio.get_running_loop() + except RuntimeError: + return False + return True + + +def can_block_current_thread() -> bool: + return threading.current_thread() is threading.main_thread() and not is_event_loop_running() + + def run_async_function(async_function, *args, **kwargs): """ Helper utility to run an async function in a sync context. diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 563826c2b93..c27bb83437e 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -9,6 +9,7 @@ import httpx from pydantic import JsonValue, TypeAdapter, ValidationError from litellm._logging import verbose_logger +from litellm.litellm_core_utils.asyncify import can_block_current_thread from litellm.llms.custom_httpx.http_handler import _get_httpx_client from .common_utils import ( @@ -25,6 +26,7 @@ from .common_utils import ( ) TOKEN_EXPIRY_SKEW_SECONDS: Final = 60 +TOKEN_REFRESH_TIMEOUT_SECONDS: Final = 30 DEVICE_CODE_TIMEOUT_SECONDS: Final = 15 * 60 DEVICE_CODE_COOLDOWN_SECONDS: Final = 5 * 60 DEVICE_CODE_POLL_SLEEP_SECONDS: Final = 5 @@ -66,6 +68,18 @@ class Authenticator: except RefreshAccessTokenError as exc: verbose_logger.warning("ChatGPT refresh token failed, re-login required: %s", exc) + if not can_block_current_thread(): + raise GetAccessTokenError( + message=( + "ChatGPT device-code login needs a human and cannot run inside a running event loop " + "or a worker thread (for example the LiteLLM proxy). Log in once outside the proxy with " + '`python -c "from litellm.llms.chatgpt.authenticator import Authenticator; ' + 'Authenticator().get_access_token()"` and mount the resulting auth.json into the proxy, ' + "or set CHATGPT_TOKEN_DIR to a directory that already holds it." + ), + status_code=401, + ) + cooldown_remaining: Final = self._get_device_code_cooldown_remaining(auth_data) if cooldown_remaining > 0: token: Final = self._wait_for_access_token(cooldown_remaining) @@ -309,6 +323,7 @@ class Authenticator: "refresh_token": refresh_token, "scope": "openid profile email", }, + timeout=TOKEN_REFRESH_TIMEOUT_SECONDS, ) resp.raise_for_status() data: Final = _JSON_OBJECT_ADAPTER.validate_python(resp.json()) diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index 80fd4f755e7..7821756bc16 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -7,6 +7,7 @@ from typing import Any, Final import httpx from litellm._logging import verbose_logger +from litellm.litellm_core_utils.asyncify import can_block_current_thread from litellm.llms.custom_httpx.http_handler import _get_httpx_client from .common_utils import ( @@ -57,6 +58,18 @@ class Authenticator: except OSError: verbose_logger.warning("No existing access token found or error reading file") + if not can_block_current_thread(): + raise GetAccessTokenError( + message=( + "GitHub Copilot device-code login needs a human and cannot run inside a running event loop " + "or a worker thread (for example the LiteLLM proxy). Log in once outside the proxy with " + '`python -c "from litellm.llms.github_copilot.authenticator import Authenticator; ' + 'Authenticator().get_access_token()"` and mount the resulting access-token file into ' + "the proxy, or set GITHUB_COPILOT_TOKEN_DIR to a directory that already holds it." + ), + status_code=401, + ) + for attempt in range(3): verbose_logger.debug("Access token acquisition attempt %s/3", attempt + 1) try: diff --git a/tests/test_litellm/llms/chatgpt/test_chatgpt_authenticator.py b/tests/test_litellm/llms/chatgpt/test_chatgpt_authenticator.py index a9ced2afcf9..833220e8a67 100644 --- a/tests/test_litellm/llms/chatgpt/test_chatgpt_authenticator.py +++ b/tests/test_litellm/llms/chatgpt/test_chatgpt_authenticator.py @@ -1,11 +1,16 @@ import base64 import json import time -from unittest.mock import mock_open, patch +from concurrent.futures import ThreadPoolExecutor +from unittest.mock import MagicMock, mock_open, patch import pytest -from litellm.llms.chatgpt.authenticator import Authenticator +from litellm.llms.chatgpt.authenticator import ( + TOKEN_REFRESH_TIMEOUT_SECONDS, + Authenticator, +) +from litellm.llms.chatgpt.common_utils import GetAccessTokenError def _make_jwt(payload: dict) -> str: @@ -20,9 +25,9 @@ def _make_jwt(payload: dict) -> str: class TestChatGPTAuthenticator: @pytest.fixture - def authenticator(self): - with patch("os.path.exists", return_value=True): - return Authenticator() + def authenticator(self, tmp_path, monkeypatch): + monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path)) + return Authenticator() def test_get_access_token_from_file(self, authenticator): future_time = time.time() + 3600 @@ -54,10 +59,84 @@ class TestChatGPTAuthenticator: token = authenticator.get_access_token() assert token == "token-new" + def test_refresh_tokens_uses_bounded_timeout(self, authenticator): + client = MagicMock() + response = MagicMock() + response.json.return_value = { + "access_token": "token-new", + "id_token": "id-123", + } + client.post.return_value = response + + with patch( # test-quality-ok: requested seam for asserting timeout propagation + "litellm.llms.chatgpt.authenticator._get_httpx_client", return_value=client + ): + refreshed = authenticator._refresh_tokens("refresh-123") + + assert refreshed["access_token"] == "token-new" + assert client.post.call_args.kwargs["timeout"] == TOKEN_REFRESH_TIMEOUT_SECONDS + + @pytest.mark.asyncio + async def test_get_access_token_refuses_device_code_login_in_event_loop(self, authenticator): + with ( + patch("builtins.open", side_effect=FileNotFoundError), + patch.object(authenticator, "_login_device_code") as mock_login, + patch.object(authenticator, "_wait_for_access_token") as mock_wait, + ): + with pytest.raises(GetAccessTokenError) as exc: + authenticator.get_access_token() + + assert exc.value.status_code == 401 + assert "event loop" in str(exc.value) + assert authenticator.auth_file not in str(exc.value) + mock_login.assert_not_called() + mock_wait.assert_not_called() + + @pytest.mark.asyncio + async def test_get_access_token_refuses_cooldown_wait_in_event_loop(self, authenticator): + auth_data = json.dumps({"device_code_requested_at": time.time()}) + + with ( + patch("builtins.open", mock_open(read_data=auth_data)), + patch.object(authenticator, "_login_device_code") as mock_login, + patch.object(authenticator, "_wait_for_access_token") as mock_wait, + ): + with pytest.raises(GetAccessTokenError) as exc: + authenticator.get_access_token() + + assert exc.value.status_code == 401 + assert "event loop" in str(exc.value) + assert authenticator.auth_file not in str(exc.value) + mock_login.assert_not_called() + mock_wait.assert_not_called() + + def test_get_access_token_refuses_device_code_login_in_worker_thread(self, authenticator): + with ( + patch("builtins.open", side_effect=FileNotFoundError), + patch.object(authenticator, "_login_device_code") as mock_login, + patch.object(authenticator, "_wait_for_access_token") as mock_wait, + ): + with ThreadPoolExecutor(max_workers=1) as pool: + with pytest.raises(GetAccessTokenError) as exc: + pool.submit(authenticator.get_access_token).result() + + assert exc.value.status_code == 401 + assert "worker thread" in str(exc.value) + assert authenticator.auth_file not in str(exc.value) + mock_login.assert_not_called() + mock_wait.assert_not_called() + + def test_get_access_token_device_code_login_without_event_loop(self, authenticator): + with ( + patch("builtins.open", side_effect=FileNotFoundError), + patch.object(authenticator, "_login_device_code", return_value={"access_token": "tok"}), + ): + token = authenticator.get_access_token() + + assert token == "tok" + def test_get_account_id_from_id_token(self, authenticator): - id_token = _make_jwt( - {"https://api.openai.com/auth": {"chatgpt_account_id": "acct-123"}} - ) + id_token = _make_jwt({"https://api.openai.com/auth": {"chatgpt_account_id": "acct-123"}}) auth_data = json.dumps({"id_token": id_token}) with ( diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py index 6c846a90c71..a49a4b44b74 100644 --- a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py +++ b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py @@ -1,6 +1,7 @@ import json import os import time +from concurrent.futures import ThreadPoolExecutor from datetime import datetime, timedelta from unittest.mock import MagicMock, mock_open, patch @@ -89,6 +90,34 @@ class TestGitHubCopilotAuthenticator: assert token == mock_token authenticator._login.assert_called_once() + @pytest.mark.asyncio + async def test_get_access_token_refuses_device_code_login_in_event_loop(self, authenticator): + with ( + patch("builtins.open", side_effect=FileNotFoundError), + patch.object(authenticator, "_login") as mock_login, + ): + with pytest.raises(GetAccessTokenError) as exc: + authenticator.get_access_token() + + assert exc.value.status_code == 401 + assert "event loop" in str(exc.value) + assert authenticator.access_token_file not in str(exc.value) + mock_login.assert_not_called() + + def test_get_access_token_refuses_device_code_login_in_worker_thread(self, authenticator): + with ( + patch("builtins.open", side_effect=FileNotFoundError), + patch.object(authenticator, "_login") as mock_login, + ): + with ThreadPoolExecutor(max_workers=1) as pool: + with pytest.raises(GetAccessTokenError) as exc: + pool.submit(authenticator.get_access_token).result() + + assert exc.value.status_code == 401 + assert "worker thread" in str(exc.value) + assert authenticator.access_token_file not in str(exc.value) + mock_login.assert_not_called() + def test_get_access_token_failure(self, authenticator): """Test that an exception is raised after multiple login failures.""" with (