mirror of
https://github.com/DayuanJiang/next-ai-draw-io.git
synced 2026-10-04 00:37:48 +08:00
fix(chat): close credential leaks and harden the chat route
- Vertex: a client-supplied base URL only works with the client's own Vertex key - Accept only data: URLs for file parts in every message, so the server never downloads them - Output budget retry accounts for the thinking budget Bedrock/Anthropic add, and reads Volcengine, DashScope, SGLang and vLLM rejections; falls back to 16000 once - x-max-output-tokens can only lower the budget on server credentials - On server credentials only server models or AI_MODEL entries can be used - Drop tool results together with the invalid tool calls they belong to - Count quota tokens as input + output (cached tokens were counted twice) - Private-URL check for custom base URLs, end Langfuse traces on error/abort/early return - Fix repairToolCall ordering and placeholder, align edit_diagram prompt with operations - Panel Bedrock keys are read from ADMIN_AWS_*; forward the access code to EdgeOne - isMinimalDiagram only treats root cells as an empty canvas
This commit is contained in:
+87
-96
@@ -12,13 +12,17 @@ import fs from "fs/promises"
|
||||
import { jsonrepair } from "jsonrepair"
|
||||
import path from "path"
|
||||
import { z } from "zod"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
import {
|
||||
getAIModel,
|
||||
SINGLE_SYSTEM_PROVIDERS,
|
||||
supportsPromptCaching,
|
||||
usesServerCredentials,
|
||||
} from "@/lib/ai-providers"
|
||||
import { findCachedResponse } from "@/lib/cached-responses"
|
||||
import {
|
||||
dropInvalidToolCalls,
|
||||
fixToolInputJson,
|
||||
isMinimalDiagram,
|
||||
replaceHistoricalToolInputs,
|
||||
validateFileParts,
|
||||
@@ -29,6 +33,7 @@ import {
|
||||
recordTokenUsage,
|
||||
} from "@/lib/dynamo-quota-manager"
|
||||
import {
|
||||
endTrace,
|
||||
getTelemetryConfig,
|
||||
setTraceInput,
|
||||
setTraceOutput,
|
||||
@@ -38,7 +43,11 @@ import {
|
||||
resolveMaxOutputTokens,
|
||||
withOutputTokenLimitFallback,
|
||||
} from "@/lib/output-token-limit"
|
||||
import { findServerModelById } from "@/lib/server-model-config"
|
||||
import {
|
||||
type FlattenedServerModel,
|
||||
findServerModelById,
|
||||
} from "@/lib/server-model-config"
|
||||
import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
|
||||
import { getSystemPrompt } from "@/lib/system-prompts"
|
||||
import { getUserIdFromRequest } from "@/lib/user-id"
|
||||
|
||||
@@ -76,24 +85,14 @@ function createCachedStreamResponse(xml: string): Response {
|
||||
return createUIMessageStreamResponse({ stream })
|
||||
}
|
||||
|
||||
// Responses streamed from the model, whose trace streamText's callbacks end
|
||||
const modelStreamResponses = new WeakSet<Response>()
|
||||
|
||||
// Inner handler function
|
||||
async function handleChatRequest(req: Request): Promise<Response> {
|
||||
// Check for access code
|
||||
const accessCodes =
|
||||
process.env.ACCESS_CODE_LIST?.split(",")
|
||||
.map((code) => code.trim())
|
||||
.filter(Boolean) || []
|
||||
if (accessCodes.length > 0) {
|
||||
const accessCodeHeader = req.headers.get("x-access-code")
|
||||
if (!accessCodeHeader || !accessCodes.includes(accessCodeHeader)) {
|
||||
return Response.json(
|
||||
{
|
||||
error: "Invalid or missing access code. Please configure it in Settings.",
|
||||
},
|
||||
{ status: 401 },
|
||||
)
|
||||
}
|
||||
}
|
||||
const accessDenied = checkAccessCode(req)
|
||||
if (accessDenied) return accessDenied
|
||||
|
||||
const body = await req.json()
|
||||
const { messages, xml, previousXml, sessionId } = body
|
||||
@@ -192,6 +191,15 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
baseUrl = `${origin}/api/edgeai`
|
||||
}
|
||||
|
||||
// Same rule as validate-model: with ALLOW_PRIVATE_URLS=false a request may
|
||||
// not point the server at a private or internal address
|
||||
if (baseUrl && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) {
|
||||
return Response.json(
|
||||
{ error: "Private or internal base URLs are not allowed." },
|
||||
{ status: 400 },
|
||||
)
|
||||
}
|
||||
|
||||
// Get cookie header for EdgeOne authentication (eo_token, eo_time)
|
||||
const cookieHeader = req.headers.get("cookie")
|
||||
|
||||
@@ -201,8 +209,9 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
baseUrlEnv?: string
|
||||
provider?: string
|
||||
} = {}
|
||||
let serverModel: FlattenedServerModel | null = null
|
||||
if (selectedModelId?.startsWith("server:")) {
|
||||
const serverModel = await findServerModelById(selectedModelId)
|
||||
serverModel = await findServerModelById(selectedModelId)
|
||||
console.log(
|
||||
`[Server Model Lookup] ID: ${selectedModelId}, Found: ${!!serverModel}, Provider: ${serverModel?.provider}`,
|
||||
)
|
||||
@@ -221,7 +230,8 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
provider: serverModelConfig.provider || provider,
|
||||
baseUrl,
|
||||
apiKey: req.headers.get("x-ai-api-key"),
|
||||
modelId: req.headers.get("x-ai-model"),
|
||||
// A server model runs the model it was configured with, whatever the header says
|
||||
modelId: serverModel?.modelId || req.headers.get("x-ai-model"),
|
||||
// AWS Bedrock credentials
|
||||
awsAccessKeyId: req.headers.get("x-aws-access-key-id"),
|
||||
awsSecretAccessKey: req.headers.get("x-aws-secret-access-key"),
|
||||
@@ -231,11 +241,14 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
...serverModelConfig,
|
||||
// Vertex AI credentials (Express Mode)
|
||||
vertexApiKey: req.headers.get("x-vertex-api-key"),
|
||||
// Pass cookies for EdgeOne Pages authentication
|
||||
...(provider === "edgeone" &&
|
||||
cookieHeader && {
|
||||
headers: { cookie: cookieHeader },
|
||||
}),
|
||||
// Pass cookies for EdgeOne Pages authentication, and the access code,
|
||||
// which the EdgeOne function checks too
|
||||
...(provider === "edgeone" && {
|
||||
headers: {
|
||||
...(cookieHeader && { cookie: cookieHeader }),
|
||||
"x-access-code": req.headers.get("x-access-code") || "",
|
||||
},
|
||||
}),
|
||||
}
|
||||
|
||||
// Read minimal style preference from header
|
||||
@@ -254,12 +267,32 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
provider: resolvedProvider,
|
||||
} = getAIModel(clientOverrides)
|
||||
|
||||
// On the server's own keys, only run models the server offers: a server
|
||||
// model picked by id (its model name is fixed above) or one in AI_MODEL.
|
||||
// With their own key, users can run any model.
|
||||
const onServerCredentials = usesServerCredentials(
|
||||
resolvedProvider,
|
||||
clientOverrides,
|
||||
)
|
||||
const envModels =
|
||||
process.env.AI_MODEL?.split(",").map((m) => m.trim()) || []
|
||||
if (onServerCredentials && !serverModel && !envModels.includes(modelId)) {
|
||||
return Response.json(
|
||||
{
|
||||
error: `Model "${modelId}" is not available on this server. Add your own API key in Settings to use it.`,
|
||||
},
|
||||
{ status: 400 },
|
||||
)
|
||||
}
|
||||
|
||||
// Retry with a smaller budget if the provider rejects the requested one
|
||||
const model = withOutputTokenLimitFallback(baseModel)
|
||||
|
||||
// User setting wins over server env, so desktop users can raise it themselves
|
||||
// The user setting can raise the budget only on their own key (desktop users
|
||||
// can still raise it themselves); on the server's keys it can only lower it
|
||||
const maxOutputTokens = resolveMaxOutputTokens(
|
||||
req.headers.get("x-max-output-tokens"),
|
||||
onServerCredentials,
|
||||
)
|
||||
console.log(`[maxOutputTokens] ${maxOutputTokens}`)
|
||||
|
||||
@@ -340,32 +373,9 @@ ${userInputText}
|
||||
)
|
||||
|
||||
// Filter out tool-calls with invalid inputs (from failed repair or interrupted streaming)
|
||||
// Bedrock API rejects messages where toolUse.input is not a valid JSON object
|
||||
enhancedMessages = enhancedMessages
|
||||
.map((msg: any) => {
|
||||
if (msg.role !== "assistant" || !Array.isArray(msg.content)) {
|
||||
return msg
|
||||
}
|
||||
const filteredContent = msg.content.filter((part: any) => {
|
||||
if (part.type === "tool-call") {
|
||||
// Check if input is a valid object (not null, undefined, or empty)
|
||||
if (
|
||||
!part.input ||
|
||||
typeof part.input !== "object" ||
|
||||
Object.keys(part.input).length === 0
|
||||
) {
|
||||
console.warn(
|
||||
`[route.ts] Filtering out tool-call with invalid input:`,
|
||||
{ toolName: part.toolName, input: part.input },
|
||||
)
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
return { ...msg, content: filteredContent }
|
||||
})
|
||||
.filter((msg: any) => msg.content && msg.content.length > 0)
|
||||
// and their results. Bedrock API rejects messages where toolUse.input is not a valid
|
||||
// JSON object, and every provider rejects a tool result whose call is gone.
|
||||
enhancedMessages = dropInvalidToolCalls(enhancedMessages)
|
||||
|
||||
// DEBUG: Log modelMessages structure (what's being sent to AI)
|
||||
console.log("[route.ts] Model messages count:", enhancedMessages.length)
|
||||
@@ -410,7 +420,7 @@ ${userInputText}
|
||||
contentParts.push({
|
||||
type: "image",
|
||||
image: filePart.url,
|
||||
mimeType: filePart.mediaType,
|
||||
mediaType: filePart.mediaType,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -471,7 +481,7 @@ ${previousXml}
|
||||
${xml || ""}
|
||||
"""
|
||||
|
||||
IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed. When using edit_diagram, COPY search patterns exactly from the CURRENT XML - attribute order matters!`
|
||||
IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed.`
|
||||
|
||||
const systemMessages = isSingleSystemProvider
|
||||
? [
|
||||
@@ -528,23 +538,11 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
|
||||
error.name === "AI_InvalidToolInputError"
|
||||
) {
|
||||
try {
|
||||
// Pre-process to fix common LLM JSON errors that jsonrepair can't handle
|
||||
let inputToRepair = toolCall.input
|
||||
if (typeof inputToRepair === "string") {
|
||||
// Fix `:=` instead of `: ` (LLM sometimes generates this)
|
||||
inputToRepair = inputToRepair.replace(/:=/g, ": ")
|
||||
// Fix `= "` instead of `: "`
|
||||
inputToRepair = inputToRepair.replace(/=\s*"/g, ': "')
|
||||
// Fix inconsistent quote escaping in XML attributes within JSON strings
|
||||
// Pattern: attribute="value\" where opening quote is unescaped but closing is escaped
|
||||
// Example: y="-20\" should be y=\"-20\"
|
||||
inputToRepair = inputToRepair.replace(
|
||||
/(\w+)="([^"]*?)\\"/g,
|
||||
'$1=\\"$2\\"',
|
||||
)
|
||||
}
|
||||
// Use jsonrepair to fix truncated JSON
|
||||
const repairedInput = jsonrepair(inputToRepair)
|
||||
// Pre-process to fix common LLM JSON errors that jsonrepair can't handle,
|
||||
// then use jsonrepair to fix truncated JSON
|
||||
const repairedInput = jsonrepair(
|
||||
fixToolInputJson(toolCall.input),
|
||||
)
|
||||
console.log(
|
||||
`[repairToolCall] Repaired truncated JSON for tool: ${toolCall.toolName}`,
|
||||
)
|
||||
@@ -554,26 +552,8 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
|
||||
`[repairToolCall] Failed to repair JSON for tool: ${toolCall.toolName}`,
|
||||
repairError,
|
||||
)
|
||||
// Return a placeholder input to avoid API errors in multi-step
|
||||
// The tool will fail gracefully on client side
|
||||
if (toolCall.toolName === "edit_diagram") {
|
||||
return {
|
||||
...toolCall,
|
||||
input: {
|
||||
operations: [],
|
||||
_error: "JSON repair failed - no operations to apply",
|
||||
},
|
||||
}
|
||||
}
|
||||
if (toolCall.toolName === "display_diagram") {
|
||||
return {
|
||||
...toolCall,
|
||||
input: {
|
||||
xml: "",
|
||||
_error: "JSON repair failed - empty diagram",
|
||||
},
|
||||
}
|
||||
}
|
||||
// Keep the original error, so the model and the client see why
|
||||
// the input was rejected and the model can retry the call
|
||||
return null
|
||||
}
|
||||
}
|
||||
@@ -596,7 +576,7 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
|
||||
|
||||
// Record token usage for server-side quota tracking (if enabled)
|
||||
// Use totalUsage (cumulative across all steps) instead of usage (final step only)
|
||||
// Include all 4 token types: input, output, cache read, cache write
|
||||
// inputTokens already includes cache reads and writes in AI SDK 6
|
||||
if (
|
||||
isQuotaEnabled() &&
|
||||
!hasOwnApiKey &&
|
||||
@@ -605,12 +585,16 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
|
||||
) {
|
||||
const totalTokens =
|
||||
(totalUsage.inputTokens || 0) +
|
||||
(totalUsage.outputTokens || 0) +
|
||||
(totalUsage.cachedInputTokens || 0) +
|
||||
(totalUsage.inputTokenDetails?.cacheWriteTokens || 0)
|
||||
(totalUsage.outputTokens || 0)
|
||||
recordTokenUsage(userId, totalTokens)
|
||||
}
|
||||
},
|
||||
// onFinish is skipped when the stream fails or is aborted, so end the trace here
|
||||
onError: ({ error }) => {
|
||||
console.error(error) // what AI SDK does without an onError
|
||||
endTrace()
|
||||
},
|
||||
onAbort: () => endTrace(),
|
||||
tools: {
|
||||
// Client-side tool that will be executed on the client
|
||||
display_diagram: {
|
||||
@@ -782,7 +766,7 @@ Call this tool to get shape names and usage syntax for a specific library.`,
|
||||
}),
|
||||
})
|
||||
|
||||
return result.toUIMessageStreamResponse({
|
||||
const response = result.toUIMessageStreamResponse({
|
||||
sendReasoning: true,
|
||||
messageMetadata: ({ part }) => {
|
||||
if (part.type === "finish") {
|
||||
@@ -796,6 +780,8 @@ Call this tool to get shape names and usage syntax for a specific library.`,
|
||||
return undefined
|
||||
},
|
||||
})
|
||||
modelStreamResponses.add(response)
|
||||
return response
|
||||
}
|
||||
|
||||
// Helper to categorize errors and return appropriate response
|
||||
@@ -862,11 +848,16 @@ function handleError(error: unknown): Response {
|
||||
|
||||
// Wrap handler with error handling
|
||||
async function safeHandler(req: Request): Promise<Response> {
|
||||
let response: Response
|
||||
try {
|
||||
return await handleChatRequest(req)
|
||||
response = await handleChatRequest(req)
|
||||
} catch (error) {
|
||||
return handleError(error)
|
||||
response = handleError(error)
|
||||
}
|
||||
// Early returns, cache hits and errors never reach streamText's callbacks,
|
||||
// so their Langfuse trace has to be ended here
|
||||
if (!modelStreamResponses.has(response)) endTrace()
|
||||
return response
|
||||
}
|
||||
|
||||
// Wrap with Langfuse observe (if configured)
|
||||
|
||||
+2
-1
@@ -12,7 +12,8 @@ AI_PROVIDER=bedrock
|
||||
AI_MODEL=global.anthropic.claude-sonnet-4-5-20250929-v1:0
|
||||
|
||||
# Output limit, all providers (default: 64000). Shared by reasoning and the diagram XML,
|
||||
# so a thinking model can spend it all before the tool call. Users can override it in Settings.
|
||||
# so a thinking model can spend it all before the tool call. Users can lower it in Settings,
|
||||
# and raise it only when they use their own API key, so this also caps cost on server keys.
|
||||
# If a model's own ceiling is lower, the request is retried with that ceiling automatically.
|
||||
# MAX_OUTPUT_TOKENS=64000
|
||||
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
/**
|
||||
* Check the x-access-code header against ACCESS_CODE_LIST.
|
||||
* Returns a 401 response to send back when the check fails, or null when the
|
||||
* request may continue (including when no access codes are configured).
|
||||
*/
|
||||
export function checkAccessCode(req: Request): Response | null {
|
||||
const accessCodes =
|
||||
process.env.ACCESS_CODE_LIST?.split(",")
|
||||
.map((code) => code.trim())
|
||||
.filter(Boolean) || []
|
||||
if (accessCodes.length === 0) return null
|
||||
|
||||
const accessCodeHeader = req.headers.get("x-access-code")
|
||||
if (accessCodeHeader && accessCodes.includes(accessCodeHeader)) return null
|
||||
|
||||
return Response.json(
|
||||
{
|
||||
error: "Invalid or missing access code. Please configure it in Settings.",
|
||||
},
|
||||
{ status: 401 },
|
||||
)
|
||||
}
|
||||
+85
-16
@@ -10,6 +10,10 @@ import { aihubmix, createAihubmix } from "@aihubmix/ai-sdk-provider"
|
||||
import { fromNodeProviderChain } from "@aws-sdk/credential-providers"
|
||||
import { createOpenRouter } from "@openrouter/ai-sdk-provider"
|
||||
import { createOllama, ollama } from "ollama-ai-provider-v2"
|
||||
import {
|
||||
adminProvidersToConfig,
|
||||
loadAdminProviders,
|
||||
} from "@/lib/admin/providers"
|
||||
import { PROVIDER_INFO, type ProviderName } from "@/lib/types/model-config"
|
||||
|
||||
export type { ProviderName }
|
||||
@@ -824,8 +828,16 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
// Use client-provided credentials if available, otherwise fall back to IAM/env vars
|
||||
const hasClientCredentials =
|
||||
overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey
|
||||
// Keys from the admin panel. The ADMIN_ names keep them out of the
|
||||
// default AWS credential chain, which other clients such as the
|
||||
// DynamoDB quota manager use with their own credentials.
|
||||
const adminAccessKeyId = process.env.ADMIN_AWS_ACCESS_KEY_ID
|
||||
const adminSecretAccessKey = process.env.ADMIN_AWS_SECRET_ACCESS_KEY
|
||||
const bedrockRegion =
|
||||
overrides?.awsRegion || process.env.AWS_REGION || "us-west-2"
|
||||
overrides?.awsRegion ||
|
||||
process.env.ADMIN_AWS_REGION ||
|
||||
process.env.AWS_REGION ||
|
||||
"us-west-2"
|
||||
|
||||
const bedrockProvider = hasClientCredentials
|
||||
? createAmazonBedrock({
|
||||
@@ -836,10 +848,16 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
sessionToken: overrides.awsSessionToken,
|
||||
}),
|
||||
})
|
||||
: createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
credentialProvider: fromNodeProviderChain(),
|
||||
})
|
||||
: adminAccessKeyId && adminSecretAccessKey
|
||||
? createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
accessKeyId: adminAccessKeyId,
|
||||
secretAccessKey: adminSecretAccessKey,
|
||||
})
|
||||
: createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
credentialProvider: fromNodeProviderChain(),
|
||||
})
|
||||
model = bedrockProvider(modelId)
|
||||
// Add Anthropic beta options if using Claude models via Bedrock
|
||||
if (modelId.includes("anthropic.claude")) {
|
||||
@@ -872,8 +890,9 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
// for compatibility (most proxies don't support /responses endpoint)
|
||||
const customOpenAI = createOpenAI({ apiKey, baseURL })
|
||||
model = customOpenAI.chat(modelId)
|
||||
} else if (overrides?.apiKey) {
|
||||
// Custom API key but official OpenAI endpoint, use Responses API
|
||||
} 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)
|
||||
@@ -928,7 +947,9 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
overrides?.baseUrl,
|
||||
serverBaseUrl,
|
||||
)
|
||||
if (baseURL || overrides?.apiKey) {
|
||||
// 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 }),
|
||||
@@ -941,8 +962,11 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
}
|
||||
case "vertexai": {
|
||||
// Express Mode: Use API key for authentication
|
||||
const vertexApiKey =
|
||||
overrides?.vertexApiKey || process.env.GOOGLE_VERTEX_API_KEY
|
||||
// SECURITY: a client base URL only ever gets the client's key, so the
|
||||
// server's GOOGLE_VERTEX_API_KEY is never sent to a client-chosen host
|
||||
const vertexApiKey = overrides?.baseUrl
|
||||
? overrides.vertexApiKey
|
||||
: overrides?.vertexApiKey || process.env.GOOGLE_VERTEX_API_KEY
|
||||
|
||||
if (!vertexApiKey) {
|
||||
throw new Error(
|
||||
@@ -951,9 +975,13 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
)
|
||||
}
|
||||
|
||||
// Support custom base URL from env or client override
|
||||
const baseURL =
|
||||
overrides?.baseUrl || process.env.GOOGLE_VERTEX_BASE_URL
|
||||
// Support custom base URL from env or client override.
|
||||
// A client key only goes to the client's URL or the official one.
|
||||
const baseURL = resolveBaseURL(
|
||||
overrides?.vertexApiKey,
|
||||
overrides?.baseUrl,
|
||||
process.env.GOOGLE_VERTEX_BASE_URL,
|
||||
)
|
||||
|
||||
const vertexProvider = createVertex({
|
||||
apiKey: vertexApiKey,
|
||||
@@ -1079,7 +1107,7 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
overrides?.baseUrl,
|
||||
serverBaseUrl,
|
||||
)
|
||||
if (baseURL || overrides?.apiKey) {
|
||||
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
|
||||
const customDeepSeek = createDeepSeek({
|
||||
apiKey,
|
||||
...(baseURL && { baseURL }),
|
||||
@@ -1241,7 +1269,7 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
)
|
||||
// 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) {
|
||||
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
|
||||
const customGateway = createGateway({
|
||||
apiKey,
|
||||
...(baseURL && { baseURL }),
|
||||
@@ -1430,6 +1458,36 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
return { model, providerOptions, headers, modelId, provider }
|
||||
}
|
||||
|
||||
/**
|
||||
* Whether the call is paid for by the server's own credentials (env keys or
|
||||
* IAM role) rather than credentials sent with the request. Mirrors which key
|
||||
* each branch of getAIModel ends up using.
|
||||
*/
|
||||
export function usesServerCredentials(
|
||||
provider: ProviderName,
|
||||
overrides?: ClientOverrides,
|
||||
): boolean {
|
||||
switch (provider) {
|
||||
case "bedrock":
|
||||
return !(overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey)
|
||||
case "vertexai":
|
||||
return !overrides?.vertexApiKey
|
||||
case "edgeone":
|
||||
// The platform's own endpoint, no key involved
|
||||
return false
|
||||
case "ollama":
|
||||
// Only a server key costs money; a keyless local server or the
|
||||
// client's own server does not
|
||||
return (
|
||||
!overrides?.baseUrl &&
|
||||
!overrides?.apiKey &&
|
||||
!!(overrides?.apiKeyEnv || process.env.OLLAMA_API_KEY)
|
||||
)
|
||||
default:
|
||||
return !overrides?.apiKey
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if a model supports prompt caching.
|
||||
* Currently only Claude models on Bedrock support prompt caching.
|
||||
@@ -1464,6 +1522,17 @@ export function getValidationModel(): ReturnType<typeof getAIModel>["model"] {
|
||||
)
|
||||
}
|
||||
|
||||
const { model } = getAIModel({ modelId })
|
||||
// A default set in the admin panel becomes AI_PROVIDER/AI_MODEL, but its key
|
||||
// lives in an ADMIN_-prefixed env var. Point at it the way the chat route
|
||||
// does for server models, or the standard env var is required instead.
|
||||
const panelDefault = adminProvidersToConfig(
|
||||
loadAdminProviders(),
|
||||
).providers.find((p) => p.default && p.provider === process.env.AI_PROVIDER)
|
||||
|
||||
const { model } = getAIModel({
|
||||
modelId,
|
||||
apiKeyEnv: panelDefault?.apiKeyEnv,
|
||||
baseUrlEnv: panelDefault?.baseUrlEnv,
|
||||
})
|
||||
return model
|
||||
}
|
||||
|
||||
+96
-43
@@ -6,25 +6,37 @@ export const MAX_FILE_SIZE = 2 * 1024 * 1024 // 2MB
|
||||
export const MAX_FILES = 5
|
||||
|
||||
// Helper function to validate file parts in messages
|
||||
// Checks every message, since history is sent to the model too
|
||||
export function validateFileParts(messages: any[]): {
|
||||
valid: boolean
|
||||
error?: string
|
||||
} {
|
||||
const lastMessage = messages[messages.length - 1]
|
||||
const fileParts =
|
||||
lastMessage?.parts?.filter((p: any) => p.type === "file") || []
|
||||
for (const message of messages) {
|
||||
const fileParts =
|
||||
message?.parts?.filter((p: any) => p.type === "file") || []
|
||||
|
||||
if (fileParts.length > MAX_FILES) {
|
||||
return {
|
||||
valid: false,
|
||||
error: `Too many files. Maximum ${MAX_FILES} allowed.`,
|
||||
if (fileParts.length > MAX_FILES) {
|
||||
return {
|
||||
valid: false,
|
||||
error: `Too many files. Maximum ${MAX_FILES} allowed.`,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for (const filePart of fileParts) {
|
||||
// Data URLs format: data:image/png;base64,<data>
|
||||
// Base64 increases size by ~33%, so we check the decoded size
|
||||
if (filePart.url?.startsWith("data:")) {
|
||||
for (const filePart of fileParts) {
|
||||
// The client sends files inline. Any other URL would be downloaded
|
||||
// by the server (AI SDK does that for models without URL support).
|
||||
if (
|
||||
typeof filePart.url !== "string" ||
|
||||
!filePart.url.startsWith("data:")
|
||||
) {
|
||||
return {
|
||||
valid: false,
|
||||
error: "Files must be uploaded inline as data URLs.",
|
||||
}
|
||||
}
|
||||
|
||||
// Data URLs format: data:image/png;base64,<data>
|
||||
// Base64 increases size by ~33%, so we check the decoded size
|
||||
const base64Data = filePart.url.split(",")[1]
|
||||
if (base64Data) {
|
||||
const sizeInBytes = Math.ceil((base64Data.length * 3) / 4)
|
||||
@@ -42,48 +54,89 @@ export function validateFileParts(messages: any[]): {
|
||||
}
|
||||
|
||||
// Helper function to check if diagram is minimal/empty
|
||||
// Empty means no mxCell besides the root cells "0" and "1". Cells drawn in
|
||||
// draw.io get random ids, so checking for id="2" is not enough.
|
||||
export function isMinimalDiagram(xml: string): boolean {
|
||||
const stripped = xml.replace(/\s/g, "")
|
||||
return !stripped.includes('id="2"')
|
||||
return !/<mxCell\b[^>]*\bid="(?![01]")/.test(xml)
|
||||
}
|
||||
|
||||
// A tool-call input providers accept: a non-empty JSON object
|
||||
function isValidToolInput(input: unknown): boolean {
|
||||
return !!input && typeof input === "object" && Object.keys(input).length > 0
|
||||
}
|
||||
|
||||
// Helper function to replace historical tool call XML with placeholders
|
||||
// This reduces token usage and forces LLM to rely on the current diagram XML (source of truth)
|
||||
// Also fixes invalid/undefined inputs from interrupted streaming
|
||||
// Tool calls with invalid inputs are left for dropInvalidToolCalls to remove
|
||||
export function replaceHistoricalToolInputs(messages: any[]): any[] {
|
||||
return messages.map((msg) => {
|
||||
if (msg.role !== "assistant" || !Array.isArray(msg.content)) {
|
||||
return msg
|
||||
}
|
||||
const replacedContent = msg.content
|
||||
.map((part: any) => {
|
||||
if (part.type === "tool-call") {
|
||||
const toolName = part.toolName
|
||||
// Fix invalid/undefined inputs from interrupted streaming
|
||||
if (
|
||||
!part.input ||
|
||||
typeof part.input !== "object" ||
|
||||
Object.keys(part.input).length === 0
|
||||
) {
|
||||
// Skip tool calls with invalid inputs entirely
|
||||
return null
|
||||
}
|
||||
if (
|
||||
toolName === "display_diagram" ||
|
||||
toolName === "edit_diagram"
|
||||
) {
|
||||
return {
|
||||
...part,
|
||||
input: {
|
||||
placeholder:
|
||||
"[XML content replaced - see current diagram XML in system context]",
|
||||
},
|
||||
}
|
||||
}
|
||||
const replacedContent = msg.content.map((part: any) => {
|
||||
if (
|
||||
part.type === "tool-call" &&
|
||||
isValidToolInput(part.input) &&
|
||||
(part.toolName === "display_diagram" ||
|
||||
part.toolName === "edit_diagram")
|
||||
) {
|
||||
return {
|
||||
...part,
|
||||
input: {
|
||||
placeholder:
|
||||
"[XML content replaced - see current diagram XML in system context]",
|
||||
},
|
||||
}
|
||||
return part
|
||||
})
|
||||
.filter(Boolean) // Remove null entries (invalid tool calls)
|
||||
}
|
||||
return part
|
||||
})
|
||||
return { ...msg, content: replacedContent }
|
||||
})
|
||||
}
|
||||
|
||||
// Remove tool-calls with invalid inputs (from failed repair or interrupted streaming),
|
||||
// together with their tool-results: providers reject a result whose call is missing.
|
||||
// Messages left empty are removed too (Bedrock rejects empty content arrays).
|
||||
export function dropInvalidToolCalls(messages: any[]): any[] {
|
||||
const droppedIds = new Set<string>()
|
||||
return messages
|
||||
.map((msg) => {
|
||||
if (!Array.isArray(msg.content)) return msg
|
||||
const content = msg.content.filter((part: any) => {
|
||||
if (
|
||||
msg.role === "assistant" &&
|
||||
part.type === "tool-call" &&
|
||||
!isValidToolInput(part.input)
|
||||
) {
|
||||
console.warn(
|
||||
`[chat-helpers] Dropping tool-call with invalid input:`,
|
||||
{ toolName: part.toolName, input: part.input },
|
||||
)
|
||||
droppedIds.add(part.toolCallId)
|
||||
return false
|
||||
}
|
||||
// Results always come after their call, so the id is known by now
|
||||
return !(
|
||||
part.type === "tool-result" &&
|
||||
droppedIds.has(part.toolCallId)
|
||||
)
|
||||
})
|
||||
return { ...msg, content }
|
||||
})
|
||||
.filter((msg) => !Array.isArray(msg.content) || msg.content.length > 0)
|
||||
}
|
||||
|
||||
// Fix common LLM JSON mistakes in tool-call input before jsonrepair runs
|
||||
export function fixToolInputJson(input: string): string {
|
||||
return (
|
||||
input
|
||||
// Inconsistent quote escaping in XML attributes inside JSON strings:
|
||||
// y="-20\" (opening quote unescaped, closing escaped) becomes y=\"-20\".
|
||||
// Must run before the key fix below, which would rewrite the `="`.
|
||||
.replace(/(\w+)="([^"]*?)\\"/g, '$1=\\"$2\\"')
|
||||
// `:=` instead of `: `
|
||||
.replace(/:=/g, ": ")
|
||||
// `"key"= "` instead of `"key": "`, only for JSON keys
|
||||
.replace(/"(\w+)"\s*=\s*"/g, '"$1": "')
|
||||
)
|
||||
}
|
||||
|
||||
+8
-1
@@ -51,8 +51,15 @@ export function setTraceOutput(output: string) {
|
||||
if (!isLangfuseEnabled()) return
|
||||
|
||||
updateActiveTrace({ output })
|
||||
endTrace()
|
||||
}
|
||||
|
||||
// End the observe() wrapper span (AI SDK creates its own child spans with usage).
|
||||
// It uses endOnExit: false, so every request path has to end it, or the trace
|
||||
// is never exported: stream finish, stream error/abort, and early returns.
|
||||
export function endTrace() {
|
||||
if (!isLangfuseEnabled()) return
|
||||
|
||||
// End the observe() wrapper span (AI SDK creates its own child spans with usage)
|
||||
const activeSpan = api.trace.getActiveSpan()
|
||||
if (activeSpan) {
|
||||
activeSpan.end()
|
||||
|
||||
+117
-38
@@ -22,6 +22,12 @@ export const MAX_OUTPUT_TOKENS_LIMIT = 200000
|
||||
*/
|
||||
const MIN_USABLE_OUTPUT_TOKENS = 1024
|
||||
|
||||
/**
|
||||
* Retry budget when a rejection names the budget parameter but no number we can
|
||||
* read. It is the default from before 64000, which these providers ran with.
|
||||
*/
|
||||
const FALLBACK_OUTPUT_TOKENS = 16000
|
||||
|
||||
/** Status codes that can carry a complaint about the requested budget. */
|
||||
const BUDGET_REJECTION_STATUSES = new Set([400, 422])
|
||||
|
||||
@@ -29,24 +35,8 @@ function usableLimit(value: number): number | null {
|
||||
return value >= MIN_USABLE_OUTPUT_TOKENS ? value : null
|
||||
}
|
||||
|
||||
/**
|
||||
* A budget this large exceeds what some models accept. Providers reject it with a
|
||||
* 400 that names the real limit, so we parse the number out and retry once
|
||||
* instead of failing the turn.
|
||||
*
|
||||
* Formats seen in the wild:
|
||||
* - Bedrock: "The maximum tokens you requested exceeds the model limit of 4096."
|
||||
* - OpenRouter: "This endpoint's maximum context length is 64000 tokens. However,
|
||||
* you requested about 64025 tokens (25 of text input, 64000 in the output)."
|
||||
* Note this one is an input+output ceiling, so the input has to be subtracted.
|
||||
* - Anthropic: "max_tokens: 200000 > 64000, which is the maximum allowed..."
|
||||
* - OpenAI: "This model supports at most 16384 completion tokens"
|
||||
*
|
||||
* Every pattern names tokens explicitly. A generic one (an earlier draft matched
|
||||
* "lower than N") would reinterpret unrelated failures, and retrying on a bogus
|
||||
* number turns a readable error into an empty diagram.
|
||||
*/
|
||||
export function parseOutputTokenLimit(error: unknown): number | null {
|
||||
/** Message and body of an error that may be about the budget, or null. */
|
||||
function rejectionText(error: unknown): string | null {
|
||||
const err = error as {
|
||||
message?: unknown
|
||||
responseBody?: unknown
|
||||
@@ -66,24 +56,109 @@ export function parseOutputTokenLimit(error: unknown): number | null {
|
||||
typeof err?.responseBody === "string" ? err.responseBody : "",
|
||||
].join(" ")
|
||||
|
||||
if (!text) return null
|
||||
return text.trim() ? text : null
|
||||
}
|
||||
|
||||
/**
|
||||
* A budget this large exceeds what some models accept. Providers reject it with a
|
||||
* 400 that names the real limit, so we parse the number out and retry once
|
||||
* instead of failing the turn.
|
||||
*
|
||||
* Formats seen in the wild:
|
||||
* - Bedrock: "The maximum tokens you requested exceeds the model limit of 4096."
|
||||
* - OpenRouter: "This endpoint's maximum context length is 64000 tokens. However,
|
||||
* you requested about 64025 tokens (25 of text input, 64000 in the output)."
|
||||
* Note this one is an input+output ceiling, so the input has to be subtracted.
|
||||
* vLLM and SGLang send the same kind of ceiling, with the input written as
|
||||
* "6000 in the messages", "has 6000 input tokens" or "6000 tokens from the input".
|
||||
* - Anthropic: "max_tokens: 200000 > 64000, which is the maximum allowed..."
|
||||
* - OpenAI: "This model supports at most 16384 completion tokens"
|
||||
* - Volcengine Ark: "The parameter `max_tokens` specified in the request are not
|
||||
* valid: integer above maximum value, expected a value <= 32768, but got 64000"
|
||||
* - DashScope: "Range of max_tokens should be [1, 8192]"
|
||||
*
|
||||
* Every pattern names tokens explicitly. A generic one (an earlier draft matched
|
||||
* "lower than N") would reinterpret unrelated failures, and retrying on a bogus
|
||||
* number turns a readable error into an empty diagram.
|
||||
*/
|
||||
function readCeiling(text: string): number | null {
|
||||
// Combined input+output ceiling: subtract the input the provider counted,
|
||||
// plus a small margin because its estimate is approximate.
|
||||
const context = text.match(/maximum context length is (\d+)/i)
|
||||
const context = text.match(/maximum context length (?:is|of) (\d+)/i)
|
||||
if (context) {
|
||||
const input = text.match(/(\d+) of text input/i)
|
||||
return usableLimit(
|
||||
Number(context[1]) - (input ? Number(input[1]) : 0) - 1024,
|
||||
)
|
||||
const input =
|
||||
text.match(/(\d+) of text input/i) ||
|
||||
text.match(/(\d+) in the messages/i) ||
|
||||
text.match(/(\d+) tokens from the input/i) ||
|
||||
text.match(/(\d+) input tokens/i)
|
||||
return Number(context[1]) - (input ? Number(input[1]) : 0) - 1024
|
||||
}
|
||||
|
||||
const output =
|
||||
text.match(/model limit of (\d+)/i) ||
|
||||
text.match(/> (\d+), which is the maximum/i) ||
|
||||
text.match(/at most (\d+) completion tokens/i)
|
||||
text.match(/at most (\d+) completion tokens/i) ||
|
||||
text.match(/max_\w*tokens.*?expected a value (?:<=|\\u003c=) (\d+)/i) ||
|
||||
text.match(/Range of max_tokens should be \[1,\s*(\d+)\]/i)
|
||||
|
||||
return output ? usableLimit(Number(output[1])) : null
|
||||
return output ? Number(output[1]) : null
|
||||
}
|
||||
|
||||
/** The usable output ceiling named in a rejection, or null. */
|
||||
export function parseOutputTokenLimit(error: unknown): number | null {
|
||||
const text = rejectionText(error)
|
||||
const ceiling = text ? readCeiling(text) : null
|
||||
return ceiling === null ? null : usableLimit(ceiling)
|
||||
}
|
||||
|
||||
/**
|
||||
* Thinking budget the provider adds on top of maxOutputTokens. Bedrock and
|
||||
* Anthropic send maxOutputTokens + budgetTokens as max_tokens, so a ceiling in
|
||||
* their rejection covers both.
|
||||
*/
|
||||
function thinkingBudget(providerOptions: unknown): number {
|
||||
const options = providerOptions as
|
||||
| {
|
||||
bedrock?: {
|
||||
reasoningConfig?: { type?: string; budgetTokens?: unknown }
|
||||
}
|
||||
anthropic?: {
|
||||
thinking?: { type?: string; budgetTokens?: unknown }
|
||||
}
|
||||
}
|
||||
| undefined
|
||||
const config =
|
||||
options?.bedrock?.reasoningConfig ?? options?.anthropic?.thinking
|
||||
return config?.type === "enabled" && typeof config.budgetTokens === "number"
|
||||
? config.budgetTokens
|
||||
: 0
|
||||
}
|
||||
|
||||
/**
|
||||
* The budget to retry with after a rejection, or null to surface the error.
|
||||
*/
|
||||
export function retryOutputTokens(
|
||||
error: unknown,
|
||||
params: { maxOutputTokens?: number; providerOptions?: unknown },
|
||||
): number | null {
|
||||
const requested = params.maxOutputTokens
|
||||
const text = rejectionText(error)
|
||||
if (!requested || !text) return null
|
||||
|
||||
const ceiling = readCeiling(text)
|
||||
if (ceiling !== null) {
|
||||
// The ceiling applies to what was actually sent, thinking included,
|
||||
// so the retry has to leave room for the thinking too.
|
||||
const thinking = thinkingBudget(params.providerOptions)
|
||||
if (ceiling >= requested + thinking) return null
|
||||
return usableLimit(ceiling - thinking)
|
||||
}
|
||||
|
||||
// Names the budget parameter, but in a format we cannot read a number from
|
||||
if (/max_\w*tokens/i.test(text) && requested > FALLBACK_OUTPUT_TOKENS) {
|
||||
return FALLBACK_OUTPUT_TOKENS
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -103,17 +178,15 @@ export function withOutputTokenLimitFallback(
|
||||
try {
|
||||
return await doStream()
|
||||
} catch (error) {
|
||||
const limit = parseOutputTokenLimit(error)
|
||||
const requested = params.maxOutputTokens
|
||||
|
||||
if (!limit || !requested || limit >= requested) throw error
|
||||
const retry = retryOutputTokens(error, params)
|
||||
if (!retry) throw error
|
||||
|
||||
console.warn(
|
||||
`[maxOutputTokens] ${requested} rejected, retrying with ${limit}`,
|
||||
`[maxOutputTokens] ${params.maxOutputTokens} rejected, retrying with ${retry}`,
|
||||
)
|
||||
return await inner.doStream({
|
||||
...params,
|
||||
maxOutputTokens: limit,
|
||||
maxOutputTokens: retry,
|
||||
})
|
||||
}
|
||||
},
|
||||
@@ -135,11 +208,17 @@ function validBudget(value: string | null | undefined): number | null {
|
||||
* desktop app too), then server env, then the default. Both sources go through
|
||||
* the same validation, so a typo in either falls back instead of reaching the
|
||||
* provider.
|
||||
*
|
||||
* On the server's credentials the user setting can only lower the server value,
|
||||
* so MAX_OUTPUT_TOKENS keeps capping what the server pays for.
|
||||
*/
|
||||
export function resolveMaxOutputTokens(headerValue: string | null): number {
|
||||
return (
|
||||
validBudget(headerValue) ??
|
||||
validBudget(process.env.MAX_OUTPUT_TOKENS) ??
|
||||
DEFAULT_MAX_OUTPUT_TOKENS
|
||||
)
|
||||
export function resolveMaxOutputTokens(
|
||||
headerValue: string | null,
|
||||
usesServerCredentials: boolean,
|
||||
): number {
|
||||
const header = validBudget(headerValue)
|
||||
const server =
|
||||
validBudget(process.env.MAX_OUTPUT_TOKENS) ?? DEFAULT_MAX_OUTPUT_TOKENS
|
||||
if (header === null) return server
|
||||
return usesServerCredentials ? Math.min(header, server) : header
|
||||
}
|
||||
|
||||
@@ -41,7 +41,7 @@ parameters: {
|
||||
tool name: edit_diagram
|
||||
description: Edit specific parts of the EXISTING diagram. Use this when making small targeted changes like adding/removing elements, changing labels, or adjusting properties. This is more efficient than regenerating the entire diagram.
|
||||
parameters: {
|
||||
edits: Array<{search: string, replace: string}>
|
||||
operations: Array<{operation: "update" | "add" | "delete", cell_id: string, new_xml?: string}>
|
||||
}
|
||||
---Tool3---
|
||||
tool name: append_diagram
|
||||
|
||||
@@ -0,0 +1,274 @@
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
import {
|
||||
getAIModel,
|
||||
getValidationModel,
|
||||
usesServerCredentials,
|
||||
} from "@/lib/ai-providers"
|
||||
|
||||
const settings = vi.hoisted(() => ({ values: {} as Record<string, string> }))
|
||||
|
||||
vi.mock("@/lib/admin/settings", () => ({
|
||||
loadSettings: () => settings.values,
|
||||
}))
|
||||
|
||||
vi.mock("@ai-sdk/google-vertex", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createVertex: vi.fn(() => mockProviderFn) }
|
||||
})
|
||||
|
||||
vi.mock("@ai-sdk/openai", () => {
|
||||
const mockModel = { modelId: "test-model" }
|
||||
const mockProviderFn = vi.fn(() => mockModel) as any
|
||||
mockProviderFn.chat = vi.fn(() => mockModel)
|
||||
return {
|
||||
createOpenAI: vi.fn(() => mockProviderFn),
|
||||
openai: vi.fn(() => mockModel),
|
||||
}
|
||||
})
|
||||
|
||||
vi.mock("@ai-sdk/amazon-bedrock", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createAmazonBedrock: vi.fn(() => mockProviderFn) }
|
||||
})
|
||||
|
||||
vi.mock("@aws-sdk/credential-providers", () => ({
|
||||
fromNodeProviderChain: vi.fn(() => "node-chain"),
|
||||
}))
|
||||
|
||||
vi.mock("@openrouter/ai-sdk-provider", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createOpenRouter: vi.fn(() => mockProviderFn) }
|
||||
})
|
||||
|
||||
const ENV_KEYS = [
|
||||
"GOOGLE_VERTEX_API_KEY",
|
||||
"GOOGLE_VERTEX_BASE_URL",
|
||||
"OPENAI_API_KEY",
|
||||
"OPENAI_BASE_URL",
|
||||
"OPENROUTER_API_KEY",
|
||||
"ADMIN_OPENAI_API_KEY",
|
||||
"ADMIN_OPENROUTER_API_KEY",
|
||||
"OLLAMA_API_KEY",
|
||||
"ADMIN_AWS_ACCESS_KEY_ID",
|
||||
"ADMIN_AWS_SECRET_ACCESS_KEY",
|
||||
"ADMIN_AWS_REGION",
|
||||
"AWS_REGION",
|
||||
"AI_PROVIDER",
|
||||
"AI_MODEL",
|
||||
"VALIDATION_MODEL",
|
||||
]
|
||||
const savedEnv: Record<string, string | undefined> = {}
|
||||
|
||||
beforeEach(() => {
|
||||
for (const key of ENV_KEYS) {
|
||||
savedEnv[key] = process.env[key]
|
||||
delete process.env[key]
|
||||
}
|
||||
settings.values = {}
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
for (const key of ENV_KEYS) {
|
||||
if (savedEnv[key] === undefined) delete process.env[key]
|
||||
else process.env[key] = savedEnv[key]
|
||||
}
|
||||
})
|
||||
|
||||
describe("Vertex AI key security", () => {
|
||||
it("never sends the server key to a client base URL", () => {
|
||||
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
|
||||
|
||||
// Any x-ai-api-key passes the outer guard; the branch must still refuse
|
||||
expect(() =>
|
||||
getAIModel({
|
||||
provider: "vertexai",
|
||||
apiKey: "x",
|
||||
baseUrl: "https://attacker.example",
|
||||
modelId: "gemini-2.5-flash",
|
||||
}),
|
||||
).toThrow("Vertex AI requires an API key")
|
||||
})
|
||||
|
||||
it("sends the client key to the client base URL", async () => {
|
||||
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
|
||||
const { createVertex } = await import("@ai-sdk/google-vertex")
|
||||
|
||||
getAIModel({
|
||||
provider: "vertexai",
|
||||
vertexApiKey: "client-key",
|
||||
baseUrl: "https://my-proxy.example",
|
||||
modelId: "gemini-2.5-flash",
|
||||
})
|
||||
|
||||
expect(createVertex).toHaveBeenCalledWith({
|
||||
apiKey: "client-key",
|
||||
baseURL: "https://my-proxy.example",
|
||||
})
|
||||
})
|
||||
|
||||
it("does not send the client key to the server's base URL", async () => {
|
||||
process.env.GOOGLE_VERTEX_BASE_URL = "https://server-proxy.internal"
|
||||
const { createVertex } = await import("@ai-sdk/google-vertex")
|
||||
|
||||
getAIModel({
|
||||
provider: "vertexai",
|
||||
vertexApiKey: "client-key",
|
||||
modelId: "gemini-2.5-flash",
|
||||
})
|
||||
|
||||
expect(createVertex).toHaveBeenCalledWith({ apiKey: "client-key" })
|
||||
})
|
||||
|
||||
it("still uses the server key and base URL without client overrides", async () => {
|
||||
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
|
||||
process.env.GOOGLE_VERTEX_BASE_URL = "https://server-proxy.internal"
|
||||
const { createVertex } = await import("@ai-sdk/google-vertex")
|
||||
|
||||
getAIModel({ provider: "vertexai", modelId: "gemini-2.5-flash" })
|
||||
|
||||
expect(createVertex).toHaveBeenCalledWith({
|
||||
apiKey: "server-vertex-key",
|
||||
baseURL: "https://server-proxy.internal",
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe("Bedrock admin panel credentials", () => {
|
||||
it("uses the ADMIN_AWS_* keys when the client sends none", async () => {
|
||||
process.env.ADMIN_AWS_ACCESS_KEY_ID = "panel-id"
|
||||
process.env.ADMIN_AWS_SECRET_ACCESS_KEY = "panel-secret"
|
||||
process.env.ADMIN_AWS_REGION = "eu-west-1"
|
||||
process.env.AWS_REGION = "us-east-1"
|
||||
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
|
||||
|
||||
getAIModel({ provider: "bedrock", modelId: "amazon.nova-lite-v1:0" })
|
||||
|
||||
expect(createAmazonBedrock).toHaveBeenCalledWith({
|
||||
region: "eu-west-1",
|
||||
accessKeyId: "panel-id",
|
||||
secretAccessKey: "panel-secret",
|
||||
})
|
||||
})
|
||||
|
||||
it("prefers the client's keys and region", async () => {
|
||||
process.env.ADMIN_AWS_ACCESS_KEY_ID = "panel-id"
|
||||
process.env.ADMIN_AWS_SECRET_ACCESS_KEY = "panel-secret"
|
||||
process.env.ADMIN_AWS_REGION = "eu-west-1"
|
||||
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
|
||||
|
||||
getAIModel({
|
||||
provider: "bedrock",
|
||||
modelId: "amazon.nova-lite-v1:0",
|
||||
awsAccessKeyId: "client-id",
|
||||
awsSecretAccessKey: "client-secret",
|
||||
awsRegion: "ap-northeast-1",
|
||||
})
|
||||
|
||||
expect(createAmazonBedrock).toHaveBeenCalledWith({
|
||||
region: "ap-northeast-1",
|
||||
accessKeyId: "client-id",
|
||||
secretAccessKey: "client-secret",
|
||||
})
|
||||
})
|
||||
|
||||
it("falls back to the default AWS credential chain", async () => {
|
||||
process.env.AWS_REGION = "us-east-1"
|
||||
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
|
||||
|
||||
getAIModel({ provider: "bedrock", modelId: "amazon.nova-lite-v1:0" })
|
||||
|
||||
expect(createAmazonBedrock).toHaveBeenCalledWith({
|
||||
region: "us-east-1",
|
||||
credentialProvider: "node-chain",
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe("usesServerCredentials", () => {
|
||||
it("is true when no key comes with the request", () => {
|
||||
expect(usesServerCredentials("openai", {})).toBe(true)
|
||||
expect(usesServerCredentials("openai", { apiKey: "k" })).toBe(false)
|
||||
})
|
||||
|
||||
it("looks at the credential each provider actually uses", () => {
|
||||
// A stray x-ai-api-key does not replace the IAM role or Vertex key
|
||||
expect(usesServerCredentials("bedrock", { apiKey: "x" })).toBe(true)
|
||||
expect(
|
||||
usesServerCredentials("bedrock", {
|
||||
awsAccessKeyId: "id",
|
||||
awsSecretAccessKey: "secret",
|
||||
}),
|
||||
).toBe(false)
|
||||
expect(usesServerCredentials("vertexai", { apiKey: "x" })).toBe(true)
|
||||
expect(usesServerCredentials("vertexai", { vertexApiKey: "k" })).toBe(
|
||||
false,
|
||||
)
|
||||
})
|
||||
|
||||
it("treats keyless EdgeOne and local Ollama as free", () => {
|
||||
expect(usesServerCredentials("edgeone", {})).toBe(false)
|
||||
expect(usesServerCredentials("ollama", {})).toBe(false)
|
||||
expect(
|
||||
usesServerCredentials("ollama", {
|
||||
baseUrl: "http://localhost:11434",
|
||||
}),
|
||||
).toBe(false)
|
||||
|
||||
process.env.OLLAMA_API_KEY = "server-ollama-key"
|
||||
expect(usesServerCredentials("ollama", {})).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
describe("server model apiKeyEnv", () => {
|
||||
it("uses the custom env var on the official OpenAI endpoint", async () => {
|
||||
process.env.ADMIN_OPENAI_API_KEY = "panel-key"
|
||||
const { createOpenAI, openai } = await import("@ai-sdk/openai")
|
||||
|
||||
getAIModel({
|
||||
provider: "openai",
|
||||
modelId: "gpt-4o",
|
||||
apiKeyEnv: "ADMIN_OPENAI_API_KEY",
|
||||
})
|
||||
|
||||
// The default instance would read OPENAI_API_KEY instead
|
||||
expect(openai).not.toHaveBeenCalled()
|
||||
expect(createOpenAI).toHaveBeenCalledWith({ apiKey: "panel-key" })
|
||||
})
|
||||
})
|
||||
|
||||
describe("getValidationModel", () => {
|
||||
it("uses the admin panel default's ADMIN_ key", async () => {
|
||||
settings.values = {
|
||||
ADMIN_PROVIDERS: JSON.stringify([
|
||||
{
|
||||
id: "p1",
|
||||
provider: "openrouter",
|
||||
name: "My OpenRouter",
|
||||
apiKey: "panel-key",
|
||||
models: ["openai/gpt-4o"],
|
||||
isDefault: true,
|
||||
},
|
||||
]),
|
||||
}
|
||||
// What deriveEnvUpdates writes for that panel config
|
||||
process.env.AI_PROVIDER = "openrouter"
|
||||
process.env.AI_MODEL = "openai/gpt-4o"
|
||||
process.env.ADMIN_OPENROUTER_API_KEY = "panel-key"
|
||||
const { createOpenRouter } = await import("@openrouter/ai-sdk-provider")
|
||||
|
||||
expect(() => getValidationModel()).not.toThrow()
|
||||
expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "panel-key" })
|
||||
})
|
||||
|
||||
it("uses the standard env vars without a panel default", async () => {
|
||||
process.env.AI_PROVIDER = "openrouter"
|
||||
process.env.AI_MODEL = "openai/gpt-4o"
|
||||
process.env.OPENROUTER_API_KEY = "env-key"
|
||||
const { createOpenRouter } = await import("@openrouter/ai-sdk-provider")
|
||||
|
||||
getValidationModel()
|
||||
|
||||
expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "env-key" })
|
||||
})
|
||||
})
|
||||
@@ -1,6 +1,11 @@
|
||||
// @vitest-environment node
|
||||
|
||||
import { convertToModelMessages } from "ai"
|
||||
import { jsonrepair } from "jsonrepair"
|
||||
import { describe, expect, it } from "vitest"
|
||||
import {
|
||||
dropInvalidToolCalls,
|
||||
fixToolInputJson,
|
||||
isMinimalDiagram,
|
||||
replaceHistoricalToolInputs,
|
||||
validateFileParts,
|
||||
@@ -65,6 +70,29 @@ describe("validateFileParts", () => {
|
||||
expect(result.valid).toBe(false)
|
||||
expect(result.error).toContain("exceeds")
|
||||
})
|
||||
|
||||
it("rejects file URLs the server would have to download", () => {
|
||||
for (const url of [
|
||||
"http://10.0.0.5/secret.png",
|
||||
"https://example.com/a.png",
|
||||
undefined,
|
||||
]) {
|
||||
const messages = [{ role: "user", parts: [{ type: "file", url }] }]
|
||||
expect(validateFileParts(messages).valid).toBe(false)
|
||||
}
|
||||
})
|
||||
|
||||
it("checks files in earlier messages too", () => {
|
||||
const messages = [
|
||||
{
|
||||
role: "user",
|
||||
parts: [{ type: "file", url: "http://169.254.169.254/x" }],
|
||||
},
|
||||
{ role: "assistant", parts: [{ type: "text", text: "ok" }] },
|
||||
{ role: "user", parts: [{ type: "text", text: "hello" }] },
|
||||
]
|
||||
expect(validateFileParts(messages).valid).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe("isMinimalDiagram", () => {
|
||||
@@ -83,6 +111,18 @@ describe("isMinimalDiagram", () => {
|
||||
const xml = ' <mxCell id="0"/> <mxCell id="1" parent="0"/> '
|
||||
expect(isMinimalDiagram(xml)).toBe(true)
|
||||
})
|
||||
|
||||
it("returns false for a shape drawn in draw.io with a random id", () => {
|
||||
const xml =
|
||||
'<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="xY3kQ9-1" value="" style="rounded=0;" vertex="1" parent="1"><mxGeometry x="10" y="10" width="120" height="60" as="geometry"/></mxCell></root></mxGraphModel>'
|
||||
expect(isMinimalDiagram(xml)).toBe(false)
|
||||
})
|
||||
|
||||
it("does not mistake ids that start with 0 or 1 for root cells", () => {
|
||||
const xml =
|
||||
'<mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="10"/>'
|
||||
expect(isMinimalDiagram(xml)).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe("replaceHistoricalToolInputs", () => {
|
||||
@@ -124,7 +164,7 @@ describe("replaceHistoricalToolInputs", () => {
|
||||
)
|
||||
})
|
||||
|
||||
it("removes tool calls with invalid inputs", () => {
|
||||
it("leaves tool calls with invalid inputs for dropInvalidToolCalls", () => {
|
||||
const messages = [
|
||||
{
|
||||
role: "assistant",
|
||||
@@ -143,7 +183,7 @@ describe("replaceHistoricalToolInputs", () => {
|
||||
},
|
||||
]
|
||||
const result = replaceHistoricalToolInputs(messages)
|
||||
expect(result[0].content).toHaveLength(0)
|
||||
expect(result[0].content).toEqual(messages[0].content)
|
||||
})
|
||||
|
||||
it("preserves non-assistant messages", () => {
|
||||
@@ -169,3 +209,123 @@ describe("replaceHistoricalToolInputs", () => {
|
||||
expect(result[0].content[0].input).toEqual({ foo: "bar" })
|
||||
})
|
||||
})
|
||||
|
||||
describe("dropInvalidToolCalls", () => {
|
||||
it("drops an invalid tool-call together with its tool-result", () => {
|
||||
const messages = [
|
||||
{ role: "user", content: [{ type: "text", text: "draw" }] },
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{
|
||||
type: "tool-call",
|
||||
toolCallId: "call-1",
|
||||
toolName: "display_diagram",
|
||||
input: undefined,
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
role: "tool",
|
||||
content: [
|
||||
{
|
||||
type: "tool-result",
|
||||
toolCallId: "call-1",
|
||||
toolName: "display_diagram",
|
||||
output: { type: "error-text", value: "Stopped" },
|
||||
},
|
||||
],
|
||||
},
|
||||
{ role: "user", content: [{ type: "text", text: "again" }] },
|
||||
]
|
||||
const result = dropInvalidToolCalls(messages)
|
||||
expect(result.map((m) => m.role)).toEqual(["user", "user"])
|
||||
})
|
||||
|
||||
it("keeps valid calls and results in the same messages", () => {
|
||||
const messages = [
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{ type: "text", text: "Here you go" },
|
||||
{
|
||||
type: "tool-call",
|
||||
toolCallId: "bad",
|
||||
toolName: "edit_diagram",
|
||||
input: "{broken",
|
||||
},
|
||||
{
|
||||
type: "tool-call",
|
||||
toolCallId: "good",
|
||||
toolName: "display_diagram",
|
||||
input: { xml: "<mxCell/>" },
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
role: "tool",
|
||||
content: [
|
||||
{ type: "tool-result", toolCallId: "bad", output: {} },
|
||||
{ type: "tool-result", toolCallId: "good", output: {} },
|
||||
],
|
||||
},
|
||||
]
|
||||
const result = dropInvalidToolCalls(messages)
|
||||
expect(result[0].content.map((p: any) => p.toolCallId)).toEqual([
|
||||
undefined,
|
||||
"good",
|
||||
])
|
||||
expect(result[1].content.map((p: any) => p.toolCallId)).toEqual([
|
||||
"good",
|
||||
])
|
||||
})
|
||||
|
||||
it("cleans up a tool call the user stopped before its input arrived", async () => {
|
||||
// handleStop turns a still-streaming call into output-error with no input
|
||||
const modelMessages = await convertToModelMessages([
|
||||
{ role: "user", parts: [{ type: "text", text: "draw" }] },
|
||||
{
|
||||
role: "assistant",
|
||||
parts: [
|
||||
{
|
||||
type: "tool-display_diagram",
|
||||
toolCallId: "call-1",
|
||||
state: "output-error",
|
||||
input: undefined,
|
||||
errorText: "Stopped by user",
|
||||
} as any,
|
||||
],
|
||||
},
|
||||
{ role: "user", parts: [{ type: "text", text: "again" }] },
|
||||
])
|
||||
const result = dropInvalidToolCalls(modelMessages)
|
||||
expect(result.map((m) => m.role)).toEqual(["user", "user"])
|
||||
})
|
||||
|
||||
it("leaves messages with string content alone", () => {
|
||||
const messages = [{ role: "system", content: "You are..." }]
|
||||
expect(dropInvalidToolCalls(messages)).toEqual(messages)
|
||||
})
|
||||
})
|
||||
|
||||
describe("fixToolInputJson", () => {
|
||||
it("fixes an attribute whose closing quote alone is escaped", () => {
|
||||
const input =
|
||||
'{"xml": "<mxCell id=\\"2\\" vertex=\\"1\\"><mxGeometry x=\\"10\\" y="-20\\" as=\\"geometry\\"/></mxCell>"}'
|
||||
const parsed = JSON.parse(jsonrepair(fixToolInputJson(input)))
|
||||
expect(parsed.xml).toContain('y="-20"')
|
||||
expect(parsed.xml).toContain('id="2"')
|
||||
})
|
||||
|
||||
it("fixes = used instead of : after a JSON key", () => {
|
||||
const input = '{"xml"= "<mxCell id=\\"2\\"/>"}'
|
||||
const parsed = JSON.parse(jsonrepair(fixToolInputJson(input)))
|
||||
expect(parsed.xml).toBe('<mxCell id="2"/>')
|
||||
})
|
||||
|
||||
it("leaves well-formed input unchanged", () => {
|
||||
const input =
|
||||
'{"operations": [{"operation": "add", "cell_id": "a", "new_xml": "<mxCell id=\\"a\\" value=\\"x=1\\"/>"}]}'
|
||||
expect(fixToolInputJson(input)).toBe(input)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -3,6 +3,7 @@ import {
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
parseOutputTokenLimit,
|
||||
resolveMaxOutputTokens,
|
||||
retryOutputTokens,
|
||||
withOutputTokenLimitFallback,
|
||||
} from "@/lib/output-token-limit"
|
||||
|
||||
@@ -98,34 +99,204 @@ describe("parseOutputTokenLimit", () => {
|
||||
}
|
||||
expect(parseOutputTokenLimit(error)).toBeNull()
|
||||
})
|
||||
|
||||
it("reads the ceiling from a Volcengine Ark rejection", () => {
|
||||
const error = {
|
||||
message:
|
||||
"The parameter `max_tokens` specified in the request are not valid: integer above maximum value, expected a value <= 32768, but got 64000 instead.",
|
||||
statusCode: 400,
|
||||
}
|
||||
expect(parseOutputTokenLimit(error)).toBe(32768)
|
||||
// Same message JSON-escaped in the response body
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message: "Bad request",
|
||||
responseBody:
|
||||
'{"error":{"message":"The parameter `max_tokens` specified in the request are not valid: integer above maximum value, expected a value \\u003c= 16384, but got 64000 instead."}}',
|
||||
}),
|
||||
).toBe(16384)
|
||||
})
|
||||
|
||||
it("reads the ceiling from a DashScope rejection", () => {
|
||||
const error = {
|
||||
message:
|
||||
"<400> InternalError.Algo.InvalidParameter: Range of max_tokens should be [1, 8192]",
|
||||
}
|
||||
expect(parseOutputTokenLimit(error)).toBe(8192)
|
||||
})
|
||||
|
||||
it("subtracts the input in SGLang and vLLM context rejections", () => {
|
||||
// SGLang
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message:
|
||||
"Requested token count exceeds the model's maximum context length of 32768 tokens. You requested a total of 70000 tokens: 6000 tokens from the input messages and 64000 tokens for the completion.",
|
||||
}),
|
||||
).toBe(32768 - 6000 - 1024)
|
||||
// vLLM, older wording
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message:
|
||||
"This model's maximum context length is 32768 tokens. However, you requested 70000 tokens (6000 in the messages, 64000 in the completion).",
|
||||
}),
|
||||
).toBe(32768 - 6000 - 1024)
|
||||
// vLLM, newer wording
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message:
|
||||
"This model's maximum context length is 32768 tokens and your request has 6000 input tokens (64000 > 32768 - 6000).",
|
||||
}),
|
||||
).toBe(32768 - 6000 - 1024)
|
||||
})
|
||||
})
|
||||
|
||||
describe("retryOutputTokens", () => {
|
||||
const bedrockLimit = Object.assign(
|
||||
new Error(
|
||||
"The maximum tokens you requested exceeds the model limit of 64000.",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
|
||||
it("leaves room for the Bedrock thinking budget the provider adds", () => {
|
||||
// 64000 + 12000 thinking was sent, so the ceiling of 64000 is below it
|
||||
expect(
|
||||
retryOutputTokens(bedrockLimit, {
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
bedrock: {
|
||||
reasoningConfig: {
|
||||
type: "enabled",
|
||||
budgetTokens: 12000,
|
||||
},
|
||||
},
|
||||
},
|
||||
}),
|
||||
).toBe(52000)
|
||||
})
|
||||
|
||||
it("leaves room for the Anthropic thinking budget the provider adds", () => {
|
||||
const error = Object.assign(
|
||||
new Error(
|
||||
"max_tokens: 76000 > 64000, which is the maximum allowed number of output tokens",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
expect(
|
||||
retryOutputTokens(error, {
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
anthropic: {
|
||||
thinking: { type: "enabled", budgetTokens: 12000 },
|
||||
},
|
||||
},
|
||||
}),
|
||||
).toBe(52000)
|
||||
})
|
||||
|
||||
it("does not retry when the ceiling covers what was sent", () => {
|
||||
expect(
|
||||
retryOutputTokens(bedrockLimit, { maxOutputTokens: 64000 }),
|
||||
).toBeNull()
|
||||
})
|
||||
|
||||
it("does not retry when the thinking budget leaves no usable room", () => {
|
||||
const error = Object.assign(new Error("model limit of 16000"), {
|
||||
statusCode: 400,
|
||||
})
|
||||
expect(
|
||||
retryOutputTokens(error, {
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
bedrock: {
|
||||
reasoningConfig: {
|
||||
type: "enabled",
|
||||
budgetTokens: 15500,
|
||||
},
|
||||
},
|
||||
},
|
||||
}),
|
||||
).toBeNull()
|
||||
})
|
||||
|
||||
it("falls back to 16000 when the budget is named but no number can be read", () => {
|
||||
const error = Object.assign(
|
||||
new Error("max_tokens (64000) exceeds the limit for this model"),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBe(16000)
|
||||
// Nothing to gain when the request was already that small
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 16000 })).toBeNull()
|
||||
})
|
||||
|
||||
it("does not fall back for errors that do not name the budget", () => {
|
||||
const error = Object.assign(new Error("temperature must be <= 2"), {
|
||||
statusCode: 400,
|
||||
})
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBeNull()
|
||||
// Not a bad request, even though it names the budget
|
||||
const auth = Object.assign(new Error("max_tokens: invalid API key"), {
|
||||
statusCode: 401,
|
||||
})
|
||||
expect(retryOutputTokens(auth, { maxOutputTokens: 64000 })).toBeNull()
|
||||
})
|
||||
|
||||
it("does not fall back when a ceiling was found but is too small", () => {
|
||||
const error = Object.assign(
|
||||
new Error(
|
||||
"max_tokens: 64000 > 512, which is the maximum allowed number of output tokens",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBeNull()
|
||||
})
|
||||
})
|
||||
|
||||
describe("resolveMaxOutputTokens", () => {
|
||||
it("uses a valid header value", () => {
|
||||
expect(resolveMaxOutputTokens("32000")).toBe(32000)
|
||||
expect(resolveMaxOutputTokens("32000", false)).toBe(32000)
|
||||
expect(resolveMaxOutputTokens("32000", true)).toBe(32000)
|
||||
})
|
||||
|
||||
it("falls back to the default for missing or bogus values", () => {
|
||||
expect(resolveMaxOutputTokens(null)).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("abc")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("0")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("-5")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("1.5")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
// Above the sanity ceiling, e.g. an extra zero
|
||||
expect(resolveMaxOutputTokens("640000")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
for (const value of [null, "", "abc", "0", "-5", "1.5", "640000"]) {
|
||||
// "640000" is above the sanity ceiling, e.g. an extra zero
|
||||
expect(resolveMaxOutputTokens(value, false)).toBe(
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
}
|
||||
})
|
||||
|
||||
it("uses the env value when no header is sent, and validates it too", () => {
|
||||
const original = process.env.MAX_OUTPUT_TOKENS
|
||||
try {
|
||||
process.env.MAX_OUTPUT_TOKENS = "24000"
|
||||
expect(resolveMaxOutputTokens(null)).toBe(24000)
|
||||
// Header still wins
|
||||
expect(resolveMaxOutputTokens("8000")).toBe(8000)
|
||||
expect(resolveMaxOutputTokens(null, true)).toBe(24000)
|
||||
// A lower header still wins
|
||||
expect(resolveMaxOutputTokens("8000", true)).toBe(8000)
|
||||
|
||||
process.env.MAX_OUTPUT_TOKENS = "-1"
|
||||
expect(resolveMaxOutputTokens(null)).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens(null, true)).toBe(
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
} finally {
|
||||
if (original === undefined) delete process.env.MAX_OUTPUT_TOKENS
|
||||
else process.env.MAX_OUTPUT_TOKENS = original
|
||||
}
|
||||
})
|
||||
|
||||
it("lets the header raise the budget only on the user's own credentials", () => {
|
||||
const original = process.env.MAX_OUTPUT_TOKENS
|
||||
try {
|
||||
process.env.MAX_OUTPUT_TOKENS = "16000"
|
||||
expect(resolveMaxOutputTokens("200000", true)).toBe(16000)
|
||||
expect(resolveMaxOutputTokens("200000", false)).toBe(200000)
|
||||
|
||||
// Without MAX_OUTPUT_TOKENS the default is the cap
|
||||
delete process.env.MAX_OUTPUT_TOKENS
|
||||
expect(resolveMaxOutputTokens("100000", true)).toBe(
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
} finally {
|
||||
if (original === undefined) delete process.env.MAX_OUTPUT_TOKENS
|
||||
else process.env.MAX_OUTPUT_TOKENS = original
|
||||
@@ -178,6 +349,34 @@ describe("withOutputTokenLimitFallback", () => {
|
||||
expect(calls.map((c) => c.maxOutputTokens)).toEqual([64000, 4096])
|
||||
})
|
||||
|
||||
it("subtracts the thinking budget from the retry", async () => {
|
||||
const [model, calls] = fakeModel([
|
||||
() =>
|
||||
Promise.reject(
|
||||
Object.assign(
|
||||
new Error(
|
||||
"The maximum tokens you requested exceeds the model limit of 64000.",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
),
|
||||
),
|
||||
() => Promise.resolve(STREAM_OK),
|
||||
])
|
||||
|
||||
const wrapped = withOutputTokenLimitFallback(model)
|
||||
await wrapped.doStream({
|
||||
prompt: [],
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
bedrock: {
|
||||
reasoningConfig: { type: "enabled", budgetTokens: 12000 },
|
||||
},
|
||||
},
|
||||
} as any)
|
||||
|
||||
expect(calls.map((c) => c.maxOutputTokens)).toEqual([64000, 52000])
|
||||
})
|
||||
|
||||
it("does not retry an error it cannot attribute to the budget", async () => {
|
||||
const [model, calls] = fakeModel([
|
||||
() =>
|
||||
|
||||
Reference in New Issue
Block a user