mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 7d00e16c26 into c39bf62936
This commit is contained in:
commit
0b6bdd355d
5 changed files with 228 additions and 95 deletions
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import Any, Final
|
||||
|
|
@ -37,36 +38,43 @@ class Authenticator:
|
|||
os.getenv("GITHUB_COPILOT_ACCESS_TOKEN_FILE", "access-token"),
|
||||
)
|
||||
self.api_key_file = os.path.join(self.token_dir, os.getenv("GITHUB_COPILOT_API_KEY_FILE", "api-key.json"))
|
||||
self._ensure_token_dir()
|
||||
self._session_cache: dict[str, dict[str, Any]] = {} # mutable-ok: session cache
|
||||
|
||||
def get_access_token(self) -> str:
|
||||
def get_access_token(self, access_token: str | None = None) -> str:
|
||||
"""
|
||||
Login to Copilot with retry 3 times.
|
||||
|
||||
Args:
|
||||
access_token: Optional access token passed via config/request params.
|
||||
|
||||
Returns:
|
||||
str: The GitHub access token.
|
||||
|
||||
Raises:
|
||||
GetAccessTokenError: If unable to obtain an access token after retries.
|
||||
"""
|
||||
if access_token:
|
||||
return access_token
|
||||
|
||||
try:
|
||||
with open(self.access_token_file, "r") as f:
|
||||
access_token = f.read().strip()
|
||||
if access_token:
|
||||
return access_token
|
||||
saved_token: Final = f.read().strip()
|
||||
if saved_token:
|
||||
return saved_token
|
||||
except OSError:
|
||||
verbose_logger.warning("No existing access token found or error reading file")
|
||||
|
||||
for attempt in range(3):
|
||||
verbose_logger.debug("Access token acquisition attempt %s/3", attempt + 1)
|
||||
try:
|
||||
access_token = self._login()
|
||||
new_token = self._login() # rebind-ok: loop variable
|
||||
try:
|
||||
self._ensure_token_dir()
|
||||
with open(self.access_token_file, "w") as f:
|
||||
f.write(access_token)
|
||||
f.write(new_token)
|
||||
except OSError:
|
||||
verbose_logger.error("Error saving access token to file")
|
||||
return access_token
|
||||
return new_token
|
||||
except (GetDeviceCodeError, GetAccessTokenError, RefreshAPIKeyError) as e:
|
||||
verbose_logger.warning("Failed attempt %s: %s", attempt + 1, e)
|
||||
continue
|
||||
|
|
@ -76,21 +84,52 @@ class Authenticator:
|
|||
status_code=401,
|
||||
)
|
||||
|
||||
def get_api_key(self) -> str:
|
||||
def get_api_key(self, access_token: str | None = None) -> str:
|
||||
"""
|
||||
Get the API key, refreshing if necessary.
|
||||
|
||||
Args:
|
||||
access_token: Optional GitHub access token passed via config/request params.
|
||||
|
||||
Returns:
|
||||
str: The GitHub Copilot API key.
|
||||
|
||||
Raises:
|
||||
GetAPIKeyError: If unable to obtain an API key.
|
||||
"""
|
||||
# When an explicit access_token is provided, isolate in memory to prevent sharing a cached key across callers
|
||||
if access_token:
|
||||
cached_info: Final = self._session_cache.get(access_token)
|
||||
if cached_info and cached_info.get("expires_at", 0) > time.time():
|
||||
cached_token: Final = cached_info.get("token")
|
||||
if isinstance(cached_token, str):
|
||||
return cached_token
|
||||
|
||||
try:
|
||||
session_key_info: Final = self._refresh_api_key(access_token)
|
||||
self._session_cache[access_token] = session_key_info
|
||||
refreshed_token: Final = session_key_info.get("token")
|
||||
if isinstance(refreshed_token, str):
|
||||
return refreshed_token
|
||||
else:
|
||||
raise GetAPIKeyError(
|
||||
message="API key response missing token",
|
||||
status_code=401,
|
||||
)
|
||||
except RefreshAPIKeyError as e:
|
||||
raise GetAPIKeyError(
|
||||
message=f"Failed to refresh API key: {e}",
|
||||
status_code=401,
|
||||
)
|
||||
|
||||
# Fallback to single-user local file storage when no explicit token is passed
|
||||
try:
|
||||
with open(self.api_key_file, "r") as f:
|
||||
api_key_info = json.load(f)
|
||||
if api_key_info.get("expires_at", 0) > datetime.now().timestamp():
|
||||
return api_key_info.get("token")
|
||||
file_key_info: Final = json.load(f)
|
||||
if isinstance(file_key_info, dict) and file_key_info.get("expires_at", 0) > datetime.now().timestamp():
|
||||
file_token: Final = file_key_info.get("token")
|
||||
if isinstance(file_token, str):
|
||||
return file_token
|
||||
else:
|
||||
verbose_logger.warning("API key expired, refreshing")
|
||||
raise APIKeyExpiredError(
|
||||
|
|
@ -105,23 +144,21 @@ class Authenticator:
|
|||
pass # Already logged in the try block
|
||||
|
||||
try:
|
||||
api_key_info = self._refresh_api_key()
|
||||
with open(self.api_key_file, "w") as f:
|
||||
json.dump(api_key_info, f)
|
||||
token: Final = api_key_info.get("token")
|
||||
if token:
|
||||
return token
|
||||
refreshed_key_info: Final = self._refresh_api_key()
|
||||
try:
|
||||
self._ensure_token_dir()
|
||||
with open(self.api_key_file, "w") as f:
|
||||
json.dump(refreshed_key_info, f)
|
||||
except OSError as e:
|
||||
verbose_logger.warning("Error saving API key to file (continuing with in-memory token): %s", e)
|
||||
new_key_token: Final = refreshed_key_info.get("token")
|
||||
if isinstance(new_key_token, str):
|
||||
return new_key_token
|
||||
else:
|
||||
raise GetAPIKeyError(
|
||||
message="API key response missing token",
|
||||
status_code=401,
|
||||
)
|
||||
except OSError as e:
|
||||
verbose_logger.error("Error saving API key to file: %s", e)
|
||||
raise GetAPIKeyError(
|
||||
message=f"Failed to save API key: {e}",
|
||||
status_code=500,
|
||||
)
|
||||
except RefreshAPIKeyError as e:
|
||||
raise GetAPIKeyError(
|
||||
message=f"Failed to refresh API key: {e}",
|
||||
|
|
@ -138,37 +175,44 @@ class Authenticator:
|
|||
try:
|
||||
with open(self.api_key_file, "r") as f:
|
||||
api_key_info: Final = json.load(f)
|
||||
endpoints: Final = api_key_info.get("endpoints", {})
|
||||
api_endpoint: Final = endpoints.get("api")
|
||||
return api_endpoint
|
||||
if isinstance(api_key_info, dict):
|
||||
endpoints: Final = api_key_info.get("endpoints", {})
|
||||
if isinstance(endpoints, dict):
|
||||
api_endpoint: Final = endpoints.get("api")
|
||||
if isinstance(api_endpoint, str):
|
||||
return api_endpoint
|
||||
return None
|
||||
except (OSError, json.JSONDecodeError, KeyError) as e:
|
||||
verbose_logger.warning("Error reading API endpoint from file: %s", e)
|
||||
return None
|
||||
|
||||
def _refresh_api_key(self) -> dict[str, Any]:
|
||||
def _refresh_api_key(self, access_token: str | None = None) -> dict[str, Any]:
|
||||
"""
|
||||
Refresh the API key using the access token.
|
||||
|
||||
Args:
|
||||
access_token: Optional access token. If not provided, will call get_access_token().
|
||||
|
||||
Returns:
|
||||
Dict[str, Any]: The API key information including token and expiration.
|
||||
|
||||
Raises:
|
||||
RefreshAPIKeyError: If unable to refresh the API key.
|
||||
"""
|
||||
access_token: Final = self.get_access_token()
|
||||
headers: Final = self._get_github_headers(access_token)
|
||||
resolved_token: Final = access_token or self.get_access_token()
|
||||
headers: Final = self._get_github_headers(resolved_token)
|
||||
api_key_url: Final = os.getenv("GITHUB_COPILOT_API_KEY_URL", DEFAULT_GITHUB_API_KEY_URL)
|
||||
|
||||
max_retries: Final = 3
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
sync_client = _get_httpx_client()
|
||||
response = sync_client.get(api_key_url, headers=headers)
|
||||
sync_client = _get_httpx_client() # rebind-ok: loop variable
|
||||
response = sync_client.get(api_key_url, headers=headers) # rebind-ok: loop variable
|
||||
response.raise_for_status()
|
||||
|
||||
response_json = response.json()
|
||||
response_json = response.json() # rebind-ok: loop variable
|
||||
|
||||
if "token" in response_json:
|
||||
if isinstance(response_json, dict) and "token" in response_json:
|
||||
return response_json
|
||||
else:
|
||||
verbose_logger.warning("API key response missing token: %s", response_json)
|
||||
|
|
@ -183,9 +227,28 @@ class Authenticator:
|
|||
)
|
||||
|
||||
def _ensure_token_dir(self) -> None:
|
||||
"""Ensure the token directory exists."""
|
||||
if not os.path.exists(self.token_dir):
|
||||
os.makedirs(self.token_dir, exist_ok=True)
|
||||
"""Ensure the token directory exists, falling back to secure user-isolated temp directory on permission errors."""
|
||||
try:
|
||||
if not os.path.exists(self.token_dir):
|
||||
os.makedirs(self.token_dir, mode=0o700, exist_ok=True)
|
||||
except (PermissionError, OSError) as e:
|
||||
verbose_logger.warning(
|
||||
"Cannot create token directory at %s (%s). Falling back to temp directory.",
|
||||
self.token_dir,
|
||||
e,
|
||||
)
|
||||
try:
|
||||
self.token_dir = tempfile.mkdtemp(prefix="litellm_copilot_")
|
||||
except OSError:
|
||||
pass
|
||||
self.access_token_file = os.path.join(
|
||||
self.token_dir,
|
||||
os.getenv("GITHUB_COPILOT_ACCESS_TOKEN_FILE", "access-token"),
|
||||
)
|
||||
self.api_key_file = os.path.join(
|
||||
self.token_dir,
|
||||
os.getenv("GITHUB_COPILOT_API_KEY_FILE", "api-key.json"),
|
||||
)
|
||||
|
||||
def _get_github_headers(self, access_token: str | None = None) -> dict[str, str]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ class GithubCopilotConfig(OpenAIConfig):
|
|||
or DEFAULT_GITHUB_COPILOT_API_BASE
|
||||
)
|
||||
try:
|
||||
dynamic_api_key: Final = self.authenticator.get_api_key()
|
||||
dynamic_api_key: Final = self.authenticator.get_api_key(access_token=api_key)
|
||||
except GetAPIKeyError as e:
|
||||
raise AuthenticationError(
|
||||
model=model,
|
||||
|
|
@ -94,7 +94,7 @@ class GithubCopilotConfig(OpenAIConfig):
|
|||
|
||||
# Add Copilot-specific headers (editor-version, user-agent, etc.)
|
||||
try:
|
||||
copilot_api_key: Final = self.authenticator.get_api_key()
|
||||
copilot_api_key: Final = self.authenticator.get_api_key(access_token=api_key)
|
||||
copilot_headers: Final = get_copilot_default_headers(copilot_api_key)
|
||||
validated_headers = {**copilot_headers, **validated_headers}
|
||||
except GetAPIKeyError:
|
||||
|
|
|
|||
|
|
@ -60,9 +60,9 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig):
|
|||
"""
|
||||
try:
|
||||
# Get GitHub Copilot API key via OAuth
|
||||
api_key = self.authenticator.get_api_key()
|
||||
copilot_api_key: Final = self.authenticator.get_api_key(access_token=api_key)
|
||||
|
||||
if not api_key:
|
||||
if not copilot_api_key:
|
||||
raise AuthenticationError(
|
||||
model=model,
|
||||
llm_provider="github_copilot",
|
||||
|
|
@ -70,7 +70,7 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig):
|
|||
)
|
||||
|
||||
# Get default headers
|
||||
default_headers: Final = get_copilot_default_headers(api_key)
|
||||
default_headers: Final = get_copilot_default_headers(copilot_api_key)
|
||||
|
||||
# Merge with existing headers (user's extra_headers take priority)
|
||||
merged_headers: Final = {**default_headers, **headers}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
"limit": 744
|
||||
},
|
||||
"TQ002": {
|
||||
"limit": 742
|
||||
"limit": 741
|
||||
},
|
||||
"TQ003": {
|
||||
"limit": 62
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import json
|
||||
import os
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
|
|
@ -8,7 +7,6 @@ import pytest
|
|||
|
||||
from litellm.llms.github_copilot.authenticator import Authenticator
|
||||
from litellm.llms.github_copilot.common_utils import (
|
||||
APIKeyExpiredError,
|
||||
GetAccessTokenError,
|
||||
GetAPIKeyError,
|
||||
GetDeviceCodeError,
|
||||
|
|
@ -19,13 +17,8 @@ from litellm.llms.github_copilot.common_utils import (
|
|||
class TestGitHubCopilotAuthenticator:
|
||||
@pytest.fixture
|
||||
def authenticator(self):
|
||||
with (
|
||||
patch("os.path.exists", return_value=False),
|
||||
patch("os.makedirs") as mock_makedirs,
|
||||
):
|
||||
auth = Authenticator()
|
||||
mock_makedirs.assert_called_once()
|
||||
return auth
|
||||
auth = Authenticator()
|
||||
return auth
|
||||
|
||||
@pytest.fixture
|
||||
def mock_http_client(self):
|
||||
|
|
@ -38,24 +31,63 @@ class TestGitHubCopilotAuthenticator:
|
|||
|
||||
def test_init(self):
|
||||
"""Test the initialization of the authenticator."""
|
||||
with (
|
||||
patch("os.path.exists", return_value=False),
|
||||
patch("os.makedirs") as mock_makedirs,
|
||||
):
|
||||
with patch("os.makedirs") as mock_makedirs:
|
||||
auth = Authenticator()
|
||||
assert auth.token_dir.endswith("/github_copilot")
|
||||
assert auth.access_token_file.endswith("/access-token")
|
||||
assert auth.api_key_file.endswith("/api-key.json")
|
||||
mock_makedirs.assert_called_once()
|
||||
assert os.path.basename(auth.token_dir) == "github_copilot"
|
||||
assert os.path.basename(auth.access_token_file) == "access-token"
|
||||
assert os.path.basename(auth.api_key_file) == "api-key.json"
|
||||
mock_makedirs.assert_not_called()
|
||||
|
||||
def test_ensure_token_dir(self):
|
||||
def test_ensure_token_dir(self, tmp_path):
|
||||
"""Test that the token directory is created if it doesn't exist."""
|
||||
test_dir = str(tmp_path / "new_copilot_dir")
|
||||
auth = Authenticator()
|
||||
auth.token_dir = test_dir
|
||||
auth._ensure_token_dir()
|
||||
assert os.path.exists(test_dir)
|
||||
|
||||
def test_ensure_token_dir_permission_error_fallback(self):
|
||||
"""Test that _ensure_token_dir falls back to temp directory on PermissionError."""
|
||||
auth = Authenticator()
|
||||
original_dir = auth.token_dir
|
||||
with (
|
||||
patch("os.path.exists", return_value=False),
|
||||
patch("os.makedirs") as mock_makedirs,
|
||||
patch("os.makedirs", side_effect=PermissionError("Permission denied")),
|
||||
):
|
||||
auth = Authenticator()
|
||||
mock_makedirs.assert_called_once_with(auth.token_dir, exist_ok=True)
|
||||
auth._ensure_token_dir()
|
||||
assert auth.token_dir != original_dir
|
||||
assert "litellm" in auth.token_dir
|
||||
|
||||
def test_get_api_key_with_explicit_token_isolation(self, authenticator):
|
||||
"""Test that explicit tokens use isolated in-memory caching and do not read stale disk files."""
|
||||
mock_data = {"token": "token-b-session-key", "expires_at": (datetime.now() + timedelta(hours=1)).timestamp()}
|
||||
with (
|
||||
patch.object(authenticator, "_refresh_api_key", return_value=mock_data) as mock_refresh,
|
||||
patch("builtins.open") as mock_file_open,
|
||||
):
|
||||
api_key = authenticator.get_api_key(access_token="user-b-custom-token")
|
||||
assert api_key == "token-b-session-key"
|
||||
mock_refresh.assert_called_once_with("user-b-custom-token")
|
||||
mock_file_open.assert_not_called()
|
||||
|
||||
# Second call with the same token should hit in-memory cache without calling _refresh_api_key again
|
||||
cached_key = authenticator.get_api_key(access_token="user-b-custom-token")
|
||||
assert cached_key == "token-b-session-key"
|
||||
assert mock_refresh.call_count == 1
|
||||
|
||||
def test_get_api_key_with_explicit_token_missing_token_in_response(self, authenticator):
|
||||
"""Test that get_api_key raises GetAPIKeyError when API response lacks token."""
|
||||
with patch.object(authenticator, "_refresh_api_key", return_value={}):
|
||||
with pytest.raises(GetAPIKeyError):
|
||||
authenticator.get_api_key(access_token="token-without-key")
|
||||
|
||||
def test_get_api_key_with_explicit_token_refresh_error(self, authenticator):
|
||||
"""Test that get_api_key handles RefreshAPIKeyError when refreshing explicit token."""
|
||||
with patch.object(
|
||||
authenticator, "_refresh_api_key", side_effect=RefreshAPIKeyError(message="Refresh failed", status_code=401)
|
||||
):
|
||||
with pytest.raises(GetAPIKeyError):
|
||||
authenticator.get_api_key(access_token="failing-token")
|
||||
|
||||
def test_get_github_headers(self, authenticator):
|
||||
"""Test that GitHub headers are correctly generated."""
|
||||
|
|
@ -82,8 +114,7 @@ class TestGitHubCopilotAuthenticator:
|
|||
|
||||
with (
|
||||
patch.object(authenticator, "_login", return_value=mock_token),
|
||||
patch("builtins.open", mock_open()),
|
||||
patch("builtins.open", side_effect=IOError) as mock_read,
|
||||
patch("builtins.open", side_effect=IOError),
|
||||
):
|
||||
token = authenticator.get_access_token()
|
||||
assert token == mock_token
|
||||
|
|
@ -106,9 +137,7 @@ class TestGitHubCopilotAuthenticator:
|
|||
def test_get_api_key_from_file(self, authenticator):
|
||||
"""Test retrieving an API key from a file."""
|
||||
future_time = (datetime.now() + timedelta(hours=1)).timestamp()
|
||||
mock_api_key_data = json.dumps(
|
||||
{"token": "mock-api-key", "expires_at": future_time}
|
||||
)
|
||||
mock_api_key_data = json.dumps({"token": "mock-api-key", "expires_at": future_time})
|
||||
|
||||
with patch("builtins.open", mock_open(read_data=mock_api_key_data)):
|
||||
api_key = authenticator.get_api_key()
|
||||
|
|
@ -117,9 +146,7 @@ class TestGitHubCopilotAuthenticator:
|
|||
def test_get_api_key_expired(self, authenticator):
|
||||
"""Test refreshing an expired API key."""
|
||||
past_time = (datetime.now() - timedelta(hours=1)).timestamp()
|
||||
mock_expired_data = json.dumps(
|
||||
{"token": "expired-api-key", "expires_at": past_time}
|
||||
)
|
||||
mock_expired_data = json.dumps({"token": "expired-api-key", "expires_at": past_time})
|
||||
mock_new_data = {
|
||||
"token": "new-api-key",
|
||||
"expires_at": (datetime.now() + timedelta(hours=1)).timestamp(),
|
||||
|
|
@ -128,7 +155,7 @@ class TestGitHubCopilotAuthenticator:
|
|||
with (
|
||||
patch("builtins.open", mock_open(read_data=mock_expired_data)),
|
||||
patch.object(authenticator, "_refresh_api_key", return_value=mock_new_data),
|
||||
patch("json.dump") as mock_json_dump,
|
||||
patch("json.dump"),
|
||||
):
|
||||
api_key = authenticator.get_api_key()
|
||||
assert api_key == "new-api-key"
|
||||
|
|
@ -217,20 +244,14 @@ class TestGitHubCopilotAuthenticator:
|
|||
mock_token = "mock-access-token"
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
authenticator, "_get_device_code", return_value=mock_device_code_data
|
||||
),
|
||||
patch.object(
|
||||
authenticator, "_poll_for_access_token", return_value=mock_token
|
||||
),
|
||||
patch.object(authenticator, "_get_device_code", return_value=mock_device_code_data),
|
||||
patch.object(authenticator, "_poll_for_access_token", return_value=mock_token),
|
||||
patch("builtins.print") as mock_print,
|
||||
):
|
||||
result = authenticator._login()
|
||||
assert result == mock_token
|
||||
authenticator._get_device_code.assert_called_once()
|
||||
authenticator._poll_for_access_token.assert_called_once_with(
|
||||
"mock-device-code"
|
||||
)
|
||||
authenticator._poll_for_access_token.assert_called_once_with("mock-device-code")
|
||||
mock_print.assert_called_once()
|
||||
|
||||
def test_get_api_base_from_file(self, authenticator):
|
||||
|
|
@ -255,8 +276,10 @@ class TestGitHubCopilotAuthenticator:
|
|||
"user_code": "UC",
|
||||
"verification_uri": "https://example.com",
|
||||
}
|
||||
with patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}), \
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client):
|
||||
with (
|
||||
patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}),
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client),
|
||||
):
|
||||
authenticator._get_device_code()
|
||||
assert mock_client.post.call_args[0][0] == custom_url
|
||||
|
||||
|
|
@ -269,8 +292,10 @@ class TestGitHubCopilotAuthenticator:
|
|||
"user_code": "UC",
|
||||
"verification_uri": "https://example.com",
|
||||
}
|
||||
with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client):
|
||||
with (
|
||||
patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}),
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client),
|
||||
):
|
||||
authenticator._get_device_code()
|
||||
assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id
|
||||
|
||||
|
|
@ -279,9 +304,11 @@ class TestGitHubCopilotAuthenticator:
|
|||
mock_client, mock_response = mock_http_client
|
||||
custom_url = "https://custom.example.com/token"
|
||||
mock_response.json.return_value = {"access_token": "tok"}
|
||||
with patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}), \
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
|
||||
patch("time.sleep"):
|
||||
with (
|
||||
patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}),
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client),
|
||||
patch("time.sleep"),
|
||||
):
|
||||
authenticator._poll_for_access_token("dc")
|
||||
assert mock_client.post.call_args[0][0] == custom_url
|
||||
|
||||
|
|
@ -290,9 +317,11 @@ class TestGitHubCopilotAuthenticator:
|
|||
mock_client, mock_response = mock_http_client
|
||||
custom_id = "custom_client_id"
|
||||
mock_response.json.return_value = {"access_token": "tok"}
|
||||
with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
|
||||
patch("time.sleep"):
|
||||
with (
|
||||
patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}),
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client),
|
||||
patch("time.sleep"),
|
||||
):
|
||||
authenticator._poll_for_access_token("dc")
|
||||
assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id
|
||||
|
||||
|
|
@ -301,9 +330,50 @@ class TestGitHubCopilotAuthenticator:
|
|||
mock_client, mock_response = mock_http_client
|
||||
custom_url = "https://custom.example.com/api-key"
|
||||
mock_response.json.return_value = {"token": "api-tok", "expires_at": 9999999999}
|
||||
with patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}), \
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \
|
||||
patch.object(authenticator, "get_access_token", return_value="access-tok"):
|
||||
with (
|
||||
patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}),
|
||||
patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client),
|
||||
patch.object(authenticator, "get_access_token", return_value="access-tok"),
|
||||
):
|
||||
authenticator._refresh_api_key()
|
||||
assert mock_client.get.call_args[0][0] == custom_url
|
||||
|
||||
def test_get_api_key_fallback_refresh_missing_token(self, authenticator):
|
||||
"""Test fallback flow when refreshed API key is missing token."""
|
||||
with (
|
||||
patch("builtins.open", side_effect=OSError),
|
||||
patch.object(authenticator, "_refresh_api_key", return_value={}),
|
||||
):
|
||||
with pytest.raises(GetAPIKeyError):
|
||||
authenticator.get_api_key()
|
||||
|
||||
def test_get_api_key_fallback_refresh_error(self, authenticator):
|
||||
"""Test fallback flow when _refresh_api_key raises RefreshAPIKeyError."""
|
||||
with (
|
||||
patch("builtins.open", side_effect=OSError),
|
||||
patch.object(
|
||||
authenticator, "_refresh_api_key", side_effect=RefreshAPIKeyError(message="Error", status_code=401)
|
||||
),
|
||||
):
|
||||
with pytest.raises(GetAPIKeyError):
|
||||
authenticator.get_api_key()
|
||||
|
||||
def test_get_api_key_fallback_save_os_error(self, authenticator):
|
||||
"""Test fallback flow continues when saving API key raises OSError."""
|
||||
mock_new_data = {
|
||||
"token": "in-memory-token",
|
||||
"expires_at": (datetime.now() + timedelta(hours=1)).timestamp(),
|
||||
}
|
||||
with (
|
||||
patch("builtins.open", side_effect=[OSError, OSError]),
|
||||
patch.object(authenticator, "_refresh_api_key", return_value=mock_new_data),
|
||||
patch.object(authenticator, "_ensure_token_dir", side_effect=OSError),
|
||||
):
|
||||
api_key = authenticator.get_api_key()
|
||||
assert api_key == "in-memory-token"
|
||||
|
||||
def test_get_api_base_non_dict_or_missing(self, authenticator):
|
||||
"""Test get_api_base returns None for non-dict or missing endpoints."""
|
||||
with patch("builtins.open", mock_open(read_data=json.dumps({"endpoints": "invalid"}))):
|
||||
assert authenticator.get_api_base() is None
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue