mirror of
https://github.com/DayuanJiang/next-ai-draw-io.git
synced 2026-10-09 03:07:46 +08:00
fix(server): keep users' keys at their own endpoints, and more fixes from the third review
- Bedrock: a user's AWS keys no longer go to an endpoint the server sets in AWS_ENDPOINT_URL_BEDROCK_RUNTIME / AWS_ENDPOINT_URL (read by the upgraded SDK), and admin panel keys win over AWS_BEARER_TOKEN_BEDROCK, as the Test button checks them. Checked with Bedrock. - Ollama: a server key without a base URL (admin panel, OLLAMA_API_KEY) goes to Ollama Cloud, as env.example says, instead of 127.0.0.1. - Quota: EdgeOne counts whatever key header comes along, keyless Ollama at a private address counts, and their provider texts stay in the log. - An EdgeOne server model (admin panel, ai-models.json) works: the route checked the raw provider header, which holds the name's slug. - parse-url ends downloads it does not read (too large, PDF, errors). - Desktop app: the port follows where the chats are (IndexedDB per origin) instead of a remembered port, which could hide them for good; a same-port restart tells the page to refetch the server models; a failed preset switch no longer undoes a newer choice; a presets file removed after a failed read can be saved again; .env values quoted from start to end keep their inner quotes, as dotenv reads them.
This commit is contained in:
+36
-24
@@ -164,25 +164,6 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
let baseUrl = req.headers.get("x-ai-base-url")
|
||||
const selectedModelId = req.headers.get("x-selected-model-id")
|
||||
|
||||
// For EdgeOne provider, construct full URL from request origin
|
||||
// because createOpenAI needs absolute URL, not relative path
|
||||
if (provider === "edgeone" && !baseUrl) {
|
||||
const origin = req.headers.get("origin") || new URL(req.url).origin
|
||||
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")
|
||||
|
||||
// Check if this is a server model with custom env var names
|
||||
let serverModelConfig: {
|
||||
apiKeyEnv?: string | string[]
|
||||
@@ -205,6 +186,29 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
}
|
||||
}
|
||||
|
||||
// A server model's provider comes from its config: for one set up in
|
||||
// the admin panel the header holds the provider name's slug
|
||||
const isEdgeOne = (serverModelConfig.provider || provider) === "edgeone"
|
||||
|
||||
// For EdgeOne provider, construct full URL from request origin
|
||||
// because createOpenAI needs absolute URL, not relative path
|
||||
if (isEdgeOne && !baseUrl) {
|
||||
const origin = req.headers.get("origin") || new URL(req.url).origin
|
||||
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")
|
||||
|
||||
const clientOverrides = {
|
||||
// Server model provider takes precedence over client header
|
||||
provider: serverModelConfig.provider || provider,
|
||||
@@ -223,7 +227,7 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
vertexApiKey: req.headers.get("x-vertex-api-key"),
|
||||
// Pass cookies for EdgeOne Pages authentication, and the access code,
|
||||
// which the EdgeOne function checks too
|
||||
...(provider === "edgeone" && {
|
||||
...(isEdgeOne && {
|
||||
headers: {
|
||||
...(cookieHeader && { cookie: cookieHeader }),
|
||||
"x-access-code": req.headers.get("x-access-code") || "",
|
||||
@@ -270,10 +274,16 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
// Quota is opt-in (DYNAMODB_QUOTA_TABLE) and counts what runs on the
|
||||
// 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.
|
||||
// EdgeOne never reads one; keyless Ollama at a private address is the
|
||||
// server's own network.
|
||||
const clientBaseUrl = normalizeBaseUrl(
|
||||
req.headers.get("x-ai-base-url") ?? "",
|
||||
)
|
||||
const onServerEndpoint =
|
||||
(resolvedProvider === "ollama" || resolvedProvider === "edgeone") &&
|
||||
!clientOverrides.apiKey &&
|
||||
!normalizeBaseUrl(req.headers.get("x-ai-base-url") ?? "")
|
||||
(resolvedProvider === "edgeone" && !clientBaseUrl) ||
|
||||
(resolvedProvider === "ollama" &&
|
||||
!clientOverrides.apiKey &&
|
||||
(!clientBaseUrl || (await isPrivateUrl(clientBaseUrl))))
|
||||
const countsQuota =
|
||||
isQuotaEnabled() &&
|
||||
(onServerCredentials || onServerEndpoint) &&
|
||||
@@ -726,7 +736,9 @@ Call this tool to get shape names and usage syntax for a specific library.`,
|
||||
|
||||
const response = result.toUIMessageStreamResponse({
|
||||
sendReasoning: true,
|
||||
onError: (error) => streamErrorText(error, onServerCredentials),
|
||||
// The provider's text can name the server's account or hosts
|
||||
onError: (error) =>
|
||||
streamErrorText(error, onServerCredentials || onServerEndpoint),
|
||||
messageMetadata: ({ part }) => {
|
||||
if (part.type === "finish") {
|
||||
const usage = (part as any).totalUsage
|
||||
|
||||
@@ -154,6 +154,9 @@ export async function POST(req: Request) {
|
||||
)
|
||||
} finally {
|
||||
clearTimeout(timeoutId)
|
||||
// Ends a download left unread (too large, PDF, error status);
|
||||
// a body already read is not affected
|
||||
controller.abort()
|
||||
}
|
||||
|
||||
// extractFromHtml throws (not returns null) on empty/non-HTML bodies,
|
||||
|
||||
Vendored
+5
@@ -70,6 +70,11 @@ declare global {
|
||||
>
|
||||
/** Set user's preferred locale */
|
||||
setUserLocale: (locale: string) => Promise<SetUserLocaleResult>
|
||||
/**
|
||||
* Call back after the server restarted on the same port (another
|
||||
* preset); returns a function that stops the calls
|
||||
*/
|
||||
onServerRestarted?: (callback: () => void) => () => void
|
||||
}
|
||||
|
||||
/** Settings window Electron API */
|
||||
|
||||
@@ -60,6 +60,13 @@ export async function switchPreset(
|
||||
console.error("Failed to restart server:", error)
|
||||
const reason = error instanceof Error ? error.message : String(error)
|
||||
|
||||
// Another preset was chosen meanwhile: its own restart follows
|
||||
if (getCurrentPresetId() !== id) {
|
||||
throw new Error(
|
||||
`The server could not be restarted.\n\nError: ${reason}`,
|
||||
)
|
||||
}
|
||||
|
||||
// Revert to previous preset on failure
|
||||
if (!previousPresetId || !applyPresetToEnv(previousPresetId)) {
|
||||
setCurrentPreset(null)
|
||||
|
||||
@@ -171,6 +171,8 @@ export function loadPresets(): ConfigPresetsFile {
|
||||
const configPath = getConfigFilePath()
|
||||
|
||||
if (!existsSync(configPath)) {
|
||||
// Nothing left that a save could overwrite
|
||||
presetsUnreadable = false
|
||||
return {
|
||||
version: 1,
|
||||
currentPresetId: null,
|
||||
|
||||
@@ -51,7 +51,11 @@ function loadEnvFromFile(filePath: string): void {
|
||||
const quote = value[0]
|
||||
const closingQuote =
|
||||
quote === '"' || quote === "'" ? value.indexOf(quote, 1) : -1
|
||||
if (closingQuote > 0) {
|
||||
if (value.length > 1 && value.endsWith(quote) && closingQuote > 0) {
|
||||
// Quoted from start to end: the quotes inside belong to the
|
||||
// value (JSON with an apostrophe), as dotenv reads it
|
||||
value = value.slice(1, -1)
|
||||
} else if (closingQuote > 0) {
|
||||
// Quoted value: keep what's inside the quotes and drop
|
||||
// anything after them (e.g. a comment)
|
||||
value = value.slice(1, closingQuote)
|
||||
|
||||
@@ -6,7 +6,6 @@ import {
|
||||
getAllocatedPort,
|
||||
getServerUrl,
|
||||
isPortAvailable,
|
||||
saveServerPort,
|
||||
} from "./port-manager"
|
||||
import { setAppUrl } from "./window-manager"
|
||||
|
||||
@@ -145,7 +144,6 @@ async function startServer(): Promise<string> {
|
||||
const url = getServerUrl()
|
||||
await waitForServer(url)
|
||||
console.log(`Next.js server started at ${url}`)
|
||||
saveServerPort(port)
|
||||
|
||||
return url
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { readFileSync, writeFileSync } from "node:fs"
|
||||
import { existsSync } from "node:fs"
|
||||
import net from "node:net"
|
||||
import path from "node:path"
|
||||
import { app } from "electron"
|
||||
@@ -26,84 +26,42 @@ const PORT_CONFIG = {
|
||||
let allocatedPort: number | null = null
|
||||
|
||||
/**
|
||||
* File that remembers the production port from the last launch, so the app
|
||||
* keeps the same origin (and its localStorage) instead of switching between
|
||||
* the legacy and new port depending on which one is free at startup
|
||||
* Whether chats are saved under http://127.0.0.1:<port>: Electron keeps
|
||||
* each origin's IndexedDB in its own folder
|
||||
*/
|
||||
function getSavedPortPath(): string {
|
||||
return path.join(app.getPath("userData"), "server-port.json")
|
||||
}
|
||||
|
||||
function loadSavedPort(): number | null {
|
||||
try {
|
||||
const { port } = JSON.parse(readFileSync(getSavedPortPath(), "utf-8"))
|
||||
return Number.isInteger(port) ? port : null
|
||||
} catch {
|
||||
return null
|
||||
}
|
||||
function hasStoredData(port: number): boolean {
|
||||
return existsSync(
|
||||
path.join(
|
||||
app.getPath("userData"),
|
||||
"IndexedDB",
|
||||
`http_127.0.0.1_${port}.indexeddb.leveldb`,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* 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
|
||||
* Check if a specific port is available
|
||||
*/
|
||||
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. 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) {
|
||||
console.error("Failed to save server port:", error)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Try to listen on a port. Resolves to null when it is free, else to the
|
||||
* error code.
|
||||
*/
|
||||
function portError(port: number): Promise<string | null> {
|
||||
export function isPortAvailable(port: number): Promise<boolean> {
|
||||
return new Promise((resolve) => {
|
||||
const server = net.createServer()
|
||||
server.once("error", (err: NodeJS.ErrnoException) => {
|
||||
console.warn(`Port ${port} unavailable: ${err.code}`)
|
||||
resolve(err.code ?? "unknown")
|
||||
resolve(false)
|
||||
})
|
||||
server.once("listening", () => {
|
||||
server.close()
|
||||
resolve(null)
|
||||
resolve(true)
|
||||
})
|
||||
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)
|
||||
* - In production: uses the port from the last launch, then the legacy
|
||||
* port (61337), then 13370, to preserve localStorage
|
||||
* - In production: uses the legacy port (61337), then 13370, to preserve
|
||||
* localStorage; 13370 first when only it has saved chats
|
||||
* - Falls back to sequential ports if preferred port is unavailable
|
||||
* - Last resort: lets the OS assign a port (port 0)
|
||||
*
|
||||
@@ -128,36 +86,22 @@ export async function findAvailablePort(reuseExisting = true): Promise<number> {
|
||||
allocatedPort = null
|
||||
}
|
||||
|
||||
// In production, use the port from the last launch first
|
||||
if (!isDev) {
|
||||
const savedPort = loadSavedPort()
|
||||
if (savedPort !== null) {
|
||||
if (await isPortAvailable(savedPort)) {
|
||||
allocatedPort = savedPort
|
||||
return savedPort
|
||||
}
|
||||
console.warn(
|
||||
`Port ${savedPort} from the last launch is unavailable. Data saved under it will not show on the new port.`,
|
||||
)
|
||||
// In production, try the legacy port first to preserve existing users'
|
||||
// data, unless only the new port has data: their app started on 13370
|
||||
// while Windows reserved 61337, and 61337 being free now would hide it
|
||||
const candidates = isDev
|
||||
? [preferredPort]
|
||||
: hasStoredData(PORT_CONFIG.production) &&
|
||||
!hasStoredData(PORT_CONFIG.legacyProduction)
|
||||
? [PORT_CONFIG.production, PORT_CONFIG.legacyProduction]
|
||||
: [PORT_CONFIG.legacyProduction, PORT_CONFIG.production]
|
||||
for (const port of candidates) {
|
||||
if (await isPortAvailable(port)) {
|
||||
allocatedPort = port
|
||||
return port
|
||||
}
|
||||
}
|
||||
|
||||
// In production, try legacy port first to preserve existing users' localStorage
|
||||
if (!isDev) {
|
||||
const legacyPort = PORT_CONFIG.legacyProduction
|
||||
legacyPortError = await portError(legacyPort)
|
||||
if (legacyPortError === null) {
|
||||
allocatedPort = legacyPort
|
||||
return legacyPort
|
||||
}
|
||||
}
|
||||
|
||||
// Try preferred port
|
||||
if (await isPortAvailable(preferredPort)) {
|
||||
allocatedPort = preferredPort
|
||||
return preferredPort
|
||||
}
|
||||
|
||||
console.warn(
|
||||
`Preferred port ${preferredPort} is in use, finding alternative...`,
|
||||
)
|
||||
|
||||
@@ -107,12 +107,13 @@ export function getAppUrl(): string | null {
|
||||
|
||||
/**
|
||||
* 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.
|
||||
* different port). On the same port the page fetches the new preset's
|
||||
* server models instead of sending the old preset's choice; it is not
|
||||
* reloaded, which would drop unsent attachments.
|
||||
*/
|
||||
export function setAppUrl(url: string): void {
|
||||
if (url === appUrl) {
|
||||
mainWindow?.webContents.reload()
|
||||
mainWindow?.webContents.send("server-restarted")
|
||||
return
|
||||
}
|
||||
appUrl = url
|
||||
|
||||
@@ -27,4 +27,13 @@ contextBridge.exposeInMainWorld("electronAPI", {
|
||||
getUserLocale: () => ipcRenderer.invoke("get-user-locale"),
|
||||
setUserLocale: (locale: string) =>
|
||||
ipcRenderer.invoke("set-user-locale", locale),
|
||||
|
||||
// The server restarted on the same port (another preset)
|
||||
onServerRestarted: (callback: () => void) => {
|
||||
const listener = () => callback()
|
||||
ipcRenderer.on("server-restarted", listener)
|
||||
return () => {
|
||||
ipcRenderer.removeListener("server-restarted", listener)
|
||||
}
|
||||
},
|
||||
})
|
||||
|
||||
+31
-7
@@ -617,6 +617,20 @@ function validateProviderCredentials(
|
||||
}
|
||||
}
|
||||
|
||||
/** AWS's Bedrock endpoint for a region, as the Bedrock SDK builds it */
|
||||
function bedrockRuntimeUrl(region: string): string {
|
||||
const suffix =
|
||||
[
|
||||
["cn-", "amazonaws.com.cn"],
|
||||
["us-iso-", "c2s.ic.gov"],
|
||||
["us-isob-", "sc2s.sgov.gov"],
|
||||
["eu-isoe-", "cloud.adc-e.uk"],
|
||||
["us-isof-", "csp.hci.ic.gov"],
|
||||
["eusc-", "amazonaws.eu"],
|
||||
].find(([prefix]) => region.startsWith(prefix))?.[1] ?? "amazonaws.com"
|
||||
return `https://bedrock-runtime.${region}.${suffix}`
|
||||
}
|
||||
|
||||
/**
|
||||
* Providers whose SDK has the official endpoint built in. The others are
|
||||
* OpenAI-compatible APIs (or Anthropic) that are called at
|
||||
@@ -903,12 +917,17 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
// Without an apiKey the SDK reads the server's
|
||||
// AWS_BEARER_TOKEN_BEDROCK, which wins over the keys
|
||||
apiKey: "",
|
||||
// Without a baseURL it reads the server's
|
||||
// AWS_ENDPOINT_URL_BEDROCK_RUNTIME / AWS_ENDPOINT_URL
|
||||
baseURL: bedrockRuntimeUrl(bedrockRegion),
|
||||
})
|
||||
: adminAccessKeyId && adminSecretAccessKey
|
||||
? createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
accessKeyId: adminAccessKeyId,
|
||||
secretAccessKey: adminSecretAccessKey,
|
||||
// The keys the admin panel's Test button checked
|
||||
apiKey: "",
|
||||
})
|
||||
: createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
@@ -959,19 +978,24 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig {
|
||||
}
|
||||
|
||||
case "ollama": {
|
||||
// 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.
|
||||
const apiKey = overrides?.baseUrl
|
||||
? overrides?.apiKey || undefined
|
||||
: resolveApiKey(overrides, "OLLAMA_API_KEY")
|
||||
// Like other providers, a user's key never goes to the server's
|
||||
// base URL. A key without a base URL is an Ollama Cloud key
|
||||
// (local Ollama has no keys); without either, the SDK's local
|
||||
// default.
|
||||
const baseURL =
|
||||
overrides?.baseUrl ||
|
||||
(overrides?.apiKey
|
||||
? PROVIDER_INFO.ollama.defaultBaseUrl
|
||||
: process.env.OLLAMA_BASE_URL ||
|
||||
(apiKey
|
||||
? PROVIDER_INFO.ollama.defaultBaseUrl
|
||||
: undefined))
|
||||
model = createOllama({
|
||||
...(baseURL && { baseURL }),
|
||||
...(apiKey && {
|
||||
|
||||
@@ -159,6 +159,8 @@ describe("Bedrock admin panel credentials", () => {
|
||||
region: "eu-west-1",
|
||||
accessKeyId: "panel-id",
|
||||
secretAccessKey: "panel-secret",
|
||||
// The keys the Test button checked, not AWS_BEARER_TOKEN_BEDROCK
|
||||
apiKey: "",
|
||||
})
|
||||
})
|
||||
|
||||
@@ -182,6 +184,8 @@ describe("Bedrock admin panel credentials", () => {
|
||||
secretAccessKey: "client-secret",
|
||||
// The SDK would otherwise use the server's AWS_BEARER_TOKEN_BEDROCK
|
||||
apiKey: "",
|
||||
// and the server's AWS_ENDPOINT_URL_BEDROCK_RUNTIME
|
||||
baseURL: "https://bedrock-runtime.ap-northeast-1.amazonaws.com",
|
||||
})
|
||||
})
|
||||
|
||||
@@ -349,6 +353,23 @@ describe("whose keys a request uses", () => {
|
||||
)
|
||||
})
|
||||
|
||||
it("sends the server's Ollama key without a base URL to Ollama Cloud", async () => {
|
||||
// An Ollama key is an Ollama Cloud key: local Ollama has none.
|
||||
// The admin panel saves it as OLLAMA_API_KEY, without a base URL.
|
||||
process.env.OLLAMA_API_KEY = "server-key"
|
||||
const { createOllama } = await import("ollama-ai-provider-v2")
|
||||
getAIModel({ provider: "ollama", modelId: "m" })
|
||||
expect(createOllama).toHaveBeenLastCalledWith(
|
||||
expect.objectContaining({ baseURL: "https://ollama.com/api" }),
|
||||
)
|
||||
// Without a key: the SDK's local default
|
||||
delete process.env.OLLAMA_API_KEY
|
||||
getAIModel({ provider: "ollama", modelId: "m" })
|
||||
expect(vi.mocked(createOllama).mock.lastCall?.[0]).not.toHaveProperty(
|
||||
"baseURL",
|
||||
)
|
||||
})
|
||||
|
||||
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"
|
||||
|
||||
@@ -462,9 +462,11 @@ describe("Ollama API key security", () => {
|
||||
|
||||
expect(createOllamaMock).toHaveBeenCalledTimes(1)
|
||||
const callArgs = createOllamaMock.mock.calls[0][0]
|
||||
expect(callArgs).not.toHaveProperty("baseURL")
|
||||
// As env.example says: without OLLAMA_BASE_URL, Ollama Cloud (the
|
||||
// SDK's default is the local server, which has no keys)
|
||||
expect(callArgs).toEqual(
|
||||
expect.objectContaining({
|
||||
baseURL: "https://ollama.com/api",
|
||||
headers: { Authorization: "Bearer server-key" },
|
||||
}),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
// @vitest-environment node
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||
|
||||
vi.mock("electron", () => ({
|
||||
app: {
|
||||
isPackaged: true,
|
||||
getName: () => "app",
|
||||
getVersion: () => "1",
|
||||
getLocale: () => "en",
|
||||
},
|
||||
BrowserWindow: { getFocusedWindow: () => null },
|
||||
dialog: {},
|
||||
Menu: { buildFromTemplate: () => ({}), setApplicationMenu: () => {} },
|
||||
shell: {},
|
||||
}))
|
||||
|
||||
// The saved current preset, and restarts that wait until the test ends them
|
||||
const state = vi.hoisted(() => ({
|
||||
current: "A" as string | null,
|
||||
restarts: [] as Array<{ resolve: () => void; reject: (e: Error) => void }>,
|
||||
}))
|
||||
vi.mock("@/electron/main/config-manager", () => ({
|
||||
applyPresetToEnv: (id: string) => {
|
||||
state.current = id
|
||||
return { AI_PROVIDER: id }
|
||||
},
|
||||
getAllPresets: () => [],
|
||||
getCurrentPresetId: () => state.current,
|
||||
setCurrentPreset: (id: string | null) => {
|
||||
state.current = id
|
||||
return true
|
||||
},
|
||||
}))
|
||||
vi.mock("@/electron/main/next-server", () => ({
|
||||
restartNextServer: () =>
|
||||
new Promise<void>((resolve, reject) =>
|
||||
state.restarts.push({ resolve, reject }),
|
||||
),
|
||||
}))
|
||||
vi.mock("@/electron/main/menu-i18n", () => ({
|
||||
getMenuTranslations: () => new Proxy({}, { get: () => "x" }),
|
||||
getPreferredLocale: () => "en",
|
||||
}))
|
||||
vi.mock("@/electron/main/settings-window", () => ({
|
||||
showSettingsWindow: () => {},
|
||||
}))
|
||||
|
||||
import { switchPreset } from "@/electron/main/app-menu"
|
||||
|
||||
beforeEach(() => {
|
||||
state.current = "A"
|
||||
state.restarts = []
|
||||
})
|
||||
|
||||
describe("switchPreset", () => {
|
||||
it("keeps a preset chosen while a failed switch was restarting", async () => {
|
||||
const toB = switchPreset("B").catch(() => {})
|
||||
const toC = switchPreset("C")
|
||||
// B's restart fails after the user already picked C
|
||||
state.restarts[0].reject(new Error("timed out"))
|
||||
await new Promise((r) => setTimeout(r, 0))
|
||||
for (const r of state.restarts.slice(1)) r.resolve()
|
||||
await toB
|
||||
await toC
|
||||
expect(state.current).toBe("C")
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,79 @@
|
||||
// @vitest-environment node
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
|
||||
vi.mock("@/lib/dynamo-quota-manager", () => ({
|
||||
isQuotaEnabled: () => false,
|
||||
checkAndIncrementRequest: async () => ({ allowed: true }),
|
||||
recordTokenUsage: async () => {},
|
||||
}))
|
||||
|
||||
import { POST as chat } from "@/app/api/chat/route"
|
||||
|
||||
const ENV = ["AI_MODELS_CONFIG", "AI_PROVIDER", "AI_MODEL"]
|
||||
const saved: Record<string, string | undefined> = {}
|
||||
const calls: Array<{ url: string; headers: Headers }> = []
|
||||
|
||||
beforeEach(() => {
|
||||
for (const k of ENV) saved[k] = process.env[k]
|
||||
delete process.env.AI_PROVIDER
|
||||
delete process.env.AI_MODEL
|
||||
calls.length = 0
|
||||
vi.stubGlobal(
|
||||
"fetch",
|
||||
vi.fn(async (url: string, init: RequestInit) => {
|
||||
calls.push({ url: String(url), headers: new Headers(init.headers) })
|
||||
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()
|
||||
})
|
||||
|
||||
describe("EdgeOne as a server model", () => {
|
||||
it("calls the site's Edge AI function with the cookies", async () => {
|
||||
// An admin panel or ai-models.json provider: the client sends the
|
||||
// provider name's slug, not "edgeone"
|
||||
process.env.AI_MODELS_CONFIG = JSON.stringify({
|
||||
providers: [
|
||||
{
|
||||
name: "Edge Pages",
|
||||
provider: "edgeone",
|
||||
models: ["@tx/deepseek-ai/deepseek-v3-0324"],
|
||||
},
|
||||
],
|
||||
})
|
||||
const res = await chat(
|
||||
new Request("http://localhost/api/chat", {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
"x-ai-provider": "edge-pages",
|
||||
"x-selected-model-id":
|
||||
"server:edge-pages:@tx/deepseek-ai/deepseek-v3-0324",
|
||||
cookie: "eo_token=t; eo_time=1",
|
||||
},
|
||||
body: JSON.stringify({
|
||||
messages: [
|
||||
{
|
||||
id: "u1",
|
||||
role: "user",
|
||||
parts: [{ type: "text", text: "Draw two boxes" }],
|
||||
},
|
||||
],
|
||||
xml: "",
|
||||
}),
|
||||
}),
|
||||
)
|
||||
await res.text()
|
||||
expect(calls[0]?.url).toBe(
|
||||
"http://localhost/api/edgeai/chat/completions",
|
||||
)
|
||||
expect(calls[0]?.headers.get("cookie")).toBe("eo_token=t; eo_time=1")
|
||||
})
|
||||
})
|
||||
@@ -18,6 +18,13 @@ vi.mock("@/lib/dynamo-quota-manager", () => ({
|
||||
recordTokenUsage: async () => {},
|
||||
}))
|
||||
|
||||
// No DNS in tests: only loopback addresses are private
|
||||
vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({
|
||||
...(await importOriginal<typeof import("@/lib/ssrf-protection")>()),
|
||||
isPrivateUrl: async (url: string) =>
|
||||
/^https?:\/\/(127\.0\.0\.1|localhost)\b/.test(url),
|
||||
}))
|
||||
|
||||
import { POST as chat } from "@/app/api/chat/route"
|
||||
|
||||
const ENV = [
|
||||
@@ -113,6 +120,26 @@ describe("chat quota", () => {
|
||||
expect(quota.checks).toBe(2)
|
||||
})
|
||||
|
||||
it("counts EdgeOne with a key header it never reads", async () => {
|
||||
const res = await send({
|
||||
"x-ai-provider": "edgeone",
|
||||
"x-ai-api-key": "ignored",
|
||||
"x-ai-model": "@tx/deepseek-ai/deepseek-v3-0324",
|
||||
})
|
||||
expect(res.status).toBe(429)
|
||||
expect(quota.checks).toBe(1)
|
||||
})
|
||||
|
||||
it("counts Ollama at a private address, the server's network", async () => {
|
||||
const res = await send({
|
||||
"x-ai-provider": "ollama",
|
||||
"x-ai-base-url": "http://127.0.0.1:11434/api",
|
||||
"x-ai-model": "llama3.2",
|
||||
})
|
||||
expect(res.status).toBe(429)
|
||||
expect(quota.checks).toBe(1)
|
||||
})
|
||||
|
||||
it("does not count Ollama on the user's own server", async () => {
|
||||
const res = await send({
|
||||
"x-ai-provider": "ollama",
|
||||
|
||||
@@ -1,5 +1,11 @@
|
||||
// @vitest-environment node
|
||||
import { existsSync, mkdtempSync, readdirSync, writeFileSync } from "node:fs"
|
||||
import {
|
||||
existsSync,
|
||||
mkdtempSync,
|
||||
readdirSync,
|
||||
rmSync,
|
||||
writeFileSync,
|
||||
} from "node:fs"
|
||||
import { tmpdir } from "node:os"
|
||||
import { join } from "node:path"
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||
@@ -54,6 +60,16 @@ describe("config presets file", () => {
|
||||
expect(loadPresets().presets.map((p) => p.name)).toEqual(["Mine"])
|
||||
})
|
||||
|
||||
it("saves again once a file it could not read is gone", () => {
|
||||
createPreset({ name: "Mine", config: { AI_PROVIDER: "openai" } })
|
||||
readFails.next = true
|
||||
loadPresets()
|
||||
// The user removes the file to start over
|
||||
rmSync(presetsFile())
|
||||
createPreset({ name: "New", config: { AI_PROVIDER: "openai" } })
|
||||
expect(loadPresets().presets.map((p) => p.name)).toEqual(["New"])
|
||||
})
|
||||
|
||||
it("moves a file that is not JSON aside", () => {
|
||||
writeFileSync(presetsFile(), "{not json")
|
||||
expect(loadPresets().presets).toEqual([])
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
// @vitest-environment node
|
||||
import { mkdtempSync, writeFileSync } from "node:fs"
|
||||
import { tmpdir } from "node:os"
|
||||
import { join } from "node:path"
|
||||
import { afterEach, describe, expect, it, vi } from "vitest"
|
||||
|
||||
const dir = vi.hoisted(() => ({ path: "" }))
|
||||
vi.mock("electron", () => ({
|
||||
app: {
|
||||
getPath: (name: string) =>
|
||||
name === "exe" ? `${dir.path}/app/exe` : dir.path,
|
||||
getAppPath: () => `${dir.path}/app`,
|
||||
},
|
||||
}))
|
||||
|
||||
import { loadEnvFile } from "@/electron/main/env-loader"
|
||||
|
||||
const KEYS = ["T_JSON", "T_COMMENT", "T_PLAIN", "T_DOUBLE"]
|
||||
afterEach(() => {
|
||||
for (const k of KEYS) delete process.env[k]
|
||||
})
|
||||
|
||||
describe("loadEnvFile", () => {
|
||||
it("reads quoted values like dotenv", () => {
|
||||
dir.path = mkdtempSync(join(tmpdir(), "env-loader-"))
|
||||
writeFileSync(
|
||||
join(dir.path, ".env"),
|
||||
[
|
||||
// An apostrophe inside a single-quoted JSON value
|
||||
`T_JSON='{"name":"Team's models"}'`,
|
||||
`T_COMMENT="value" # a comment`,
|
||||
"T_PLAIN=plain # a comment",
|
||||
`T_DOUBLE="say "hi""`,
|
||||
].join("\n"),
|
||||
)
|
||||
loadEnvFile()
|
||||
expect(process.env.T_JSON).toBe(`{"name":"Team's models"}`)
|
||||
expect(process.env.T_COMMENT).toBe("value")
|
||||
expect(process.env.T_PLAIN).toBe("plain")
|
||||
expect(process.env.T_DOUBLE).toBe(`say "hi"`)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,41 @@
|
||||
// @vitest-environment node
|
||||
import { afterEach, describe, expect, it, vi } from "vitest"
|
||||
import { POST as parseUrl } from "@/app/api/parse-url/route"
|
||||
|
||||
// Treat every URL as public so no test hits DNS
|
||||
vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({
|
||||
...(await importOriginal<typeof import("@/lib/ssrf-protection")>()),
|
||||
isPrivateUrl: async () => false,
|
||||
}))
|
||||
|
||||
afterEach(() => {
|
||||
vi.unstubAllGlobals()
|
||||
})
|
||||
|
||||
describe("POST /api/parse-url", () => {
|
||||
it("stops the download of a page announced as too large", async () => {
|
||||
let signal: AbortSignal | undefined
|
||||
vi.stubGlobal(
|
||||
"fetch",
|
||||
vi.fn(async (_url: string, init: RequestInit) => {
|
||||
signal = init.signal ?? undefined
|
||||
// A body that never ends unless the request is aborted
|
||||
return new Response(new ReadableStream(), {
|
||||
headers: {
|
||||
"content-type": "text/html",
|
||||
"content-length": String(50 * 1024 * 1024),
|
||||
},
|
||||
})
|
||||
}),
|
||||
)
|
||||
const res = await parseUrl(
|
||||
new Request("http://localhost/api/parse-url", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({ url: "https://example.com/huge" }),
|
||||
}),
|
||||
)
|
||||
expect(res.status).toBe(413)
|
||||
expect(signal?.aborted).toBe(true)
|
||||
})
|
||||
})
|
||||
@@ -1,5 +1,5 @@
|
||||
// @vitest-environment node
|
||||
import { existsSync, mkdtempSync, readFileSync, writeFileSync } from "node:fs"
|
||||
import { mkdirSync, mkdtempSync } from "node:fs"
|
||||
import { tmpdir } from "node:os"
|
||||
import { join } from "node:path"
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||
@@ -35,18 +35,19 @@ vi.mock("node:net", () => ({
|
||||
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
|
||||
}
|
||||
/** Chats saved under http://127.0.0.1:<port>, as Electron stores them */
|
||||
const storeData = (port: number) =>
|
||||
mkdirSync(
|
||||
join(
|
||||
userData.dir,
|
||||
"IndexedDB",
|
||||
`http_127.0.0.1_${port}.indexeddb.leveldb`,
|
||||
),
|
||||
{ recursive: true },
|
||||
)
|
||||
const launch = () => findAvailablePort(false)
|
||||
|
||||
beforeEach(() => {
|
||||
userData.dir = mkdtempSync(join(tmpdir(), "port-manager-"))
|
||||
@@ -54,44 +55,45 @@ beforeEach(() => {
|
||||
resetAllocatedPort()
|
||||
})
|
||||
|
||||
describe("saveServerPort", () => {
|
||||
it("remembers the legacy port of the first launch", async () => {
|
||||
describe("findAvailablePort", () => {
|
||||
it("uses the legacy port first, as main does", async () => {
|
||||
expect(await launch()).toBe(61337)
|
||||
storeData(61337)
|
||||
storeData(13370)
|
||||
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"
|
||||
it("uses 13370 when only it has the user's chats", async () => {
|
||||
// Windows reserved 61337 when they started using the app
|
||||
storeData(13370)
|
||||
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
|
||||
it("goes back to the port with the chats once it is free", async () => {
|
||||
// The previous version still quitting after an update
|
||||
storeData(61337)
|
||||
busy.ports[61337] = "EADDRINUSE"
|
||||
expect(await launch()).toBe(13370)
|
||||
expect(existsSync(portFile())).toBe(false)
|
||||
|
||||
storeData(13370)
|
||||
busy.ports = {}
|
||||
expect(await launch()).toBe(61337)
|
||||
expect(savedPort()).toBe(61337)
|
||||
})
|
||||
|
||||
it("does not remember a fallback port", async () => {
|
||||
it("does not hide the chats for good after one reserved launch", async () => {
|
||||
// Windows reserves port ranges per boot
|
||||
storeData(61337)
|
||||
busy.ports[61337] = "EACCES"
|
||||
busy.ports[13370] = "EADDRINUSE"
|
||||
expect(await launch()).toBe(13371)
|
||||
expect(existsSync(portFile())).toBe(false)
|
||||
expect(await launch()).toBe(13370)
|
||||
storeData(13370)
|
||||
busy.ports = {}
|
||||
expect(await launch()).toBe(61337)
|
||||
})
|
||||
|
||||
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(portFile(), JSON.stringify({ port: 61337 }))
|
||||
saveServerPort(13371)
|
||||
expect(savedPort()).toBe(61337)
|
||||
it("falls back to the next ports", async () => {
|
||||
storeData(13370)
|
||||
busy.ports[13370] = "EADDRINUSE"
|
||||
expect(await launch()).toBe(61337)
|
||||
busy.ports[61337] = "EACCES"
|
||||
expect(await launch()).toBe(13371)
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user