mirror of
https://github.com/DayuanJiang/next-ai-draw-io.git
synced 2026-10-07 18:27:47 +08:00
fix(server): count quota by the key actually used, and more review fixes
Found by the PR review, each with a test that failed first: - Quota: any key header skipped it, even one the provider never reads (x-aws-access-key-id with OpenAI), so a request ran on the server's key without being counted. The check now runs after the model is resolved and uses usesServerCredentials. On main already. - usesServerCredentials read the raw base URL; "/" cleans up to none, so an Ollama request ran on the server's key past the server-model check. - SGLang's default 127.0.0.1:8000 only fills the settings form. Chat and the model list used it as a real address, so the server called its own machine even with private URLs blocked. Now a base URL is required. - With a user's OpenAI key and no base URL, the SDK read the server's OPENAI_BASE_URL. The official endpoint is now passed. On main already. - The Test button refused nothing on the server's keys (Ollama Cloud), and a 15 s timeout reported "connected, no tool call". - The model list for Ollama without a base URL came from ollama.com while chat went to the server's Ollama. - Bedrock's "Too many tokens, please wait" counted as context too long. - On the server's keys the provider's error text stays in the server log; it can name the server's AWS account, role or internal hosts. - Desktop app: the preset keys are the user's own (NEXT_AI_DRAWIO_DESKTOP), so Max Output Tokens can be raised and keyless models in settings work again. A launch that found the remembered port taken no longer replaces it, which hid the user's chats and settings for good.
This commit is contained in:
+31
-39
@@ -133,36 +133,6 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
userId: userId,
|
||||
})
|
||||
|
||||
// === SERVER-SIDE QUOTA CHECK START ===
|
||||
// Quota is opt-in: only enabled when DYNAMODB_QUOTA_TABLE env var is set
|
||||
const hasOwnApiKey = !!(
|
||||
req.headers.get("x-ai-provider") &&
|
||||
(req.headers.get("x-ai-api-key") ||
|
||||
req.headers.get("x-aws-access-key-id") ||
|
||||
req.headers.get("x-vertex-api-key"))
|
||||
)
|
||||
|
||||
// Skip quota check if: quota disabled, user has own API key, or is anonymous
|
||||
if (isQuotaEnabled() && !hasOwnApiKey && userId !== "anonymous") {
|
||||
const quotaCheck = await checkAndIncrementRequest(userId, {
|
||||
requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10,
|
||||
tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000,
|
||||
tpm: Number(process.env.TPM_LIMIT) || 20000,
|
||||
})
|
||||
if (!quotaCheck.allowed) {
|
||||
return Response.json(
|
||||
{
|
||||
error: quotaCheck.error,
|
||||
type: quotaCheck.type,
|
||||
used: quotaCheck.used,
|
||||
limit: quotaCheck.limit,
|
||||
},
|
||||
{ status: 429 },
|
||||
)
|
||||
}
|
||||
}
|
||||
// === SERVER-SIDE QUOTA CHECK END ===
|
||||
|
||||
// === FILE VALIDATION START ===
|
||||
const fileValidation = validateFileParts(messages)
|
||||
if (!fileValidation.valid) {
|
||||
@@ -292,14 +262,41 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
)
|
||||
}
|
||||
|
||||
// === SERVER-SIDE QUOTA CHECK START ===
|
||||
// Quota is opt-in (DYNAMODB_QUOTA_TABLE) and counts what runs on the
|
||||
// server's keys. Decided by the key actually used: a key header the
|
||||
// provider never reads must not skip it.
|
||||
const countsQuota =
|
||||
isQuotaEnabled() && onServerCredentials && userId !== "anonymous"
|
||||
if (countsQuota) {
|
||||
const quotaCheck = await checkAndIncrementRequest(userId, {
|
||||
requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10,
|
||||
tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000,
|
||||
tpm: Number(process.env.TPM_LIMIT) || 20000,
|
||||
})
|
||||
if (!quotaCheck.allowed) {
|
||||
return Response.json(
|
||||
{
|
||||
error: quotaCheck.error,
|
||||
type: quotaCheck.type,
|
||||
used: quotaCheck.used,
|
||||
limit: quotaCheck.limit,
|
||||
},
|
||||
{ status: 429 },
|
||||
)
|
||||
}
|
||||
}
|
||||
// === SERVER-SIDE QUOTA CHECK END ===
|
||||
|
||||
// Retry once if the provider rejects the requested budget, or (newer
|
||||
// Claude models) the sampling or thinking settings
|
||||
const model = withOutputTokenLimitFallback(
|
||||
withDeprecatedParamsFallback(baseModel),
|
||||
)
|
||||
|
||||
// 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
|
||||
// The user setting can raise the budget only on their own key (in the
|
||||
// desktop app every key is the user's); on the server's keys it can only
|
||||
// lower it
|
||||
const maxOutputTokens = resolveMaxOutputTokens(
|
||||
req.headers.get("x-max-output-tokens"),
|
||||
onServerCredentials,
|
||||
@@ -587,12 +584,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)
|
||||
// inputTokens already includes cache reads and writes in AI SDK 6
|
||||
if (
|
||||
isQuotaEnabled() &&
|
||||
!hasOwnApiKey &&
|
||||
userId !== "anonymous" &&
|
||||
totalUsage
|
||||
) {
|
||||
if (countsQuota && totalUsage) {
|
||||
const totalTokens =
|
||||
(totalUsage.inputTokens || 0) +
|
||||
(totalUsage.outputTokens || 0)
|
||||
@@ -724,7 +716,7 @@ Call this tool to get shape names and usage syntax for a specific library.`,
|
||||
|
||||
const response = result.toUIMessageStreamResponse({
|
||||
sendReasoning: true,
|
||||
onError: streamErrorText,
|
||||
onError: (error) => streamErrorText(error, onServerCredentials),
|
||||
messageMetadata: ({ part }) => {
|
||||
if (part.type === "finish") {
|
||||
const usage = (part as any).totalUsage
|
||||
|
||||
@@ -2,14 +2,15 @@ import { streamText, tool } from "ai"
|
||||
import { NextResponse } from "next/server"
|
||||
import { z } from "zod"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
import { getAIModel } from "@/lib/ai-providers"
|
||||
import { getAIModel, usesServerCredentials } from "@/lib/ai-providers"
|
||||
import { classifyLLMError } from "@/lib/llm-errors"
|
||||
import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
|
||||
import type { ProviderName } from "@/lib/types/model-config"
|
||||
|
||||
export const runtime = "nodejs"
|
||||
|
||||
interface ValidateRequest {
|
||||
provider: string
|
||||
provider: ProviderName
|
||||
apiKey: string
|
||||
baseUrl?: string
|
||||
modelId: string
|
||||
@@ -93,6 +94,22 @@ export async function POST(req: Request) {
|
||||
{ status: 400 },
|
||||
)
|
||||
}
|
||||
// The Test button checks the user's own provider. On the server's
|
||||
// keys (Ollama Cloud without a key or URL) anyone could run any model.
|
||||
if (
|
||||
usesServerCredentials(provider, {
|
||||
apiKey,
|
||||
baseUrl,
|
||||
awsAccessKeyId,
|
||||
awsSecretAccessKey,
|
||||
vertexApiKey,
|
||||
})
|
||||
) {
|
||||
return NextResponse.json(
|
||||
{ valid: false, error: "API key is required" },
|
||||
{ status: 400 },
|
||||
)
|
||||
}
|
||||
|
||||
// The same model the chat would use. A client base URL makes it
|
||||
// refuse redirects to internal hosts.
|
||||
@@ -130,6 +147,14 @@ export async function POST(req: Request) {
|
||||
let finishReason: string | undefined
|
||||
for await (const part of result.fullStream) {
|
||||
if (part.type === "error") throw part.error
|
||||
// The timeout ends the stream with an abort part, not an error
|
||||
if (part.type === "abort") {
|
||||
const timeout = new Error(
|
||||
`The model did not answer within ${TEST_TIMEOUT_MS / 1000} s.`,
|
||||
)
|
||||
timeout.name = "TimeoutError"
|
||||
throw timeout
|
||||
}
|
||||
if (part.type === "tool-call") {
|
||||
calledTool = true
|
||||
break
|
||||
|
||||
@@ -87,6 +87,8 @@ async function startServer(): Promise<string> {
|
||||
HOSTNAME: "127.0.0.1",
|
||||
// Enable Node.js built-in proxy support for fetch (Node.js 24+)
|
||||
NODE_USE_ENV_PROXY: "1",
|
||||
// The preset keys are the user's own, not a server's
|
||||
NEXT_AI_DRAWIO_DESKTOP: "1",
|
||||
}
|
||||
|
||||
// Keep requests to local model servers (e.g. Ollama) off the proxy
|
||||
|
||||
@@ -44,10 +44,12 @@ function loadSavedPort(): number | null {
|
||||
}
|
||||
|
||||
/**
|
||||
* Remember the port the production server started on
|
||||
* Remember the port of the first production launch. A later launch that
|
||||
* found it taken keeps it remembered: the user's data lives under that
|
||||
* origin, and the next launch goes back to it once it is free.
|
||||
*/
|
||||
export function saveServerPort(port: number): void {
|
||||
if (!app.isPackaged || port === loadSavedPort()) {
|
||||
if (!app.isPackaged || loadSavedPort() !== null) {
|
||||
return
|
||||
}
|
||||
try {
|
||||
|
||||
+23
-4
@@ -682,7 +682,8 @@ function createModel(
|
||||
// A custom base URL is usually a proxy that only has Chat
|
||||
// Completions; the official endpoint uses the Responses API,
|
||||
// which returns reasoning for the o-series and gpt-5 or later
|
||||
return e.baseURL
|
||||
return e.baseURL &&
|
||||
e.baseURL !== PROVIDER_INFO.openai.defaultBaseUrl
|
||||
? openaiProvider.chat(modelId)
|
||||
: openaiProvider(modelId)
|
||||
}
|
||||
@@ -994,14 +995,28 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
provider === "gateway"
|
||||
? "AI_GATEWAY_BASE_URL"
|
||||
: `${provider.toUpperCase()}_BASE_URL`
|
||||
// A local default (SGLang's 127.0.0.1) only fills the settings
|
||||
// form; the server must not call its own machine for it. With a
|
||||
// user's key the OpenAI SDK would read the server's
|
||||
// OPENAI_BASE_URL, so name the official endpoint.
|
||||
const defaultUrl = PROVIDER_INFO[provider].defaultBaseUrl
|
||||
const publicDefault = defaultUrl?.startsWith("https://")
|
||||
? defaultUrl
|
||||
: undefined
|
||||
const baseURL = resolveBaseURL(
|
||||
overrides?.apiKey,
|
||||
overrides?.baseUrl,
|
||||
resolveBaseUrlEnv(overrides, baseUrlEnv),
|
||||
SDK_KNOWS_ENDPOINT.has(provider)
|
||||
SDK_KNOWS_ENDPOINT.has(provider) &&
|
||||
!(provider === "openai" && overrides?.apiKey)
|
||||
? undefined
|
||||
: PROVIDER_INFO[provider].defaultBaseUrl,
|
||||
: publicDefault,
|
||||
)
|
||||
if (!baseURL && !SDK_KNOWS_ENDPOINT.has(provider)) {
|
||||
throw new Error(
|
||||
`${PROVIDER_INFO[provider].label} needs a base URL. Add it in the model settings.`,
|
||||
)
|
||||
}
|
||||
model = createModel(provider, modelId, {
|
||||
apiKey,
|
||||
baseURL,
|
||||
@@ -1032,6 +1047,10 @@ export function usesServerCredentials(
|
||||
provider: ProviderName,
|
||||
overrides?: ClientOverrides,
|
||||
): boolean {
|
||||
// The desktop app's local server holds the user's own preset keys
|
||||
if (process.env.NEXT_AI_DRAWIO_DESKTOP === "1") return false
|
||||
// Cleaned like getAIModel does: "/" means no base URL
|
||||
const baseUrl = normalizeBaseUrl(overrides?.baseUrl ?? "")
|
||||
switch (provider) {
|
||||
case "bedrock":
|
||||
return !(overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey)
|
||||
@@ -1044,7 +1063,7 @@ export function usesServerCredentials(
|
||||
// Only a server key costs money; a keyless local server or the
|
||||
// client's own server does not
|
||||
return (
|
||||
!overrides?.baseUrl &&
|
||||
!baseUrl &&
|
||||
!overrides?.apiKey &&
|
||||
!!(overrides?.apiKeyEnv || process.env.OLLAMA_API_KEY)
|
||||
)
|
||||
|
||||
+11
-3
@@ -36,7 +36,8 @@ export interface LLMError {
|
||||
// error can come as 403 or 429, a context or image error as a plain 400
|
||||
const SPECIFIC_TEXTS: Array<[RegExp, LLMErrorCode]> = [
|
||||
[
|
||||
/context length|context window|maximum context|prompt is too long|input is too long|too many (?:input )?tokens/i,
|
||||
// Not "too many tokens": that is Bedrock's throttling message
|
||||
/context length|context window|maximum context|prompt is too long|input is too long|too many input tokens/i,
|
||||
"context_too_long",
|
||||
],
|
||||
[
|
||||
@@ -110,12 +111,19 @@ function problemDetail(body: string): string | undefined {
|
||||
/**
|
||||
* The error text for the chat stream: what went wrong with the provider as
|
||||
* JSON for the hint, or the text the model must read to fix a tool call.
|
||||
* On the server's keys the provider's own text stays in the server log:
|
||||
* it can name the server's account, role or internal hosts.
|
||||
*/
|
||||
export function streamErrorText(error: unknown): string {
|
||||
export function streamErrorText(error: unknown, hideDetails = false): string {
|
||||
// The SDK passes an invalid tool call's error as a plain string
|
||||
if (typeof error === "string") return error
|
||||
if (isToolCallError(error)) return (error as Error).message
|
||||
return JSON.stringify(classifyLLMError(error))
|
||||
const classified = classifyLLMError(error)
|
||||
if (hideDetails) {
|
||||
console.error("[chat] Provider error:", error)
|
||||
classified.message = "The provider returned an error."
|
||||
}
|
||||
return JSON.stringify(classified)
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
+19
-3
@@ -75,6 +75,19 @@ async function getJson(
|
||||
return response.json()
|
||||
}
|
||||
|
||||
/**
|
||||
* Where to list from without the user's base URL: where chat goes then. For
|
||||
* Ollama that is the server's Ollama, else the SDK's local default; a local
|
||||
* default in PROVIDER_INFO (SGLang's) only fills the settings form.
|
||||
*/
|
||||
function listFallbackUrl(provider: ProviderName): string {
|
||||
if (provider === "ollama") {
|
||||
return process.env.OLLAMA_BASE_URL || "http://127.0.0.1:11434/api"
|
||||
}
|
||||
const url = PROVIDER_INFO[provider].defaultBaseUrl
|
||||
return url?.startsWith("https://") ? url : ""
|
||||
}
|
||||
|
||||
/**
|
||||
* The provider's chat models, with tool support from the provider's own
|
||||
* data or else models.dev. Only the client's key is used, so the server's
|
||||
@@ -85,9 +98,7 @@ export async function listProviderModels(
|
||||
{ apiKey, baseUrl }: { apiKey?: string; baseUrl?: string },
|
||||
fetchFn: typeof fetch = fetch,
|
||||
): Promise<ListedModel[]> {
|
||||
const base = normalizeBaseUrl(
|
||||
baseUrl || PROVIDER_INFO[provider].defaultBaseUrl || "",
|
||||
)
|
||||
const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider))
|
||||
const bearer: Record<string, string> = apiKey
|
||||
? { Authorization: `Bearer ${apiKey}` }
|
||||
: {}
|
||||
@@ -170,6 +181,11 @@ export async function listProviderModels(
|
||||
break
|
||||
}
|
||||
default: {
|
||||
if (!base) {
|
||||
throw new Error(
|
||||
`${PROVIDER_INFO[provider].label} needs a base URL to list its models.`,
|
||||
)
|
||||
}
|
||||
const data = await getJson(`${base}/models`, bearer, fetchFn)
|
||||
models = (data.data ?? [])
|
||||
.map((m: { id: string }) => ({ id: m.id }))
|
||||
|
||||
@@ -146,7 +146,8 @@ function killProcess(proc) {
|
||||
* Start Next.js dev server with preset environment
|
||||
*/
|
||||
function startNextServer(presetEnv) {
|
||||
const env = { ...process.env }
|
||||
// The preset keys are the user's own, not a server's
|
||||
const env = { ...process.env, NEXT_AI_DRAWIO_DESKTOP: "1" }
|
||||
|
||||
// Apply preset environment variables
|
||||
if (presetEnv) {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import { createOpenAI } from "@ai-sdk/openai"
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
import {
|
||||
getAIModel,
|
||||
@@ -56,6 +57,9 @@ const ENV_KEYS = [
|
||||
"AI_PROVIDER",
|
||||
"AI_MODEL",
|
||||
"VALIDATION_MODEL",
|
||||
"NEXT_AI_DRAWIO_DESKTOP",
|
||||
"SGLANG_API_KEY",
|
||||
"SGLANG_BASE_URL",
|
||||
]
|
||||
const savedEnv: Record<string, string | undefined> = {}
|
||||
|
||||
@@ -272,3 +276,45 @@ describe("getValidationModel", () => {
|
||||
expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "env-key" })
|
||||
})
|
||||
})
|
||||
|
||||
describe("whose keys a request uses", () => {
|
||||
it("cleans the base URL like the request does", () => {
|
||||
// "/" and a pasted path clean up to no base URL: the server's Ollama
|
||||
process.env.OLLAMA_API_KEY = "server-ollama-key"
|
||||
expect(usesServerCredentials("ollama", { baseUrl: "/" })).toBe(true)
|
||||
expect(
|
||||
usesServerCredentials("ollama", { baseUrl: "/chat/completions" }),
|
||||
).toBe(true)
|
||||
})
|
||||
|
||||
it("counts the desktop app's keys as the user's own", () => {
|
||||
// Electron passes the user's preset keys as server env vars
|
||||
process.env.NEXT_AI_DRAWIO_DESKTOP = "1"
|
||||
expect(usesServerCredentials("openai", {})).toBe(false)
|
||||
})
|
||||
|
||||
it("sends a user's OpenAI key to the official endpoint", () => {
|
||||
// The SDK would otherwise read the server's OPENAI_BASE_URL
|
||||
process.env.OPENAI_BASE_URL = "https://operator-proxy.example.com/v1"
|
||||
getAIModel({
|
||||
provider: "openai",
|
||||
apiKey: "user-key",
|
||||
modelId: "gpt-5.5",
|
||||
})
|
||||
expect(createOpenAI).toHaveBeenLastCalledWith(
|
||||
expect.objectContaining({
|
||||
apiKey: "user-key",
|
||||
baseURL: "https://api.openai.com/v1",
|
||||
}),
|
||||
)
|
||||
// Still the Responses API, like without a base URL
|
||||
const provider = vi.mocked(createOpenAI).mock.results.at(-1)?.value
|
||||
expect(provider.chat).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("needs a base URL for SGLang instead of using 127.0.0.1", () => {
|
||||
expect(() =>
|
||||
getAIModel({ provider: "sglang", apiKey: "k", modelId: "m" }),
|
||||
).toThrow(/base URL/)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
// @vitest-environment node
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
|
||||
// Quota on, and every check answers that the daily limit is used up
|
||||
const quota = vi.hoisted(() => ({ checks: 0 }))
|
||||
vi.mock("@/lib/dynamo-quota-manager", () => ({
|
||||
isQuotaEnabled: () => true,
|
||||
checkAndIncrementRequest: async () => {
|
||||
quota.checks++
|
||||
return {
|
||||
allowed: false,
|
||||
error: "Daily limit reached",
|
||||
type: "request",
|
||||
used: 10,
|
||||
limit: 10,
|
||||
}
|
||||
},
|
||||
recordTokenUsage: async () => {},
|
||||
}))
|
||||
|
||||
import { POST as chat } from "@/app/api/chat/route"
|
||||
|
||||
const ENV = ["AI_PROVIDER", "AI_MODEL", "OPENAI_API_KEY"]
|
||||
const saved: Record<string, string | undefined> = {}
|
||||
|
||||
beforeEach(() => {
|
||||
for (const k of ENV) saved[k] = process.env[k]
|
||||
process.env.AI_PROVIDER = "openai"
|
||||
process.env.AI_MODEL = "gpt-5.5"
|
||||
process.env.OPENAI_API_KEY = "server-key"
|
||||
quota.checks = 0
|
||||
// No request may reach a provider
|
||||
vi.stubGlobal(
|
||||
"fetch",
|
||||
vi.fn(async () => {
|
||||
throw new Error("no network in tests")
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
for (const k of ENV) {
|
||||
if (saved[k] === undefined) delete process.env[k]
|
||||
else process.env[k] = saved[k]
|
||||
}
|
||||
vi.unstubAllGlobals()
|
||||
})
|
||||
|
||||
const send = (headers: Record<string, string>) =>
|
||||
chat(
|
||||
new Request("http://localhost/api/chat", {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
"x-forwarded-for": "203.0.113.7",
|
||||
...headers,
|
||||
},
|
||||
body: JSON.stringify({
|
||||
messages: [
|
||||
{
|
||||
id: "u1",
|
||||
role: "user",
|
||||
parts: [{ type: "text", text: "Draw two boxes" }],
|
||||
},
|
||||
],
|
||||
xml: "",
|
||||
}),
|
||||
}),
|
||||
)
|
||||
|
||||
describe("chat quota", () => {
|
||||
it("counts a request whose key header the provider never reads", async () => {
|
||||
// OpenAI ignores the AWS key, so this runs on the server's key
|
||||
const res = await send({
|
||||
"x-ai-provider": "openai",
|
||||
"x-aws-access-key-id": "x",
|
||||
})
|
||||
expect(res.status).toBe(429)
|
||||
expect(quota.checks).toBe(1)
|
||||
})
|
||||
|
||||
it("does not count a request on the user's own key", async () => {
|
||||
const res = await send({
|
||||
"x-ai-provider": "openai",
|
||||
"x-ai-api-key": "user-key",
|
||||
"x-ai-model": "gpt-5.5",
|
||||
})
|
||||
expect(res.status).not.toBe(429)
|
||||
expect(quota.checks).toBe(0)
|
||||
})
|
||||
})
|
||||
@@ -110,6 +110,14 @@ describe("classifyLLMError", () => {
|
||||
expect(classifyLLMError(error).code).toBe("model_not_found")
|
||||
})
|
||||
|
||||
it("reads Bedrock's token throttling as a rate limit", () => {
|
||||
const error = apiError(
|
||||
429,
|
||||
"Too many tokens, please wait before trying again.",
|
||||
)
|
||||
expect(classifyLLMError(error).code).toBe("rate_limited")
|
||||
})
|
||||
|
||||
it("names a network error the SDK wrapped", () => {
|
||||
const error = new APICallError({
|
||||
message:
|
||||
@@ -201,6 +209,19 @@ describe("streamErrorText", () => {
|
||||
}
|
||||
})
|
||||
|
||||
it("hides the provider's text on the server's keys", () => {
|
||||
const error = apiError(
|
||||
403,
|
||||
"User: arn:aws:sts::123456789012:assumed-role/app/s is not authorized to perform: bedrock:InvokeModel",
|
||||
)
|
||||
const hidden = JSON.parse(streamErrorText(error, true))
|
||||
expect(hidden.code).toBe("forbidden")
|
||||
expect(hidden.message).not.toMatch(/arn:aws|123456789012/)
|
||||
expect(JSON.parse(streamErrorText(error)).message).toMatch(
|
||||
/not authorized/,
|
||||
)
|
||||
})
|
||||
|
||||
it("classifies a provider error", () => {
|
||||
expect(JSON.parse(streamErrorText(apiError(401, "bad key")))).toEqual({
|
||||
type: "provider",
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
// @vitest-environment node
|
||||
import { mkdtempSync, readFileSync, writeFileSync } from "node:fs"
|
||||
import { tmpdir } from "node:os"
|
||||
import { join } from "node:path"
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||
|
||||
const userData = vi.hoisted(() => ({ dir: "" }))
|
||||
vi.mock("electron", () => ({
|
||||
app: { isPackaged: true, getPath: () => userData.dir },
|
||||
}))
|
||||
|
||||
import { saveServerPort } from "@/electron/main/port-manager"
|
||||
|
||||
const savedPort = () =>
|
||||
JSON.parse(readFileSync(join(userData.dir, "server-port.json"), "utf-8"))
|
||||
.port
|
||||
|
||||
beforeEach(() => {
|
||||
userData.dir = mkdtempSync(join(tmpdir(), "port-manager-"))
|
||||
})
|
||||
|
||||
describe("saveServerPort", () => {
|
||||
it("remembers the port of the first launch", () => {
|
||||
saveServerPort(13370)
|
||||
expect(savedPort()).toBe(13370)
|
||||
})
|
||||
|
||||
it("keeps the remembered port when a launch had to use another", () => {
|
||||
// The app's data lives under the remembered port's origin; going
|
||||
// back to it once it is free brings the chats and settings back
|
||||
writeFileSync(
|
||||
join(userData.dir, "server-port.json"),
|
||||
JSON.stringify({ port: 61337 }),
|
||||
)
|
||||
saveServerPort(13371)
|
||||
expect(savedPort()).toBe(61337)
|
||||
})
|
||||
})
|
||||
@@ -103,6 +103,29 @@ describe("listProviderModels", () => {
|
||||
])
|
||||
})
|
||||
|
||||
it("lists Ollama from where chat goes without a base URL", async () => {
|
||||
const { fn, calls } = answer({ models: [{ name: "llama3.2" }] })
|
||||
process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434"
|
||||
try {
|
||||
await listProviderModels("ollama", {}, fn)
|
||||
} finally {
|
||||
delete process.env.OLLAMA_BASE_URL
|
||||
}
|
||||
await listProviderModels("ollama", {}, fn)
|
||||
expect(calls.map((c) => c.url)).toEqual([
|
||||
"http://ollama.internal:11434/api/tags",
|
||||
"http://127.0.0.1:11434/api/tags",
|
||||
])
|
||||
})
|
||||
|
||||
it("does not use SGLang's local address as a default", async () => {
|
||||
const { fn, calls } = answer({ data: [] })
|
||||
await expect(
|
||||
listProviderModels("sglang", { apiKey: "k" }, fn),
|
||||
).rejects.toThrow(/base URL/)
|
||||
expect(calls).toHaveLength(0)
|
||||
})
|
||||
|
||||
it("turns a failed request into an error with its status", async () => {
|
||||
const { fn } = answer({ error: "bad key" }, 401)
|
||||
await expect(
|
||||
|
||||
@@ -74,6 +74,50 @@ describe("POST /api/validate-model", () => {
|
||||
expect(typeof data.responseTime).toBe("number")
|
||||
})
|
||||
|
||||
it("reports a model that did not answer in time", async () => {
|
||||
// The 15 s timeout has fired: the SDK ends the stream with an
|
||||
// abort part instead of throwing
|
||||
const timedOut = AbortSignal.abort(
|
||||
new DOMException("The operation timed out.", "TimeoutError"),
|
||||
)
|
||||
const timeout = vi
|
||||
.spyOn(AbortSignal, "timeout")
|
||||
.mockReturnValue(timedOut)
|
||||
vi.stubGlobal(
|
||||
"fetch",
|
||||
vi.fn(async () => {
|
||||
throw timedOut.reason
|
||||
}),
|
||||
)
|
||||
try {
|
||||
const data = await testGlm()
|
||||
expect(data.valid).toBe(false)
|
||||
expect(data.code).toBe("timeout")
|
||||
} finally {
|
||||
timeout.mockRestore()
|
||||
}
|
||||
})
|
||||
|
||||
it("does not run on the server's keys", async () => {
|
||||
process.env.OLLAMA_API_KEY = "server-ollama-key"
|
||||
try {
|
||||
const res = await validateModel(
|
||||
new Request("http://localhost/api/validate-model", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({
|
||||
provider: "ollama",
|
||||
modelId: "any-cloud-model",
|
||||
}),
|
||||
}),
|
||||
)
|
||||
expect(res.status).toBe(400)
|
||||
expect((await res.json()).error).toMatch(/API key/)
|
||||
} finally {
|
||||
delete process.env.OLLAMA_API_KEY
|
||||
}
|
||||
})
|
||||
|
||||
it("warns when the model answers without a tool call", async () => {
|
||||
streamReply({ role: "assistant", content: "OK" })
|
||||
const data = await testGlm()
|
||||
|
||||
Reference in New Issue
Block a user