From b7432a4c447e9faec212ad15382d46805cbb01f9 Mon Sep 17 00:00:00 2001 From: Hunter Wittenborn Date: Mon, 30 Mar 2026 05:02:35 -0500 Subject: [PATCH] feat: add ChatGPT device code OAuth + fix token exchange across all endpoints MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add ChatGPT credential-based auth (device code OAuth flow) matching the GitHub Copilot pattern. The refresh token is stored as the credential's api_key and exchanged for an access token (JWT) at request time via OpenAI's OAuth endpoint — same flow as openai/codex. Key changes: - ChatGPT SSO endpoints (/credentials/chatgpt/initiate and /status) - ChatGPT authenticator credential mode with module-level token cache - Token exchange in responses validate_environment (single auth point for all endpoints: /responses, /chat/completions, /v1/messages, Gemini) - Raw httpx.Client for OAuth calls to avoid litellm wrapper interference - Generic device code UI dispatch for both GitHub Copilot and ChatGPT - Simplified main.py header blocks to static headers only (no auth) - GitHub Copilot chat validate_environment uses api_key directly Tested: all 8 endpoint/streaming combinations work (responses, chat, Anthropic, Gemini × streaming/non-streaming). Co-Authored-By: Claude Opus 4.6 (1M context) --- litellm/llms/chatgpt/authenticator.py | 171 ++++++++++-- litellm/llms/chatgpt/chat/transformation.py | 49 ++-- litellm/llms/chatgpt/common_utils.py | 23 +- .../llms/chatgpt/responses/transformation.py | 26 +- .../github_copilot/chat/transformation.py | 13 +- .../embedding/transformation.py | 7 +- .../responses/transformation.py | 9 +- litellm/main.py | 27 +- .../proxy/credential_endpoints/chatgpt_sso.py | 252 ++++++++++++++++++ .../github_copilot_sso.py | 202 ++++++++++++++ litellm/proxy/proxy_server.py | 4 + .../provider_create_fields.json | 8 + .../test_chatgpt_responses_transformation.py | 28 +- ...github_copilot_embedding_transformation.py | 28 +- ...github_copilot_responses_transformation.py | 23 +- .../proxy/credential_endpoints/__init__.py | 0 .../test_github_copilot_sso.py | 166 ++++++++++++ .../ModelsAndEndpointsView.tsx | 3 +- .../src/components/add_model/AddModelForm.tsx | 231 +++++++++++++++- .../components/add_model/add_model_tab.tsx | 3 + .../model_add/AddCredentialModal.tsx | 121 ++++++--- .../src/components/model_add/credentials.tsx | 85 ++++-- .../src/components/networking.tsx | 95 +++++++ .../src/components/provider_info_helpers.tsx | 3 + 24 files changed, 1372 insertions(+), 205 deletions(-) create mode 100644 litellm/proxy/credential_endpoints/chatgpt_sso.py create mode 100644 litellm/proxy/credential_endpoints/github_copilot_sso.py create mode 100644 tests/test_litellm/proxy/credential_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/credential_endpoints/test_github_copilot_sso.py diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index e35b04a3fb3..0308c5f76a9 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -3,6 +3,7 @@ import json import os import time from typing import Any, Dict, Optional +from urllib.parse import quote import httpx @@ -27,17 +28,34 @@ DEVICE_CODE_TIMEOUT_SECONDS = 15 * 60 DEVICE_CODE_COOLDOWN_SECONDS = 5 * 60 DEVICE_CODE_POLL_SLEEP_SECONDS = 5 +# Module-level cache for credential mode. +# Key = refresh_token, value = dict with access_token, expires_at, account_id. +_credential_token_cache: Dict[str, Dict[str, Any]] = {} + class Authenticator: - def __init__(self) -> None: - self.token_dir = os.getenv( - "CHATGPT_TOKEN_DIR", - os.path.expanduser("~/.config/litellm/chatgpt"), - ) - self.auth_file = os.path.join( - self.token_dir, os.getenv("CHATGPT_AUTH_FILE", "auth.json") - ) - self._ensure_token_dir() + def __init__(self, refresh_token: Optional[str] = None) -> None: + """Initialize the ChatGPT authenticator. + + Args: + refresh_token: If provided, the authenticator operates in + *credential mode* — it uses this refresh token to obtain + access tokens via the OpenAI OAuth refresh flow. When + ``None`` (the default), the existing file-based behaviour + is preserved. + """ + self._injected_refresh_token = refresh_token + + if refresh_token is None: + # File-based mode (backward compatible) + self.token_dir = os.getenv( + "CHATGPT_TOKEN_DIR", + os.path.expanduser("~/.config/litellm/chatgpt"), + ) + self.auth_file = os.path.join( + self.token_dir, os.getenv("CHATGPT_AUTH_FILE", "auth.json") + ) + self._ensure_token_dir() def get_api_base(self) -> str: return ( @@ -47,6 +65,18 @@ class Authenticator: ) def get_access_token(self) -> str: + """Get a valid access token, refreshing if necessary. + + In credential mode, uses the injected refresh_token with a + module-level cache to avoid redundant OAuth calls. + + In file-based mode, reads from the auth file and refreshes if + expired. Never triggers the interactive device code flow — + raises GetAccessTokenError with a docs link instead. + """ + if self._injected_refresh_token is not None: + return self._get_access_token_credential_mode() + auth_data = self._read_auth_file() if auth_data: access_token = auth_data.get("access_token") @@ -62,16 +92,27 @@ class Authenticator: "ChatGPT refresh token failed, re-login required: %s", exc ) - cooldown_remaining = self._get_device_code_cooldown_remaining(auth_data) - if cooldown_remaining > 0: - token = self._wait_for_access_token(cooldown_remaining) - if token: - return token - - tokens = self._login_device_code() - return tokens["access_token"] + raise GetAccessTokenError( + message=( + "No ChatGPT access token configured. " + "Use a named credential via the LiteLLM proxy or UI before making requests. " + "See: https://docs.litellm.ai/docs/providers/chatgpt" + ), + status_code=401, + ) def get_account_id(self) -> Optional[str]: + """Get the ChatGPT account ID. + + In credential mode, derives from the cached access/id token. + In file-based mode, reads from the auth file. + """ + if self._injected_refresh_token is not None: + cached = _credential_token_cache.get(self._injected_refresh_token) + if cached: + return cached.get("account_id") + return None + auth_data = self._read_auth_file() if not auth_data: return None @@ -86,6 +127,79 @@ class Authenticator: self._write_auth_file(auth_data) return derived + def _get_access_token_credential_mode(self) -> str: + """Get access token in credential mode using the injected refresh token.""" + cached = _credential_token_cache.get(self._injected_refresh_token) # type: ignore[arg-type] + if cached: + access_token = cached.get("access_token") + expires_at = cached.get("expires_at") + if ( + access_token + and expires_at + and time.time() < float(expires_at) - TOKEN_EXPIRY_SKEW_SECONDS + ): + return access_token + + try: + refreshed = self._refresh_tokens_credential_mode(self._injected_refresh_token) # type: ignore[arg-type] + access_token = refreshed["access_token"] + _credential_token_cache[self._injected_refresh_token] = { # type: ignore[index] + "access_token": access_token, + "expires_at": self._get_expires_at(access_token), + "account_id": self._extract_account_id( + refreshed.get("id_token") or access_token + ), + } + return access_token + except RefreshAccessTokenError as exc: + raise GetAccessTokenError( + message=f"Failed to refresh ChatGPT access token: {exc}", + status_code=401, + ) + + def _refresh_tokens_credential_mode(self, refresh_token: str) -> Dict[str, str]: + """Refresh tokens without writing to disk (credential mode). + + Uses a raw httpx.Client to avoid litellm's wrapper which can add + extra headers or interfere with OAuth token exchange calls. + """ + try: + with httpx.Client() as client: + resp = client.post( + CHATGPT_OAUTH_TOKEN_URL, + json={ + "client_id": CHATGPT_CLIENT_ID, + "grant_type": "refresh_token", + "refresh_token": refresh_token, + "scope": "openid profile email", + }, + ) + resp.raise_for_status() + data = resp.json() + except httpx.HTTPStatusError as exc: + raise RefreshAccessTokenError( + message=f"Refresh token failed: {exc}", + status_code=exc.response.status_code, + ) + except Exception as exc: + raise RefreshAccessTokenError( + message=f"Refresh token failed: {exc}", + status_code=400, + ) + + access_token = data.get("access_token") + if not access_token: + raise RefreshAccessTokenError( + message=f"Refresh response missing access_token: {data}", + status_code=400, + ) + + return { + "access_token": access_token, + "refresh_token": data.get("refresh_token", refresh_token), + "id_token": data.get("id_token", ""), + } + def _ensure_token_dir(self) -> None: if not os.path.exists(self.token_dir): os.makedirs(self.token_dir, exist_ok=True) @@ -125,7 +239,8 @@ class Authenticator: return int(exp) return None - def _decode_jwt_claims(self, token: str) -> Dict[str, Any]: + @staticmethod + def _decode_jwt_claims(token: str) -> Dict[str, Any]: try: parts = token.split(".") if len(parts) < 2: @@ -137,10 +252,11 @@ class Authenticator: except Exception: return {} - def _extract_account_id(self, token: Optional[str]) -> Optional[str]: + @staticmethod + def _extract_account_id(token: Optional[str]) -> Optional[str]: if not token: return None - claims = self._decode_jwt_claims(token) + claims = Authenticator._decode_jwt_claims(token) auth_claims = claims.get("https://api.openai.com/auth") if isinstance(auth_claims, dict): account_id = auth_claims.get("chatgpt_account_id") @@ -149,6 +265,7 @@ class Authenticator: return None def _login_device_code(self) -> Dict[str, str]: + """Interactive device code login (SDK CLI only, never called automatically).""" cooldown_remaining = self._get_device_code_cooldown_remaining( self._read_auth_file() ) @@ -172,7 +289,8 @@ class Authenticator: self._write_auth_file(auth_data) return tokens - def _request_device_code(self) -> Dict[str, str]: + @staticmethod + def _request_device_code() -> Dict[str, str]: try: client = _get_httpx_client() resp = client.post( @@ -257,16 +375,17 @@ class Authenticator: status_code=408, ) - def _exchange_code_for_tokens(self, code_data: Dict[str, str]) -> Dict[str, str]: + @staticmethod + def _exchange_code_for_tokens(code_data: Dict[str, str]) -> Dict[str, str]: try: client = _get_httpx_client() redirect_uri = f"{CHATGPT_AUTH_BASE}/deviceauth/callback" body = ( "grant_type=authorization_code" - f"&code={code_data['authorization_code']}" - f"&redirect_uri={redirect_uri}" - f"&client_id={CHATGPT_CLIENT_ID}" - f"&code_verifier={code_data['code_verifier']}" + f"&code={quote(code_data['authorization_code'], safe='')}" + f"&redirect_uri={quote(redirect_uri, safe='')}" + f"&client_id={quote(CHATGPT_CLIENT_ID, safe='')}" + f"&code_verifier={quote(code_data['code_verifier'], safe='')}" ) resp = client.post( CHATGPT_OAUTH_TOKEN_URL, diff --git a/litellm/llms/chatgpt/chat/transformation.py b/litellm/llms/chatgpt/chat/transformation.py index e6480398c7e..a6675fb053c 100644 --- a/litellm/llms/chatgpt/chat/transformation.py +++ b/litellm/llms/chatgpt/chat/transformation.py @@ -6,9 +6,12 @@ from litellm.types.llms.openai import AllMessageValues from ..authenticator import Authenticator from ..common_utils import ( + CHATGPT_API_BASE, GetAccessTokenError, + RefreshAccessTokenError, ensure_chatgpt_session_id, get_chatgpt_default_headers, + get_chatgpt_static_headers, ) from .streaming_utils import ChatGPTToolCallNormalizer @@ -21,7 +24,6 @@ class ChatGPTConfig(OpenAIConfig): custom_llm_provider: str = "openai", ) -> None: super().__init__() - self.authenticator = Authenticator() def _get_openai_compatible_provider_info( self, @@ -30,16 +32,15 @@ class ChatGPTConfig(OpenAIConfig): api_key: Optional[str], custom_llm_provider: str, ) -> Tuple[Optional[str], Optional[str], str]: - dynamic_api_base = self.authenticator.get_api_base() - try: - dynamic_api_key = self.authenticator.get_access_token() - except GetAccessTokenError as e: - raise AuthenticationError( - model=model, - llm_provider=custom_llm_provider, - message=str(e), - ) - return dynamic_api_base, dynamic_api_key, custom_llm_provider + if not api_key: + return CHATGPT_API_BASE, None, custom_llm_provider + # Pass the refresh token through without exchanging. ChatGPT models + # use mode="responses", so /chat/completions gets bridged to the + # responses API (via responses_api_bridge) which uses the token + # directly. Exchanging here would fail because get_llm_provider() + # runs before the bridge check. + dynamic_api_base = api_base or CHATGPT_API_BASE + return dynamic_api_base, api_key, custom_llm_provider def validate_environment( self, @@ -55,12 +56,26 @@ class ChatGPTConfig(OpenAIConfig): headers, model, messages, optional_params, litellm_params, api_key, api_base ) - account_id = self.authenticator.get_account_id() - session_id = ensure_chatgpt_session_id(litellm_params) - default_headers = get_chatgpt_default_headers( - api_key or "", account_id, session_id - ) - return {**default_headers, **validated_headers} + # Always add static ChatGPT headers (originator, user-agent, etc.) + validated_headers = {**get_chatgpt_static_headers(), **validated_headers} + + # api_key is the refresh token. Exchange for access token (JWT) + # via OpenAI OAuth — same flow as openai/codex. Cached at module + # level by the authenticator. + if api_key: + try: + authenticator = Authenticator(refresh_token=api_key) + access_token = authenticator.get_access_token() + account_id = authenticator.get_account_id() + session_id = ensure_chatgpt_session_id(litellm_params) + default_headers = get_chatgpt_default_headers( + access_token, account_id, session_id + ) + validated_headers = {**default_headers, **validated_headers} + except (GetAccessTokenError, RefreshAccessTokenError): + pass + + return validated_headers def post_stream_processing(self, stream: Any) -> Any: return ChatGPTToolCallNormalizer(stream) diff --git a/litellm/llms/chatgpt/common_utils.py b/litellm/llms/chatgpt/common_utils.py index 9cbcd6a4f46..44425f66305 100644 --- a/litellm/llms/chatgpt/common_utils.py +++ b/litellm/llms/chatgpt/common_utils.py @@ -230,19 +230,30 @@ def get_chatgpt_user_agent(originator: str) -> str: return _safe_header_value(candidate) or DEFAULT_USER_AGENT +def get_chatgpt_static_headers() -> dict: + """Get static headers required by the ChatGPT API. + + These headers (originator, user-agent, etc.) must be present on every + request regardless of whether the access token has been resolved yet. + """ + originator = get_chatgpt_originator() + user_agent = get_chatgpt_user_agent(originator) + return { + "content-type": "application/json", + "accept": "text/event-stream", + "originator": originator, + "user-agent": user_agent, + } + + def get_chatgpt_default_headers( access_token: str, account_id: Optional[str], session_id: Optional[str] = None, ) -> dict: - originator = get_chatgpt_originator() - user_agent = get_chatgpt_user_agent(originator) headers = { + **get_chatgpt_static_headers(), "Authorization": f"Bearer {access_token}", - "content-type": "application/json", - "accept": "text/event-stream", - "originator": originator, - "user-agent": user_agent, } if session_id: headers["session_id"] = session_id diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index 3c59ca16581..ad173e7de6c 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -21,6 +21,7 @@ from ..authenticator import Authenticator from ..common_utils import ( CHATGPT_API_BASE, GetAccessTokenError, + RefreshAccessTokenError, ensure_chatgpt_session_id, get_chatgpt_default_headers, get_chatgpt_default_instructions, @@ -30,7 +31,6 @@ from ..common_utils import ( class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): def __init__(self) -> None: super().__init__() - self.authenticator = Authenticator() @property def custom_llm_provider(self) -> LlmProviders: @@ -42,16 +42,32 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): model: str, litellm_params: Optional[GenericLiteLLMParams], ) -> dict: + 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="chatgpt", + message="ChatGPT API key (refresh token) is required. Please authenticate via the LiteLLM UI.", + ) + + # _api_key is the refresh token from the credential. Exchange it + # for an access token (JWT) via OpenAI's OAuth endpoint — same + # flow as openai/codex. The authenticator caches the access token + # at module level so subsequent requests reuse it until expiry. try: - access_token = self.authenticator.get_access_token() - except GetAccessTokenError as e: + authenticator = Authenticator(refresh_token=_api_key) + access_token = authenticator.get_access_token() + account_id = authenticator.get_account_id() + except (GetAccessTokenError, RefreshAccessTokenError) as e: raise AuthenticationError( model=model, llm_provider="chatgpt", message=str(e), ) - account_id = self.authenticator.get_account_id() session_id = ensure_chatgpt_session_id(litellm_params) default_headers = get_chatgpt_default_headers( access_token, account_id, session_id @@ -197,7 +213,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): api_base: Optional[str], litellm_params: dict, ) -> str: - api_base = api_base or self.authenticator.get_api_base() or CHATGPT_API_BASE + api_base = api_base or CHATGPT_API_BASE api_base = api_base.rstrip("/") return f"{api_base}/responses" diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 8e35ea5f6c7..1f8829bebbe 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -94,15 +94,12 @@ class GithubCopilotConfig(OpenAIConfig): # 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. + # api_key at this point is already the resolved copilot inference + # token (exchanged by _get_openai_compatible_provider_info or + # main.py). Just use it directly — no re-authentication needed. 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 + copilot_headers = get_copilot_default_headers(api_key) + validated_headers = {**copilot_headers, **validated_headers} # Add X-Initiator header based on message roles initiator = self._determine_initiator(messages) diff --git a/litellm/llms/github_copilot/embedding/transformation.py b/litellm/llms/github_copilot/embedding/transformation.py index 63ee6ca1e81..68ce93550cb 100644 --- a/litellm/llms/github_copilot/embedding/transformation.py +++ b/litellm/llms/github_copilot/embedding/transformation.py @@ -64,10 +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 + # api_key at this point is already the resolved copilot inference + # token (exchanged by the provider info step or main.py). + # Just use it directly — no re-authentication needed. default_headers = get_copilot_default_headers(api_key) # Merge with existing headers (user's extra_headers take priority) diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index ee5f95baacd..fb31920f5bb 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -114,11 +114,10 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): 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) + # api_key at this point is already the resolved copilot inference + # token (exchanged by _get_openai_compatible_provider_info or + # main.py). Just use it directly — no re-authentication needed. + default_headers = get_copilot_default_headers(_api_key) # Merge with existing headers (user's extra_headers take priority) merged_headers = {**default_headers, **headers} diff --git a/litellm/main.py b/litellm/main.py index 97b24c4ad81..56b46c09015 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2605,29 +2605,30 @@ def completion( # type: ignore # noqa: PLR0915 headers = headers or litellm.headers - # Add GitHub Copilot headers (same as /responses endpoint does) + # Add static headers for OAuth providers. Token exchange is + # already done by _get_openai_compatible_provider_info; auth + # headers are set by validate_environment. These blocks only + # ensure the provider-required non-auth headers are present. 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, ) - # 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 + if custom_llm_provider == "chatgpt": + from litellm.llms.chatgpt.common_utils import ( + get_chatgpt_static_headers, + ) + + chatgpt_headers = get_chatgpt_static_headers() + if extra_headers: + chatgpt_headers.update(extra_headers) + extra_headers = chatgpt_headers + if extra_headers is not None: optional_params["extra_headers"] = extra_headers diff --git a/litellm/proxy/credential_endpoints/chatgpt_sso.py b/litellm/proxy/credential_endpoints/chatgpt_sso.py new file mode 100644 index 00000000000..0b9e7a70106 --- /dev/null +++ b/litellm/proxy/credential_endpoints/chatgpt_sso.py @@ -0,0 +1,252 @@ +""" +ChatGPT SSO endpoints for device code OAuth flow. + +The backend is fully stateless during the auth flow. The client holds the +device_auth_id and user_code between requests — nothing is stored in the DB +until the tokens are successfully obtained. + +Flow: + 1. POST /credentials/chatgpt/initiate + Returns device_auth_id + user_code + verification_uri + poll_interval_ms. + No DB write. + + 2. POST /credentials/chatgpt/status (client calls repeatedly) + Body: { device_auth_id, user_code } + Polls OpenAI's device token endpoint. When authorization_code is received, + exchanges it for tokens automatically. + Returns: { status: "pending"|"complete"|"failed", refresh_token?, account_id?, error? } + On complete, returns the refresh_token — client stores it as a named credential. + No DB write — client decides what to do with the token. + +Design: no session_id, no pending/failed DB rows, no cancel endpoint. +Multi-worker safe: device codes are validated by OpenAI, not by LiteLLM. +""" + +from typing import Literal, Optional +from urllib.parse import quote + +import httpx as _httpx +from fastapi import APIRouter, Depends, HTTPException, Request, Response +from pydantic import BaseModel + +from litellm._logging import verbose_proxy_logger +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.types.llms.custom_http import httpxSpecialProvider +from litellm.llms.chatgpt.authenticator import Authenticator +from litellm.llms.chatgpt.common_utils import ( + CHATGPT_AUTH_BASE, + CHATGPT_CLIENT_ID, + CHATGPT_DEVICE_CODE_URL, + CHATGPT_DEVICE_TOKEN_URL, + CHATGPT_DEVICE_VERIFY_URL, + CHATGPT_OAUTH_TOKEN_URL, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +# --------------------------------------------------------------------------- +# Request / response models +# --------------------------------------------------------------------------- + + +class InitiateResponse(BaseModel): + device_auth_id: str + user_code: str + verification_uri: str + poll_interval_ms: int + + +class StatusRequest(BaseModel): + device_auth_id: str + user_code: str + + +class StatusResponse(BaseModel): + status: Literal["pending", "complete", "failed"] + refresh_token: Optional[str] = None # complete only + account_id: Optional[str] = None # complete only + error: Optional[str] = None # failed only + + +# --------------------------------------------------------------------------- +# Endpoints +# --------------------------------------------------------------------------- + + +@router.post( + "/credentials/chatgpt/initiate", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=InitiateResponse, +) +async def chatgpt_initiate( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Start the ChatGPT Device Code flow. + + Calls OpenAI's device code endpoint and returns device_auth_id, user_code, + verification_uri, and poll_interval_ms to the client. Nothing is stored + in the database. + """ + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER) + try: + resp = await async_client.post( + CHATGPT_DEVICE_CODE_URL, + json={"client_id": CHATGPT_CLIENT_ID}, + ) + resp.raise_for_status() + except Exception as e: + verbose_proxy_logger.error(f"ChatGPT device code request failed: {e}") + raise HTTPException( + status_code=502, + detail={"error": f"ChatGPT device code request failed: {e}"}, + ) + resp_json = resp.json() + + device_auth_id = resp_json.get("device_auth_id") + user_code = resp_json.get("user_code") or resp_json.get("usercode") + poll_interval_seconds = int(resp_json.get("interval", 5)) + + if not all([device_auth_id, user_code]): + raise HTTPException( + status_code=502, + detail={"error": "OpenAI response missing required fields"}, + ) + + return InitiateResponse( + device_auth_id=device_auth_id, + user_code=user_code, + verification_uri=CHATGPT_DEVICE_VERIFY_URL, + poll_interval_ms=poll_interval_seconds * 1000, + ) + + +@router.post( + "/credentials/chatgpt/status", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=StatusResponse, +) +async def chatgpt_status( + request: Request, + fastapi_response: Response, + body: StatusRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Poll the ChatGPT Device Code flow for completion. + + Makes a single attempt to check if the user has authorized the device code. + OpenAI's device token endpoint returns 403/404 while pending. + When the authorization_code is received, this endpoint automatically + exchanges it for tokens and returns the refresh_token. + + - 403/404 → {"status": "pending"} + - success → exchanges code → {"status": "complete", "refresh_token": "...", "account_id": "..."} + - any other error → {"status": "failed", "error": "..."} + + No DB write is made — the client holds the refresh_token and decides + what to do with it (store as named credential or use inline). + """ + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER) + + # Step 1: Poll for authorization code + # OpenAI returns 403/404 while the user hasn't authorized yet. + # The litellm httpx wrapper calls raise_for_status() automatically, + # so we catch HTTPStatusError and check the status code ourselves. + try: + resp = await async_client.post( + CHATGPT_DEVICE_TOKEN_URL, + json={ + "device_auth_id": body.device_auth_id, + "user_code": body.user_code, + }, + ) + except _httpx.HTTPStatusError as e: + if e.response.status_code in (403, 404): + verbose_proxy_logger.debug("ChatGPT: authorization_pending") + return StatusResponse(status="pending") + verbose_proxy_logger.error(f"ChatGPT token poll failed: {e}") + return StatusResponse(status="failed", error=str(e)) + except Exception as e: + verbose_proxy_logger.error(f"ChatGPT token poll failed: {e}") + return StatusResponse(status="failed", error=str(e)) + + if resp.status_code in (403, 404): + verbose_proxy_logger.debug("ChatGPT: authorization_pending") + return StatusResponse(status="pending") + + if resp.status_code != 200: + error_text = resp.text[:200] if resp.text else f"HTTP {resp.status_code}" + verbose_proxy_logger.warning(f"ChatGPT device code poll unexpected response: {error_text}") + return StatusResponse(status="failed", error=error_text) + + resp_json = resp.json() + verbose_proxy_logger.debug(f"ChatGPT token poll response keys: {list(resp_json.keys())}") + + authorization_code = resp_json.get("authorization_code") + code_verifier = resp_json.get("code_verifier") + if not authorization_code or not code_verifier: + verbose_proxy_logger.debug("ChatGPT: response 200 but missing authorization_code/code_verifier") + return StatusResponse(status="pending") + + # Step 2: Exchange authorization code for tokens + # Use a plain httpx client for the exchange — the litellm wrapper adds + # raise_for_status hooks that swallow the error body we need for debugging, + # and may add headers that interfere with the OAuth token endpoint. + import httpx as _raw_httpx + + try: + redirect_uri = f"{CHATGPT_AUTH_BASE}/deviceauth/callback" + exchange_body = ( + "grant_type=authorization_code" + f"&code={quote(authorization_code, safe='')}" + f"&redirect_uri={quote(redirect_uri, safe='')}" + f"&client_id={quote(CHATGPT_CLIENT_ID, safe='')}" + f"&code_verifier={quote(code_verifier, safe='')}" + ) + async with _raw_httpx.AsyncClient() as raw_client: + token_resp = await raw_client.post( + CHATGPT_OAUTH_TOKEN_URL, + headers={"Content-Type": "application/x-www-form-urlencoded"}, + content=exchange_body, + ) + if not token_resp.is_success: + verbose_proxy_logger.error( + f"ChatGPT token exchange returned {token_resp.status_code}: {token_resp.text[:500]}" + ) + return StatusResponse(status="failed", error=f"Token exchange failed: HTTP {token_resp.status_code}") + token_data = token_resp.json() + except Exception as e: + verbose_proxy_logger.error(f"ChatGPT token exchange failed: {e}") + return StatusResponse(status="failed", error=f"Token exchange failed: {e}") + + refresh_token = token_data.get("refresh_token") + verbose_proxy_logger.info( + f"ChatGPT token exchange result: keys={list(token_data.keys())}, " + f"has_refresh={bool(refresh_token)}, refresh_len={len(refresh_token or '')}, " + f"has_access={bool(token_data.get('access_token'))}" + ) + if not refresh_token: + return StatusResponse( + status="failed", + error="Token exchange response missing refresh_token", + ) + + # Derive account_id from id_token or access_token + id_token = token_data.get("id_token") + access_token = token_data.get("access_token") + account_id = Authenticator._extract_account_id(id_token or access_token) + + verbose_proxy_logger.info("ChatGPT device code flow completed successfully") + return StatusResponse( + status="complete", + refresh_token=refresh_token, + account_id=account_id, + ) diff --git a/litellm/proxy/credential_endpoints/github_copilot_sso.py b/litellm/proxy/credential_endpoints/github_copilot_sso.py new file mode 100644 index 00000000000..4e5bc1bdbcc --- /dev/null +++ b/litellm/proxy/credential_endpoints/github_copilot_sso.py @@ -0,0 +1,202 @@ +""" +GitHub Copilot SSO endpoints for device code OAuth flow. + +The backend is fully stateless during the auth flow. The client holds the +device_code between requests — nothing is stored in the DB until the GitHub +token is successfully validated. + +Flow: + 1. POST /credentials/github_copilot/initiate + Returns device_code + user_code + verification_uri + poll_interval_ms + expires_in. + No DB write. + + 2. POST /credentials/github_copilot/status (client calls repeatedly) + Body: { device_code } + Returns: { status: "pending"|"complete"|"failed", access_token?, retry_after_ms?, error? } + On slow_down, returns retry_after_ms — client must wait that long before retrying. + No DB write — client decides what to do with the token. + +Design: no session_id, no pending/failed DB rows, no cancel endpoint. +Multi-worker safe: device_code is validated by GitHub, not by LiteLLM. +""" + +from typing import Literal, Optional + +from fastapi import APIRouter, Depends, HTTPException, Request, Response +from pydantic import BaseModel + +from litellm._logging import verbose_proxy_logger +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.types.llms.custom_http import httpxSpecialProvider +from litellm.llms.github_copilot.authenticator import ( + GITHUB_ACCESS_TOKEN_URL, + GITHUB_CLIENT_ID, + GITHUB_DEVICE_CODE_URL, + Authenticator, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +# --------------------------------------------------------------------------- +# Request / response models +# --------------------------------------------------------------------------- + + +class InitiateResponse(BaseModel): + device_code: str + user_code: str + verification_uri: str + poll_interval_ms: int + expires_in: int + + +class StatusRequest(BaseModel): + device_code: str + + +class StatusResponse(BaseModel): + status: Literal["pending", "complete", "failed"] + access_token: Optional[str] = None # complete only + retry_after_ms: Optional[int] = None # pending + slow_down only; client must wait this long before retrying + error: Optional[str] = None # failed only + + +# --------------------------------------------------------------------------- +# Endpoints +# --------------------------------------------------------------------------- + + +@router.post( + "/credentials/github_copilot/initiate", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=InitiateResponse, +) +async def github_copilot_initiate( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Start the GitHub Device Code flow. + + Calls GitHub and returns device_code, user_code, verification_uri, + poll_interval_ms, and expires_in to the client. Nothing is stored in + the database. + """ + headers = Authenticator.get_github_headers() + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER) + try: + resp = await async_client.post( + GITHUB_DEVICE_CODE_URL, + headers=headers, + json={"client_id": GITHUB_CLIENT_ID, "scope": "read:user"}, + ) + resp.raise_for_status() + except Exception as e: + verbose_proxy_logger.error(f"GitHub device code request failed: {e}") + raise HTTPException( + status_code=502, + detail={"error": f"GitHub device code request failed: {e}"}, + ) + resp_json = resp.json() + + device_code = resp_json.get("device_code") + user_code = resp_json.get("user_code") + verification_uri = resp_json.get("verification_uri") + poll_interval_seconds = int(resp_json.get("interval", 5)) + expires_in = int(resp_json.get("expires_in", 900)) + + if not all([device_code, user_code, verification_uri]): + raise HTTPException( + status_code=502, + detail={"error": "GitHub response missing required fields"}, + ) + + return InitiateResponse( + device_code=device_code, + user_code=user_code, + verification_uri=verification_uri, + poll_interval_ms=poll_interval_seconds * 1000, + expires_in=expires_in, + ) + + +@router.post( + "/credentials/github_copilot/status", + dependencies=[Depends(user_api_key_auth)], + tags=["credential management"], + response_model=StatusResponse, +) +async def github_copilot_status( + request: Request, + fastapi_response: Response, + body: StatusRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Poll the GitHub Device Code flow for completion. + + Makes a single attempt to exchange the device_code for an access_token. + + - authorization_pending → {"status": "pending"} + - slow_down → {"status": "pending", "retry_after_ms": } + The client MUST wait retry_after_ms before calling again, or GitHub will + keep increasing the required interval. + - success → {"status": "complete", "access_token": "ghu_xxx"} + - any other error → {"status": "failed", "error": "..."} + + No DB write is made — the client holds the access_token and decides + what to do with it (store as named credential or use inline). + """ + headers = Authenticator.get_github_headers() + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER) + try: + resp = await async_client.post( + GITHUB_ACCESS_TOKEN_URL, + headers=headers, + json={ + "client_id": GITHUB_CLIENT_ID, + "device_code": body.device_code, + "grant_type": "urn:ietf:params:oauth:grant-type:device_code", + }, + ) + resp.raise_for_status() + resp_json = resp.json() + except Exception as e: + verbose_proxy_logger.error(f"GitHub token poll failed: {e}") + return StatusResponse(status="failed", error=str(e)) + + verbose_proxy_logger.debug(f"GitHub token poll response: {resp_json}") + + if "access_token" in resp_json: + verbose_proxy_logger.info("GitHub Copilot device code flow completed successfully") + return StatusResponse(status="complete", access_token=resp_json["access_token"]) + + error_code = resp_json.get("error", "") + + if error_code == "slow_down": + interval = resp_json.get("interval") + if interval is None: + verbose_proxy_logger.warning("GitHub slow_down with no interval — treating as failed") + return StatusResponse(status="failed", error="GitHub slow_down with no interval reported") + interval_ms = int(interval) * 1000 + verbose_proxy_logger.warning( + f"GitHub slow_down — telling client to wait {interval_ms}ms before retrying" + ) + return StatusResponse(status="pending", retry_after_ms=interval_ms) + + if error_code == "authorization_pending": + verbose_proxy_logger.debug("GitHub Copilot: authorization_pending") + return StatusResponse(status="pending") + + # Any other error (expired_token, access_denied, etc.) + error_description = resp_json.get("error_description", error_code) + verbose_proxy_logger.warning( + f"GitHub Copilot device code flow failed: error={error_code!r}, " + f"description={error_description!r}" + ) + return StatusResponse(status="failed", error=error_description) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 806842a6495..142d10ccf89 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.chatgpt_sso import ( + router as chatgpt_sso_router, +) from litellm.proxy.credential_endpoints.github_copilot_sso import ( router as github_copilot_sso_router, ) @@ -13450,6 +13453,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(chatgpt_sso_router) app.include_router(github_copilot_sso_router) app.include_router(llm_passthrough_router) app.include_router(webrtc_router) diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index 3856ade5274..3b08d8db998 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -1170,6 +1170,14 @@ ], "default_model_placeholder": "gpt-3.5-turbo" }, + { + "provider": "CHATGPT", + "provider_display_name": "ChatGPT", + "litellm_provider": "chatgpt", + "auth_flow": "device_code", + "credential_fields": [], + "default_model_placeholder": "codex-mini" + }, { "provider": "GITHUB_COPILOT", "provider_display_name": "Github Copilot", diff --git a/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py b/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py index c0a0927b7d1..18c059ab207 100644 --- a/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py +++ b/tests/test_litellm/llms/chatgpt/responses/test_chatgpt_responses_transformation.py @@ -41,16 +41,11 @@ class TestChatGPTResponsesAPITransformation: assert isinstance(config, ChatGPTResponsesAPIConfig) assert config.custom_llm_provider == LlmProviders.CHATGPT - @patch("litellm.llms.chatgpt.responses.transformation.Authenticator") - def test_chatgpt_responses_endpoint_url(self, mock_authenticator_class): - mock_auth_instance = MagicMock() - mock_auth_instance.get_api_base.return_value = "https://chatgpt.example.com" - mock_authenticator_class.return_value = mock_auth_instance - + def test_chatgpt_responses_endpoint_url(self): config = ChatGPTResponsesAPIConfig() url = config.get_complete_url(api_base=None, litellm_params={}) - assert url == "https://chatgpt.example.com/responses" + assert url == "https://chatgpt.com/backend-api/codex/responses" custom_url = config.get_complete_url( api_base="https://custom.chatgpt.com", litellm_params={} @@ -64,21 +59,26 @@ class TestChatGPTResponsesAPITransformation: @patch("litellm.llms.chatgpt.responses.transformation.Authenticator") def test_validate_environment_headers(self, mock_authenticator_class): - mock_auth_instance = MagicMock() - mock_auth_instance.get_access_token.return_value = "access-123" - mock_auth_instance.get_account_id.return_value = "acct-123" - mock_authenticator_class.return_value = mock_auth_instance + mock_auth = MagicMock() + mock_auth.get_access_token.return_value = "test-access-token" + mock_auth.get_account_id.return_value = "account-123" + mock_authenticator_class.return_value = mock_auth config = ChatGPTResponsesAPIConfig() - litellm_params = GenericLiteLLMParams(litellm_session_id="session-123") + litellm_params = GenericLiteLLMParams( + api_key="test-refresh-token", + litellm_session_id="session-123", + ) headers = config.validate_environment( headers={"originator": "custom-origin"}, model="gpt-5.2", litellm_params=litellm_params, ) - assert headers["Authorization"] == "Bearer access-123" - assert headers["ChatGPT-Account-Id"] == "acct-123" + # Refresh token is exchanged for access token via Authenticator + mock_authenticator_class.assert_called_once_with(refresh_token="test-refresh-token") + mock_auth.get_access_token.assert_called_once() + assert headers["Authorization"] == "Bearer test-access-token" assert headers["originator"] == "custom-origin" assert headers["content-type"] == "application/json" assert headers["accept"] == "text/event-stream" 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 da4a9129bd3..3b1ebc1a682 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,14 +10,8 @@ from litellm.exceptions import AuthenticationError from litellm.llms.github_copilot.embedding.transformation import GithubCopilotEmbeddingConfig from litellm.llms.github_copilot.common_utils import GetAPIKeyError -@patch("litellm.llms.github_copilot.embedding.transformation.Authenticator") -def test_github_copilot_embedding_config_validate_environment(mock_authenticator_class): +def test_github_copilot_embedding_config_validate_environment(): """Test the GitHub Copilot embedding configuration environment validation.""" - mock_api_key = "gh.test-key-123456789" - mock_auth_instance = MagicMock() - mock_auth_instance.get_api_key.return_value = mock_api_key - mock_authenticator_class.return_value = mock_auth_instance - config = GithubCopilotEmbeddingConfig() model = "github_copilot/text-embedding-3-small" @@ -30,7 +24,8 @@ def test_github_copilot_embedding_config_validate_environment(mock_authenticator api_key="gh-access-token", ) - assert validated_headers["Authorization"] == f"Bearer {mock_api_key}" + # api_key is used directly — no re-exchange in validate_environment + assert validated_headers["Authorization"] == "Bearer gh-access-token" assert validated_headers["copilot-integration-id"] == "vscode-chat" assert validated_headers["editor-version"] == "vscode/1.95.0" assert "x-request-id" in validated_headers @@ -47,21 +42,8 @@ def test_github_copilot_embedding_config_validate_environment(mock_authenticator ) 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) + # Auth failures now happen upstream (in _get_openai_compatible_provider_info + # or main.py), not in validate_environment. No re-exchange test needed. @patch("litellm.llms.github_copilot.embedding.transformation.Authenticator") def test_github_copilot_embedding_config_get_complete_url(mock_authenticator_class): 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 ce5da4fc3cc..49fb2f2e186 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 @@ -82,22 +82,16 @@ class TestGithubCopilotResponsesAPITransformation: "Should handle trailing slash" ) - @patch("litellm.llms.github_copilot.responses.transformation.Authenticator") - def test_validate_environment_default_headers(self, mock_authenticator_class): + def test_validate_environment_default_headers(self): """Test that validate_environment generates correct default headers""" - # Mock the authenticator - mock_auth_instance = MagicMock() - mock_auth_instance.get_api_key.return_value = "test-api-key-123" - mock_authenticator_class.return_value = mock_auth_instance - config = GithubCopilotResponsesAPIConfig() headers = config.validate_environment( headers={}, model="gpt-5.1-codex", litellm_params={"api_key": "gh-access-token"} ) - # Check required headers - assert headers["Authorization"] == "Bearer test-api-key-123" + # api_key is used directly — no re-exchange in validate_environment + assert headers["Authorization"] == "Bearer gh-access-token" assert headers["content-type"] == "application/json" assert headers["copilot-integration-id"] == "vscode-chat" assert headers["editor-version"] == "vscode/1.95.0" @@ -107,13 +101,8 @@ class TestGithubCopilotResponsesAPITransformation: assert headers["x-github-api-version"] == "2025-04-01" assert "x-request-id" in headers - @patch("litellm.llms.github_copilot.responses.transformation.Authenticator") - def test_validate_environment_user_headers_override(self, mock_authenticator_class): + def test_validate_environment_user_headers_override(self): """Test that user-provided headers override default headers""" - mock_auth_instance = MagicMock() - mock_auth_instance.get_api_key.return_value = "test-api-key-123" - mock_authenticator_class.return_value = mock_auth_instance - config = GithubCopilotResponsesAPIConfig() custom_headers = { @@ -129,8 +118,8 @@ class TestGithubCopilotResponsesAPITransformation: assert headers["editor-version"] == "custom/2.0.0" # Custom header should be preserved assert headers["custom-header"] == "custom-value" - # Default headers should still be present - assert headers["Authorization"] == "Bearer test-api-key-123" + # api_key used directly + assert headers["Authorization"] == "Bearer gh-access-token" def test_get_initiator_with_assistant_role(self): """Test _get_initiator returns 'agent' for assistant role""" diff --git a/tests/test_litellm/proxy/credential_endpoints/__init__.py b/tests/test_litellm/proxy/credential_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/credential_endpoints/test_github_copilot_sso.py b/tests/test_litellm/proxy/credential_endpoints/test_github_copilot_sso.py new file mode 100644 index 00000000000..cc2a777d36d --- /dev/null +++ b/tests/test_litellm/proxy/credential_endpoints/test_github_copilot_sso.py @@ -0,0 +1,166 @@ +""" +Tests for the GitHub Copilot SSO (device code) credential endpoints. + +The backend is fully stateless — no DB writes during the auth flow. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +ASYNC_HTTP_PATCH = "litellm.proxy.credential_endpoints.github_copilot_sso.get_async_httpx_client" + + +def _make_async_client(json_data): + """Build a mock client whose .post() returns a single response.""" + mock_resp = AsyncMock() + mock_resp.raise_for_status = MagicMock() + mock_resp.json = MagicMock(return_value=json_data) + + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_resp) + return mock_client + + +# --------------------------------------------------------------------------- +# Tests for initiate endpoint +# --------------------------------------------------------------------------- + + +class TestGithubCopilotInitiate: + @pytest.mark.asyncio + async def test_initiate_success(self): + """POST /initiate returns device_code, user_code, verification_uri, poll_interval_ms, expires_in.""" + ctx = _make_async_client({ + "device_code": "test-device-code-123", + "user_code": "ABCD-1234", + "verification_uri": "https://github.com/login/device", + "interval": 5, + "expires_in": 900, + }) + + with patch(ASYNC_HTTP_PATCH, return_value=ctx): + from litellm.proxy.credential_endpoints.github_copilot_sso import ( + github_copilot_initiate, + ) + + result = await github_copilot_initiate(MagicMock(), MagicMock(), MagicMock()) + assert result.device_code == "test-device-code-123" + assert result.user_code == "ABCD-1234" + assert result.verification_uri == "https://github.com/login/device" + assert result.poll_interval_ms == 5000 + assert result.expires_in == 900 + + @pytest.mark.asyncio + async def test_initiate_github_error(self): + """POST /initiate returns 502 when GitHub device code API fails.""" + mock_client = AsyncMock() + mock_client.post = AsyncMock(side_effect=Exception("network error")) + + with patch(ASYNC_HTTP_PATCH, return_value=mock_client): + from fastapi import HTTPException + + from litellm.proxy.credential_endpoints.github_copilot_sso import ( + github_copilot_initiate, + ) + + with pytest.raises(HTTPException) as exc_info: + await github_copilot_initiate(MagicMock(), MagicMock(), MagicMock()) + assert exc_info.value.status_code == 502 + + +# --------------------------------------------------------------------------- +# Tests for status endpoint +# --------------------------------------------------------------------------- + + +class TestGithubCopilotStatus: + @pytest.mark.asyncio + async def test_status_pending(self): + """POST /status returns pending when GitHub says authorization_pending.""" + ctx = _make_async_client({"error": "authorization_pending"}) + + with patch(ASYNC_HTTP_PATCH, return_value=ctx): + from litellm.proxy.credential_endpoints.github_copilot_sso import ( + StatusRequest, + github_copilot_status, + ) + + result = await github_copilot_status( + MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock() + ) + assert result.status == "pending" + assert result.access_token is None + assert result.retry_after_ms is None + + @pytest.mark.asyncio + async def test_status_slow_down(self): + """POST /status returns pending with retry_after_ms when GitHub says slow_down.""" + ctx = _make_async_client({"error": "slow_down", "interval": 10}) + + with patch(ASYNC_HTTP_PATCH, return_value=ctx): + from litellm.proxy.credential_endpoints.github_copilot_sso import ( + StatusRequest, + github_copilot_status, + ) + + result = await github_copilot_status( + MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock() + ) + assert result.status == "pending" + assert result.retry_after_ms == 10_000 + + @pytest.mark.asyncio + async def test_status_complete(self): + """POST /status returns access_token on success.""" + ctx = _make_async_client({"access_token": "ghu_abc123"}) + + with patch(ASYNC_HTTP_PATCH, return_value=ctx): + from litellm.proxy.credential_endpoints.github_copilot_sso import ( + StatusRequest, + github_copilot_status, + ) + + result = await github_copilot_status( + MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock() + ) + assert result.status == "complete" + assert result.access_token == "ghu_abc123" + assert result.error is None + + @pytest.mark.asyncio + async def test_status_slow_down_missing_interval(self): + """POST /status returns failed when GitHub slow_down has no interval field.""" + ctx = _make_async_client({"error": "slow_down"}) + + with patch(ASYNC_HTTP_PATCH, return_value=ctx): + from litellm.proxy.credential_endpoints.github_copilot_sso import ( + StatusRequest, + github_copilot_status, + ) + + result = await github_copilot_status( + MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock() + ) + assert result.status == "failed" + assert result.error is not None + + @pytest.mark.asyncio + async def test_status_failed(self): + """POST /status returns failed when GitHub returns expired_token.""" + ctx = _make_async_client({ + "error": "expired_token", + "error_description": "The device code has expired.", + }) + + with patch(ASYNC_HTTP_PATCH, return_value=ctx): + from litellm.proxy.credential_endpoints.github_copilot_sso import ( + StatusRequest, + github_copilot_status, + ) + + result = await github_copilot_status( + MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock() + ) + assert result.status == "failed" + assert "expired" in (result.error or "").lower() diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 514ae673d06..f5448998a92 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -72,7 +72,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const queryClient = useQueryClient(); const { data: modelDataResponse, isLoading: isLoadingModels, refetch: refetchModels } = useModelsInfo(); const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap(); - const { data: credentialsResponse, isLoading: isLoadingCredentials } = useCredentials(); + const { data: credentialsResponse, isLoading: isLoadingCredentials, refetch: refetchCredentials } = useCredentials(); const credentialsList = credentialsResponse?.credentials || []; const { data: uiSettings, isLoading: isLoadingUISettings } = useUISettings(); @@ -421,6 +421,7 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te setShowAdvancedSettings={setShowAdvancedSettings} teams={teams} credentials={credentialsList} + refetchCredentials={refetchCredentials} accessToken={accessToken} userRole={userRole} /> diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index 2b3f23a35ae..c39136a5161 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -4,12 +4,21 @@ import { useTags } from "@/app/(dashboard)/hooks/tags/useTags"; import { all_admin_roles, isUserTeamAdminForAnyTeam } from "@/utils/roles"; import { Switch, Text } from "@tremor/react"; import type { FormInstance } from "antd"; -import { Select as AntdSelect, Button, Card, Col, Form, Modal, Row, Tooltip, Typography, Alert } from "antd"; +import { Select as AntdSelect, Button, Card, Col, Form, Input, Modal, Row, Spin, Tooltip, Typography, Alert } from "antd"; import type { UploadProps } from "antd/es/upload"; -import React, { useEffect, useMemo, useState } from "react"; +import React, { useCallback, useEffect, useMemo, useRef, useState } from "react"; import TeamDropdown from "../common_components/team_dropdown"; +import NotificationsManager from "../molecules/notifications_manager"; import type { Team } from "../key_team_helpers/key_list"; -import { type CredentialItem, type ProviderCreateInfo, modelAvailableCall } from "../networking"; +import { + type CredentialItem, + type ProviderCreateInfo, + githubCopilotInitiateAuth, + githubCopilotCheckStatus, + chatgptInitiateAuth, + chatgptCheckStatus, + modelAvailableCall, +} from "../networking"; import { Providers, providerLogoMap } from "../provider_info_helpers"; import { ProviderLogo } from "../molecules/models/ProviderLogo"; import AdvancedSettings from "./advanced_settings"; @@ -33,6 +42,7 @@ interface AddModelFormProps { setShowAdvancedSettings: (show: boolean) => void; teams: Team[] | null; credentials: CredentialItem[]; + refetchCredentials?: () => void; } const { Title, Link } = Typography; @@ -50,6 +60,7 @@ const AddModelForm: React.FC = ({ setShowAdvancedSettings, teams, credentials, + refetchCredentials, }) => { const [testMode, setTestMode] = useState("chat"); const [isResultModalVisible, setIsResultModalVisible] = useState(false); @@ -63,6 +74,214 @@ const AddModelForm: React.FC = ({ isLoading: isProviderMetadataLoading, error: providerMetadataError, } = useProviderFields(); + + // Inline device code flow (GitHub Copilot, ChatGPT, etc.) + const [ghDeviceCodeState, setGhDeviceCodeState] = useState< + | { phase: "idle" } + | { phase: "polling"; deviceCode: string; userCode: string; verificationUri: string } + | { phase: "success" } + | { phase: "error"; message: string } + >({ phase: "idle" }); + // Hold the access_token in a ref — never rendered, injected at submit time + const ghAccessTokenRef = useRef(null); + const ghPollingRef = useRef | null>(null); + + const stopGhPolling = useCallback(() => { + if (ghPollingRef.current) { + clearTimeout(ghPollingRef.current); + ghPollingRef.current = null; + } + }, []); + + useEffect(() => () => stopGhPolling(), [stopGhPolling]); + + // Determine if the selected provider uses device_code auth flow and get its metadata + const deviceCodeProviderInfo = React.useMemo(() => { + if (!providerMetadata) return null; + const info = providerMetadata.find( + (p) => + p.auth_flow === "device_code" && + (p.provider === selectedProvider || + p.provider_display_name === Providers[selectedProvider as keyof typeof Providers]), + ); + return info || null; + }, [selectedProvider, providerMetadata]); + const isDeviceCodeProvider = deviceCodeProviderInfo != null; + + // Reset device code state when provider changes away from a device_code provider + useEffect(() => { + if (!isDeviceCodeProvider) { + stopGhPolling(); + setGhDeviceCodeState({ phase: "idle" }); + ghAccessTokenRef.current = null; + } + }, [isDeviceCodeProvider, stopGhPolling]); + + const handleGhStartDeviceCode = async () => { + if (!accessToken || !deviceCodeProviderInfo) return; + const litellmProvider = deviceCodeProviderInfo.litellm_provider; + const providerLabel = deviceCodeProviderInfo.provider_display_name || litellmProvider; + + try { + let deviceId: string; + let userCode: string; + let verificationUri: string; + let pollIntervalMs: number; + + if (litellmProvider === "chatgpt") { + const result = await chatgptInitiateAuth(accessToken); + deviceId = result.device_auth_id; + userCode = result.user_code; + verificationUri = result.verification_uri; + pollIntervalMs = result.poll_interval_ms; + } else { + const result = await githubCopilotInitiateAuth(accessToken); + deviceId = result.device_code; + userCode = result.user_code; + verificationUri = result.verification_uri; + pollIntervalMs = result.poll_interval_ms; + } + + setGhDeviceCodeState({ + phase: "polling", + deviceCode: deviceId, + userCode, + verificationUri, + }); + let currentPollInterval = pollIntervalMs || 5000; + + const schedulePoll = (delayMs: number) => { + ghPollingRef.current = setTimeout(async () => { + try { + let apiKey: string | undefined; + let failed = false; + let errorMsg: string | undefined; + let retryAfterMs: number | undefined; + + if (litellmProvider === "chatgpt") { + const status = await chatgptCheckStatus(accessToken, deviceId, userCode); + console.log(`[${providerLabel} AddModel] poll response:`, status); + if (status.status === "complete" && status.refresh_token) { + apiKey = status.refresh_token; + } else if (status.status === "failed") { + failed = true; + errorMsg = status.error; + } + } else { + const status = await githubCopilotCheckStatus(accessToken, deviceId); + console.log(`[${providerLabel} AddModel] poll response:`, status); + if (status.status === "complete" && status.access_token) { + apiKey = status.access_token; + } else if (status.status === "failed") { + failed = true; + errorMsg = status.error; + } else { + retryAfterMs = status.retry_after_ms ?? undefined; + } + } + + if (apiKey) { + stopGhPolling(); + ghAccessTokenRef.current = apiKey; + form.setFieldValue("api_key", apiKey); + setGhDeviceCodeState({ phase: "success" }); + } else if (failed) { + stopGhPolling(); + setGhDeviceCodeState({ phase: "error", message: errorMsg || "Authorization failed" }); + } else { + if (retryAfterMs != null) { + currentPollInterval = retryAfterMs; + } + schedulePoll(currentPollInterval); + } + } catch (e) { + console.error(`[${providerLabel} AddModel] poll error:`, e); + stopGhPolling(); + setGhDeviceCodeState({ phase: "error", message: "Failed to check authorization status" }); + } + }, delayMs); + }; + schedulePoll(currentPollInterval); + } catch { + NotificationsManager.error(`Failed to start ${providerLabel} authorization`); + setGhDeviceCodeState({ phase: "error", message: `Failed to start ${providerLabel} authorization` }); + } + }; + + const handleGhCancel = () => { + stopGhPolling(); + setGhDeviceCodeState({ phase: "idle" }); + ghAccessTokenRef.current = null; + }; + + const dcProviderLabel = deviceCodeProviderInfo?.provider_display_name || "Provider"; + + const renderGhDeviceCodeFlow = () => { + switch (ghDeviceCodeState.phase) { + case "idle": + return ( +
+ +
+ ); + case "polling": + return ( +
+ Enter this code to authorize: +
+ {ghDeviceCodeState.userCode} +
+
+ +
+ + Waiting for {dcProviderLabel} authorization... + +
+ ); + case "success": + return ( +
+ + ✓ Authorization complete! Submit the form to add the model. + +
+ ); + case "error": + return ( +
+ {ghDeviceCodeState.message} + + +
+ ); + } + }; const { data: guardrailsList, isLoading: isGuardrailsLoading, error: guardrailsError } = useGuardrails(); const { data: tagsList, isLoading: isTagsLoading, error: tagsError } = useTags(); @@ -270,7 +489,11 @@ const AddModelForm: React.FC = ({ OR
- + {isDeviceCodeProvider ? ( + renderGhDeviceCodeFlow() + ) : ( + + )} ); } diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx index f9b6533ac60..75d02f13fde 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx @@ -23,6 +23,7 @@ interface AddModelTabProps { setShowAdvancedSettings: (show: boolean) => void; teams: Team[] | null; credentials: CredentialItem[]; + refetchCredentials?: () => void; accessToken: string; userRole: string; } @@ -40,6 +41,7 @@ const AddModelTab: React.FC = ({ setShowAdvancedSettings, teams, credentials, + refetchCredentials, accessToken, userRole, }) => { @@ -79,6 +81,7 @@ const AddModelTab: React.FC = ({ setShowAdvancedSettings={setShowAdvancedSettings} teams={teams} credentials={credentials} + refetchCredentials={refetchCredentials} /> diff --git a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx index 7270ef0577b..8ad4831d336 100644 --- a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx @@ -6,6 +6,8 @@ import { credentialCreateCall, githubCopilotInitiateAuth, githubCopilotCheckStatus, + chatgptInitiateAuth, + chatgptCheckStatus, } from "@/components/networking"; import NotificationsManager from "../molecules/notifications_manager"; import ProviderSpecificFields from "../add_model/provider_specific_fields"; @@ -41,16 +43,18 @@ const AddCredentialsModal: React.FC = ({ open, onCance 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; + // Determine if the selected provider uses device_code auth flow and get its litellm_provider + const deviceCodeProviderInfo = React.useMemo(() => { + if (!providerMetadata) return null; const info = providerMetadata.find( (p) => p.provider === selectedProvider || p.provider_display_name === Providers[selectedProvider as keyof typeof Providers], ); - return info?.auth_flow === "device_code"; + if (info?.auth_flow === "device_code") return info; + return null; }, [selectedProvider, providerMetadata]); + const isDeviceCodeProvider = deviceCodeProviderInfo != null; // Cleanup polling on unmount or modal close const stopPolling = useCallback(() => { @@ -89,59 +93,106 @@ const AddCredentialsModal: React.FC = ({ open, onCance form.validateFields(["credential_name"]); return; } - if (!accessToken) return; + if (!accessToken || !deviceCodeProviderInfo) return; + + const litellmProvider = deviceCodeProviderInfo.litellm_provider; + const providerLabel = deviceCodeProviderInfo.provider_display_name || litellmProvider; try { - const result = await githubCopilotInitiateAuth(accessToken); + // Initiate — dispatch to the right provider + let deviceId: string; + let userCode: string; + let verificationUri: string; + let pollIntervalMs: number; + + if (litellmProvider === "chatgpt") { + const result = await chatgptInitiateAuth(accessToken); + deviceId = result.device_auth_id; + userCode = result.user_code; + verificationUri = result.verification_uri; + pollIntervalMs = result.poll_interval_ms; + } else { + // Default: GitHub Copilot + const result = await githubCopilotInitiateAuth(accessToken); + deviceId = result.device_code; + userCode = result.user_code; + verificationUri = result.verification_uri; + pollIntervalMs = result.poll_interval_ms; + } + setDeviceCodeState({ phase: "polling", - deviceCode: result.device_code, - userCode: result.user_code, - verificationUri: result.verification_uri, + deviceCode: deviceId, + userCode, + verificationUri, }); - 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; + if (!pollIntervalMs) throw new Error(`${providerLabel} initiate response missing poll_interval_ms`); + // Mutable baseline — ratchets up on slow_down so that subsequent + // normal "pending" responses keep using the increased interval. + let currentPollInterval = pollIntervalMs; // setTimeout-based loop so each poll fires only after the previous one - // completes, and slow_down's retry_after_ms is respected exactly. + // completes, and slow_down / rate-limit intervals are 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) { + let apiKey: string | undefined; + let failed = false; + let errorMsg: string | undefined; + let retryAfterMs: number | undefined; + + if (litellmProvider === "chatgpt") { + const status = await chatgptCheckStatus(accessToken, deviceId, userCode); + console.log(`[${providerLabel} AddCredential] poll response:`, status); + if (status.status === "complete" && status.refresh_token) { + apiKey = status.refresh_token; + } else if (status.status === "failed") { + failed = true; + errorMsg = status.error; + } + } else { + const status = await githubCopilotCheckStatus(accessToken, deviceId); + console.log(`[${providerLabel} AddCredential] poll response:`, status); + if (status.status === "complete" && status.access_token) { + apiKey = status.access_token; + } else if (status.status === "failed") { + failed = true; + errorMsg = status.error; + } else { + retryAfterMs = status.retry_after_ms ?? undefined; + } + } + + if (apiKey) { stopPolling(); - accessTokenRef.current = status.access_token; - // Store as named credential + accessTokenRef.current = apiKey; try { await credentialCreateCall(accessToken, { credential_name: credentialName, - credential_values: { api_key: status.access_token }, - credential_info: { custom_llm_provider: "github_copilot" }, + credential_values: { api_key: apiKey }, + credential_info: { custom_llm_provider: litellmProvider }, }); setDeviceCodeState({ phase: "success", credentialName }); } catch (e) { - console.error("[GH Copilot AddCredential] credentialCreateCall failed:", e); + console.error(`[${providerLabel} 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") { + } else if (failed) { stopPolling(); - setDeviceCodeState({ phase: "error", message: status.error || "Authorization failed" }); + setDeviceCodeState({ phase: "error", message: errorMsg || "Authorization failed" }); } else { - // pending — ratchet up the baseline if GitHub requested slower - if (status.retry_after_ms != null) { - currentPollInterval = status.retry_after_ms; + // pending — ratchet up the baseline if provider requested slower + if (retryAfterMs != null) { + currentPollInterval = retryAfterMs; } schedulePoll(currentPollInterval); } } catch (e) { - console.error("[GH Copilot AddCredential] poll error:", e); + console.error(`[${providerLabel} AddCredential] poll error:`, e); stopPolling(); setDeviceCodeState({ phase: "error", message: "Failed to check authorization status" }); } @@ -149,7 +200,7 @@ const AddCredentialsModal: React.FC = ({ open, onCance }; schedulePoll(currentPollInterval); } catch { - setDeviceCodeState({ phase: "error", message: "Failed to start GitHub authorization" }); + setDeviceCodeState({ phase: "error", message: `Failed to start ${providerLabel} authorization` }); } }; @@ -161,16 +212,18 @@ const AddCredentialsModal: React.FC = ({ open, onCance form.resetFields(); }; + const providerDisplayName = deviceCodeProviderInfo?.provider_display_name || "Provider"; + const renderDeviceCodeFlow = () => { switch (deviceCodeState.phase) { case "idle": return (
- GitHub Copilot uses OAuth Device Code authorization. Click below to start. + {providerDisplayName} uses OAuth Device Code authorization. Click below to start.
); @@ -178,7 +231,7 @@ const AddCredentialsModal: React.FC = ({ open, onCance return (
- Enter this code on GitHub: + Enter this code to authorize:
= ({ open, onCance
- Waiting for GitHub authorization... + Waiting for {providerDisplayName} authorization...
@@ -215,7 +268,7 @@ const AddCredentialsModal: React.FC = ({ open, onCance return (
- GitHub Copilot credential "{deviceCodeState.credentialName}" created successfully! + {providerDisplayName} credential "{deviceCodeState.credentialName}" created successfully!