From de31fc4471ef10227f17596fcbe7a4a550c6f3b3 Mon Sep 17 00:00:00 2001 From: Hunter Wittenborn Date: Mon, 30 Mar 2026 05:19:56 -0500 Subject: [PATCH] refactor: DRY up device code OAuth flow, fix quality/efficiency issues MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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) --- litellm/llms/chatgpt/authenticator.py | 30 +- litellm/llms/github_copilot/authenticator.py | 10 +- litellm/main.py | 30 +- .../proxy/credential_endpoints/endpoints.py | 22 +- litellm/proxy/proxy_cli.py | 2 +- .../test_github_copilot_authenticator.py | 6 +- .../ModelsAndEndpointsView.tsx | 3 +- .../src/components/add_model/AddModelForm.tsx | 222 ++------------ .../components/add_model/add_model_tab.tsx | 3 - .../model_add/AddCredentialModal.tsx | 289 ++++-------------- .../src/components/model_add/credentials.tsx | 14 +- .../src/components/networking.tsx | 98 ++---- .../src/hooks/useDeviceCodeFlow.tsx | 198 ++++++++++++ 13 files changed, 350 insertions(+), 577 deletions(-) create mode 100644 ui/litellm-dashboard/src/hooks/useDeviceCodeFlow.tsx diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 0308c5f76a9..930ac9f51d4 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -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}", diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index f17cd659393..8ecd0b63298 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -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, diff --git a/litellm/main.py b/litellm/main.py index 56b46c09015..035f3dd00a1 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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 diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 42cbd45aa23..5348d6011ca 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -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, } ) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index e843916595a..a30b36b8961 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -328,7 +328,7 @@ class ProxyInitializationHelpers: import uvloop # noqa: F401 return "uvloop" - except (ImportError, Exception): + except ImportError: return "asyncio" @staticmethod diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py index 92c892faad6..a1500662015 100644 --- a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py +++ b/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py @@ -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): diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index f5448998a92..514ae673d06 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -72,7 +72,7 @@ const ModelsAndEndpointsView: React.FC = ({ 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 = ({ premiumUser, te setShowAdvancedSettings={setShowAdvancedSettings} teams={teams} credentials={credentialsList} - refetchCredentials={refetchCredentials} accessToken={accessToken} userRole={userRole} /> diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index c39136a5161..07070557fce 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -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 = ({ setShowAdvancedSettings, teams, credentials, - refetchCredentials, }) => { const [testMode, setTestMode] = useState("chat"); const [isResultModalVisible, setIsResultModalVisible] = useState(false); @@ -75,28 +69,8 @@ const AddModelForm: React.FC = ({ 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(null); - const ghPollingRef = useRef | 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 = ({ }, [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 ( -
- -
- ); - case "polling": - return ( -
- Enter this code to authorize: -
- {ghDeviceCodeState.userCode} -
-
- -
- - Waiting for {dcProviderLabel} authorization... - -
- ); - case "success": - return ( -
- - ✓ Authorization complete! Submit the form to add the model. - -
- ); - case "error": - return ( -
- {ghDeviceCodeState.message} - - -
- ); - } - }; const { data: guardrailsList, isLoading: isGuardrailsLoading, error: guardrailsError } = useGuardrails(); const { data: tagsList, isLoading: isTagsLoading, error: tagsError } = useTags(); @@ -490,7 +306,7 @@ const AddModelForm: React.FC = ({
{isDeviceCodeProvider ? ( - renderGhDeviceCodeFlow() + renderDeviceCodeFlow() ) : ( )} diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx index 75d02f13fde..f9b6533ac60 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx @@ -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 = ({ setShowAdvancedSettings, teams, credentials, - refetchCredentials, accessToken, userRole, }) => { @@ -81,7 +79,6 @@ const AddModelTab: React.FC = ({ setShowAdvancedSettings={setShowAdvancedSettings} teams={teams} credentials={credentials} - refetchCredentials={refetchCredentials} /> diff --git a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx index 8ad4831d336..83cbafe72f9 100644 --- a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx @@ -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 = ({ open, onCancel, onAddCredential, uploadProps, initialCredentialName, initialProvider }) => { const [form] = Form.useForm(); const [selectedProvider, setSelectedProvider] = useState(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(null); - const pollingRef = useRef | 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 = ({ 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 = ({ 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 ( -
- - {providerDisplayName} uses OAuth Device Code authorization. Click below to start. - - -
- ); - case "polling": - return ( -
- - Enter this code to authorize: - -
- {deviceCodeState.userCode} -
-
- -
- - - Waiting for {providerDisplayName} authorization... - - -
- ); - case "success": - return ( -
- - {providerDisplayName} credential "{deviceCodeState.credentialName}" created successfully! - - -
- ); - case "error": - return ( -
- - {deviceCodeState.message} - - - -
- ); + const renderDeviceCodeSection = () => { + if (deviceCodeState.phase === "idle") { + return ( +
+ + {deviceCodeProviderInfo?.provider_display_name || "Provider"} uses OAuth Device Code authorization. Click below to start. + + +
+ ); } + if (deviceCodeState.phase === "success") { + return ( +
+ + Credential created successfully! + + +
+ ); + } + return renderDeviceCodeFlow(); }; return ( @@ -323,10 +151,7 @@ const AddCredentialsModal: React.FC = ({ 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 = ({ open, onCance {isDeviceCodeProvider ? ( - renderDeviceCodeFlow() + renderDeviceCodeSection() ) : ( <> diff --git a/ui/litellm-dashboard/src/components/model_add/credentials.tsx b/ui/litellm-dashboard/src/components/model_add/credentials.tsx index 315d2fc01d2..92a793ce03d 100644 --- a/ui/litellm-dashboard/src/components/model_add/credentials.tsx +++ b/ui/litellm-dashboard/src/components/model_add/credentials.tsx @@ -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 = ({ uploadProps }) => { {credential.credential_name} {githubLogin && ( - ( - + ( {githubLogin} - - ) + ) )} diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 47884e4091b..68d3ab3049b 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -3569,18 +3569,19 @@ export const credentialDeleteCall = async (accessToken: string, credentialName: } }; -export const githubCopilotInitiateAuth = async ( +async function authenticatedPost( + 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, +): Promise { + 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, diff --git a/ui/litellm-dashboard/src/hooks/useDeviceCodeFlow.tsx b/ui/litellm-dashboard/src/hooks/useDeviceCodeFlow.tsx new file mode 100644 index 00000000000..01f60861cdd --- /dev/null +++ b/ui/litellm-dashboard/src/hooks/useDeviceCodeFlow.tsx @@ -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; +} + +export function useDeviceCodeFlow({ accessToken, providerInfo, onSuccess }: UseDeviceCodeFlowOptions) { + const [state, setState] = useState({ phase: "idle" }); + const tokenRef = useRef(null); + const timerRef = useRef | 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 ( +
+ +
+ ); + case "polling": + return ( +
+ Enter this code to authorize: +
+ {state.userCode} +
+
+ +
+ + + Waiting for {providerLabel} authorization... + + +
+ ); + case "success": + return ( +
+ + Authorization complete! + +
+ ); + case "error": + return ( +
+ + {state.message} + + + +
+ ); + } + }, [state, start, reset, providerLabel]); + + return { state, start, reset, tokenRef, renderUI }; +}