fix(github-copilot): harden per-deployment credentials

This commit is contained in:
Jacky Lam 2026-08-27 20:09:16 +08:00
parent 28af05ec19
commit 2dbaacad2d
13 changed files with 194 additions and 99 deletions

View file

@ -63,13 +63,31 @@ class GithubCopilotConfig(OpenAIConfig):
message=str(e),
)
dynamic_api_base: Final = (
api_base
or authenticator.get_api_base()
or os.getenv("GITHUB_COPILOT_API_BASE")
or DEFAULT_GITHUB_COPILOT_API_BASE
authenticator.get_api_base() or os.getenv("GITHUB_COPILOT_API_BASE") or DEFAULT_GITHUB_COPILOT_API_BASE
)
return dynamic_api_base, dynamic_api_key, custom_llm_provider
def resolve_request_credentials(
self,
model: str,
api_key: str | None,
headers: Mapping[str, str] | None,
litellm_params: Mapping[str, object] | None,
) -> tuple[str, dict[str, str]]: # mutable-ok: OpenAI request headers are merged by callers
authenticator: Final = get_authenticator_for_litellm_params(
default_authenticator=self.authenticator,
litellm_params=litellm_params,
)
try:
resolved_api_key: Final = api_key or authenticator.get_api_key()
except GetAPIKeyError as e:
raise AuthenticationError(
model=model,
llm_provider="github_copilot",
message=str(e),
)
return resolved_api_key, {**get_copilot_default_headers(resolved_api_key), **dict(headers or {})}
def _transform_messages(
self,
messages,
@ -106,22 +124,21 @@ class GithubCopilotConfig(OpenAIConfig):
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
# Get base headers from parent
validated_headers = super().validate_environment(
headers, model, messages, optional_params, litellm_params, api_key, api_base
copilot_api_key, copilot_headers = self.resolve_request_credentials(
model=model,
api_key=api_key,
headers=headers,
litellm_params=litellm_params,
)
validated_headers = super().validate_environment(
copilot_headers,
model,
messages,
optional_params,
litellm_params,
copilot_api_key,
api_base,
)
# Add Copilot-specific headers (editor-version, user-agent, etc.)
try:
authenticator: Final = get_authenticator_for_litellm_params(
default_authenticator=self.authenticator,
litellm_params=litellm_params,
)
copilot_api_key: Final = api_key or authenticator.get_api_key()
copilot_headers: Final = get_copilot_default_headers(copilot_api_key)
validated_headers = {**copilot_headers, **validated_headers}
except GetAPIKeyError:
pass # Will be handled later in the request flow
# Add X-Initiator header based on message roles
initiator: Final = self._determine_initiator(messages)

View file

@ -17,11 +17,6 @@ API_VERSION: Final = "2025-04-01"
DEFAULT_GITHUB_COPILOT_API_BASE: Final = "https://api.githubcopilot.com"
def is_default_copilot_api_base(api_base: str | None) -> bool:
"""Return whether ``api_base`` is absent or LiteLLM's generic Copilot base."""
return api_base is None or api_base.rstrip("/") == DEFAULT_GITHUB_COPILOT_API_BASE
class GithubCopilotError(BaseLLMException):
def __init__(
self,

View file

@ -24,7 +24,6 @@ from ..common_utils import (
DEFAULT_GITHUB_COPILOT_API_BASE,
GetAPIKeyError,
get_copilot_default_headers,
is_default_copilot_api_base,
)
if TYPE_CHECKING:
@ -106,16 +105,8 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig):
default_authenticator=self.authenticator,
litellm_params=litellm_params,
)
authenticated_api_base: Final = authenticator.get_api_base()
# embedding() can pre-populate api_base with the generic Copilot host;
# prefer the endpoint bound to the selected account in that case.
effective_api_base: Final = (
authenticated_api_base
if authenticated_api_base and is_default_copilot_api_base(api_base)
else api_base
or authenticated_api_base
or os.getenv("GITHUB_COPILOT_API_BASE")
or DEFAULT_GITHUB_COPILOT_API_BASE
authenticator.get_api_base() or os.getenv("GITHUB_COPILOT_API_BASE") or DEFAULT_GITHUB_COPILOT_API_BASE
).rstrip("/")
# Return the embeddings endpoint

View file

@ -30,7 +30,6 @@ from ..common_utils import (
DEFAULT_GITHUB_COPILOT_API_BASE,
GetAPIKeyError,
get_copilot_default_headers,
is_default_copilot_api_base,
)
if TYPE_CHECKING:
@ -264,18 +263,8 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig):
default_authenticator=self.authenticator,
litellm_params=litellm_params,
)
authenticated_api_base: Final = authenticator.get_api_base()
# responses() pre-populates api_base with the generic Copilot host. That
# must not hide the account-specific endpoint returned alongside this
# token (for example, the business Copilot host). Preserve genuinely
# custom bases for backwards compatibility.
effective_api_base: Final = (
authenticated_api_base
if authenticated_api_base and is_default_copilot_api_base(api_base)
else api_base
or authenticated_api_base
or os.getenv("GITHUB_COPILOT_API_BASE")
or DEFAULT_GITHUB_COPILOT_API_BASE
authenticator.get_api_base() or os.getenv("GITHUB_COPILOT_API_BASE") or DEFAULT_GITHUB_COPILOT_API_BASE
).rstrip("/")
# Return the responses endpoint

View file

@ -2494,7 +2494,6 @@ def _complete_custom_openai(
acompletion: Final = ctx.acompletion
api_base = ctx.api_base
api_key = ctx.api_key
explicit_api_key: Final = ctx.api_key
client: Final = _dispatch_client_http(ctx)
custom_llm_provider: Final = ctx.custom_llm_provider
custom_prompt_dict: Final = ctx.custom_prompt_dict
@ -2535,24 +2534,17 @@ def _complete_custom_openai(
# Add GitHub Copilot headers (same as /responses endpoint does)
if custom_llm_provider == "github_copilot":
from litellm.llms.github_copilot.authenticator import (
Authenticator,
get_authenticator_for_litellm_params,
)
from litellm.llms.github_copilot.common_utils import (
get_copilot_default_headers,
)
from litellm.llms.github_copilot.chat.transformation import GithubCopilotConfig
copilot_auth: Final = get_authenticator_for_litellm_params(
default_authenticator=Authenticator(),
copilot_config: Final = (
provider_config if isinstance(provider_config, GithubCopilotConfig) else GithubCopilotConfig()
)
api_key, extra_headers = copilot_config.resolve_request_credentials(
model=model,
api_key=ctx.api_key,
headers=extra_headers,
litellm_params=litellm_params,
)
copilot_api_key: Final = explicit_api_key or copilot_auth.get_api_key()
api_key = copilot_api_key
copilot_headers: Final = get_copilot_default_headers(copilot_api_key)
if extra_headers:
copilot_headers.update(extra_headers)
extra_headers = copilot_headers
if extra_headers is not None:
optional_params["extra_headers"] = extra_headers
@ -6168,12 +6160,13 @@ def embedding(
non_default_params: Final = {
k: v for k, v in kwargs.items() if k not in default_params
} # model-specific params - pass them straight to the model/provider
provider_litellm_params: Final = GenericLiteLLMParams(
github_copilot_token_dir=(
kwargs["github_copilot_token_dir"]
if "github_copilot_token_dir" in kwargs and isinstance(kwargs["github_copilot_token_dir"], str)
else None
)
github_copilot_token_dir: Final = (
kwargs["github_copilot_token_dir"]
if "github_copilot_token_dir" in kwargs and isinstance(kwargs["github_copilot_token_dir"], str)
else None
)
provider_litellm_params: Final = GenericLiteLLMParams.model_validate(
{"github_copilot_token_dir": github_copilot_token_dir} if github_copilot_token_dir is not None else {}
)
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(

View file

@ -1431,6 +1431,22 @@ class ModelManagementAuthChecks:
Common auth checks for model management endpoints
"""
@staticmethod
def require_proxy_admin_for_github_copilot_token_dir(
model_params: Deployment | updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
) -> Literal[True]:
litellm_params: Final = model_params.litellm_params
has_token_dir: Final = (
litellm_params is not None and "github_copilot_token_dir" in litellm_params.model_fields_set
)
if has_token_dir and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403,
detail={"error": "Only proxy admins can configure GitHub Copilot token directories."},
)
return True
@staticmethod
def can_user_make_team_model_call(
team_id: str,
@ -1461,6 +1477,10 @@ class ModelManagementAuthChecks:
prisma_client: PrismaClient,
premium_user: bool,
) -> Literal[True]:
ModelManagementAuthChecks.require_proxy_admin_for_github_copilot_token_dir(
model_params=model_params,
user_api_key_dict=user_api_key_dict,
)
if model_params.model_info is None or model_params.model_info.team_id is None:
return True
if model_params.model_info.team_id is not None and premium_user is not True:
@ -1496,6 +1516,10 @@ class ModelManagementAuthChecks:
premium_user: bool,
allow_missing_team: bool = False,
) -> Literal[True]:
ModelManagementAuthChecks.require_proxy_admin_for_github_copilot_token_dir(
model_params=model_params,
user_api_key_dict=user_api_key_dict,
)
## Check team model auth
if model_params.model_info is not None and model_params.model_info.team_id is not None:
team_obj_row: Final = await _repo_team_table(prisma_client).find_unique(

View file

@ -233,7 +233,6 @@ class CredentialLiteLLMParams(BaseModel):
api_key: str | None = None
api_base: str | None = None
api_version: str | None = None
github_copilot_token_dir: str | None = None
## AZURE OAUTH ##
# Without this field, ``get_deployment_credentials_with_provider``
# round-trips ``litellm_params`` through a strict Pydantic dump and

View file

@ -3,7 +3,7 @@
"limit": 744
},
"TQ002": {
"limit": 742
"limit": 741
},
"TQ003": {
"limit": 62

View file

@ -79,10 +79,13 @@ class TestGetLitellmParamsKwargsExtraction:
assert get_non_default_completion_params({"github_copilot_token_dir": token_dir}) == {}
assert "github_copilot_token_dir" in all_litellm_params
normalized = CredentialLiteLLMParams.model_validate(
GenericLiteLLMParams(github_copilot_token_dir=token_dir).model_dump(exclude_none=True)
internal_params = GenericLiteLLMParams.model_validate(
{"github_copilot_token_dir": token_dir}
).model_dump(exclude_none=True)
assert normalized["github_copilot_token_dir"] == token_dir
public_credentials = CredentialLiteLLMParams.model_validate(internal_params).model_dump(exclude_none=True)
assert internal_params["github_copilot_token_dir"] == token_dir
assert "github_copilot_token_dir" not in public_credentials
class TestGetLitellmParamsBaseModel:

View file

@ -2,8 +2,11 @@ import json
from datetime import datetime, timedelta
from unittest.mock import MagicMock
import httpx
import pytest
import respx
import litellm
from litellm.exceptions import AuthenticationError
from litellm.llms.github_copilot.common_utils import GetAPIKeyError
from litellm.llms.github_copilot.embedding.transformation import (
@ -84,8 +87,8 @@ 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
# Caller-controlled bases must not receive the Copilot bearer token.
config.authenticator.get_api_base.return_value = "https://api.enterprise.githubcopilot.com"
url = config.get_complete_url(
api_base="https://custom.api.com",
api_key=None,
@ -93,7 +96,7 @@ def test_github_copilot_embedding_config_get_complete_url():
optional_params={},
litellm_params={},
)
assert url == "https://custom.api.com/embeddings"
assert url == "https://api.enterprise.githubcopilot.com/embeddings"
def test_github_copilot_embedding_uses_per_deployment_token_directory(tmp_path):
@ -119,7 +122,7 @@ def test_github_copilot_embedding_uses_per_deployment_token_directory(tmp_path):
litellm_params=params,
)
url = config.get_complete_url(
api_base="https://api.githubcopilot.com/",
api_base="https://attacker.example",
api_key=None,
model="github_copilot/text-embedding-3-small",
optional_params={},
@ -130,6 +133,42 @@ def test_github_copilot_embedding_uses_per_deployment_token_directory(tmp_path):
assert url == "https://embedding-account.example/embeddings"
@respx.mock
def test_embedding_forwards_selected_account_credentials(tmp_path):
token_dir = tmp_path / "embedding-account"
token_dir.mkdir()
(token_dir / "api-key.json").write_text(
json.dumps(
{
"token": "embedding-account-key",
"expires_at": (datetime.now() + timedelta(hours=1)).timestamp(),
"endpoints": {"api": "https://embedding-account.example"},
}
)
)
route = respx.post("https://embedding-account.example/embeddings").mock(
return_value=httpx.Response(
200,
json={
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
},
)
)
response = litellm.embedding(
model="github_copilot/text-embedding-3-small",
input=["hello"],
github_copilot_token_dir=str(token_dir),
)
assert route.called
assert route.calls.last.request.headers["authorization"] == "Bearer embedding-account-key"
assert response.data[0]["embedding"] == [0.1, 0.2]
def test_github_copilot_embedding_config_transform_request():
"""Test the GitHub Copilot embedding request transformation."""
config = GithubCopilotEmbeddingConfig()

View file

@ -67,9 +67,9 @@ class TestGithubCopilotResponsesAPITransformation:
f"Expected GitHub Copilot responses endpoint, got {url}"
)
# Test with custom api_base (overrides authenticator)
# Caller-controlled bases must not receive the Copilot bearer token.
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}"
assert custom_url == "https://api.individual.githubcopilot.com/responses"
# The generic base injected by responses() must not override the base
# bound to the selected account's token.
@ -141,7 +141,7 @@ class TestGithubCopilotResponsesAPITransformation:
headers = config.validate_environment(headers={}, model="gpt-5.3-codex", litellm_params=params)
url = config.get_complete_url(
api_base="https://api.githubcopilot.com/",
api_base="https://attacker.example",
litellm_params=dict(params),
)

View file

@ -4,6 +4,7 @@ from unittest.mock import MagicMock, patch
import httpx
import pytest
import respx
import litellm
@ -71,14 +72,14 @@ def test_github_copilot_config_resolves_per_deployment_token_directory(tmp_path)
config = GithubCopilotConfig()
resolved_a = config._get_openai_compatible_provider_info(
model="github_copilot/gpt-4",
api_base=None,
api_base="https://attacker.example",
api_key=None,
custom_llm_provider="github_copilot",
litellm_params={"github_copilot_token_dir": str(account_a)},
)
resolved_b = config._get_openai_compatible_provider_info(
model="github_copilot/gpt-4",
api_base=None,
api_base="https://attacker.example",
api_key=None,
custom_llm_provider="github_copilot",
litellm_params={"github_copilot_token_dir": str(account_b)},
@ -184,11 +185,8 @@ def test_completion_github_copilot_mock_response(
assert kwargs.get("messages") == messages
@patch("litellm.main.openai_chat_completions.completion")
@patch("litellm.llms.openai.openai.OpenAIChatCompletion.completion")
def test_completion_github_copilot_uses_selected_account_directory(
mock_class_completion, mock_instance_completion, tmp_path, monkeypatch
):
@respx.mock
def test_completion_github_copilot_uses_selected_account_directory(tmp_path, monkeypatch):
monkeypatch.delenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", raising=False)
token_dir = tmp_path / "selected-account"
token_dir.mkdir()
@ -201,23 +199,35 @@ def test_completion_github_copilot_uses_selected_account_directory(
}
)
)
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
mock_class_completion.return_value = mock_response
mock_instance_completion.return_value = mock_response
route = respx.post("https://selected-account.example/chat/completions").mock(
return_value=httpx.Response(
200,
json={
"id": "chatcmpl-selected-account",
"object": "chat.completion",
"created": 1,
"model": "gpt-4",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "selected-account-ok"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
},
)
)
completion(
response = completion(
model="github_copilot/gpt-4",
messages=[{"role": "user", "content": "Hello"}],
github_copilot_token_dir=str(token_dir),
)
invoked = [mock for mock in (mock_class_completion, mock_instance_completion) if mock.called]
assert len(invoked) == 1
_, kwargs = invoked[0].call_args
assert kwargs["api_base"] == "https://selected-account.example"
assert kwargs["api_key"] == "selected-account-key"
assert kwargs["optional_params"]["extra_headers"]["Authorization"] == "Bearer selected-account-key"
assert route.called
assert route.calls.last.request.headers["authorization"] == "Bearer selected-account-key"
assert response.choices[0].message.content == "selected-account-ok"
def test_transform_messages_disable_copilot_system_to_assistant(monkeypatch):

View file

@ -181,6 +181,40 @@ class TestModelManagementAuthChecks:
)
assert result is True
@pytest.mark.asyncio
async def test_team_admin_cannot_add_github_copilot_token_directory(self):
model_params = Deployment(
model_name="copilot-pool",
litellm_params=LiteLLM_Params(
model="github_copilot/gpt-5.3-codex",
github_copilot_token_dir="/server/copilot/account-a",
),
model_info={"team_id": "test_team"},
)
with pytest.raises(Exception, match="Only proxy admins can configure GitHub Copilot token directories"):
await ModelManagementAuthChecks.can_user_make_model_call(
model_params=model_params,
user_api_key_dict=self.team_admin_user,
prisma_client=MockPrismaClient(team_exists=True),
premium_user=True,
)
@pytest.mark.asyncio
async def test_team_admin_cannot_patch_github_copilot_token_directory(self):
model_params = updateDeployment(
litellm_params={"github_copilot_token_dir": "/server/copilot/account-a"},
model_info={"team_id": "test_team"},
)
with pytest.raises(Exception, match="Only proxy admins can configure GitHub Copilot token directories"):
await ModelManagementAuthChecks.allow_team_model_action(
model_params=model_params,
user_api_key_dict=self.team_admin_user,
prisma_client=MockPrismaClient(team_exists=True),
premium_user=True,
)
@pytest.mark.asyncio
async def test_allow_team_model_action_non_premium_fails(self):
"""Test team model action fails for non-premium users"""
@ -227,7 +261,8 @@ class TestModelManagementAuthChecks:
model_params = Deployment(
model_name="test_model",
litellm_params=LiteLLM_Params(
model="test_model",
model="github_copilot/gpt-5.3-codex",
github_copilot_token_dir="/server/copilot/account-a",
),
model_info={"team_id": "test_team"},
)