diff --git a/src/components/providers/utils.ts b/src/components/providers/utils.ts index 3e92cfe8..75dd4470 100644 --- a/src/components/providers/utils.ts +++ b/src/components/providers/utils.ts @@ -10,6 +10,7 @@ import { } from '@/utils/recentRequests'; const DISABLE_ALL_MODELS_RULE = '*'; +const DEFAULT_GEMINI_BASE_URL = 'https://generativelanguage.googleapis.com'; export const hasDisableAllModelsRule = (models?: string[]) => Array.isArray(models) && @@ -52,6 +53,33 @@ const normalizeClaudeBaseUrl = (baseUrl: string): string => { return trimmed; }; +const normalizeGeminiBaseUrl = (baseUrl: string): string => { + let trimmed = String(baseUrl || '').trim(); + if (!trimmed) { + return DEFAULT_GEMINI_BASE_URL; + } + trimmed = trimmed.replace(/\/?v0\/management\/?$/i, ''); + trimmed = trimmed.replace(/\/+$/g, ''); + if (!/^https?:\/\//i.test(trimmed)) { + trimmed = `http://${trimmed}`; + } + return trimmed; +}; + +const buildGeminiModelResource = (model: string): string => { + const trimmed = String(model || '') + .trim() + .replace(/^\/+/g, '') + .replace(/:generateContent$/i, ''); + if (!trimmed) return ''; + + if (/^(models|tunedModels)\//i.test(trimmed)) { + return trimmed.split('/').map(encodeURIComponent).join('/'); + } + + return `models/${encodeURIComponent(trimmed)}`; +}; + export const buildOpenAIChatCompletionsEndpoint = (baseUrl: string): string => { const trimmed = normalizeOpenAIBaseUrl(baseUrl); if (!trimmed) return ''; @@ -73,6 +101,30 @@ export const buildClaudeMessagesEndpoint = (baseUrl: string): string => { return `${trimmed}/v1/messages`; }; +export const buildGeminiGenerateContentEndpoint = ( + baseUrl: string, + model: string +): string => { + const resource = buildGeminiModelResource(model); + if (!resource) return ''; + + const trimmed = normalizeGeminiBaseUrl(baseUrl); + if (!trimmed) return ''; + if (/:generateContent$/i.test(trimmed)) { + return trimmed; + } + + let root = trimmed.replace(/\/+$/g, ''); + if (/\/v1beta\/models$/i.test(root)) { + root = root.replace(/\/models$/i, ''); + } else if (!/\/v1beta$/i.test(root)) { + root = root.replace(/\/v1beta(?:\/.*)?$/i, ''); + root = `${root}/v1beta`; + } + + return `${root}/${resource}:generateContent`; +}; + export type ProviderRecentUsageMap = Map>; const EMPTY_RECENT_USAGE_ENTRY: RecentRequestUsageEntry = { diff --git a/src/features/providers/descriptors.ts b/src/features/providers/descriptors.ts index f023f669..672da9df 100644 --- a/src/features/providers/descriptors.ts +++ b/src/features/providers/descriptors.ts @@ -36,7 +36,7 @@ export const PROVIDER_DESCRIPTORS: Record = { supportsHeaders: true, supportsExcludedModels: true, supportsPriority: true, - supportsTestModel: false, + supportsTestModel: true, supportsWebsockets: false, supportsCloak: false, supportsApiKeyEntries: false, diff --git a/src/features/providers/sheets/forms/BaseProviderForm.tsx b/src/features/providers/sheets/forms/BaseProviderForm.tsx index c3d50c4c..d4a2881e 100644 --- a/src/features/providers/sheets/forms/BaseProviderForm.tsx +++ b/src/features/providers/sheets/forms/BaseProviderForm.tsx @@ -76,7 +76,10 @@ function buildInitialForm( websockets: brand === 'codex' ? false : undefined, cloak: brand === 'claude' ? { mode: '', strictMode: false, sensitiveWordsText: '' } : undefined, - testModel: brand === 'openaiCompatibility' || brand === 'claude' ? '' : undefined, + testModel: + brand === 'openaiCompatibility' || brand === 'claude' || brand === 'gemini' + ? '' + : undefined, apiKeyEntries: brand === 'openaiCompatibility' ? [emptyApiKeyEntry()] : undefined, }; } @@ -152,7 +155,7 @@ function buildInitialForm( sensitiveWordsText: (cfg as ProviderKeyConfig).cloak?.sensitiveWords?.join('\n') ?? '', } : undefined, - testModel: brand === 'claude' ? '' : undefined, + testModel: brand === 'claude' || brand === 'gemini' ? '' : undefined, }; } @@ -415,6 +418,12 @@ export function BaseProviderForm({ [form.apiKeyEntries] ); const actualApiKeyEntries = form.apiKeyEntries ?? []; + const singleConnectivity = + brand === 'gemini' + ? { status: connectivity.geminiStatus, run: connectivity.runGemini } + : brand === 'claude' + ? { status: connectivity.claudeStatus, run: connectivity.runClaude } + : null; const removeApiKeyEntry = (removeIdx: number) => { setShowPasswords((prev) => { @@ -578,7 +587,7 @@ export function BaseProviderForm({
) : null} diff --git a/src/features/providers/sheets/forms/useConnectivityTest.ts b/src/features/providers/sheets/forms/useConnectivityTest.ts index bf2d0768..07b1de5a 100644 --- a/src/features/providers/sheets/forms/useConnectivityTest.ts +++ b/src/features/providers/sheets/forms/useConnectivityTest.ts @@ -2,6 +2,7 @@ import { useCallback, useEffect, useMemo, useRef, useState } from 'react'; import { apiCallApi, getApiCallErrorMessage } from '@/services/api'; import { buildClaudeMessagesEndpoint, + buildGeminiGenerateContentEndpoint, buildOpenAIChatCompletionsEndpoint, } from '@/components/providers/utils'; import { buildHeaderObject, hasHeader } from '@/utils/headers'; @@ -25,6 +26,18 @@ const errorMessage = (err: unknown): string => { return ''; }; +const requestFailureMessage = (err: unknown, messages: ConnectivityErrorMessages): string => { + const raw = errorMessage(err); + const isTimeout = + (typeof err === 'object' && + err !== null && + 'code' in err && + String((err as { code?: string }).code) === 'ECONNABORTED') || + raw.toLowerCase().includes('timeout'); + + return isTimeout ? messages.timeout(DEFAULT_TIMEOUT_MS / 1000) : raw || messages.requestFailed; +}; + const pickModel = (testModel: string | undefined, models: ModelEntryInput[]): string => { const trimmed = (testModel ?? '').trim(); if (trimmed) return trimmed; @@ -65,10 +78,12 @@ export interface ConnectivityErrorMessages { export interface UseConnectivityTestResult { openaiStatuses: ConnectivityStatus[]; + geminiStatus: ConnectivityStatus; claudeStatus: ConnectivityStatus; isTestingAny: boolean; runOpenAIKey: (idx: number) => Promise; runOpenAIAllKeys: () => Promise; + runGemini: () => Promise; runClaude: () => Promise; } @@ -93,6 +108,7 @@ export function useConnectivityTest( const [openaiStatuses, setOpenaiStatuses] = useState(() => Array.from({ length: entriesCount }, () => IDLE) ); + const [geminiStatus, setGeminiStatus] = useState(IDLE); const [claudeStatus, setClaudeStatus] = useState(IDLE); const [inFlight, setInFlight] = useState(0); @@ -133,14 +149,23 @@ export function useConnectivityTest( const signature = useMemo(() => { const h = formHeaders.map((it) => `${it.key}:${it.value}`).join('|'); const m = models.map((it) => `${it.name}:${it.alias ?? ''}`).join('|'); - return `${baseUrl}||${(testModel ?? '').trim()}||${h}||${m}`; - }, [baseUrl, testModel, formHeaders, models]); + return [ + baseUrl, + (testModel ?? '').trim(), + apiKey ?? '', + fallbackApiKey ?? '', + authIndex ?? '', + h, + m, + ].join('||'); + }, [apiKey, authIndex, baseUrl, fallbackApiKey, testModel, formHeaders, models]); const lastSignatureRef = useRef(signature); useEffect(() => { if (lastSignatureRef.current === signature) return; lastSignatureRef.current = signature; setOpenaiStatuses((prev) => prev.map(() => IDLE)); + setGeminiStatus(IDLE); setClaudeStatus(IDLE); }, [signature]); @@ -228,18 +253,9 @@ export function useConnectivityTest( updateOpenaiStatus(idx, { state: 'success', message: '' }); return true; } catch (err) { - const raw = errorMessage(err); - const isTimeout = - (typeof err === 'object' && - err !== null && - 'code' in err && - String((err as { code?: string }).code) === 'ECONNABORTED') || - raw.toLowerCase().includes('timeout'); updateOpenaiStatus(idx, { state: 'error', - message: isTimeout - ? messages.timeout(DEFAULT_TIMEOUT_MS / 1000) - : raw || messages.requestFailed, + message: requestFailureMessage(err, messages), }); return false; } finally { @@ -266,6 +282,75 @@ export function useConnectivityTest( await Promise.all(entries.map((_, idx) => runOpenAIKey(idx))); }, [apiKeyEntries, brand, runOpenAIKey]); + const runGemini = useCallback(async (): Promise => { + if (brand !== 'gemini') return; + + const model = pickModel(testModel, models); + if (!model) { + setGeminiStatus({ state: 'error', message: messages.modelRequired }); + return; + } + + const endpoint = buildGeminiGenerateContentEndpoint(baseUrl ?? '', model); + if (!endpoint) { + setGeminiStatus({ state: 'error', message: messages.endpointInvalid }); + return; + } + + const customHeaders = buildHeaderObject(formHeaders); + const explicitKey = (apiKey ?? '').trim(); + const persistedKey = (fallbackApiKey ?? '').trim(); + const hasApiKeyHeader = hasHeader(customHeaders, 'x-goog-api-key'); + const resolvedKey = explicitKey || persistedKey; + const resolvedAuthIndex = (authIndex ?? '').trim() || undefined; + + if (!resolvedKey && !hasApiKeyHeader && !resolvedAuthIndex) { + setGeminiStatus({ state: 'error', message: messages.apiKeyRequired }); + return; + } + + const headerObj: Record = { + 'Content-Type': 'application/json', + ...customHeaders, + }; + if (!hasHeader(headerObj, 'x-goog-api-key')) { + if (resolvedKey) { + headerObj['x-goog-api-key'] = resolvedKey; + } else if (resolvedAuthIndex) { + headerObj['x-goog-api-key'] = '$TOKEN$'; + } + } + + setGeminiStatus({ state: 'loading', message: '' }); + setInFlight((n) => n + 1); + try { + const result = await apiCallApi.request( + { + authIndex: resolvedAuthIndex, + method: 'POST', + url: endpoint, + header: headerObj, + data: JSON.stringify({ + contents: [{ parts: [{ text: 'Hi' }] }], + generationConfig: { maxOutputTokens: 8 }, + }), + }, + { timeout: DEFAULT_TIMEOUT_MS } + ); + if (result.statusCode < 200 || result.statusCode >= 300) { + throw new Error(getApiCallErrorMessage(result)); + } + setGeminiStatus({ state: 'success', message: '' }); + } catch (err) { + setGeminiStatus({ + state: 'error', + message: requestFailureMessage(err, messages), + }); + } finally { + setInFlight((n) => n - 1); + } + }, [apiKey, authIndex, baseUrl, brand, fallbackApiKey, formHeaders, messages, models, testModel]); + const runClaude = useCallback(async (): Promise => { if (brand !== 'claude') return; @@ -328,18 +413,9 @@ export function useConnectivityTest( } setClaudeStatus({ state: 'success', message: '' }); } catch (err) { - const raw = errorMessage(err); - const isTimeout = - (typeof err === 'object' && - err !== null && - 'code' in err && - String((err as { code?: string }).code) === 'ECONNABORTED') || - raw.toLowerCase().includes('timeout'); setClaudeStatus({ state: 'error', - message: isTimeout - ? messages.timeout(DEFAULT_TIMEOUT_MS / 1000) - : raw || messages.requestFailed, + message: requestFailureMessage(err, messages), }); } finally { setInFlight((n) => n - 1); @@ -348,10 +424,12 @@ export function useConnectivityTest( return { openaiStatuses, + geminiStatus, claudeStatus, isTestingAny: inFlight > 0, runOpenAIKey, runOpenAIAllKeys, + runGemini, runClaude, }; } diff --git a/src/features/providers/types.ts b/src/features/providers/types.ts index c0e16b52..7b7329a1 100644 --- a/src/features/providers/types.ts +++ b/src/features/providers/types.ts @@ -121,7 +121,7 @@ export interface ProviderEntryFormInput { websockets?: boolean; /** Claude 专属 */ cloak?: CloakInput; - /** OpenAI 专属 */ + /** OpenAI persists this; Gemini/Claude use it for one-off connectivity tests. */ testModel?: string; apiKeyEntries?: ApiKeyEntryInput[]; }