mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
de31fc4471
commit
fc0bf4bd30
8 changed files with 62 additions and 79 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 };
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue