From c24aae6de020dd2a659f78b397476b5688b40636 Mon Sep 17 00:00:00 2001 From: "dayuan.jiang" Date: Sun, 4 Oct 2026 13:31:47 +0900 Subject: [PATCH] refactor(providers): one model factory for chat and the settings Test button - getAIModel resolves credentials (client key, server env vars, the existing SSRF rules) and createModel builds the model by SDK. The provider-by-provider switch shrinks from 24 cases to the few that differ (lib/ai-providers.ts 1531 -> 1106 lines) - /api/validate-model calls getAIModel instead of its own 24-case switch (503 -> 175 lines), which had drifted from the chat: it built Azure with createOpenAI, Kimi and MiMo with createOpenAI instead of createDeepSeek, and the official OpenAI endpoint with Chat Completions. A passing test now means the chat works - Plain OpenAI-compatible providers (SiliconFlow, SGLang, ModelScope, GLM, Qwen, Qiniu, Novita, Atlas Cloud, EdgeOne, Doubao, MiniMax in OpenAI mode, AIHubMix on a custom URL) use @ai-sdk/openai-compatible, which reads reasoning_content, so their reasoning shows, and accepts SGLang's stream as is (its 95-line stream rewrite is gone). includeUsage keeps token usage for quotas. tags in their text become reasoning (extractReasoningMiddleware) - SGLang without a base URL used OpenAI's endpoint; it now defaults to http://127.0.0.1:8000/v1 like the Test button did - Chat requests to a client base URL refuse redirects, as the Test button already did (redirectGuardedFetch moves to lib/ssrf-protection) - The Test button streams like the chat (the ModelScope special case is gone), times out after 15 s, does not retry, asks the model to call a ping tool and warns when it answers without one, and tests all models at once. The time each test took shows on its check mark - Unknown provider names are rejected with Object.hasOwn, and the error texts list providers from PROVIDER_INFO instead of hand-kept lists --- app/api/validate-model/route.ts | 440 ++--------- components/model-config-dialog.tsx | 185 ++--- lib/ai-providers.ts | 871 ++++++---------------- lib/ssrf-protection.ts | 16 + lib/types/model-config.ts | 2 + tests/e2e/model-test.spec.ts | 67 ++ tests/unit/ai-providers.test.ts | 26 +- tests/unit/validate-diagram-route.test.ts | 6 +- tests/unit/validate-model-route.test.ts | 116 +++ 9 files changed, 609 insertions(+), 1120 deletions(-) create mode 100644 tests/e2e/model-test.spec.ts create mode 100644 tests/unit/validate-model-route.test.ts diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts index 6bef18eb..1bebcb14 100644 --- a/app/api/validate-model/route.ts +++ b/app/api/validate-model/route.ts @@ -1,22 +1,9 @@ -import { createAmazonBedrock } from "@ai-sdk/amazon-bedrock" -import { createAnthropic } from "@ai-sdk/anthropic" -import { createDeepSeek, deepseek } from "@ai-sdk/deepseek" -import { createGoogleGenerativeAI } from "@ai-sdk/google" -import { createVertex } from "@ai-sdk/google-vertex" -import { createOpenAI } from "@ai-sdk/openai" -import { createAihubmix } from "@aihubmix/ai-sdk-provider" -import { createOpenRouter } from "@openrouter/ai-sdk-provider" -import { createGateway, generateText } from "ai" +import { streamText, tool } from "ai" import { NextResponse } from "next/server" -import { createOllama } from "ollama-ai-provider-v2" +import { z } from "zod" import { checkAccessCode } from "@/lib/access-code" -import { - AIHUBMIX_APP_CODE, - isAihubmixStandardBaseURL, - normalizeMiniMaxBaseURL, -} from "@/lib/ai-providers" +import { getAIModel } from "@/lib/ai-providers" import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection" -import { PROVIDER_INFO, type ProviderName } from "@/lib/types/model-config" export const runtime = "nodejs" @@ -33,18 +20,16 @@ interface ValidateRequest { vertexApiKey?: string // Express Mode API key } -// With private URLs blocked, a public baseUrl could still redirect the -// request to an internal host, so redirects are refused in that case. -function redirectGuardedFetch(): typeof fetch | undefined { - if (allowPrivateUrls()) return undefined - return async (input, init) => { - const response = await fetch(input, { ...init, redirect: "manual" }) - if (response.status >= 300 && response.status < 400) { - throw new Error("Redirects are not allowed for custom base URLs") - } - return response - } -} +const TEST_TIMEOUT_MS = 15_000 + +// Drawing works through tool calls, so the test asks for one +const PING_TOOL = tool({ + description: "Report that the connection works.", + inputSchema: z.object({}), +}) + +const NO_TOOL_CALL_WARNING = + "Connected, but the model answered without calling a tool. It may not support tool calls, which drawing needs." export async function POST(req: Request) { // Lets the server send requests to arbitrary URLs, so require the access code @@ -108,364 +93,55 @@ export async function POST(req: Request) { ) } - const guardedFetch = redirectGuardedFetch() - let model: any - - switch (provider) { - case "openai": { - const openai = createOpenAI({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - fetch: guardedFetch, - }) - model = openai.chat(modelId) - break - } - - case "anthropic": { - const anthropic = createAnthropic({ - apiKey, - baseURL: baseUrl || "https://api.anthropic.com/v1", - fetch: guardedFetch, - }) - model = anthropic(modelId) - break - } - - case "google": { - const google = createGoogleGenerativeAI({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - fetch: guardedFetch, - }) - model = google(modelId) - break - } - - case "vertexai": { - const vertex = createVertex({ - apiKey: vertexApiKey, - ...(baseUrl && { baseURL: baseUrl }), - fetch: guardedFetch, - }) - model = vertex(modelId) - break - } - - case "azure": { - const azure = createOpenAI({ - apiKey, - baseURL: baseUrl, - fetch: guardedFetch, - }) - model = azure.chat(modelId) - break - } - - case "bedrock": { - const bedrock = createAmazonBedrock({ - accessKeyId: awsAccessKeyId, - secretAccessKey: awsSecretAccessKey, - region: awsRegion, - }) - model = bedrock(modelId) - break - } - - case "openrouter": { - const openrouter = createOpenRouter({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - fetch: guardedFetch, - }) - model = openrouter(modelId) - break - } - - case "aihubmix": { - const defaultBaseURL = PROVIDER_INFO.aihubmix.defaultBaseUrl - - if ( - isAihubmixStandardBaseURL(baseUrl) || - baseUrl === defaultBaseURL - ) { - const aihubmix = createAihubmix({ - apiKey, - appCode: AIHUBMIX_APP_CODE, - }) - model = aihubmix(modelId) - } else { - const aihubmixCompatible = createOpenAI({ - apiKey, - baseURL: baseUrl, - fetch: guardedFetch, - }) - model = aihubmixCompatible.chat(modelId) - } - break - } - - case "deepseek": { - if (baseUrl || apiKey) { - const ds = createDeepSeek({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - fetch: guardedFetch, - }) - model = ds(modelId) - } else { - model = deepseek(modelId) - } - break - } - - case "siliconflow": { - const sf = createOpenAI({ - apiKey, - baseURL: baseUrl || "https://api.siliconflow.cn/v1", - fetch: guardedFetch, - }) - model = sf.chat(modelId) - break - } - - case "ollama": { - // SECURITY: Mirror ai-providers.ts guard — only use server - // OLLAMA_API_KEY when the URL is also from server config. - const ollamaApiKey = baseUrl - ? apiKey || undefined - : apiKey || process.env.OLLAMA_API_KEY || undefined - const ollamaProvider = createOllama({ - baseURL: - baseUrl || - process.env.OLLAMA_BASE_URL || - "https://ollama.com/api", - fetch: guardedFetch, - ...(ollamaApiKey && { - headers: { Authorization: `Bearer ${ollamaApiKey}` }, - }), - }) - model = ollamaProvider(modelId) - break - } - - case "gateway": { - const gw = createGateway({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - fetch: guardedFetch, - }) - model = gw(modelId) - break - } - - case "edgeone": { - // EdgeOne uses OpenAI-compatible API via Edge Functions - // Need to pass cookies for EdgeOne Pages authentication, - // and the access code, which the edge function also checks - const cookieHeader = req.headers.get("cookie") || "" - const edgeone = createOpenAI({ - apiKey: "edgeone", // EdgeOne doesn't require API key - baseURL: baseUrl || "/api/edgeai", - fetch: guardedFetch, - headers: { - cookie: cookieHeader, - "x-access-code": req.headers.get("x-access-code") || "", - }, - }) - model = edgeone.chat(modelId) - break - } - - case "sglang": { - // SGLang is OpenAI-compatible - const sglang = createOpenAI({ - apiKey: apiKey || "not-needed", - baseURL: baseUrl || "http://127.0.0.1:8000/v1", - fetch: guardedFetch, - }) - model = sglang.chat(modelId) - break - } - - case "doubao": { - // ByteDance Doubao: use DeepSeek for DeepSeek/Kimi models, OpenAI for others - const doubaoBaseUrl = - baseUrl || "https://ark.cn-beijing.volces.com/api/v3" - const lowerModelId = modelId.toLowerCase() - if ( - lowerModelId.includes("deepseek") || - lowerModelId.includes("kimi") - ) { - const doubao = createDeepSeek({ - apiKey, - baseURL: doubaoBaseUrl, - fetch: guardedFetch, - }) - model = doubao(modelId) - } else { - const doubao = createOpenAI({ - apiKey, - baseURL: doubaoBaseUrl, - fetch: guardedFetch, - }) - model = doubao.chat(modelId) - } - break - } - - case "modelscope": { - const baseURL = - baseUrl || "https://api-inference.modelscope.cn/v1" - const startTime = Date.now() - - try { - // Initiate a streaming request (required for QwQ-32B and certain Qwen3 models) - const response = await (guardedFetch ?? fetch)( - `${baseURL}/chat/completions`, - { - method: "POST", - headers: { - "Content-Type": "application/json", - Authorization: `Bearer ${apiKey}`, - }, - body: JSON.stringify({ - model: modelId, - messages: [ - { role: "user", content: "Say 'OK'" }, - ], - max_tokens: 20, - stream: true, - enable_thinking: false, - }), - }, - ) - - if (!response.ok) { - // Log the body but return only the status: the - // caller chooses baseUrl, so the body may come from - // any host the server can reach - console.error( - "[validate-model] ModelScope error body:", - await response.text(), - ) - throw new Error( - `ModelScope API error (${response.status})`, - ) - } - - const contentType = - response.headers.get("content-type") || "" - const isValidStreamingResponse = - response.status === 200 && - (contentType.includes("text/event-stream") || - contentType.includes("application/json")) - - if (!isValidStreamingResponse) { - throw new Error( - `Unexpected response format: ${contentType}`, - ) - } - - const responseTime = Date.now() - startTime - - if (response.body) { - response.body.cancel().catch(() => { - /* Ignore cancellation errors */ - }) - } - - return NextResponse.json({ - valid: true, - responseTime, - note: "ModelScope model validated (using streaming API)", - }) - } catch (error) { - console.error( - "[validate-model] ModelScope validation failed:", - error, - ) - throw error - } - } - - case "minimax": { - const rawUrl = - baseUrl || - PROVIDER_INFO.minimax?.defaultBaseUrl || - "https://api.minimaxi.com/anthropic" - const { baseURL: minimaxBaseUrl, isAnthropicCompatible } = - normalizeMiniMaxBaseURL(rawUrl) - - if (isAnthropicCompatible) { - const minimax = createAnthropic({ - apiKey, - baseURL: minimaxBaseUrl, - fetch: guardedFetch, - }) - model = minimax.chat(modelId) - } else { - const minimax = createOpenAI({ - apiKey, - baseURL: minimaxBaseUrl, - fetch: guardedFetch, - }) - model = minimax.chat(modelId) - } - break - } - - // GLM, Qwen, Kimi, Qiniu, Novita, MiMo, Atlas Cloud - OpenAI compatible - case "glm": - case "qwen": - case "kimi": - case "qiniu": - case "novita": - case "atlascloud": - case "mimo": { - const baseURL = - baseUrl || - PROVIDER_INFO[provider as ProviderName]?.defaultBaseUrl || - "" - - if (!baseURL) { - return NextResponse.json( - { - valid: false, - error: `No base URL configured for provider: ${provider}`, - }, - { status: 400 }, - ) - } - - const openai = createOpenAI({ - apiKey, - baseURL, - fetch: guardedFetch, - }) - model = openai.chat(modelId) - break - } - - default: - return NextResponse.json( - { valid: false, error: `Unknown provider: ${provider}` }, - { status: 400 }, - ) - } - - // Make a minimal test request - const startTime = Date.now() - await generateText({ - model, - prompt: "Say 'OK'", - maxOutputTokens: 20, + // The same model the chat would use. A client base URL makes it + // refuse redirects to internal hosts. + const { model } = getAIModel({ + provider, + modelId, + apiKey, + baseUrl, + awsAccessKeyId, + awsSecretAccessKey, + awsRegion, + vertexApiKey, + // EdgeOne checks the Pages cookies and the access code + ...(provider === "edgeone" && { + headers: { + cookie: req.headers.get("cookie") || "", + "x-access-code": req.headers.get("x-access-code") || "", + }, + }), }) + + // Streaming, like the chat (some models only stream). Stop at the + // first tool call; a reasoning model that runs out of tokens first + // proves the connection but not tool support. + const startTime = Date.now() + const result = streamText({ + model, + prompt: "Call the ping tool.", + tools: { ping: PING_TOOL }, + maxOutputTokens: 1024, + maxRetries: 0, + abortSignal: AbortSignal.timeout(TEST_TIMEOUT_MS), + }) + let calledTool = false + let finishReason: string | undefined + for await (const part of result.fullStream) { + if (part.type === "error") throw part.error + if (part.type === "tool-call") { + calledTool = true + break + } + if (part.type === "finish") finishReason = part.finishReason + } const responseTime = Date.now() - startTime return NextResponse.json({ valid: true, responseTime, + ...(!calledTool && + finishReason !== "length" && { warning: NO_TOOL_CALL_WARNING }), }) } catch (error) { console.error("[validate-model] Error:", error) @@ -473,7 +149,9 @@ export async function POST(req: Request) { let errorMessage = "Validation failed" if (error instanceof Error) { // Extract meaningful error message - if ( + if (error.name === "TimeoutError") { + errorMessage = `No answer within ${TEST_TIMEOUT_MS / 1000} seconds` + } else if ( error.message.includes("401") || error.message.includes("Unauthorized") ) { diff --git a/components/model-config-dialog.tsx b/components/model-config-dialog.tsx index d908f04a..b1c2a35a 100644 --- a/components/model-config-dialog.tsx +++ b/components/model-config-dialog.tsx @@ -4,7 +4,6 @@ import { AlertCircle, Check, ChevronRight, - Clock, Eye, EyeOff, Key, @@ -57,7 +56,11 @@ import type { UseModelConfigReturn } from "@/hooks/use-model-config" import { getApiEndpoint } from "@/lib/base-path" import { formatMessage } from "@/lib/i18n/utils" import { STORAGE_KEYS } from "@/lib/storage" -import type { ProviderConfig, ProviderName } from "@/lib/types/model-config" +import type { + ModelConfig, + ProviderConfig, + ProviderName, +} from "@/lib/types/model-config" import { PROVIDER_INFO, SUGGESTED_MODELS } from "@/lib/types/model-config" import { cn } from "@/lib/utils" @@ -126,9 +129,10 @@ export function ModelConfigDialog({ > | null>(null) const [deleteConfirmOpen, setDeleteConfirmOpen] = useState(false) const [deleteConfirmText, setDeleteConfirmText] = useState("") - const [validatingModelIndex, setValidatingModelIndex] = useState< - number | null - >(null) + // Models whose test is running (they are all tested at once) + const [validatingModelIds, setValidatingModelIds] = useState>( + () => new Set(), + ) const [duplicateError, setDuplicateError] = useState("") const [editError, setEditError] = useState<{ modelId: string @@ -281,12 +285,14 @@ export function ModelConfigDialog({ if (credentialFields.includes(field)) { credentialsVersionRef.current++ setValidationStatus("idle") - setValidatingModelIndex(null) + setValidatingModelIds(new Set()) updates.validated = false updates.models = selectedProvider.models.map((m) => ({ ...m, validated: undefined, validationError: undefined, + validationWarning: undefined, + responseTime: undefined, })) } updateProvider(selectedProviderId, updates) @@ -361,75 +367,82 @@ export function ModelConfigDialog({ let errorCount = 0 const credentialsVersion = credentialsVersionRef.current - // Validate each model - for (let i = 0; i < selectedProvider.models.length; i++) { - const model = selectedProvider.models[i] - setValidatingModelIndex(i) + // For EdgeOne, construct baseUrl from current origin + const baseUrl = isEdgeOne + ? `${window.location.origin}${getApiEndpoint("/api/edgeai")}` + : selectedProvider.baseUrl - try { - // For EdgeOne, construct baseUrl from current origin - const baseUrl = isEdgeOne - ? `${window.location.origin}${getApiEndpoint("/api/edgeai")}` - : selectedProvider.baseUrl - - const response = await fetch( - getApiEndpoint("/api/validate-model"), - { - method: "POST", - headers: { - "Content-Type": "application/json", - "x-access-code": - localStorage.getItem(STORAGE_KEYS.accessCode) || - "", + // Test every model at once; each row updates when its answer arrives + setValidatingModelIds(new Set(selectedProvider.models.map((m) => m.id))) + await Promise.all( + selectedProvider.models.map(async (model) => { + let update: Partial + try { + const response = await fetch( + getApiEndpoint("/api/validate-model"), + { + method: "POST", + headers: { + "Content-Type": "application/json", + "x-access-code": + localStorage.getItem( + STORAGE_KEYS.accessCode, + ) || "", + }, + body: JSON.stringify({ + provider: selectedProvider.provider, + apiKey: selectedProvider.apiKey, + baseUrl, + modelId: model.modelId, + // AWS Bedrock credentials + awsAccessKeyId: selectedProvider.awsAccessKeyId, + awsSecretAccessKey: + selectedProvider.awsSecretAccessKey, + awsRegion: selectedProvider.awsRegion, + // Vertex AI credentials (Express Mode) + vertexApiKey: selectedProvider.vertexApiKey, + }), }, - body: JSON.stringify({ - provider: selectedProvider.provider, - apiKey: selectedProvider.apiKey, - baseUrl, - modelId: model.modelId, - // AWS Bedrock credentials - awsAccessKeyId: selectedProvider.awsAccessKeyId, - awsSecretAccessKey: - selectedProvider.awsSecretAccessKey, - awsRegion: selectedProvider.awsRegion, - // Vertex AI credentials (Express Mode) - vertexApiKey: selectedProvider.vertexApiKey, - }), - }, - ) - const data = await response.json().catch(() => ({})) - // Credentials changed during the test: drop the results + ) + const data = await response.json().catch(() => ({})) + update = data.valid + ? { + validated: true, + validationError: undefined, + validationWarning: data.warning, + responseTime: data.responseTime, + } + : { + validated: false, + validationError: + data.error || + (response.ok + ? "Validation failed" + : `Request failed (${response.status})`), + validationWarning: undefined, + } + } catch { + update = { + validated: false, + validationError: "Network error", + validationWarning: undefined, + } + } + // Credentials changed during the test: drop the result if (credentialsVersionRef.current !== credentialsVersion) return - - if (data.valid) { - updateModel(selectedProviderId, model.id, { - validated: true, - validationError: undefined, - }) - } else { + if (update.validated === false) { allValid = false errorCount++ - updateModel(selectedProviderId, model.id, { - validated: false, - validationError: - data.error || - (response.ok - ? "Validation failed" - : `Request failed (${response.status})`), - }) } - } catch { - if (credentialsVersionRef.current !== credentialsVersion) return - allValid = false - errorCount++ - updateModel(selectedProviderId, model.id, { - validated: false, - validationError: "Network error", + updateModel(selectedProviderId, model.id, update) + setValidatingModelIds((prev) => { + const next = new Set(prev) + next.delete(model.id) + return next }) - } - } - - setValidatingModelIndex(null) + }), + ) + if (credentialsVersionRef.current !== credentialsVersion) return if (allValid) { setValidationStatus("success") @@ -992,28 +1005,24 @@ export function ModelConfigDialog({
{/* Status icon */}
- {validatingModelIndex !== - null && - index === - validatingModelIndex ? ( + {validatingModelIds.has( + model.id, + ) ? ( // Currently validating
- ) : validatingModelIndex !== - null && - index > - validatingModelIndex && - model.validated === - undefined ? ( - // Queued -
- -
) : model.validated === true ? ( - // Valid -
+ // Valid, with the time the test took +
) : model.validated === @@ -1219,6 +1228,14 @@ export function ModelConfigDialog({ }

)} + {model.validated && + model.validationWarning && ( +

+ { + model.validationWarning + } +

+ )} {/* Show edit error inline */} {editError?.modelId === model.id && ( diff --git a/lib/ai-providers.ts b/lib/ai-providers.ts index bf8f71fe..b74bdc44 100644 --- a/lib/ai-providers.ts +++ b/lib/ai-providers.ts @@ -1,24 +1,27 @@ import { createAmazonBedrock } from "@ai-sdk/amazon-bedrock" import { createAnthropic } from "@ai-sdk/anthropic" -import { azure, createAzure } from "@ai-sdk/azure" -import { createDeepSeek, deepseek } from "@ai-sdk/deepseek" -import { createGoogleGenerativeAI, google } from "@ai-sdk/google" +import { createAzure } from "@ai-sdk/azure" +import { createDeepSeek } from "@ai-sdk/deepseek" +import { createGoogleGenerativeAI } from "@ai-sdk/google" import { createVertex } from "@ai-sdk/google-vertex" -import { createOpenAI, openai } from "@ai-sdk/openai" -import { aihubmix, createAihubmix } from "@aihubmix/ai-sdk-provider" +import { createOpenAI } from "@ai-sdk/openai" +import { createOpenAICompatible } from "@ai-sdk/openai-compatible" +import { createAihubmix } from "@aihubmix/ai-sdk-provider" import { fromNodeProviderChain } from "@aws-sdk/credential-providers" import { createOpenRouter } from "@openrouter/ai-sdk-provider" import { createGateway, defaultSettingsMiddleware, - gateway, + extractReasoningMiddleware, + type LanguageModel, wrapLanguageModel, } from "ai" -import { createOllama, ollama } from "ollama-ai-provider-v2" +import { createOllama } from "ollama-ai-provider-v2" import { adminProvidersToConfig, loadAdminProviders, } from "@/lib/admin/providers" +import { redirectGuardedFetch } from "@/lib/ssrf-protection" import { PROVIDER_INFO, type ProviderName } from "@/lib/types/model-config" export type { ProviderName } @@ -101,34 +104,6 @@ export interface ClientOverrides { baseUrlEnv?: string } -// Providers that can be selected from client settings -const ALLOWED_CLIENT_PROVIDERS: ProviderName[] = [ - "openai", - "anthropic", - "google", - "vertexai", - "azure", - "bedrock", - "openrouter", - "aihubmix", - "deepseek", - "siliconflow", - "sglang", - "gateway", - "edgeone", - "ollama", - "doubao", - "modelscope", - "glm", - "qwen", - "qiniu", - "kimi", - "minimax", - "novita", - "mimo", - "atlascloud", -] - // Bedrock provider options for Anthropic beta features const BEDROCK_ANTHROPIC_BETA = { bedrock: { @@ -511,27 +486,6 @@ function buildProviderOptions( break } - case "deepseek": - case "openrouter": - case "aihubmix": - case "siliconflow": - case "sglang": - case "gateway": - case "modelscope": - case "doubao": - case "minimax": - case "glm": - case "qwen": - case "kimi": - case "qiniu": - case "novita": - case "atlascloud": - case "mimo": { - // These providers don't have reasoning configs in AI SDK yet - // Gateway passes through to underlying providers which handle their own configs - break - } - default: break } @@ -665,30 +619,152 @@ function validateProviderCredentials( } /** - * Get the AI model based on environment variables - * - * Environment variables: - * - AI_PROVIDER: The provider to use (bedrock, openai, anthropic, google, azure, ollama, openrouter, aihubmix, deepseek, siliconflow, sglang, gateway, modelscope) - * - AI_MODEL: The model ID/name for the selected provider - * - * Provider-specific env vars: - * - OPENAI_API_KEY: OpenAI API key - * - OPENAI_BASE_URL: Custom OpenAI-compatible endpoint (optional) - * - ANTHROPIC_API_KEY: Anthropic API key - * - GOOGLE_GENERATIVE_AI_API_KEY: Google API key - * - AZURE_RESOURCE_NAME, AZURE_API_KEY: Azure OpenAI credentials - * - AWS_REGION, AWS_ACCESS_KEY_ID, AWS_SECRET_ACCESS_KEY: AWS Bedrock credentials - * - OLLAMA_BASE_URL: Ollama server URL (optional, defaults to https://ollama.com/api) - * - OPENROUTER_API_KEY: OpenRouter API key - * - AIHUBMIX_API_KEY: AIHubMix API key - * - DEEPSEEK_API_KEY: DeepSeek API key - * - DEEPSEEK_BASE_URL: DeepSeek endpoint (optional) - * - SILICONFLOW_API_KEY: SiliconFlow API key - * - SILICONFLOW_BASE_URL: SiliconFlow endpoint (optional, defaults to https://api.siliconflow.cn/v1) - * - SGLANG_API_KEY: SGLang API key - * - SGLANG_BASE_URL: SGLang endpoint (optional) - * - MODELSCOPE_API_KEY: ModelScope API key - * - MODELSCOPE_BASE_URL: ModelScope endpoint (optional) + * Providers whose SDK has the official endpoint built in. The others are + * OpenAI-compatible APIs (or Anthropic) that are called at + * PROVIDER_INFO.defaultBaseUrl unless a base URL is configured. + */ +const SDK_KNOWS_ENDPOINT = new Set([ + "openai", + "google", + "azure", + "openrouter", + "gateway", + "deepseek", +]) + +/** Where and how to call a provider, once credentials are resolved */ +interface Endpoint { + apiKey?: string + baseURL?: string + headers?: Record + fetch?: typeof fetch + authToken?: string // Anthropic Bearer auth + resourceName?: string // Azure +} + +/** + * An OpenAI-compatible chat model. includeUsage asks for token usage in the + * stream, which quota tracking needs. Some of these models write their + * reasoning inside tags; that text becomes reasoning, not reply. + */ +function compatibleModel( + provider: ProviderName, + modelId: string, + e: Endpoint, +): LanguageModel { + const model = createOpenAICompatible({ + name: provider, + apiKey: e.apiKey, + baseURL: e.baseURL ?? "", + ...(e.headers && { headers: e.headers }), + ...(e.fetch && { fetch: e.fetch }), + includeUsage: true, + })(modelId) + return wrapLanguageModel({ + model, + middleware: extractReasoningMiddleware({ tagName: "think" }), + }) +} + +/** Create the model for a provider. Credentials are already resolved. */ +function createModel( + provider: ProviderName, + modelId: string, + e: Endpoint, +): LanguageModel { + const opts = { + apiKey: e.apiKey, + ...(e.baseURL && { baseURL: e.baseURL }), + ...(e.fetch && { fetch: e.fetch }), + } + switch (provider) { + case "openai": { + const openaiProvider = createOpenAI(opts) + // A custom base URL is usually a proxy that only has Chat + // Completions; the official endpoint uses the Responses API, + // which returns reasoning for gpt-5 and the o-series + return e.baseURL + ? openaiProvider.chat(modelId) + : openaiProvider(modelId) + } + case "anthropic": + // The provider streams tool input per tool (eager_input_streaming), + // which replaced the fine-grained-tool-streaming beta header + return createAnthropic({ + ...(e.authToken + ? { authToken: e.authToken } + : { apiKey: e.apiKey }), + baseURL: e.baseURL, + ...(e.fetch && { fetch: e.fetch }), + })(modelId) + case "google": { + const model = createGoogleGenerativeAI(opts)(modelId) + const sampling = googleSamplingSettings() + return Object.keys(sampling).length > 0 + ? wrapLanguageModel({ + model, + middleware: defaultSettingsMiddleware({ + settings: sampling, + }), + }) + : model + } + case "azure": + // baseURL takes precedence over resourceName per SDK behavior + return createAzure({ + ...opts, + ...(!e.baseURL && + e.resourceName && { resourceName: e.resourceName }), + })(modelId) + case "openrouter": + return createOpenRouter(opts)(modelId) + case "gateway": + // Without a key or URL the SDK uses Vercel's endpoint and OIDC + return createGateway(opts)(modelId) + case "deepseek": + case "kimi": + case "mimo": + // Kimi and MiMo return reasoning_content like DeepSeek and need it + // passed back in multi-turn tool calls (MiMo answers 400 otherwise) + return createDeepSeek(opts)(modelId) + case "doubao": { + // DeepSeek and Kimi models on Doubao use reasoning_content too + const lower = modelId.toLowerCase() + return lower.includes("deepseek") || lower.includes("kimi") + ? createDeepSeek(opts)(modelId) + : compatibleModel(provider, modelId, e) + } + case "aihubmix": + return isAihubmixStandardBaseURL(e.baseURL) + ? createAihubmix({ + apiKey: e.apiKey, + appCode: AIHUBMIX_APP_CODE, + })(modelId) + : compatibleModel(provider, modelId, e) + case "minimax": { + const { baseURL, isAnthropicCompatible } = normalizeMiniMaxBaseURL( + e.baseURL as string, + ) + return isAnthropicCompatible + ? createAnthropic({ + apiKey: e.apiKey, + baseURL, + ...(e.fetch && { fetch: e.fetch }), + })(modelId) + : compatibleModel(provider, modelId, { ...e, baseURL }) + } + default: + // siliconflow, sglang, modelscope, glm, qwen, qiniu, novita, + // atlascloud, edgeone + return compatibleModel(provider, modelId, e) + } +} + +/** + * Get the AI model for a chat request: the client's own provider and + * credentials, or the server's (AI_PROVIDER, AI_MODEL and each provider's + * _API_KEY / _BASE_URL, see env.example). The settings test + * button uses the same function, so a passing test means the chat works. */ export function getAIModel(overrides?: ClientOverrides): ModelConfig { // SECURITY: Prevent SSRF attacks (GHSA-9qf7-mprq-9qgm) @@ -737,13 +813,9 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig { let provider: ProviderName if (overrides?.provider) { // Validate client-provided provider - if ( - !ALLOWED_CLIENT_PROVIDERS.includes( - overrides.provider as ProviderName, - ) - ) { + if (!Object.hasOwn(PROVIDER_INFO, overrides.provider)) { throw new Error( - `Invalid provider: ${overrides.provider}. Allowed providers: ${ALLOWED_CLIENT_PROVIDERS.join(", ")}`, + `Invalid provider: ${overrides.provider}. Allowed providers: ${Object.keys(PROVIDER_INFO).join(", ")}`, ) } provider = overrides.provider as ProviderName @@ -761,30 +833,27 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig { .map(([p]) => p) if (configured.length === 0) { + const keys = Object.entries(PROVIDER_ENV_VARS) + .filter(([, envVar]) => envVar) + .map(([p, envVar]) => `- ${envVar} for ${p}`) throw new Error( `No AI provider configured. Please set one of the following API keys in your .env.local file:\n` + - `- AI_GATEWAY_API_KEY for Vercel AI Gateway\n` + - `- DEEPSEEK_API_KEY for DeepSeek\n` + - `- OPENAI_API_KEY for OpenAI\n` + - `- ANTHROPIC_API_KEY for Anthropic\n` + - `- GOOGLE_GENERATIVE_AI_API_KEY for Google\n` + - `- AWS_ACCESS_KEY_ID for Bedrock\n` + - `- OPENROUTER_API_KEY for OpenRouter\n` + - `- AIHUBMIX_API_KEY for AIHubMix\n` + - `- AZURE_API_KEY for Azure\n` + - `- SILICONFLOW_API_KEY for SiliconFlow\n` + - `- SGLANG_API_KEY for SGLang\n` + - `- MODELSCOPE_API_KEY for ModelScope\n` + + `${keys.join("\n")}\n` + + `- AWS_ACCESS_KEY_ID for bedrock\n` + `Or set AI_PROVIDER=ollama for local Ollama.`, ) - } else { - throw new Error( - `Multiple AI providers configured (${configured.join(", ")}). ` + - `Please set AI_PROVIDER to specify which one to use.`, - ) } + throw new Error( + `Multiple AI providers configured (${configured.join(", ")}). ` + + `Please set AI_PROVIDER to specify which one to use.`, + ) } } + if (!Object.hasOwn(PROVIDER_INFO, provider)) { + throw new Error( + `Unknown AI provider: ${provider}. Supported providers: ${Object.keys(PROVIDER_INFO).join(", ")}`, + ) + } // Only validate server credentials if client isn't providing their own API key if (!isClientOverride) { @@ -793,11 +862,11 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig { console.log(`[AI Provider] Initializing ${provider} with model: ${modelId}`) - let model: any - let providerOptions: any - + // Requests to a base URL the client chose must not follow redirects + const guardedFetch = overrides?.baseUrl ? redirectGuardedFetch() : undefined // Build provider-specific options from environment variables - const customProviderOptions = buildProviderOptions(provider, modelId) + let providerOptions = buildProviderOptions(provider, modelId) + let model: LanguageModel switch (provider) { case "bedrock": { @@ -841,109 +910,13 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig { providerOptions = { bedrock: { ...BEDROCK_ANTHROPIC_BETA.bedrock, - ...(customProviderOptions?.bedrock || {}), + ...(providerOptions?.bedrock || {}), }, } - } else if (customProviderOptions) { - providerOptions = customProviderOptions } break } - case "openai": { - const apiKey = resolveApiKey(overrides, "OPENAI_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "OPENAI_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - ) - if (baseURL) { - // Custom base URL = third-party proxy, use Chat Completions API - // for compatibility (most proxies don't support /responses endpoint) - const customOpenAI = createOpenAI({ apiKey, baseURL }) - model = customOpenAI.chat(modelId) - } else if (overrides?.apiKey || overrides?.apiKeyEnv) { - // Custom API key (the client's, or a server model's own env var) - // but official OpenAI endpoint, use Responses API - // to support reasoning for gpt-5, o1, o3, o4 models - const customOpenAI = createOpenAI({ apiKey }) - model = customOpenAI(modelId) - } else { - model = openai(modelId) - } - break - } - - case "anthropic": { - const apiKey = resolveApiKey(overrides, "ANTHROPIC_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "ANTHROPIC_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - "https://api.anthropic.com/v1", - ) - // Anthropic supports two auth methods (mutually exclusive): - // - apiKey: sends as `x-api-key` header - // - authToken: sends as `Authorization: Bearer ` header - // Prefer apiKey if present (including client overrides); fall back - // to ANTHROPIC_AUTH_TOKEN env var only when no apiKey is available. - const authToken = !apiKey - ? process.env.ANTHROPIC_AUTH_TOKEN - : undefined - // The provider streams tool input per tool (eager_input_streaming), - // which replaced the fine-grained-tool-streaming beta header - const customProvider = createAnthropic({ - ...(authToken ? { authToken } : { apiKey }), - baseURL, - }) - model = customProvider(modelId) - break - } - - case "google": { - const apiKey = resolveApiKey( - overrides, - "GOOGLE_GENERATIVE_AI_API_KEY", - ) - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "GOOGLE_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - ) - // The default instance only reads GOOGLE_GENERATIVE_AI_API_KEY, so a - // server model's own env var (apiKeyEnv) needs a custom instance too - if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) { - const customGoogle = createGoogleGenerativeAI({ - apiKey, - ...(baseURL && { baseURL }), - }) - model = customGoogle(modelId) - } else { - model = google(modelId) - } - const sampling = googleSamplingSettings() - if (Object.keys(sampling).length > 0) { - model = wrapLanguageModel({ - model, - middleware: defaultSettingsMiddleware({ - settings: sampling, - }), - }) - } - break - } case "vertexai": { // Express Mode: Use API key for authentication // SECURITY: a client base URL only ever gets the client's key, so the @@ -966,40 +939,11 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig { overrides?.baseUrl, process.env.GOOGLE_VERTEX_BASE_URL, ) - - const vertexProvider = createVertex({ + model = createVertex({ apiKey: vertexApiKey, ...(baseURL && { baseURL }), - }) - model = vertexProvider(modelId) - break - } - - case "azure": { - const apiKey = resolveApiKey(overrides, "AZURE_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv(overrides, "AZURE_BASE_URL") - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - ) - // Only use server's resourceName if user is NOT providing their own API key - const resourceName = overrides?.apiKey - ? undefined - : process.env.AZURE_RESOURCE_NAME - // Azure requires either baseURL or resourceName to construct the endpoint - // resourceName constructs: https://{resourceName}.openai.azure.com/openai/v1{path} - if (baseURL || resourceName || overrides?.apiKey) { - const customAzure = createAzure({ - apiKey, - // baseURL takes precedence over resourceName per SDK behavior - ...(baseURL && { baseURL }), - ...(!baseURL && resourceName && { resourceName }), - }) - model = customAzure(modelId) - } else { - model = azure(modelId) - } + ...(guardedFetch && { fetch: guardedFetch }), + })(modelId) break } @@ -1011,432 +955,63 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig { const apiKey = overrides?.baseUrl ? overrides?.apiKey || undefined : resolveApiKey(overrides, "OLLAMA_API_KEY") - if (baseURL || apiKey) { - const customOllama = createOllama({ - ...(baseURL && { baseURL }), - ...(apiKey && { - headers: { Authorization: `Bearer ${apiKey}` }, - }), - }) - model = customOllama(modelId) - } else { - model = ollama(modelId) - } - break - } - - case "openrouter": { - const apiKey = resolveApiKey(overrides, "OPENROUTER_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "OPENROUTER_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - ) - const openrouter = createOpenRouter({ - apiKey, + model = createOllama({ ...(baseURL && { baseURL }), + ...(apiKey && { + headers: { Authorization: `Bearer ${apiKey}` }, + }), + ...(guardedFetch && { fetch: guardedFetch }), + })(modelId) + break + } + + case "edgeone": + // EdgeOne Pages Edge AI, an OpenAI-compatible API without a key. + // The SDK appends /chat/completions to the base URL. Cookies + // (eo_token, eo_time) and the access code authenticate the call. + model = compatibleModel(provider, modelId, { + apiKey: "edgeone", + baseURL: overrides?.baseUrl || "/api/edgeai", + headers: overrides?.headers, + fetch: guardedFetch, }) - model = openrouter(modelId) break - } - case "aihubmix": { - const apiKey = resolveApiKey(overrides, "AIHUBMIX_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( + default: { + // Every other provider takes an API key and a base URL from + // _API_KEY / _BASE_URL (or a server model's apiKeyEnv) + const apiKey = resolveApiKey( overrides, - "AIHUBMIX_BASE_URL", + PROVIDER_ENV_VARS[provider] as string, ) + const baseUrlEnv = + provider === "gateway" + ? "AI_GATEWAY_BASE_URL" + : `${provider.toUpperCase()}_BASE_URL` const baseURL = resolveBaseURL( overrides?.apiKey, overrides?.baseUrl, - serverBaseUrl, - PROVIDER_INFO.aihubmix.defaultBaseUrl, + resolveBaseUrlEnv(overrides, baseUrlEnv), + SDK_KNOWS_ENDPOINT.has(provider) + ? undefined + : PROVIDER_INFO[provider].defaultBaseUrl, ) - const defaultBaseURL = PROVIDER_INFO.aihubmix.defaultBaseUrl - - if ( - isAihubmixStandardBaseURL(baseURL) || - baseURL === defaultBaseURL - ) { - const aihubmixProvider = - overrides?.apiKey || apiKey - ? createAihubmix({ - apiKey, - appCode: AIHUBMIX_APP_CODE, - }) - : aihubmix - model = aihubmixProvider(modelId) - } else { - const aihubmixCompatibleProvider = createOpenAI({ - apiKey, - baseURL, - }) - model = aihubmixCompatibleProvider.chat(modelId) - } - break - } - - case "deepseek": { - const apiKey = resolveApiKey(overrides, "DEEPSEEK_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "DEEPSEEK_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - ) - if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) { - const customDeepSeek = createDeepSeek({ - apiKey, - ...(baseURL && { baseURL }), - }) - model = customDeepSeek(modelId) - } else { - model = deepseek(modelId) - } - break - } - - case "siliconflow": { - const apiKey = resolveApiKey(overrides, "SILICONFLOW_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "SILICONFLOW_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - "https://api.siliconflow.cn/v1", - ) - const siliconflowProvider = createOpenAI({ + model = createModel(provider, modelId, { apiKey, baseURL, + fetch: guardedFetch, + // Bearer auth for Anthropic when there is no API key + authToken: + provider === "anthropic" && !apiKey + ? process.env.ANTHROPIC_AUTH_TOKEN + : undefined, + // Only the server's own resource; a client key needs its URL + resourceName: + provider === "azure" && !overrides?.apiKey + ? process.env.AZURE_RESOURCE_NAME + : undefined, }) - model = siliconflowProvider.chat(modelId) - break } - - case "sglang": { - const apiKey = resolveApiKey(overrides, "SGLANG_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "SGLANG_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - ) - - const sglangProvider = createOpenAI({ - apiKey, - ...(baseURL && { baseURL }), - // Add a custom fetch wrapper to intercept and fix the stream from sglang - fetch: async (url, options) => { - const response = await fetch(url, options) - if (!response.body) { - return response - } - - // Create a transform stream to fix the non-compliant sglang stream - let buffer = "" - const decoder = new TextDecoder() - - const transformStream = new TransformStream({ - transform(chunk, controller) { - buffer += decoder.decode(chunk, { stream: true }) - // Process all complete messages in the buffer - let messageEndPos - while ( - (messageEndPos = buffer.indexOf("\n\n")) !== -1 - ) { - const message = buffer.substring( - 0, - messageEndPos, - ) - buffer = buffer.substring(messageEndPos + 2) // Move past the '\n\n' - - if (message.startsWith("data: ")) { - const jsonStr = message.substring(6).trim() - if (jsonStr === "[DONE]") { - controller.enqueue( - new TextEncoder().encode( - message + "\n\n", - ), - ) - continue - } - try { - const data = JSON.parse(jsonStr) - const delta = data.choices?.[0]?.delta - - if (delta) { - // Fix 1: remove invalid empty role - if (delta.role === "") { - delete delta.role - } - // Fix 2: remove non-standard reasoning_content field - if ("reasoning_content" in delta) { - delete delta.reasoning_content - } - } - - // Re-serialize and forward the corrected data with the correct SSE format - controller.enqueue( - new TextEncoder().encode( - `data: ${JSON.stringify(data)}\n\n`, - ), - ) - } catch (_e) { - // If parsing fails, forward the original message to avoid breaking the stream. - controller.enqueue( - new TextEncoder().encode( - message + "\n\n", - ), - ) - } - } else if (message.trim() !== "") { - // Pass through other message types (e.g., 'event: ...') - controller.enqueue( - new TextEncoder().encode( - message + "\n\n", - ), - ) - } - } - }, - flush(controller) { - // If there's anything left in the buffer, forward it. - if (buffer.trim()) { - controller.enqueue( - new TextEncoder().encode(buffer), - ) - } - }, - }) - - const transformedBody = - response.body.pipeThrough(transformStream) - - // Return a new response with the transformed body - return new Response(transformedBody, { - status: response.status, - statusText: response.statusText, - headers: response.headers, - }) - }, - }) - model = sglangProvider.chat(modelId) - break - } - - case "gateway": { - // Vercel AI Gateway - unified access to multiple AI providers - // Model format: "provider/model" e.g., "openai/gpt-4o", "anthropic/claude-sonnet-4-5" - // See: https://vercel.com/ai-gateway - const apiKey = resolveApiKey(overrides, "AI_GATEWAY_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "AI_GATEWAY_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - ) - // Only use custom configuration if explicitly set (local dev or custom Gateway) - // Otherwise undefined → AI SDK uses Vercel default (https://ai-gateway.vercel.sh/v1/ai) + OIDC - if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) { - const customGateway = createGateway({ - apiKey, - ...(baseURL && { baseURL }), - }) - model = customGateway(modelId) - } else { - model = gateway(modelId) - } - break - } - - case "edgeone": { - // EdgeOne Pages Edge AI - uses OpenAI-compatible API - // AI SDK appends /chat/completions to baseURL - // /api/edgeai + /chat/completions = /api/edgeai/chat/completions - const baseURL = overrides?.baseUrl || "/api/edgeai" - const edgeoneProvider = createOpenAI({ - apiKey: "edgeone", // Dummy key - EdgeOne doesn't require API key - baseURL, - // Pass cookies for EdgeOne Pages authentication (eo_token, eo_time) - ...(overrides?.headers && { headers: overrides.headers }), - }) - model = edgeoneProvider.chat(modelId) - break - } - - case "doubao": { - const apiKey = resolveApiKey(overrides, "DOUBAO_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "DOUBAO_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - "https://ark.cn-beijing.volces.com/api/v3", - ) - const lowerModelId = modelId.toLowerCase() - // Use DeepSeek provider for DeepSeek/Kimi models, OpenAI for others (multimodal support) - if ( - lowerModelId.includes("deepseek") || - lowerModelId.includes("kimi") - ) { - const doubaoProvider = createDeepSeek({ - apiKey, - baseURL, - }) - model = doubaoProvider(modelId) - } else { - const doubaoProvider = createOpenAI({ - apiKey, - baseURL, - }) - model = doubaoProvider.chat(modelId) - } - break - } - - case "modelscope": { - const apiKey = resolveApiKey(overrides, "MODELSCOPE_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "MODELSCOPE_BASE_URL", - ) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - "https://api-inference.modelscope.cn/v1", - ) - const modelscopeProvider = createOpenAI({ - apiKey, - baseURL, - }) - model = modelscopeProvider.chat(modelId) - break - } - - case "minimax": { - const apiKey = resolveApiKey(overrides, "MINIMAX_API_KEY") - const serverBaseUrl = resolveBaseUrlEnv( - overrides, - "MINIMAX_BASE_URL", - ) - const rawBaseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - serverBaseUrl, - PROVIDER_INFO.minimax.defaultBaseUrl, - ) - - if (!rawBaseURL) { - throw new Error( - "MiniMax base URL could not be resolved. Set MINIMAX_BASE_URL or configure a base URL in settings.", - ) - } - - const { baseURL, isAnthropicCompatible } = - normalizeMiniMaxBaseURL(rawBaseURL) - - if (isAnthropicCompatible) { - const minimax = createAnthropic({ apiKey, baseURL }) - model = minimax.chat(modelId) - } else { - const minimax = createOpenAI({ apiKey, baseURL }) - model = minimax.chat(modelId) - } - break - } - - case "mimo": { - const apiKey = resolveApiKey(overrides, "MIMO_API_KEY") - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - resolveBaseUrlEnv(overrides, "MIMO_BASE_URL"), - PROVIDER_INFO.mimo?.defaultBaseUrl, - ) - // Use createDeepSeek to properly handle reasoning_content for MiMo - // thinking models (e.g., mimo-v2.5-pro). MiMo's API requires - // reasoning_content to be passed back during multi-turn tool calls - // (returns 400 otherwise), same convention as DeepSeek and Kimi. - const mimoProvider = createDeepSeek({ apiKey, baseURL }) - model = mimoProvider(modelId) - break - } - - case "glm": - case "qwen": - case "qiniu": - case "novita": - case "atlascloud": { - const envVar = PROVIDER_ENV_VARS[provider] - if (!envVar) { - throw new Error( - `API key environment variable not defined for provider: ${provider}`, - ) - } - const apiKey = resolveApiKey(overrides, envVar) - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - resolveBaseUrlEnv( - overrides, - `${provider.toUpperCase()}_BASE_URL`, - ), - PROVIDER_INFO[provider]?.defaultBaseUrl, - ) - const customProvider = createOpenAI({ - apiKey, - baseURL, - }) - model = customProvider.chat(modelId) - break - } - - case "kimi": { - const apiKey = resolveApiKey(overrides, "KIMI_API_KEY") - const baseURL = resolveBaseURL( - overrides?.apiKey, - overrides?.baseUrl, - resolveBaseUrlEnv(overrides, "KIMI_BASE_URL"), - PROVIDER_INFO.kimi?.defaultBaseUrl, - ) - // Use createDeepSeek to properly handle reasoning_content for Kimi - // thinking models (e.g., kimi-k2.6). Kimi's API uses the same - // reasoning_content field as DeepSeek, so this provider correctly - // captures and replays reasoning in multi-turn conversations. - const customProvider = createDeepSeek({ apiKey, baseURL }) - model = customProvider(modelId) - break - } - - default: - throw new Error( - `Unknown AI provider: ${provider}. Supported providers: bedrock, openai, anthropic, google, azure, ollama, openrouter, aihubmix, deepseek, siliconflow, sglang, gateway, edgeone, doubao, modelscope, glm, qwen, qiniu, kimi, minimax, novita, mimo, atlascloud`, - ) - } - - // Apply provider-specific options for all providers except bedrock (which has special handling) - if (customProviderOptions && provider !== "bedrock" && !providerOptions) { - providerOptions = customProviderOptions } return { model, providerOptions, modelId, provider } diff --git a/lib/ssrf-protection.ts b/lib/ssrf-protection.ts index 6b43ddf2..88e4952e 100644 --- a/lib/ssrf-protection.ts +++ b/lib/ssrf-protection.ts @@ -115,3 +115,19 @@ export async function isPrivateUrl(urlString: string): Promise { export function allowPrivateUrls(): boolean { return process.env.ALLOW_PRIVATE_URLS !== "false" } + +/** + * A fetch for requests to a base URL the client chose. With private URLs + * blocked, a public URL could still redirect the request to an internal + * host, so redirects are refused. Undefined when private URLs are allowed. + */ +export function redirectGuardedFetch(): typeof fetch | undefined { + if (allowPrivateUrls()) return undefined + return async (input, init) => { + const response = await fetch(input, { ...init, redirect: "manual" }) + if (response.status >= 300 && response.status < 400) { + throw new Error("Redirects are not allowed for custom base URLs") + } + return response + } +} diff --git a/lib/types/model-config.ts b/lib/types/model-config.ts index a1d9603a..3c2e44e0 100644 --- a/lib/types/model-config.ts +++ b/lib/types/model-config.ts @@ -32,6 +32,8 @@ export interface ModelConfig { modelId: string // e.g., "gpt-4o", "claude-sonnet-4-5" validated?: boolean // Has this model been validated validationError?: string // Error message if validation failed + validationWarning?: string // Passed, but e.g. did not call a tool + responseTime?: number // Milliseconds the last test took } // Provider configuration diff --git a/tests/e2e/model-test.spec.ts b/tests/e2e/model-test.spec.ts new file mode 100644 index 00000000..cf3101a8 --- /dev/null +++ b/tests/e2e/model-test.spec.ts @@ -0,0 +1,67 @@ +import { expect, test } from "@playwright/test" +import { getIframe } from "./lib/fixtures" + +// One provider with three models; the test endpoint answers each differently +const CONFIG = { + version: 1, + providers: [ + { + id: "p1", + provider: "glm", + apiKey: "test-key", + models: [ + { id: "m1", modelId: "model-ok" }, + { id: "m2", modelId: "model-no-tools" }, + { id: "m3", modelId: "model-broken" }, + ], + }, + ], +} + +test("the Test button checks all models at once and shows each result", async ({ + page, +}) => { + await page.addInitScript((config) => { + localStorage.setItem( + "next-ai-draw-io-model-configs", + JSON.stringify(config), + ) + }, CONFIG) + const started: string[] = [] + await page.route("**/api/validate-model", async (route) => { + const { modelId } = route.request().postDataJSON() + started.push(modelId) + // Answer only once all three requests arrived: they run in parallel + while (started.length < 3) await new Promise((r) => setTimeout(r, 50)) + const answers: Record = { + "model-ok": { valid: true, responseTime: 1234 }, + "model-no-tools": { + valid: true, + responseTime: 800, + warning: + "Connected, but the model answered without calling a tool.", + }, + "model-broken": { valid: false, error: "Model not found" }, + } + await route.fulfill({ json: answers[modelId] }) + }) + await page.goto("/", { waitUntil: "networkidle" }) + await getIframe(page).waitFor({ state: "visible", timeout: 30000 }) + + await page.locator("button:has(svg.lucide-bot)").first().click() + await page.getByText("Configure Models...").click() + const dialog = page.locator('[role="dialog"]') + await dialog.getByText("GLM (Zhipu)").first().click() + await dialog.getByRole("button", { name: "Test", exact: true }).click() + + await expect( + dialog.getByText("answered without calling a tool"), + ).toBeVisible({ timeout: 15000 }) + await expect(dialog.getByText("Model not found")).toBeVisible() + await expect(dialog.locator('[title="1.2 s"]')).toBeVisible() + expect(started.sort()).toEqual([ + "model-broken", + "model-no-tools", + "model-ok", + ]) +}) diff --git a/tests/unit/ai-providers.test.ts b/tests/unit/ai-providers.test.ts index bba0ea10..78fa29b7 100644 --- a/tests/unit/ai-providers.test.ts +++ b/tests/unit/ai-providers.test.ts @@ -272,8 +272,15 @@ describe("AIHubMix provider", () => { }) }) +vi.mock("@ai-sdk/openai-compatible", () => { + const mockModel = { specificationVersion: "v3", modelId: "test-model" } + const mockProviderFn = vi.fn(() => mockModel) + const mockCreate = vi.fn(() => mockProviderFn) + return { createOpenAICompatible: mockCreate } +}) + describe("Atlas Cloud provider", () => { - let createOpenAIMock: ReturnType + let createCompatibleMock: ReturnType const savedEnv: Record = {} beforeEach(async () => { @@ -281,9 +288,11 @@ describe("Atlas Cloud provider", () => { savedEnv.ATLASCLOUD_BASE_URL = process.env.ATLASCLOUD_BASE_URL delete process.env.ATLASCLOUD_BASE_URL - const mod = await import("@ai-sdk/openai") - createOpenAIMock = mod.createOpenAI as ReturnType - createOpenAIMock.mockClear() + const mod = await import("@ai-sdk/openai-compatible") + createCompatibleMock = mod.createOpenAICompatible as ReturnType< + typeof vi.fn + > + createCompatibleMock.mockClear() }) afterEach(() => { @@ -299,9 +308,12 @@ describe("Atlas Cloud provider", () => { modelId: "qwen/qwen3.5-flash", }) - expect(createOpenAIMock).toHaveBeenCalledWith({ + // An OpenAI-compatible API; includeUsage keeps quota tracking working + expect(createCompatibleMock).toHaveBeenCalledWith({ + name: "atlascloud", apiKey: "server-atlas-key", baseURL: "https://api.atlascloud.ai/v1", + includeUsage: true, }) }) @@ -313,9 +325,11 @@ describe("Atlas Cloud provider", () => { modelId: "deepseek-ai/deepseek-v4-pro", }) - expect(createOpenAIMock).toHaveBeenCalledWith({ + expect(createCompatibleMock).toHaveBeenCalledWith({ + name: "atlascloud", apiKey: "client-atlas-key", baseURL: "https://proxy.example.com/v1", + includeUsage: true, }) }) }) diff --git a/tests/unit/validate-diagram-route.test.ts b/tests/unit/validate-diagram-route.test.ts index 8332e9ae..24f8d3fa 100644 --- a/tests/unit/validate-diagram-route.test.ts +++ b/tests/unit/validate-diagram-route.test.ts @@ -27,7 +27,11 @@ vi.mock("@/lib/ai-providers", () => ({ chunks: [ { type: "text-start", id: "t" }, ...[json.slice(0, 20), json.slice(20)].map( - (delta) => ({ type: "text-delta", id: "t", delta }), + (delta) => ({ + type: "text-delta", + id: "t", + delta, + }), ), { type: "text-end", id: "t" }, { diff --git a/tests/unit/validate-model-route.test.ts b/tests/unit/validate-model-route.test.ts new file mode 100644 index 00000000..8f8d271f --- /dev/null +++ b/tests/unit/validate-model-route.test.ts @@ -0,0 +1,116 @@ +// @vitest-environment node +import { streamText } from "ai" +import { afterEach, describe, expect, it, vi } from "vitest" +import { POST as validateModel } from "@/app/api/validate-model/route" +import { getAIModel } from "@/lib/ai-providers" + +// Treat every URL as public so no test hits DNS +vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({ + ...(await importOriginal()), + isPrivateUrl: async () => false, +})) + +afterEach(() => { + delete process.env.ALLOW_PRIVATE_URLS + vi.unstubAllGlobals() +}) + +/** An OpenAI-compatible streaming reply made of the given deltas */ +function streamReply(...deltas: object[]) { + const chunk = (delta: object, finish: string | null) => + `data: ${JSON.stringify({ + id: "c1", + object: "chat.completion.chunk", + created: 1, + model: "m", + choices: [{ index: 0, delta, finish_reason: finish }], + })}\n\n` + const body = + deltas.map((d) => chunk(d, null)).join("") + + chunk({}, "stop") + + "data: [DONE]\n\n" + vi.stubGlobal( + "fetch", + vi.fn( + async () => + new Response(body, { + headers: { "content-type": "text/event-stream" }, + }), + ), + ) +} + +const testGlm = async () => { + const res = await validateModel( + new Request("http://localhost/api/validate-model", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + provider: "glm", + apiKey: "key", + modelId: "glm-5", + }), + }), + ) + return res.json() +} + +describe("POST /api/validate-model", () => { + it("passes when the model calls the test tool", async () => { + streamReply({ + role: "assistant", + tool_calls: [ + { + index: 0, + id: "call_1", + type: "function", + function: { name: "ping", arguments: "{}" }, + }, + ], + }) + const data = await testGlm() + expect(data.valid).toBe(true) + expect(data.warning).toBeUndefined() + expect(typeof data.responseTime).toBe("number") + }) + + it("warns when the model answers without a tool call", async () => { + streamReply({ role: "assistant", content: "OK" }) + const data = await testGlm() + expect(data.valid).toBe(true) + expect(data.warning).toMatch(/without calling a tool/) + }) +}) + +describe("chat requests to a client base URL", () => { + it("refuse redirects when private URLs are blocked", async () => { + process.env.ALLOW_PRIVATE_URLS = "false" + vi.stubGlobal( + "fetch", + vi.fn( + async () => + new Response(null, { + status: 302, + headers: { location: "http://169.254.169.254/" }, + }), + ), + ) + const { model } = getAIModel({ + provider: "glm", + apiKey: "key", + baseUrl: "https://attacker.example/v1", + modelId: "glm-5", + }) + let error: unknown + const result = streamText({ + model, + prompt: "hi", + maxRetries: 0, + onError: ({ error: e }) => { + error = e + }, + }) + await result.consumeStream() + expect(String(error)).toMatch(/Redirects are not allowed/) + }) +})