mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat: add ChatGPT device code OAuth + fix token exchange across all endpoints
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) <noreply@anthropic.com>
This commit is contained in:
parent
3cd86f3982
commit
b7432a4c44
24 changed files with 1372 additions and 205 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
252
litellm/proxy/credential_endpoints/chatgpt_sso.py
Normal file
252
litellm/proxy/credential_endpoints/chatgpt_sso.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
202
litellm/proxy/credential_endpoints/github_copilot_sso.py
Normal file
202
litellm/proxy/credential_endpoints/github_copilot_sso.py
Normal file
|
|
@ -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": <GitHub-reported interval in 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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -72,7 +72,7 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({ 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<ModelDashboardProps> = ({ premiumUser, te
|
|||
setShowAdvancedSettings={setShowAdvancedSettings}
|
||||
teams={teams}
|
||||
credentials={credentialsList}
|
||||
refetchCredentials={refetchCredentials}
|
||||
accessToken={accessToken}
|
||||
userRole={userRole}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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<AddModelFormProps> = ({
|
|||
setShowAdvancedSettings,
|
||||
teams,
|
||||
credentials,
|
||||
refetchCredentials,
|
||||
}) => {
|
||||
const [testMode, setTestMode] = useState<string>("chat");
|
||||
const [isResultModalVisible, setIsResultModalVisible] = useState<boolean>(false);
|
||||
|
|
@ -63,6 +74,214 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
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<string | null>(null);
|
||||
const ghPollingRef = useRef<ReturnType<typeof setTimeout> | 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 (
|
||||
<div className="mt-2 text-center">
|
||||
<Button
|
||||
type="primary"
|
||||
onClick={handleGhStartDeviceCode}
|
||||
>
|
||||
Authorize with {dcProviderLabel}
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
case "polling":
|
||||
return (
|
||||
<div className="text-center py-2">
|
||||
<Typography.Text className="block mb-2">Enter this code to authorize:</Typography.Text>
|
||||
<div
|
||||
style={{
|
||||
fontSize: "1.8rem",
|
||||
fontWeight: "bold",
|
||||
fontFamily: "monospace",
|
||||
letterSpacing: "0.3em",
|
||||
margin: "12px 0",
|
||||
padding: "10px 20px",
|
||||
background: "#f5f5f5",
|
||||
borderRadius: 8,
|
||||
display: "inline-block",
|
||||
userSelect: "all",
|
||||
}}
|
||||
>
|
||||
{ghDeviceCodeState.userCode}
|
||||
</div>
|
||||
<div className="mb-3">
|
||||
<Button type="link" onClick={() => window.open(ghDeviceCodeState.verificationUri, "_blank")}>
|
||||
Open {ghDeviceCodeState.verificationUri}
|
||||
</Button>
|
||||
</div>
|
||||
<Spin />
|
||||
<Typography.Text className="block mt-2 mb-3" type="secondary">Waiting for {dcProviderLabel} authorization...</Typography.Text>
|
||||
<Button onClick={handleGhCancel}>Cancel</Button>
|
||||
</div>
|
||||
);
|
||||
case "success":
|
||||
return (
|
||||
<div className="text-center py-2">
|
||||
<Typography.Text type="success" className="block mb-2">
|
||||
✓ Authorization complete! Submit the form to add the model.
|
||||
</Typography.Text>
|
||||
</div>
|
||||
);
|
||||
case "error":
|
||||
return (
|
||||
<div className="text-center py-2">
|
||||
<Typography.Text type="danger" className="block mb-3">{ghDeviceCodeState.message}</Typography.Text>
|
||||
<Button
|
||||
style={{ marginRight: 8 }}
|
||||
onClick={() => setGhDeviceCodeState({ phase: "idle" })}
|
||||
>
|
||||
Retry
|
||||
</Button>
|
||||
<Button onClick={handleGhCancel}>Cancel</Button>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
};
|
||||
const { data: guardrailsList, isLoading: isGuardrailsLoading, error: guardrailsError } = useGuardrails();
|
||||
const { data: tagsList, isLoading: isTagsLoading, error: tagsError } = useTags();
|
||||
|
||||
|
|
@ -270,7 +489,11 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
<span className="px-4 text-gray-500 text-sm">OR</span>
|
||||
<div className="flex-grow border-t border-gray-200"></div>
|
||||
</div>
|
||||
<ProviderSpecificFields selectedProvider={selectedProvider} uploadProps={uploadProps} />
|
||||
{isDeviceCodeProvider ? (
|
||||
renderGhDeviceCodeFlow()
|
||||
) : (
|
||||
<ProviderSpecificFields selectedProvider={selectedProvider} uploadProps={uploadProps} />
|
||||
)}
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<AddModelTabProps> = ({
|
|||
setShowAdvancedSettings,
|
||||
teams,
|
||||
credentials,
|
||||
refetchCredentials,
|
||||
accessToken,
|
||||
userRole,
|
||||
}) => {
|
||||
|
|
@ -79,6 +81,7 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
|
|||
setShowAdvancedSettings={setShowAdvancedSettings}
|
||||
teams={teams}
|
||||
credentials={credentials}
|
||||
refetchCredentials={refetchCredentials}
|
||||
/>
|
||||
</TabPanel>
|
||||
<TabPanel>
|
||||
|
|
|
|||
|
|
@ -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<AddCredentialsModalProps> = ({ open, onCance
|
|||
const accessTokenRef = useRef<string | null>(null);
|
||||
const pollingRef = useRef<ReturnType<typeof setInterval> | 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<AddCredentialsModalProps> = ({ 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<AddCredentialsModalProps> = ({ 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<AddCredentialsModalProps> = ({ open, onCance
|
|||
form.resetFields();
|
||||
};
|
||||
|
||||
const providerDisplayName = deviceCodeProviderInfo?.provider_display_name || "Provider";
|
||||
|
||||
const renderDeviceCodeFlow = () => {
|
||||
switch (deviceCodeState.phase) {
|
||||
case "idle":
|
||||
return (
|
||||
<div className="text-center py-4">
|
||||
<Text className="block mb-4">
|
||||
GitHub Copilot uses OAuth Device Code authorization. Click below to start.
|
||||
{providerDisplayName} uses OAuth Device Code authorization. Click below to start.
|
||||
</Text>
|
||||
<Button type="primary" onClick={handleStartDeviceCode}>
|
||||
Start GitHub Authorization
|
||||
Start {providerDisplayName} Authorization
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
|
|
@ -178,7 +231,7 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
|
|||
return (
|
||||
<div className="text-center py-4">
|
||||
<Text className="block mb-2">
|
||||
Enter this code on GitHub:
|
||||
Enter this code to authorize:
|
||||
</Text>
|
||||
<div
|
||||
style={{
|
||||
|
|
@ -206,7 +259,7 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
|
|||
</div>
|
||||
<Spin />
|
||||
<Text className="block mt-2 mb-4" type="secondary">
|
||||
Waiting for GitHub authorization...
|
||||
Waiting for {providerDisplayName} authorization...
|
||||
</Text>
|
||||
<Button onClick={handleCancel}>Cancel</Button>
|
||||
</div>
|
||||
|
|
@ -215,7 +268,7 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
|
|||
return (
|
||||
<div className="text-center py-4">
|
||||
<Text className="block mb-4" type="success" style={{ fontSize: "1.1rem" }}>
|
||||
GitHub Copilot credential "{deviceCodeState.credentialName}" created successfully!
|
||||
{providerDisplayName} credential "{deviceCodeState.credentialName}" created successfully!
|
||||
</Text>
|
||||
<Button type="primary" onClick={handleSuccessClose}>
|
||||
Done
|
||||
|
|
|
|||
|
|
@ -154,37 +154,63 @@ const CredentialsPanel: React.FC<CredentialsPanelProps> = ({ uploadProps }) => {
|
|||
<TableBody>
|
||||
{!credentialList || credentialList.length === 0 ? (
|
||||
<TableRow>
|
||||
<TableCell colSpan={4} className="text-center py-4 text-gray-500">
|
||||
<TableCell colSpan={3} className="text-center py-4 text-gray-500">
|
||||
No credentials configured
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
) : (
|
||||
credentialList.map((credential: CredentialItem, index: number) => (
|
||||
<TableRow key={index}>
|
||||
<TableCell>{credential.credential_name}</TableCell>
|
||||
<TableCell>
|
||||
{renderProviderBadge((credential.credential_info?.custom_llm_provider as string) || "-")}
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
<Button
|
||||
icon={PencilAltIcon}
|
||||
variant="light"
|
||||
size="sm"
|
||||
onClick={() => {
|
||||
setSelectedCredential(credential);
|
||||
setIsUpdateModalOpen(true);
|
||||
}}
|
||||
/>
|
||||
<Button
|
||||
icon={TrashIcon}
|
||||
variant="light"
|
||||
size="sm"
|
||||
onClick={() => openDeleteModal(credential)}
|
||||
className="ml-2"
|
||||
/>
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
))
|
||||
credentialList.map((credential: CredentialItem, index: number) => {
|
||||
const githubLogin = credential.credential_info?.github_login;
|
||||
return (
|
||||
<TableRow key={index}>
|
||||
<TableCell>
|
||||
{credential.credential_name}
|
||||
{githubLogin && (
|
||||
<span className="ml-2 text-gray-500 text-sm inline-flex items-center gap-1">
|
||||
(
|
||||
<svg
|
||||
viewBox="0 0 16 16"
|
||||
className="w-4 h-4 inline-block"
|
||||
aria-hidden="true"
|
||||
fill="currentColor"
|
||||
>
|
||||
<path d="M8 0C3.58 0 0 3.58 0 8c0 3.54 2.29 6.53 5.47 7.59.4.07.55-.17.55-.38 0-.19-.01-.82-.01-1.49-2.01.37-2.53-.49-2.69-.94-.09-.23-.48-.94-.82-1.13-.28-.15-.68-.52-.01-.53.63-.01 1.08.58 1.23.82.72 1.21 1.87.87 2.33.66.07-.52.28-.87.51-1.07-1.78-.2-3.64-.89-3.64-3.95 0-.87.31-1.59.82-2.15-.08-.2-.36-1.02.08-2.12 0 0 .67-.21 2.2.82.64-.18 1.32-.27 2-.27.68 0 1.36.09 2 .27 1.53-1.04 2.2-.82 2.2-.82.44 1.1.16 1.92.08 2.12.51.56.82 1.27.82 2.15 0 3.07-1.87 3.75-3.65 3.95.29.25.54.73.54 1.48 0 1.07-.01 1.93-.01 2.2 0 .21.15.46.55.38A8.013 8.013 0 0 0 16 8c0-4.42-3.58-8-8-8z" />
|
||||
</svg>
|
||||
<a
|
||||
href={`https://github.com/${githubLogin}`}
|
||||
target="_blank"
|
||||
rel="noopener noreferrer"
|
||||
>
|
||||
{githubLogin}
|
||||
</a>
|
||||
)
|
||||
</span>
|
||||
)}
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
{renderProviderBadge((credential.credential_info?.custom_llm_provider as string) || "-")}
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
<Button
|
||||
icon={PencilAltIcon}
|
||||
variant="light"
|
||||
size="sm"
|
||||
onClick={() => {
|
||||
setSelectedCredential(credential);
|
||||
setIsUpdateModalOpen(true);
|
||||
}}
|
||||
/>
|
||||
<Button
|
||||
icon={TrashIcon}
|
||||
variant="light"
|
||||
size="sm"
|
||||
onClick={() => openDeleteModal(credential)}
|
||||
className="ml-2"
|
||||
/>
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
);
|
||||
})
|
||||
)}
|
||||
</TableBody>
|
||||
</Table>
|
||||
|
|
@ -194,7 +220,10 @@ const CredentialsPanel: React.FC<CredentialsPanelProps> = ({ uploadProps }) => {
|
|||
<AddCredentialsTab
|
||||
onAddCredential={handleAddCredential}
|
||||
open={isAddModalOpen}
|
||||
onCancel={() => setIsAddModalOpen(false)}
|
||||
onCancel={() => {
|
||||
setIsAddModalOpen(false);
|
||||
refetchCredentials();
|
||||
}}
|
||||
uploadProps={uploadProps}
|
||||
/>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -259,6 +259,7 @@ export interface CredentialItem {
|
|||
custom_llm_provider?: string;
|
||||
description?: string;
|
||||
required?: boolean;
|
||||
github_login?: string;
|
||||
};
|
||||
}
|
||||
|
||||
|
|
@ -278,6 +279,7 @@ export interface ProviderCreateInfo {
|
|||
provider_display_name: string;
|
||||
litellm_provider: string;
|
||||
default_model_placeholder?: string | null;
|
||||
auth_flow?: string | null;
|
||||
credential_fields: ProviderCredentialFieldMetadata[];
|
||||
}
|
||||
|
||||
|
|
@ -3567,6 +3569,99 @@ export const credentialDeleteCall = async (accessToken: string, credentialName:
|
|||
}
|
||||
};
|
||||
|
||||
export const githubCopilotInitiateAuth = async (
|
||||
accessToken: string,
|
||||
): Promise<{ device_code: string; user_code: string; verification_uri: string; poll_interval_ms: number; expires_in: number }> => {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/credentials/github_copilot/initiate`
|
||||
: `/credentials/github_copilot/initiate`;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
return response.json();
|
||||
};
|
||||
|
||||
export const githubCopilotCheckStatus = async (
|
||||
accessToken: string,
|
||||
deviceCode: string,
|
||||
): Promise<{ status: string; access_token?: string; retry_after_ms?: number; error?: string }> => {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/credentials/github_copilot/status`
|
||||
: `/credentials/github_copilot/status`;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify({ device_code: deviceCode }),
|
||||
});
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
return response.json();
|
||||
};
|
||||
|
||||
export const chatgptInitiateAuth = async (
|
||||
accessToken: string,
|
||||
): Promise<{ device_auth_id: string; user_code: string; verification_uri: string; poll_interval_ms: number }> => {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/credentials/chatgpt/initiate`
|
||||
: `/credentials/chatgpt/initiate`;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
return response.json();
|
||||
};
|
||||
|
||||
export const chatgptCheckStatus = async (
|
||||
accessToken: string,
|
||||
deviceAuthId: string,
|
||||
userCode: string,
|
||||
): Promise<{ status: string; refresh_token?: string; account_id?: string; error?: string }> => {
|
||||
const url = proxyBaseUrl
|
||||
? `${proxyBaseUrl}/credentials/chatgpt/status`
|
||||
: `/credentials/chatgpt/status`;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify({ device_auth_id: deviceAuthId, user_code: userCode }),
|
||||
});
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
return response.json();
|
||||
};
|
||||
|
||||
export const credentialUpdateCall = async (
|
||||
accessToken: string,
|
||||
credentialName: string,
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ export enum Providers {
|
|||
FireworksAI = "Fireworks AI",
|
||||
FRIENDLIAI = "Friendliai",
|
||||
GALADRIEL = "Galadriel",
|
||||
CHATGPT = "ChatGPT",
|
||||
GITHUB_COPILOT = "Github Copilot",
|
||||
Google_AI_Studio = "Google AI Studio",
|
||||
GradientAI = "GradientAI",
|
||||
|
|
@ -146,6 +147,7 @@ export const provider_map: Record<string, string> = {
|
|||
FireworksAI: "fireworks_ai",
|
||||
FRIENDLIAI: "friendliai",
|
||||
GALADRIEL: "galadriel",
|
||||
CHATGPT: "chatgpt",
|
||||
GITHUB_COPILOT: "github_copilot",
|
||||
Google_AI_Studio: "gemini",
|
||||
GradientAI: "gradient_ai",
|
||||
|
|
@ -247,6 +249,7 @@ export const providerLogoMap: Record<string, string> = {
|
|||
[Providers.FEATHERLESS_AI]: `${asset_logos_folder}featherless.svg`,
|
||||
[Providers.FireworksAI]: `${asset_logos_folder}fireworks.svg`,
|
||||
[Providers.FRIENDLIAI]: `${asset_logos_folder}friendli.svg`,
|
||||
[Providers.CHATGPT]: `${asset_logos_folder}openai_small.svg`,
|
||||
[Providers.GITHUB_COPILOT]: `${asset_logos_folder}github_copilot.svg`,
|
||||
[Providers.Google_AI_Studio]: `${asset_logos_folder}google.svg`,
|
||||
[Providers.GradientAI]: `${asset_logos_folder}gradientai.svg`,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue