From 8d384cf651ba982321b22b8b134b17e664244840 Mon Sep 17 00:00:00 2001 From: codgician <15964984+codgician@users.noreply.github.com> Date: Tue, 28 Jul 2026 13:20:33 +0800 Subject: [PATCH] refactor(github-copilot): use unified OAuth token flow --- litellm/llms/github_copilot/authenticator.py | 125 ++---------------- .../github_copilot/chat/transformation.py | 12 +- litellm/llms/github_copilot/common_utils.py | 36 ++--- .../embedding/transformation.py | 9 +- .../responses/transformation.py | 9 +- ...github_copilot_embedding_transformation.py | 10 +- ...github_copilot_responses_transformation.py | 10 +- .../test_github_copilot_authenticator.py | 120 +++-------------- .../test_github_copilot_transformation.py | 43 +++--- 9 files changed, 74 insertions(+), 300 deletions(-) diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index 17ef4d6e7c8..27a4651f154 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -1,8 +1,7 @@ import json import os import time -from datetime import datetime -from typing import Any, Final +from typing import Final import httpx @@ -10,23 +9,16 @@ from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import _get_httpx_client from .common_utils import ( - APIKeyExpiredError, GetAccessTokenError, GetAPIKeyError, GetDeviceCodeError, - RefreshAPIKeyError, - get_copilot_default_headers, + get_copilot_auth_headers, ) # Constants (default values — overridable via environment variables at call time) DEFAULT_GITHUB_CLIENT_ID: Final = "Iv1.b507a08c87ecfe98" DEFAULT_GITHUB_DEVICE_CODE_URL: Final = "https://github.com/login/device/code" DEFAULT_GITHUB_ACCESS_TOKEN_URL: Final = "https://github.com/login/oauth/access_token" -DEFAULT_GITHUB_API_KEY_URL: Final = "https://api.github.com/copilot_internal/v2/token" - - -def _use_oauth_token() -> bool: - return os.getenv("GITHUB_COPILOT_USE_OAUTH_TOKEN", "").strip().lower() in {"1", "true", "yes", "on"} class Authenticator: @@ -41,7 +33,6 @@ class Authenticator: 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")) self._ensure_token_dir() def get_access_token(self) -> str: @@ -72,7 +63,7 @@ class Authenticator: except OSError: verbose_logger.error("Error saving access token to file") return access_token - except (GetDeviceCodeError, GetAccessTokenError, RefreshAPIKeyError) as e: + except (GetDeviceCodeError, GetAccessTokenError) as e: verbose_logger.warning("Failed attempt %s: %s", attempt + 1, e) continue @@ -82,122 +73,24 @@ class Authenticator: ) def get_api_key(self) -> str: - """ - Get the API key, refreshing if necessary. - - Returns: - str: The GitHub Copilot API key. - - Raises: - GetAPIKeyError: If unable to obtain an API key. - """ - if _use_oauth_token(): + try: return self.get_access_token() - 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") - else: - verbose_logger.warning("API key expired, refreshing") - raise APIKeyExpiredError( - message="API key expired", - status_code=401, - ) - except OSError: - verbose_logger.warning("No API key file found or error opening file") - except (json.JSONDecodeError, KeyError) as e: - verbose_logger.warning("Error reading API key from file: %s", e) - except APIKeyExpiredError: - 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 - 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) + except GetAccessTokenError as 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}", + message=f"Failed to get OAuth access token: {str(e)}", status_code=401, ) def get_api_base(self) -> str | None: - """ - Get the API endpoint from the api-key.json file. - - Returns: - Optional[str]: The GitHub Copilot API endpoint, or None if not found. - """ - if _use_oauth_token(): - return os.getenv("GITHUB_COPILOT_API_BASE") - 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 - 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]: - """ - Refresh the API key using the 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) - 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) - response.raise_for_status() - - response_json = response.json() - - if "token" in response_json: - return response_json - else: - verbose_logger.warning("API key response missing token: %s", response_json) - except httpx.HTTPStatusError as e: - verbose_logger.error("HTTP error refreshing API key (attempt %s/%s): %s", attempt + 1, max_retries, e) - except Exception as e: - verbose_logger.error("Unexpected error refreshing API key: %s", e) - - raise RefreshAPIKeyError( - message="Failed to refresh API key after maximum retries", - status_code=401, - ) + return os.getenv("GITHUB_COPILOT_API_BASE") 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) - def _get_github_headers(self, access_token: str | None = None) -> dict[str, str]: - return get_copilot_default_headers(access_token=access_token) + def _get_github_headers(self) -> dict[str, str]: + return get_copilot_auth_headers() def _get_device_code(self) -> dict[str, str]: """ diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 169b9a037a5..66a239e598b 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -1,5 +1,4 @@ import json -import os from typing import TYPE_CHECKING, Any, Final import httpx @@ -31,13 +30,8 @@ class GithubCopilotConfig(OpenAIConfig): super().__init__() self.authenticator = Authenticator() - def api_base_without_login(self, api_base: str | None = None) -> str: - return ( - api_base - or self.authenticator.get_api_base() - or os.getenv("GITHUB_COPILOT_API_BASE") - or DEFAULT_GITHUB_COPILOT_API_BASE - ) + def api_base_without_login(self) -> str: + return self.authenticator.get_api_base() or DEFAULT_GITHUB_COPILOT_API_BASE def _get_openai_compatible_provider_info( self, @@ -46,7 +40,7 @@ class GithubCopilotConfig(OpenAIConfig): api_key: str | None, custom_llm_provider: str, ) -> tuple[str | None, str | None, str]: - dynamic_api_base: Final = self.api_base_without_login(api_base) + dynamic_api_base: Final = self.api_base_without_login() try: dynamic_api_key: Final = self.authenticator.get_api_key() except GetAPIKeyError as e: diff --git a/litellm/llms/github_copilot/common_utils.py b/litellm/llms/github_copilot/common_utils.py index 2ea87671466..ad5d05184b2 100644 --- a/litellm/llms/github_copilot/common_utils.py +++ b/litellm/llms/github_copilot/common_utils.py @@ -20,13 +20,15 @@ DEFAULT_COPILOT_EDITOR_VERSION: Final = "vscode/1.115.0" DEFAULT_COPILOT_EDITOR_PLUGIN_VERSION: Final = EDITOR_PLUGIN_VERSION DEFAULT_COPILOT_USER_AGENT: Final = USER_AGENT -_COPILOT_HEADER_CONFIG = ( +_COPILOT_AUTH_HEADER_CONFIG = ( ("accept", "GITHUB_COPILOT_ACCEPT", "application/json"), ("content-type", "GITHUB_COPILOT_CONTENT_TYPE", "application/json"), ("copilot-integration-id", "GITHUB_COPILOT_INTEGRATION_ID", DEFAULT_COPILOT_INTEGRATION_ID), ("editor-version", "GITHUB_COPILOT_EDITOR_VERSION", DEFAULT_COPILOT_EDITOR_VERSION), ("editor-plugin-version", "GITHUB_COPILOT_EDITOR_PLUGIN_VERSION", DEFAULT_COPILOT_EDITOR_PLUGIN_VERSION), ("user-agent", "GITHUB_COPILOT_USER_AGENT", DEFAULT_COPILOT_USER_AGENT), +) +_COPILOT_REQUEST_HEADER_CONFIG = _COPILOT_AUTH_HEADER_CONFIG + ( ("openai-intent", "GITHUB_COPILOT_OPENAI_INTENT", None), ("x-github-api-version", "GITHUB_COPILOT_API_VERSION", None), ( @@ -65,14 +67,6 @@ class GetAccessTokenError(GithubCopilotError): pass -class APIKeyExpiredError(GithubCopilotError): - pass - - -class RefreshAPIKeyError(GithubCopilotError): - pass - - class GetAPIKeyError(GithubCopilotError): pass @@ -84,16 +78,22 @@ def _get_copilot_header_value(environment_variable: str, default: str | None) -> return value or None -def get_copilot_default_headers( - api_key: str | None = None, - *, - access_token: str | None = None, +def _get_configured_copilot_headers( + config: tuple[tuple[str, str, str | None], ...], ) -> dict[str, str]: - configured_headers = { + return { header: value - for header, environment_variable, default in _COPILOT_HEADER_CONFIG + for header, environment_variable, default in config if (value := _get_copilot_header_value(environment_variable, default)) is not None } - authorization = f"token {access_token}" if access_token else f"Bearer {api_key}" if api_key else None - authorization_header = {"Authorization": authorization} if authorization else {} - return {**configured_headers, **authorization_header} + + +def get_copilot_auth_headers() -> dict[str, str]: + return _get_configured_copilot_headers(_COPILOT_AUTH_HEADER_CONFIG) + + +def get_copilot_default_headers(api_key: str) -> dict[str, str]: + return { + **_get_configured_copilot_headers(_COPILOT_REQUEST_HEADER_CONFIG), + "Authorization": f"Bearer {api_key}", + } diff --git a/litellm/llms/github_copilot/embedding/transformation.py b/litellm/llms/github_copilot/embedding/transformation.py index 7ea7a89b4ca..e5e35059ce2 100644 --- a/litellm/llms/github_copilot/embedding/transformation.py +++ b/litellm/llms/github_copilot/embedding/transformation.py @@ -7,7 +7,6 @@ Implementation based on analysis of the copilot-api project by caozhiyuan: https://github.com/caozhiyuan/copilot-api """ -import os from typing import TYPE_CHECKING, Any, Final import httpx @@ -98,13 +97,7 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): """ Get the complete URL for GitHub Copilot Embedding API endpoint. """ - # Use provided api_base or fall back to authenticator's base or default - effective_api_base = ( - api_base - or self.authenticator.get_api_base() - or os.getenv("GITHUB_COPILOT_API_BASE") - or DEFAULT_GITHUB_COPILOT_API_BASE - ) + effective_api_base = self.authenticator.get_api_base() or DEFAULT_GITHUB_COPILOT_API_BASE # Remove trailing slashes effective_api_base = effective_api_base.rstrip("/") diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 8b85b668cba..51873b3203d 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -8,7 +8,6 @@ Implementation based on analysis of the copilot-api project by caozhiyuan: https://github.com/caozhiyuan/copilot-api """ -import os from typing import TYPE_CHECKING, Any, Final import litellm @@ -249,13 +248,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): """ Get the complete URL for GitHub Copilot Responses API endpoint. """ - # Use provided api_base or fall back to authenticator's base or default - effective_api_base = ( - api_base - or self.authenticator.get_api_base() - or os.getenv("GITHUB_COPILOT_API_BASE") - or DEFAULT_GITHUB_COPILOT_API_BASE - ) + effective_api_base = self.authenticator.get_api_base() or DEFAULT_GITHUB_COPILOT_API_BASE # Remove trailing slashes effective_api_base = effective_api_base.rstrip("/") diff --git a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py index c5b1aba2599..40f9e06d4b8 100644 --- a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py +++ b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py @@ -73,9 +73,7 @@ def test_github_copilot_embedding_config_get_complete_url(): assert url == "https://api.githubcopilot.com/embeddings" # Test with custom API base from authenticator - config.authenticator.get_api_base.return_value = ( - "https://api.enterprise.githubcopilot.com" - ) + config.authenticator.get_api_base.return_value = "https://api.enterprise.githubcopilot.com" url = config.get_complete_url( api_base=None, api_key=None, @@ -85,16 +83,14 @@ def test_github_copilot_embedding_config_get_complete_url(): ) assert url == "https://api.enterprise.githubcopilot.com/embeddings" - # Test with custom API base from params - config.authenticator.get_api_base.return_value = None url = config.get_complete_url( - api_base="https://custom.api.com", + api_base="https://untrusted.example.com", api_key=None, model="github_copilot/text-embedding-3-small", optional_params={}, litellm_params={}, ) - assert url == "https://custom.api.com/embeddings" + assert url == "https://api.enterprise.githubcopilot.com/embeddings" def test_github_copilot_embedding_config_transform_request(): diff --git a/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py b/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py index 1ac4100626b..45f03c5cacb 100644 --- a/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py +++ b/tests/unit/llms/github_copilot/responses/test_github_copilot_responses_transformation.py @@ -64,13 +64,11 @@ class TestGithubCopilotResponsesAPITransformation: f"Expected GitHub Copilot responses endpoint, got {url}" ) - # Test with custom api_base (overrides authenticator) - custom_url = config.get_complete_url(api_base="https://custom.githubcopilot.com", litellm_params={}) - assert custom_url == "https://custom.githubcopilot.com/responses", f"Expected custom endpoint, got {custom_url}" + custom_url = config.get_complete_url(api_base="https://untrusted.example.com", litellm_params={}) + assert custom_url == "https://api.individual.githubcopilot.com/responses" - # Test with trailing slash - url_with_slash = config.get_complete_url(api_base="https://api.githubcopilot.com/", litellm_params={}) - assert url_with_slash == "https://api.githubcopilot.com/responses", "Should handle trailing slash" + url_with_slash = config.get_complete_url(api_base="https://untrusted.example.com/", litellm_params={}) + assert url_with_slash == "https://api.individual.githubcopilot.com/responses" @patch("litellm.llms.github_copilot.responses.transformation.Authenticator") def test_validate_environment_default_headers(self, mock_authenticator_class, monkeypatch): diff --git a/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py b/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py index c378402324d..7b42a774b43 100644 --- a/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py +++ b/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py @@ -1,18 +1,14 @@ -import json import os -import time -from datetime import datetime, timedelta from unittest.mock import MagicMock, mock_open, patch import pytest from litellm.llms.github_copilot.authenticator import Authenticator from litellm.llms.github_copilot.common_utils import ( - APIKeyExpiredError, GetAccessTokenError, GetAPIKeyError, GetDeviceCodeError, - RefreshAPIKeyError, + get_copilot_default_headers, ) @@ -45,7 +41,6 @@ class TestGitHubCopilotAuthenticator: 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() def test_ensure_token_dir(self): @@ -68,9 +63,6 @@ class TestGitHubCopilotAuthenticator: "user-agent": "GitHubCopilotChat/0.44.0", } - headers_with_token = authenticator._get_github_headers("test-token") - assert headers_with_token["Authorization"] == "token test-token" - def test_auth_requests_support_opencode_identity(self, authenticator, mock_http_client): mock_client, mock_response = mock_http_client mock_response.json.side_effect = ( @@ -87,7 +79,8 @@ class TestGitHubCopilotAuthenticator: "GITHUB_COPILOT_INTEGRATION_ID": "", "GITHUB_COPILOT_EDITOR_VERSION": "", "GITHUB_COPILOT_EDITOR_PLUGIN_VERSION": "", - "GITHUB_COPILOT_USE_OAUTH_TOKEN": "true", + "GITHUB_COPILOT_API_VERSION": "2026-06-01", + "GITHUB_COPILOT_OPENAI_INTENT": "conversation-edits", "GITHUB_COPILOT_API_BASE": "https://api.githubcopilot.com", } @@ -103,27 +96,34 @@ class TestGitHubCopilotAuthenticator: assert authenticator._poll_for_access_token("dc") == "opencode-oauth-token" assert authenticator.get_api_key() == "opencode-oauth-token" assert authenticator.get_api_base() == "https://api.githubcopilot.com" + request_headers = get_copilot_default_headers("opencode-oauth-token") - expected_headers = { + expected_auth_headers = { "accept": "application/json", "content-type": "application/json", "user-agent": "opencode/1.18.7", } assert mock_client.post.call_args_list[0].kwargs == { - "headers": expected_headers, + "headers": expected_auth_headers, "json": { "client_id": "Ov23li8tweQw6odWQebz", "scope": "read:user", }, } assert mock_client.post.call_args_list[1].kwargs == { - "headers": expected_headers, + "headers": expected_auth_headers, "json": { "client_id": "Ov23li8tweQw6odWQebz", "device_code": "dc", "grant_type": "urn:ietf:params:oauth:grant-type:device_code", }, } + assert request_headers == { + **expected_auth_headers, + "openai-intent": "conversation-edits", + "x-github-api-version": "2026-06-01", + "Authorization": "Bearer opencode-oauth-token", + } mock_client.get.assert_not_called() def test_get_access_token_from_file(self, authenticator): @@ -161,68 +161,14 @@ class TestGitHubCopilotAuthenticator: authenticator.get_access_token() assert authenticator._login.call_count == 3 - 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}) - - with patch("builtins.open", mock_open(read_data=mock_api_key_data)): - api_key = authenticator.get_api_key() - assert api_key == "mock-api-key" - - 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_new_data = { - "token": "new-api-key", - "expires_at": (datetime.now() + timedelta(hours=1)).timestamp(), - } - - 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, + def test_get_api_key_maps_access_token_failure(self, authenticator): + with patch.object( + authenticator, + "get_access_token", + side_effect=GetAccessTokenError(message="OAuth failed", status_code=401), ): - api_key = authenticator.get_api_key() - assert api_key == "new-api-key" - authenticator._refresh_api_key.assert_called_once() - - def test_refresh_api_key(self, authenticator, mock_http_client): - """Test refreshing an API key.""" - mock_client, mock_response = mock_http_client - mock_token = "mock-access-token" - mock_api_key_data = {"token": "new-api-key", "expires_at": 12345} - - with ( - patch.object(authenticator, "get_access_token", return_value=mock_token), - patch( - "litellm.llms.github_copilot.authenticator._get_httpx_client", - return_value=mock_client, - ), - patch.object(mock_response, "json", return_value=mock_api_key_data), - ): - result = authenticator._refresh_api_key() - assert result == mock_api_key_data - mock_client.get.assert_called_once() - authenticator.get_access_token.assert_called_once() - - def test_refresh_api_key_failure(self, authenticator, mock_http_client): - """Test failure to refresh an API key.""" - mock_client, mock_response = mock_http_client - mock_token = "mock-access-token" - - with ( - patch.object(authenticator, "get_access_token", return_value=mock_token), - patch( - "litellm.llms.github_copilot.authenticator._get_httpx_client", - return_value=mock_client, - ), - patch.object(mock_response, "json", return_value={}), - ): - with pytest.raises(RefreshAPIKeyError): - authenticator._refresh_api_key() - assert mock_client.get.call_count == 3 + with pytest.raises(GetAPIKeyError, match="Failed to get OAuth access token"): + authenticator.get_api_key() def test_get_device_code(self, authenticator, mock_http_client): """Test getting a device code.""" @@ -281,19 +227,6 @@ class TestGitHubCopilotAuthenticator: 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): - """Test retrieving the API base endpoint from a file.""" - mock_api_key_data = json.dumps( - { - "token": "mock-api-key", - "expires_at": (datetime.now() + timedelta(hours=1)).timestamp(), - "endpoints": {"api": "https://api.enterprise.githubcopilot.com"}, - } - ) - with patch("builtins.open", mock_open(read_data=mock_api_key_data)): - api_base = authenticator.get_api_base() - assert api_base == "https://api.enterprise.githubcopilot.com" - def test_get_device_code_with_custom_url(self, authenticator, mock_http_client): """GITHUB_COPILOT_DEVICE_CODE_URL env var must be used by _get_device_code at call time.""" mock_client, mock_response = mock_http_client @@ -351,16 +284,3 @@ class TestGitHubCopilotAuthenticator: ): authenticator._poll_for_access_token("dc") assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id - - def test_refresh_api_key_with_custom_url(self, authenticator, mock_http_client): - """GITHUB_COPILOT_API_KEY_URL env var must be used by _refresh_api_key at call time.""" - 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"), - ): - authenticator._refresh_api_key() - assert mock_client.get.call_args[0][0] == custom_url diff --git a/tests/unit/llms/github_copilot/test_github_copilot_transformation.py b/tests/unit/llms/github_copilot/test_github_copilot_transformation.py index b1d7e49fbac..da68a984f33 100644 --- a/tests/unit/llms/github_copilot/test_github_copilot_transformation.py +++ b/tests/unit/llms/github_copilot/test_github_copilot_transformation.py @@ -19,11 +19,9 @@ from litellm.exceptions import AuthenticationError from litellm.llms.github_copilot.authenticator import Authenticator from litellm.llms.github_copilot.chat.transformation import GithubCopilotConfig from litellm.llms.github_copilot.common_utils import ( - APIKeyExpiredError, GetAccessTokenError, GetAPIKeyError, GetDeviceCodeError, - RefreshAPIKeyError, ) @@ -37,9 +35,7 @@ def test_github_copilot_config_get_openai_compatible_provider_info(): config.authenticator = MagicMock() config.authenticator.get_api_key.return_value = mock_api_key # Test with dynamic endpoint - config.authenticator.get_api_base.return_value = ( - "https://api.enterprise.githubcopilot.com" - ) + config.authenticator.get_api_base.return_value = "https://api.enterprise.githubcopilot.com" # Test with default values model = "github_copilot/gpt-4" @@ -57,6 +53,13 @@ def test_github_copilot_config_get_openai_compatible_provider_info(): assert api_base == "https://api.enterprise.githubcopilot.com" assert dynamic_api_key == mock_api_key assert custom_llm_provider == "github_copilot" + api_base, _, _ = config._get_openai_compatible_provider_info( + model=model, + api_base="https://untrusted.example.com", + api_key=None, + custom_llm_provider="github_copilot", + ) + assert api_base == "https://api.enterprise.githubcopilot.com" # Test fallback to default if no dynamic endpoint config.authenticator.get_api_base.return_value = None @@ -158,25 +161,19 @@ def test_transform_messages_disable_copilot_system_to_assistant(monkeypatch): {"role": "system", "content": "System message."}, {"role": "user", "content": "User message."}, ] - out = config._transform_messages( - [m.copy() for m in messages], model="github_copilot/gpt-4" - ) + out = config._transform_messages([m.copy() for m in messages], model="github_copilot/gpt-4") assert out[0]["role"] == "assistant" assert out[1]["role"] == "user" # Case 2: Flag is True (conversion does not happen) litellm.disable_copilot_system_to_assistant = True - out = config._transform_messages( - [m.copy() for m in messages], model="github_copilot/gpt-4" - ) + out = config._transform_messages([m.copy() for m in messages], model="github_copilot/gpt-4") assert out[0]["role"] == "system" assert out[1]["role"] == "user" # Case 3: Flag is False again (conversion happens) litellm.disable_copilot_system_to_assistant = False - out = config._transform_messages( - [m.copy() for m in messages], model="github_copilot/gpt-4" - ) + out = config._transform_messages([m.copy() for m in messages], model="github_copilot/gpt-4") assert out[0]["role"] == "assistant" assert out[1]["role"] == "user" finally: @@ -374,7 +371,6 @@ def test_x_initiator_header_system_only_messages(): - def test_copilot_vision_request_header_with_image(): """Test that Copilot-Vision-Request header is added when messages contain images""" config = GithubCopilotConfig() @@ -699,13 +695,8 @@ class TestGithubCopilotTransformResponse: assert result.choices[0].message.tool_calls is not None assert len(result.choices[0].message.tool_calls) == 1 assert result.choices[0].message.tool_calls[0]["id"] == "toolu_01ABC" - assert ( - result.choices[0].message.tool_calls[0]["function"]["name"] == "get_weather" - ) - assert ( - '"Boston, MA"' - in result.choices[0].message.tool_calls[0]["function"]["arguments"] - ) + assert result.choices[0].message.tool_calls[0]["function"]["name"] == "get_weather" + assert '"Boston, MA"' in result.choices[0].message.tool_calls[0]["function"]["arguments"] def test_transform_response_anthropic_native_multiple_text_blocks(self): """All text blocks must be concatenated, not only the first.""" @@ -873,12 +864,8 @@ class TestGithubCopilotTransformParsedResponseDict: @patch("litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client") -@patch( - "litellm.llms.openai.openai.OpenAIChatCompletion.make_sync_openai_chat_completion_request" -) -def test_openai_handler_repairs_github_copilot_empty_choices( - mock_request, mock_get_client -): +@patch("litellm.llms.openai.openai.OpenAIChatCompletion.make_sync_openai_chat_completion_request") +def test_openai_handler_repairs_github_copilot_empty_choices(mock_request, mock_get_client): """ The OpenAI SDK handler calls convert_to_model_response_object directly on the SDK's parsed output, bypassing transform_response. convert raises APIError on