mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge 94f34b7032 into 9071ca503e
This commit is contained in:
commit
d12371e679
5 changed files with 157 additions and 8 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue