refactor: DRY up device code OAuth flow, fix quality/efficiency issues

- Extract shared useDeviceCodeFlow hook (eliminates ~200 lines of duplicated
  polling logic between AddModelForm and AddCredentialModal)
- Extract authenticatedPost helper in networking.tsx (4 OAuth functions → 4 one-liners)
- Replace inline SVG with GithubOutlined icon (CLAUDE.md violation)
- Remove unused refetchCredentials prop from AddModelForm
- Fix sync HTTP blocking in async get_credentials endpoint (N+1 GitHub API calls)
- Use cached _get_httpx_client() in ChatGPT credential-mode refresh
- Remove _get_github_headers backward-compat alias (only used internally)
- Fix redundant except (ImportError, Exception) in proxy_cli.py
- Unify static-header blocks in main.py with dispatch table

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Hunter Wittenborn 2026-03-30 05:19:56 -05:00
parent b7432a4c44
commit de31fc4471
13 changed files with 350 additions and 577 deletions

View file

@ -158,24 +158,20 @@ class Authenticator:
)
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.
"""
"""Refresh tokens without writing to disk (credential mode)."""
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()
client = _get_httpx_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}",

View file

@ -207,7 +207,7 @@ class Authenticator:
RefreshAPIKeyError: If unable to refresh the API key.
"""
access_token = self.get_access_token()
headers = self._get_github_headers(access_token)
headers = Authenticator.get_github_headers(access_token)
max_retries = 3
for attempt in range(max_retries):
@ -271,10 +271,6 @@ class Authenticator:
return headers
# Backward-compatible instance alias
def _get_github_headers(self, access_token: Optional[str] = None) -> Dict[str, str]:
return Authenticator.get_github_headers(access_token)
def _get_device_code(self) -> Dict[str, str]:
"""
Get a device code for GitHub authentication.
@ -289,7 +285,7 @@ class Authenticator:
sync_client = _get_httpx_client()
resp = sync_client.post(
GITHUB_DEVICE_CODE_URL,
headers=self._get_github_headers(),
headers=Authenticator.get_github_headers(),
json={"client_id": GITHUB_CLIENT_ID, "scope": "read:user"},
)
resp.raise_for_status()
@ -343,7 +339,7 @@ class Authenticator:
try:
resp = sync_client.post(
GITHUB_ACCESS_TOKEN_URL,
headers=self._get_github_headers(),
headers=Authenticator.get_github_headers(),
json={
"client_id": GITHUB_CLIENT_ID,
"device_code": device_code,

View file

@ -2609,25 +2609,21 @@ def completion( # type: ignore # noqa: PLR0915
# 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.common_utils import (
get_copilot_static_headers,
)
_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
copilot_headers = get_copilot_static_headers()
_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:
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
provider_headers.update(extra_headers)
extra_headers = provider_headers
if extra_headers is not None:
optional_params["extra_headers"] = extra_headers

View file

@ -113,16 +113,19 @@ async def create_credential(
raise handle_exception_on_proxy(e)
def _fetch_github_login(api_key: str) -> Optional[str]:
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_httpx_client
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
try:
sync_client = _get_httpx_client()
resp = sync_client.get(
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.SSO_HANDLER
)
resp = await async_client.get(
"https://api.github.com/user",
headers={
"Authorization": f"token {api_key}",
@ -159,13 +162,18 @@ async def get_credentials(
if credential_info.get("custom_llm_provider") == "github_copilot":
api_key = (credential.credential_values or {}).get("api_key")
if api_key:
github_login = _fetch_github_login(api_key)
github_login = await _fetch_github_login(api_key)
if github_login:
credential_info = {**credential_info, "github_login": github_login}
credential_info = {
**credential_info,
"github_login": github_login,
}
masked_credentials.append(
{
"credential_name": credential.credential_name,
"credential_values": _get_masked_values(credential.credential_values),
"credential_values": _get_masked_values(
credential.credential_values
),
"credential_info": credential_info,
}
)

View file

@ -328,7 +328,7 @@ class ProxyInitializationHelpers:
import uvloop # noqa: F401
return "uvloop"
except (ImportError, Exception):
except ImportError:
return "asyncio"
@staticmethod

View file

@ -50,13 +50,13 @@ class TestGitHubCopilotAuthenticator:
def test_get_github_headers(self, authenticator):
"""Test that GitHub headers are correctly generated."""
headers = authenticator._get_github_headers()
headers = Authenticator.get_github_headers()
assert "accept" in headers
assert "editor-version" in headers
assert "user-agent" in headers
assert "content-type" in headers
headers_with_token = authenticator._get_github_headers("test-token")
headers_with_token = Authenticator.get_github_headers("test-token")
assert headers_with_token["authorization"] == "token test-token"
def test_get_access_token_from_file(self, authenticator):

View file

@ -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, refetch: refetchCredentials } = useCredentials();
const { data: credentialsResponse, isLoading: isLoadingCredentials } = useCredentials();
const credentialsList = credentialsResponse?.credentials || [];
const { data: uiSettings, isLoading: isLoadingUISettings } = useUISettings();
@ -421,7 +421,6 @@ const ModelsAndEndpointsView: React.FC<ModelDashboardProps> = ({ premiumUser, te
setShowAdvancedSettings={setShowAdvancedSettings}
teams={teams}
credentials={credentialsList}
refetchCredentials={refetchCredentials}
accessToken={accessToken}
userRole={userRole}
/>

View file

@ -4,23 +4,19 @@ 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, Input, Modal, Row, Spin, Tooltip, Typography, Alert } from "antd";
import { Select as AntdSelect, Button, Card, Col, Form, Input, Modal, Row, Tooltip, Typography, Alert } from "antd";
import type { UploadProps } from "antd/es/upload";
import React, { useCallback, useEffect, useMemo, useRef, useState } from "react";
import React, { useCallback, useEffect, useMemo, 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,
githubCopilotInitiateAuth,
githubCopilotCheckStatus,
chatgptInitiateAuth,
chatgptCheckStatus,
modelAvailableCall,
} from "../networking";
import { Providers, providerLogoMap } from "../provider_info_helpers";
import { ProviderLogo } from "../molecules/models/ProviderLogo";
import { useDeviceCodeFlow } from "@/hooks/useDeviceCodeFlow";
import AdvancedSettings from "./advanced_settings";
import ConditionalPublicModelName from "./conditional_public_model_name";
import LiteLLMModelNameField from "./litellm_model_name";
@ -42,7 +38,6 @@ interface AddModelFormProps {
setShowAdvancedSettings: (show: boolean) => void;
teams: Team[] | null;
credentials: CredentialItem[];
refetchCredentials?: () => void;
}
const { Title, Link } = Typography;
@ -60,7 +55,6 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
setShowAdvancedSettings,
teams,
credentials,
refetchCredentials,
}) => {
const [testMode, setTestMode] = useState<string>("chat");
const [isResultModalVisible, setIsResultModalVisible] = useState<boolean>(false);
@ -75,28 +69,8 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
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(() => {
// Determine if the selected provider uses device_code auth flow
const deviceCodeProviderInfo = useMemo(() => {
if (!providerMetadata) return null;
const info = providerMetadata.find(
(p) =>
@ -108,180 +82,22 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
}, [selectedProvider, providerMetadata]);
const isDeviceCodeProvider = deviceCodeProviderInfo != null;
const handleDeviceCodeSuccess = useCallback(
(apiKey: string) => { form.setFieldValue("api_key", apiKey); },
[form],
);
const { reset: resetDeviceCode, renderUI: renderDeviceCodeFlow } = useDeviceCodeFlow({
accessToken,
providerInfo: deviceCodeProviderInfo,
onSuccess: handleDeviceCodeSuccess,
});
// 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]);
if (!isDeviceCodeProvider) resetDeviceCode();
}, [isDeviceCodeProvider, resetDeviceCode]);
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();
@ -490,7 +306,7 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
<div className="flex-grow border-t border-gray-200"></div>
</div>
{isDeviceCodeProvider ? (
renderGhDeviceCodeFlow()
renderDeviceCodeFlow()
) : (
<ProviderSpecificFields selectedProvider={selectedProvider} uploadProps={uploadProps} />
)}

View file

@ -23,7 +23,6 @@ interface AddModelTabProps {
setShowAdvancedSettings: (show: boolean) => void;
teams: Team[] | null;
credentials: CredentialItem[];
refetchCredentials?: () => void;
accessToken: string;
userRole: string;
}
@ -41,7 +40,6 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
setShowAdvancedSettings,
teams,
credentials,
refetchCredentials,
accessToken,
userRole,
}) => {
@ -81,7 +79,6 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
setShowAdvancedSettings={setShowAdvancedSettings}
teams={teams}
credentials={credentials}
refetchCredentials={refetchCredentials}
/>
</TabPanel>
<TabPanel>

View file

@ -1,20 +1,15 @@
import { TextInput } from "@tremor/react";
import { Select as AntdSelect, Button, Form, Modal, Spin, Tooltip, Typography } from "antd";
import { Select as AntdSelect, Button, Form, Modal, Tooltip, Typography } from "antd";
import type { UploadProps } from "antd/es/upload";
import React, { useCallback, useEffect, useRef, useState } from "react";
import {
credentialCreateCall,
githubCopilotInitiateAuth,
githubCopilotCheckStatus,
chatgptInitiateAuth,
chatgptCheckStatus,
} from "@/components/networking";
import React, { useCallback, useMemo, useState } from "react";
import { credentialCreateCall } from "@/components/networking";
import NotificationsManager from "../molecules/notifications_manager";
import ProviderSpecificFields from "../add_model/provider_specific_fields";
import { Providers, providerLogoMap } from "../provider_info_helpers";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useProviderFields } from "@/app/(dashboard)/hooks/providers/useProviderFields";
const { Link, Text } = Typography;
import { useDeviceCodeFlow } from "@/hooks/useDeviceCodeFlow";
const { Link } = Typography;
interface AddCredentialsModalProps {
open: boolean;
@ -25,26 +20,13 @@ interface AddCredentialsModalProps {
initialProvider?: string;
}
const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCancel, onAddCredential, uploadProps, initialCredentialName, initialProvider }) => {
const [form] = Form.useForm();
const [selectedProvider, setSelectedProvider] = useState<Providers>(Providers.OpenAI);
const { accessToken } = useAuthorized();
const { data: providerMetadata } = useProviderFields();
// Device code flow state
const [deviceCodeState, setDeviceCodeState] = useState<
| { phase: "idle" }
| { phase: "polling"; deviceCode: string; userCode: string; verificationUri: string }
| { phase: "success"; credentialName: string }
| { phase: "error"; message: string }
>({ phase: "idle" });
// Hold access_token in a ref — never rendered, never put in form fields
const accessTokenRef = useRef<string | null>(null);
const pollingRef = useRef<ReturnType<typeof setInterval> | null>(null);
// Determine if the selected provider uses device_code auth flow and get its litellm_provider
const deviceCodeProviderInfo = React.useMemo(() => {
const deviceCodeProviderInfo = useMemo(() => {
if (!providerMetadata) return null;
const info = providerMetadata.find(
(p) =>
@ -56,22 +38,31 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
}, [selectedProvider, providerMetadata]);
const isDeviceCodeProvider = deviceCodeProviderInfo != null;
// Cleanup polling on unmount or modal close
const stopPolling = useCallback(() => {
if (pollingRef.current) {
clearInterval(pollingRef.current);
pollingRef.current = null;
}
}, []);
const handleDeviceCodeSuccess = useCallback(
async (apiKey: string, litellmProvider: string) => {
const credentialName = form.getFieldValue("credential_name");
if (!credentialName) {
form.validateFields(["credential_name"]);
throw new Error("Credential name required");
}
if (!accessToken) throw new Error("No access token");
await credentialCreateCall(accessToken, {
credential_name: credentialName,
credential_values: { api_key: apiKey },
credential_info: { custom_llm_provider: litellmProvider },
});
},
[form, accessToken],
);
useEffect(() => {
return stopPolling;
}, [stopPolling]);
const { state: deviceCodeState, start: startDeviceCode, reset: resetDeviceCode, renderUI: renderDeviceCodeFlow } = useDeviceCodeFlow({
accessToken,
providerInfo: deviceCodeProviderInfo,
onSuccess: handleDeviceCodeSuccess,
});
const handleCancel = () => {
stopPolling();
setDeviceCodeState({ phase: "idle" });
accessTokenRef.current = null;
resetDeviceCode();
onCancel();
form.resetFields();
};
@ -93,204 +84,41 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
form.validateFields(["credential_name"]);
return;
}
if (!accessToken || !deviceCodeProviderInfo) return;
const litellmProvider = deviceCodeProviderInfo.litellm_provider;
const providerLabel = deviceCodeProviderInfo.provider_display_name || litellmProvider;
try {
// 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: deviceId,
userCode,
verificationUri,
});
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 / rate-limit intervals are respected exactly.
const schedulePoll = (delayMs: number) => {
pollingRef.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} 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 = apiKey;
try {
await credentialCreateCall(accessToken, {
credential_name: credentialName,
credential_values: { api_key: apiKey },
credential_info: { custom_llm_provider: litellmProvider },
});
setDeviceCodeState({ phase: "success", credentialName });
} catch (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 (failed) {
stopPolling();
setDeviceCodeState({ phase: "error", message: errorMsg || "Authorization failed" });
} else {
// pending — ratchet up the baseline if provider requested slower
if (retryAfterMs != null) {
currentPollInterval = retryAfterMs;
}
schedulePoll(currentPollInterval);
}
} catch (e) {
console.error(`[${providerLabel} AddCredential] poll error:`, e);
stopPolling();
setDeviceCodeState({ phase: "error", message: "Failed to check authorization status" });
}
}, delayMs);
};
schedulePoll(currentPollInterval);
} catch {
setDeviceCodeState({ phase: "error", message: `Failed to start ${providerLabel} authorization` });
}
await startDeviceCode();
};
const handleSuccessClose = () => {
stopPolling();
setDeviceCodeState({ phase: "idle" });
accessTokenRef.current = null;
resetDeviceCode();
onCancel();
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">
{providerDisplayName} uses OAuth Device Code authorization. Click below to start.
</Text>
<Button type="primary" onClick={handleStartDeviceCode}>
Start {providerDisplayName} Authorization
</Button>
</div>
);
case "polling":
return (
<div className="text-center py-4">
<Text className="block mb-2">
Enter this code to authorize:
</Text>
<div
style={{
fontSize: "2rem",
fontWeight: "bold",
fontFamily: "monospace",
letterSpacing: "0.3em",
margin: "16px 0",
padding: "12px 24px",
background: "#f5f5f5",
borderRadius: 8,
display: "inline-block",
userSelect: "all",
}}
>
{deviceCodeState.userCode}
</div>
<div className="mb-4">
<Button
type="link"
onClick={() => window.open(deviceCodeState.verificationUri, "_blank")}
>
Open {deviceCodeState.verificationUri}
</Button>
</div>
<Spin />
<Text className="block mt-2 mb-4" type="secondary">
Waiting for {providerDisplayName} authorization...
</Text>
<Button onClick={handleCancel}>Cancel</Button>
</div>
);
case "success":
return (
<div className="text-center py-4">
<Text className="block mb-4" type="success" style={{ fontSize: "1.1rem" }}>
{providerDisplayName} credential &quot;{deviceCodeState.credentialName}&quot; created successfully!
</Text>
<Button type="primary" onClick={handleSuccessClose}>
Done
</Button>
</div>
);
case "error":
return (
<div className="text-center py-4">
<Text className="block mb-4" type="danger">
{deviceCodeState.message}
</Text>
<Button
onClick={() => setDeviceCodeState({ phase: "idle" })}
style={{ marginRight: 8 }}
>
Retry
</Button>
<Button onClick={handleCancel}>Cancel</Button>
</div>
);
const renderDeviceCodeSection = () => {
if (deviceCodeState.phase === "idle") {
return (
<div className="text-center py-4">
<Typography.Text className="block mb-4">
{deviceCodeProviderInfo?.provider_display_name || "Provider"} uses OAuth Device Code authorization. Click below to start.
</Typography.Text>
<Button type="primary" onClick={handleStartDeviceCode}>
Start {deviceCodeProviderInfo?.provider_display_name || "Provider"} Authorization
</Button>
</div>
);
}
if (deviceCodeState.phase === "success") {
return (
<div className="text-center py-4">
<Typography.Text className="block mb-4" type="success" style={{ fontSize: "1.1rem" }}>
Credential created successfully!
</Typography.Text>
<Button type="primary" onClick={handleSuccessClose}>
Done
</Button>
</div>
);
}
return renderDeviceCodeFlow();
};
return (
@ -323,10 +151,7 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
onChange={(value) => {
setSelectedProvider(value as Providers);
form.setFieldValue("custom_llm_provider", value);
// Reset device code state when provider changes
stopPolling();
setDeviceCodeState({ phase: "idle" });
accessTokenRef.current = null;
resetDeviceCode();
}}
>
{Object.entries(Providers).map(([providerEnum, providerDisplayName]) => (
@ -356,7 +181,7 @@ const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({ open, onCance
</Form.Item>
{isDeviceCodeProvider ? (
renderDeviceCodeFlow()
renderDeviceCodeSection()
) : (
<>
<ProviderSpecificFields selectedProvider={selectedProvider} uploadProps={uploadProps} />

View file

@ -4,6 +4,7 @@ import {
CredentialItem,
credentialUpdateCall,
} from "@/components/networking"; // Assume this is your networking function
import { GithubOutlined } from "@ant-design/icons";
import { PencilAltIcon, TrashIcon } from "@heroicons/react/outline";
import {
Badge,
@ -167,23 +168,14 @@ const CredentialsPanel: React.FC<CredentialsPanelProps> = ({ uploadProps }) => {
{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>
(<GithubOutlined className="w-4 h-4" />
<a
href={`https://github.com/${githubLogin}`}
target="_blank"
rel="noopener noreferrer"
>
{githubLogin}
</a>
)
</a>)
</span>
)}
</TableCell>

View file

@ -3569,18 +3569,19 @@ export const credentialDeleteCall = async (accessToken: string, credentialName:
}
};
export const githubCopilotInitiateAuth = async (
async function authenticatedPost<T>(
endpoint: string,
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`;
body?: Record<string, unknown>,
): Promise<T> {
const url = proxyBaseUrl ? `${proxyBaseUrl}${endpoint}` : endpoint;
const response = await fetch(url, {
method: "POST",
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
...(body && { body: JSON.stringify(body) }),
});
if (!response.ok) {
const errorData = await response.json();
@ -3589,78 +3590,27 @@ export const githubCopilotInitiateAuth = async (
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 githubCopilotInitiateAuth = (accessToken: string) =>
authenticatedPost<{ device_code: string; user_code: string; verification_uri: string; poll_interval_ms: number; expires_in: number }>(
"/credentials/github_copilot/initiate", accessToken,
);
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 githubCopilotCheckStatus = (accessToken: string, deviceCode: string) =>
authenticatedPost<{ status: string; access_token?: string; retry_after_ms?: number; error?: string }>(
"/credentials/github_copilot/status", accessToken, { device_code: deviceCode },
);
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 chatgptInitiateAuth = (accessToken: string) =>
authenticatedPost<{ device_auth_id: string; user_code: string; verification_uri: string; poll_interval_ms: number }>(
"/credentials/chatgpt/initiate", accessToken,
);
export const chatgptCheckStatus = (accessToken: string, deviceAuthId: string, userCode: string) =>
authenticatedPost<{ status: string; refresh_token?: string; account_id?: string; error?: string }>(
"/credentials/chatgpt/status", accessToken, { device_auth_id: deviceAuthId, user_code: userCode },
);
export const credentialUpdateCall = async (
accessToken: string,

View file

@ -0,0 +1,198 @@
import { useCallback, useEffect, useRef, useState } from "react";
import { Button, Spin, Typography } from "antd";
import {
type ProviderCreateInfo,
githubCopilotInitiateAuth,
githubCopilotCheckStatus,
chatgptInitiateAuth,
chatgptCheckStatus,
} from "@/components/networking";
const { Text } = Typography;
export type DeviceCodeState =
| { phase: "idle" }
| { phase: "polling"; deviceCode: string; userCode: string; verificationUri: string }
| { phase: "success" }
| { phase: "error"; message: string };
interface UseDeviceCodeFlowOptions {
accessToken: string | null;
providerInfo: ProviderCreateInfo | null;
/** Called with the resolved api key and litellm provider name on successful auth. */
onSuccess: (apiKey: string, litellmProvider: string) => void | Promise<void>;
}
export function useDeviceCodeFlow({ accessToken, providerInfo, onSuccess }: UseDeviceCodeFlowOptions) {
const [state, setState] = useState<DeviceCodeState>({ phase: "idle" });
const tokenRef = useRef<string | null>(null);
const timerRef = useRef<ReturnType<typeof setTimeout> | null>(null);
const stopPolling = useCallback(() => {
if (timerRef.current) {
clearTimeout(timerRef.current);
timerRef.current = null;
}
}, []);
useEffect(() => () => stopPolling(), [stopPolling]);
const reset = useCallback(() => {
stopPolling();
setState({ phase: "idle" });
tokenRef.current = null;
}, [stopPolling]);
const start = useCallback(async () => {
if (!accessToken || !providerInfo) return;
const litellmProvider = providerInfo.litellm_provider;
const label = providerInfo.provider_display_name || litellmProvider;
try {
let deviceId: string;
let userCode: string;
let verificationUri: string;
let pollIntervalMs: number;
if (litellmProvider === "chatgpt") {
const r = await chatgptInitiateAuth(accessToken);
deviceId = r.device_auth_id;
userCode = r.user_code;
verificationUri = r.verification_uri;
pollIntervalMs = r.poll_interval_ms;
} else {
const r = await githubCopilotInitiateAuth(accessToken);
deviceId = r.device_code;
userCode = r.user_code;
verificationUri = r.verification_uri;
pollIntervalMs = r.poll_interval_ms;
}
setState({ phase: "polling", deviceCode: deviceId, userCode, verificationUri });
let interval = pollIntervalMs || 5000;
const schedulePoll = (delayMs: number) => {
timerRef.current = setTimeout(async () => {
try {
let apiKey: string | undefined;
let failed = false;
let errorMsg: string | undefined;
let retryAfterMs: number | undefined;
if (litellmProvider === "chatgpt") {
const s = await chatgptCheckStatus(accessToken, deviceId, userCode);
if (s.status === "complete" && s.refresh_token) {
apiKey = s.refresh_token;
} else if (s.status === "failed") {
failed = true;
errorMsg = s.error;
}
} else {
const s = await githubCopilotCheckStatus(accessToken, deviceId);
if (s.status === "complete" && s.access_token) {
apiKey = s.access_token;
} else if (s.status === "failed") {
failed = true;
errorMsg = s.error;
} else {
retryAfterMs = s.retry_after_ms ?? undefined;
}
}
if (apiKey) {
stopPolling();
tokenRef.current = apiKey;
try {
await onSuccess(apiKey, litellmProvider);
setState({ phase: "success" });
} catch {
setState({ phase: "error", message: "Authorization succeeded but callback failed" });
}
} else if (failed) {
stopPolling();
setState({ phase: "error", message: errorMsg || "Authorization failed" });
} else {
if (retryAfterMs != null) interval = retryAfterMs;
schedulePoll(interval);
}
} catch {
stopPolling();
setState({ phase: "error", message: "Failed to check authorization status" });
}
}, delayMs);
};
schedulePoll(interval);
} catch {
setState({ phase: "error", message: `Failed to start ${label} authorization` });
}
}, [accessToken, providerInfo, onSuccess, stopPolling]);
const providerLabel = providerInfo?.provider_display_name || "Provider";
const renderUI = useCallback(() => {
switch (state.phase) {
case "idle":
return (
<div className="text-center py-2">
<Button type="primary" onClick={start}>
Authorize with {providerLabel}
</Button>
</div>
);
case "polling":
return (
<div className="text-center py-2">
<Text className="block mb-2">Enter this code to authorize:</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",
}}
>
{state.userCode}
</div>
<div className="mb-3">
<Button type="link" onClick={() => window.open(state.verificationUri, "_blank")}>
Open {state.verificationUri}
</Button>
</div>
<Spin />
<Text className="block mt-2 mb-3" type="secondary">
Waiting for {providerLabel} authorization...
</Text>
<Button onClick={reset}>Cancel</Button>
</div>
);
case "success":
return (
<div className="text-center py-2">
<Text type="success" className="block mb-2">
Authorization complete!
</Text>
</div>
);
case "error":
return (
<div className="text-center py-2">
<Text type="danger" className="block mb-3">
{state.message}
</Text>
<Button style={{ marginRight: 8 }} onClick={() => setState({ phase: "idle" })}>
Retry
</Button>
<Button onClick={reset}>Cancel</Button>
</div>
);
}
}, [state, start, reset, providerLabel]);
return { state, start, reset, tokenRef, renderUI };
}