diff --git a/docs/my-website/docs/providers/github_copilot.md b/docs/my-website/docs/providers/github_copilot.md index e9fd3444f5f..414672572d2 100644 --- a/docs/my-website/docs/providers/github_copilot.md +++ b/docs/my-website/docs/providers/github_copilot.md @@ -20,14 +20,27 @@ https://docs.github.com/en/copilot ## Authentication -GitHub Copilot uses OAuth device flow for authentication. On first use, you'll be prompted to authenticate via GitHub: +GitHub Copilot uses OAuth Device Code flow for authentication. LiteLLM supports two authentication modes: -1. LiteLLM will display a device code and verification URL -2. Visit the URL and enter the code to authenticate -3. Your credentials will be stored locally for future use +- **LiteLLM Proxy**: Use named credentials created via the credential API or UI. The interactive device flow is handled by the proxy on your behalf — credentials are stored and reused automatically. +- **Python SDK**: Credentials are read from `~/.config/litellm/github_copilot/access-token` on disk (file-based, for backward compatibility). + +:::info + +If you hit a GitHub Copilot model without a configured credential (proxy) or access-token file (SDK), LiteLLM returns an `AuthenticationError` pointing to this page. + +::: ## Usage - LiteLLM Python SDK +### Setup (first time) + +Authenticate once by running the login command. This stores your access token in `~/.config/litellm/github_copilot/access-token` for future use. + +```bash showLineNumbers title="Authenticate via device flow (SDK only)" +litellm --login github_copilot +``` + ### Chat Completion ```python showLineNumbers title="GitHub Copilot Chat Completion" @@ -87,6 +100,8 @@ print(response) ## Usage - LiteLLM Proxy +The proxy requires a named credential. See [Credential-Based Authentication](#credential-based-authentication-proxy) below. + Add the following to your LiteLLM Proxy configuration file: ```yaml showLineNumbers title="config.yaml" @@ -94,16 +109,19 @@ model_list: - model_name: github_copilot/gpt-4 litellm_params: model: github_copilot/gpt-4 + litellm_credential_name: my-copilot # named credential (required for proxy) - model_name: github_copilot/gpt-5.1-codex model_info: mode: responses litellm_params: model: github_copilot/gpt-5.1-codex + litellm_credential_name: my-copilot - model_name: github_copilot/text-embedding-ada-002 model_info: mode: embedding litellm_params: model: github_copilot/text-embedding-ada-002 + litellm_credential_name: my-copilot ``` Start your LiteLLM Proxy server: @@ -170,18 +188,115 @@ curl http://localhost:4000/v1/chat/completions \ -## Getting Started +## Credential-Based Authentication (Proxy) -1. Ensure you have GitHub Copilot access (paid GitHub subscription required) -2. Run your first LiteLLM request - you'll be prompted to authenticate -3. Follow the device flow authentication process -4. Start making requests to GitHub Copilot through LiteLLM +The LiteLLM Proxy uses a stateless OAuth Device Code flow. Nothing is stored in the database until the GitHub token is successfully obtained — the client holds the `device_code` between API calls. + +You can complete the flow via the **LiteLLM UI** (Models → Credentials → Add Credential → GitHub Copilot) or with the curl steps below. + +### Step 1: Initiate the Device Code Flow + +```bash showLineNumbers title="Start GitHub OAuth" +curl -X POST http://localhost:4000/credentials/github_copilot/initiate \ + -H "Authorization: Bearer your-proxy-api-key" +``` + +Response: + +```json +{ + "device_code": "xxx", + "user_code": "ABCD-1234", + "verification_uri": "https://github.com/login/device", + "poll_interval_ms": 5000, + "expires_in": 900 +} +``` + +### Step 2: Authorize on GitHub + +Visit the `verification_uri` and enter the `user_code` to authorize LiteLLM. + +### Step 3: Poll for Completion + +Poll the status endpoint until the flow completes. Use `device_code` from step 1. +Use `poll_interval_ms` from the initiate response as your default polling interval. + +```bash showLineNumbers title="Check authorization status" +curl -X POST http://localhost:4000/credentials/github_copilot/status \ + -H "Authorization: Bearer your-proxy-api-key" \ + -H "Content-Type: application/json" \ + -d '{"device_code": "xxx"}' +``` + +Possible responses: + +```json title="Still waiting for user" +{"status": "pending"} +``` + +```json title="GitHub is being polled too fast — wait before retrying" +{"status": "pending", "retry_after_ms": 10000} +``` + +If `retry_after_ms` is present, you **must** wait that many milliseconds before calling `/status` again. Ignoring it causes GitHub to keep increasing the required interval. + +```json title="User authorized successfully" +{"status": "complete", "access_token": "ghu_xxx"} +``` + +```json title="Flow expired or denied" +{"status": "failed", "error": "The device code has expired."} +``` + +### Step 4: Store as a Named Credential + +Once you have the `access_token`, store it as a named credential: + +```bash showLineNumbers title="Save credential" +curl -X POST http://localhost:4000/credentials \ + -H "Authorization: Bearer your-proxy-api-key" \ + -H "Content-Type: application/json" \ + -d '{ + "credential_name": "my-copilot", + "credential_values": {"api_key": "ghu_xxx"}, + "credential_info": {"custom_llm_provider": "github_copilot"} + }' +``` + +### Step 5: Attach Credential to a Model + +Reference the credential in your config (or via the UI under Models → Add Model → Existing Credentials): + +```yaml showLineNumbers title="config.yaml" +model_list: + - model_name: copilot-gpt4 + litellm_params: + model: github_copilot/gpt-4 + litellm_credential_name: my-copilot +``` + +### Multiple Accounts + +Create multiple credentials with different names to support different GitHub accounts: + +```yaml showLineNumbers title="config.yaml - Multiple accounts" +model_list: + - model_name: copilot-team-a + litellm_params: + model: github_copilot/gpt-4 + litellm_credential_name: team-a-copilot + - model_name: copilot-team-b + litellm_params: + model: github_copilot/gpt-4 + litellm_credential_name: team-b-copilot +``` ## Configuration ### Environment Variables -You can customize token storage locations: +You can customize token storage locations (SDK / file-based mode): ```bash showLineNumbers title="Environment Variables" # Optional: Custom token directory @@ -208,4 +323,3 @@ extra_headers = { "user-agent": "GithubCopilot/1.155.0" # User agent } ``` - diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index 85c22516f95..f17cd659393 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -19,27 +19,44 @@ from .common_utils import ( # Constants GITHUB_CLIENT_ID = "Iv1.b507a08c87ecfe98" + +# Module-level cache for copilot inference tokens in credential mode. +# Key = GitHub access_token, value = api_key_info dict (token + expires_at). +# Mirrors the file-based api-key.json pattern but kept in memory so that +# per-request Authenticator instances share the cached token. +_credential_api_key_cache: Dict[str, Dict[str, Any]] = {} GITHUB_DEVICE_CODE_URL = "https://github.com/login/device/code" GITHUB_ACCESS_TOKEN_URL = "https://github.com/login/oauth/access_token" GITHUB_API_KEY_URL = "https://api.github.com/copilot_internal/v2/token" class Authenticator: - def __init__(self) -> None: - """Initialize the GitHub Copilot authenticator with configurable token paths.""" - # Token storage paths - self.token_dir = os.getenv( - "GITHUB_COPILOT_TOKEN_DIR", - os.path.expanduser("~/.config/litellm/github_copilot"), - ) - 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") - ) - self._ensure_token_dir() + def __init__(self, access_token: Optional[str] = None) -> None: + """Initialize the GitHub Copilot authenticator. + + Args: + access_token: If provided, the authenticator operates in + *credential mode* — it uses this token directly instead of + the file-based device-code flow. When ``None`` (the + default), the existing file-based behaviour is preserved. + """ + self._injected_access_token = access_token + + if access_token is None: + # File-based mode (backward compatible) + self.token_dir = os.getenv( + "GITHUB_COPILOT_TOKEN_DIR", + os.path.expanduser("~/.config/litellm/github_copilot"), + ) + 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"), + ) + self._ensure_token_dir() def get_access_token(self) -> str: """ @@ -51,32 +68,23 @@ class Authenticator: Raises: GetAccessTokenError: If unable to obtain an access token after retries. """ + if self._injected_access_token is not None: + return self._injected_access_token + try: with open(self.access_token_file, "r") as f: access_token = f.read().strip() if access_token: return access_token except IOError: - verbose_logger.warning( - "No existing access token found or error reading file" - ) - - for attempt in range(3): - verbose_logger.debug(f"Access token acquisition attempt {attempt + 1}/3") - try: - access_token = self._login() - try: - with open(self.access_token_file, "w") as f: - f.write(access_token) - except IOError: - verbose_logger.error("Error saving access token to file") - return access_token - except (GetDeviceCodeError, GetAccessTokenError, RefreshAPIKeyError) as e: - verbose_logger.warning(f"Failed attempt {attempt + 1}: {str(e)}") - continue + pass # No file — fall through to auth error below raise GetAccessTokenError( - message="Failed to get access token after 3 attempts", + message=( + "No GitHub Copilot access token configured. " + "Use a named credential via the LiteLLM proxy or UI before making requests. " + "See: https://docs.litellm.ai/docs/providers/github_copilot" + ), status_code=401, ) @@ -90,6 +98,9 @@ class Authenticator: Raises: GetAPIKeyError: If unable to obtain an API key. """ + if self._injected_access_token is not None: + return self._get_api_key_credential_mode() + try: with open(self.api_key_file, "r") as f: api_key_info = json.load(f) @@ -139,6 +150,13 @@ class Authenticator: Returns: Optional[str]: The GitHub Copilot API endpoint, or None if not found. """ + if self._injected_access_token is not None: + cached = _credential_api_key_cache.get(self._injected_access_token) + if cached: + endpoints = cached.get("endpoints", {}) + return endpoints.get("api") + return None + try: with open(self.api_key_file, "r") as f: api_key_info = json.load(f) @@ -149,6 +167,35 @@ class Authenticator: verbose_logger.warning(f"Error reading API endpoint from file: {str(e)}") return None + def _get_api_key_credential_mode(self) -> str: + """Get API key when operating in credential mode (injected access token). + + Uses a module-level cache keyed by access_token so that multiple + per-request Authenticator instances share the same cached copilot + inference token and avoid redundant GitHub API calls. + """ + cached = _credential_api_key_cache.get(self._injected_access_token) # type: ignore[arg-type] + if cached and cached.get("expires_at", 0) > datetime.now().timestamp(): + token = cached.get("token") + if token: + return token + + try: + api_key_info = self._refresh_api_key() + _credential_api_key_cache[self._injected_access_token] = api_key_info # type: ignore[index] + token = api_key_info.get("token") + if not token: + raise GetAPIKeyError( + message="API key response missing token", + status_code=401, + ) + return token + except RefreshAPIKeyError as e: + raise GetAPIKeyError( + message=f"Failed to refresh API key: {str(e)}", + status_code=401, + ) + def _refresh_api_key(self) -> Dict[str, Any]: """ Refresh the API key using the access token. @@ -177,6 +224,8 @@ class Authenticator: verbose_logger.warning( f"API key response missing token: {response_json}" ) + except GetAccessTokenError: + raise # Re-raise with the original helpful message (docs link etc.) except httpx.HTTPStatusError as e: verbose_logger.error( f"HTTP error refreshing API key (attempt {attempt+1}/{max_retries}): {str(e)}" @@ -194,10 +243,14 @@ class Authenticator: if not os.path.exists(self.token_dir): os.makedirs(self.token_dir, exist_ok=True) - def _get_github_headers(self, access_token: Optional[str] = None) -> Dict[str, str]: + @staticmethod + def get_github_headers(access_token: Optional[str] = None) -> Dict[str, str]: """ Generate standard GitHub headers for API requests. + This is a static method so it can be imported and used by the SSO + endpoint module without instantiating an Authenticator. + Args: access_token: Optional access token to include in the headers. @@ -210,16 +263,18 @@ class Authenticator: "editor-plugin-version": "copilot/1.155.0", "user-agent": "GithubCopilot/1.155.0", "accept-encoding": "gzip,deflate,br", + "content-type": "application/json", } if access_token: headers["authorization"] = f"token {access_token}" - if "content-type" not in headers: - headers["content-type"] = "application/json" - return headers + # Backward-compatible instance alias + def _get_github_headers(self, access_token: Optional[str] = None) -> Dict[str, str]: + return Authenticator.get_github_headers(access_token) + def _get_device_code(self) -> Dict[str, str]: """ Get a device code for GitHub authentication. diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index be8ad7d0877..8e35ea5f6c7 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -8,8 +8,10 @@ from litellm.types.llms.openai import AllMessageValues from ..authenticator import Authenticator from ..common_utils import ( GITHUB_COPILOT_API_BASE, + GetAccessTokenError, GetAPIKeyError, get_copilot_default_headers, + get_copilot_static_headers, ) @@ -21,7 +23,6 @@ class GithubCopilotConfig(OpenAIConfig): custom_llm_provider: str = "openai", ) -> None: super().__init__() - self.authenticator = Authenticator() def _get_openai_compatible_provider_info( self, @@ -30,10 +31,17 @@ class GithubCopilotConfig(OpenAIConfig): api_key: Optional[str], custom_llm_provider: str, ) -> Tuple[Optional[str], Optional[str], str]: - dynamic_api_base = self.authenticator.get_api_base() or GITHUB_COPILOT_API_BASE + # If no api_key is provided we're being called at router registration time + # (before litellm_credential_name has been resolved). Return the default + # base and defer auth to request time rather than raising here and causing + # the deployment to be silently dropped from the router. + if not api_key: + return GITHUB_COPILOT_API_BASE, None, custom_llm_provider + authenticator = Authenticator(access_token=api_key) + dynamic_api_base = authenticator.get_api_base() or GITHUB_COPILOT_API_BASE try: - dynamic_api_key = self.authenticator.get_api_key() - except GetAPIKeyError as e: + dynamic_api_key = authenticator.get_api_key() + except (GetAPIKeyError, GetAccessTokenError) as e: raise AuthenticationError( model=model, llm_provider=custom_llm_provider, @@ -82,13 +90,19 @@ class GithubCopilotConfig(OpenAIConfig): headers, model, messages, optional_params, litellm_params, api_key, api_base ) - # Add Copilot-specific headers (editor-version, user-agent, etc.) - try: - copilot_api_key = self.authenticator.get_api_key() - copilot_headers = get_copilot_default_headers(copilot_api_key) - validated_headers = {**copilot_headers, **validated_headers} - except GetAPIKeyError: - pass # Will be handled later in the request flow + # Always add static Copilot headers (editor-version, user-agent, etc.) + # These are required by the GitHub Copilot API on every request. + validated_headers = {**get_copilot_static_headers(), **validated_headers} + + # If we have an api_key (GitHub access token), exchange it for a + # copilot inference token and set the Authorization header. + if api_key: + try: + copilot_api_key = Authenticator(access_token=api_key).get_api_key() + copilot_headers = get_copilot_default_headers(copilot_api_key) + validated_headers = {**copilot_headers, **validated_headers} + except (GetAPIKeyError, GetAccessTokenError): + pass # Will be handled later in the request flow # Add X-Initiator header based on message roles initiator = self._determine_initiator(messages) diff --git a/litellm/llms/github_copilot/common_utils.py b/litellm/llms/github_copilot/common_utils.py index d3169e3ca94..105297761d2 100644 --- a/litellm/llms/github_copilot/common_utils.py +++ b/litellm/llms/github_copilot/common_utils.py @@ -56,14 +56,14 @@ class GetAPIKeyError(GithubCopilotError): pass -def get_copilot_default_headers(api_key: str) -> dict: +def get_copilot_static_headers() -> dict: """ - Get default headers for GitHub Copilot Responses API. + Get static headers required by the GitHub Copilot API. - Based on copilot-api's header configuration. + These headers (editor-version, user-agent, etc.) must be present on every + request regardless of whether the API key has been resolved yet. """ return { - "Authorization": f"Bearer {api_key}", "content-type": "application/json", "copilot-integration-id": "vscode-chat", "editor-version": "vscode/1.95.0", # Fixed version for stability @@ -74,3 +74,15 @@ def get_copilot_default_headers(api_key: str) -> dict: "x-request-id": str(uuid4()), "x-vscode-user-agent-library-version": "electron-fetch", } + + +def get_copilot_default_headers(api_key: str) -> dict: + """ + Get default headers for GitHub Copilot Responses API. + + Based on copilot-api's header configuration. + """ + return { + **get_copilot_static_headers(), + "Authorization": f"Bearer {api_key}", + } diff --git a/litellm/llms/github_copilot/embedding/transformation.py b/litellm/llms/github_copilot/embedding/transformation.py index fa7bd4e3223..63ee6ca1e81 100644 --- a/litellm/llms/github_copilot/embedding/transformation.py +++ b/litellm/llms/github_copilot/embedding/transformation.py @@ -19,6 +19,7 @@ from litellm.utils import convert_to_model_response_object from ..authenticator import Authenticator from ..common_utils import ( + GetAccessTokenError, GetAPIKeyError, GITHUB_COPILOT_API_BASE, get_copilot_default_headers, @@ -41,7 +42,6 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): def __init__(self) -> None: super().__init__() - self.authenticator = Authenticator() def validate_environment( self, @@ -57,9 +57,6 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): Validate environment and set up headers for GitHub Copilot API. """ try: - # Get GitHub Copilot API key via OAuth - api_key = self.authenticator.get_api_key() - if not api_key: raise AuthenticationError( model=model, @@ -67,6 +64,9 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): message="GitHub Copilot API key is required. Please authenticate via OAuth Device Flow.", ) + # Get GitHub Copilot API key via OAuth + api_key = Authenticator(access_token=api_key).get_api_key() + # Get default headers default_headers = get_copilot_default_headers(api_key) @@ -79,7 +79,7 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): return merged_headers - except GetAPIKeyError as e: + except (GetAPIKeyError, GetAccessTokenError) as e: raise AuthenticationError( model=model, llm_provider="github_copilot", @@ -98,10 +98,11 @@ 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 - api_base = ( - self.authenticator.get_api_base() or api_base or GITHUB_COPILOT_API_BASE - ) + # Use provided api_base or fall back to credential-resolved base or default + if api_key: + api_base = Authenticator(access_token=api_key).get_api_base() or api_base or GITHUB_COPILOT_API_BASE + else: + api_base = api_base or GITHUB_COPILOT_API_BASE # Remove trailing slashes api_base = api_base.rstrip("/") diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 46efc124b1d..ee5f95baacd 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -23,6 +23,7 @@ from litellm.types.utils import LlmProviders from ..authenticator import Authenticator from ..common_utils import ( GITHUB_COPILOT_API_BASE, + GetAccessTokenError, GetAPIKeyError, get_copilot_default_headers, ) @@ -54,7 +55,6 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): def __init__(self) -> None: super().__init__() - self.authenticator = Authenticator() @property def custom_llm_provider(self) -> LlmProviders: @@ -103,16 +103,20 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): - User-provided extra_headers (merged with priority) """ try: - # Get GitHub Copilot API key via OAuth - api_key = self.authenticator.get_api_key() - - if not api_key: + if isinstance(litellm_params, dict): + _api_key = litellm_params.get("api_key") + else: + _api_key = getattr(litellm_params, "api_key", None) if litellm_params else None + if not _api_key: raise AuthenticationError( model=model, llm_provider="github_copilot", message="GitHub Copilot API key is required. Please authenticate via OAuth Device Flow.", ) + # Get GitHub Copilot API key via OAuth + api_key = Authenticator(access_token=_api_key).get_api_key() + # Get default headers (from copilot-api configuration) default_headers = get_copilot_default_headers(api_key) @@ -143,7 +147,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): return merged_headers - except GetAPIKeyError as e: + except (GetAPIKeyError, GetAccessTokenError) as e: raise AuthenticationError( model=model, llm_provider="github_copilot", @@ -164,10 +168,12 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): Business/enterprise accounts (api.business.githubcopilot.com) can be added in the future by detecting account type. """ - # Use provided api_base or fall back to authenticator's base or default - api_base = ( - api_base or self.authenticator.get_api_base() or GITHUB_COPILOT_API_BASE - ) + # Use provided api_base or fall back to credential-resolved base or default + _api_key = litellm_params.get("api_key") if isinstance(litellm_params, dict) else None + if _api_key: + api_base = api_base or Authenticator(access_token=_api_key).get_api_base() or GITHUB_COPILOT_API_BASE + else: + api_base = api_base or GITHUB_COPILOT_API_BASE # Remove trailing slashes api_base = api_base.rstrip("/") diff --git a/litellm/main.py b/litellm/main.py index 112fef44e55..97b24c4ad81 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2609,12 +2609,21 @@ def completion( # type: ignore # noqa: PLR0915 if custom_llm_provider == "github_copilot": from litellm.llms.github_copilot.authenticator import Authenticator from litellm.llms.github_copilot.common_utils import ( + GetAccessTokenError, + GetAPIKeyError, get_copilot_default_headers, + get_copilot_static_headers, ) - copilot_auth = Authenticator() - copilot_api_key = copilot_auth.get_api_key() - copilot_headers = get_copilot_default_headers(copilot_api_key) + # Always add static headers (editor-version, user-agent, etc.) + # — the Copilot API requires these on every request. + copilot_headers = get_copilot_static_headers() + if api_key: + try: + copilot_api_key = Authenticator(access_token=api_key).get_api_key() + copilot_headers = get_copilot_default_headers(copilot_api_key) + except (GetAPIKeyError, GetAccessTokenError): + pass # auth failure handled downstream if extra_headers: copilot_headers.update(extra_headers) extra_headers = copilot_headers diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 64f860fc4f1..42cbd45aa23 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -113,6 +113,29 @@ async def create_credential( raise handle_exception_on_proxy(e) +def _fetch_github_login(api_key: str) -> Optional[str]: + """ + Call GET https://api.github.com/user with the given GitHub access token + and return the login name, or None if the call fails. + """ + from litellm.llms.custom_httpx.http_handler import _get_httpx_client + + try: + sync_client = _get_httpx_client() + resp = sync_client.get( + "https://api.github.com/user", + headers={ + "Authorization": f"token {api_key}", + "Accept": "application/json", + }, + ) + if resp.status_code == 200: + return resp.json().get("login") + except Exception as e: + verbose_proxy_logger.warning(f"Could not fetch GitHub user info: {e}") + return None + + @router.get( "/credentials", dependencies=[Depends(user_api_key_auth)], @@ -127,14 +150,25 @@ async def get_credentials( [BETA] endpoint. This might change unexpectedly. """ try: - masked_credentials = [ - { - "credential_name": credential.credential_name, - "credential_values": _get_masked_values(credential.credential_values), - "credential_info": credential.credential_info, - } - for credential in litellm.credential_list - ] + masked_credentials = [] + for credential in litellm.credential_list: + credential_info = dict(credential.credential_info or {}) + # For GitHub Copilot credentials, inject runtime github_login from the API. + # The login is NOT stored in the DB — it's fetched live and added to the + # response only so the UI can display it. + if credential_info.get("custom_llm_provider") == "github_copilot": + api_key = (credential.credential_values or {}).get("api_key") + if api_key: + github_login = _fetch_github_login(api_key) + if github_login: + credential_info = {**credential_info, "github_login": github_login} + masked_credentials.append( + { + "credential_name": credential.credential_name, + "credential_values": _get_masked_values(credential.credential_values), + "credential_info": credential_info, + } + ) return {"success": True, "credentials": masked_credentials} except Exception as e: return handle_exception_on_proxy(e) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index c638e294268..e843916595a 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -324,7 +324,12 @@ class ProxyInitializationHelpers: """Helper function to determine the event loop type based on platform""" if sys.platform in ("win32", "cygwin", "cli"): return None # Let uvicorn choose the default loop on Windows - return "uvloop" + try: + import uvloop # noqa: F401 + + return "uvloop" + except (ImportError, Exception): + return "asyncio" @staticmethod def _maybe_setup_prometheus_multiproc_dir( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e982c934aa6..806842a6495 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -319,6 +319,9 @@ from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob from litellm.proxy.common_utils.swagger_utils import ERROR_RESPONSES from litellm.proxy.container_endpoints.endpoints import router as container_router from litellm.proxy.credential_endpoints.endpoints import router as credential_router +from litellm.proxy.credential_endpoints.github_copilot_sso import ( + router as github_copilot_sso_router, +) from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router @@ -13447,6 +13450,7 @@ app.include_router(vector_store_router) app.include_router(vector_store_management_router) app.include_router(vector_store_files_router) app.include_router(credential_router) +app.include_router(github_copilot_sso_router) app.include_router(llm_passthrough_router) app.include_router(webrtc_router) app.include_router(mcp_management_router) diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index dda8e49d4c8..3856ade5274 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -1174,28 +1174,8 @@ "provider": "GITHUB_COPILOT", "provider_display_name": "Github Copilot", "litellm_provider": "github_copilot", - "credential_fields": [ - { - "key": "api_base", - "label": "API Base", - "placeholder": null, - "tooltip": null, - "required": false, - "field_type": "text", - "options": null, - "default_value": null - }, - { - "key": "api_key", - "label": "API Key", - "placeholder": null, - "tooltip": null, - "required": false, - "field_type": "password", - "options": null, - "default_value": null - } - ], + "auth_flow": "device_code", + "credential_fields": [], "default_model_placeholder": "gpt-3.5-turbo" }, { diff --git a/litellm/router.py b/litellm/router.py index 36046ebf302..9b2fdd2696c 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1455,6 +1455,7 @@ class Router: self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) kwargs.pop("silent_model", None) # Ensure it's not in kwargs either + model_name = litellm_params["model"] potential_model_client = self._get_client( deployment=deployment, kwargs=kwargs @@ -2365,6 +2366,17 @@ class Router: existing_tags.append(credential_tag) kwargs[metadata_variable_name]["tags"] = existing_tags + ## EARLY CREDENTIAL RESOLUTION + # Resolve api_key from litellm_credential_name before client selection + # or function invocation. Without this, deployments using named + # credentials pass api_key=None to cached-client checks and to + # downstream litellm functions (whose @client decorator skips + # load_credentials_from_list for async requests). + if credential_name and not kwargs.get("api_key"): + _cred = CredentialAccessor.get_credential_values(credential_name) + if _cred.get("api_key"): + kwargs["api_key"] = _cred["api_key"] + kwargs["model_info"] = model_info kwargs["timeout"] = self._get_timeout( diff --git a/litellm/types/proxy/public_endpoints/public_endpoints.py b/litellm/types/proxy/public_endpoints/public_endpoints.py index caa9a978530..9945d799534 100644 --- a/litellm/types/proxy/public_endpoints/public_endpoints.py +++ b/litellm/types/proxy/public_endpoints/public_endpoints.py @@ -28,6 +28,7 @@ class ProviderCreateInfo(BaseModel): provider: str provider_display_name: str litellm_provider: str + auth_flow: Optional[Literal["device_code"]] = None credential_fields: List[ProviderCredentialField] default_model_placeholder: Optional[str] = None diff --git a/tests/test_litellm/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py b/tests/test_litellm/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py index f4440aff9d1..da4a9129bd3 100644 --- a/tests/test_litellm/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py +++ b/tests/test_litellm/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py @@ -10,26 +10,24 @@ from litellm.exceptions import AuthenticationError from litellm.llms.github_copilot.embedding.transformation import GithubCopilotEmbeddingConfig from litellm.llms.github_copilot.common_utils import GetAPIKeyError -def test_github_copilot_embedding_config_validate_environment(): +@patch("litellm.llms.github_copilot.embedding.transformation.Authenticator") +def test_github_copilot_embedding_config_validate_environment(mock_authenticator_class): """Test the GitHub Copilot embedding configuration environment validation.""" - config = GithubCopilotEmbeddingConfig() - - # Mock the authenticator mock_api_key = "gh.test-key-123456789" - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = mock_api_key + mock_auth_instance = MagicMock() + mock_auth_instance.get_api_key.return_value = mock_api_key + mock_authenticator_class.return_value = mock_auth_instance - # Test with valid API key - headers = {} + config = GithubCopilotEmbeddingConfig() model = "github_copilot/text-embedding-3-small" - + validated_headers = config.validate_environment( - headers=headers, + headers={}, model=model, messages=[], optional_params={}, litellm_params={}, - api_key=None, + api_key="gh-access-token", ) assert validated_headers["Authorization"] == f"Bearer {mock_api_key}" @@ -37,12 +35,7 @@ def test_github_copilot_embedding_config_validate_environment(): assert validated_headers["editor-version"] == "vscode/1.95.0" assert "x-request-id" in validated_headers - # Test with authentication failure - config.authenticator.get_api_key.side_effect = GetAPIKeyError( - message="Failed to get API key", - status_code=401, - ) - + # Test with no api_key → immediate AuthenticationError with pytest.raises(AuthenticationError) as excinfo: config.validate_environment( headers={}, @@ -52,16 +45,33 @@ def test_github_copilot_embedding_config_validate_environment(): litellm_params={}, api_key=None, ) + assert "required" in str(excinfo.value).lower() + # Test with authentication failure from GitHub + mock_auth_instance.get_api_key.side_effect = GetAPIKeyError( + message="Failed to get API key", + status_code=401, + ) + with pytest.raises(AuthenticationError) as excinfo: + config.validate_environment( + headers={}, + model=model, + messages=[], + optional_params={}, + litellm_params={}, + api_key="gh-access-token", + ) assert "Failed to get API key" in str(excinfo.value) -def test_github_copilot_embedding_config_get_complete_url(): +@patch("litellm.llms.github_copilot.embedding.transformation.Authenticator") +def test_github_copilot_embedding_config_get_complete_url(mock_authenticator_class): """Test the GitHub Copilot embedding configuration URL generation.""" + mock_auth_instance = MagicMock() + mock_authenticator_class.return_value = mock_auth_instance + config = GithubCopilotEmbeddingConfig() - config.authenticator = MagicMock() - - # Test with default API base - config.authenticator.get_api_base.return_value = None + + # No api_key → always default base url = config.get_complete_url( api_base=None, api_key=None, @@ -71,19 +81,18 @@ 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" + # api_key + authenticator returns custom base + mock_auth_instance.get_api_base.return_value = "https://api.enterprise.githubcopilot.com" url = config.get_complete_url( api_base=None, - api_key=None, + api_key="gh-access-token", model="github_copilot/text-embedding-3-small", optional_params={}, litellm_params={}, ) assert url == "https://api.enterprise.githubcopilot.com/embeddings" - # Test with custom API base from params - config.authenticator.get_api_base.return_value = None + # Explicit api_base always wins url = config.get_complete_url( api_base="https://custom.api.com", api_key=None, diff --git a/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py b/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py index 1feb0244dbb..ce5da4fc3cc 100644 --- a/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py +++ b/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py @@ -44,7 +44,6 @@ class TestGithubCopilotResponsesAPITransformation: @patch("litellm.llms.github_copilot.responses.transformation.Authenticator") def test_github_copilot_responses_endpoint_url(self, mock_authenticator_class): """Test that get_complete_url returns correct GitHub Copilot endpoint""" - # Mock authenticator to return default base mock_auth_instance = MagicMock() mock_auth_instance.get_api_base.return_value = ( "https://api.individual.githubcopilot.com" @@ -53,13 +52,21 @@ class TestGithubCopilotResponsesAPITransformation: config = GithubCopilotResponsesAPIConfig() - # Test with default GitHub Copilot API base (from authenticator) + # No api_key in litellm_params → default base url = config.get_complete_url(api_base=None, litellm_params={}) - assert url == "https://api.individual.githubcopilot.com/responses", ( - f"Expected GitHub Copilot responses endpoint, got {url}" + assert url == "https://api.githubcopilot.com/responses", ( + f"Expected default endpoint when no api_key, got {url}" ) - # Test with custom api_base (overrides authenticator) + # api_key present → authenticator resolves custom base + url = config.get_complete_url( + api_base=None, litellm_params={"api_key": "gh-access-token"} + ) + assert url == "https://api.individual.githubcopilot.com/responses", ( + f"Expected authenticator-resolved endpoint, got {url}" + ) + + # Explicit api_base always wins regardless of api_key custom_url = config.get_complete_url( api_base="https://custom.githubcopilot.com", litellm_params={} ) @@ -67,7 +74,7 @@ class TestGithubCopilotResponsesAPITransformation: f"Expected custom endpoint, got {custom_url}" ) - # Test with trailing slash + # Trailing slash stripped url_with_slash = config.get_complete_url( api_base="https://api.githubcopilot.com/", litellm_params={} ) @@ -86,7 +93,7 @@ class TestGithubCopilotResponsesAPITransformation: config = GithubCopilotResponsesAPIConfig() headers = config.validate_environment( - headers={}, model="gpt-5.1-codex", litellm_params={} + headers={}, model="gpt-5.1-codex", litellm_params={"api_key": "gh-access-token"} ) # Check required headers @@ -115,7 +122,7 @@ class TestGithubCopilotResponsesAPITransformation: } headers = config.validate_environment( - headers=custom_headers, model="gpt-5.1-codex", litellm_params={} + headers=custom_headers, model="gpt-5.1-codex", litellm_params={"api_key": "gh-access-token"} ) # User header should override default diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py index c6ae2b9c4e1..92c892faad6 100644 --- a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py +++ b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py @@ -67,24 +67,18 @@ class TestGitHubCopilotAuthenticator: token = authenticator.get_access_token() assert token == mock_token - def test_get_access_token_login(self, authenticator): - """Test logging in to get an access token.""" - mock_token = "mock-access-token" - - with patch.object(authenticator, "_login", return_value=mock_token), \ - patch("builtins.open", mock_open()), \ - patch("builtins.open", side_effect=IOError) as mock_read: - token = authenticator.get_access_token() - assert token == mock_token - authenticator._login.assert_called_once() + def test_get_access_token_no_file_raises(self, authenticator): + """Test that GetAccessTokenError is raised when no access-token file exists.""" + with patch("builtins.open", side_effect=IOError): + with pytest.raises(GetAccessTokenError) as exc_info: + authenticator.get_access_token() + assert "https://docs.litellm.ai/docs/providers/github_copilot" in str(exc_info.value) - def test_get_access_token_failure(self, authenticator): - """Test that an exception is raised after multiple login failures.""" - with patch.object(authenticator, "_login", side_effect=GetDeviceCodeError(message="Test error", status_code=400)), \ - patch("builtins.open", side_effect=IOError): + def test_get_access_token_empty_file_raises(self, authenticator): + """Test that GetAccessTokenError is raised when the access-token file is empty.""" + with patch("builtins.open", mock_open(read_data="")): with pytest.raises(GetAccessTokenError): 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.""" @@ -179,6 +173,16 @@ class TestGitHubCopilotAuthenticator: authenticator._poll_for_access_token.assert_called_once_with("mock-device-code") mock_print.assert_called_once() + def test_get_github_headers_static(self): + """Test that get_github_headers works as a static method.""" + headers = Authenticator.get_github_headers() + assert "accept" in headers + assert "content-type" in headers + assert "authorization" not in headers + + headers_with_token = Authenticator.get_github_headers("my-token") + assert headers_with_token["authorization"] == "token my-token" + def test_get_api_base_from_file(self, authenticator): """Test retrieving the API base endpoint from a file.""" mock_api_key_data = json.dumps({ @@ -189,3 +193,108 @@ class TestGitHubCopilotAuthenticator: 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" + + +class TestAuthenticatorCredentialMode: + """Tests for credential mode (injected access token).""" + + def test_init_credential_mode_no_file_io(self): + """Credential mode should not create any directories.""" + auth = Authenticator(access_token="test-token") + assert auth._injected_access_token == "test-token" + assert not hasattr(auth, "token_dir") + + def test_get_access_token_returns_injected(self): + """get_access_token returns the injected token directly.""" + auth = Authenticator(access_token="my-github-token") + assert auth.get_access_token() == "my-github-token" + + def test_get_api_key_credential_mode(self): + """get_api_key in credential mode calls _refresh_api_key and caches.""" + import litellm.llms.github_copilot.authenticator as auth_module + + access_token = "my-github-token-caching-test" + # Clear any stale cache entry before the test + auth_module._credential_api_key_cache.pop(access_token, None) + + auth = Authenticator(access_token=access_token) + future_time = (datetime.now() + timedelta(hours=1)).timestamp() + mock_api_key_info = {"token": "copilot-api-key", "expires_at": future_time} + + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.raise_for_status.return_value = None + mock_response.json.return_value = mock_api_key_info + mock_client.get.return_value = mock_response + + with patch( + "litellm.llms.github_copilot.authenticator._get_httpx_client", + return_value=mock_client, + ): + api_key = auth.get_api_key() + assert api_key == "copilot-api-key" + + # Second call should use cache (no additional HTTP calls) + api_key2 = auth.get_api_key() + assert api_key2 == "copilot-api-key" + assert mock_client.get.call_count == 1 # Only called once + + # Cleanup + auth_module._credential_api_key_cache.pop(access_token, None) + + def test_get_api_key_credential_mode_expired_cache(self): + """get_api_key re-fetches when cached token is expired.""" + import litellm.llms.github_copilot.authenticator as auth_module + + past_time = (datetime.now() - timedelta(hours=1)).timestamp() + future_time = (datetime.now() + timedelta(hours=1)).timestamp() + access_token = "my-github-token-expired-test" + + auth = Authenticator(access_token=access_token) + # Pre-populate module-level cache with an expired entry + auth_module._credential_api_key_cache[access_token] = { + "token": "old-key", + "expires_at": past_time, + } + + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.raise_for_status.return_value = None + mock_response.json.return_value = {"token": "new-key", "expires_at": future_time} + mock_client.get.return_value = mock_response + + with patch( + "litellm.llms.github_copilot.authenticator._get_httpx_client", + return_value=mock_client, + ): + api_key = auth.get_api_key() + assert api_key == "new-key" + assert mock_client.get.call_count == 1 + + # Cleanup + auth_module._credential_api_key_cache.pop(access_token, None) + + def test_get_api_base_credential_mode_no_cache(self): + """get_api_base returns None when no cache is available.""" + import litellm.llms.github_copilot.authenticator as auth_module + + access_token = "my-github-token-no-cache-test" + auth = Authenticator(access_token=access_token) + # Ensure no stale cache + auth_module._credential_api_key_cache.pop(access_token, None) + assert auth.get_api_base() is None + + def test_get_api_base_credential_mode_with_cache(self): + """get_api_base returns endpoint from module-level cache.""" + import litellm.llms.github_copilot.authenticator as auth_module + + access_token = "my-github-token-cache-test" + auth = Authenticator(access_token=access_token) + auth_module._credential_api_key_cache[access_token] = { + "token": "test", + "endpoints": {"api": "https://custom.copilot.api"}, + } + assert auth.get_api_base() == "https://custom.copilot.api" + + # Cleanup + auth_module._credential_api_key_cache.pop(access_token, None) diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py b/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py index a1b6ff7c509..8532fd69683 100644 --- a/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py +++ b/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py @@ -31,63 +31,62 @@ from litellm.llms.github_copilot.common_utils import ( ) -def test_github_copilot_config_get_openai_compatible_provider_info(): +@patch("litellm.llms.github_copilot.chat.transformation.Authenticator") +def test_github_copilot_config_get_openai_compatible_provider_info(mock_authenticator_class): """Test the GitHub Copilot configuration provider info retrieval.""" + mock_api_key = "gh.test-key-123456789" + mock_auth_instance = MagicMock() + mock_auth_instance.get_api_key.return_value = mock_api_key + mock_auth_instance.get_api_base.return_value = "https://api.enterprise.githubcopilot.com" + mock_authenticator_class.return_value = mock_auth_instance config = GithubCopilotConfig() - - # Mock the authenticator to avoid actual API calls - mock_api_key = "gh.test-key-123456789" - 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" - - # Test with default values model = "github_copilot/gpt-4" - ( - api_base, - dynamic_api_key, - custom_llm_provider, - ) = config._get_openai_compatible_provider_info( - model=model, - api_base=None, - api_key=None, - custom_llm_provider="github_copilot", - ) - assert api_base == "https://api.enterprise.githubcopilot.com" - assert dynamic_api_key == mock_api_key - assert custom_llm_provider == "github_copilot" - - # Test fallback to default if no dynamic endpoint - config.authenticator.get_api_base.return_value = None - ( - api_base, - dynamic_api_key, - custom_llm_provider, - ) = config._get_openai_compatible_provider_info( + # No api_key → returns defaults immediately without calling authenticator + api_base, dynamic_api_key, custom_llm_provider = config._get_openai_compatible_provider_info( model=model, api_base=None, api_key=None, custom_llm_provider="github_copilot", ) assert api_base == "https://api.githubcopilot.com" + assert dynamic_api_key is None + assert custom_llm_provider == "github_copilot" - # Test with authentication failure - config.authenticator.get_api_key.side_effect = GetAPIKeyError( + # With api_key → uses authenticator's dynamic base + api_base, dynamic_api_key, custom_llm_provider = config._get_openai_compatible_provider_info( + model=model, + api_base=None, + api_key="gh-access-token", + custom_llm_provider="github_copilot", + ) + assert api_base == "https://api.enterprise.githubcopilot.com" + assert dynamic_api_key == mock_api_key + assert custom_llm_provider == "github_copilot" + + # Fallback to default when authenticator returns no base + mock_auth_instance.get_api_base.return_value = None + api_base, _, _ = config._get_openai_compatible_provider_info( + model=model, + api_base=None, + api_key="gh-access-token", + custom_llm_provider="github_copilot", + ) + assert api_base == "https://api.githubcopilot.com" + + # Authentication failure + mock_auth_instance.get_api_key.side_effect = GetAPIKeyError( message="Failed to get API key", status_code=401, ) - with pytest.raises(AuthenticationError) as excinfo: config._get_openai_compatible_provider_info( model=model, api_base=None, - api_key=None, + api_key="gh-access-token", custom_llm_provider="github_copilot", ) - assert "Failed to get API key" in str(excinfo.value) @@ -121,6 +120,7 @@ def test_completion_github_copilot_mock_response(mock_completion, mock_get_api_k response = completion( model="github_copilot/gpt-4", messages=messages, + api_key="gh-access-token", extra_headers=headers, ) @@ -181,10 +181,6 @@ def test_x_initiator_header_user_request(): """Test that user-only messages result in X-Initiator: user header""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ {"role": "system", "content": "You are an assistant."}, @@ -208,10 +204,6 @@ def test_x_initiator_header_agent_request_with_assistant(): """Test that messages with assistant role result in X-Initiator: agent header""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ {"role": "system", "content": "You are an assistant."}, @@ -235,10 +227,6 @@ def test_x_initiator_header_agent_request_with_tool(): """Test that messages with tool role result in X-Initiator: agent header""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ {"role": "system", "content": "You are an assistant."}, @@ -261,11 +249,6 @@ def test_x_initiator_header_agent_request_with_tool(): def test_x_initiator_header_mixed_messages_with_agent_roles(): """Test that mixed messages with agent roles (assistant/tool) result in X-Initiator: agent header""" config = GithubCopilotConfig() - - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ {"role": "user", "content": "Hello"}, @@ -289,11 +272,6 @@ def test_x_initiator_header_mixed_messages_with_agent_roles(): def test_x_initiator_header_user_only_messages(): """Test that user + system only messages result in X-Initiator: user header""" config = GithubCopilotConfig() - - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ {"role": "system", "content": "You are an assistant."}, @@ -318,10 +296,6 @@ def test_x_initiator_header_empty_messages(): """Test that empty messages result in X-Initiator: user header""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [] @@ -342,10 +316,6 @@ def test_x_initiator_header_system_only_messages(): """Test that system-only messages result in X-Initiator: user header""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ {"role": "system", "content": "You are an assistant."}, @@ -419,10 +389,6 @@ def test_copilot_vision_request_header_with_image(): """Test that Copilot-Vision-Request header is added when messages contain images""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ { @@ -455,10 +421,6 @@ def test_copilot_vision_request_header_text_only(): """Test that Copilot-Vision-Request header is not added for text-only messages""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ {"role": "user", "content": "Just a text message"}, @@ -482,10 +444,6 @@ def test_copilot_vision_request_header_with_type_image_url(): """Test that Copilot-Vision-Request header is added for content with type: image_url""" config = GithubCopilotConfig() - # Mock the authenticator - config.authenticator = MagicMock() - config.authenticator.get_api_key.return_value = "gh.test-key-123" - config.authenticator.get_api_base.return_value = None messages = [ { diff --git a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx index 694a98201c6..7270ef0577b 100644 --- a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx @@ -1,21 +1,76 @@ import { TextInput } from "@tremor/react"; -import { Select as AntdSelect, Button, Form, Modal, Tooltip, Typography } from "antd"; +import { Select as AntdSelect, Button, Form, Modal, Spin, Tooltip, Typography } from "antd"; import type { UploadProps } from "antd/es/upload"; -import React, { useState } from "react"; +import React, { useCallback, useEffect, useRef, useState } from "react"; +import { + credentialCreateCall, + githubCopilotInitiateAuth, + githubCopilotCheckStatus, +} from "@/components/networking"; +import NotificationsManager from "../molecules/notifications_manager"; import ProviderSpecificFields from "../add_model/provider_specific_fields"; import { Providers, providerLogoMap } from "../provider_info_helpers"; -const { Link } = Typography; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useProviderFields } from "@/app/(dashboard)/hooks/providers/useProviderFields"; +const { Link, Text } = Typography; interface AddCredentialsModalProps { open: boolean; onCancel: () => void; onAddCredential: (values: any) => void; uploadProps: UploadProps; + initialCredentialName?: string; + initialProvider?: string; } -const AddCredentialsModal: React.FC = ({ open, onCancel, onAddCredential, uploadProps }) => { + +const AddCredentialsModal: React.FC = ({ open, onCancel, onAddCredential, uploadProps, initialCredentialName, initialProvider }) => { const [form] = Form.useForm(); const [selectedProvider, setSelectedProvider] = useState(Providers.OpenAI); + const { accessToken } = useAuthorized(); + const { data: providerMetadata } = useProviderFields(); + + // Device code flow state + const [deviceCodeState, setDeviceCodeState] = useState< + | { phase: "idle" } + | { phase: "polling"; deviceCode: string; userCode: string; verificationUri: string } + | { phase: "success"; credentialName: string } + | { phase: "error"; message: string } + >({ phase: "idle" }); + // Hold access_token in a ref — never rendered, never put in form fields + const accessTokenRef = useRef(null); + const pollingRef = useRef | null>(null); + + // Determine if the selected provider uses device_code auth flow + const isDeviceCodeProvider = React.useMemo(() => { + if (!providerMetadata) return false; + const info = providerMetadata.find( + (p) => + p.provider === selectedProvider || + p.provider_display_name === Providers[selectedProvider as keyof typeof Providers], + ); + return info?.auth_flow === "device_code"; + }, [selectedProvider, providerMetadata]); + + // Cleanup polling on unmount or modal close + const stopPolling = useCallback(() => { + if (pollingRef.current) { + clearInterval(pollingRef.current); + pollingRef.current = null; + } + }, []); + + useEffect(() => { + return stopPolling; + }, [stopPolling]); + + const handleCancel = () => { + stopPolling(); + setDeviceCodeState({ phase: "idle" }); + accessTokenRef.current = null; + onCancel(); + form.resetFields(); + }; const handleSubmit = (values: any) => { const filteredValues = Object.entries(values).reduce((acc, [key, value]) => { @@ -28,14 +83,168 @@ const AddCredentialsModal: React.FC = ({ open, onCance form.resetFields(); }; + const handleStartDeviceCode = async () => { + const credentialName = form.getFieldValue("credential_name"); + if (!credentialName) { + form.validateFields(["credential_name"]); + return; + } + if (!accessToken) return; + + try { + const result = await githubCopilotInitiateAuth(accessToken); + setDeviceCodeState({ + phase: "polling", + deviceCode: result.device_code, + userCode: result.user_code, + verificationUri: result.verification_uri, + }); + + if (!result.poll_interval_ms) throw new Error("GitHub initiate response missing poll_interval_ms"); + // Mutable baseline — ratchets up when GitHub sends slow_down so that + // subsequent normal "pending" responses keep using the increased interval. + let currentPollInterval = result.poll_interval_ms; + + // setTimeout-based loop so each poll fires only after the previous one + // completes, and slow_down's retry_after_ms is respected exactly. + const schedulePoll = (delayMs: number) => { + pollingRef.current = setTimeout(async () => { + try { + const status = await githubCopilotCheckStatus(accessToken, result.device_code); + console.log("[GH Copilot AddCredential] poll response:", status); + if (status.status === "complete" && status.access_token) { + stopPolling(); + accessTokenRef.current = status.access_token; + // Store as named credential + try { + await credentialCreateCall(accessToken, { + credential_name: credentialName, + credential_values: { api_key: status.access_token }, + credential_info: { custom_llm_provider: "github_copilot" }, + }); + setDeviceCodeState({ phase: "success", credentialName }); + } catch (e) { + console.error("[GH Copilot AddCredential] credentialCreateCall failed:", e); + NotificationsManager.error( + `Failed to save credential: ${e instanceof Error ? e.message : "Unknown error"}`, + ); + setDeviceCodeState({ phase: "error", message: "Failed to save credential" }); + } + } else if (status.status === "failed") { + stopPolling(); + setDeviceCodeState({ phase: "error", message: status.error || "Authorization failed" }); + } else { + // pending — ratchet up the baseline if GitHub requested slower + if (status.retry_after_ms != null) { + currentPollInterval = status.retry_after_ms; + } + schedulePoll(currentPollInterval); + } + } catch (e) { + console.error("[GH Copilot AddCredential] poll error:", e); + stopPolling(); + setDeviceCodeState({ phase: "error", message: "Failed to check authorization status" }); + } + }, delayMs); + }; + schedulePoll(currentPollInterval); + } catch { + setDeviceCodeState({ phase: "error", message: "Failed to start GitHub authorization" }); + } + }; + + const handleSuccessClose = () => { + stopPolling(); + setDeviceCodeState({ phase: "idle" }); + accessTokenRef.current = null; + onCancel(); + form.resetFields(); + }; + + const renderDeviceCodeFlow = () => { + switch (deviceCodeState.phase) { + case "idle": + return ( +
+ + GitHub Copilot uses OAuth Device Code authorization. Click below to start. + + +
+ ); + case "polling": + return ( +
+ + Enter this code on GitHub: + +
+ {deviceCodeState.userCode} +
+
+ +
+ + + Waiting for GitHub authorization... + + +
+ ); + case "success": + return ( +
+ + GitHub Copilot credential "{deviceCodeState.credentialName}" created successfully! + + +
+ ); + case "error": + return ( +
+ + {deviceCodeState.message} + + + +
+ ); + } + }; + return ( { - onCancel(); - form.resetFields(); - }} + onCancel={handleCancel} footer={null} width={600} > @@ -61,6 +270,10 @@ const AddCredentialsModal: React.FC = ({ open, onCance onChange={(value) => { setSelectedProvider(value as Providers); form.setFieldValue("custom_llm_provider", value); + // Reset device code state when provider changes + stopPolling(); + setDeviceCodeState({ phase: "idle" }); + accessTokenRef.current = null; }} > {Object.entries(Providers).map(([providerEnum, providerDisplayName]) => ( @@ -89,27 +302,30 @@ const AddCredentialsModal: React.FC = ({ open, onCance - + {isDeviceCodeProvider ? ( + renderDeviceCodeFlow() + ) : ( + <> + - {/* Modal Footer */} -
- - Need Help? - + {/* Modal Footer */} +
+ + Need Help? + -
- - -
-
+
+ + +
+
+ + )}
);