mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
b7432a4c44
commit
de31fc4471
13 changed files with 350 additions and 577 deletions
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -328,7 +328,7 @@ class ProxyInitializationHelpers:
|
|||
import uvloop # noqa: F401
|
||||
|
||||
return "uvloop"
|
||||
except (ImportError, Exception):
|
||||
except ImportError:
|
||||
return "asyncio"
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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} />
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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 "{deviceCodeState.credentialName}" 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} />
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
198
ui/litellm-dashboard/src/hooks/useDeviceCodeFlow.tsx
Normal file
198
ui/litellm-dashboard/src/hooks/useDeviceCodeFlow.tsx
Normal 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 };
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue