mirror of
https://github.com/DayuanJiang/next-ai-draw-io.git
synced 2026-10-09 19:19:50 +08:00
fix(server): use the keys the user sent, and more review fixes
Found by the second PR review: - With AWS_BEARER_TOKEN_BEDROCK set on the server, a request with the user's AWS keys ran on the server's token: the Bedrock SDK prefers it. Checked with Bedrock: invalid user keys used to get an answer. - An OpenAI key with the official URL filled in (the settings form does that) went to the Responses API. Back to main's rule: a configured base URL uses Chat Completions. - A user's Ollama key went to the server's OLLAMA_BASE_URL, for chat and for the model list. Like every other provider, it goes to the user's base URL or Ollama Cloud. - The server's keyless Ollama and EdgeOne were not counted in the quota. - AI_MODEL models on the server's keys ran on any provider with a server key, not only on AI_PROVIDER. - A user's Azure key without a base URL used the server's resource name. - The admin panel's Test button failed whenever access codes were set. - DeepSeek's errors in the stream (plain text) were shown as they were, without a hint and also on the server's keys. Bedrock's throttling in the stream was not recognised as a rate limit. - The EdgeOne function accepted text/plain; x=application/json, which other sites can send without a CORS preflight. - Desktop app: a launch that found the old port taken for a moment (the previous version still quitting after an update) remembered the new port for good. The new port is kept only when Windows reserves the old one. A failed read of the presets file moved it aside as corrupt, and a save could then replace the presets. Switching presets on the same port now reloads the page. The dev launcher no longer misses a preset change made before or during a restart.
This commit is contained in:
@@ -50,7 +50,11 @@ export async function POST(req: Request) {
|
||||
return validateModel(
|
||||
new Request(new URL("/api/validate-model", req.url), {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
// Checked again there, in place of an access code
|
||||
"x-admin-password": req.headers.get("x-admin-password") || "",
|
||||
},
|
||||
body: JSON.stringify({
|
||||
provider: resolved.provider,
|
||||
apiKey: resolved.apiKey,
|
||||
|
||||
+16
-6
@@ -14,6 +14,7 @@ import { checkAccessCode } from "@/lib/access-code"
|
||||
import {
|
||||
CACHE_POINT,
|
||||
getAIModel,
|
||||
getServerProvider,
|
||||
SINGLE_SYSTEM_PROVIDERS,
|
||||
supportsPromptCaching,
|
||||
usesServerCredentials,
|
||||
@@ -49,6 +50,7 @@ import {
|
||||
} from "@/lib/server-model-config"
|
||||
import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
|
||||
import { getSystemPrompt } from "@/lib/system-prompts"
|
||||
import { normalizeBaseUrl } from "@/lib/types/model-config"
|
||||
import { getUserIdFromRequest } from "@/lib/user-id"
|
||||
import { hasCells } from "@/packages/mcp-server/src/pages.ts"
|
||||
import {
|
||||
@@ -245,15 +247,17 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
} = 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.
|
||||
// model picked by id (its model name is fixed above) or one in AI_MODEL
|
||||
// on AI_PROVIDER. 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)) {
|
||||
const offeredInEnv =
|
||||
envModels.includes(modelId) && resolvedProvider === getServerProvider()
|
||||
if (onServerCredentials && !serverModel && !offeredInEnv) {
|
||||
return Response.json(
|
||||
{
|
||||
error: `Model "${modelId}" is not available on this server. Add your own API key in Settings to use it.`,
|
||||
@@ -264,10 +268,16 @@ 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.
|
||||
// server's keys, or on its keyless Ollama or EdgeOne. Decided by the key
|
||||
// actually used: a key header the provider never reads must not skip it.
|
||||
const onServerEndpoint =
|
||||
(resolvedProvider === "ollama" || resolvedProvider === "edgeone") &&
|
||||
!clientOverrides.apiKey &&
|
||||
!normalizeBaseUrl(req.headers.get("x-ai-base-url") ?? "")
|
||||
const countsQuota =
|
||||
isQuotaEnabled() && onServerCredentials && userId !== "anonymous"
|
||||
isQuotaEnabled() &&
|
||||
(onServerCredentials || onServerEndpoint) &&
|
||||
userId !== "anonymous"
|
||||
if (countsQuota) {
|
||||
const quotaCheck = await checkAndIncrementRequest(userId, {
|
||||
requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10,
|
||||
|
||||
@@ -2,6 +2,7 @@ import { streamText, tool } from "ai"
|
||||
import { NextResponse } from "next/server"
|
||||
import { z } from "zod"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
import { checkAdminAuth } from "@/lib/admin/auth"
|
||||
import { getAIModel, usesServerCredentials } from "@/lib/ai-providers"
|
||||
import { classifyLLMError } from "@/lib/llm-errors"
|
||||
import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
|
||||
@@ -34,9 +35,10 @@ const NO_TOOL_CALL_WARNING =
|
||||
"Connected, but the model answered without calling a tool. It may not support tool calls, which drawing needs."
|
||||
|
||||
export async function POST(req: Request) {
|
||||
// Lets the server send requests to arbitrary URLs, so require the access code
|
||||
// Lets the server send requests to arbitrary URLs, so require the access
|
||||
// code, or the admin password (the admin panel's Test button)
|
||||
const accessError = checkAccessCode(req)
|
||||
if (accessError) return accessError
|
||||
if (accessError && checkAdminAuth(req)) return accessError
|
||||
|
||||
try {
|
||||
const body: ValidateRequest = await req.json()
|
||||
|
||||
@@ -97,11 +97,13 @@ function hasValidAccessCode(request: Request, env: any): boolean {
|
||||
|
||||
export async function onRequest({ request, env }: any) {
|
||||
// Requiring JSON also makes any cross-site browser request need a CORS
|
||||
// preflight, which fails without CORS headers
|
||||
if (
|
||||
request.method !== "POST" ||
|
||||
!request.headers.get("content-type")?.includes("application/json")
|
||||
) {
|
||||
// preflight, which fails without CORS headers. Only the type before any
|
||||
// parameters counts: "text/plain; x=application/json" needs none.
|
||||
const mediaType = (request.headers.get("content-type") ?? "")
|
||||
.split(";")[0]
|
||||
.trim()
|
||||
.toLowerCase()
|
||||
if (request.method !== "POST" || mediaType !== "application/json") {
|
||||
return createResponse(
|
||||
{
|
||||
error: {
|
||||
|
||||
@@ -159,6 +159,10 @@ function getConfigFilePath(): string {
|
||||
return path.join(userDataPath, CONFIG_FILE_NAME)
|
||||
}
|
||||
|
||||
// The presets file exists but the last read failed: a save now would
|
||||
// replace the user's presets with the empty list that read returned
|
||||
let presetsUnreadable = false
|
||||
|
||||
/**
|
||||
* Load presets from the config file
|
||||
* Decrypts sensitive fields automatically
|
||||
@@ -175,8 +179,24 @@ export function loadPresets(): ConfigPresetsFile {
|
||||
}
|
||||
}
|
||||
|
||||
let content: string
|
||||
try {
|
||||
content = readFileSync(configPath, "utf-8")
|
||||
presetsUnreadable = false
|
||||
} catch (error) {
|
||||
// Often only for now (on Windows an antivirus scanner can hold the
|
||||
// file): keep the file, and refuse saves based on this empty list
|
||||
console.error("Failed to read config presets:", error)
|
||||
presetsUnreadable = true
|
||||
return {
|
||||
version: 1,
|
||||
currentPresetId: null,
|
||||
presets: [],
|
||||
userLocale: undefined,
|
||||
}
|
||||
}
|
||||
|
||||
try {
|
||||
const content = readFileSync(configPath, "utf-8")
|
||||
const data = JSON.parse(content) as ConfigPresetsFile
|
||||
|
||||
// Decrypt sensitive fields in each preset
|
||||
@@ -211,6 +231,11 @@ export function loadPresets(): ConfigPresetsFile {
|
||||
* Encrypts sensitive fields automatically
|
||||
*/
|
||||
export function savePresets(data: ConfigPresetsFile): void {
|
||||
if (presetsUnreadable) {
|
||||
throw new Error(
|
||||
"The presets file could not be read, so it was not overwritten. Please try again.",
|
||||
)
|
||||
}
|
||||
const configPath = getConfigFilePath()
|
||||
const userDataPath = app.getPath("userData")
|
||||
|
||||
|
||||
@@ -43,15 +43,29 @@ function loadSavedPort(): number | null {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Why the legacy port was not used at this launch: "EACCES" when the system
|
||||
* reserves it (Windows excludes port ranges for Hyper-V, which can change
|
||||
* on each boot), "EADDRINUSE" when another process holds it for now
|
||||
*/
|
||||
let legacyPortError: string | null = null
|
||||
|
||||
/**
|
||||
* 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.
|
||||
* origin, and the next launch goes back to it once it is free. Only the two
|
||||
* fixed ports count, and 13370 only when the system reserves the legacy
|
||||
* port: while another process holds it (such as the previous version still
|
||||
* quitting after an update), the next launch tries it again.
|
||||
*/
|
||||
export function saveServerPort(port: number): void {
|
||||
if (!app.isPackaged || loadSavedPort() !== null) {
|
||||
return
|
||||
}
|
||||
const fixed =
|
||||
port === PORT_CONFIG.legacyProduction ||
|
||||
(port === PORT_CONFIG.production && legacyPortError === "EACCES")
|
||||
if (!fixed) return
|
||||
try {
|
||||
writeFileSync(getSavedPortPath(), JSON.stringify({ port }), "utf-8")
|
||||
} catch (error) {
|
||||
@@ -60,23 +74,31 @@ export function saveServerPort(port: number): void {
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if a specific port is available
|
||||
* Try to listen on a port. Resolves to null when it is free, else to the
|
||||
* error code.
|
||||
*/
|
||||
export function isPortAvailable(port: number): Promise<boolean> {
|
||||
function portError(port: number): Promise<string | null> {
|
||||
return new Promise((resolve) => {
|
||||
const server = net.createServer()
|
||||
server.once("error", (err: NodeJS.ErrnoException) => {
|
||||
console.warn(`Port ${port} unavailable: ${err.code}`)
|
||||
resolve(false)
|
||||
resolve(err.code ?? "unknown")
|
||||
})
|
||||
server.once("listening", () => {
|
||||
server.close()
|
||||
resolve(true)
|
||||
resolve(null)
|
||||
})
|
||||
server.listen(port, "127.0.0.1")
|
||||
})
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if a specific port is available
|
||||
*/
|
||||
export async function isPortAvailable(port: number): Promise<boolean> {
|
||||
return (await portError(port)) === null
|
||||
}
|
||||
|
||||
/**
|
||||
* Find an available port
|
||||
* - In development: uses fixed port (6002)
|
||||
@@ -123,7 +145,8 @@ export async function findAvailablePort(reuseExisting = true): Promise<number> {
|
||||
// In production, try legacy port first to preserve existing users' localStorage
|
||||
if (!isDev) {
|
||||
const legacyPort = PORT_CONFIG.legacyProduction
|
||||
if (await isPortAvailable(legacyPort)) {
|
||||
legacyPortError = await portError(legacyPort)
|
||||
if (legacyPortError === null) {
|
||||
allocatedPort = legacyPort
|
||||
return legacyPort
|
||||
}
|
||||
|
||||
@@ -106,11 +106,13 @@ export function getAppUrl(): string | null {
|
||||
}
|
||||
|
||||
/**
|
||||
* Point the main window at a new app server URL
|
||||
* (the restarted server can come up on a different port)
|
||||
* Point the main window at the restarted app server (it can come up on a
|
||||
* different port). On the same port the page reloads, so it fetches the new
|
||||
* preset's server models instead of sending the old preset's choice.
|
||||
*/
|
||||
export function setAppUrl(url: string): void {
|
||||
if (url === appUrl) {
|
||||
mainWindow?.webContents.reload()
|
||||
return
|
||||
}
|
||||
appUrl = url
|
||||
|
||||
+36
-12
@@ -639,6 +639,8 @@ interface Endpoint {
|
||||
fetch?: typeof fetch
|
||||
authToken?: string // Anthropic Bearer auth
|
||||
resourceName?: string // Azure
|
||||
// baseURL comes from the settings or env, not the provider's default
|
||||
configuredBaseURL?: boolean
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -679,11 +681,10 @@ function createModel(
|
||||
switch (provider) {
|
||||
case "openai": {
|
||||
const openaiProvider = createOpenAI(opts)
|
||||
// A custom base URL is usually a proxy that only has Chat
|
||||
// Completions; the official endpoint uses the Responses API,
|
||||
// which returns reasoning for the o-series and gpt-5 or later
|
||||
return e.baseURL &&
|
||||
e.baseURL !== PROVIDER_INFO.openai.defaultBaseUrl
|
||||
// A configured base URL is usually a proxy that only has Chat
|
||||
// Completions; without one the Responses API is used, which
|
||||
// returns reasoning for the o-series and gpt-5 or later
|
||||
return e.configuredBaseURL
|
||||
? openaiProvider.chat(modelId)
|
||||
: openaiProvider(modelId)
|
||||
}
|
||||
@@ -899,6 +900,9 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
...(overrides?.awsSessionToken && {
|
||||
sessionToken: overrides.awsSessionToken,
|
||||
}),
|
||||
// Without an apiKey the SDK reads the server's
|
||||
// AWS_BEARER_TOKEN_BEDROCK, which wins over the keys
|
||||
apiKey: "",
|
||||
})
|
||||
: adminAccessKeyId && adminSecretAccessKey
|
||||
? createAmazonBedrock({
|
||||
@@ -955,7 +959,13 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
}
|
||||
|
||||
case "ollama": {
|
||||
const baseURL = overrides?.baseUrl || process.env.OLLAMA_BASE_URL
|
||||
// Like other providers, a user's key never goes to the server's
|
||||
// base URL; without a base URL it is an Ollama Cloud key
|
||||
const baseURL =
|
||||
overrides?.baseUrl ||
|
||||
(overrides?.apiKey
|
||||
? PROVIDER_INFO.ollama.defaultBaseUrl
|
||||
: process.env.OLLAMA_BASE_URL)
|
||||
// SECURITY: When client provides a custom base URL, only use
|
||||
// client-provided API key. Never fall back to server OLLAMA_API_KEY
|
||||
// to prevent leaking server credentials to user-controlled endpoints.
|
||||
@@ -1003,16 +1013,24 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
const publicDefault = defaultUrl?.startsWith("https://")
|
||||
? defaultUrl
|
||||
: undefined
|
||||
const baseURL = resolveBaseURL(
|
||||
const configuredBaseURL = resolveBaseURL(
|
||||
overrides?.apiKey,
|
||||
overrides?.baseUrl,
|
||||
resolveBaseUrlEnv(overrides, baseUrlEnv),
|
||||
SDK_KNOWS_ENDPOINT.has(provider) &&
|
||||
!(provider === "openai" && overrides?.apiKey)
|
||||
? undefined
|
||||
: publicDefault,
|
||||
)
|
||||
if (!baseURL && !SDK_KNOWS_ENDPOINT.has(provider)) {
|
||||
const baseURL =
|
||||
configuredBaseURL ||
|
||||
(SDK_KNOWS_ENDPOINT.has(provider) &&
|
||||
!(provider === "openai" && overrides?.apiKey)
|
||||
? undefined
|
||||
: publicDefault)
|
||||
// With a user's Azure key the SDK would read the server's
|
||||
// AZURE_RESOURCE_NAME
|
||||
if (
|
||||
!baseURL &&
|
||||
(!SDK_KNOWS_ENDPOINT.has(provider) ||
|
||||
(provider === "azure" && overrides?.apiKey))
|
||||
) {
|
||||
throw new Error(
|
||||
`${PROVIDER_INFO[provider].label} needs a base URL. Add it in the model settings.`,
|
||||
)
|
||||
@@ -1020,6 +1038,7 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
model = createModel(provider, modelId, {
|
||||
apiKey,
|
||||
baseURL,
|
||||
configuredBaseURL: !!configuredBaseURL,
|
||||
fetch: guardedFetch,
|
||||
// Bearer auth for Anthropic when there is no API key
|
||||
authToken:
|
||||
@@ -1038,6 +1057,11 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
return { model, providerOptions, modelId, provider }
|
||||
}
|
||||
|
||||
/** The provider of the server's own config: AI_PROVIDER, or the one with a key */
|
||||
export function getServerProvider(): ProviderName | null {
|
||||
return (process.env.AI_PROVIDER as ProviderName) || detectProvider()
|
||||
}
|
||||
|
||||
/**
|
||||
* 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
|
||||
|
||||
+12
-3
@@ -80,7 +80,8 @@ const GENERAL_TEXTS: Array<[RegExp, LLMErrorCode]> = [
|
||||
/invalid[_ ]api[_ ]key|incorrect api key|unauthorized/i,
|
||||
"invalid_api_key",
|
||||
],
|
||||
[/rate limit|too many requests/i, "rate_limited"],
|
||||
// "too many tokens": Bedrock's throttling
|
||||
[/rate limit|too many requests|too many tokens/i, "rate_limited"],
|
||||
[
|
||||
/Cannot connect to API|ECONNREFUSED|ENOTFOUND|ECONNRESET|ETIMEDOUT|fetch failed/i,
|
||||
"cannot_connect",
|
||||
@@ -115,8 +116,16 @@ function problemDetail(body: string): string | undefined {
|
||||
* it can name the server's account, role or internal hosts.
|
||||
*/
|
||||
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
|
||||
// The SDK passes an invalid tool call's error as a plain string. Other
|
||||
// strings come from providers (DeepSeek's SDK sends stream errors so).
|
||||
if (
|
||||
typeof error === "string" &&
|
||||
/^(Invalid input for tool|Model tried to call unavailable tool)/.test(
|
||||
error,
|
||||
)
|
||||
) {
|
||||
return error
|
||||
}
|
||||
if (isToolCallError(error)) return (error as Error).message
|
||||
const classified = classifyLLMError(error)
|
||||
if (hideDetails) {
|
||||
|
||||
@@ -77,11 +77,12 @@ async function getJson(
|
||||
|
||||
/**
|
||||
* 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.
|
||||
* Ollama without a key 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") {
|
||||
function listFallbackUrl(provider: ProviderName, apiKey?: string): string {
|
||||
if (provider === "ollama" && !apiKey) {
|
||||
return process.env.OLLAMA_BASE_URL || "http://127.0.0.1:11434/api"
|
||||
}
|
||||
const url = PROVIDER_INFO[provider].defaultBaseUrl
|
||||
@@ -98,7 +99,7 @@ export async function listProviderModels(
|
||||
{ apiKey, baseUrl }: { apiKey?: string; baseUrl?: string },
|
||||
fetchFn: typeof fetch = fetch,
|
||||
): Promise<ListedModel[]> {
|
||||
const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider))
|
||||
const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider, apiKey))
|
||||
const bearer: Record<string, string> = apiKey
|
||||
? { Authorization: `Bearer ${apiKey}` }
|
||||
: {}
|
||||
|
||||
+37
-36
@@ -224,6 +224,39 @@ async function main() {
|
||||
let configWatcher = null
|
||||
let restartPending = false
|
||||
|
||||
// Restart Next.js when the preset env vars really changed
|
||||
async function applyPresetChange() {
|
||||
if (restartPending) return
|
||||
const newContent = readPresetEnvFile()
|
||||
if (newContent === null || newContent === presetEnvContent) return
|
||||
|
||||
restartPending = true
|
||||
presetEnvContent = newContent
|
||||
console.log(
|
||||
"\n🔄 Preset configuration changed, restarting Next.js server...",
|
||||
)
|
||||
|
||||
// Kill current Next.js process
|
||||
killProcess(nextProcess)
|
||||
|
||||
// Wait a bit for process to die
|
||||
await new Promise((r) => setTimeout(r, 1000))
|
||||
|
||||
// Reload preset and restart
|
||||
nextProcess = startNextServer(loadPresetEnv(newContent))
|
||||
|
||||
try {
|
||||
await waitForServer(NEXT_URL)
|
||||
console.log("✅ Next.js server restarted with new configuration\n")
|
||||
} catch (err) {
|
||||
console.error("❌ Failed to restart Next.js:", err.message)
|
||||
}
|
||||
|
||||
restartPending = false
|
||||
// A change written during the restart was skipped above
|
||||
applyPresetChange()
|
||||
}
|
||||
|
||||
function setupConfigWatcher() {
|
||||
if (!existsSync(userDataPath)) {
|
||||
// Directory doesn't exist yet, check again later
|
||||
@@ -236,45 +269,13 @@ async function main() {
|
||||
configWatcher = watch(
|
||||
userDataPath,
|
||||
{ persistent: false },
|
||||
async (_eventType, filename) => {
|
||||
if (filename !== PRESET_ENV_FILE || restartPending) return
|
||||
|
||||
// Only restart when the preset env vars really changed
|
||||
const newContent = readPresetEnvFile()
|
||||
if (newContent === null || newContent === presetEnvContent)
|
||||
return
|
||||
|
||||
restartPending = true
|
||||
presetEnvContent = newContent
|
||||
console.log(
|
||||
"\n🔄 Preset configuration changed, restarting Next.js server...",
|
||||
)
|
||||
|
||||
// Kill current Next.js process
|
||||
killProcess(nextProcess)
|
||||
|
||||
// Wait a bit for process to die
|
||||
await new Promise((r) => setTimeout(r, 1000))
|
||||
|
||||
// Reload preset and restart
|
||||
nextProcess = startNextServer(loadPresetEnv(newContent))
|
||||
|
||||
try {
|
||||
await waitForServer(NEXT_URL)
|
||||
console.log(
|
||||
"✅ Next.js server restarted with new configuration\n",
|
||||
)
|
||||
} catch (err) {
|
||||
console.error(
|
||||
"❌ Failed to restart Next.js:",
|
||||
err.message,
|
||||
)
|
||||
}
|
||||
|
||||
restartPending = false
|
||||
(_eventType, filename) => {
|
||||
if (filename === PRESET_ENV_FILE) applyPresetChange()
|
||||
},
|
||||
)
|
||||
console.log("👀 Watching for preset configuration changes...")
|
||||
// Electron may have written its preset before the watch started
|
||||
applyPresetChange()
|
||||
} catch (_err) {
|
||||
// Directory might not be ready yet, try again later
|
||||
setTimeout(setupConfigWatcher, 5000)
|
||||
|
||||
@@ -36,6 +36,11 @@ vi.mock("@aws-sdk/credential-providers", () => ({
|
||||
fromNodeProviderChain: vi.fn(() => "node-chain"),
|
||||
}))
|
||||
|
||||
vi.mock("ollama-ai-provider-v2", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createOllama: vi.fn(() => mockProviderFn) }
|
||||
})
|
||||
|
||||
vi.mock("@openrouter/ai-sdk-provider", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createOpenRouter: vi.fn(() => mockProviderFn) }
|
||||
@@ -60,6 +65,8 @@ const ENV_KEYS = [
|
||||
"NEXT_AI_DRAWIO_DESKTOP",
|
||||
"SGLANG_API_KEY",
|
||||
"SGLANG_BASE_URL",
|
||||
"AZURE_RESOURCE_NAME",
|
||||
"OLLAMA_BASE_URL",
|
||||
]
|
||||
const savedEnv: Record<string, string | undefined> = {}
|
||||
|
||||
@@ -173,6 +180,8 @@ describe("Bedrock admin panel credentials", () => {
|
||||
region: "ap-northeast-1",
|
||||
accessKeyId: "client-id",
|
||||
secretAccessKey: "client-secret",
|
||||
// The SDK would otherwise use the server's AWS_BEARER_TOKEN_BEDROCK
|
||||
apiKey: "",
|
||||
})
|
||||
})
|
||||
|
||||
@@ -312,6 +321,42 @@ describe("whose keys a request uses", () => {
|
||||
expect(provider.chat).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("uses Chat Completions for any configured base URL", () => {
|
||||
// The settings form fills in the official URL for a new provider
|
||||
getAIModel({
|
||||
provider: "openai",
|
||||
apiKey: "user-key",
|
||||
baseUrl: "https://api.openai.com/v1",
|
||||
modelId: "gpt-5.5",
|
||||
})
|
||||
const provider = vi.mocked(createOpenAI).mock.results.at(-1)?.value
|
||||
expect(provider.chat).toHaveBeenCalledWith("gpt-5.5")
|
||||
})
|
||||
|
||||
it("sends a user's Ollama key to Ollama Cloud, not the server's Ollama", async () => {
|
||||
process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434/api"
|
||||
const { createOllama } = await import("ollama-ai-provider-v2")
|
||||
getAIModel({ provider: "ollama", apiKey: "user-key", modelId: "m" })
|
||||
expect(createOllama).toHaveBeenLastCalledWith(
|
||||
expect.objectContaining({ baseURL: "https://ollama.com/api" }),
|
||||
)
|
||||
// Without a key: the server's Ollama
|
||||
getAIModel({ provider: "ollama", modelId: "m" })
|
||||
expect(createOllama).toHaveBeenLastCalledWith(
|
||||
expect.objectContaining({
|
||||
baseURL: "http://ollama.internal:11434/api",
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
it("needs a base URL with a user's Azure key", () => {
|
||||
// The SDK would otherwise read the server's AZURE_RESOURCE_NAME
|
||||
process.env.AZURE_RESOURCE_NAME = "operator-resource"
|
||||
expect(() =>
|
||||
getAIModel({ provider: "azure", apiKey: "k", modelId: "gpt-4o" }),
|
||||
).toThrow(/base URL/)
|
||||
})
|
||||
|
||||
it("needs a base URL for SGLang instead of using 127.0.0.1", () => {
|
||||
expect(() =>
|
||||
getAIModel({ provider: "sglang", apiKey: "k", modelId: "m" }),
|
||||
|
||||
@@ -20,7 +20,15 @@ vi.mock("@/lib/dynamo-quota-manager", () => ({
|
||||
|
||||
import { POST as chat } from "@/app/api/chat/route"
|
||||
|
||||
const ENV = ["AI_PROVIDER", "AI_MODEL", "OPENAI_API_KEY"]
|
||||
const ENV = [
|
||||
"AI_PROVIDER",
|
||||
"AI_MODEL",
|
||||
"OPENAI_API_KEY",
|
||||
"OLLAMA_BASE_URL",
|
||||
"OLLAMA_API_KEY",
|
||||
"AI_GATEWAY_API_KEY",
|
||||
"ALLOW_PRIVATE_URLS",
|
||||
]
|
||||
const saved: Record<string, string | undefined> = {}
|
||||
|
||||
beforeEach(() => {
|
||||
@@ -88,4 +96,44 @@ describe("chat quota", () => {
|
||||
expect(res.status).not.toBe(429)
|
||||
expect(quota.checks).toBe(0)
|
||||
})
|
||||
|
||||
it("counts the server's keyless Ollama and EdgeOne", async () => {
|
||||
process.env.AI_PROVIDER = "ollama"
|
||||
process.env.AI_MODEL = "llama3.2"
|
||||
process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434/api"
|
||||
expect((await send({})).status).toBe(429)
|
||||
expect(
|
||||
(
|
||||
await send({
|
||||
"x-ai-provider": "edgeone",
|
||||
"x-ai-model": "@tx/deepseek-ai/deepseek-v3-0324",
|
||||
})
|
||||
).status,
|
||||
).toBe(429)
|
||||
expect(quota.checks).toBe(2)
|
||||
})
|
||||
|
||||
it("does not count Ollama on the user's own server", async () => {
|
||||
const res = await send({
|
||||
"x-ai-provider": "ollama",
|
||||
"x-ai-base-url": "https://ollama.example.com/api",
|
||||
"x-ai-model": "llama3.2",
|
||||
})
|
||||
expect(res.status).not.toBe(429)
|
||||
expect(quota.checks).toBe(0)
|
||||
})
|
||||
})
|
||||
|
||||
describe("server model allowlist", () => {
|
||||
it("runs AI_MODEL only on the server's AI_PROVIDER", async () => {
|
||||
// Another provider's server key must not run it
|
||||
process.env.AI_GATEWAY_API_KEY = "server-gateway-key"
|
||||
const res = await send({
|
||||
"x-ai-provider": "gateway",
|
||||
"x-ai-model": "gpt-5.5",
|
||||
})
|
||||
expect(res.status).toBe(400)
|
||||
expect(await res.text()).toMatch(/not available on this server/)
|
||||
expect(quota.checks).toBe(0)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
// @vitest-environment node
|
||||
import { existsSync, mkdtempSync, readdirSync, 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: { getPath: () => userData.dir },
|
||||
safeStorage: { isEncryptionAvailable: () => false },
|
||||
}))
|
||||
|
||||
// Make the next read of the presets file fail, like a file an antivirus
|
||||
// scanner holds on Windows
|
||||
const readFails = vi.hoisted(() => ({ next: false }))
|
||||
vi.mock("node:fs", async (importOriginal) => {
|
||||
const fs = await importOriginal<typeof import("node:fs")>()
|
||||
return {
|
||||
...fs,
|
||||
readFileSync: ((...args: Parameters<typeof fs.readFileSync>) => {
|
||||
if (readFails.next && String(args[0]).endsWith(".json")) {
|
||||
readFails.next = false
|
||||
throw Object.assign(new Error("EBUSY: resource busy"), {
|
||||
code: "EBUSY",
|
||||
})
|
||||
}
|
||||
return fs.readFileSync(...args)
|
||||
}) as typeof fs.readFileSync,
|
||||
}
|
||||
})
|
||||
|
||||
import { createPreset, loadPresets } from "@/electron/main/config-manager"
|
||||
|
||||
const presetsFile = () => join(userData.dir, "config-presets.json")
|
||||
|
||||
beforeEach(() => {
|
||||
userData.dir = mkdtempSync(join(tmpdir(), "config-manager-"))
|
||||
readFails.next = false
|
||||
})
|
||||
|
||||
describe("config presets file", () => {
|
||||
it("keeps a file it could not read for now", () => {
|
||||
createPreset({ name: "Mine", config: { AI_PROVIDER: "openai" } })
|
||||
|
||||
readFails.next = true
|
||||
expect(loadPresets().presets).toEqual([])
|
||||
expect(existsSync(presetsFile())).toBe(true)
|
||||
|
||||
// A save based on that empty read must not replace the presets
|
||||
readFails.next = true
|
||||
expect(() =>
|
||||
createPreset({ name: "New", config: { AI_PROVIDER: "openai" } }),
|
||||
).toThrow()
|
||||
expect(loadPresets().presets.map((p) => p.name)).toEqual(["Mine"])
|
||||
})
|
||||
|
||||
it("moves a file that is not JSON aside", () => {
|
||||
writeFileSync(presetsFile(), "{not json")
|
||||
expect(loadPresets().presets).toEqual([])
|
||||
expect(existsSync(presetsFile())).toBe(false)
|
||||
expect(
|
||||
readdirSync(userData.dir).some((f) =>
|
||||
f.startsWith("config-presets.json.corrupt-"),
|
||||
),
|
||||
).toBe(true)
|
||||
})
|
||||
})
|
||||
@@ -26,6 +26,21 @@ describe("EdgeOne chat completions function", () => {
|
||||
env: {},
|
||||
})
|
||||
expect(res.status).toBe(400)
|
||||
// A plain-text type that only mentions JSON needs no CORS preflight
|
||||
const disguised = await onRequest({
|
||||
request: request({
|
||||
"Content-Type": "text/plain; x=application/json",
|
||||
}),
|
||||
env: {},
|
||||
})
|
||||
expect(disguised.status).toBe(400)
|
||||
const withCharset = await onRequest({
|
||||
request: request({
|
||||
"Content-Type": "application/json; charset=utf-8",
|
||||
}),
|
||||
env: {},
|
||||
})
|
||||
expect(withCharset.status).toBe(200)
|
||||
})
|
||||
|
||||
it("checks the access code when ACCESS_CODE_LIST is set", async () => {
|
||||
|
||||
@@ -229,6 +229,27 @@ describe("streamErrorText", () => {
|
||||
message: "bad key",
|
||||
})
|
||||
})
|
||||
|
||||
it("classifies a provider error sent as plain text", () => {
|
||||
// DeepSeek's SDK sends errors in the stream as a string
|
||||
const text = "Insufficient Balance for account 42"
|
||||
expect(JSON.parse(streamErrorText(text))).toEqual({
|
||||
type: "provider",
|
||||
code: "insufficient_quota",
|
||||
message: text,
|
||||
})
|
||||
expect(JSON.parse(streamErrorText(text, true)).message).not.toMatch(
|
||||
/account 42/,
|
||||
)
|
||||
})
|
||||
|
||||
it("classifies Bedrock's throttling sent in the stream", () => {
|
||||
// Bedrock's ThrottlingException as a plain object, not an API error
|
||||
const throttled = {
|
||||
message: "Too many tokens, please wait before trying again.",
|
||||
}
|
||||
expect(JSON.parse(streamErrorText(throttled)).code).toBe("rate_limited")
|
||||
})
|
||||
})
|
||||
|
||||
describe("isToolCallError", () => {
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
// @vitest-environment node
|
||||
import { mkdtempSync, readFileSync, writeFileSync } from "node:fs"
|
||||
import { existsSync, mkdtempSync, readFileSync, writeFileSync } from "node:fs"
|
||||
import { tmpdir } from "node:os"
|
||||
import { join } from "node:path"
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||
@@ -9,29 +9,88 @@ vi.mock("electron", () => ({
|
||||
app: { isPackaged: true, getPath: () => userData.dir },
|
||||
}))
|
||||
|
||||
import { saveServerPort } from "@/electron/main/port-manager"
|
||||
// Ports that fail to listen, with the error code
|
||||
const busy = vi.hoisted(() => ({ ports: {} as Record<number, string> }))
|
||||
vi.mock("node:net", () => ({
|
||||
default: {
|
||||
createServer: () => {
|
||||
const handlers: Record<string, (arg?: unknown) => void> = {}
|
||||
const server = {
|
||||
once: (event: string, cb: (arg?: unknown) => void) => {
|
||||
handlers[event] = cb
|
||||
return server
|
||||
},
|
||||
listen: (port: number) => {
|
||||
const code = busy.ports[port]
|
||||
if (code) handlers.error?.({ code })
|
||||
else handlers.listening?.()
|
||||
},
|
||||
close: () => {},
|
||||
}
|
||||
return server
|
||||
},
|
||||
},
|
||||
}))
|
||||
|
||||
const savedPort = () =>
|
||||
JSON.parse(readFileSync(join(userData.dir, "server-port.json"), "utf-8"))
|
||||
.port
|
||||
import {
|
||||
findAvailablePort,
|
||||
resetAllocatedPort,
|
||||
saveServerPort,
|
||||
} from "@/electron/main/port-manager"
|
||||
|
||||
const portFile = () => join(userData.dir, "server-port.json")
|
||||
const savedPort = () => JSON.parse(readFileSync(portFile(), "utf-8")).port
|
||||
|
||||
/** One launch: pick a port, start on it, remember it if it should be */
|
||||
async function launch() {
|
||||
const port = await findAvailablePort(false)
|
||||
saveServerPort(port)
|
||||
return port
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
userData.dir = mkdtempSync(join(tmpdir(), "port-manager-"))
|
||||
busy.ports = {}
|
||||
resetAllocatedPort()
|
||||
})
|
||||
|
||||
describe("saveServerPort", () => {
|
||||
it("remembers the port of the first launch", () => {
|
||||
saveServerPort(13370)
|
||||
it("remembers the legacy port of the first launch", async () => {
|
||||
expect(await launch()).toBe(61337)
|
||||
expect(savedPort()).toBe(61337)
|
||||
})
|
||||
|
||||
it("remembers 13370 when the system reserves the legacy port", async () => {
|
||||
// Windows excludes port ranges for Hyper-V, which can change on
|
||||
// each boot
|
||||
busy.ports[61337] = "EACCES"
|
||||
expect(await launch()).toBe(13370)
|
||||
expect(savedPort()).toBe(13370)
|
||||
})
|
||||
|
||||
it("does not remember a port used while the legacy port was in use", async () => {
|
||||
// For example the previous version still quitting after an update:
|
||||
// the user's data is under 61337, so the next launch tries it again
|
||||
busy.ports[61337] = "EADDRINUSE"
|
||||
expect(await launch()).toBe(13370)
|
||||
expect(existsSync(portFile())).toBe(false)
|
||||
|
||||
busy.ports = {}
|
||||
expect(await launch()).toBe(61337)
|
||||
expect(savedPort()).toBe(61337)
|
||||
})
|
||||
|
||||
it("does not remember a fallback port", async () => {
|
||||
busy.ports[61337] = "EACCES"
|
||||
busy.ports[13370] = "EADDRINUSE"
|
||||
expect(await launch()).toBe(13371)
|
||||
expect(existsSync(portFile())).toBe(false)
|
||||
})
|
||||
|
||||
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 }),
|
||||
)
|
||||
writeFileSync(portFile(), JSON.stringify({ port: 61337 }))
|
||||
saveServerPort(13371)
|
||||
expect(savedPort()).toBe(61337)
|
||||
})
|
||||
|
||||
@@ -118,6 +118,18 @@ describe("listProviderModels", () => {
|
||||
])
|
||||
})
|
||||
|
||||
it("lists Ollama Cloud with a user's key, like chat", async () => {
|
||||
// The user's key must not go to the server's Ollama
|
||||
const { fn, calls } = answer({ models: [{ name: "gpt-oss:120b" }] })
|
||||
process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434"
|
||||
try {
|
||||
await listProviderModels("ollama", { apiKey: "user-key" }, fn)
|
||||
} finally {
|
||||
delete process.env.OLLAMA_BASE_URL
|
||||
}
|
||||
expect(calls[0].url).toBe("https://ollama.com/api/tags")
|
||||
})
|
||||
|
||||
it("does not use SGLang's local address as a default", async () => {
|
||||
const { fn, calls } = answer({ data: [] })
|
||||
await expect(
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
// @vitest-environment node
|
||||
import { streamText } from "ai"
|
||||
import { afterEach, describe, expect, it, vi } from "vitest"
|
||||
import { POST as testModel } from "@/app/api/admin/test-model/route"
|
||||
import { POST as validateModel } from "@/app/api/validate-model/route"
|
||||
import { getAIModel } from "@/lib/ai-providers"
|
||||
|
||||
// No saved admin providers
|
||||
vi.mock("@/lib/admin/settings", () => ({ loadSettings: () => ({}) }))
|
||||
|
||||
// Treat every URL as public so no test hits DNS
|
||||
vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({
|
||||
...(await importOriginal<typeof import("@/lib/ssrf-protection")>()),
|
||||
@@ -158,3 +162,37 @@ describe("chat requests to a client base URL", () => {
|
||||
expect(String(error)).toMatch(/Redirects are not allowed/)
|
||||
})
|
||||
})
|
||||
|
||||
describe("the admin panel's Test button", () => {
|
||||
it("works when access codes are set", async () => {
|
||||
// The admin password stands in for the visitor access code
|
||||
process.env.ACCESS_CODE_LIST = "visitor-code"
|
||||
process.env.ADMIN_PASSWORD = "admin-pw"
|
||||
try {
|
||||
streamReply({ role: "assistant", content: "OK" })
|
||||
const res = await testModel(
|
||||
new Request("http://localhost/api/admin/test-model", {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
"x-admin-password": "admin-pw",
|
||||
},
|
||||
body: JSON.stringify({
|
||||
provider: {
|
||||
id: "p1",
|
||||
provider: "glm",
|
||||
apiKey: "key",
|
||||
models: ["glm-5"],
|
||||
},
|
||||
modelId: "glm-5",
|
||||
}),
|
||||
}),
|
||||
)
|
||||
expect(res.status).toBe(200)
|
||||
expect((await res.json()).valid).toBe(true)
|
||||
} finally {
|
||||
delete process.env.ACCESS_CODE_LIST
|
||||
delete process.env.ADMIN_PASSWORD
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user