diff --git a/packages/stage-ui/src/libs/providers/validators/openai-compatible.test.ts b/packages/stage-ui/src/libs/providers/validators/openai-compatible.test.ts index dde3cb3e7..b1fa16d55 100644 --- a/packages/stage-ui/src/libs/providers/validators/openai-compatible.test.ts +++ b/packages/stage-ui/src/libs/providers/validators/openai-compatible.test.ts @@ -1,4 +1,8 @@ -import { beforeEach, describe, expect, it, vi } from 'vitest' +import type { ComposerTranslation } from 'vue-i18n' + +import type { ProviderExtraMethods, ProviderInstance } from '../types' + +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' import { createOpenAICompatibleValidators } from './openai-compatible' @@ -18,28 +22,68 @@ vi.mock('@xsai/model', () => ({ listModels: listModelsMock, })) +const mockT = vi.fn((key: string) => key) as unknown as ComposerTranslation + function getProviderValidators(options?: Parameters[0]) { const validators = createOpenAICompatibleValidators(options) - return (validators?.validateProvider || []).map(create => create({ - t: (input: string) => input, - } as any)) + return (validators?.validateProvider || []).map(create => create({ t: mockT })) } +interface TestConfig { apiKey?: string, baseUrl?: string } + describe('createOpenAICompatibleValidators', () => { - const config = { + const config: TestConfig = { apiKey: 'test-key', baseUrl: 'https://example.com/v1/', } - const provider = { + const provider: ProviderInstance = { model: () => ({ apiKey: config.apiKey, baseURL: config.baseUrl, }), - } + } as ProviderInstance + const providerExtra: ProviderExtraMethods = {} + + let fetchMock: ReturnType beforeEach(() => { vi.clearAllMocks() + fetchMock = vi.fn().mockResolvedValue(new Response('{}', { status: 200 })) + vi.stubGlobal('fetch', fetchMock) + }) + + afterEach(() => { + vi.unstubAllGlobals() + }) + + it('connectivity check uses lightweight fetch instead of generateText', async () => { + const [connectivityValidator] = getProviderValidators({ + checks: ['connectivity'], + }) + + const result = await connectivityValidator.validator(config, provider, providerExtra, { t: mockT }) + + expect(result.valid).toBe(true) + expect(generateTextMock).not.toHaveBeenCalled() + expect(fetchMock).toHaveBeenCalledWith( + 'https://example.com/v1/models', + expect.objectContaining({ method: 'GET' }), + ) + }) + + it('connectivity check fails on network error', async () => { + fetchMock.mockRejectedValue(new TypeError('fetch failed')) + + const [connectivityValidator] = getProviderValidators({ + checks: ['connectivity'], + }) + + const result = await connectivityValidator.validator(config, provider, providerExtra, { t: mockT }) + + expect(result.valid).toBe(false) + expect(result.reason).toContain('Connectivity check failed') + expect(generateTextMock).not.toHaveBeenCalled() }) it('does not probe chat completions with a synthetic fallback model', async () => { @@ -49,11 +93,10 @@ describe('createOpenAICompatibleValidators', () => { checks: ['connectivity', 'chat_completions'], }) - const connectivityResult = await connectivityValidator.validator(config, provider as any, undefined as any, undefined as any) - const chatResult = await chatValidator.validator(config, provider as any, undefined as any, undefined as any) + const connectivityResult = await connectivityValidator.validator(config, provider, providerExtra, { t: mockT }) + const chatResult = await chatValidator.validator(config, provider, providerExtra, { t: mockT }) - expect(connectivityResult.valid).toBe(false) - expect(connectivityResult.reason).toContain('No model available for validation.') + expect(connectivityResult.valid).toBe(true) expect(chatResult.valid).toBe(false) expect(chatResult.reason).toContain('No model available for validation.') expect(generateTextMock).not.toHaveBeenCalled() @@ -67,11 +110,20 @@ describe('createOpenAICompatibleValidators', () => { allowValidationWithoutModel: true, }) - const connectivityResult = await connectivityValidator.validator(config, provider as any, undefined as any, undefined as any) - const chatResult = await chatValidator.validator(config, provider as any, undefined as any, undefined as any) + const connectivityResult = await connectivityValidator.validator(config, provider, providerExtra, { t: mockT }) + const chatResult = await chatValidator.validator(config, provider, providerExtra, { t: mockT }) expect(connectivityResult.valid).toBe(true) expect(chatResult.valid).toBe(true) expect(generateTextMock).not.toHaveBeenCalled() }) + + it('default checks do not include chat_completions', () => { + const validators = getProviderValidators() + const ids = validators.map(v => v.id) + + expect(ids).toContain('openai-compatible:check-connectivity') + expect(ids).toContain('openai-compatible:check-model-list') + expect(ids).not.toContain('openai-compatible:check-chat-completions') + }) }) diff --git a/packages/stage-ui/src/libs/providers/validators/openai-compatible.ts b/packages/stage-ui/src/libs/providers/validators/openai-compatible.ts index 71bce86d2..c33ad9e30 100644 --- a/packages/stage-ui/src/libs/providers/validators/openai-compatible.ts +++ b/packages/stage-ui/src/libs/providers/validators/openai-compatible.ts @@ -103,7 +103,7 @@ async function pickValidationModel( options?: OpenAICompatibleValidationOptions, ): ProviderDefinition['validators'] { - const checks = options?.checks ?? ['connectivity', 'model_list', 'chat_completions'] + const checks = options?.checks ?? ['connectivity', 'model_list'] const additionalHeaders = options?.additionalHeaders const missingValidationModelReason = 'No model available for validation. Configure a model manually and try again.' @@ -240,21 +240,41 @@ export function createOpenAICompatibleValidators { + validator: async (config) => { const errors: Array<{ error: unknown }> = [] - const result = await getChatCheckResult( - config, - provider, - providerExtra, - contextOptions as { validationCache?: Map } | undefined, - ) - if (!result.connectivityOk) { - const errorMessage = result.errorMessage || 'Unknown error.' + const baseUrl = String(config.baseUrl ?? '') + const modelsUrl = baseUrl.endsWith('/') ? `${baseUrl}models` : `${baseUrl}/models` + const controller = new AbortController() + const timeout = setTimeout(() => controller.abort(), 10_000) + + try { + const response = await fetch(modelsUrl, { + method: 'GET', + headers: { + ...(config.apiKey ? { Authorization: `Bearer ${config.apiKey}` } : {}), + ...additionalHeaders, + }, + signal: controller.signal, + }) + + if (response.status >= 500) { + const errorMessage = `Server error: HTTP ${response.status}` + const reason = options?.connectivityFailureReason + ? options.connectivityFailureReason({ config, error: new Error(errorMessage), errorMessage }) + : `Connectivity check failed: ${errorMessage}` + errors.push({ error: new Error(reason) }) + } + } + catch (e) { + const errorMessage = errorMessageFrom(e) || 'Unknown error.' const reason = options?.connectivityFailureReason - ? options.connectivityFailureReason({ config, error: result.error, errorMessage }) + ? options.connectivityFailureReason({ config, error: e, errorMessage }) : `Connectivity check failed: ${errorMessage}` errors.push({ error: new Error(reason) }) } + finally { + clearTimeout(timeout) + } return { errors,