mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix(github-copilot): support direct OAuth tokens
This commit is contained in:
parent
e264b893bf
commit
29bae6e21a
2 changed files with 72 additions and 65 deletions
|
|
@ -25,6 +25,10 @@ DEFAULT_GITHUB_ACCESS_TOKEN_URL: Final = "https://github.com/login/oauth/access_
|
|||
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:
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the GitHub Copilot authenticator with configurable token paths."""
|
||||
|
|
@ -87,6 +91,8 @@ class Authenticator:
|
|||
Raises:
|
||||
GetAPIKeyError: If unable to obtain an API key.
|
||||
"""
|
||||
if _use_oauth_token():
|
||||
return self.get_access_token()
|
||||
try:
|
||||
with open(self.api_key_file, "r") as f:
|
||||
api_key_info = json.load(f)
|
||||
|
|
@ -136,6 +142,8 @@ class Authenticator:
|
|||
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)
|
||||
|
|
|
|||
|
|
@ -71,60 +71,60 @@ class TestGitHubCopilotAuthenticator:
|
|||
headers_with_token = authenticator._get_github_headers("test-token")
|
||||
assert headers_with_token["Authorization"] == "token test-token"
|
||||
|
||||
def test_auth_requests_use_custom_copilot_headers(self, authenticator, mock_http_client):
|
||||
def test_auth_requests_support_opencode_identity(self, authenticator, mock_http_client):
|
||||
mock_client, mock_response = mock_http_client
|
||||
mock_response.json.side_effect = (
|
||||
{
|
||||
"device_code": "dc",
|
||||
"user_code": "UC",
|
||||
"verification_uri": "https://example.com",
|
||||
"verification_uri": "https://github.com/login/device",
|
||||
},
|
||||
{"access_token": "access-token"},
|
||||
{"token": "api-token", "expires_at": 9999999999},
|
||||
{"access_token": "opencode-oauth-token"},
|
||||
)
|
||||
environment = {
|
||||
"GITHUB_COPILOT_ACCEPT": "application/vnd.github+json",
|
||||
"GITHUB_COPILOT_CONTENT_TYPE": "application/custom+json",
|
||||
"GITHUB_COPILOT_INTEGRATION_ID": "custom-integration",
|
||||
"GITHUB_COPILOT_EDITOR_VERSION": "custom-editor/1.0",
|
||||
"GITHUB_COPILOT_EDITOR_PLUGIN_VERSION": "custom-plugin/2.0",
|
||||
"GITHUB_COPILOT_USER_AGENT": "CustomAgent/2.0",
|
||||
"GITHUB_COPILOT_OPENAI_INTENT": "custom-intent",
|
||||
"GITHUB_COPILOT_API_VERSION": "2099-01-01",
|
||||
"GITHUB_COPILOT_USER_AGENT_LIBRARY_VERSION": "custom-library",
|
||||
"GITHUB_COPILOT_CLIENT_ID": "Ov23li8tweQw6odWQebz",
|
||||
"GITHUB_COPILOT_USER_AGENT": "opencode/1.18.7",
|
||||
"GITHUB_COPILOT_INTEGRATION_ID": "",
|
||||
"GITHUB_COPILOT_EDITOR_VERSION": "",
|
||||
"GITHUB_COPILOT_EDITOR_PLUGIN_VERSION": "",
|
||||
"GITHUB_COPILOT_USE_OAUTH_TOKEN": "true",
|
||||
"GITHUB_COPILOT_API_BASE": "https://api.githubcopilot.com",
|
||||
}
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, environment),
|
||||
patch.dict(os.environ, environment, clear=True),
|
||||
patch(
|
||||
"litellm.llms.github_copilot.authenticator._get_httpx_client",
|
||||
return_value=mock_client,
|
||||
),
|
||||
patch.object(authenticator, "get_access_token", return_value="github-token"),
|
||||
patch.object(authenticator, "get_access_token", return_value="opencode-oauth-token"),
|
||||
):
|
||||
authenticator._get_device_code()
|
||||
authenticator._poll_for_access_token("dc")
|
||||
authenticator._refresh_api_key()
|
||||
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 = (
|
||||
mock_client.post.call_args_list[0].kwargs["headers"],
|
||||
mock_client.post.call_args_list[1].kwargs["headers"],
|
||||
mock_client.get.call_args.kwargs["headers"],
|
||||
)
|
||||
for headers in request_headers:
|
||||
assert headers["accept"] == "application/vnd.github+json"
|
||||
assert headers["content-type"] == "application/custom+json"
|
||||
assert headers["copilot-integration-id"] == "custom-integration"
|
||||
assert headers["editor-version"] == "custom-editor/1.0"
|
||||
assert headers["editor-plugin-version"] == "custom-plugin/2.0"
|
||||
assert headers["user-agent"] == "CustomAgent/2.0"
|
||||
assert headers["openai-intent"] == "custom-intent"
|
||||
assert headers["x-github-api-version"] == "2099-01-01"
|
||||
assert headers["x-vscode-user-agent-library-version"] == "custom-library"
|
||||
|
||||
assert "Authorization" not in request_headers[0]
|
||||
assert "Authorization" not in request_headers[1]
|
||||
assert request_headers[2]["Authorization"] == "token github-token"
|
||||
expected_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,
|
||||
"json": {
|
||||
"client_id": "Ov23li8tweQw6odWQebz",
|
||||
"scope": "read:user",
|
||||
},
|
||||
}
|
||||
assert mock_client.post.call_args_list[1].kwargs == {
|
||||
"headers": expected_headers,
|
||||
"json": {
|
||||
"client_id": "Ov23li8tweQw6odWQebz",
|
||||
"device_code": "dc",
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:device_code",
|
||||
},
|
||||
}
|
||||
mock_client.get.assert_not_called()
|
||||
|
||||
def test_get_access_token_from_file(self, authenticator):
|
||||
"""Test retrieving an access token from a file."""
|
||||
|
|
@ -164,9 +164,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()
|
||||
|
|
@ -175,9 +173,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(),
|
||||
|
|
@ -275,20 +271,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):
|
||||
|
|
@ -313,8 +303,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
|
||||
|
||||
|
|
@ -327,8 +319,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
|
||||
|
||||
|
|
@ -337,9 +331,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
|
||||
|
||||
|
|
@ -348,9 +344,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
|
||||
|
||||
|
|
@ -359,9 +357,10 @@ 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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue