mirror of
https://github.com/DayuanJiang/next-ai-draw-io.git
synced 2026-10-10 03:29:50 +08:00
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. <think> 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
This commit is contained in:
+59
-381
@@ -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")
|
||||
) {
|
||||
|
||||
@@ -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<Set<string>>(
|
||||
() => new Set(),
|
||||
)
|
||||
const [duplicateError, setDuplicateError] = useState<string>("")
|
||||
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<ModelConfig>
|
||||
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({
|
||||
<div className="flex items-center gap-3 p-3 min-w-0">
|
||||
{/* Status icon */}
|
||||
<div className="flex items-center justify-center w-8 h-8 rounded-lg flex-shrink-0">
|
||||
{validatingModelIndex !==
|
||||
null &&
|
||||
index ===
|
||||
validatingModelIndex ? (
|
||||
{validatingModelIds.has(
|
||||
model.id,
|
||||
) ? (
|
||||
// Currently validating
|
||||
<div className="w-full h-full rounded-lg bg-blue-500/10 flex items-center justify-center">
|
||||
<Loader2 className="h-4 w-4 text-blue-500 animate-spin" />
|
||||
</div>
|
||||
) : validatingModelIndex !==
|
||||
null &&
|
||||
index >
|
||||
validatingModelIndex &&
|
||||
model.validated ===
|
||||
undefined ? (
|
||||
// Queued
|
||||
<div className="w-full h-full rounded-lg bg-muted flex items-center justify-center">
|
||||
<Clock className="h-4 w-4 text-muted-foreground" />
|
||||
</div>
|
||||
) : model.validated ===
|
||||
true ? (
|
||||
// Valid
|
||||
<div className="w-full h-full rounded-lg bg-success-muted flex items-center justify-center">
|
||||
// Valid, with the time the test took
|
||||
<div
|
||||
className="w-full h-full rounded-lg bg-success-muted flex items-center justify-center"
|
||||
title={
|
||||
model.responseTime
|
||||
? `${(model.responseTime / 1000).toFixed(1)} s`
|
||||
: undefined
|
||||
}
|
||||
>
|
||||
<Check className="h-4 w-4 text-success" />
|
||||
</div>
|
||||
) : model.validated ===
|
||||
@@ -1219,6 +1228,14 @@ export function ModelConfigDialog({
|
||||
}
|
||||
</p>
|
||||
)}
|
||||
{model.validated &&
|
||||
model.validationWarning && (
|
||||
<p className="text-[11px] text-amber-600 dark:text-amber-400 px-3 pb-2 pl-14">
|
||||
{
|
||||
model.validationWarning
|
||||
}
|
||||
</p>
|
||||
)}
|
||||
{/* Show edit error inline */}
|
||||
{editError?.modelId ===
|
||||
model.id && (
|
||||
|
||||
+223
-648
@@ -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<ProviderName>([
|
||||
"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<string, string>
|
||||
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 <think> 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
|
||||
* <NAME>_API_KEY / <NAME>_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 <token>` 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
|
||||
// <NAME>_API_KEY / <NAME>_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 }
|
||||
|
||||
@@ -115,3 +115,19 @@ export async function isPrivateUrl(urlString: string): Promise<boolean> {
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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<string, object> = {
|
||||
"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",
|
||||
])
|
||||
})
|
||||
@@ -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<typeof vi.fn>
|
||||
let createCompatibleMock: ReturnType<typeof vi.fn>
|
||||
const savedEnv: Record<string, string | undefined> = {}
|
||||
|
||||
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<typeof vi.fn>
|
||||
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,
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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" },
|
||||
{
|
||||
|
||||
@@ -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<typeof import("@/lib/ssrf-protection")>()),
|
||||
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/)
|
||||
})
|
||||
})
|
||||
Reference in New Issue
Block a user