fix: clean up OAuth device code flow for reviewer clarity

Remove double Depends(user_api_key_auth) that ran auth twice per SSO
request, raise AuthenticationError instead of silently swallowing ChatGPT
token exchange failures, parallelize N+1 GitHub API calls in GET
/credentials with asyncio.gather, replace importlib dynamic dispatch with
direct imports, eliminate redundant static header generation, move inline
imports to module level, and remove unused parameters/exports.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Hunter Wittenborn 2026-03-30 05:31:16 -05:00
parent de31fc4471
commit fc0bf4bd30
8 changed files with 62 additions and 79 deletions

View file

@ -56,24 +56,27 @@ class ChatGPTConfig(OpenAIConfig):
headers, model, messages, optional_params, litellm_params, api_key, api_base
)
# 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:
# api_key is the refresh token. Exchange for access token (JWT)
# via OpenAI OAuth. Cached at module level by the authenticator.
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)
# get_chatgpt_default_headers already includes static headers.
default_headers = get_chatgpt_default_headers(
access_token, account_id, session_id
)
validated_headers = {**default_headers, **validated_headers}
except (GetAccessTokenError, RefreshAccessTokenError):
pass
except (GetAccessTokenError, RefreshAccessTokenError) as e:
raise AuthenticationError(
message=f"ChatGPT token exchange failed: {e}",
llm_provider="chatgpt",
model=model,
)
else:
validated_headers = {**get_chatgpt_static_headers(), **validated_headers}
return validated_headers

View file

@ -90,16 +90,12 @@ class GithubCopilotConfig(OpenAIConfig):
headers, model, messages, optional_params, litellm_params, api_key, api_base
)
# Always add static Copilot headers (editor-version, user-agent, etc.)
# These are required by the GitHub Copilot API on every request.
validated_headers = {**get_copilot_static_headers(), **validated_headers}
# 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:
# get_copilot_default_headers already includes static headers.
copilot_headers = get_copilot_default_headers(api_key)
validated_headers = {**copilot_headers, **validated_headers}
else:
validated_headers = {**get_copilot_static_headers(), **validated_headers}
# Add X-Initiator header based on message roles
initiator = self._determine_initiator(messages)

View file

@ -93,6 +93,12 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_content_from_model_response,
)
from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig
from litellm.llms.chatgpt.common_utils import (
get_chatgpt_static_headers,
)
from litellm.llms.github_copilot.common_utils import (
get_copilot_static_headers,
)
from litellm.llms.base_llm.base_model_iterator import (
convert_model_response_to_streaming,
)
@ -2605,25 +2611,12 @@ def completion( # type: ignore # noqa: PLR0915
headers = headers or litellm.headers
# 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.
_static_header_getters = {
"github_copilot": "litellm.llms.github_copilot.common_utils.get_copilot_static_headers",
"chatgpt": "litellm.llms.chatgpt.common_utils.get_chatgpt_static_headers",
}
if custom_llm_provider in _static_header_getters:
import importlib
_mod_path, _func_name = _static_header_getters[
custom_llm_provider
].rsplit(".", 1)
_get_static = getattr(importlib.import_module(_mod_path), _func_name)
provider_headers = _get_static()
if extra_headers:
provider_headers.update(extra_headers)
extra_headers = provider_headers
if custom_llm_provider == "github_copilot":
provider_headers = get_copilot_static_headers()
extra_headers = {**provider_headers, **(extra_headers or {})}
elif custom_llm_provider == "chatgpt":
provider_headers = get_chatgpt_static_headers()
extra_headers = {**provider_headers, **(extra_headers or {})}
if extra_headers is not None:
optional_params["extra_headers"] = extra_headers

View file

@ -26,7 +26,7 @@ from typing import Literal, Optional
from urllib.parse import quote
import httpx as _httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
@ -78,13 +78,10 @@ class StatusResponse(BaseModel):
@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),
):
"""
@ -129,13 +126,10 @@ async def chatgpt_initiate(
@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),
):
@ -156,7 +150,6 @@ async def chatgpt_status(
"""
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.
@ -196,12 +189,9 @@ async def chatgpt_status(
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 = (
@ -211,7 +201,7 @@ async def chatgpt_status(
f"&client_id={quote(CHATGPT_CLIENT_ID, safe='')}"
f"&code_verifier={quote(code_verifier, safe='')}"
)
async with _raw_httpx.AsyncClient() as raw_client:
async with _httpx.AsyncClient() as raw_client:
token_resp = await raw_client.post(
CHATGPT_OAUTH_TOKEN_URL,
headers={"Content-Type": "application/x-www-form-urlencoded"},
@ -228,9 +218,9 @@ async def chatgpt_status(
return StatusResponse(status="failed", error=f"Token exchange failed: {e}")
refresh_token = token_data.get("refresh_token")
verbose_proxy_logger.info(
verbose_proxy_logger.debug(
f"ChatGPT token exchange result: keys={list(token_data.keys())}, "
f"has_refresh={bool(refresh_token)}, refresh_len={len(refresh_token or '')}, "
f"has_refresh={bool(refresh_token)}, "
f"has_access={bool(token_data.get('access_token'))}"
)
if not refresh_token:

View file

@ -2,12 +2,15 @@
CRUD endpoints for storing reusable credentials.
"""
import asyncio
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Request, Response, Path
import litellm
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.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.litellm_logging import _get_masked_values
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
@ -118,9 +121,6 @@ async def _fetch_github_login(api_key: str) -> Optional[str]:
Call GET https://api.github.com/user with the given GitHub access token
and return the login name, or None if the call fails.
"""
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
try:
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.SSO_HANDLER
@ -153,21 +153,28 @@ async def get_credentials(
[BETA] endpoint. This might change unexpectedly.
"""
try:
masked_credentials = []
for credential in litellm.credential_list:
credential_info = dict(credential.credential_info or {})
# For GitHub Copilot credentials, inject runtime github_login from the API.
# The login is NOT stored in the DB — it's fetched live and added to the
# response only so the UI can display it.
if credential_info.get("custom_llm_provider") == "github_copilot":
# Batch-fetch GitHub logins in parallel to avoid N+1 sequential HTTP calls.
copilot_keys: list[tuple[int, str]] = []
for i, credential in enumerate(litellm.credential_list):
info = credential.credential_info or {}
if info.get("custom_llm_provider") == "github_copilot":
api_key = (credential.credential_values or {}).get("api_key")
if api_key:
github_login = await _fetch_github_login(api_key)
if github_login:
credential_info = {
**credential_info,
"github_login": github_login,
}
copilot_keys.append((i, api_key))
logins: dict[int, Optional[str]] = {}
if copilot_keys:
results = await asyncio.gather(
*(_fetch_github_login(key) for _, key in copilot_keys)
)
logins = {idx: login for (idx, _), login in zip(copilot_keys, results)}
masked_credentials = []
for i, credential in enumerate(litellm.credential_list):
credential_info = dict(credential.credential_info or {})
github_login = logins.get(i)
if github_login:
credential_info = {**credential_info, "github_login": github_login}
masked_credentials.append(
{
"credential_name": credential.credential_name,

View file

@ -22,7 +22,7 @@ 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 fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
@ -71,13 +71,10 @@ class StatusResponse(BaseModel):
@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),
):
"""
@ -127,13 +124,10 @@ async def github_copilot_initiate(
@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),
):

View file

@ -44,7 +44,7 @@ class TestGithubCopilotInitiate:
github_copilot_initiate,
)
result = await github_copilot_initiate(MagicMock(), MagicMock(), MagicMock())
result = await github_copilot_initiate(MagicMock())
assert result.device_code == "test-device-code-123"
assert result.user_code == "ABCD-1234"
assert result.verification_uri == "https://github.com/login/device"
@ -65,7 +65,7 @@ class TestGithubCopilotInitiate:
)
with pytest.raises(HTTPException) as exc_info:
await github_copilot_initiate(MagicMock(), MagicMock(), MagicMock())
await github_copilot_initiate(MagicMock())
assert exc_info.value.status_code == 502
@ -87,7 +87,7 @@ class TestGithubCopilotStatus:
)
result = await github_copilot_status(
MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock()
StatusRequest(device_code="test-dc"), MagicMock()
)
assert result.status == "pending"
assert result.access_token is None
@ -105,7 +105,7 @@ class TestGithubCopilotStatus:
)
result = await github_copilot_status(
MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock()
StatusRequest(device_code="test-dc"), MagicMock()
)
assert result.status == "pending"
assert result.retry_after_ms == 10_000
@ -122,7 +122,7 @@ class TestGithubCopilotStatus:
)
result = await github_copilot_status(
MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock()
StatusRequest(device_code="test-dc"), MagicMock()
)
assert result.status == "complete"
assert result.access_token == "ghu_abc123"
@ -140,7 +140,7 @@ class TestGithubCopilotStatus:
)
result = await github_copilot_status(
MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock()
StatusRequest(device_code="test-dc"), MagicMock()
)
assert result.status == "failed"
assert result.error is not None
@ -160,7 +160,7 @@ class TestGithubCopilotStatus:
)
result = await github_copilot_status(
MagicMock(), MagicMock(), StatusRequest(device_code="test-dc"), MagicMock()
StatusRequest(device_code="test-dc"), MagicMock()
)
assert result.status == "failed"
assert "expired" in (result.error or "").lower()

View file

@ -194,5 +194,5 @@ export function useDeviceCodeFlow({ accessToken, providerInfo, onSuccess }: UseD
}
}, [state, start, reset, providerLabel]);
return { state, start, reset, tokenRef, renderUI };
return { state, start, reset, renderUI };
}