mirror of
https://github.com/DayuanJiang/next-ai-draw-io.git
synced 2026-10-04 00:37:48 +08:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
43924d5850 |
@@ -35,7 +35,6 @@ import { useDictionary } from "@/hooks/use-dictionary"
|
||||
import { formatMessage } from "@/lib/i18n/utils"
|
||||
import {
|
||||
FIXED_CRED_PROVIDERS,
|
||||
generateId,
|
||||
PROVIDER_INFO,
|
||||
type ProviderName,
|
||||
SUGGESTED_MODELS,
|
||||
@@ -226,7 +225,6 @@ function ProviderDetail({
|
||||
</Button>
|
||||
{suggestions.length > 0 && (
|
||||
<Select
|
||||
value=""
|
||||
disabled={disabled}
|
||||
onValueChange={(v) => addModel(v)}
|
||||
>
|
||||
@@ -392,14 +390,12 @@ function ProviderDetail({
|
||||
export function ModelsSection({
|
||||
providers,
|
||||
envProviders,
|
||||
envHasDefaultModel,
|
||||
disabled,
|
||||
password,
|
||||
onChange,
|
||||
}: {
|
||||
providers: AdminProvider[]
|
||||
envProviders: EnvProvider[]
|
||||
envHasDefaultModel: boolean
|
||||
disabled: boolean
|
||||
password: string
|
||||
onChange: (providers: AdminProvider[]) => void
|
||||
@@ -413,16 +409,10 @@ export function ModelsSection({
|
||||
|
||||
const addProvider = (provider: ProviderName) => {
|
||||
const newProvider: AdminProvider = {
|
||||
// generateId works over plain HTTP; crypto.randomUUID needs HTTPS
|
||||
id: generateId(),
|
||||
id: crypto.randomUUID(),
|
||||
provider,
|
||||
models: [],
|
||||
// Only the very first provider becomes the default, and only when
|
||||
// the env config has no default that it would replace on save
|
||||
isDefault:
|
||||
providers.length === 0 &&
|
||||
!envProviders.some((p) => p.isDefault) &&
|
||||
!envHasDefaultModel,
|
||||
isDefault: providers.length === 0,
|
||||
}
|
||||
onChange([...providers, newProvider])
|
||||
setSelectedId(newProvider.id)
|
||||
@@ -506,9 +496,7 @@ export function ModelsSection({
|
||||
))}
|
||||
</div>
|
||||
<div className="border-t p-2">
|
||||
{/* Always empty so picking the same type again still fires */}
|
||||
<Select
|
||||
value=""
|
||||
disabled={disabled}
|
||||
onValueChange={(v) => addProvider(v as ProviderName)}
|
||||
>
|
||||
|
||||
+14
-49
@@ -37,19 +37,6 @@ import { SettingField } from "./setting-field"
|
||||
|
||||
const NAV_GROUP_IDS = ["models", ...SETTING_GROUPS.map((g) => g.id)]
|
||||
|
||||
// For each toggleable group, whether any of its settings has a value (from
|
||||
// the settings file or the environment)
|
||||
function groupsWithValues(map: SettingsMap): Record<string, boolean> {
|
||||
const result: Record<string, boolean> = {}
|
||||
for (const group of SETTING_GROUPS) {
|
||||
if (!group.toggleable) continue
|
||||
result[group.id] = !!SETTINGS_BY_GROUP.get(group.id)?.some(
|
||||
(d) => map[d.key]?.source !== "default",
|
||||
)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
export default function AdminPage() {
|
||||
const dict = useDictionary()
|
||||
// Localized group title/description, keyed by group id
|
||||
@@ -75,8 +62,6 @@ export default function AdminPage() {
|
||||
// Models section state
|
||||
const [providers, setProviders] = useState<AdminProvider[]>([])
|
||||
const [envProviders, setEnvProviders] = useState<EnvProvider[]>([])
|
||||
// Whether .env itself sets AI_MODEL (a default the panel would override)
|
||||
const [envHasDefaultModel, setEnvHasDefaultModel] = useState(false)
|
||||
const [savedProviders, setSavedProviders] = useState<string>("[]")
|
||||
const providersDirty = JSON.stringify(providers) !== savedProviders
|
||||
|
||||
@@ -103,13 +88,15 @@ export default function AdminPage() {
|
||||
const map: SettingsMap = {}
|
||||
for (const s of data.settings) map[s.key] = s
|
||||
setSettings(map)
|
||||
// A group stays on while it still has values (e.g. from env vars
|
||||
// that saving can't remove); a user's explicit "on" for a group
|
||||
// with no values yet is kept across saves
|
||||
// Seed each toggle once from whether the group has configured
|
||||
// values; don't stomp a user's explicit toggle on later saves
|
||||
setEnabledGroups((prev) => {
|
||||
const next = groupsWithValues(map)
|
||||
for (const id of Object.keys(next)) {
|
||||
next[id] = next[id] || !!prev[id]
|
||||
const next = { ...prev }
|
||||
for (const group of SETTING_GROUPS) {
|
||||
if (!group.toggleable || group.id in next) continue
|
||||
next[group.id] = !!SETTINGS_BY_GROUP.get(group.id)?.some(
|
||||
(d) => map[d.key]?.source !== "default",
|
||||
)
|
||||
}
|
||||
return next
|
||||
})
|
||||
@@ -121,12 +108,10 @@ export default function AdminPage() {
|
||||
(data: {
|
||||
providers: AdminProvider[]
|
||||
envProviders?: EnvProvider[]
|
||||
envHasDefaultModel?: boolean
|
||||
}) => {
|
||||
setProviders(data.providers)
|
||||
setSavedProviders(JSON.stringify(data.providers))
|
||||
setEnvProviders(data.envProviders ?? [])
|
||||
setEnvHasDefaultModel(!!data.envHasDefaultModel)
|
||||
},
|
||||
[],
|
||||
)
|
||||
@@ -196,9 +181,8 @@ export default function AdminPage() {
|
||||
return () => observer.disconnect()
|
||||
}, [authedPassword])
|
||||
|
||||
// value undefined drops the pending change (back to the saved value)
|
||||
const handleChange = useCallback(
|
||||
(key: string, value: string | null | undefined) => {
|
||||
(key: string, value: string | null) => {
|
||||
setSaveMessage(null)
|
||||
setErrors((prev) => {
|
||||
if (!(key in prev)) return prev
|
||||
@@ -217,7 +201,7 @@ export default function AdminPage() {
|
||||
value === "" &&
|
||||
(!state || state.source !== "file") &&
|
||||
!isSecretValue(state?.value)
|
||||
if (value === undefined || isRevert || isNoop) {
|
||||
if (isRevert || isNoop) {
|
||||
const next = { ...prev }
|
||||
delete next[key]
|
||||
return next
|
||||
@@ -241,10 +225,9 @@ export default function AdminPage() {
|
||||
const next = { ...prev }
|
||||
for (const key of keys) {
|
||||
if (!enabled) {
|
||||
// Stage deletion of saved values; drop unsaved input
|
||||
if (settings[key]?.source === "default")
|
||||
delete next[key]
|
||||
else next[key] = null
|
||||
// Stage deletion only for values currently set
|
||||
if (settings[key]?.source !== "default")
|
||||
next[key] = null
|
||||
} else if (next[key] === null) {
|
||||
delete next[key]
|
||||
}
|
||||
@@ -464,7 +447,6 @@ export default function AdminPage() {
|
||||
<ModelsSection
|
||||
providers={providers}
|
||||
envProviders={envProviders}
|
||||
envHasDefaultModel={envHasDefaultModel}
|
||||
disabled={!writable || saving}
|
||||
password={authedPassword}
|
||||
onChange={(next) => {
|
||||
@@ -480,11 +462,6 @@ export default function AdminPage() {
|
||||
const defs = SETTINGS_BY_GROUP.get(group.id) ?? []
|
||||
const groupOff =
|
||||
group.toggleable && !enabledGroups[group.id]
|
||||
// Values from env vars can't be removed here, so the
|
||||
// group can't be turned off from the panel
|
||||
const envLocked = defs.some(
|
||||
(d) => settings[d.key]?.source === "env",
|
||||
)
|
||||
const fieldsDisabled = !writable || saving || !!groupOff
|
||||
const gt = groupText(group.id)
|
||||
const title = gt?.title ?? group.title
|
||||
@@ -503,11 +480,6 @@ export default function AdminPage() {
|
||||
</h2>
|
||||
{group.toggleable && (
|
||||
<label
|
||||
title={
|
||||
envLocked
|
||||
? dict.admin.sourceEnvTitle
|
||||
: undefined
|
||||
}
|
||||
className={cn(
|
||||
"flex cursor-pointer items-center gap-2 rounded-full border px-3 py-1.5 text-xs font-medium transition-colors motion-reduce:transition-none",
|
||||
enabledGroups[group.id]
|
||||
@@ -522,11 +494,7 @@ export default function AdminPage() {
|
||||
checked={
|
||||
!!enabledGroups[group.id]
|
||||
}
|
||||
disabled={
|
||||
!writable ||
|
||||
saving ||
|
||||
envLocked
|
||||
}
|
||||
disabled={!writable || saving}
|
||||
aria-label={formatMessage(
|
||||
dict.admin.enableGroup,
|
||||
{ group: title },
|
||||
@@ -611,9 +579,6 @@ export default function AdminPage() {
|
||||
setPending({})
|
||||
setErrors({})
|
||||
setProviders(JSON.parse(savedProviders))
|
||||
setEnabledGroups(
|
||||
groupsWithValues(settings),
|
||||
)
|
||||
}}
|
||||
>
|
||||
{dict.admin.discard}
|
||||
|
||||
@@ -73,10 +73,8 @@ export function SecretInput({
|
||||
}) {
|
||||
const dict = useDictionary()
|
||||
const [show, setShow] = useState(false)
|
||||
// The stored marker to revert to on empty. Refreshed whenever the parent
|
||||
// passes server state (a marker or nothing), e.g. after a save.
|
||||
const [original, setOriginal] = useState(value)
|
||||
if (typeof value !== "string" && value !== original) setOriginal(value)
|
||||
// The stored marker as it was at mount, to revert to on empty
|
||||
const [original] = useState(value)
|
||||
const hadStored = isSecretValue(original)
|
||||
const text = typeof value === "string" ? value : ""
|
||||
const placeholder = isSecretValue(value)
|
||||
@@ -148,8 +146,7 @@ export function SettingField({
|
||||
pendingValue: string | null | undefined
|
||||
error?: string
|
||||
disabled: boolean
|
||||
// undefined drops the pending change (back to the saved value)
|
||||
onChange: (value: string | null | undefined) => void
|
||||
onChange: (value: string | null) => void
|
||||
}) {
|
||||
const dict = useDictionary()
|
||||
const isDirty = pendingValue !== undefined
|
||||
@@ -229,18 +226,16 @@ export function SettingField({
|
||||
case "secret":
|
||||
control = (
|
||||
<div className="w-full max-w-md">
|
||||
{/* Clearing a saved secret reverts to it; the X button deletes */}
|
||||
<SecretInput
|
||||
id={inputId}
|
||||
keepOnEmpty={source === "file"}
|
||||
value={
|
||||
isDirty
|
||||
? (pendingValue ?? "")
|
||||
: (secretState ?? undefined)
|
||||
: (secretState ?? currentValue)
|
||||
}
|
||||
disabled={disabled}
|
||||
onChange={(v) =>
|
||||
onChange(typeof v === "string" ? v : undefined)
|
||||
onChange(typeof v === "string" ? v : "")
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
+17
-12
@@ -37,6 +37,7 @@ export default function Home() {
|
||||
)
|
||||
|
||||
const chatPanelRef = useRef<ImperativePanelHandle>(null)
|
||||
const isMobileRef = useRef(false)
|
||||
|
||||
// Load preferences from localStorage after mount
|
||||
useEffect(() => {
|
||||
@@ -47,9 +48,7 @@ export default function Home() {
|
||||
const currentLocale = pathParts[0]
|
||||
if (currentLocale !== savedLocale) {
|
||||
pathParts[0] = savedLocale
|
||||
// Keep the query (e.g. ?session=) and hash
|
||||
const { search, hash } = window.location
|
||||
router.replace(`/${pathParts.join("/")}${search}${hash}`)
|
||||
router.replace(`/${pathParts.join("/")}`)
|
||||
return // Wait for redirect
|
||||
}
|
||||
}
|
||||
@@ -107,17 +106,27 @@ export default function Home() {
|
||||
resetDrawioReady()
|
||||
}
|
||||
|
||||
// Check mobile. The draw.io iframe is not remounted when crossing the
|
||||
// breakpoint (only the chat panel is), so its ready state stays as is.
|
||||
// Check mobile - reset draw.io before crossing breakpoint
|
||||
const isInitialRenderRef = useRef(true)
|
||||
useEffect(() => {
|
||||
const checkMobile = () => {
|
||||
setIsMobile(window.innerWidth < 768)
|
||||
const newIsMobile = window.innerWidth < 768
|
||||
if (
|
||||
!isInitialRenderRef.current &&
|
||||
newIsMobile !== isMobileRef.current
|
||||
) {
|
||||
setIsDrawioReady(false)
|
||||
resetDrawioReady()
|
||||
}
|
||||
isMobileRef.current = newIsMobile
|
||||
isInitialRenderRef.current = false
|
||||
setIsMobile(newIsMobile)
|
||||
}
|
||||
|
||||
checkMobile()
|
||||
window.addEventListener("resize", checkMobile)
|
||||
return () => window.removeEventListener("resize", checkMobile)
|
||||
}, [])
|
||||
}, [resetDrawioReady])
|
||||
|
||||
const toggleChatPanel = () => {
|
||||
const panel = chatPanelRef.current
|
||||
@@ -184,11 +193,7 @@ export default function Home() {
|
||||
noExitBtn: true,
|
||||
dark:
|
||||
darkMode || drawioUi === "dark",
|
||||
// draw.io names Traditional Chinese "zh-tw"
|
||||
lang:
|
||||
currentLang === "zh-Hant"
|
||||
? "zh-tw"
|
||||
: currentLang,
|
||||
lang: currentLang,
|
||||
// Enable offline mode in Electron to disable external service calls
|
||||
...(isElectron && {
|
||||
offline: true,
|
||||
|
||||
@@ -7,11 +7,7 @@ import {
|
||||
mergeSecrets,
|
||||
validateAdminProviders,
|
||||
} from "@/lib/admin/providers"
|
||||
import {
|
||||
getEnvFallback,
|
||||
isSettingsWritable,
|
||||
saveSettings,
|
||||
} from "@/lib/admin/settings"
|
||||
import { isSettingsWritable, saveSettings } from "@/lib/admin/settings"
|
||||
import { loadEnvServerModelsConfig } from "@/lib/server-model-config"
|
||||
|
||||
export const runtime = "nodejs"
|
||||
@@ -37,9 +33,6 @@ async function payload() {
|
||||
models: p.models,
|
||||
isDefault: !!p.default && !adminHasDefault,
|
||||
})) ?? [],
|
||||
// Whether .env sets a default model. getEnvFallback skips the value
|
||||
// the panel overlays onto process.env, so a panel default doesn't count.
|
||||
envHasDefaultModel: !!getEnvFallback("AI_MODEL"),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+96
-87
@@ -12,17 +12,13 @@ import fs from "fs/promises"
|
||||
import { jsonrepair } from "jsonrepair"
|
||||
import path from "path"
|
||||
import { z } from "zod"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
import {
|
||||
getAIModel,
|
||||
SINGLE_SYSTEM_PROVIDERS,
|
||||
supportsPromptCaching,
|
||||
usesServerCredentials,
|
||||
} from "@/lib/ai-providers"
|
||||
import { findCachedResponse } from "@/lib/cached-responses"
|
||||
import {
|
||||
dropInvalidToolCalls,
|
||||
fixToolInputJson,
|
||||
isMinimalDiagram,
|
||||
replaceHistoricalToolInputs,
|
||||
validateFileParts,
|
||||
@@ -33,7 +29,6 @@ import {
|
||||
recordTokenUsage,
|
||||
} from "@/lib/dynamo-quota-manager"
|
||||
import {
|
||||
endTrace,
|
||||
getTelemetryConfig,
|
||||
setTraceInput,
|
||||
setTraceOutput,
|
||||
@@ -43,11 +38,7 @@ import {
|
||||
resolveMaxOutputTokens,
|
||||
withOutputTokenLimitFallback,
|
||||
} from "@/lib/output-token-limit"
|
||||
import {
|
||||
type FlattenedServerModel,
|
||||
findServerModelById,
|
||||
} from "@/lib/server-model-config"
|
||||
import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
|
||||
import { findServerModelById } from "@/lib/server-model-config"
|
||||
import { getSystemPrompt } from "@/lib/system-prompts"
|
||||
import { getUserIdFromRequest } from "@/lib/user-id"
|
||||
|
||||
@@ -85,14 +76,24 @@ function createCachedStreamResponse(xml: string): Response {
|
||||
return createUIMessageStreamResponse({ stream })
|
||||
}
|
||||
|
||||
// Responses streamed from the model, whose trace streamText's callbacks end
|
||||
const modelStreamResponses = new WeakSet<Response>()
|
||||
|
||||
// Inner handler function
|
||||
async function handleChatRequest(req: Request): Promise<Response> {
|
||||
// Check for access code
|
||||
const accessDenied = checkAccessCode(req)
|
||||
if (accessDenied) return accessDenied
|
||||
const accessCodes =
|
||||
process.env.ACCESS_CODE_LIST?.split(",")
|
||||
.map((code) => code.trim())
|
||||
.filter(Boolean) || []
|
||||
if (accessCodes.length > 0) {
|
||||
const accessCodeHeader = req.headers.get("x-access-code")
|
||||
if (!accessCodeHeader || !accessCodes.includes(accessCodeHeader)) {
|
||||
return Response.json(
|
||||
{
|
||||
error: "Invalid or missing access code. Please configure it in Settings.",
|
||||
},
|
||||
{ status: 401 },
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
const body = await req.json()
|
||||
const { messages, xml, previousXml, sessionId } = body
|
||||
@@ -191,15 +192,6 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
baseUrl = `${origin}/api/edgeai`
|
||||
}
|
||||
|
||||
// Same rule as validate-model: with ALLOW_PRIVATE_URLS=false a request may
|
||||
// not point the server at a private or internal address
|
||||
if (baseUrl && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) {
|
||||
return Response.json(
|
||||
{ error: "Private or internal base URLs are not allowed." },
|
||||
{ status: 400 },
|
||||
)
|
||||
}
|
||||
|
||||
// Get cookie header for EdgeOne authentication (eo_token, eo_time)
|
||||
const cookieHeader = req.headers.get("cookie")
|
||||
|
||||
@@ -209,9 +201,8 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
baseUrlEnv?: string
|
||||
provider?: string
|
||||
} = {}
|
||||
let serverModel: FlattenedServerModel | null = null
|
||||
if (selectedModelId?.startsWith("server:")) {
|
||||
serverModel = await findServerModelById(selectedModelId)
|
||||
const serverModel = await findServerModelById(selectedModelId)
|
||||
console.log(
|
||||
`[Server Model Lookup] ID: ${selectedModelId}, Found: ${!!serverModel}, Provider: ${serverModel?.provider}`,
|
||||
)
|
||||
@@ -230,8 +221,7 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
provider: serverModelConfig.provider || provider,
|
||||
baseUrl,
|
||||
apiKey: req.headers.get("x-ai-api-key"),
|
||||
// A server model runs the model it was configured with, whatever the header says
|
||||
modelId: serverModel?.modelId || req.headers.get("x-ai-model"),
|
||||
modelId: req.headers.get("x-ai-model"),
|
||||
// AWS Bedrock credentials
|
||||
awsAccessKeyId: req.headers.get("x-aws-access-key-id"),
|
||||
awsSecretAccessKey: req.headers.get("x-aws-secret-access-key"),
|
||||
@@ -241,14 +231,11 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
...serverModelConfig,
|
||||
// Vertex AI credentials (Express Mode)
|
||||
vertexApiKey: req.headers.get("x-vertex-api-key"),
|
||||
// Pass cookies for EdgeOne Pages authentication, and the access code,
|
||||
// which the EdgeOne function checks too
|
||||
...(provider === "edgeone" && {
|
||||
headers: {
|
||||
...(cookieHeader && { cookie: cookieHeader }),
|
||||
"x-access-code": req.headers.get("x-access-code") || "",
|
||||
},
|
||||
}),
|
||||
// Pass cookies for EdgeOne Pages authentication
|
||||
...(provider === "edgeone" &&
|
||||
cookieHeader && {
|
||||
headers: { cookie: cookieHeader },
|
||||
}),
|
||||
}
|
||||
|
||||
// Read minimal style preference from header
|
||||
@@ -267,32 +254,12 @@ async function handleChatRequest(req: Request): Promise<Response> {
|
||||
provider: resolvedProvider,
|
||||
} = getAIModel(clientOverrides)
|
||||
|
||||
// On the server's own keys, only run models the server offers: a server
|
||||
// model picked by id (its model name is fixed above) or one in AI_MODEL.
|
||||
// With their own key, users can run any model.
|
||||
const onServerCredentials = usesServerCredentials(
|
||||
resolvedProvider,
|
||||
clientOverrides,
|
||||
)
|
||||
const envModels =
|
||||
process.env.AI_MODEL?.split(",").map((m) => m.trim()) || []
|
||||
if (onServerCredentials && !serverModel && !envModels.includes(modelId)) {
|
||||
return Response.json(
|
||||
{
|
||||
error: `Model "${modelId}" is not available on this server. Add your own API key in Settings to use it.`,
|
||||
},
|
||||
{ status: 400 },
|
||||
)
|
||||
}
|
||||
|
||||
// Retry with a smaller budget if the provider rejects the requested one
|
||||
const model = withOutputTokenLimitFallback(baseModel)
|
||||
|
||||
// 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
|
||||
// User setting wins over server env, so desktop users can raise it themselves
|
||||
const maxOutputTokens = resolveMaxOutputTokens(
|
||||
req.headers.get("x-max-output-tokens"),
|
||||
onServerCredentials,
|
||||
)
|
||||
console.log(`[maxOutputTokens] ${maxOutputTokens}`)
|
||||
|
||||
@@ -373,9 +340,32 @@ ${userInputText}
|
||||
)
|
||||
|
||||
// Filter out tool-calls with invalid inputs (from failed repair or interrupted streaming)
|
||||
// and their results. Bedrock API rejects messages where toolUse.input is not a valid
|
||||
// JSON object, and every provider rejects a tool result whose call is gone.
|
||||
enhancedMessages = dropInvalidToolCalls(enhancedMessages)
|
||||
// Bedrock API rejects messages where toolUse.input is not a valid JSON object
|
||||
enhancedMessages = enhancedMessages
|
||||
.map((msg: any) => {
|
||||
if (msg.role !== "assistant" || !Array.isArray(msg.content)) {
|
||||
return msg
|
||||
}
|
||||
const filteredContent = msg.content.filter((part: any) => {
|
||||
if (part.type === "tool-call") {
|
||||
// Check if input is a valid object (not null, undefined, or empty)
|
||||
if (
|
||||
!part.input ||
|
||||
typeof part.input !== "object" ||
|
||||
Object.keys(part.input).length === 0
|
||||
) {
|
||||
console.warn(
|
||||
`[route.ts] Filtering out tool-call with invalid input:`,
|
||||
{ toolName: part.toolName, input: part.input },
|
||||
)
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
return { ...msg, content: filteredContent }
|
||||
})
|
||||
.filter((msg: any) => msg.content && msg.content.length > 0)
|
||||
|
||||
// DEBUG: Log modelMessages structure (what's being sent to AI)
|
||||
console.log("[route.ts] Model messages count:", enhancedMessages.length)
|
||||
@@ -420,7 +410,7 @@ ${userInputText}
|
||||
contentParts.push({
|
||||
type: "image",
|
||||
image: filePart.url,
|
||||
mediaType: filePart.mediaType,
|
||||
mimeType: filePart.mediaType,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -481,7 +471,7 @@ ${previousXml}
|
||||
${xml || ""}
|
||||
"""
|
||||
|
||||
IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed.`
|
||||
IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed. When using edit_diagram, COPY search patterns exactly from the CURRENT XML - attribute order matters!`
|
||||
|
||||
const systemMessages = isSingleSystemProvider
|
||||
? [
|
||||
@@ -538,11 +528,23 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
|
||||
error.name === "AI_InvalidToolInputError"
|
||||
) {
|
||||
try {
|
||||
// Pre-process to fix common LLM JSON errors that jsonrepair can't handle,
|
||||
// then use jsonrepair to fix truncated JSON
|
||||
const repairedInput = jsonrepair(
|
||||
fixToolInputJson(toolCall.input),
|
||||
)
|
||||
// Pre-process to fix common LLM JSON errors that jsonrepair can't handle
|
||||
let inputToRepair = toolCall.input
|
||||
if (typeof inputToRepair === "string") {
|
||||
// Fix `:=` instead of `: ` (LLM sometimes generates this)
|
||||
inputToRepair = inputToRepair.replace(/:=/g, ": ")
|
||||
// Fix `= "` instead of `: "`
|
||||
inputToRepair = inputToRepair.replace(/=\s*"/g, ': "')
|
||||
// Fix inconsistent quote escaping in XML attributes within JSON strings
|
||||
// Pattern: attribute="value\" where opening quote is unescaped but closing is escaped
|
||||
// Example: y="-20\" should be y=\"-20\"
|
||||
inputToRepair = inputToRepair.replace(
|
||||
/(\w+)="([^"]*?)\\"/g,
|
||||
'$1=\\"$2\\"',
|
||||
)
|
||||
}
|
||||
// Use jsonrepair to fix truncated JSON
|
||||
const repairedInput = jsonrepair(inputToRepair)
|
||||
console.log(
|
||||
`[repairToolCall] Repaired truncated JSON for tool: ${toolCall.toolName}`,
|
||||
)
|
||||
@@ -552,8 +554,26 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
|
||||
`[repairToolCall] Failed to repair JSON for tool: ${toolCall.toolName}`,
|
||||
repairError,
|
||||
)
|
||||
// Keep the original error, so the model and the client see why
|
||||
// the input was rejected and the model can retry the call
|
||||
// Return a placeholder input to avoid API errors in multi-step
|
||||
// The tool will fail gracefully on client side
|
||||
if (toolCall.toolName === "edit_diagram") {
|
||||
return {
|
||||
...toolCall,
|
||||
input: {
|
||||
operations: [],
|
||||
_error: "JSON repair failed - no operations to apply",
|
||||
},
|
||||
}
|
||||
}
|
||||
if (toolCall.toolName === "display_diagram") {
|
||||
return {
|
||||
...toolCall,
|
||||
input: {
|
||||
xml: "",
|
||||
_error: "JSON repair failed - empty diagram",
|
||||
},
|
||||
}
|
||||
}
|
||||
return null
|
||||
}
|
||||
}
|
||||
@@ -576,7 +596,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
|
||||
// Include all 4 token types: input, output, cache read, cache write
|
||||
if (
|
||||
isQuotaEnabled() &&
|
||||
!hasOwnApiKey &&
|
||||
@@ -585,16 +605,12 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
|
||||
) {
|
||||
const totalTokens =
|
||||
(totalUsage.inputTokens || 0) +
|
||||
(totalUsage.outputTokens || 0)
|
||||
(totalUsage.outputTokens || 0) +
|
||||
(totalUsage.cachedInputTokens || 0) +
|
||||
(totalUsage.inputTokenDetails?.cacheWriteTokens || 0)
|
||||
recordTokenUsage(userId, totalTokens)
|
||||
}
|
||||
},
|
||||
// onFinish is skipped when the stream fails or is aborted, so end the trace here
|
||||
onError: ({ error }) => {
|
||||
console.error(error) // what AI SDK does without an onError
|
||||
endTrace()
|
||||
},
|
||||
onAbort: () => endTrace(),
|
||||
tools: {
|
||||
// Client-side tool that will be executed on the client
|
||||
display_diagram: {
|
||||
@@ -766,7 +782,7 @@ Call this tool to get shape names and usage syntax for a specific library.`,
|
||||
}),
|
||||
})
|
||||
|
||||
const response = result.toUIMessageStreamResponse({
|
||||
return result.toUIMessageStreamResponse({
|
||||
sendReasoning: true,
|
||||
messageMetadata: ({ part }) => {
|
||||
if (part.type === "finish") {
|
||||
@@ -780,8 +796,6 @@ Call this tool to get shape names and usage syntax for a specific library.`,
|
||||
return undefined
|
||||
},
|
||||
})
|
||||
modelStreamResponses.add(response)
|
||||
return response
|
||||
}
|
||||
|
||||
// Helper to categorize errors and return appropriate response
|
||||
@@ -848,16 +862,11 @@ function handleError(error: unknown): Response {
|
||||
|
||||
// Wrap handler with error handling
|
||||
async function safeHandler(req: Request): Promise<Response> {
|
||||
let response: Response
|
||||
try {
|
||||
response = await handleChatRequest(req)
|
||||
return await handleChatRequest(req)
|
||||
} catch (error) {
|
||||
response = handleError(error)
|
||||
return handleError(error)
|
||||
}
|
||||
// Early returns, cache hits and errors never reach streamText's callbacks,
|
||||
// so their Langfuse trace has to be ended here
|
||||
if (!modelStreamResponses.has(response)) endTrace()
|
||||
return response
|
||||
}
|
||||
|
||||
// Wrap with Langfuse observe (if configured)
|
||||
|
||||
@@ -1,11 +1,9 @@
|
||||
import { extractFromHtml } from "@extractus/article-extractor"
|
||||
import { NextResponse } from "next/server"
|
||||
import TurndownService from "turndown"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
import { isPrivateUrl } from "@/lib/ssrf-protection"
|
||||
|
||||
const MAX_CONTENT_LENGTH = 150000 // Match PDF limit
|
||||
const MAX_RESPONSE_BYTES = 5 * 1024 * 1024
|
||||
const EXTRACT_TIMEOUT_MS = 15000
|
||||
const USER_AGENT = "Mozilla/5.0 (compatible; NextAIDrawio/1.0)"
|
||||
|
||||
@@ -34,36 +32,7 @@ function detectCharset(
|
||||
}
|
||||
}
|
||||
|
||||
// Read the response body, giving up once it passes MAX_RESPONSE_BYTES so a
|
||||
// huge download can't exhaust server memory. Returns null when too large.
|
||||
async function readLimitedBody(
|
||||
response: Response,
|
||||
): Promise<ArrayBuffer | null> {
|
||||
if (Number(response.headers.get("content-length")) > MAX_RESPONSE_BYTES) {
|
||||
return null
|
||||
}
|
||||
if (!response.body) return new ArrayBuffer(0)
|
||||
|
||||
const reader = response.body.getReader()
|
||||
const chunks: Uint8Array[] = []
|
||||
let total = 0
|
||||
while (true) {
|
||||
const { done, value } = await reader.read()
|
||||
if (done) break
|
||||
total += value.byteLength
|
||||
if (total > MAX_RESPONSE_BYTES) {
|
||||
await reader.cancel()
|
||||
return null
|
||||
}
|
||||
chunks.push(value)
|
||||
}
|
||||
return new Blob(chunks as BlobPart[]).arrayBuffer()
|
||||
}
|
||||
|
||||
export async function POST(req: Request) {
|
||||
const accessError = checkAccessCode(req)
|
||||
if (accessError) return accessError
|
||||
|
||||
try {
|
||||
const { url } = await req.json()
|
||||
|
||||
@@ -128,15 +97,7 @@ export async function POST(req: Request) {
|
||||
)
|
||||
}
|
||||
|
||||
const buffer = await readLimitedBody(response)
|
||||
if (!buffer) {
|
||||
return NextResponse.json(
|
||||
{
|
||||
error: `Page exceeds the ${MAX_RESPONSE_BYTES / 1024 / 1024} MB download limit`,
|
||||
},
|
||||
{ status: 413 },
|
||||
)
|
||||
}
|
||||
const buffer = await response.arrayBuffer()
|
||||
const charset = detectCharset(contentType, buffer)
|
||||
html = new TextDecoder(charset).decode(buffer)
|
||||
} catch (err: any) {
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
*/
|
||||
|
||||
import { streamObject } from "ai"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
import { getValidationModel } from "@/lib/ai-providers"
|
||||
import { VALIDATION_SYSTEM_PROMPT } from "@/lib/validation-prompts"
|
||||
import {
|
||||
@@ -14,9 +13,6 @@ import {
|
||||
|
||||
export const maxDuration = 30
|
||||
|
||||
// Data URL length cap (~3.75 MB of PNG), well above a normal diagram capture
|
||||
const MAX_IMAGE_DATA_LENGTH = 5 * 1024 * 1024
|
||||
|
||||
interface ValidateDiagramRequest {
|
||||
imageData: string // Base64 PNG data URL
|
||||
sessionId?: string
|
||||
@@ -48,10 +44,6 @@ function createStreamingResponse(result: ValidationResult): Response {
|
||||
}
|
||||
|
||||
export async function POST(req: Request): Promise<Response> {
|
||||
// Uses the server's model credentials, so require the access code
|
||||
const accessError = checkAccessCode(req)
|
||||
if (accessError) return accessError
|
||||
|
||||
try {
|
||||
// Check if VLM validation is enabled (default: true)
|
||||
const enableValidation = process.env.ENABLE_VLM_VALIDATION !== "false"
|
||||
@@ -80,13 +72,6 @@ export async function POST(req: Request): Promise<Response> {
|
||||
)
|
||||
}
|
||||
|
||||
if (imageData.length > MAX_IMAGE_DATA_LENGTH) {
|
||||
return Response.json(
|
||||
{ error: "Image data too large" },
|
||||
{ status: 413 },
|
||||
)
|
||||
}
|
||||
|
||||
// Get the validation model
|
||||
let model
|
||||
try {
|
||||
|
||||
@@ -10,7 +10,6 @@ import { createOpenRouter } from "@openrouter/ai-sdk-provider"
|
||||
import { generateText } from "ai"
|
||||
import { NextResponse } from "next/server"
|
||||
import { createOllama } from "ollama-ai-provider-v2"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
import {
|
||||
AIHUBMIX_APP_CODE,
|
||||
isAihubmixStandardBaseURL,
|
||||
@@ -34,24 +33,7 @@ interface ValidateRequest {
|
||||
vertexApiKey?: string // Express Mode API key
|
||||
}
|
||||
|
||||
// With private URLs blocked, a public baseUrl could still redirect the
|
||||
// request to an internal host, so redirects are refused in that case.
|
||||
function redirectGuardedFetch(): typeof fetch | undefined {
|
||||
if (allowPrivateUrls()) return undefined
|
||||
return async (input, init) => {
|
||||
const response = await fetch(input, { ...init, redirect: "manual" })
|
||||
if (response.status >= 300 && response.status < 400) {
|
||||
throw new Error("Redirects are not allowed for custom base URLs")
|
||||
}
|
||||
return response
|
||||
}
|
||||
}
|
||||
|
||||
export async function POST(req: Request) {
|
||||
// Lets the server send requests to arbitrary URLs, so require the access code
|
||||
const accessError = checkAccessCode(req)
|
||||
if (accessError) return accessError
|
||||
|
||||
try {
|
||||
const body: ValidateRequest = await req.json()
|
||||
const {
|
||||
@@ -109,7 +91,6 @@ export async function POST(req: Request) {
|
||||
)
|
||||
}
|
||||
|
||||
const guardedFetch = redirectGuardedFetch()
|
||||
let model: any
|
||||
|
||||
switch (provider) {
|
||||
@@ -117,7 +98,6 @@ export async function POST(req: Request) {
|
||||
const openai = createOpenAI({
|
||||
apiKey,
|
||||
...(baseUrl && { baseURL: baseUrl }),
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = openai.chat(modelId)
|
||||
break
|
||||
@@ -127,7 +107,6 @@ export async function POST(req: Request) {
|
||||
const anthropic = createAnthropic({
|
||||
apiKey,
|
||||
baseURL: baseUrl || "https://api.anthropic.com/v1",
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = anthropic(modelId)
|
||||
break
|
||||
@@ -137,7 +116,6 @@ export async function POST(req: Request) {
|
||||
const google = createGoogleGenerativeAI({
|
||||
apiKey,
|
||||
...(baseUrl && { baseURL: baseUrl }),
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = google(modelId)
|
||||
break
|
||||
@@ -147,7 +125,6 @@ export async function POST(req: Request) {
|
||||
const vertex = createVertex({
|
||||
apiKey: vertexApiKey,
|
||||
...(baseUrl && { baseURL: baseUrl }),
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = vertex(modelId)
|
||||
break
|
||||
@@ -157,7 +134,6 @@ export async function POST(req: Request) {
|
||||
const azure = createOpenAI({
|
||||
apiKey,
|
||||
baseURL: baseUrl,
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = azure.chat(modelId)
|
||||
break
|
||||
@@ -177,7 +153,6 @@ export async function POST(req: Request) {
|
||||
const openrouter = createOpenRouter({
|
||||
apiKey,
|
||||
...(baseUrl && { baseURL: baseUrl }),
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = openrouter(modelId)
|
||||
break
|
||||
@@ -199,7 +174,6 @@ export async function POST(req: Request) {
|
||||
const aihubmixCompatible = createOpenAI({
|
||||
apiKey,
|
||||
baseURL: baseUrl,
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = aihubmixCompatible.chat(modelId)
|
||||
}
|
||||
@@ -211,7 +185,6 @@ export async function POST(req: Request) {
|
||||
const ds = createDeepSeek({
|
||||
apiKey,
|
||||
...(baseUrl && { baseURL: baseUrl }),
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = ds(modelId)
|
||||
} else {
|
||||
@@ -224,7 +197,6 @@ export async function POST(req: Request) {
|
||||
const sf = createOpenAI({
|
||||
apiKey,
|
||||
baseURL: baseUrl || "https://api.siliconflow.cn/v1",
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = sf.chat(modelId)
|
||||
break
|
||||
@@ -241,7 +213,6 @@ export async function POST(req: Request) {
|
||||
baseUrl ||
|
||||
process.env.OLLAMA_BASE_URL ||
|
||||
"https://ollama.com/api",
|
||||
fetch: guardedFetch,
|
||||
...(ollamaApiKey && {
|
||||
headers: { Authorization: `Bearer ${ollamaApiKey}` },
|
||||
}),
|
||||
@@ -254,7 +225,6 @@ export async function POST(req: Request) {
|
||||
const gw = createGateway({
|
||||
apiKey,
|
||||
...(baseUrl && { baseURL: baseUrl }),
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = gw(modelId)
|
||||
break
|
||||
@@ -262,16 +232,13 @@ export async function POST(req: Request) {
|
||||
|
||||
case "edgeone": {
|
||||
// EdgeOne uses OpenAI-compatible API via Edge Functions
|
||||
// Need to pass cookies for EdgeOne Pages authentication,
|
||||
// and the access code, which the edge function also checks
|
||||
// Need to pass cookies for EdgeOne Pages authentication
|
||||
const cookieHeader = req.headers.get("cookie") || ""
|
||||
const edgeone = createOpenAI({
|
||||
apiKey: "edgeone", // EdgeOne doesn't require API key
|
||||
baseURL: baseUrl || "/api/edgeai",
|
||||
fetch: guardedFetch,
|
||||
headers: {
|
||||
cookie: cookieHeader,
|
||||
"x-access-code": req.headers.get("x-access-code") || "",
|
||||
},
|
||||
})
|
||||
model = edgeone.chat(modelId)
|
||||
@@ -283,7 +250,6 @@ export async function POST(req: Request) {
|
||||
const sglang = createOpenAI({
|
||||
apiKey: apiKey || "not-needed",
|
||||
baseURL: baseUrl || "http://127.0.0.1:8000/v1",
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = sglang.chat(modelId)
|
||||
break
|
||||
@@ -301,14 +267,12 @@ export async function POST(req: Request) {
|
||||
const doubao = createDeepSeek({
|
||||
apiKey,
|
||||
baseURL: doubaoBaseUrl,
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = doubao(modelId)
|
||||
} else {
|
||||
const doubao = createOpenAI({
|
||||
apiKey,
|
||||
baseURL: doubaoBaseUrl,
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = doubao.chat(modelId)
|
||||
}
|
||||
@@ -322,7 +286,7 @@ export async function POST(req: Request) {
|
||||
|
||||
try {
|
||||
// Initiate a streaming request (required for QwQ-32B and certain Qwen3 models)
|
||||
const response = await (guardedFetch ?? fetch)(
|
||||
const response = await fetch(
|
||||
`${baseURL}/chat/completions`,
|
||||
{
|
||||
method: "POST",
|
||||
@@ -343,15 +307,9 @@ export async function POST(req: Request) {
|
||||
)
|
||||
|
||||
if (!response.ok) {
|
||||
// Log the body but return only the status: the
|
||||
// caller chooses baseUrl, so the body may come from
|
||||
// any host the server can reach
|
||||
console.error(
|
||||
"[validate-model] ModelScope error body:",
|
||||
await response.text(),
|
||||
)
|
||||
const errorText = await response.text()
|
||||
throw new Error(
|
||||
`ModelScope API error (${response.status})`,
|
||||
`ModelScope API error (${response.status}): ${errorText}`,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -402,14 +360,12 @@ export async function POST(req: Request) {
|
||||
const minimax = createAnthropic({
|
||||
apiKey,
|
||||
baseURL: minimaxBaseUrl,
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = minimax.chat(modelId)
|
||||
} else {
|
||||
const minimax = createOpenAI({
|
||||
apiKey,
|
||||
baseURL: minimaxBaseUrl,
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = minimax.chat(modelId)
|
||||
}
|
||||
@@ -442,7 +398,6 @@ export async function POST(req: Request) {
|
||||
const openai = createOpenAI({
|
||||
apiKey,
|
||||
baseURL,
|
||||
fetch: guardedFetch,
|
||||
})
|
||||
model = openai.chat(modelId)
|
||||
break
|
||||
|
||||
@@ -1,9 +1,29 @@
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
|
||||
export async function POST(req: Request) {
|
||||
if (checkAccessCode(req)) {
|
||||
const accessCodes =
|
||||
process.env.ACCESS_CODE_LIST?.split(",")
|
||||
.map((code) => code.trim())
|
||||
.filter(Boolean) || []
|
||||
|
||||
// If no access codes configured, verification always passes
|
||||
if (accessCodes.length === 0) {
|
||||
return Response.json({
|
||||
valid: true,
|
||||
message: "No access code required",
|
||||
})
|
||||
}
|
||||
|
||||
const accessCodeHeader = req.headers.get("x-access-code")
|
||||
|
||||
if (!accessCodeHeader) {
|
||||
return Response.json(
|
||||
{ valid: false, message: "Invalid or missing access code" },
|
||||
{ valid: false, message: "Access code is required" },
|
||||
{ status: 401 },
|
||||
)
|
||||
}
|
||||
|
||||
if (!accessCodes.includes(accessCodeHeader)) {
|
||||
return Response.json(
|
||||
{ valid: false, message: "Invalid access code" },
|
||||
{ status: 401 },
|
||||
)
|
||||
}
|
||||
|
||||
+36
-68
@@ -11,9 +11,7 @@ import {
|
||||
} from "lucide-react"
|
||||
import type React from "react"
|
||||
import {
|
||||
type Dispatch,
|
||||
forwardRef,
|
||||
type SetStateAction,
|
||||
useCallback,
|
||||
useEffect,
|
||||
useImperativeHandle,
|
||||
@@ -43,20 +41,9 @@ import { FilePreviewList } from "./file-preview-list"
|
||||
|
||||
const MAX_IMAGE_SIZE = 2 * 1024 * 1024 // 2MB
|
||||
const MAX_FILES = 5
|
||||
// Image formats every supported model provider accepts (SVG is read as text)
|
||||
const SUPPORTED_IMAGE_TYPES = [
|
||||
"image/png",
|
||||
"image/jpeg",
|
||||
"image/gif",
|
||||
"image/webp",
|
||||
]
|
||||
|
||||
function isValidFileType(file: File): boolean {
|
||||
return (
|
||||
SUPPORTED_IMAGE_TYPES.includes(file.type) ||
|
||||
isPdfFile(file) ||
|
||||
isTextFile(file)
|
||||
)
|
||||
return file.type.startsWith("image/") || isPdfFile(file) || isTextFile(file)
|
||||
}
|
||||
|
||||
function formatFileSize(bytes: number): string {
|
||||
@@ -177,7 +164,7 @@ interface ChatInputProps {
|
||||
{ text: string; charCount: number; isExtracting: boolean }
|
||||
>
|
||||
urlData?: Map<string, UrlData>
|
||||
onUrlChange?: Dispatch<SetStateAction<Map<string, UrlData>>>
|
||||
onUrlChange?: (data: Map<string, UrlData>) => void
|
||||
|
||||
sessionId?: string
|
||||
error?: Error | null
|
||||
@@ -257,11 +244,6 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
// Allow retry when there's an error (even if status is still "streaming" or "submitted")
|
||||
const isDisabled =
|
||||
(status === "streaming" || status === "submitted") && !error
|
||||
// Block sending until attached files and URLs have their text, otherwise
|
||||
// their content would be silently dropped
|
||||
const isExtractingAttachments =
|
||||
files.some((file) => pdfData.get(file)?.isExtracting) ||
|
||||
Array.from(urlData?.values() ?? []).some((d) => d.isExtracting)
|
||||
|
||||
const adjustTextareaHeight = useCallback(() => {
|
||||
const textarea = textareaRef.current
|
||||
@@ -299,9 +281,6 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
}
|
||||
|
||||
const handleKeyDown = (e: React.KeyboardEvent) => {
|
||||
// Enter that confirms an IME candidate must not send the message
|
||||
if (e.nativeEvent.isComposing || e.keyCode === 229) return
|
||||
|
||||
const shouldSend =
|
||||
sendShortcut === "enter"
|
||||
? e.key === "Enter" &&
|
||||
@@ -313,12 +292,7 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
if (shouldSend) {
|
||||
e.preventDefault()
|
||||
const form = e.currentTarget.closest("form")
|
||||
if (
|
||||
form &&
|
||||
input.trim() &&
|
||||
!isDisabled &&
|
||||
!isExtractingAttachments
|
||||
) {
|
||||
if (form && input.trim() && !isDisabled) {
|
||||
form.requestSubmit()
|
||||
}
|
||||
}
|
||||
@@ -406,9 +380,13 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
|
||||
if (isDisabled) return
|
||||
|
||||
// Let validateFiles show a toast for unsupported types
|
||||
const droppedFiles = e.dataTransfer.files
|
||||
const supportedFiles = Array.from(droppedFiles).filter((file) =>
|
||||
isValidFileType(file),
|
||||
)
|
||||
|
||||
const { validFiles, errors } = validateFiles(
|
||||
Array.from(e.dataTransfer.files),
|
||||
supportedFiles,
|
||||
files.length,
|
||||
dict,
|
||||
)
|
||||
@@ -423,34 +401,33 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
|
||||
setIsExtractingUrl(true)
|
||||
|
||||
// Use functional updates so a removal or send made while extracting
|
||||
// is not overwritten when the request finishes
|
||||
try {
|
||||
onUrlChange((prev) =>
|
||||
new Map(prev).set(url, {
|
||||
url,
|
||||
title: url,
|
||||
content: "",
|
||||
charCount: 0,
|
||||
isExtracting: true,
|
||||
}),
|
||||
)
|
||||
const existing = urlData
|
||||
? new Map(urlData)
|
||||
: new Map<string, UrlData>()
|
||||
existing.set(url, {
|
||||
url,
|
||||
title: url,
|
||||
content: "",
|
||||
charCount: 0,
|
||||
isExtracting: true,
|
||||
})
|
||||
onUrlChange(existing)
|
||||
|
||||
const data = await extractUrlContent(url)
|
||||
|
||||
// Skip if the URL was removed while extracting
|
||||
onUrlChange((prev) =>
|
||||
prev.has(url) ? new Map(prev).set(url, data) : prev,
|
||||
)
|
||||
const newUrlData = new Map(existing)
|
||||
newUrlData.set(url, data)
|
||||
onUrlChange(newUrlData)
|
||||
|
||||
setShowUrlDialog(false)
|
||||
} catch (error) {
|
||||
// Remove the URL from the data map on error
|
||||
onUrlChange((prev) => {
|
||||
const next = new Map(prev)
|
||||
next.delete(url)
|
||||
return next
|
||||
})
|
||||
const newUrlData = urlData
|
||||
? new Map(urlData)
|
||||
: new Map<string, UrlData>()
|
||||
newUrlData.delete(url)
|
||||
onUrlChange(newUrlData)
|
||||
showErrorToast(
|
||||
<span className="text-muted-foreground">
|
||||
{error instanceof Error
|
||||
@@ -486,12 +463,11 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
urlData={urlData}
|
||||
onRemoveUrl={
|
||||
onUrlChange
|
||||
? (url) =>
|
||||
onUrlChange((prev) => {
|
||||
const next = new Map(prev)
|
||||
next.delete(url)
|
||||
return next
|
||||
})
|
||||
? (url) => {
|
||||
const next = new Map(urlData)
|
||||
next.delete(url)
|
||||
onUrlChange(next)
|
||||
}
|
||||
: undefined
|
||||
}
|
||||
/>
|
||||
@@ -583,7 +559,7 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
ref={fileInputRef}
|
||||
className="hidden"
|
||||
onChange={handleFileChange}
|
||||
accept="image/png,image/jpeg,image/gif,image/webp,.svg,.pdf,application/pdf,text/*,.md,.markdown,.json,.csv,.xml,.yaml,.yml,.toml"
|
||||
accept="image/*,.pdf,application/pdf,text/*,.md,.markdown,.json,.csv,.xml,.yaml,.yml,.toml"
|
||||
multiple
|
||||
disabled={isDisabled}
|
||||
/>
|
||||
@@ -612,11 +588,7 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
) : (
|
||||
<Button
|
||||
type="submit"
|
||||
disabled={
|
||||
isDisabled ||
|
||||
isExtractingAttachments ||
|
||||
!input.trim()
|
||||
}
|
||||
disabled={isDisabled || !input.trim()}
|
||||
size="sm"
|
||||
className="h-8 px-4 rounded-xl font-medium shadow-sm"
|
||||
aria-label={dict.chat.send}
|
||||
@@ -657,11 +629,7 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
|
||||
<TemplateCreateDialog
|
||||
open={showSaveAsTemplate}
|
||||
onOpenChange={setShowSaveAsTemplate}
|
||||
onSuccess={() => {
|
||||
setShowSaveAsTemplate(false)
|
||||
// Let the template list in the lobby reload
|
||||
window.dispatchEvent(new Event("templatesChanged"))
|
||||
}}
|
||||
onSuccess={() => setShowSaveAsTemplate(false)}
|
||||
initialPrompt={input.trim()}
|
||||
/>
|
||||
</form>
|
||||
|
||||
@@ -129,14 +129,12 @@ const getMessageTextContent = (message: UIMessage): string => {
|
||||
.join("\n")
|
||||
}
|
||||
|
||||
// Matches the [PDF: ...], [File: ...] and [URL: ...] sections appended to the user's text
|
||||
export const APPENDED_FILE_SECTIONS_PATTERN =
|
||||
/\n\n\[(PDF|File|URL):\s*[^\]]+\]\n[\s\S]*$/
|
||||
|
||||
// Get only the user's original text, excluding appended file content
|
||||
const getUserOriginalText = (message: UIMessage): string => {
|
||||
const fullText = getMessageTextContent(message)
|
||||
return fullText.replace(APPENDED_FILE_SECTIONS_PATTERN, "").trim()
|
||||
// Strip out [PDF: ...], [File: ...], and [URL: ...] sections that were appended
|
||||
const filePattern = /\n\n\[(PDF|File|URL):\s*[^\]]+\]\n[\s\S]*$/
|
||||
return fullText.replace(filePattern, "").trim()
|
||||
}
|
||||
|
||||
interface SessionMetadata {
|
||||
@@ -460,11 +458,6 @@ export function ChatMessageDisplay({
|
||||
messages.length > 0 ? [messages[messages.length - 1]] : []
|
||||
|
||||
messagesToProcess.forEach((message) => {
|
||||
// Messages restored from a saved session were applied before it was
|
||||
// saved; the saved diagram is authoritative, so don't replay them
|
||||
const isRestoredMessage =
|
||||
loadedMessageIdsRef?.current.has(message.id) ?? false
|
||||
|
||||
if (message.parts) {
|
||||
message.parts.forEach((part) => {
|
||||
if (part.type?.startsWith("tool-")) {
|
||||
@@ -482,8 +475,6 @@ export function ChatMessageDisplay({
|
||||
})
|
||||
}
|
||||
|
||||
if (isRestoredMessage) return
|
||||
|
||||
if (
|
||||
part.type === "tool-display_diagram" &&
|
||||
input?.xml
|
||||
@@ -550,32 +541,6 @@ export function ChatMessageDisplay({
|
||||
part.type === "tool-edit_diagram" &&
|
||||
input?.operations
|
||||
) {
|
||||
// Failed or stopped: drop the queued preview. If the original
|
||||
// XML is still stored, the tool handler never ran (user pressed
|
||||
// stop), so undo the streamed preview here.
|
||||
if (state === "output-error") {
|
||||
if (
|
||||
pendingEditRef.current?.toolCallId ===
|
||||
toolCallId &&
|
||||
editDebounceTimeoutRef.current
|
||||
) {
|
||||
clearTimeout(editDebounceTimeoutRef.current)
|
||||
editDebounceTimeoutRef.current = null
|
||||
pendingEditRef.current = null
|
||||
}
|
||||
const originalXml =
|
||||
editDiagramOriginalXmlRef.current.get(
|
||||
toolCallId,
|
||||
)
|
||||
if (originalXml) {
|
||||
editDiagramOriginalXmlRef.current.delete(
|
||||
toolCallId,
|
||||
)
|
||||
onDisplayChart(originalXml, true)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
const completeOps = getCompleteOperations(
|
||||
input.operations as DiagramOperation[],
|
||||
)
|
||||
@@ -645,10 +610,9 @@ export function ChatMessageDisplay({
|
||||
origXml,
|
||||
pending.operations,
|
||||
)
|
||||
// Load the full document so other pages stay intact
|
||||
onDisplayChart(
|
||||
handleDisplayChart(
|
||||
editedXml,
|
||||
true,
|
||||
false,
|
||||
)
|
||||
lastProcessedXmlRef.current.set(
|
||||
pending.toolCallId +
|
||||
|
||||
+50
-103
@@ -32,7 +32,6 @@ import { useSessionManager } from "@/hooks/use-session-manager"
|
||||
import { useValidateDiagram } from "@/hooks/use-validate-diagram"
|
||||
import { getApiEndpoint } from "@/lib/base-path"
|
||||
import { findCachedResponse } from "@/lib/cached-responses"
|
||||
import { isMinimalDiagram } from "@/lib/chat-helpers"
|
||||
import type { DrawioTheme } from "@/lib/drawio-themes"
|
||||
import { formatMessage } from "@/lib/i18n/utils"
|
||||
import { isPdfFile, isTextFile } from "@/lib/pdf-utils"
|
||||
@@ -41,12 +40,9 @@ import { STORAGE_KEYS } from "@/lib/storage"
|
||||
import type { UrlData } from "@/lib/url-utils"
|
||||
import { type FileData, useFileProcessor } from "@/lib/use-file-processor"
|
||||
import { useQuotaManager } from "@/lib/use-quota-manager"
|
||||
import { cn, formatXML, isRealDiagram, wrapWithMxFile } from "@/lib/utils"
|
||||
import { cn, formatXML, isRealDiagram } from "@/lib/utils"
|
||||
import type { ValidationState } from "./chat/ValidationCard"
|
||||
import {
|
||||
APPENDED_FILE_SECTIONS_PATTERN,
|
||||
ChatMessageDisplay,
|
||||
} from "./chat-message-display"
|
||||
import { ChatMessageDisplay } from "./chat-message-display"
|
||||
import { DevXmlSimulator } from "./dev-xml-simulator"
|
||||
|
||||
// localStorage keys for persistence
|
||||
@@ -111,18 +107,6 @@ function hasToolErrors(messages: ChatMessage[]): boolean {
|
||||
return lastToolPart?.state === TOOL_ERROR_STATE
|
||||
}
|
||||
|
||||
/**
|
||||
* Snapshots keep the full multi-page document, but the model only sees and
|
||||
* edits the first page, so give it the first page's mxGraphModel.
|
||||
* Older snapshots already hold a single mxGraphModel and are returned as is.
|
||||
*/
|
||||
function getFirstPageXml(xml: string): string {
|
||||
if (!xml.includes("<mxfile")) return xml
|
||||
const doc = new DOMParser().parseFromString(xml, "text/xml")
|
||||
const model = doc.querySelector("diagram")?.querySelector("mxGraphModel")
|
||||
return model ? formatXML(new XMLSerializer().serializeToString(model)) : xml
|
||||
}
|
||||
|
||||
export default function ChatPanel({
|
||||
isVisible,
|
||||
onToggleVisibility,
|
||||
@@ -352,8 +336,19 @@ export default function ChatPanel({
|
||||
localStorage.setItem(STORAGE_KEYS.maxOutputTokens, digitsOnly)
|
||||
}, [])
|
||||
|
||||
// Failed VLM validations in the current user turn (reset on user action)
|
||||
const validationRetryCountRef = useRef(0)
|
||||
// Ref to store the sendMessage function for use in callbacks
|
||||
const sendMessageRef = useRef<typeof sendMessage | null>(null)
|
||||
|
||||
// Callback to improve diagram with validation suggestions
|
||||
const handleImproveWithSuggestions = useCallback((feedback: string) => {
|
||||
if (sendMessageRef.current) {
|
||||
// Send the feedback as a new user message to trigger regeneration
|
||||
sendMessageRef.current({
|
||||
role: "user",
|
||||
parts: [{ type: "text", text: feedback }],
|
||||
})
|
||||
}
|
||||
}, [])
|
||||
|
||||
// VLM validation hook using AI SDK's useObject
|
||||
const { validateWithFallback } = useValidateDiagram()
|
||||
@@ -362,7 +357,6 @@ export default function ChatPanel({
|
||||
const { handleToolCall } = useDiagramToolHandlers({
|
||||
partialXmlRef,
|
||||
editDiagramOriginalXmlRef,
|
||||
validationRetryCountRef,
|
||||
chartXMLRef,
|
||||
onDisplayChart,
|
||||
onFetchChart,
|
||||
@@ -524,6 +518,11 @@ export default function ChatPanel({
|
||||
},
|
||||
})
|
||||
|
||||
// Store sendMessage in ref for use in callbacks (like handleImproveWithSuggestions)
|
||||
useEffect(() => {
|
||||
sendMessageRef.current = sendMessage
|
||||
}, [sendMessage])
|
||||
|
||||
// Ref to track latest messages for unload persistence
|
||||
const messagesRef = useRef(messages)
|
||||
useEffect(() => {
|
||||
@@ -532,9 +531,6 @@ export default function ChatPanel({
|
||||
|
||||
// Track last synced session ID to detect external changes (e.g., URL back/forward)
|
||||
const lastSyncedSessionIdRef = useRef<string | null>(null)
|
||||
// Messages array from our latest save. A session holding this exact array was
|
||||
// created by our own save, so it must not be treated as an external switch.
|
||||
const lastSavedMessagesRef = useRef<unknown[] | null>(null)
|
||||
|
||||
// Helper: Sync UI state with session data (eliminates duplication)
|
||||
// Track message IDs that are being loaded from session (to skip animations/scroll)
|
||||
@@ -601,10 +597,8 @@ export default function ChatPanel({
|
||||
thumbnailDataUrl = latestSvgRef.current
|
||||
}
|
||||
}
|
||||
const messages = sanitizeMessages(messagesRef.current)
|
||||
lastSavedMessagesRef.current = messages
|
||||
return {
|
||||
messages,
|
||||
messages: sanitizeMessages(messagesRef.current),
|
||||
xmlSnapshots: Array.from(xmlSnapshotsRef.current.entries()),
|
||||
diagramXml: currentDiagramXml,
|
||||
thumbnailDataUrl,
|
||||
@@ -657,13 +651,8 @@ export default function ChatPanel({
|
||||
// Skip if session ID hasn't changed (our own saves don't change the ID)
|
||||
if (newSessionId === lastSyncedSessionIdRef.current) return
|
||||
|
||||
// Our own save created this session; the UI already shows its content
|
||||
const isOwnNewSession =
|
||||
newSession?.messages === lastSavedMessagesRef.current
|
||||
|
||||
// Update last synced ID
|
||||
lastSyncedSessionIdRef.current = newSessionId
|
||||
if (isOwnNewSession) return
|
||||
|
||||
// Sync UI with new session
|
||||
if (newSession) {
|
||||
@@ -804,23 +793,12 @@ export default function ChatPanel({
|
||||
const onFormSubmit = async (e: React.FormEvent<HTMLFormElement>) => {
|
||||
e.preventDefault()
|
||||
const isProcessing = status === "streaming" || status === "submitted"
|
||||
// Attachments still extracting have no text yet. Template sends call
|
||||
// requestSubmit() and skip the disabled send button, so check here too.
|
||||
const isExtracting =
|
||||
files.some((f) => pdfData.get(f)?.isExtracting) ||
|
||||
Array.from(urlData.values()).some((d) => d.isExtracting)
|
||||
if (input.trim() && !isProcessing && !isExtracting) {
|
||||
// Check if input matches a cached example (only when no messages
|
||||
// yet and the canvas is empty, same rule as the server)
|
||||
if (
|
||||
messages.length === 0 &&
|
||||
isMinimalDiagram(chartXMLRef.current || "")
|
||||
) {
|
||||
// Pass the file name so a user's own file never matches an example
|
||||
if (input.trim() && !isProcessing) {
|
||||
// Check if input matches a cached example (only when no messages yet)
|
||||
if (messages.length === 0) {
|
||||
const cached = findCachedResponse(
|
||||
input.trim(),
|
||||
files.length > 0,
|
||||
files.length === 1 ? files[0].name : undefined,
|
||||
)
|
||||
if (cached) {
|
||||
// Add user message and fake assistant response to messages
|
||||
@@ -856,11 +834,6 @@ export default function ChatPanel({
|
||||
],
|
||||
},
|
||||
] as any)
|
||||
// Snapshot the canvas before the example so editing this message works
|
||||
xmlSnapshotsRef.current.set(
|
||||
0,
|
||||
chartXMLRef.current || wrapWithMxFile(""),
|
||||
)
|
||||
setInput("")
|
||||
sessionStorage.removeItem(SESSION_STORAGE_INPUT_KEY)
|
||||
setFiles([])
|
||||
@@ -870,6 +843,9 @@ export default function ChatPanel({
|
||||
}
|
||||
|
||||
try {
|
||||
let chartXml = await onFetchChart()
|
||||
chartXml = formatXML(chartXml)
|
||||
|
||||
// Build user text by concatenating input with pre-extracted text
|
||||
// (Backend only reads first text part, so we must combine them)
|
||||
const parts: any[] = []
|
||||
@@ -884,7 +860,20 @@ export default function ChatPanel({
|
||||
// Add the combined text as the first part
|
||||
parts.unshift({ type: "text", text: userText })
|
||||
|
||||
await sendWithCurrentDiagram(parts)
|
||||
// Get previous XML from the last snapshot (before this message)
|
||||
const snapshotKeys = Array.from(
|
||||
xmlSnapshotsRef.current.keys(),
|
||||
).sort((a, b) => b - a)
|
||||
const previousXml =
|
||||
snapshotKeys.length > 0
|
||||
? xmlSnapshotsRef.current.get(snapshotKeys[0]) || ""
|
||||
: ""
|
||||
|
||||
// Save XML snapshot for this message (will be at index = current messages.length)
|
||||
const messageIndex = messages.length
|
||||
xmlSnapshotsRef.current.set(messageIndex, chartXml)
|
||||
|
||||
sendChatMessage(parts, chartXml, previousXml, sessionId)
|
||||
|
||||
// Token count is tracked in onFinish with actual server usage
|
||||
setInput("")
|
||||
@@ -893,37 +882,10 @@ export default function ChatPanel({
|
||||
setUrlData(new Map())
|
||||
} catch (error) {
|
||||
console.error("Error fetching chart data:", error)
|
||||
toast.error(dict.errors.failedToExport)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Export the current diagram, snapshot it for this message, and send
|
||||
const sendWithCurrentDiagram = async (parts: any[]) => {
|
||||
const chartXml = formatXML(await onFetchChart())
|
||||
const previousXml = getPreviousXml(messages.length)
|
||||
|
||||
// Snapshot the full multi-page document (kept fresh by autosave) so
|
||||
// regenerate/edit can restore every page; the model gets page 1 only
|
||||
xmlSnapshotsRef.current.set(
|
||||
messages.length,
|
||||
chartXMLRef.current || chartXml,
|
||||
)
|
||||
|
||||
sendChatMessage(parts, chartXml, previousXml, sessionId)
|
||||
}
|
||||
|
||||
// Send VLM validation feedback as a new user message through the normal send path
|
||||
const handleImproveWithSuggestions = async (feedback: string) => {
|
||||
if (status === "streaming" || status === "submitted") return
|
||||
try {
|
||||
await sendWithCurrentDiagram([{ type: "text", text: feedback }])
|
||||
} catch (error) {
|
||||
console.error("Error fetching chart data:", error)
|
||||
toast.error(dict.errors.failedToExport)
|
||||
}
|
||||
}
|
||||
|
||||
// Handle session switching from history dropdown
|
||||
const handleSelectSession = useCallback(
|
||||
async (sessionId: string) => {
|
||||
@@ -1027,9 +989,10 @@ export default function ChatPanel({
|
||||
// Handle sending a template directly (called from TemplatePanel)
|
||||
const handleSendTemplate = useCallback(
|
||||
async (template: { prompt: string }) => {
|
||||
// Keep attachments: they are sent along with the template prompt
|
||||
flushSync(() => {
|
||||
setInput(template.prompt)
|
||||
setFiles([])
|
||||
setUrlData(new Map())
|
||||
})
|
||||
|
||||
const formElement = document.getElementById(
|
||||
@@ -1039,7 +1002,7 @@ export default function ChatPanel({
|
||||
formElement.requestSubmit()
|
||||
}
|
||||
},
|
||||
[setInput],
|
||||
[setInput, setFiles, setUrlData],
|
||||
)
|
||||
|
||||
const handleInputChange = (
|
||||
@@ -1054,15 +1017,13 @@ export default function ChatPanel({
|
||||
}
|
||||
|
||||
// Helper functions for message actions (regenerate/edit)
|
||||
// Extract previous XML snapshot (first page, as sent to the model) before a given message index
|
||||
// Extract previous XML snapshot before a given message index
|
||||
const getPreviousXml = (beforeIndex: number): string => {
|
||||
const snapshotKeys = Array.from(xmlSnapshotsRef.current.keys())
|
||||
.filter((k) => k < beforeIndex)
|
||||
.sort((a, b) => b - a)
|
||||
return snapshotKeys.length > 0
|
||||
? getFirstPageXml(
|
||||
xmlSnapshotsRef.current.get(snapshotKeys[0]) || "",
|
||||
)
|
||||
? xmlSnapshotsRef.current.get(snapshotKeys[0]) || ""
|
||||
: ""
|
||||
}
|
||||
|
||||
@@ -1114,7 +1075,6 @@ export default function ChatPanel({
|
||||
// Reset all retry/continuation state on user-initiated message
|
||||
autoRetryCountRef.current = 0
|
||||
continuationRetryCountRef.current = 0
|
||||
validationRetryCountRef.current = 0
|
||||
partialXmlRef.current = ""
|
||||
|
||||
const config = getSelectedAIConfig()
|
||||
@@ -1263,12 +1223,7 @@ export default function ChatPanel({
|
||||
})
|
||||
|
||||
// Now send the message after state is guaranteed to be updated
|
||||
sendChatMessage(
|
||||
userParts,
|
||||
getFirstPageXml(savedXml),
|
||||
previousXml,
|
||||
sessionId,
|
||||
)
|
||||
sendChatMessage(userParts, savedXml, previousXml, sessionId)
|
||||
}
|
||||
|
||||
const handleEditMessage = async (messageIndex: number, newText: string) => {
|
||||
@@ -1295,13 +1250,10 @@ export default function ChatPanel({
|
||||
// Clean up snapshots for messages after the user message (they will be removed)
|
||||
cleanupSnapshotsAfter(messageIndex)
|
||||
|
||||
// Create new parts with updated text. The edit box only shows the typed
|
||||
// text, so keep the appended PDF/file/URL content
|
||||
// Create new parts with updated text
|
||||
const newParts = message.parts?.map((part: any) => {
|
||||
if (part.type === "text") {
|
||||
const appended =
|
||||
part.text.match(APPENDED_FILE_SECTIONS_PATTERN)?.[0] ?? ""
|
||||
return { ...part, text: newText + appended }
|
||||
return { ...part, text: newText }
|
||||
}
|
||||
return part
|
||||
}) || [{ type: "text", text: newText }]
|
||||
@@ -1314,12 +1266,7 @@ export default function ChatPanel({
|
||||
})
|
||||
|
||||
// Now send the edited message after state is guaranteed to be updated
|
||||
sendChatMessage(
|
||||
newParts,
|
||||
getFirstPageXml(savedXml),
|
||||
previousXml,
|
||||
sessionId,
|
||||
)
|
||||
sendChatMessage(newParts, savedXml, previousXml, sessionId)
|
||||
}
|
||||
|
||||
// Collapsed view (desktop only)
|
||||
|
||||
@@ -194,8 +194,6 @@ export function ChatLobby({
|
||||
className="group w-full flex items-center gap-3 p-3 rounded-xl border border-border/60 bg-card hover:bg-accent/50 hover:border-primary/30 transition-all duration-200 cursor-pointer text-left"
|
||||
onClick={() => onSelectSession(session.id)}
|
||||
onKeyDown={(e) => {
|
||||
// Ignore keys bubbling up from the delete button
|
||||
if (e.target !== e.currentTarget) return
|
||||
if (
|
||||
e.key === "Enter" ||
|
||||
e.key === " "
|
||||
|
||||
@@ -55,9 +55,6 @@ export function TemplateCreateDialog({
|
||||
|
||||
const handleSubmit = async (e: React.FormEvent) => {
|
||||
e.preventDefault()
|
||||
// React submit events bubble through the portal; keep them away from
|
||||
// the chat form this dialog may be rendered in
|
||||
e.stopPropagation()
|
||||
|
||||
const trimmedPrompt = prompt.trim()
|
||||
if (!trimmedPrompt) {
|
||||
|
||||
@@ -39,16 +39,16 @@ export function TemplateEditDialog({
|
||||
const [isSubmitting, setIsSubmitting] = useState(false)
|
||||
const [error, setError] = useState<string | null>(null)
|
||||
|
||||
// Populate form each time the dialog opens, dropping any cancelled edits
|
||||
// Populate form when template changes
|
||||
useEffect(() => {
|
||||
if (open && template) {
|
||||
if (template) {
|
||||
setTitle(template.title || "")
|
||||
setDescription(template.description || "")
|
||||
setPrompt(template.prompt || "")
|
||||
setPinned(template.pinned || false)
|
||||
setError(null)
|
||||
}
|
||||
}, [open, template])
|
||||
}, [template])
|
||||
|
||||
const handleOpenChange = (newOpen: boolean) => {
|
||||
if (!newOpen) {
|
||||
@@ -59,9 +59,6 @@ export function TemplateEditDialog({
|
||||
|
||||
const handleSubmit = async (e: React.FormEvent) => {
|
||||
e.preventDefault()
|
||||
// React submit events bubble through the portal; keep them away from
|
||||
// any form this dialog may be rendered in
|
||||
e.stopPropagation()
|
||||
|
||||
if (!template) return
|
||||
|
||||
|
||||
@@ -110,10 +110,6 @@ export function TemplatePanel({
|
||||
|
||||
useEffect(() => {
|
||||
loadTemplates()
|
||||
// Reload when a template is saved elsewhere, e.g. from the chat input
|
||||
window.addEventListener("templatesChanged", loadTemplates)
|
||||
return () =>
|
||||
window.removeEventListener("templatesChanged", loadTemplates)
|
||||
}, [loadTemplates])
|
||||
|
||||
const handleCreateSuccess = () => {
|
||||
@@ -306,28 +302,6 @@ export function TemplatePanel({
|
||||
}
|
||||
}
|
||||
|
||||
// Shared by the empty state and the list, so import works in both
|
||||
const importInput = (
|
||||
<input
|
||||
ref={fileInputRef}
|
||||
type="file"
|
||||
accept="application/json,.json"
|
||||
onChange={handleImport}
|
||||
className="hidden"
|
||||
/>
|
||||
)
|
||||
const importMessageBox = importMessage && (
|
||||
<div
|
||||
className={`text-xs px-3 py-2 rounded-lg ${
|
||||
importMessage.type === "success"
|
||||
? "bg-green-100 text-green-800 dark:bg-green-900/30 dark:text-green-400"
|
||||
: "bg-red-100 text-red-800 dark:bg-red-900/30 dark:text-red-400"
|
||||
}`}
|
||||
>
|
||||
{importMessage.text}
|
||||
</div>
|
||||
)
|
||||
|
||||
// Empty state: no templates at all
|
||||
if (!loading && templates.length === 0) {
|
||||
return (
|
||||
@@ -358,18 +332,6 @@ export function TemplatePanel({
|
||||
<Plus className="w-4 h-4" />
|
||||
{dict.templates.createFirst}
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => fileInputRef.current?.click()}
|
||||
className="mt-2 inline-flex items-center gap-1.5 px-3 py-1.5 rounded-md text-xs font-medium text-muted-foreground hover:text-foreground hover:bg-muted transition-colors"
|
||||
>
|
||||
<Upload className="w-3.5 h-3.5" />
|
||||
{dict.templates.importTemplates}
|
||||
</button>
|
||||
{importInput}
|
||||
{importMessageBox && (
|
||||
<div className="mt-3">{importMessageBox}</div>
|
||||
)}
|
||||
|
||||
<TemplateCreateDialog
|
||||
open={createDialogOpen}
|
||||
@@ -427,11 +389,27 @@ export function TemplatePanel({
|
||||
<Upload className="w-3.5 h-3.5" />
|
||||
{dict.templates.importTemplates}
|
||||
</button>
|
||||
{importInput}
|
||||
<input
|
||||
ref={fileInputRef}
|
||||
type="file"
|
||||
accept="application/json,.json"
|
||||
onChange={handleImport}
|
||||
className="hidden"
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* Import message */}
|
||||
{importMessageBox}
|
||||
{importMessage && (
|
||||
<div
|
||||
className={`text-xs px-3 py-2 rounded-lg ${
|
||||
importMessage.type === "success"
|
||||
? "bg-green-100 text-green-800 dark:bg-green-900/30 dark:text-green-400"
|
||||
: "bg-red-100 text-red-800 dark:bg-red-900/30 dark:text-red-400"
|
||||
}`}
|
||||
>
|
||||
{importMessage.text}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div className="space-y-2">
|
||||
{loading
|
||||
@@ -469,8 +447,6 @@ export function TemplatePanel({
|
||||
handleTemplateClick(template)
|
||||
}
|
||||
onKeyDown={(e) => {
|
||||
// Ignore keys bubbling up from the action buttons
|
||||
if (e.target !== e.currentTarget) return
|
||||
if (
|
||||
e.key === "Enter" ||
|
||||
e.key === " "
|
||||
|
||||
@@ -66,7 +66,7 @@ export function ToolCallCard({
|
||||
dict,
|
||||
}: ToolCallCardProps) {
|
||||
const callId = part.toolCallId
|
||||
const { state, input, output, errorText } = part
|
||||
const { state, input, output } = part
|
||||
// Default to expanded for all states (user can manually collapse if needed)
|
||||
const isExpanded = expandedTools[callId] ?? true
|
||||
const toolName = part.type?.replace("tool-", "")
|
||||
@@ -92,14 +92,6 @@ export function ToolCallCard({
|
||||
}
|
||||
}
|
||||
|
||||
// Incomplete XML means the output hit the length limit, unless the user
|
||||
// stopped the generation themselves
|
||||
const isTruncated =
|
||||
state === "output-error" &&
|
||||
errorText !== "Stopped by user" &&
|
||||
(toolName === "display_diagram" || toolName === "append_diagram") &&
|
||||
!isMxCellXmlComplete(input?.xml)
|
||||
|
||||
const handleCopy = () => {
|
||||
let textToCopy = ""
|
||||
|
||||
@@ -169,15 +161,22 @@ export function ToolCallCard({
|
||||
</>
|
||||
)}
|
||||
{state === "output-error" &&
|
||||
(isTruncated ? (
|
||||
<span className="text-xs font-medium text-yellow-600 bg-yellow-50 px-2 py-0.5 rounded-full">
|
||||
Truncated
|
||||
</span>
|
||||
) : (
|
||||
<span className="text-xs font-medium text-red-600 bg-red-50 px-2 py-0.5 rounded-full">
|
||||
Error
|
||||
</span>
|
||||
))}
|
||||
(() => {
|
||||
// Check if this is a truncation (incomplete XML) vs real error
|
||||
const isTruncated =
|
||||
(toolName === "display_diagram" ||
|
||||
toolName === "append_diagram") &&
|
||||
!isMxCellXmlComplete(input?.xml)
|
||||
return isTruncated ? (
|
||||
<span className="text-xs font-medium text-yellow-600 bg-yellow-50 px-2 py-0.5 rounded-full">
|
||||
Truncated
|
||||
</span>
|
||||
) : (
|
||||
<span className="text-xs font-medium text-red-600 bg-red-50 px-2 py-0.5 rounded-full">
|
||||
Error
|
||||
</span>
|
||||
)
|
||||
})()}
|
||||
{input && Object.keys(input).length > 0 && (
|
||||
<button
|
||||
type="button"
|
||||
@@ -225,16 +224,23 @@ export function ToolCallCard({
|
||||
) : null}
|
||||
</div>
|
||||
)}
|
||||
{/* AI SDK stores tool errors in errorText */}
|
||||
{state === "output-error" && (errorText || output) && (
|
||||
<div
|
||||
className={`px-4 py-3 border-t border-border/40 text-sm whitespace-pre-wrap break-words ${isTruncated ? "text-yellow-600" : "text-red-600"}`}
|
||||
>
|
||||
{isTruncated
|
||||
? "Output truncated due to length limits. Try a simpler request or increase Max Output Tokens in Settings."
|
||||
: (errorText ?? output)}
|
||||
</div>
|
||||
)}
|
||||
{output &&
|
||||
state === "output-error" &&
|
||||
(() => {
|
||||
const isTruncated =
|
||||
(toolName === "display_diagram" ||
|
||||
toolName === "append_diagram") &&
|
||||
!isMxCellXmlComplete(input?.xml)
|
||||
return (
|
||||
<div
|
||||
className={`px-4 py-3 border-t border-border/40 text-sm ${isTruncated ? "text-yellow-600" : "text-red-600"}`}
|
||||
>
|
||||
{isTruncated
|
||||
? "Output truncated due to length limits. Try a simpler request or increase the maxOutputLength."
|
||||
: output}
|
||||
</div>
|
||||
)
|
||||
})()}
|
||||
{/* Show get_shape_library output on success */}
|
||||
{output &&
|
||||
toolName === "get_shape_library" &&
|
||||
|
||||
@@ -13,5 +13,4 @@ export interface ToolPartLike {
|
||||
operations?: DiagramOperation[]
|
||||
} & Record<string, unknown>
|
||||
output?: string
|
||||
errorText?: string
|
||||
}
|
||||
|
||||
@@ -56,7 +56,6 @@ import { useDictionary } from "@/hooks/use-dictionary"
|
||||
import type { UseModelConfigReturn } from "@/hooks/use-model-config"
|
||||
import { getApiEndpoint } from "@/lib/base-path"
|
||||
import { formatMessage } from "@/lib/i18n/utils"
|
||||
import { STORAGE_KEYS } from "@/lib/storage"
|
||||
import type { ProviderConfig, ProviderName } from "@/lib/types/model-config"
|
||||
import { PROVIDER_INFO, SUGGESTED_MODELS } from "@/lib/types/model-config"
|
||||
import { cn } from "@/lib/utils"
|
||||
@@ -134,14 +133,6 @@ export function ModelConfigDialog({
|
||||
modelId: string
|
||||
message: string
|
||||
} | null>(null)
|
||||
// Model ID being typed; written to the config only when valid on blur
|
||||
const [modelIdDraft, setModelIdDraft] = useState<{
|
||||
id: string
|
||||
value: string
|
||||
} | null>(null)
|
||||
// Bumped on every credential edit so a running test can tell that its
|
||||
// results belong to the old credentials
|
||||
const credentialsVersionRef = useRef(0)
|
||||
const [dynamicSuggestedModels, setDynamicSuggestedModels] = useState<
|
||||
Partial<Record<ProviderName, string[]>>
|
||||
>({})
|
||||
@@ -166,11 +157,6 @@ export function ModelConfigDialog({
|
||||
(p) => p.id === selectedProviderId,
|
||||
)
|
||||
|
||||
// Discard an unfinished model ID edit when the dialog closes
|
||||
useEffect(() => {
|
||||
if (!open) setModelIdDraft(null)
|
||||
}, [open])
|
||||
|
||||
// Cleanup validation reset timeout on unmount
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
@@ -267,9 +253,9 @@ export function ModelConfigDialog({
|
||||
field: keyof ProviderConfig,
|
||||
value: string | boolean,
|
||||
) => {
|
||||
if (!selectedProviderId || !selectedProvider) return
|
||||
const updates: Partial<ProviderConfig> = { [field]: value }
|
||||
// Reset validation of the provider and its models when credentials change
|
||||
if (!selectedProviderId) return
|
||||
updateProvider(selectedProviderId, { [field]: value })
|
||||
// Reset validation when credentials change
|
||||
const credentialFields = [
|
||||
"apiKey",
|
||||
"baseUrl",
|
||||
@@ -279,17 +265,9 @@ export function ModelConfigDialog({
|
||||
"vertexApiKey",
|
||||
]
|
||||
if (credentialFields.includes(field)) {
|
||||
credentialsVersionRef.current++
|
||||
setValidationStatus("idle")
|
||||
setValidatingModelIndex(null)
|
||||
updates.validated = false
|
||||
updates.models = selectedProvider.models.map((m) => ({
|
||||
...m,
|
||||
validated: undefined,
|
||||
validationError: undefined,
|
||||
}))
|
||||
updateProvider(selectedProviderId, { validated: false })
|
||||
}
|
||||
updateProvider(selectedProviderId, updates)
|
||||
}
|
||||
|
||||
// Handle adding a model to current provider
|
||||
@@ -359,7 +337,6 @@ export function ModelConfigDialog({
|
||||
|
||||
let allValid = true
|
||||
let errorCount = 0
|
||||
const credentialsVersion = credentialsVersionRef.current
|
||||
|
||||
// Validate each model
|
||||
for (let i = 0; i < selectedProvider.models.length; i++) {
|
||||
@@ -369,37 +346,26 @@ export function ModelConfigDialog({
|
||||
try {
|
||||
// For EdgeOne, construct baseUrl from current origin
|
||||
const baseUrl = isEdgeOne
|
||||
? `${window.location.origin}${getApiEndpoint("/api/edgeai")}`
|
||||
? `${window.location.origin}/api/edgeai`
|
||||
: selectedProvider.baseUrl
|
||||
|
||||
const response = await fetch(
|
||||
getApiEndpoint("/api/validate-model"),
|
||||
{
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
"x-access-code":
|
||||
localStorage.getItem(STORAGE_KEYS.accessCode) ||
|
||||
"",
|
||||
},
|
||||
body: JSON.stringify({
|
||||
provider: selectedProvider.provider,
|
||||
apiKey: selectedProvider.apiKey,
|
||||
baseUrl,
|
||||
modelId: model.modelId,
|
||||
// AWS Bedrock credentials
|
||||
awsAccessKeyId: selectedProvider.awsAccessKeyId,
|
||||
awsSecretAccessKey:
|
||||
selectedProvider.awsSecretAccessKey,
|
||||
awsRegion: selectedProvider.awsRegion,
|
||||
// Vertex AI credentials (Express Mode)
|
||||
vertexApiKey: selectedProvider.vertexApiKey,
|
||||
}),
|
||||
},
|
||||
)
|
||||
const data = await response.json().catch(() => ({}))
|
||||
// Credentials changed during the test: drop the results
|
||||
if (credentialsVersionRef.current !== credentialsVersion) return
|
||||
const response = await fetch("/api/validate-model", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({
|
||||
provider: selectedProvider.provider,
|
||||
apiKey: selectedProvider.apiKey,
|
||||
baseUrl,
|
||||
modelId: model.modelId,
|
||||
// AWS Bedrock credentials
|
||||
awsAccessKeyId: selectedProvider.awsAccessKeyId,
|
||||
awsSecretAccessKey: selectedProvider.awsSecretAccessKey,
|
||||
awsRegion: selectedProvider.awsRegion,
|
||||
// Vertex AI credentials (Express Mode)
|
||||
vertexApiKey: selectedProvider.vertexApiKey,
|
||||
}),
|
||||
})
|
||||
const data = await response.json()
|
||||
|
||||
if (data.valid) {
|
||||
updateModel(selectedProviderId, model.id, {
|
||||
@@ -411,15 +377,10 @@ export function ModelConfigDialog({
|
||||
errorCount++
|
||||
updateModel(selectedProviderId, model.id, {
|
||||
validated: false,
|
||||
validationError:
|
||||
data.error ||
|
||||
(response.ok
|
||||
? "Validation failed"
|
||||
: `Request failed (${response.status})`),
|
||||
validationError: data.error || "Validation failed",
|
||||
})
|
||||
}
|
||||
} catch {
|
||||
if (credentialsVersionRef.current !== credentialsVersion) return
|
||||
allValid = false
|
||||
errorCount++
|
||||
updateModel(selectedProviderId, model.id, {
|
||||
@@ -654,9 +615,7 @@ export function ModelConfigDialog({
|
||||
|
||||
{/* Add Provider */}
|
||||
<div className="p-3 border-t border-border-subtle">
|
||||
{/* Always empty so picking the same type again still fires */}
|
||||
<Select
|
||||
value=""
|
||||
onValueChange={(v) =>
|
||||
handleAddProvider(v as ProviderName)
|
||||
}
|
||||
@@ -878,7 +837,6 @@ export function ModelConfigDialog({
|
||||
<Plus className="h-3.5 w-3.5" />
|
||||
</Button>
|
||||
<Select
|
||||
value=""
|
||||
onValueChange={(value) => {
|
||||
if (value) {
|
||||
handleAddModel(
|
||||
@@ -1031,10 +989,7 @@ export function ModelConfigDialog({
|
||||
</div>
|
||||
<Input
|
||||
value={
|
||||
modelIdDraft?.id ===
|
||||
model.id
|
||||
? modelIdDraft.value
|
||||
: model.modelId
|
||||
model.modelId
|
||||
}
|
||||
title={
|
||||
model.modelId
|
||||
@@ -1052,14 +1007,24 @@ export function ModelConfigDialog({
|
||||
null,
|
||||
)
|
||||
}
|
||||
setModelIdDraft(
|
||||
{
|
||||
id: model.id,
|
||||
value: e
|
||||
.target
|
||||
.value,
|
||||
},
|
||||
)
|
||||
if (
|
||||
selectedProviderId
|
||||
) {
|
||||
updateModel(
|
||||
selectedProviderId,
|
||||
model.id,
|
||||
{
|
||||
modelId:
|
||||
e
|
||||
.target
|
||||
.value,
|
||||
validated:
|
||||
undefined,
|
||||
validationError:
|
||||
undefined,
|
||||
},
|
||||
)
|
||||
}
|
||||
}}
|
||||
onKeyDown={(
|
||||
e,
|
||||
@@ -1076,10 +1041,6 @@ export function ModelConfigDialog({
|
||||
) => {
|
||||
const newModelId =
|
||||
e.target.value.trim()
|
||||
// Drop the draft; an invalid ID falls back to the saved one
|
||||
setModelIdDraft(
|
||||
null,
|
||||
)
|
||||
|
||||
// Helper to show error with shake
|
||||
const showError =
|
||||
@@ -1174,24 +1135,6 @@ export function ModelConfigDialog({
|
||||
setEditError(
|
||||
null,
|
||||
)
|
||||
if (
|
||||
selectedProviderId &&
|
||||
newModelId !==
|
||||
model.modelId
|
||||
) {
|
||||
updateModel(
|
||||
selectedProviderId,
|
||||
model.id,
|
||||
{
|
||||
modelId:
|
||||
newModelId,
|
||||
validated:
|
||||
undefined,
|
||||
validationError:
|
||||
undefined,
|
||||
},
|
||||
)
|
||||
}
|
||||
}}
|
||||
className="flex-1 min-w-0 font-mono text-sm h-8 border-0 bg-transparent focus-visible:bg-background focus-visible:ring-1"
|
||||
/>
|
||||
|
||||
@@ -264,13 +264,9 @@ export function ModelSelector({
|
||||
(model) => (
|
||||
<ModelSelectorItem
|
||||
key={model.id}
|
||||
// Unique value so same-named models highlight
|
||||
// separately; keywords keep search by name
|
||||
value={model.id}
|
||||
keywords={[
|
||||
model.modelId,
|
||||
providerLabel,
|
||||
]}
|
||||
value={
|
||||
model.modelId
|
||||
}
|
||||
onSelect={() =>
|
||||
handleSelect(
|
||||
model.id,
|
||||
@@ -355,11 +351,9 @@ export function ModelSelector({
|
||||
(model) => (
|
||||
<ModelSelectorItem
|
||||
key={model.id}
|
||||
value={model.id}
|
||||
keywords={[
|
||||
model.modelId,
|
||||
providerLabel,
|
||||
]}
|
||||
value={
|
||||
model.modelId
|
||||
}
|
||||
onSelect={() =>
|
||||
handleSelect(
|
||||
model.id,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"use client"
|
||||
|
||||
import type React from "react"
|
||||
import { createContext, useContext, useRef, useState } from "react"
|
||||
import { createContext, useContext, useEffect, useRef, useState } from "react"
|
||||
import type { DrawIoEmbedRef, EventExport } from "react-drawio"
|
||||
import { toast } from "sonner"
|
||||
import type { ExportFormat } from "@/components/save-dialog"
|
||||
@@ -42,12 +42,6 @@ interface DiagramContextType {
|
||||
|
||||
const DiagramContext = createContext<DiagramContextType | undefined>(undefined)
|
||||
|
||||
// Exports for thumbnails, validation PNGs and file saves carry a tag in the
|
||||
// request's `message` field. draw.io echoes the request back in the export
|
||||
// event, so each result reaches its own caller; untagged exports (chat-panel's
|
||||
// onFetchChart) resolve resolverRef.
|
||||
type ExportTag = "thumbnail" | "validation"
|
||||
|
||||
export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
const [chartXML, setChartXML] = useState<string>("")
|
||||
const [latestSvg, setLatestSvg] = useState<string>("")
|
||||
@@ -59,10 +53,8 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
const hasCalledOnLoadRef = useRef(false)
|
||||
const drawioRef = useRef<DrawIoEmbedRef | null>(null)
|
||||
const resolverRef = useRef<((value: string) => void) | null>(null)
|
||||
// Pending thumbnail and validation PNG exports, keyed by their export tag
|
||||
const taggedResolversRef = useRef<
|
||||
Partial<Record<ExportTag, (value: string) => void>>
|
||||
>({})
|
||||
// Resolver for PNG export (used for VLM validation)
|
||||
const pngResolverRef = useRef<((value: string) => void) | null>(null)
|
||||
// Track if we're expecting an export for history (user-initiated)
|
||||
const expectHistoryExportRef = useRef<boolean>(false)
|
||||
// Track latest chartXML for restoration after remount
|
||||
@@ -84,12 +76,10 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
setIsDrawioReady(false)
|
||||
}
|
||||
|
||||
// Update chartXML and its ref together, so callbacks that read the ref
|
||||
// (export handler, autosave) see the new value right away
|
||||
const updateChartXML = (xml: string) => {
|
||||
chartXMLRef.current = xml
|
||||
setChartXML(xml)
|
||||
}
|
||||
// Keep chartXMLRef in sync with state for restoration after remount
|
||||
useEffect(() => {
|
||||
chartXMLRef.current = chartXML
|
||||
}, [chartXML])
|
||||
|
||||
// Track if we're expecting an export for file save (stores raw export data)
|
||||
const saveResolverRef = useRef<{
|
||||
@@ -116,52 +106,64 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
}
|
||||
}
|
||||
|
||||
// Export with a tag in `message` (draw.io echoes it back in the export
|
||||
// event) and wait for that result. Resolves to null on timeout, which is
|
||||
// expected occasionally.
|
||||
const requestTaggedExport = (
|
||||
tag: ExportTag,
|
||||
format: "xmlsvg" | "png",
|
||||
timeoutMs: number,
|
||||
) =>
|
||||
new Promise<string | null>((resolve) => {
|
||||
const finish = (value: string | null) => {
|
||||
clearTimeout(timer)
|
||||
if (taggedResolversRef.current[tag] === finish) {
|
||||
delete taggedResolversRef.current[tag]
|
||||
}
|
||||
resolve(value)
|
||||
}
|
||||
const timer = setTimeout(() => finish(null), timeoutMs)
|
||||
taggedResolversRef.current[tag] = finish
|
||||
drawioRef.current?.exportDiagram({ format, message: tag })
|
||||
})
|
||||
|
||||
// Get current diagram as SVG for thumbnail (used by session storage)
|
||||
const getThumbnailSvg = async (): Promise<string | null> => {
|
||||
if (!drawioRef.current) return null
|
||||
// Don't export if diagram is empty
|
||||
if (!isRealDiagram(chartXMLRef.current)) return null
|
||||
if (!isRealDiagram(chartXML)) return null
|
||||
|
||||
try {
|
||||
const svgData = await Promise.race([
|
||||
new Promise<string>((resolve) => {
|
||||
resolverRef.current = resolve
|
||||
drawioRef.current?.exportDiagram({ format: "xmlsvg" })
|
||||
}),
|
||||
new Promise<string>((_, reject) =>
|
||||
setTimeout(() => reject(new Error("Export timeout")), 3000),
|
||||
),
|
||||
])
|
||||
|
||||
// xmlsvg exports return an SVG data URL
|
||||
const svgData = await requestTaggedExport("thumbnail", "xmlsvg", 3000)
|
||||
if (svgData?.startsWith("data:image/svg")) {
|
||||
// Update latestSvg so it's available for future saves
|
||||
setLatestSvg(svgData)
|
||||
return svgData
|
||||
if (svgData?.includes("<svg")) {
|
||||
setLatestSvg(svgData)
|
||||
return svgData
|
||||
}
|
||||
return null
|
||||
} catch {
|
||||
// Timeout is expected occasionally - don't log as error
|
||||
return null
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
// Capture current diagram as PNG for VLM validation
|
||||
const captureValidationPng = async (): Promise<string | null> => {
|
||||
if (!drawioRef.current) return null
|
||||
// Don't export if diagram is empty
|
||||
if (!isRealDiagram(chartXMLRef.current)) return null
|
||||
if (!isRealDiagram(chartXML)) return null
|
||||
|
||||
const pngData = await requestTaggedExport("validation", "png", 5000)
|
||||
// PNG data should be a base64 data URL
|
||||
return pngData?.startsWith("data:image/png") ? pngData : null
|
||||
try {
|
||||
const pngData = await Promise.race([
|
||||
new Promise<string>((resolve) => {
|
||||
pngResolverRef.current = resolve
|
||||
drawioRef.current?.exportDiagram({ format: "png" })
|
||||
}),
|
||||
new Promise<string>((_, reject) =>
|
||||
setTimeout(
|
||||
() => reject(new Error("PNG export timeout")),
|
||||
5000,
|
||||
),
|
||||
),
|
||||
])
|
||||
|
||||
// PNG data should be a base64 data URL
|
||||
if (pngData?.startsWith("data:image/png")) {
|
||||
return pngData
|
||||
}
|
||||
return null
|
||||
} catch {
|
||||
// Timeout is expected occasionally - don't log as error
|
||||
return null
|
||||
}
|
||||
}
|
||||
|
||||
const loadDiagram = (
|
||||
@@ -191,7 +193,7 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
}
|
||||
|
||||
// Keep chartXML in sync even when diagrams are injected (e.g., display_diagram tool)
|
||||
updateChartXML(xmlToLoad)
|
||||
setChartXML(xmlToLoad)
|
||||
|
||||
if (drawioRef.current) {
|
||||
drawioRef.current.load({
|
||||
@@ -203,17 +205,24 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
}
|
||||
|
||||
const handleDiagramExport = (data: EventExport) => {
|
||||
// Tagged exports (thumbnail, validation PNG, file save) go only to
|
||||
// their own caller, so they never take the result meant for resolverRef
|
||||
const tag = data.message?.message
|
||||
if (tag === "thumbnail" || tag === "validation") {
|
||||
taggedResolversRef.current[tag]?.(data.data)
|
||||
// Handle PNG export for VLM validation
|
||||
if (pngResolverRef.current && data.data?.startsWith("data:image/png")) {
|
||||
pngResolverRef.current(data.data)
|
||||
pngResolverRef.current = null
|
||||
return
|
||||
}
|
||||
if (tag === "save") {
|
||||
saveResolverRef.current.resolver?.(data.data, data.xml)
|
||||
|
||||
// Handle save to file if requested (process raw data before extraction)
|
||||
if (saveResolverRef.current.resolver) {
|
||||
const format = saveResolverRef.current.format
|
||||
saveResolverRef.current.resolver(data.data, data.xml)
|
||||
saveResolverRef.current = { resolver: null, format: null }
|
||||
return
|
||||
// For non-xmlsvg formats, skip XML extraction as it will fail
|
||||
// Only drawio (which uses xmlsvg internally) has the content attribute
|
||||
// xmlsvg is saved directly as SVG file, no need for extraction
|
||||
if (format === "png" || format === "svg" || format === "xmlsvg") {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Don't write chartXML here: exports don't change the diagram, and
|
||||
@@ -227,15 +236,12 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
// Limit to 20 entries to prevent memory leaks during long sessions
|
||||
const MAX_HISTORY_SIZE = 20
|
||||
if (expectHistoryExportRef.current) {
|
||||
// Store the full multi-page document (extractedXML is only the
|
||||
// first page), so restoring a version keeps every page
|
||||
const historyXml = chartXMLRef.current || extractedXML
|
||||
setDiagramHistory((prev) => {
|
||||
const newHistory = [
|
||||
...prev,
|
||||
{
|
||||
svg: data.data,
|
||||
xml: historyXml,
|
||||
xml: extractedXML,
|
||||
},
|
||||
]
|
||||
// Keep only the last MAX_HISTORY_SIZE entries (circular buffer)
|
||||
@@ -250,16 +256,14 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
}
|
||||
}
|
||||
|
||||
// react-drawio registers this callback once per iframe mount, so it must
|
||||
// read refs: state captured in its closure would stay stale after a remount
|
||||
const handleDiagramAutoSave = (data: { xml?: string }) => {
|
||||
if (!data?.xml) return
|
||||
// Don't overwrite a pending restore - if we have a real diagram but
|
||||
// DrawIO hasn't loaded yet, it means we're waiting to restore
|
||||
if (!hasCalledOnLoadRef.current && isRealDiagram(chartXMLRef.current)) {
|
||||
// Don't overwrite a pending restore - if we have a real diagram in state
|
||||
// but DrawIO isn't ready yet, it means we're waiting to restore
|
||||
if (!isDrawioReady && isRealDiagram(chartXML)) {
|
||||
return
|
||||
}
|
||||
updateChartXML(data.xml)
|
||||
setChartXML(data.xml)
|
||||
}
|
||||
|
||||
const clearDiagram = () => {
|
||||
@@ -361,10 +365,7 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
|
||||
}
|
||||
|
||||
// Export diagram - callback will be handled in handleDiagramExport
|
||||
drawioRef.current.exportDiagram({
|
||||
format: drawioFormat,
|
||||
message: "save",
|
||||
})
|
||||
drawioRef.current.exportDiagram({ format: drawioFormat })
|
||||
}
|
||||
|
||||
// Log save event to Langfuse (just flags the trace, doesn't send content)
|
||||
|
||||
@@ -67,62 +67,41 @@ const MODEL_ALIASES: Record<string, string> = {
|
||||
"deepseek-v3-0324": "@tx/deepseek-ai/deepseek-v3-0324",
|
||||
}
|
||||
|
||||
const CORS_HEADERS = {
|
||||
"Access-Control-Allow-Origin": "*",
|
||||
"Access-Control-Allow-Methods": "POST, OPTIONS",
|
||||
"Access-Control-Allow-Headers": "Content-Type, Authorization",
|
||||
}
|
||||
|
||||
/**
|
||||
* Create standardized JSON response
|
||||
* Create standardized response with CORS headers
|
||||
*/
|
||||
function createResponse(body: any, status = 200, extraHeaders = {}): Response {
|
||||
return new Response(JSON.stringify(body), {
|
||||
status,
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
...CORS_HEADERS,
|
||||
...extraHeaders,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Only the app's own server (/api/chat, /api/validate-model) calls this
|
||||
// function, so no CORS headers are sent: other sites' pages can't call it
|
||||
// from a browser and spend the deployment's Edge AI quota.
|
||||
// Same rule as lib/access-code.ts, but reading the edge function's env.
|
||||
// No codes configured (or env unavailable) means no check.
|
||||
function hasValidAccessCode(request: Request, env: any): boolean {
|
||||
const accessCodes: string[] =
|
||||
env?.ACCESS_CODE_LIST?.split(",")
|
||||
.map((code: string) => code.trim())
|
||||
.filter(Boolean) || []
|
||||
if (accessCodes.length === 0) return true
|
||||
const accessCode = request.headers.get("x-access-code")
|
||||
return !!accessCode && accessCodes.includes(accessCode)
|
||||
/**
|
||||
* Handle OPTIONS request for CORS preflight
|
||||
*/
|
||||
function handleOptionsRequest(): Response {
|
||||
return new Response(null, {
|
||||
headers: {
|
||||
...CORS_HEADERS,
|
||||
"Access-Control-Max-Age": "86400",
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
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")
|
||||
) {
|
||||
return createResponse(
|
||||
{
|
||||
error: {
|
||||
message: "Expected a POST request with a JSON body",
|
||||
type: "invalid_request_error",
|
||||
},
|
||||
},
|
||||
400,
|
||||
)
|
||||
}
|
||||
|
||||
if (!hasValidAccessCode(request, env)) {
|
||||
return createResponse(
|
||||
{
|
||||
error: {
|
||||
message: "Invalid or missing access code",
|
||||
type: "invalid_request_error",
|
||||
},
|
||||
},
|
||||
401,
|
||||
)
|
||||
export async function onRequest({ request, env: _env }: any) {
|
||||
if (request.method === "OPTIONS") {
|
||||
return handleOptionsRequest()
|
||||
}
|
||||
|
||||
request.headers.delete("accept-encoding")
|
||||
@@ -174,7 +153,7 @@ export async function onRequest({ request, env }: any) {
|
||||
type: "invalid_request_error",
|
||||
},
|
||||
},
|
||||
400,
|
||||
429,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -237,6 +216,7 @@ export async function onRequest({ request, env }: any) {
|
||||
"Cache-Control": "no-cache, no-store, no-transform",
|
||||
"X-Accel-Buffering": "no",
|
||||
Connection: "keep-alive",
|
||||
...CORS_HEADERS,
|
||||
},
|
||||
})
|
||||
} catch (error: any) {
|
||||
|
||||
+26
-57
@@ -32,55 +32,6 @@ export function rebuildAppMenu(): void {
|
||||
buildAppMenu()
|
||||
}
|
||||
|
||||
/**
|
||||
* Apply a preset and restart the server so it takes effect.
|
||||
* If the restart fails, go back to the previous preset and restart again,
|
||||
* so the running server always matches the saved current preset.
|
||||
* Throws an error describing the outcome on failure.
|
||||
*/
|
||||
export async function switchPreset(
|
||||
id: string,
|
||||
): Promise<Record<string, string>> {
|
||||
const previousPresetId = getCurrentPresetId()
|
||||
const env = applyPresetToEnv(id)
|
||||
if (!env) {
|
||||
throw new Error("Preset not found")
|
||||
}
|
||||
rebuildAppMenu()
|
||||
|
||||
// In development, scripts/electron-dev.mjs restarts the Next.js dev server
|
||||
if (!app.isPackaged) {
|
||||
return env
|
||||
}
|
||||
|
||||
try {
|
||||
await restartNextServer()
|
||||
return env
|
||||
} catch (error) {
|
||||
console.error("Failed to restart server:", error)
|
||||
const reason = error instanceof Error ? error.message : String(error)
|
||||
|
||||
// Revert to previous preset on failure
|
||||
if (!previousPresetId || !applyPresetToEnv(previousPresetId)) {
|
||||
setCurrentPreset(null)
|
||||
}
|
||||
// Rebuild menu to restore previous checkmark state
|
||||
rebuildAppMenu()
|
||||
|
||||
try {
|
||||
await restartNextServer()
|
||||
} catch (retryError) {
|
||||
console.error("Failed to restart server again:", retryError)
|
||||
throw new Error(
|
||||
`The server could not be restarted.\n\nPlease restart the app.\n\nError: ${reason}`,
|
||||
)
|
||||
}
|
||||
throw new Error(
|
||||
`The server could not be restarted.\n\nThe previous configuration has been restored.\n\nError: ${reason}`,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Get the menu template with translations
|
||||
*/
|
||||
@@ -241,14 +192,32 @@ function buildConfigMenu(
|
||||
type: "radio",
|
||||
checked: preset.id === currentPresetId,
|
||||
click: async () => {
|
||||
try {
|
||||
await switchPreset(preset.id)
|
||||
} catch (error) {
|
||||
// Show error dialog to notify user
|
||||
dialog.showErrorBox(
|
||||
"Configuration Error",
|
||||
`Failed to apply preset "${preset.name}". ${error instanceof Error ? error.message : String(error)}`,
|
||||
)
|
||||
const previousPresetId = getCurrentPresetId()
|
||||
const env = applyPresetToEnv(preset.id)
|
||||
|
||||
if (env) {
|
||||
try {
|
||||
await restartNextServer()
|
||||
rebuildAppMenu() // Rebuild menu to update checkmarks
|
||||
} catch (error) {
|
||||
console.error("Failed to restart server:", error)
|
||||
|
||||
// Revert to previous preset on failure
|
||||
if (previousPresetId) {
|
||||
applyPresetToEnv(previousPresetId)
|
||||
} else {
|
||||
setCurrentPreset(null)
|
||||
}
|
||||
|
||||
// Rebuild menu to restore previous checkmark state
|
||||
rebuildAppMenu()
|
||||
|
||||
// Show error dialog to notify user
|
||||
dialog.showErrorBox(
|
||||
"Configuration Error",
|
||||
`Failed to apply preset "${preset.name}". The server could not be restarted.\n\nThe previous configuration has been restored.\n\nError: ${error instanceof Error ? error.message : String(error)}`,
|
||||
)
|
||||
}
|
||||
}
|
||||
},
|
||||
}))
|
||||
|
||||
+69
-123
@@ -1,11 +1,5 @@
|
||||
import { randomUUID } from "node:crypto"
|
||||
import {
|
||||
existsSync,
|
||||
mkdirSync,
|
||||
readFileSync,
|
||||
renameSync,
|
||||
writeFileSync,
|
||||
} from "node:fs"
|
||||
import { existsSync, mkdirSync, readFileSync, writeFileSync } from "node:fs"
|
||||
import path from "node:path"
|
||||
import { app, safeStorage } from "electron"
|
||||
|
||||
@@ -36,9 +30,7 @@ let hasWarnedAboutPlaintext = false
|
||||
* Warns if encryption is not available (API key stored in plaintext)
|
||||
*/
|
||||
function encryptValue(value: string): string {
|
||||
// Already encrypted (a value that could not be decrypted): keep it as is
|
||||
// instead of wrapping it in a second layer of encryption
|
||||
if (!value || value.startsWith(ENCRYPTED_PREFIX)) {
|
||||
if (!value) {
|
||||
return value
|
||||
}
|
||||
|
||||
@@ -69,7 +61,6 @@ function encryptValue(value: string): string {
|
||||
/**
|
||||
* Decrypt a sensitive value using safeStorage
|
||||
* Returns the original value if it's not encrypted or decryption fails
|
||||
* (so saving writes the stored ciphertext back unchanged)
|
||||
*/
|
||||
function decryptValue(value: string): string {
|
||||
if (!value || !value.startsWith(ENCRYPTED_PREFIX)) {
|
||||
@@ -188,15 +179,6 @@ export function loadPresets(): ConfigPresetsFile {
|
||||
return data
|
||||
} catch (error) {
|
||||
console.error("Failed to load config presets:", error)
|
||||
// Move the unreadable file aside so the next save can't overwrite
|
||||
// the user's presets with an empty list
|
||||
const backupPath = `${configPath}.corrupt-${Date.now()}`
|
||||
try {
|
||||
renameSync(configPath, backupPath)
|
||||
console.error(`Unreadable config presets moved to ${backupPath}`)
|
||||
} catch (renameError) {
|
||||
console.error("Failed to back up config presets:", renameError)
|
||||
}
|
||||
return {
|
||||
version: 1,
|
||||
currentPresetId: null,
|
||||
@@ -229,11 +211,7 @@ export function savePresets(data: ConfigPresetsFile): void {
|
||||
}
|
||||
|
||||
try {
|
||||
// Write a temp file and rename it, so a crash mid-write can't leave
|
||||
// a truncated config file
|
||||
const tempPath = `${configPath}.tmp`
|
||||
writeFileSync(tempPath, JSON.stringify(dataToSave, null, 2), "utf-8")
|
||||
renameSync(tempPath, configPath)
|
||||
writeFileSync(configPath, JSON.stringify(dataToSave, null, 2), "utf-8")
|
||||
} catch (error) {
|
||||
console.error("Failed to save config presets:", error)
|
||||
throw error
|
||||
@@ -329,10 +307,9 @@ export function deletePreset(id: string): boolean {
|
||||
|
||||
data.presets.splice(index, 1)
|
||||
|
||||
// Clear current preset (and its env vars) if it was deleted
|
||||
// Clear current preset if it was deleted
|
||||
if (data.currentPresetId === id) {
|
||||
data.currentPresetId = null
|
||||
setPresetEnv(null)
|
||||
}
|
||||
|
||||
savePresets(data)
|
||||
@@ -345,15 +322,13 @@ export function deletePreset(id: string): boolean {
|
||||
export function setCurrentPreset(id: string | null): boolean {
|
||||
const data = loadPresets()
|
||||
|
||||
let preset: ConfigPreset | null = null
|
||||
if (id !== null) {
|
||||
preset = data.presets.find((p) => p.id === id) || null
|
||||
const preset = data.presets.find((p) => p.id === id)
|
||||
if (!preset) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
setPresetEnv(preset)
|
||||
data.currentPresetId = id
|
||||
savePresets(data)
|
||||
return true
|
||||
@@ -390,23 +365,78 @@ const PROVIDER_ENV_MAP: Record<string, { apiKey: string; baseUrl: string }> = {
|
||||
}
|
||||
|
||||
/**
|
||||
* Map a preset's config to environment variables
|
||||
* Apply preset environment variables to the current process
|
||||
* Returns the environment variables that were applied
|
||||
*/
|
||||
export function applyPresetToEnv(id: string): Record<string, string> | null {
|
||||
const data = loadPresets()
|
||||
const preset = data.presets.find((p) => p.id === id)
|
||||
|
||||
if (!preset) {
|
||||
return null
|
||||
}
|
||||
|
||||
const appliedEnv: Record<string, string> = {}
|
||||
const provider = preset.config.AI_PROVIDER?.toLowerCase()
|
||||
|
||||
for (const [key, value] of Object.entries(preset.config)) {
|
||||
if (value !== undefined && value !== "") {
|
||||
// Map generic AI_API_KEY to provider-specific key
|
||||
if (
|
||||
key === "AI_API_KEY" &&
|
||||
provider &&
|
||||
PROVIDER_ENV_MAP[provider]
|
||||
) {
|
||||
const providerApiKey = PROVIDER_ENV_MAP[provider].apiKey
|
||||
if (providerApiKey) {
|
||||
process.env[providerApiKey] = value
|
||||
appliedEnv[providerApiKey] = value
|
||||
}
|
||||
}
|
||||
// Map generic AI_BASE_URL to provider-specific key
|
||||
else if (
|
||||
key === "AI_BASE_URL" &&
|
||||
provider &&
|
||||
PROVIDER_ENV_MAP[provider]
|
||||
) {
|
||||
const providerBaseUrl = PROVIDER_ENV_MAP[provider].baseUrl
|
||||
if (providerBaseUrl) {
|
||||
process.env[providerBaseUrl] = value
|
||||
appliedEnv[providerBaseUrl] = value
|
||||
}
|
||||
}
|
||||
// Apply other env vars directly
|
||||
else {
|
||||
process.env[key] = value
|
||||
appliedEnv[key] = value
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Set as current preset
|
||||
data.currentPresetId = id
|
||||
savePresets(data)
|
||||
|
||||
return appliedEnv
|
||||
}
|
||||
|
||||
/**
|
||||
* Get environment variables from current preset
|
||||
* Maps generic AI_API_KEY/AI_BASE_URL to provider-specific keys
|
||||
*/
|
||||
function presetToEnv(preset: ConfigPreset): Record<string, string> {
|
||||
export function getCurrentPresetEnv(): Record<string, string> {
|
||||
const preset = getCurrentPreset()
|
||||
if (!preset) {
|
||||
return {}
|
||||
}
|
||||
|
||||
const env: Record<string, string> = {}
|
||||
const provider = preset.config.AI_PROVIDER?.toLowerCase()
|
||||
|
||||
for (const [key, value] of Object.entries(preset.config)) {
|
||||
if (value !== undefined && value !== "") {
|
||||
// A key that could not be decrypted is useless to the server
|
||||
if (value.startsWith(ENCRYPTED_PREFIX)) {
|
||||
console.warn(
|
||||
`Preset "${preset.name}": ${key} could not be decrypted. Please enter it again in Settings.`,
|
||||
)
|
||||
}
|
||||
// Map generic AI_API_KEY to provider-specific key
|
||||
else if (
|
||||
if (
|
||||
key === "AI_API_KEY" &&
|
||||
provider &&
|
||||
PROVIDER_ENV_MAP[provider]
|
||||
@@ -436,90 +466,6 @@ function presetToEnv(preset: ConfigPreset): Record<string, string> {
|
||||
return env
|
||||
}
|
||||
|
||||
/**
|
||||
* Values that env vars had before a preset first set them
|
||||
* (from the system or .env files), and the keys the active preset set
|
||||
*/
|
||||
const originalEnv: Record<string, string | undefined> = {}
|
||||
let presetEnvKeys: string[] = []
|
||||
|
||||
/**
|
||||
* Replace the env vars of the previous preset with those of the given preset
|
||||
* (null leaves no preset applied). Restoring first means switching presets
|
||||
* never leaves the previous preset's base URL, model or key behind.
|
||||
*/
|
||||
function setPresetEnv(preset: ConfigPreset | null): Record<string, string> {
|
||||
for (const key of presetEnvKeys) {
|
||||
if (originalEnv[key] === undefined) {
|
||||
delete process.env[key]
|
||||
} else {
|
||||
process.env[key] = originalEnv[key]
|
||||
}
|
||||
}
|
||||
|
||||
const env = preset ? presetToEnv(preset) : {}
|
||||
for (const [key, value] of Object.entries(env)) {
|
||||
if (!(key in originalEnv)) {
|
||||
originalEnv[key] = process.env[key]
|
||||
}
|
||||
process.env[key] = value
|
||||
}
|
||||
presetEnvKeys = Object.keys(env)
|
||||
|
||||
writeDevPresetEnv(env)
|
||||
return env
|
||||
}
|
||||
|
||||
const DEV_ENV_FILE_NAME = "dev-preset-env.json"
|
||||
|
||||
/**
|
||||
* Development only: write the active preset's env vars (decrypted and mapped)
|
||||
* for scripts/electron-dev.mjs, which restarts the Next.js dev server when
|
||||
* this file changes. The dev server can't decrypt the config file itself.
|
||||
*/
|
||||
function writeDevPresetEnv(env: Record<string, string>): void {
|
||||
if (app.isPackaged) {
|
||||
return
|
||||
}
|
||||
try {
|
||||
const filePath = path.join(app.getPath("userData"), DEV_ENV_FILE_NAME)
|
||||
writeFileSync(filePath, JSON.stringify(env, null, 2), {
|
||||
encoding: "utf-8",
|
||||
mode: 0o600,
|
||||
})
|
||||
} catch (error) {
|
||||
console.error("Failed to write dev preset env:", error)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Apply preset environment variables to the current process
|
||||
* Returns the environment variables that were applied
|
||||
*/
|
||||
export function applyPresetToEnv(id: string): Record<string, string> | null {
|
||||
const data = loadPresets()
|
||||
const preset = data.presets.find((p) => p.id === id)
|
||||
|
||||
if (!preset) {
|
||||
return null
|
||||
}
|
||||
|
||||
const appliedEnv = setPresetEnv(preset)
|
||||
|
||||
// Set as current preset
|
||||
data.currentPresetId = id
|
||||
savePresets(data)
|
||||
|
||||
return appliedEnv
|
||||
}
|
||||
|
||||
/**
|
||||
* Apply the saved current preset's environment variables (used at startup)
|
||||
*/
|
||||
export function applyCurrentPresetToEnv(): void {
|
||||
setPresetEnv(getCurrentPreset())
|
||||
}
|
||||
|
||||
/**
|
||||
* Get user's preferred locale from config
|
||||
* Returns undefined if not set
|
||||
|
||||
@@ -48,16 +48,12 @@ function loadEnvFromFile(filePath: string): void {
|
||||
const key = trimmed.slice(0, equalIndex).trim()
|
||||
let value = trimmed.slice(equalIndex + 1).trim()
|
||||
|
||||
const quote = value[0]
|
||||
const closingQuote =
|
||||
quote === '"' || quote === "'" ? value.indexOf(quote, 1) : -1
|
||||
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)
|
||||
} else {
|
||||
// Unquoted value: drop an inline comment ("value # comment")
|
||||
value = value.replace(/\s+#.*$/, "")
|
||||
// Remove surrounding quotes
|
||||
if (
|
||||
(value.startsWith('"') && value.endsWith('"')) ||
|
||||
(value.startsWith("'") && value.endsWith("'"))
|
||||
) {
|
||||
value = value.slice(1, -1)
|
||||
}
|
||||
|
||||
// Don't override existing environment variables
|
||||
|
||||
+20
-48
@@ -1,17 +1,12 @@
|
||||
import { app, BrowserWindow, dialog, shell } from "electron"
|
||||
import { buildAppMenu } from "./app-menu"
|
||||
import { applyCurrentPresetToEnv } from "./config-manager"
|
||||
import { getCurrentPresetEnv } from "./config-manager"
|
||||
import { loadEnvFile } from "./env-loader"
|
||||
import { registerIpcHandlers } from "./ipc-handlers"
|
||||
import { startNextServer, stopNextServer } from "./next-server"
|
||||
import { applyProxyToEnv } from "./proxy-manager"
|
||||
import { registerSettingsWindowHandlers } from "./settings-window"
|
||||
import {
|
||||
createWindow,
|
||||
getAppUrl,
|
||||
getMainWindow,
|
||||
isAppUrl,
|
||||
} from "./window-manager"
|
||||
import { createWindow, getMainWindow } from "./window-manager"
|
||||
|
||||
// Single instance lock
|
||||
const gotTheLock = app.requestSingleInstanceLock()
|
||||
@@ -33,14 +28,16 @@ if (!gotTheLock) {
|
||||
// Apply proxy settings from saved config
|
||||
applyProxyToEnv()
|
||||
|
||||
const isDev = !app.isPackaged
|
||||
// Apply saved preset environment variables (overrides .env)
|
||||
const presetEnv = getCurrentPresetEnv()
|
||||
for (const [key, value] of Object.entries(presetEnv)) {
|
||||
process.env[key] = value
|
||||
}
|
||||
|
||||
const isDev = process.env.NODE_ENV === "development"
|
||||
let serverUrl: string | null = null
|
||||
|
||||
app.whenReady().then(async () => {
|
||||
// Apply saved preset environment variables (overrides .env).
|
||||
// Must run after ready: on Windows and Linux safeStorage can't
|
||||
// decrypt the API key before that.
|
||||
applyCurrentPresetToEnv()
|
||||
|
||||
// Register IPC handlers
|
||||
registerIpcHandlers()
|
||||
registerSettingsWindowHandlers()
|
||||
@@ -49,7 +46,6 @@ if (!gotTheLock) {
|
||||
buildAppMenu()
|
||||
|
||||
try {
|
||||
let serverUrl: string
|
||||
if (isDev) {
|
||||
// Development: use the dev server URL
|
||||
serverUrl =
|
||||
@@ -73,9 +69,8 @@ if (!gotTheLock) {
|
||||
|
||||
app.on("activate", () => {
|
||||
if (BrowserWindow.getAllWindows().length === 0) {
|
||||
const appUrl = getAppUrl()
|
||||
if (appUrl) {
|
||||
createWindow(appUrl)
|
||||
if (serverUrl) {
|
||||
createWindow(serverUrl)
|
||||
}
|
||||
}
|
||||
})
|
||||
@@ -92,47 +87,24 @@ if (!gotTheLock) {
|
||||
stopNextServer()
|
||||
})
|
||||
|
||||
// Pages allowed inside app windows: the app server and draw.io
|
||||
const isInAppUrl = (url: string): boolean => {
|
||||
if (isAppUrl(url)) return true
|
||||
try {
|
||||
const { hostname } = new URL(url)
|
||||
return ["diagrams.net", "draw.io"].some(
|
||||
(domain) =>
|
||||
hostname === domain || hostname.endsWith(`.${domain}`),
|
||||
)
|
||||
} catch {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
const isWebUrl = (url: string): boolean =>
|
||||
url.startsWith("http://") || url.startsWith("https://")
|
||||
|
||||
// Open external links in default browser
|
||||
app.on("web-contents-created", (_, contents) => {
|
||||
contents.setWindowOpenHandler(({ url }) => {
|
||||
if (isInAppUrl(url)) {
|
||||
// Allow diagrams.net iframe
|
||||
if (
|
||||
url.includes("diagrams.net") ||
|
||||
url.includes("draw.io") ||
|
||||
url.startsWith("http://localhost") ||
|
||||
url.startsWith("http://127.0.0.1")
|
||||
) {
|
||||
return { action: "allow" }
|
||||
}
|
||||
// Open other links in external browser
|
||||
if (isWebUrl(url)) {
|
||||
if (url.startsWith("http://") || url.startsWith("https://")) {
|
||||
shell.openExternal(url)
|
||||
return { action: "deny" }
|
||||
}
|
||||
return { action: "allow" }
|
||||
})
|
||||
|
||||
// Clicking a plain link would otherwise replace the app page with
|
||||
// an external site that keeps the preload API
|
||||
contents.on("will-navigate", (event) => {
|
||||
if (isInAppUrl(event.url)) {
|
||||
return
|
||||
}
|
||||
event.preventDefault()
|
||||
if (isWebUrl(event.url)) {
|
||||
shell.openExternal(event.url)
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,12 +1,7 @@
|
||||
import { app, BrowserWindow, dialog, ipcMain } from "electron"
|
||||
import { rebuildAppMenu } from "./app-menu"
|
||||
import {
|
||||
app,
|
||||
BrowserWindow,
|
||||
dialog,
|
||||
type IpcMainInvokeEvent,
|
||||
ipcMain,
|
||||
} from "electron"
|
||||
import { rebuildAppMenu, switchPreset } from "./app-menu"
|
||||
import {
|
||||
applyPresetToEnv,
|
||||
type ConfigPreset,
|
||||
createPreset,
|
||||
deletePreset,
|
||||
@@ -25,7 +20,6 @@ import {
|
||||
type ProxyConfig,
|
||||
saveProxyConfig,
|
||||
} from "./proxy-manager"
|
||||
import { isAppUrl } from "./window-manager"
|
||||
|
||||
/**
|
||||
* Allowed configuration keys for presets
|
||||
@@ -54,32 +48,13 @@ function sanitizePresetConfig(
|
||||
return sanitized
|
||||
}
|
||||
|
||||
/**
|
||||
* Register an IPC handler that only answers the app's own pages
|
||||
* (the main window on the app server, or the local settings page).
|
||||
* A main window that somehow ends up on an external site still gets the
|
||||
* preload API, so its calls must be rejected here.
|
||||
*/
|
||||
function handle<Args extends unknown[]>(
|
||||
channel: string,
|
||||
listener: (event: IpcMainInvokeEvent, ...args: Args) => unknown,
|
||||
): void {
|
||||
ipcMain.handle(channel, (event, ...args) => {
|
||||
const url = event.senderFrame?.url
|
||||
if (!isAppUrl(url) && !url?.startsWith("file://")) {
|
||||
throw new Error(`Blocked "${channel}" from untrusted page: ${url}`)
|
||||
}
|
||||
return listener(event, ...(args as Args))
|
||||
})
|
||||
}
|
||||
|
||||
/**
|
||||
* Register all IPC handlers
|
||||
*/
|
||||
export function registerIpcHandlers(): void {
|
||||
// ==================== App Info ====================
|
||||
|
||||
handle("get-version", () => {
|
||||
ipcMain.handle("get-version", () => {
|
||||
return app.getVersion()
|
||||
})
|
||||
|
||||
@@ -106,7 +81,7 @@ export function registerIpcHandlers(): void {
|
||||
|
||||
// ==================== File Dialogs ====================
|
||||
|
||||
handle("dialog-open-file", async (event) => {
|
||||
ipcMain.handle("dialog-open-file", async (event) => {
|
||||
const win = BrowserWindow.fromWebContents(event.sender)
|
||||
if (!win) return null
|
||||
|
||||
@@ -133,9 +108,9 @@ export function registerIpcHandlers(): void {
|
||||
}
|
||||
})
|
||||
|
||||
handle("dialog-save-file", async (event, data: string) => {
|
||||
ipcMain.handle("dialog-save-file", async (event, data: string) => {
|
||||
const win = BrowserWindow.fromWebContents(event.sender)
|
||||
if (!win || typeof data !== "string") return false
|
||||
if (!win) return false
|
||||
|
||||
const result = await dialog.showSaveDialog(win, {
|
||||
filters: [
|
||||
@@ -160,28 +135,28 @@ export function registerIpcHandlers(): void {
|
||||
|
||||
// ==================== Config Presets ====================
|
||||
|
||||
handle("config-presets:get-all", () => {
|
||||
ipcMain.handle("config-presets:get-all", () => {
|
||||
return getAllPresets()
|
||||
})
|
||||
|
||||
handle("config-presets:get-current", () => {
|
||||
ipcMain.handle("config-presets:get-current", () => {
|
||||
return getCurrentPreset()
|
||||
})
|
||||
|
||||
handle("config-presets:get-current-id", () => {
|
||||
ipcMain.handle("config-presets:get-current-id", () => {
|
||||
return getCurrentPresetId()
|
||||
})
|
||||
|
||||
handle(
|
||||
ipcMain.handle(
|
||||
"config-presets:save",
|
||||
async (
|
||||
(
|
||||
_event,
|
||||
preset: Omit<ConfigPreset, "id" | "createdAt" | "updatedAt"> & {
|
||||
id?: string
|
||||
},
|
||||
) => {
|
||||
// Validate preset name
|
||||
if (typeof preset?.name !== "string" || !preset.name.trim()) {
|
||||
if (typeof preset.name !== "string" || !preset.name.trim()) {
|
||||
throw new Error("Invalid preset name")
|
||||
}
|
||||
|
||||
@@ -190,48 +165,42 @@ export function registerIpcHandlers(): void {
|
||||
|
||||
if (preset.id) {
|
||||
// Update existing preset
|
||||
const updated = updatePreset(preset.id, {
|
||||
return updatePreset(preset.id, {
|
||||
name: preset.name.trim(),
|
||||
config: sanitizedConfig,
|
||||
})
|
||||
// Re-apply the active preset so the edit takes effect
|
||||
if (updated && updated.id === getCurrentPresetId()) {
|
||||
await switchPreset(updated.id)
|
||||
} else {
|
||||
rebuildAppMenu()
|
||||
}
|
||||
return updated
|
||||
}
|
||||
// Create new preset
|
||||
const created = createPreset({
|
||||
return createPreset({
|
||||
name: preset.name.trim(),
|
||||
config: sanitizedConfig,
|
||||
})
|
||||
rebuildAppMenu()
|
||||
return created
|
||||
},
|
||||
)
|
||||
|
||||
handle("config-presets:delete", async (_event, id: string) => {
|
||||
const wasCurrent = id === getCurrentPresetId()
|
||||
// Deleting the active preset also clears its env vars
|
||||
const deleted = deletePreset(id)
|
||||
rebuildAppMenu()
|
||||
|
||||
// Restart so the server stops using the deleted preset
|
||||
if (deleted && wasCurrent && app.isPackaged) {
|
||||
await restartNextServer()
|
||||
}
|
||||
return deleted
|
||||
ipcMain.handle("config-presets:delete", (_event, id: string) => {
|
||||
return deletePreset(id)
|
||||
})
|
||||
|
||||
handle("config-presets:apply", async (_event, id: string) => {
|
||||
ipcMain.handle("config-presets:apply", async (_event, id: string) => {
|
||||
const env = applyPresetToEnv(id)
|
||||
if (!env) {
|
||||
return { success: false, error: "Preset not found" }
|
||||
}
|
||||
|
||||
const isDev = process.env.NODE_ENV === "development"
|
||||
|
||||
if (isDev) {
|
||||
// In development mode, the config file change will trigger
|
||||
// the file watcher in electron-dev.mjs to restart Next.js
|
||||
// We just need to save the preset (already done in applyPresetToEnv)
|
||||
return { success: true, env, devMode: true }
|
||||
}
|
||||
|
||||
// Production mode: restart the Next.js server to apply new environment variables
|
||||
try {
|
||||
const env = await switchPreset(id)
|
||||
// In development mode, electron-dev.mjs restarts Next.js
|
||||
return app.isPackaged
|
||||
? { success: true, env }
|
||||
: { success: true, env, devMode: true }
|
||||
await restartNextServer()
|
||||
return { success: true, env }
|
||||
} catch (error) {
|
||||
return {
|
||||
success: false,
|
||||
@@ -243,39 +212,30 @@ export function registerIpcHandlers(): void {
|
||||
}
|
||||
})
|
||||
|
||||
handle("config-presets:set-current", (_event, id: string | null) => {
|
||||
return setCurrentPreset(id)
|
||||
})
|
||||
ipcMain.handle(
|
||||
"config-presets:set-current",
|
||||
(_event, id: string | null) => {
|
||||
return setCurrentPreset(id)
|
||||
},
|
||||
)
|
||||
|
||||
// ==================== Proxy Settings ====================
|
||||
|
||||
handle("get-proxy", () => {
|
||||
ipcMain.handle("get-proxy", () => {
|
||||
return getProxyConfig()
|
||||
})
|
||||
|
||||
handle("set-proxy", async (_event, config: ProxyConfig) => {
|
||||
const isOptionalString = (value: unknown) =>
|
||||
value === undefined || typeof value === "string"
|
||||
if (
|
||||
typeof config !== "object" ||
|
||||
config === null ||
|
||||
!isOptionalString(config.httpProxy) ||
|
||||
!isOptionalString(config.httpsProxy)
|
||||
) {
|
||||
return { success: false, error: "Invalid proxy settings" }
|
||||
}
|
||||
|
||||
ipcMain.handle("set-proxy", async (_event, config: ProxyConfig) => {
|
||||
try {
|
||||
// Save config to file
|
||||
saveProxyConfig({
|
||||
httpProxy: config.httpProxy,
|
||||
httpsProxy: config.httpsProxy,
|
||||
})
|
||||
saveProxyConfig(config)
|
||||
|
||||
// Apply to current process environment
|
||||
applyProxyToEnv()
|
||||
|
||||
if (!app.isPackaged) {
|
||||
const isDev = process.env.NODE_ENV === "development"
|
||||
|
||||
if (isDev) {
|
||||
// In development, env vars are already applied
|
||||
// Next.js dev server may need manual restart
|
||||
return { success: true, devMode: true }
|
||||
@@ -297,11 +257,11 @@ export function registerIpcHandlers(): void {
|
||||
|
||||
// ==================== User Locale ====================
|
||||
|
||||
handle("get-user-locale", () => {
|
||||
ipcMain.handle("get-user-locale", () => {
|
||||
return getUserLocale()
|
||||
})
|
||||
|
||||
handle("set-user-locale", (_event, locale: string) => {
|
||||
ipcMain.handle("set-user-locale", (_event, locale: string) => {
|
||||
// Validate locale is one of the supported values
|
||||
if (!["en", "zh", "ja", "zh-Hant"].includes(locale)) {
|
||||
return { success: false, error: "Invalid locale" }
|
||||
|
||||
@@ -6,22 +6,10 @@ import {
|
||||
getAllocatedPort,
|
||||
getServerUrl,
|
||||
isPortAvailable,
|
||||
saveServerPort,
|
||||
} from "./port-manager"
|
||||
import { setAppUrl } from "./window-manager"
|
||||
|
||||
let serverProcess: UtilityProcess | null = null
|
||||
|
||||
// Start and restart run one at a time, so overlapping calls (e.g. two quick
|
||||
// preset switches) can't leave two servers running
|
||||
let serverQueue: Promise<unknown> = Promise.resolve()
|
||||
|
||||
function runExclusive<T>(task: () => Promise<T>): Promise<T> {
|
||||
const result = serverQueue.then(task)
|
||||
serverQueue = result.catch(() => {})
|
||||
return result
|
||||
}
|
||||
|
||||
/**
|
||||
* Get the path to the standalone server resources
|
||||
* In packaged app: resources/standalone
|
||||
@@ -57,11 +45,7 @@ async function waitForServer(url: string, timeout = 30000): Promise<void> {
|
||||
* Start the Next.js standalone server using Electron's utilityProcess
|
||||
* This API is designed for running Node.js code in the background
|
||||
*/
|
||||
export function startNextServer(): Promise<string> {
|
||||
return runExclusive(startServer)
|
||||
}
|
||||
|
||||
async function startServer(): Promise<string> {
|
||||
export async function startNextServer(): Promise<string> {
|
||||
const resourcePath = getResourcePath()
|
||||
const serverPath = path.join(resourcePath, "server.js")
|
||||
|
||||
@@ -89,11 +73,6 @@ async function startServer(): Promise<string> {
|
||||
NODE_USE_ENV_PROXY: "1",
|
||||
}
|
||||
|
||||
// Keep requests to local model servers (e.g. Ollama) off the proxy
|
||||
if (!process.env.NO_PROXY && !process.env.no_proxy) {
|
||||
env.NO_PROXY = "localhost,127.0.0.1,[::1]"
|
||||
}
|
||||
|
||||
// Set cache directory to a writable location (user's app data folder)
|
||||
// This is necessary because the packaged app might be on a read-only volume
|
||||
if (app.isPackaged) {
|
||||
@@ -117,33 +96,28 @@ async function startServer(): Promise<string> {
|
||||
|
||||
// Use Electron's utilityProcess API for running Node.js in background
|
||||
// This is the recommended way to run Node.js code in Electron
|
||||
const proc = utilityProcess.fork(serverPath, [], {
|
||||
serverProcess = utilityProcess.fork(serverPath, [], {
|
||||
cwd: resourcePath,
|
||||
env,
|
||||
stdio: "pipe",
|
||||
})
|
||||
serverProcess = proc
|
||||
|
||||
proc.stdout?.on("data", (data) => {
|
||||
serverProcess.stdout?.on("data", (data) => {
|
||||
console.log(`[Next.js] ${data.toString().trim()}`)
|
||||
})
|
||||
|
||||
proc.stderr?.on("data", (data) => {
|
||||
serverProcess.stderr?.on("data", (data) => {
|
||||
console.error(`[Next.js Error] ${data.toString().trim()}`)
|
||||
})
|
||||
|
||||
proc.on("exit", (code) => {
|
||||
serverProcess.on("exit", (code) => {
|
||||
console.log(`Next.js server exited with code ${code}`)
|
||||
// An old server can exit after a new one started; keep the new one
|
||||
if (serverProcess === proc) {
|
||||
serverProcess = null
|
||||
}
|
||||
serverProcess = null
|
||||
})
|
||||
|
||||
const url = getServerUrl()
|
||||
await waitForServer(url)
|
||||
console.log(`Next.js server started at ${url}`)
|
||||
saveServerPort(port)
|
||||
|
||||
return url
|
||||
}
|
||||
@@ -152,36 +126,39 @@ async function startServer(): Promise<string> {
|
||||
* Stop the Next.js server process and wait for it to exit
|
||||
*/
|
||||
export async function stopNextServer(): Promise<void> {
|
||||
const proc = serverProcess
|
||||
if (!proc) {
|
||||
return
|
||||
}
|
||||
console.log("Stopping Next.js server...")
|
||||
serverProcess = null
|
||||
if (serverProcess) {
|
||||
console.log("Stopping Next.js server...")
|
||||
|
||||
// Resolves true when the process exits, false after the timeout
|
||||
const waitForExit = (ms: number) =>
|
||||
new Promise<boolean>((resolve) => {
|
||||
proc.once("exit", () => resolve(true))
|
||||
setTimeout(() => resolve(false), ms)
|
||||
// Create a promise that resolves when the process exits
|
||||
const exitPromise = new Promise<void>((resolve) => {
|
||||
const proc = serverProcess
|
||||
if (!proc) {
|
||||
resolve()
|
||||
return
|
||||
}
|
||||
|
||||
const onExit = () => {
|
||||
resolve()
|
||||
}
|
||||
|
||||
proc.once("exit", onExit)
|
||||
|
||||
// Timeout after 5 seconds
|
||||
setTimeout(() => {
|
||||
proc.removeListener("exit", onExit)
|
||||
resolve()
|
||||
}, 5000)
|
||||
})
|
||||
|
||||
proc.kill()
|
||||
serverProcess.kill()
|
||||
serverProcess = null
|
||||
|
||||
// Next.js waits for open requests (e.g. a streaming reply) before it
|
||||
// exits, so force kill it if it is still running after 5 seconds
|
||||
if (!(await waitForExit(5000)) && proc.pid) {
|
||||
console.warn("Next.js server did not exit in time, force killing it")
|
||||
try {
|
||||
process.kill(proc.pid, "SIGKILL")
|
||||
} catch (error) {
|
||||
console.error("Failed to force kill Next.js server:", error)
|
||||
}
|
||||
await waitForExit(2000)
|
||||
// Wait for process to exit
|
||||
await exitPromise
|
||||
|
||||
// Additional wait for OS to release port
|
||||
await new Promise((resolve) => setTimeout(resolve, 500))
|
||||
}
|
||||
|
||||
// Additional wait for OS to release port
|
||||
await new Promise((resolve) => setTimeout(resolve, 500))
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -207,19 +184,15 @@ async function waitForServerStop(timeout = 5000): Promise<void> {
|
||||
/**
|
||||
* Restart the Next.js server with new environment variables
|
||||
*/
|
||||
export function restartNextServer(): Promise<string> {
|
||||
return runExclusive(async () => {
|
||||
console.log("Restarting Next.js server...")
|
||||
export async function restartNextServer(): Promise<string> {
|
||||
console.log("Restarting Next.js server...")
|
||||
|
||||
// Stop the current server and wait for it to exit
|
||||
await stopNextServer()
|
||||
// Stop the current server and wait for it to exit
|
||||
await stopNextServer()
|
||||
|
||||
// Wait for the port to be released
|
||||
await waitForServerStop()
|
||||
// Wait for the port to be released
|
||||
await waitForServerStop()
|
||||
|
||||
// Start the server again, and follow it if it moved to another port
|
||||
const url = await startServer()
|
||||
setAppUrl(url)
|
||||
return url
|
||||
})
|
||||
// Start the server again
|
||||
return startNextServer()
|
||||
}
|
||||
|
||||
@@ -1,6 +1,4 @@
|
||||
import { readFileSync, writeFileSync } from "node:fs"
|
||||
import net from "node:net"
|
||||
import path from "node:path"
|
||||
import { app } from "electron"
|
||||
|
||||
/**
|
||||
@@ -25,38 +23,6 @@ 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
|
||||
*/
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Remember the port the production server started on
|
||||
*/
|
||||
export function saveServerPort(port: number): void {
|
||||
if (!app.isPackaged || port === loadSavedPort()) {
|
||||
return
|
||||
}
|
||||
try {
|
||||
writeFileSync(getSavedPortPath(), JSON.stringify({ port }), "utf-8")
|
||||
} catch (error) {
|
||||
console.error("Failed to save server port:", error)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if a specific port is available
|
||||
*/
|
||||
@@ -78,8 +44,7 @@ export function isPortAvailable(port: number): Promise<boolean> {
|
||||
/**
|
||||
* 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 fixed port (13370) to preserve localStorage
|
||||
* - Falls back to sequential ports if preferred port is unavailable
|
||||
* - Last resort: lets the OS assign a port (port 0)
|
||||
*
|
||||
@@ -104,20 +69,6 @@ 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 legacy port first to preserve existing users' localStorage
|
||||
if (!isDev) {
|
||||
const legacyPort = PORT_CONFIG.legacyProduction
|
||||
|
||||
@@ -13,22 +13,18 @@ function getConfigPath(): string {
|
||||
|
||||
/**
|
||||
* Load proxy configuration from JSON file
|
||||
* Returns null if the user never saved proxy settings (or the file is invalid)
|
||||
*/
|
||||
export function loadProxyConfig(): ProxyConfig | null {
|
||||
export function loadProxyConfig(): ProxyConfig {
|
||||
try {
|
||||
const configPath = getConfigPath()
|
||||
if (fs.existsSync(configPath)) {
|
||||
const data = JSON.parse(fs.readFileSync(configPath, "utf-8"))
|
||||
if (data && typeof data === "object" && !Array.isArray(data)) {
|
||||
return data as ProxyConfig
|
||||
}
|
||||
console.error("Ignoring invalid proxy config:", data)
|
||||
const data = fs.readFileSync(configPath, "utf-8")
|
||||
return JSON.parse(data) as ProxyConfig
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Failed to load proxy config:", error)
|
||||
}
|
||||
return null
|
||||
return {}
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -37,11 +33,7 @@ export function loadProxyConfig(): ProxyConfig | null {
|
||||
export function saveProxyConfig(config: ProxyConfig): void {
|
||||
try {
|
||||
const configPath = getConfigPath()
|
||||
// Write a temp file and rename it, so a crash mid-write can't leave
|
||||
// a truncated file
|
||||
const tempPath = `${configPath}.tmp`
|
||||
fs.writeFileSync(tempPath, JSON.stringify(config, null, 2), "utf-8")
|
||||
fs.renameSync(tempPath, configPath)
|
||||
fs.writeFileSync(configPath, JSON.stringify(config, null, 2), "utf-8")
|
||||
} catch (error) {
|
||||
console.error("Failed to save proxy config:", error)
|
||||
throw error
|
||||
@@ -55,11 +47,6 @@ export function saveProxyConfig(config: ProxyConfig): void {
|
||||
export function applyProxyToEnv(): void {
|
||||
const config = loadProxyConfig()
|
||||
|
||||
// No saved settings: keep proxy vars inherited from the system or .env
|
||||
if (!config) {
|
||||
return
|
||||
}
|
||||
|
||||
if (config.httpProxy) {
|
||||
process.env.HTTP_PROXY = config.httpProxy
|
||||
process.env.http_proxy = config.httpProxy
|
||||
|
||||
@@ -3,9 +3,6 @@ import { app, BrowserWindow, screen } from "electron"
|
||||
|
||||
let mainWindow: BrowserWindow | null = null
|
||||
|
||||
// URL of the app server the main window loads
|
||||
let appUrl: string | null = null
|
||||
|
||||
/**
|
||||
* Get the icon path based on platform
|
||||
* Note: electron-builder converts icon.png during packaging,
|
||||
@@ -31,7 +28,6 @@ function getIconPath(): string | undefined {
|
||||
* Create the main application window
|
||||
*/
|
||||
export function createWindow(serverUrl: string): BrowserWindow {
|
||||
appUrl = serverUrl
|
||||
const { width, height } = screen.getPrimaryDisplay().workAreaSize
|
||||
|
||||
mainWindow = new BrowserWindow({
|
||||
@@ -60,7 +56,7 @@ export function createWindow(serverUrl: string): BrowserWindow {
|
||||
})
|
||||
|
||||
// Open DevTools in development
|
||||
if (!app.isPackaged) {
|
||||
if (process.env.NODE_ENV === "development") {
|
||||
mainWindow.webContents.openDevTools()
|
||||
}
|
||||
|
||||
@@ -97,36 +93,3 @@ export function createWindow(serverUrl: string): BrowserWindow {
|
||||
export function getMainWindow(): BrowserWindow | null {
|
||||
return mainWindow
|
||||
}
|
||||
|
||||
/**
|
||||
* Get the app server URL the main window loads
|
||||
*/
|
||||
export function getAppUrl(): string | null {
|
||||
return appUrl
|
||||
}
|
||||
|
||||
/**
|
||||
* Point the main window at a new app server URL
|
||||
* (the restarted server can come up on a different port)
|
||||
*/
|
||||
export function setAppUrl(url: string): void {
|
||||
if (url === appUrl) {
|
||||
return
|
||||
}
|
||||
appUrl = url
|
||||
mainWindow?.loadURL(url)
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if a URL belongs to the app server (same origin)
|
||||
*/
|
||||
export function isAppUrl(url: string | undefined): boolean {
|
||||
if (!url || !appUrl) {
|
||||
return false
|
||||
}
|
||||
try {
|
||||
return new URL(url).origin === new URL(appUrl).origin
|
||||
} catch {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
@@ -213,9 +213,6 @@ async function savePreset() {
|
||||
}
|
||||
})
|
||||
|
||||
// closeModal() clears editingPresetId, so remember it for the toast
|
||||
const isEdit = Boolean(editingPresetId)
|
||||
|
||||
try {
|
||||
saveBtn.disabled = true
|
||||
saveBtn.innerHTML = '<span class="loading"></span>'
|
||||
@@ -223,7 +220,10 @@ async function savePreset() {
|
||||
await window.settingsAPI.savePreset(preset)
|
||||
await loadPresets()
|
||||
closeModal()
|
||||
showToast(isEdit ? "Preset updated" : "Preset created", "success")
|
||||
showToast(
|
||||
editingPresetId ? "Preset updated" : "Preset created",
|
||||
"success",
|
||||
)
|
||||
} catch (error) {
|
||||
console.error("Failed to save preset:", error)
|
||||
showToast("Failed to save preset", "error")
|
||||
@@ -265,6 +265,8 @@ async function applyPreset(id) {
|
||||
|
||||
const result = await window.settingsAPI.applyPreset(id)
|
||||
if (result.success) {
|
||||
currentPresetId = id
|
||||
renderPresets()
|
||||
showToast("Preset applied, server restarting...", "success")
|
||||
} else {
|
||||
showToast(result.error || "Failed to apply preset", "error")
|
||||
@@ -272,9 +274,6 @@ async function applyPreset(id) {
|
||||
} catch (error) {
|
||||
console.error("Failed to apply preset:", error)
|
||||
showToast("Failed to apply preset", "error")
|
||||
} finally {
|
||||
// Reload to show the active preset and reset the Apply button
|
||||
await loadPresets()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+1
-2
@@ -12,8 +12,7 @@ AI_PROVIDER=bedrock
|
||||
AI_MODEL=global.anthropic.claude-sonnet-4-5-20250929-v1:0
|
||||
|
||||
# Output limit, all providers (default: 64000). Shared by reasoning and the diagram XML,
|
||||
# so a thinking model can spend it all before the tool call. Users can lower it in Settings,
|
||||
# and raise it only when they use their own API key, so this also caps cost on server keys.
|
||||
# so a thinking model can spend it all before the tool call. Users can override it in Settings.
|
||||
# If a model's own ceiling is lower, the request is retried with that ceiling automatically.
|
||||
# MAX_OUTPUT_TOKENS=64000
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import type { MutableRefObject } from "react"
|
||||
import { useRef } from "react"
|
||||
import type { DiagramOperation } from "@/components/chat/types"
|
||||
import type {
|
||||
ValidationState,
|
||||
@@ -47,8 +48,6 @@ type ValidateDiagramFn = (
|
||||
interface UseDiagramToolHandlersParams {
|
||||
partialXmlRef: MutableRefObject<string>
|
||||
editDiagramOriginalXmlRef: MutableRefObject<Map<string, string>>
|
||||
// Failed VLM validations in the current user turn (reset on each user message)
|
||||
validationRetryCountRef: MutableRefObject<number>
|
||||
chartXMLRef: MutableRefObject<string>
|
||||
onDisplayChart: (xml: string, skipValidation?: boolean) => string | null
|
||||
onFetchChart: (saveToHistory?: boolean) => Promise<string>
|
||||
@@ -73,7 +72,6 @@ interface UseDiagramToolHandlersParams {
|
||||
export function useDiagramToolHandlers({
|
||||
partialXmlRef,
|
||||
editDiagramOriginalXmlRef,
|
||||
validationRetryCountRef,
|
||||
chartXMLRef,
|
||||
onDisplayChart,
|
||||
onFetchChart,
|
||||
@@ -84,6 +82,9 @@ export function useDiagramToolHandlers({
|
||||
sessionId,
|
||||
onValidationStateChange,
|
||||
}: UseDiagramToolHandlersParams) {
|
||||
// Track validation retry count per tool call
|
||||
const validationRetryCountRef = useRef<Map<string, number>>(new Map())
|
||||
|
||||
// Helper to update validation state
|
||||
const updateValidationState = (
|
||||
toolCallId: string,
|
||||
@@ -231,15 +232,17 @@ ${finalXml}
|
||||
)
|
||||
}
|
||||
|
||||
// Each retry is a new tool call, so count attempts per user turn
|
||||
const attempt = validationRetryCountRef.current + 1
|
||||
const retryCount =
|
||||
validationRetryCountRef.current.get(
|
||||
toolCall.toolCallId,
|
||||
) || 0
|
||||
|
||||
// Notify UI that we're validating (include the image)
|
||||
updateValidationState(
|
||||
toolCall.toolCallId,
|
||||
"validating",
|
||||
{
|
||||
attempt,
|
||||
attempt: retryCount + 1,
|
||||
maxAttempts: MAX_VALIDATION_RETRIES,
|
||||
imageData: capturedPngData,
|
||||
},
|
||||
@@ -251,14 +254,17 @@ ${finalXml}
|
||||
)
|
||||
|
||||
if (!result.valid) {
|
||||
if (attempt < MAX_VALIDATION_RETRIES) {
|
||||
validationRetryCountRef.current = attempt
|
||||
if (retryCount < MAX_VALIDATION_RETRIES) {
|
||||
validationRetryCountRef.current.set(
|
||||
toolCall.toolCallId,
|
||||
retryCount + 1,
|
||||
)
|
||||
|
||||
const feedback =
|
||||
formatValidationFeedback(result)
|
||||
if (DEBUG) {
|
||||
console.log(
|
||||
`[display_diagram] Validation failed (attempt ${attempt}/${MAX_VALIDATION_RETRIES}):`,
|
||||
`[display_diagram] Validation failed (attempt ${retryCount + 1}/${MAX_VALIDATION_RETRIES}):`,
|
||||
result.issues,
|
||||
)
|
||||
}
|
||||
@@ -268,7 +274,7 @@ ${finalXml}
|
||||
toolCall.toolCallId,
|
||||
"failed",
|
||||
{
|
||||
attempt,
|
||||
attempt: retryCount + 1,
|
||||
maxAttempts: MAX_VALIDATION_RETRIES,
|
||||
result,
|
||||
imageData: capturedPngData,
|
||||
@@ -279,17 +285,19 @@ ${finalXml}
|
||||
tool: "display_diagram",
|
||||
toolCallId: toolCall.toolCallId,
|
||||
state: "output-error",
|
||||
errorText: `[Validation attempt ${attempt}/${MAX_VALIDATION_RETRIES}]\n${feedback}`,
|
||||
errorText: `[Validation attempt ${retryCount + 1}/${MAX_VALIDATION_RETRIES}]\n${feedback}`,
|
||||
})
|
||||
return
|
||||
} else {
|
||||
// Last attempt - accept the diagram with warning
|
||||
// Max retries reached - accept the diagram with warning
|
||||
if (DEBUG) {
|
||||
console.log(
|
||||
"[display_diagram] Max validation retries reached, accepting diagram",
|
||||
)
|
||||
}
|
||||
validationRetryCountRef.current = 0
|
||||
validationRetryCountRef.current.delete(
|
||||
toolCall.toolCallId,
|
||||
)
|
||||
|
||||
// Notify UI that we're accepting with issues (include the image)
|
||||
updateValidationState(
|
||||
@@ -306,8 +314,10 @@ ${finalXml}
|
||||
return
|
||||
}
|
||||
} else {
|
||||
// Validation passed - reset retry count
|
||||
validationRetryCountRef.current = 0
|
||||
// Validation passed - clean up retry count
|
||||
validationRetryCountRef.current.delete(
|
||||
toolCall.toolCallId,
|
||||
)
|
||||
if (DEBUG) {
|
||||
console.log(
|
||||
"[display_diagram] Validation passed!",
|
||||
@@ -372,17 +382,12 @@ ${finalXml}
|
||||
}
|
||||
|
||||
let currentXml = ""
|
||||
// Use the original XML captured during streaming (shared with chat-message-display)
|
||||
// This ensures we apply operations to the same base XML that streaming used
|
||||
const originalXml = editDiagramOriginalXmlRef.current.get(
|
||||
toolCall.toolCallId,
|
||||
)
|
||||
// On failure, undo the streaming preview so the canvas matches the XML
|
||||
// reported back to the model
|
||||
const restoreOriginal = () => {
|
||||
if (originalXml) onDisplayChart(originalXml, true)
|
||||
}
|
||||
try {
|
||||
// Use the original XML captured during streaming (shared with chat-message-display)
|
||||
// This ensures we apply operations to the same base XML that streaming used
|
||||
const originalXml = editDiagramOriginalXmlRef.current.get(
|
||||
toolCall.toolCallId,
|
||||
)
|
||||
if (originalXml) {
|
||||
currentXml = originalXml
|
||||
} else {
|
||||
@@ -411,7 +416,6 @@ ${finalXml}
|
||||
)
|
||||
.join("\n")
|
||||
|
||||
restoreOriginal()
|
||||
addToolOutput({
|
||||
tool: "edit_diagram",
|
||||
toolCallId: toolCall.toolCallId,
|
||||
@@ -437,7 +441,6 @@ Please check the cell IDs and retry.`,
|
||||
"[edit_diagram] Validation error:",
|
||||
validationError,
|
||||
)
|
||||
restoreOriginal()
|
||||
addToolOutput({
|
||||
tool: "edit_diagram",
|
||||
toolCallId: toolCall.toolCallId,
|
||||
@@ -469,7 +472,6 @@ Please fix the operations to avoid structural issues.`,
|
||||
const errorMessage =
|
||||
error instanceof Error ? error.message : String(error)
|
||||
|
||||
restoreOriginal()
|
||||
addToolOutput({
|
||||
tool: "edit_diagram",
|
||||
toolCallId: toolCall.toolCallId,
|
||||
@@ -494,19 +496,6 @@ Please check cell IDs and retry, or use display_diagram to regenerate.`,
|
||||
) => {
|
||||
const { xml } = toolCall.input as { xml: string }
|
||||
|
||||
// Nothing to continue: loading the fragment alone would replace the whole diagram
|
||||
if (!partialXmlRef.current) {
|
||||
addToolOutput({
|
||||
tool: "append_diagram",
|
||||
toolCallId: toolCall.toolCallId,
|
||||
state: "output-error",
|
||||
errorText: `ERROR: There is no truncated diagram to continue, so append_diagram cannot be used now.
|
||||
|
||||
Use display_diagram to create the complete diagram, or edit_diagram to change the current one.`,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Detect if LLM incorrectly started fresh instead of continuing
|
||||
// LLM should only output bare mxCells now, so wrapper tags indicate error
|
||||
const trimmed = xml.trim()
|
||||
|
||||
+29
-57
@@ -101,15 +101,6 @@ function saveConfig(config: MultiModelConfig): void {
|
||||
localStorage.setItem(STORAGE_KEYS.modelConfigs, JSON.stringify(config))
|
||||
}
|
||||
|
||||
/**
|
||||
* Server model to fall back to: the one marked default, else the first one
|
||||
*/
|
||||
function defaultServerModelId(
|
||||
serverModels: FlattenedServerModel[],
|
||||
): string | undefined {
|
||||
return (serverModels.find((m) => m.isDefault) ?? serverModels[0])?.id
|
||||
}
|
||||
|
||||
export interface UseModelConfigReturn {
|
||||
// State
|
||||
config: MultiModelConfig
|
||||
@@ -153,16 +144,6 @@ export function useModelConfig(): UseModelConfigReturn {
|
||||
setIsLoaded(true)
|
||||
}, [])
|
||||
|
||||
// Pick up config changes saved by other tabs, so this tab neither shows a
|
||||
// stale model nor overwrites their changes on its next save
|
||||
useEffect(() => {
|
||||
const handleStorage = (e: StorageEvent) => {
|
||||
if (e.key === STORAGE_KEYS.modelConfigs) setConfig(loadConfig())
|
||||
}
|
||||
window.addEventListener("storage", handleStorage)
|
||||
return () => window.removeEventListener("storage", handleStorage)
|
||||
}, [])
|
||||
|
||||
// Load server models on mount (if any)
|
||||
useEffect(() => {
|
||||
if (typeof window === "undefined") return
|
||||
@@ -184,18 +165,17 @@ export function useModelConfig(): UseModelConfigReturn {
|
||||
setServerModels(raw)
|
||||
setServerLoaded(true)
|
||||
|
||||
// Auto-select the default server model if no model is selected,
|
||||
// or if the saved server model is gone (renamed or removed)
|
||||
// Auto-select default server model if no model is currently selected
|
||||
setConfig((prev) => {
|
||||
const id = prev.selectedModelId
|
||||
const isStale =
|
||||
id?.startsWith("server:") &&
|
||||
!raw.some((m) => m.id === id)
|
||||
if (id && !isStale) return prev
|
||||
const fallback = defaultServerModelId(raw)
|
||||
return fallback === id
|
||||
? prev
|
||||
: { ...prev, selectedModelId: fallback }
|
||||
if (!prev.selectedModelId && raw.length > 0) {
|
||||
const defaultModel = raw.find((m) => m.isDefault)
|
||||
if (defaultModel) {
|
||||
return { ...prev, selectedModelId: defaultModel.id }
|
||||
}
|
||||
// If no default marked, use first server model
|
||||
return { ...prev, selectedModelId: raw[0].id }
|
||||
}
|
||||
return prev
|
||||
})
|
||||
})
|
||||
.catch((error) => {
|
||||
@@ -280,31 +260,24 @@ export function useModelConfig(): UseModelConfigReturn {
|
||||
[],
|
||||
)
|
||||
|
||||
const deleteProvider = useCallback(
|
||||
(providerId: string) => {
|
||||
setConfig((prev) => {
|
||||
const provider = prev.providers.find((p) => p.id === providerId)
|
||||
const modelIds = provider?.models.map((m) => m.id) || []
|
||||
const deleteProvider = useCallback((providerId: string) => {
|
||||
setConfig((prev) => {
|
||||
const provider = prev.providers.find((p) => p.id === providerId)
|
||||
const modelIds = provider?.models.map((m) => m.id) || []
|
||||
|
||||
// Fall back to the default server model if the selected model
|
||||
// belongs to the deleted provider
|
||||
const newSelectedId =
|
||||
prev.selectedModelId &&
|
||||
modelIds.includes(prev.selectedModelId)
|
||||
? defaultServerModelId(serverModels)
|
||||
: prev.selectedModelId
|
||||
// Clear selected model if it belongs to deleted provider
|
||||
const newSelectedId =
|
||||
prev.selectedModelId && modelIds.includes(prev.selectedModelId)
|
||||
? undefined
|
||||
: prev.selectedModelId
|
||||
|
||||
return {
|
||||
...prev,
|
||||
providers: prev.providers.filter(
|
||||
(p) => p.id !== providerId,
|
||||
),
|
||||
selectedModelId: newSelectedId,
|
||||
}
|
||||
})
|
||||
},
|
||||
[serverModels],
|
||||
)
|
||||
return {
|
||||
...prev,
|
||||
providers: prev.providers.filter((p) => p.id !== providerId),
|
||||
selectedModelId: newSelectedId,
|
||||
}
|
||||
})
|
||||
}, [])
|
||||
|
||||
const addModel = useCallback(
|
||||
(providerId: string, modelId: string): ModelConfig => {
|
||||
@@ -361,15 +334,14 @@ export function useModelConfig(): UseModelConfigReturn {
|
||||
}
|
||||
: p,
|
||||
),
|
||||
// Fall back to the default server model if the selected model
|
||||
// was deleted
|
||||
// Clear selected model if it was deleted
|
||||
selectedModelId:
|
||||
prev.selectedModelId === modelConfigId
|
||||
? defaultServerModelId(serverModels)
|
||||
? undefined
|
||||
: prev.selectedModelId,
|
||||
}))
|
||||
},
|
||||
[serverModels],
|
||||
[],
|
||||
)
|
||||
|
||||
const resetConfig = useCallback(() => {
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
"use client"
|
||||
|
||||
import { useCallback, useEffect, useRef, useState } from "react"
|
||||
import { toast } from "sonner"
|
||||
import { useDictionary } from "@/hooks/use-dictionary"
|
||||
import {
|
||||
type ChatSession,
|
||||
createEmptySession,
|
||||
@@ -46,15 +44,6 @@ export interface UseSessionManagerReturn {
|
||||
clearCurrentSession: () => void
|
||||
}
|
||||
|
||||
// Reading the session list loads every stored session in full, and window
|
||||
// focus also fires each time the user clicks back from the draw.io iframe
|
||||
const FOCUS_REFRESH_INTERVAL_MS = 30_000
|
||||
|
||||
function notifySaveFailed(message: string) {
|
||||
// Same id, so repeated failures update one toast instead of stacking
|
||||
toast.error(message, { id: "session-save-failed", duration: 8000 })
|
||||
}
|
||||
|
||||
interface UseSessionManagerOptions {
|
||||
/** Session ID from URL param - if provided, load this session; if null, start blank */
|
||||
initialSessionId?: string | null
|
||||
@@ -64,7 +53,6 @@ export function useSessionManager(
|
||||
options: UseSessionManagerOptions = {},
|
||||
): UseSessionManagerReturn {
|
||||
const { initialSessionId } = options
|
||||
const dict = useDictionary()
|
||||
const [sessions, setSessions] = useState<SessionMetadata[]>([])
|
||||
const [currentSessionId, setCurrentSessionId] = useState<string | null>(
|
||||
null,
|
||||
@@ -175,15 +163,9 @@ export function useSessionManager(
|
||||
handleSessionIdChange()
|
||||
}, [initialSessionId, isAvailable])
|
||||
|
||||
// Refresh sessions on window focus (multi-tab sync), at most once per interval
|
||||
const lastFocusRefreshRef = useRef(0)
|
||||
// Refresh sessions on window focus (multi-tab sync)
|
||||
useEffect(() => {
|
||||
const handleFocus = () => {
|
||||
const now = Date.now()
|
||||
if (now - lastFocusRefreshRef.current < FOCUS_REFRESH_INTERVAL_MS) {
|
||||
return
|
||||
}
|
||||
lastFocusRefreshRef.current = now
|
||||
refreshSessions()
|
||||
}
|
||||
window.addEventListener("focus", handleFocus)
|
||||
@@ -256,8 +238,6 @@ export function useSessionManager(
|
||||
) {
|
||||
return
|
||||
}
|
||||
// Nothing can be stored without IndexedDB
|
||||
if (!isIndexedDBAvailable()) return
|
||||
|
||||
if (!currentSession) {
|
||||
// Create a new session if none exists
|
||||
@@ -270,12 +250,7 @@ export function useSessionManager(
|
||||
diagramHistory: data.diagramHistory,
|
||||
title: extractTitle(data.messages),
|
||||
}
|
||||
// Without a stored session, keep no session id (it would end
|
||||
// up in the URL and point to nothing after a reload)
|
||||
if (!(await saveSession(newSession))) {
|
||||
notifySaveFailed(dict.errors.sessionSaveFailed)
|
||||
return
|
||||
}
|
||||
await saveSession(newSession)
|
||||
await enforceSessionLimit()
|
||||
setCurrentSession(newSession)
|
||||
setCurrentSessionId(newSession.id)
|
||||
@@ -302,10 +277,7 @@ export function useSessionManager(
|
||||
: currentSession.title,
|
||||
}
|
||||
|
||||
if (!(await saveSession(updatedSession))) {
|
||||
notifySaveFailed(dict.errors.sessionSaveFailed)
|
||||
return
|
||||
}
|
||||
await saveSession(updatedSession)
|
||||
setCurrentSession(updatedSession)
|
||||
|
||||
// Update sessions list metadata
|
||||
@@ -326,7 +298,7 @@ export function useSessionManager(
|
||||
),
|
||||
)
|
||||
},
|
||||
[currentSession, currentSessionId, refreshSessions, dict],
|
||||
[currentSession, currentSessionId, refreshSessions],
|
||||
)
|
||||
|
||||
// Clear current session state (for starting fresh without loading another session)
|
||||
|
||||
@@ -6,7 +6,6 @@
|
||||
|
||||
import { experimental_useObject as useObject } from "@ai-sdk/react"
|
||||
import { useCallback, useRef } from "react"
|
||||
import { getSelectedAIConfig } from "@/hooks/use-model-config"
|
||||
import { getApiEndpoint } from "@/lib/base-path"
|
||||
import {
|
||||
type ValidationResult,
|
||||
@@ -40,8 +39,6 @@ export function useValidateDiagram(options: UseValidateDiagramOptions = {}) {
|
||||
const { object, submit, isLoading, error, stop } = useObject({
|
||||
api: getApiEndpoint("/api/validate-diagram"),
|
||||
schema: ValidationResultSchema,
|
||||
// Resolved per request so a changed access code is picked up
|
||||
headers: () => ({ "x-access-code": getSelectedAIConfig().accessCode }),
|
||||
onFinish: ({
|
||||
object,
|
||||
error: finishError,
|
||||
|
||||
@@ -1,22 +0,0 @@
|
||||
/**
|
||||
* Check the x-access-code header against ACCESS_CODE_LIST.
|
||||
* Returns a 401 response to send back when the check fails, or null when the
|
||||
* request may continue (including when no access codes are configured).
|
||||
*/
|
||||
export function checkAccessCode(req: Request): Response | null {
|
||||
const accessCodes =
|
||||
process.env.ACCESS_CODE_LIST?.split(",")
|
||||
.map((code) => code.trim())
|
||||
.filter(Boolean) || []
|
||||
if (accessCodes.length === 0) return null
|
||||
|
||||
const accessCodeHeader = req.headers.get("x-access-code")
|
||||
if (accessCodeHeader && accessCodes.includes(accessCodeHeader)) return null
|
||||
|
||||
return Response.json(
|
||||
{
|
||||
error: "Invalid or missing access code. Please configure it in Settings.",
|
||||
},
|
||||
{ status: 401 },
|
||||
)
|
||||
}
|
||||
+7
-19
@@ -2,7 +2,6 @@ import { z } from "zod"
|
||||
import {
|
||||
ProviderNameSchema,
|
||||
type ServerModelsConfig,
|
||||
slugify,
|
||||
} from "@/lib/server-model-config"
|
||||
import {
|
||||
FIXED_CRED_PROVIDERS,
|
||||
@@ -183,15 +182,12 @@ export function validateAdminProviders(
|
||||
return `${PROVIDER_INFO[single].label} is already configured in AI_MODELS_CONFIG / ai-models.json and shares global credentials. Manage it via the environment configuration instead.`
|
||||
}
|
||||
}
|
||||
// Server model ids are built from the slugified name, so names must
|
||||
// stay distinct after slugifying ("OpenAI" and "openai" would collide)
|
||||
const names = list.map((p) => displayName(p))
|
||||
const slugs = names.map(slugify)
|
||||
if (new Set(slugs).size !== slugs.length) {
|
||||
return "Provider display names must be unique (ignoring case and punctuation)."
|
||||
if (new Set(names).size !== names.length) {
|
||||
return "Provider display names must be unique."
|
||||
}
|
||||
const envSlugs = new Set(envProviders.map((p) => slugify(p.name)))
|
||||
const clash = names.find((_, i) => envSlugs.has(slugs[i]))
|
||||
const envNames = new Set(envProviders.map((p) => p.name))
|
||||
const clash = names.find((n) => envNames.has(n))
|
||||
if (clash) {
|
||||
return `"${clash}" is already defined in AI_MODELS_CONFIG / ai-models.json. Use a different display name.`
|
||||
}
|
||||
@@ -244,14 +240,10 @@ export function deriveEnvUpdates(
|
||||
indexByProvider.set(p.provider, index + 1)
|
||||
|
||||
if (p.provider === "bedrock") {
|
||||
// ADMIN_ names keep the standard AWS_* vars untouched, so other
|
||||
// AWS clients (e.g. the DynamoDB quota table) keep their own
|
||||
// credentials instead of picking up the panel's Bedrock keys
|
||||
if (p.awsAccessKeyId)
|
||||
updates.ADMIN_AWS_ACCESS_KEY_ID = p.awsAccessKeyId
|
||||
if (p.awsAccessKeyId) updates.AWS_ACCESS_KEY_ID = p.awsAccessKeyId
|
||||
if (p.awsSecretAccessKey)
|
||||
updates.ADMIN_AWS_SECRET_ACCESS_KEY = p.awsSecretAccessKey
|
||||
if (p.awsRegion) updates.ADMIN_AWS_REGION = p.awsRegion
|
||||
updates.AWS_SECRET_ACCESS_KEY = p.awsSecretAccessKey
|
||||
if (p.awsRegion) updates.AWS_REGION = p.awsRegion
|
||||
} else if (p.provider === "vertexai") {
|
||||
if (p.vertexApiKey) updates.GOOGLE_VERTEX_API_KEY = p.vertexApiKey
|
||||
if (p.baseUrl) updates.GOOGLE_VERTEX_BASE_URL = p.baseUrl
|
||||
@@ -292,10 +284,6 @@ function derivedEnvKeys(list: StoredAdminProvider[]): string[] {
|
||||
const index = indexByProvider.get(p.provider) ?? 0
|
||||
indexByProvider.set(p.provider, index + 1)
|
||||
if (p.provider === "bedrock") {
|
||||
keys.add("ADMIN_AWS_ACCESS_KEY_ID")
|
||||
keys.add("ADMIN_AWS_SECRET_ACCESS_KEY")
|
||||
keys.add("ADMIN_AWS_REGION")
|
||||
// Written by older versions; listed so the next save clears them
|
||||
keys.add("AWS_ACCESS_KEY_ID")
|
||||
keys.add("AWS_SECRET_ACCESS_KEY")
|
||||
keys.add("AWS_REGION")
|
||||
|
||||
+19
-34
@@ -10,27 +10,13 @@ interface SettingsFile {
|
||||
values: Record<string, string>
|
||||
}
|
||||
|
||||
interface SettingsState {
|
||||
// Original env values snapshotted before the first overlay, so removing
|
||||
// a key from the settings file restores the env default. null = was unset.
|
||||
originalEnv: Record<string, string | null>
|
||||
// Keys currently overlaid, so we can restore ones removed from the file.
|
||||
overlaidKeys: Set<string>
|
||||
cachedSettings: Record<string, string> | null
|
||||
}
|
||||
// Original env values snapshotted before the first overlay, so removing a
|
||||
// key from the settings file restores the env default. null = was unset.
|
||||
const originalEnv: Record<string, string | null> = {}
|
||||
// Keys currently overlaid, so we can restore ones removed from the file.
|
||||
let overlaidKeys = new Set<string>()
|
||||
|
||||
// Kept on globalThis because the build can load this module more than once
|
||||
// (instrumentation.ts and the API routes get separate copies); per-module
|
||||
// state would make a route forget what instrumentation overlaid at startup.
|
||||
const globalState = globalThis as typeof globalThis & {
|
||||
__adminSettingsState?: SettingsState
|
||||
}
|
||||
globalState.__adminSettingsState ??= {
|
||||
originalEnv: {},
|
||||
overlaidKeys: new Set(),
|
||||
cachedSettings: null,
|
||||
}
|
||||
const state = globalState.__adminSettingsState
|
||||
let cachedSettings: Record<string, string> | null = null
|
||||
|
||||
export function getSettingsPath(): string {
|
||||
const custom = process.env.SETTINGS_FILE
|
||||
@@ -39,7 +25,7 @@ export function getSettingsPath(): string {
|
||||
}
|
||||
|
||||
export function loadSettings(): Record<string, string> {
|
||||
if (state.cachedSettings) return state.cachedSettings
|
||||
if (cachedSettings) return cachedSettings
|
||||
try {
|
||||
const raw = fs.readFileSync(getSettingsPath(), "utf8")
|
||||
const parsed = JSON.parse(raw) as SettingsFile
|
||||
@@ -57,22 +43,21 @@ export function loadSettings(): Record<string, string> {
|
||||
for (const [key, value] of Object.entries(rawValues)) {
|
||||
if (typeof value === "string") values[key] = value
|
||||
}
|
||||
state.cachedSettings = values
|
||||
cachedSettings = values
|
||||
} catch (err: any) {
|
||||
if (err?.code !== "ENOENT") {
|
||||
console.error("[admin-settings] Failed to read settings file:", err)
|
||||
}
|
||||
state.cachedSettings = {}
|
||||
cachedSettings = {}
|
||||
}
|
||||
return state.cachedSettings
|
||||
return cachedSettings
|
||||
}
|
||||
|
||||
export function applyToEnv(): void {
|
||||
const values = loadSettings()
|
||||
const { originalEnv } = state
|
||||
|
||||
// Restore env for keys that were overlaid before but are now gone
|
||||
for (const key of state.overlaidKeys) {
|
||||
for (const key of overlaidKeys) {
|
||||
if (!(key in values)) {
|
||||
const original = originalEnv[key]
|
||||
if (original === null) delete process.env[key]
|
||||
@@ -87,12 +72,12 @@ export function applyToEnv(): void {
|
||||
process.env[key] = value
|
||||
}
|
||||
|
||||
state.overlaidKeys = new Set(Object.keys(values))
|
||||
overlaidKeys = new Set(Object.keys(values))
|
||||
}
|
||||
|
||||
// The effective env value if the file entry were removed (for fallback display)
|
||||
export function getEnvFallback(key: string): string | null {
|
||||
if (state.overlaidKeys.has(key)) return state.originalEnv[key] ?? null
|
||||
if (overlaidKeys.has(key)) return originalEnv[key] ?? null
|
||||
return process.env[key] ?? null
|
||||
}
|
||||
|
||||
@@ -116,7 +101,7 @@ export function saveSettings(updates: Record<string, string | null>): void {
|
||||
fs.writeFileSync(tmpPath, JSON.stringify(data, null, 2), { mode: 0o600 })
|
||||
fs.renameSync(tmpPath, filePath)
|
||||
|
||||
state.cachedSettings = current
|
||||
cachedSettings = current
|
||||
applyToEnv()
|
||||
}
|
||||
|
||||
@@ -137,13 +122,13 @@ export function isSettingsWritable(): boolean {
|
||||
|
||||
// Test-only: reset module state
|
||||
export function _resetForTests(): void {
|
||||
state.cachedSettings = null
|
||||
cachedSettings = null
|
||||
writableCache = null
|
||||
for (const key of state.overlaidKeys) {
|
||||
const original = state.originalEnv[key]
|
||||
for (const key of overlaidKeys) {
|
||||
const original = originalEnv[key]
|
||||
if (original === null) delete process.env[key]
|
||||
else if (original !== undefined) process.env[key] = original
|
||||
}
|
||||
state.overlaidKeys = new Set()
|
||||
state.originalEnv = {}
|
||||
overlaidKeys = new Set()
|
||||
for (const key of Object.keys(originalEnv)) delete originalEnv[key]
|
||||
}
|
||||
|
||||
+16
-85
@@ -10,10 +10,6 @@ import { aihubmix, createAihubmix } from "@aihubmix/ai-sdk-provider"
|
||||
import { fromNodeProviderChain } from "@aws-sdk/credential-providers"
|
||||
import { createOpenRouter } from "@openrouter/ai-sdk-provider"
|
||||
import { createOllama, ollama } from "ollama-ai-provider-v2"
|
||||
import {
|
||||
adminProvidersToConfig,
|
||||
loadAdminProviders,
|
||||
} from "@/lib/admin/providers"
|
||||
import { PROVIDER_INFO, type ProviderName } from "@/lib/types/model-config"
|
||||
|
||||
export type { ProviderName }
|
||||
@@ -828,16 +824,8 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
// Use client-provided credentials if available, otherwise fall back to IAM/env vars
|
||||
const hasClientCredentials =
|
||||
overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey
|
||||
// Keys from the admin panel. The ADMIN_ names keep them out of the
|
||||
// default AWS credential chain, which other clients such as the
|
||||
// DynamoDB quota manager use with their own credentials.
|
||||
const adminAccessKeyId = process.env.ADMIN_AWS_ACCESS_KEY_ID
|
||||
const adminSecretAccessKey = process.env.ADMIN_AWS_SECRET_ACCESS_KEY
|
||||
const bedrockRegion =
|
||||
overrides?.awsRegion ||
|
||||
process.env.ADMIN_AWS_REGION ||
|
||||
process.env.AWS_REGION ||
|
||||
"us-west-2"
|
||||
overrides?.awsRegion || process.env.AWS_REGION || "us-west-2"
|
||||
|
||||
const bedrockProvider = hasClientCredentials
|
||||
? createAmazonBedrock({
|
||||
@@ -848,16 +836,10 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
sessionToken: overrides.awsSessionToken,
|
||||
}),
|
||||
})
|
||||
: adminAccessKeyId && adminSecretAccessKey
|
||||
? createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
accessKeyId: adminAccessKeyId,
|
||||
secretAccessKey: adminSecretAccessKey,
|
||||
})
|
||||
: createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
credentialProvider: fromNodeProviderChain(),
|
||||
})
|
||||
: createAmazonBedrock({
|
||||
region: bedrockRegion,
|
||||
credentialProvider: fromNodeProviderChain(),
|
||||
})
|
||||
model = bedrockProvider(modelId)
|
||||
// Add Anthropic beta options if using Claude models via Bedrock
|
||||
if (modelId.includes("anthropic.claude")) {
|
||||
@@ -890,9 +872,8 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
// for compatibility (most proxies don't support /responses endpoint)
|
||||
const customOpenAI = createOpenAI({ apiKey, baseURL })
|
||||
model = customOpenAI.chat(modelId)
|
||||
} else if (overrides?.apiKey || overrides?.apiKeyEnv) {
|
||||
// Custom API key (the client's, or a server model's own env var)
|
||||
// but official OpenAI endpoint, use Responses API
|
||||
} else if (overrides?.apiKey) {
|
||||
// Custom API key but official OpenAI endpoint, use Responses API
|
||||
// to support reasoning for gpt-5, o1, o3, o4 models
|
||||
const customOpenAI = createOpenAI({ apiKey })
|
||||
model = customOpenAI(modelId)
|
||||
@@ -947,9 +928,7 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
overrides?.baseUrl,
|
||||
serverBaseUrl,
|
||||
)
|
||||
// The default instance only reads GOOGLE_GENERATIVE_AI_API_KEY, so a
|
||||
// server model's own env var (apiKeyEnv) needs a custom instance too
|
||||
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
|
||||
if (baseURL || overrides?.apiKey) {
|
||||
const customGoogle = createGoogleGenerativeAI({
|
||||
apiKey,
|
||||
...(baseURL && { baseURL }),
|
||||
@@ -962,11 +941,8 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
}
|
||||
case "vertexai": {
|
||||
// Express Mode: Use API key for authentication
|
||||
// SECURITY: a client base URL only ever gets the client's key, so the
|
||||
// server's GOOGLE_VERTEX_API_KEY is never sent to a client-chosen host
|
||||
const vertexApiKey = overrides?.baseUrl
|
||||
? overrides.vertexApiKey
|
||||
: overrides?.vertexApiKey || process.env.GOOGLE_VERTEX_API_KEY
|
||||
const vertexApiKey =
|
||||
overrides?.vertexApiKey || process.env.GOOGLE_VERTEX_API_KEY
|
||||
|
||||
if (!vertexApiKey) {
|
||||
throw new Error(
|
||||
@@ -975,13 +951,9 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
)
|
||||
}
|
||||
|
||||
// Support custom base URL from env or client override.
|
||||
// A client key only goes to the client's URL or the official one.
|
||||
const baseURL = resolveBaseURL(
|
||||
overrides?.vertexApiKey,
|
||||
overrides?.baseUrl,
|
||||
process.env.GOOGLE_VERTEX_BASE_URL,
|
||||
)
|
||||
// Support custom base URL from env or client override
|
||||
const baseURL =
|
||||
overrides?.baseUrl || process.env.GOOGLE_VERTEX_BASE_URL
|
||||
|
||||
const vertexProvider = createVertex({
|
||||
apiKey: vertexApiKey,
|
||||
@@ -1107,7 +1079,7 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
overrides?.baseUrl,
|
||||
serverBaseUrl,
|
||||
)
|
||||
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
|
||||
if (baseURL || overrides?.apiKey) {
|
||||
const customDeepSeek = createDeepSeek({
|
||||
apiKey,
|
||||
...(baseURL && { baseURL }),
|
||||
@@ -1269,7 +1241,7 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
)
|
||||
// Only use custom configuration if explicitly set (local dev or custom Gateway)
|
||||
// Otherwise undefined → AI SDK uses Vercel default (https://ai-gateway.vercel.sh/v1/ai) + OIDC
|
||||
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
|
||||
if (baseURL || overrides?.apiKey) {
|
||||
const customGateway = createGateway({
|
||||
apiKey,
|
||||
...(baseURL && { baseURL }),
|
||||
@@ -1458,36 +1430,6 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
|
||||
return { model, providerOptions, headers, modelId, provider }
|
||||
}
|
||||
|
||||
/**
|
||||
* Whether the call is paid for by the server's own credentials (env keys or
|
||||
* IAM role) rather than credentials sent with the request. Mirrors which key
|
||||
* each branch of getAIModel ends up using.
|
||||
*/
|
||||
export function usesServerCredentials(
|
||||
provider: ProviderName,
|
||||
overrides?: ClientOverrides,
|
||||
): boolean {
|
||||
switch (provider) {
|
||||
case "bedrock":
|
||||
return !(overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey)
|
||||
case "vertexai":
|
||||
return !overrides?.vertexApiKey
|
||||
case "edgeone":
|
||||
// The platform's own endpoint, no key involved
|
||||
return false
|
||||
case "ollama":
|
||||
// Only a server key costs money; a keyless local server or the
|
||||
// client's own server does not
|
||||
return (
|
||||
!overrides?.baseUrl &&
|
||||
!overrides?.apiKey &&
|
||||
!!(overrides?.apiKeyEnv || process.env.OLLAMA_API_KEY)
|
||||
)
|
||||
default:
|
||||
return !overrides?.apiKey
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if a model supports prompt caching.
|
||||
* Currently only Claude models on Bedrock support prompt caching.
|
||||
@@ -1522,17 +1464,6 @@ export function getValidationModel(): ReturnType<typeof getAIModel>["model"] {
|
||||
)
|
||||
}
|
||||
|
||||
// A default set in the admin panel becomes AI_PROVIDER/AI_MODEL, but its key
|
||||
// lives in an ADMIN_-prefixed env var. Point at it the way the chat route
|
||||
// does for server models, or the standard env var is required instead.
|
||||
const panelDefault = adminProvidersToConfig(
|
||||
loadAdminProviders(),
|
||||
).providers.find((p) => p.default && p.provider === process.env.AI_PROVIDER)
|
||||
|
||||
const { model } = getAIModel({
|
||||
modelId,
|
||||
apiKeyEnv: panelDefault?.apiKeyEnv,
|
||||
baseUrlEnv: panelDefault?.baseUrlEnv,
|
||||
})
|
||||
const { model } = getAIModel({ modelId })
|
||||
return model
|
||||
}
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
export interface CachedResponse {
|
||||
promptText: string
|
||||
hasImage: boolean
|
||||
// Name of the bundled example file the prompt is sent with
|
||||
fileName?: string
|
||||
xml: string
|
||||
}
|
||||
|
||||
@@ -256,7 +254,6 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
|
||||
{
|
||||
promptText: "Replicate this in aws style",
|
||||
hasImage: true,
|
||||
fileName: "architecture.png",
|
||||
xml: `<mxCell id="2" value="AWS" style="sketch=0;outlineConnect=0;gradientColor=none;html=1;whiteSpace=wrap;fontSize=12;fontStyle=0;container=1;pointerEvents=0;collapsible=0;recursiveResize=0;shape=mxgraph.aws4.group;grIcon=mxgraph.aws4.group_aws_cloud;strokeColor=#232F3E;fillColor=none;verticalAlign=top;align=left;spacingLeft=30;fontColor=#232F3E;dashed=0;rounded=1;arcSize=5;" vertex="1" parent="1">
|
||||
<mxGeometry x="340" y="40" width="880" height="520" as="geometry"/>
|
||||
</mxCell>
|
||||
@@ -321,7 +318,6 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
|
||||
{
|
||||
promptText: "Replicate this flowchart.",
|
||||
hasImage: true,
|
||||
fileName: "example.png",
|
||||
xml: `<mxCell id="2" value="Lamp doesn't work" style="rounded=1;whiteSpace=wrap;html=1;fillColor=#ffcccc;strokeColor=#000000;strokeWidth=2;fontSize=18;fontStyle=0;" vertex="1" parent="1">
|
||||
<mxGeometry x="140" y="40" width="180" height="60" as="geometry"/>
|
||||
</mxCell>
|
||||
@@ -383,7 +379,6 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
|
||||
{
|
||||
promptText: "Summarize this paper as a diagram",
|
||||
hasImage: true,
|
||||
fileName: "chain-of-thought.txt",
|
||||
xml: `<mxCell id="title_bg" parent="1"
|
||||
style="rounded=1;whiteSpace=wrap;html=1;fillColor=#1a237e;strokeColor=none;arcSize=8;"
|
||||
value="" vertex="1">
|
||||
@@ -884,19 +879,14 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
|
||||
},
|
||||
]
|
||||
|
||||
// Examples that come with a file only match when that exact example file is
|
||||
// attached, so a user's own file with the same prompt still goes to the model.
|
||||
// Callers that can't tell file names (the server) only get text-only examples.
|
||||
export function findCachedResponse(
|
||||
promptText: string,
|
||||
hasImage: boolean,
|
||||
fileName?: string,
|
||||
): CachedResponse | undefined {
|
||||
return CACHED_EXAMPLE_RESPONSES.find(
|
||||
(c) =>
|
||||
c.promptText === promptText &&
|
||||
c.hasImage === hasImage &&
|
||||
(!c.fileName || c.fileName === fileName) &&
|
||||
c.xml !== "",
|
||||
)
|
||||
}
|
||||
|
||||
+43
-96
@@ -6,37 +6,25 @@ export const MAX_FILE_SIZE = 2 * 1024 * 1024 // 2MB
|
||||
export const MAX_FILES = 5
|
||||
|
||||
// Helper function to validate file parts in messages
|
||||
// Checks every message, since history is sent to the model too
|
||||
export function validateFileParts(messages: any[]): {
|
||||
valid: boolean
|
||||
error?: string
|
||||
} {
|
||||
for (const message of messages) {
|
||||
const fileParts =
|
||||
message?.parts?.filter((p: any) => p.type === "file") || []
|
||||
const lastMessage = messages[messages.length - 1]
|
||||
const fileParts =
|
||||
lastMessage?.parts?.filter((p: any) => p.type === "file") || []
|
||||
|
||||
if (fileParts.length > MAX_FILES) {
|
||||
return {
|
||||
valid: false,
|
||||
error: `Too many files. Maximum ${MAX_FILES} allowed.`,
|
||||
}
|
||||
if (fileParts.length > MAX_FILES) {
|
||||
return {
|
||||
valid: false,
|
||||
error: `Too many files. Maximum ${MAX_FILES} allowed.`,
|
||||
}
|
||||
}
|
||||
|
||||
for (const filePart of fileParts) {
|
||||
// The client sends files inline. Any other URL would be downloaded
|
||||
// by the server (AI SDK does that for models without URL support).
|
||||
if (
|
||||
typeof filePart.url !== "string" ||
|
||||
!filePart.url.startsWith("data:")
|
||||
) {
|
||||
return {
|
||||
valid: false,
|
||||
error: "Files must be uploaded inline as data URLs.",
|
||||
}
|
||||
}
|
||||
|
||||
// Data URLs format: data:image/png;base64,<data>
|
||||
// Base64 increases size by ~33%, so we check the decoded size
|
||||
for (const filePart of fileParts) {
|
||||
// Data URLs format: data:image/png;base64,<data>
|
||||
// Base64 increases size by ~33%, so we check the decoded size
|
||||
if (filePart.url?.startsWith("data:")) {
|
||||
const base64Data = filePart.url.split(",")[1]
|
||||
if (base64Data) {
|
||||
const sizeInBytes = Math.ceil((base64Data.length * 3) / 4)
|
||||
@@ -54,89 +42,48 @@ export function validateFileParts(messages: any[]): {
|
||||
}
|
||||
|
||||
// Helper function to check if diagram is minimal/empty
|
||||
// Empty means no mxCell besides the root cells "0" and "1". Cells drawn in
|
||||
// draw.io get random ids, so checking for id="2" is not enough.
|
||||
export function isMinimalDiagram(xml: string): boolean {
|
||||
return !/<mxCell\b[^>]*\bid="(?![01]")/.test(xml)
|
||||
}
|
||||
|
||||
// A tool-call input providers accept: a non-empty JSON object
|
||||
function isValidToolInput(input: unknown): boolean {
|
||||
return !!input && typeof input === "object" && Object.keys(input).length > 0
|
||||
const stripped = xml.replace(/\s/g, "")
|
||||
return !stripped.includes('id="2"')
|
||||
}
|
||||
|
||||
// Helper function to replace historical tool call XML with placeholders
|
||||
// This reduces token usage and forces LLM to rely on the current diagram XML (source of truth)
|
||||
// Tool calls with invalid inputs are left for dropInvalidToolCalls to remove
|
||||
// Also fixes invalid/undefined inputs from interrupted streaming
|
||||
export function replaceHistoricalToolInputs(messages: any[]): any[] {
|
||||
return messages.map((msg) => {
|
||||
if (msg.role !== "assistant" || !Array.isArray(msg.content)) {
|
||||
return msg
|
||||
}
|
||||
const replacedContent = msg.content.map((part: any) => {
|
||||
if (
|
||||
part.type === "tool-call" &&
|
||||
isValidToolInput(part.input) &&
|
||||
(part.toolName === "display_diagram" ||
|
||||
part.toolName === "edit_diagram")
|
||||
) {
|
||||
return {
|
||||
...part,
|
||||
input: {
|
||||
placeholder:
|
||||
"[XML content replaced - see current diagram XML in system context]",
|
||||
},
|
||||
const replacedContent = msg.content
|
||||
.map((part: any) => {
|
||||
if (part.type === "tool-call") {
|
||||
const toolName = part.toolName
|
||||
// Fix invalid/undefined inputs from interrupted streaming
|
||||
if (
|
||||
!part.input ||
|
||||
typeof part.input !== "object" ||
|
||||
Object.keys(part.input).length === 0
|
||||
) {
|
||||
// Skip tool calls with invalid inputs entirely
|
||||
return null
|
||||
}
|
||||
if (
|
||||
toolName === "display_diagram" ||
|
||||
toolName === "edit_diagram"
|
||||
) {
|
||||
return {
|
||||
...part,
|
||||
input: {
|
||||
placeholder:
|
||||
"[XML content replaced - see current diagram XML in system context]",
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return part
|
||||
})
|
||||
return part
|
||||
})
|
||||
.filter(Boolean) // Remove null entries (invalid tool calls)
|
||||
return { ...msg, content: replacedContent }
|
||||
})
|
||||
}
|
||||
|
||||
// Remove tool-calls with invalid inputs (from failed repair or interrupted streaming),
|
||||
// together with their tool-results: providers reject a result whose call is missing.
|
||||
// Messages left empty are removed too (Bedrock rejects empty content arrays).
|
||||
export function dropInvalidToolCalls(messages: any[]): any[] {
|
||||
const droppedIds = new Set<string>()
|
||||
return messages
|
||||
.map((msg) => {
|
||||
if (!Array.isArray(msg.content)) return msg
|
||||
const content = msg.content.filter((part: any) => {
|
||||
if (
|
||||
msg.role === "assistant" &&
|
||||
part.type === "tool-call" &&
|
||||
!isValidToolInput(part.input)
|
||||
) {
|
||||
console.warn(
|
||||
`[chat-helpers] Dropping tool-call with invalid input:`,
|
||||
{ toolName: part.toolName, input: part.input },
|
||||
)
|
||||
droppedIds.add(part.toolCallId)
|
||||
return false
|
||||
}
|
||||
// Results always come after their call, so the id is known by now
|
||||
return !(
|
||||
part.type === "tool-result" &&
|
||||
droppedIds.has(part.toolCallId)
|
||||
)
|
||||
})
|
||||
return { ...msg, content }
|
||||
})
|
||||
.filter((msg) => !Array.isArray(msg.content) || msg.content.length > 0)
|
||||
}
|
||||
|
||||
// Fix common LLM JSON mistakes in tool-call input before jsonrepair runs
|
||||
export function fixToolInputJson(input: string): string {
|
||||
return (
|
||||
input
|
||||
// Inconsistent quote escaping in XML attributes inside JSON strings:
|
||||
// y="-20\" (opening quote unescaped, closing escaped) becomes y=\"-20\".
|
||||
// Must run before the key fix below, which would rewrite the `="`.
|
||||
.replace(/(\w+)="([^"]*?)\\"/g, '$1=\\"$2\\"')
|
||||
// `:=` instead of `: `
|
||||
.replace(/:=/g, ": ")
|
||||
// `"key"= "` instead of `"key": "`, only for JSON keys
|
||||
.replace(/"(\w+)"\s*=\s*"/g, '"$1": "')
|
||||
)
|
||||
}
|
||||
|
||||
@@ -188,8 +188,7 @@
|
||||
"failedToExport": "Error fetching chart data",
|
||||
"failedToLoadExample": "Error loading example image",
|
||||
"failedToRecordFeedback": "Failed to record your feedback. Please try again.",
|
||||
"storageUpdateFailed": "Chat cleared but browser storage could not be updated",
|
||||
"sessionSaveFailed": "Could not save this chat. Browser storage may be full: delete old chats from history and try again."
|
||||
"storageUpdateFailed": "Chat cleared but browser storage could not be updated"
|
||||
},
|
||||
"quota": {
|
||||
"dailyLimit": "Daily Quota Reached",
|
||||
|
||||
@@ -188,8 +188,7 @@
|
||||
"failedToExport": "チャートデータの取得エラー",
|
||||
"failedToLoadExample": "例の画像の読み込みエラー",
|
||||
"failedToRecordFeedback": "フィードバックの記録に失敗しました。もう一度お試しください。",
|
||||
"storageUpdateFailed": "チャットはクリアされましたが、ブラウザストレージを更新できませんでした",
|
||||
"sessionSaveFailed": "このチャットを保存できませんでした。ブラウザのストレージがいっぱいの可能性があります。履歴から古いチャットを削除して、もう一度お試しください。"
|
||||
"storageUpdateFailed": "チャットはクリアされましたが、ブラウザストレージを更新できませんでした"
|
||||
},
|
||||
"quota": {
|
||||
"dailyLimit": "1日の割当量に達しました",
|
||||
|
||||
@@ -188,8 +188,7 @@
|
||||
"failedToExport": "取得圖表資料時出錯",
|
||||
"failedToLoadExample": "載入範例圖片時出錯",
|
||||
"failedToRecordFeedback": "記錄您的回饋失敗。請重試。",
|
||||
"storageUpdateFailed": "聊天已清除,但無法更新瀏覽器儲存空間",
|
||||
"sessionSaveFailed": "無法儲存這個對話。瀏覽器儲存空間可能已滿,請在歷史紀錄裡刪除舊對話後重試。"
|
||||
"storageUpdateFailed": "聊天已清除,但無法更新瀏覽器儲存空間"
|
||||
},
|
||||
"quota": {
|
||||
"dailyLimit": "已達每日配額",
|
||||
|
||||
@@ -188,8 +188,7 @@
|
||||
"failedToExport": "获取图表数据时出错",
|
||||
"failedToLoadExample": "加载示例图片时出错",
|
||||
"failedToRecordFeedback": "记录您的反馈失败。请重试。",
|
||||
"storageUpdateFailed": "聊天已清除,但无法更新浏览器存储",
|
||||
"sessionSaveFailed": "无法保存这个对话。浏览器存储空间可能已满,请在历史记录里删除旧对话后重试。"
|
||||
"storageUpdateFailed": "聊天已清除,但无法更新浏览器存储"
|
||||
},
|
||||
"quota": {
|
||||
"dailyLimit": "已达每日配额",
|
||||
|
||||
+1
-8
@@ -51,15 +51,8 @@ export function setTraceOutput(output: string) {
|
||||
if (!isLangfuseEnabled()) return
|
||||
|
||||
updateActiveTrace({ output })
|
||||
endTrace()
|
||||
}
|
||||
|
||||
// End the observe() wrapper span (AI SDK creates its own child spans with usage).
|
||||
// It uses endOnExit: false, so every request path has to end it, or the trace
|
||||
// is never exported: stream finish, stream error/abort, and early returns.
|
||||
export function endTrace() {
|
||||
if (!isLangfuseEnabled()) return
|
||||
|
||||
// End the observe() wrapper span (AI SDK creates its own child spans with usage)
|
||||
const activeSpan = api.trace.getActiveSpan()
|
||||
if (activeSpan) {
|
||||
activeSpan.end()
|
||||
|
||||
+38
-117
@@ -22,12 +22,6 @@ export const MAX_OUTPUT_TOKENS_LIMIT = 200000
|
||||
*/
|
||||
const MIN_USABLE_OUTPUT_TOKENS = 1024
|
||||
|
||||
/**
|
||||
* Retry budget when a rejection names the budget parameter but no number we can
|
||||
* read. It is the default from before 64000, which these providers ran with.
|
||||
*/
|
||||
const FALLBACK_OUTPUT_TOKENS = 16000
|
||||
|
||||
/** Status codes that can carry a complaint about the requested budget. */
|
||||
const BUDGET_REJECTION_STATUSES = new Set([400, 422])
|
||||
|
||||
@@ -35,8 +29,24 @@ function usableLimit(value: number): number | null {
|
||||
return value >= MIN_USABLE_OUTPUT_TOKENS ? value : null
|
||||
}
|
||||
|
||||
/** Message and body of an error that may be about the budget, or null. */
|
||||
function rejectionText(error: unknown): string | null {
|
||||
/**
|
||||
* A budget this large exceeds what some models accept. Providers reject it with a
|
||||
* 400 that names the real limit, so we parse the number out and retry once
|
||||
* instead of failing the turn.
|
||||
*
|
||||
* Formats seen in the wild:
|
||||
* - Bedrock: "The maximum tokens you requested exceeds the model limit of 4096."
|
||||
* - OpenRouter: "This endpoint's maximum context length is 64000 tokens. However,
|
||||
* you requested about 64025 tokens (25 of text input, 64000 in the output)."
|
||||
* Note this one is an input+output ceiling, so the input has to be subtracted.
|
||||
* - Anthropic: "max_tokens: 200000 > 64000, which is the maximum allowed..."
|
||||
* - OpenAI: "This model supports at most 16384 completion tokens"
|
||||
*
|
||||
* Every pattern names tokens explicitly. A generic one (an earlier draft matched
|
||||
* "lower than N") would reinterpret unrelated failures, and retrying on a bogus
|
||||
* number turns a readable error into an empty diagram.
|
||||
*/
|
||||
export function parseOutputTokenLimit(error: unknown): number | null {
|
||||
const err = error as {
|
||||
message?: unknown
|
||||
responseBody?: unknown
|
||||
@@ -56,109 +66,24 @@ function rejectionText(error: unknown): string | null {
|
||||
typeof err?.responseBody === "string" ? err.responseBody : "",
|
||||
].join(" ")
|
||||
|
||||
return text.trim() ? text : null
|
||||
}
|
||||
if (!text) return null
|
||||
|
||||
/**
|
||||
* A budget this large exceeds what some models accept. Providers reject it with a
|
||||
* 400 that names the real limit, so we parse the number out and retry once
|
||||
* instead of failing the turn.
|
||||
*
|
||||
* Formats seen in the wild:
|
||||
* - Bedrock: "The maximum tokens you requested exceeds the model limit of 4096."
|
||||
* - OpenRouter: "This endpoint's maximum context length is 64000 tokens. However,
|
||||
* you requested about 64025 tokens (25 of text input, 64000 in the output)."
|
||||
* Note this one is an input+output ceiling, so the input has to be subtracted.
|
||||
* vLLM and SGLang send the same kind of ceiling, with the input written as
|
||||
* "6000 in the messages", "has 6000 input tokens" or "6000 tokens from the input".
|
||||
* - Anthropic: "max_tokens: 200000 > 64000, which is the maximum allowed..."
|
||||
* - OpenAI: "This model supports at most 16384 completion tokens"
|
||||
* - Volcengine Ark: "The parameter `max_tokens` specified in the request are not
|
||||
* valid: integer above maximum value, expected a value <= 32768, but got 64000"
|
||||
* - DashScope: "Range of max_tokens should be [1, 8192]"
|
||||
*
|
||||
* Every pattern names tokens explicitly. A generic one (an earlier draft matched
|
||||
* "lower than N") would reinterpret unrelated failures, and retrying on a bogus
|
||||
* number turns a readable error into an empty diagram.
|
||||
*/
|
||||
function readCeiling(text: string): number | null {
|
||||
// Combined input+output ceiling: subtract the input the provider counted,
|
||||
// plus a small margin because its estimate is approximate.
|
||||
const context = text.match(/maximum context length (?:is|of) (\d+)/i)
|
||||
const context = text.match(/maximum context length is (\d+)/i)
|
||||
if (context) {
|
||||
const input =
|
||||
text.match(/(\d+) of text input/i) ||
|
||||
text.match(/(\d+) in the messages/i) ||
|
||||
text.match(/(\d+) tokens from the input/i) ||
|
||||
text.match(/(\d+) input tokens/i)
|
||||
return Number(context[1]) - (input ? Number(input[1]) : 0) - 1024
|
||||
const input = text.match(/(\d+) of text input/i)
|
||||
return usableLimit(
|
||||
Number(context[1]) - (input ? Number(input[1]) : 0) - 1024,
|
||||
)
|
||||
}
|
||||
|
||||
const output =
|
||||
text.match(/model limit of (\d+)/i) ||
|
||||
text.match(/> (\d+), which is the maximum/i) ||
|
||||
text.match(/at most (\d+) completion tokens/i) ||
|
||||
text.match(/max_\w*tokens.*?expected a value (?:<=|\\u003c=) (\d+)/i) ||
|
||||
text.match(/Range of max_tokens should be \[1,\s*(\d+)\]/i)
|
||||
text.match(/at most (\d+) completion tokens/i)
|
||||
|
||||
return output ? Number(output[1]) : null
|
||||
}
|
||||
|
||||
/** The usable output ceiling named in a rejection, or null. */
|
||||
export function parseOutputTokenLimit(error: unknown): number | null {
|
||||
const text = rejectionText(error)
|
||||
const ceiling = text ? readCeiling(text) : null
|
||||
return ceiling === null ? null : usableLimit(ceiling)
|
||||
}
|
||||
|
||||
/**
|
||||
* Thinking budget the provider adds on top of maxOutputTokens. Bedrock and
|
||||
* Anthropic send maxOutputTokens + budgetTokens as max_tokens, so a ceiling in
|
||||
* their rejection covers both.
|
||||
*/
|
||||
function thinkingBudget(providerOptions: unknown): number {
|
||||
const options = providerOptions as
|
||||
| {
|
||||
bedrock?: {
|
||||
reasoningConfig?: { type?: string; budgetTokens?: unknown }
|
||||
}
|
||||
anthropic?: {
|
||||
thinking?: { type?: string; budgetTokens?: unknown }
|
||||
}
|
||||
}
|
||||
| undefined
|
||||
const config =
|
||||
options?.bedrock?.reasoningConfig ?? options?.anthropic?.thinking
|
||||
return config?.type === "enabled" && typeof config.budgetTokens === "number"
|
||||
? config.budgetTokens
|
||||
: 0
|
||||
}
|
||||
|
||||
/**
|
||||
* The budget to retry with after a rejection, or null to surface the error.
|
||||
*/
|
||||
export function retryOutputTokens(
|
||||
error: unknown,
|
||||
params: { maxOutputTokens?: number; providerOptions?: unknown },
|
||||
): number | null {
|
||||
const requested = params.maxOutputTokens
|
||||
const text = rejectionText(error)
|
||||
if (!requested || !text) return null
|
||||
|
||||
const ceiling = readCeiling(text)
|
||||
if (ceiling !== null) {
|
||||
// The ceiling applies to what was actually sent, thinking included,
|
||||
// so the retry has to leave room for the thinking too.
|
||||
const thinking = thinkingBudget(params.providerOptions)
|
||||
if (ceiling >= requested + thinking) return null
|
||||
return usableLimit(ceiling - thinking)
|
||||
}
|
||||
|
||||
// Names the budget parameter, but in a format we cannot read a number from
|
||||
if (/max_\w*tokens/i.test(text) && requested > FALLBACK_OUTPUT_TOKENS) {
|
||||
return FALLBACK_OUTPUT_TOKENS
|
||||
}
|
||||
return null
|
||||
return output ? usableLimit(Number(output[1])) : null
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -178,15 +103,17 @@ export function withOutputTokenLimitFallback(
|
||||
try {
|
||||
return await doStream()
|
||||
} catch (error) {
|
||||
const retry = retryOutputTokens(error, params)
|
||||
if (!retry) throw error
|
||||
const limit = parseOutputTokenLimit(error)
|
||||
const requested = params.maxOutputTokens
|
||||
|
||||
if (!limit || !requested || limit >= requested) throw error
|
||||
|
||||
console.warn(
|
||||
`[maxOutputTokens] ${params.maxOutputTokens} rejected, retrying with ${retry}`,
|
||||
`[maxOutputTokens] ${requested} rejected, retrying with ${limit}`,
|
||||
)
|
||||
return await inner.doStream({
|
||||
...params,
|
||||
maxOutputTokens: retry,
|
||||
maxOutputTokens: limit,
|
||||
})
|
||||
}
|
||||
},
|
||||
@@ -208,17 +135,11 @@ function validBudget(value: string | null | undefined): number | null {
|
||||
* desktop app too), then server env, then the default. Both sources go through
|
||||
* the same validation, so a typo in either falls back instead of reaching the
|
||||
* provider.
|
||||
*
|
||||
* On the server's credentials the user setting can only lower the server value,
|
||||
* so MAX_OUTPUT_TOKENS keeps capping what the server pays for.
|
||||
*/
|
||||
export function resolveMaxOutputTokens(
|
||||
headerValue: string | null,
|
||||
usesServerCredentials: boolean,
|
||||
): number {
|
||||
const header = validBudget(headerValue)
|
||||
const server =
|
||||
validBudget(process.env.MAX_OUTPUT_TOKENS) ?? DEFAULT_MAX_OUTPUT_TOKENS
|
||||
if (header === null) return server
|
||||
return usesServerCredentials ? Math.min(header, server) : header
|
||||
export function resolveMaxOutputTokens(headerValue: string | null): number {
|
||||
return (
|
||||
validBudget(headerValue) ??
|
||||
validBudget(process.env.MAX_OUTPUT_TOKENS) ??
|
||||
DEFAULT_MAX_OUTPUT_TOKENS
|
||||
)
|
||||
}
|
||||
|
||||
+3
-6
@@ -1,4 +1,4 @@
|
||||
import { extractText } from "unpdf"
|
||||
import { extractText, getDocumentProxy } from "unpdf"
|
||||
|
||||
// Maximum characters allowed for extracted text (configurable via env)
|
||||
const DEFAULT_MAX_EXTRACTED_CHARS = 150000 // 150k chars
|
||||
@@ -14,7 +14,6 @@ const TEXT_EXTENSIONS = [
|
||||
".json",
|
||||
".csv",
|
||||
".xml",
|
||||
".svg",
|
||||
".html",
|
||||
".css",
|
||||
".js",
|
||||
@@ -44,10 +43,8 @@ const TEXT_EXTENSIONS = [
|
||||
*/
|
||||
export async function extractPdfText(file: File): Promise<string> {
|
||||
const buffer = await file.arrayBuffer()
|
||||
// Pass raw bytes so unpdf destroys the PDF document when it is done
|
||||
const { text } = await extractText(new Uint8Array(buffer), {
|
||||
mergePages: true,
|
||||
})
|
||||
const pdf = await getDocumentProxy(new Uint8Array(buffer))
|
||||
const { text } = await extractText(pdf, { mergePages: true })
|
||||
return text as string
|
||||
}
|
||||
|
||||
|
||||
@@ -47,14 +47,11 @@ export interface FlattenedServerModel {
|
||||
|
||||
/**
|
||||
* Convert provider name to URL-safe slug for use in model ID
|
||||
* e.g., "OpenAI Production" → "openai-production", "主力" → "4e3b-529b"
|
||||
* Non-ASCII characters become their hex code point so CJK names stay
|
||||
* distinct; the id is sent in HTTP headers, which must be ASCII.
|
||||
* e.g., "OpenAI Production" → "openai-production"
|
||||
*/
|
||||
export function slugify(name: string): string {
|
||||
function slugify(name: string): string {
|
||||
return name
|
||||
.toLowerCase()
|
||||
.replace(/[^\p{ASCII}]/gu, (c) => `-${c.codePointAt(0)?.toString(16)}-`)
|
||||
.replace(/[^a-z0-9]+/g, "-")
|
||||
.replace(/^-|-$/g, "")
|
||||
}
|
||||
@@ -192,7 +189,6 @@ export async function loadFlattenedServerModels(): Promise<
|
||||
const defaultModelId = process.env.AI_MODEL
|
||||
|
||||
const flattened: FlattenedServerModel[] = []
|
||||
const seenIds = new Set<string>()
|
||||
|
||||
for (const p of cfg.providers) {
|
||||
const providerLabel =
|
||||
@@ -203,16 +199,6 @@ export async function loadFlattenedServerModels(): Promise<
|
||||
|
||||
for (const modelId of p.models) {
|
||||
const id = `server:${nameSlug}:${modelId}`
|
||||
// Names that differ only in case or punctuation share a slug.
|
||||
// A repeated id would always resolve to the first provider's
|
||||
// credentials, so drop it instead.
|
||||
if (seenIds.has(id)) {
|
||||
console.warn(
|
||||
`[server-model-config] Skipping duplicate model id "${id}". Provider names must differ in letters or digits.`,
|
||||
)
|
||||
continue
|
||||
}
|
||||
seenIds.add(id)
|
||||
|
||||
// Default model priority:
|
||||
// 1. From ai-models.json: first model of provider with default: true
|
||||
|
||||
+23
-31
@@ -1,6 +1,5 @@
|
||||
import { type DBSchema, type IDBPDatabase, openDB } from "idb"
|
||||
import { nanoid } from "nanoid"
|
||||
import { toast } from "sonner"
|
||||
import type { Template } from "./template-storage"
|
||||
|
||||
// Constants
|
||||
@@ -62,7 +61,6 @@ let dbPromise: Promise<IDBPDatabase<ChatSessionDB>> | null = null
|
||||
|
||||
async function getDB(): Promise<IDBPDatabase<ChatSessionDB>> {
|
||||
if (!dbPromise) {
|
||||
// A failed or lost connection is not cached: the next call reopens it
|
||||
dbPromise = openDB<ChatSessionDB>(DB_NAME, DB_VERSION, {
|
||||
upgrade(db, oldVersion) {
|
||||
if (oldVersion < 1) {
|
||||
@@ -90,28 +88,6 @@ async function getDB(): Promise<IDBPDatabase<ChatSessionDB>> {
|
||||
}
|
||||
}
|
||||
},
|
||||
blocked() {
|
||||
// An older tab keeps the DB open, so the upgrade has to wait
|
||||
toast.warning(
|
||||
"Please close other tabs of this app to finish updating chat storage.",
|
||||
{ id: "idb-upgrade-blocked", duration: 10000 },
|
||||
)
|
||||
},
|
||||
blocking(_currentVersion, _blockedVersion, event) {
|
||||
// Another tab needs to upgrade the DB: close our connection so
|
||||
// it is not stuck, and reopen on the next call
|
||||
const db = event.target as IDBDatabase
|
||||
db.close()
|
||||
dbPromise = null
|
||||
},
|
||||
terminated() {
|
||||
// The browser closed the connection (e.g. Safari after a long
|
||||
// time in the background)
|
||||
dbPromise = null
|
||||
},
|
||||
}).catch((error) => {
|
||||
dbPromise = null
|
||||
throw error
|
||||
})
|
||||
}
|
||||
return dbPromise
|
||||
@@ -169,8 +145,6 @@ export async function getSession(id: string): Promise<ChatSession | null> {
|
||||
}
|
||||
}
|
||||
|
||||
// Returns false on failure (e.g. storage quota exceeded). Other sessions are
|
||||
// never deleted automatically; the caller tells the user instead.
|
||||
export async function saveSession(session: ChatSession): Promise<boolean> {
|
||||
if (!isIndexedDBAvailable()) return false
|
||||
try {
|
||||
@@ -178,11 +152,29 @@ export async function saveSession(session: ChatSession): Promise<boolean> {
|
||||
await db.put(STORE_NAME, session)
|
||||
return true
|
||||
} catch (error) {
|
||||
console.error("Failed to save session:", error)
|
||||
// Reopen the connection next time in case it was lost (Safari reports
|
||||
// "Connection to Indexed Database server lost" without closing it)
|
||||
dbPromise = null
|
||||
return false
|
||||
// Handle quota exceeded
|
||||
if (
|
||||
error instanceof DOMException &&
|
||||
error.name === "QuotaExceededError"
|
||||
) {
|
||||
console.warn("Storage quota exceeded, deleting oldest session...")
|
||||
await deleteOldestSession()
|
||||
// Retry once
|
||||
try {
|
||||
const db = await getDB()
|
||||
await db.put(STORE_NAME, session)
|
||||
return true
|
||||
} catch (retryError) {
|
||||
console.error(
|
||||
"Failed to save session after cleanup:",
|
||||
retryError,
|
||||
)
|
||||
return false
|
||||
}
|
||||
} else {
|
||||
console.error("Failed to save session:", error)
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -41,7 +41,7 @@ parameters: {
|
||||
tool name: edit_diagram
|
||||
description: Edit specific parts of the EXISTING diagram. Use this when making small targeted changes like adding/removing elements, changing labels, or adjusting properties. This is more efficient than regenerating the entire diagram.
|
||||
parameters: {
|
||||
operations: Array<{operation: "update" | "add" | "delete", cell_id: string, new_xml?: string}>
|
||||
edits: Array<{search: string, replace: string}>
|
||||
}
|
||||
---Tool3---
|
||||
tool name: append_diagram
|
||||
|
||||
+1
-6
@@ -1,6 +1,5 @@
|
||||
import { z } from "zod"
|
||||
import { getApiEndpoint } from "@/lib/base-path"
|
||||
import { STORAGE_KEYS } from "@/lib/storage"
|
||||
|
||||
export interface UrlData {
|
||||
url: string
|
||||
@@ -19,11 +18,7 @@ const UrlResponseSchema = z.object({
|
||||
export async function extractUrlContent(url: string): Promise<UrlData> {
|
||||
const response = await fetch(getApiEndpoint("/api/parse-url"), {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
"x-access-code":
|
||||
localStorage.getItem(STORAGE_KEYS.accessCode) || "",
|
||||
},
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({ url }),
|
||||
})
|
||||
|
||||
|
||||
+61
-55
@@ -27,72 +27,78 @@ export function useFileProcessor() {
|
||||
const handleFileChange = async (newFiles: File[]) => {
|
||||
setFiles(newFiles)
|
||||
|
||||
const pending = newFiles.filter(
|
||||
(file) =>
|
||||
(isPdfFile(file) || isTextFile(file)) && !pdfData.has(file),
|
||||
)
|
||||
// Extract text immediately for new PDF/text files
|
||||
for (const file of newFiles) {
|
||||
const needsExtraction =
|
||||
(isPdfFile(file) || isTextFile(file)) && !pdfData.has(file)
|
||||
if (needsExtraction) {
|
||||
// Mark as extracting
|
||||
setPdfData((prev) => {
|
||||
const next = new Map(prev)
|
||||
next.set(file, {
|
||||
text: "",
|
||||
charCount: 0,
|
||||
isExtracting: true,
|
||||
})
|
||||
return next
|
||||
})
|
||||
|
||||
// Before any await: drop data for removed files and mark every new
|
||||
// file as extracting, so queued files also block sending
|
||||
setPdfData((prev) => {
|
||||
const next = new Map<File, FileData>()
|
||||
for (const file of newFiles) {
|
||||
const existing = prev.get(file)
|
||||
if (existing) next.set(file, existing)
|
||||
}
|
||||
for (const file of pending) {
|
||||
next.set(file, { text: "", charCount: 0, isExtracting: true })
|
||||
}
|
||||
return next
|
||||
})
|
||||
// Extract text asynchronously
|
||||
try {
|
||||
let text: string
|
||||
if (isPdfFile(file)) {
|
||||
text = await extractPdfText(file)
|
||||
} else {
|
||||
text = await extractTextFileContent(file)
|
||||
}
|
||||
|
||||
// Extract one file at a time
|
||||
for (const file of pending) {
|
||||
try {
|
||||
let text: string
|
||||
if (isPdfFile(file)) {
|
||||
text = await extractPdfText(file)
|
||||
} else {
|
||||
text = await extractTextFileContent(file)
|
||||
}
|
||||
// Check character limit
|
||||
if (text.length > MAX_EXTRACTED_CHARS) {
|
||||
const limitK = MAX_EXTRACTED_CHARS / 1000
|
||||
toast.error(
|
||||
`${file.name}: Content exceeds ${limitK}k character limit (${(text.length / 1000).toFixed(1)}k chars)`,
|
||||
)
|
||||
setPdfData((prev) => {
|
||||
const next = new Map(prev)
|
||||
next.delete(file)
|
||||
return next
|
||||
})
|
||||
// Remove the file from the list
|
||||
setFiles((prev) => prev.filter((f) => f !== file))
|
||||
continue
|
||||
}
|
||||
|
||||
// Check character limit
|
||||
if (text.length > MAX_EXTRACTED_CHARS) {
|
||||
const limitK = MAX_EXTRACTED_CHARS / 1000
|
||||
toast.error(
|
||||
`${file.name}: Content exceeds ${limitK}k character limit (${(text.length / 1000).toFixed(1)}k chars)`,
|
||||
)
|
||||
setPdfData((prev) => {
|
||||
const next = new Map(prev)
|
||||
next.set(file, {
|
||||
text,
|
||||
charCount: text.length,
|
||||
isExtracting: false,
|
||||
})
|
||||
return next
|
||||
})
|
||||
} catch (error) {
|
||||
console.error("Failed to extract text:", error)
|
||||
toast.error(`Failed to read file: ${file.name}`)
|
||||
setPdfData((prev) => {
|
||||
const next = new Map(prev)
|
||||
next.delete(file)
|
||||
return next
|
||||
})
|
||||
// Remove the file from the list
|
||||
setFiles((prev) => prev.filter((f) => f !== file))
|
||||
continue
|
||||
}
|
||||
|
||||
setPdfData((prev) => {
|
||||
// The file was removed while extracting
|
||||
if (!prev.has(file)) return prev
|
||||
const next = new Map(prev)
|
||||
next.set(file, {
|
||||
text,
|
||||
charCount: text.length,
|
||||
isExtracting: false,
|
||||
})
|
||||
return next
|
||||
})
|
||||
} catch (error) {
|
||||
console.error("Failed to extract text:", error)
|
||||
toast.error(`Failed to read file: ${file.name}`)
|
||||
setPdfData((prev) => {
|
||||
const next = new Map(prev)
|
||||
next.delete(file)
|
||||
return next
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Clean up pdfData for removed files
|
||||
setPdfData((prev) => {
|
||||
const next = new Map(prev)
|
||||
for (const key of prev.keys()) {
|
||||
if (!newFiles.includes(key)) {
|
||||
next.delete(key)
|
||||
}
|
||||
}
|
||||
return next
|
||||
})
|
||||
}
|
||||
|
||||
return {
|
||||
|
||||
+195
-264
@@ -76,17 +76,6 @@ export function isMxCellXmlComplete(xml: string | undefined | null): boolean {
|
||||
// No valid ending found at all
|
||||
if (lastValidEnd === -1) return false
|
||||
|
||||
// If the last mxCell has no </mxCell> after it, it must be self-closing.
|
||||
// Otherwise the trailing "/>" belongs to a child such as <mxGeometry .../>
|
||||
// and the output was cut off before the cell was closed.
|
||||
const lastCellStart = trimmed.lastIndexOf("<mxCell")
|
||||
if (
|
||||
lastCellStart > lastMxCellClose &&
|
||||
!/^<mxCell\b[^<]*\/>/.test(trimmed.slice(lastCellStart))
|
||||
) {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check what comes after the last valid ending
|
||||
// For />: add 2 chars, for </mxCell>: add 9 chars
|
||||
const endOffset = lastMxCellClose > lastSelfClose ? 9 : 2
|
||||
@@ -106,12 +95,36 @@ export function isMxCellXmlComplete(xml: string | undefined | null): boolean {
|
||||
export function extractCompleteMxCells(xml: string | undefined | null): string {
|
||||
if (!xml) return ""
|
||||
|
||||
// Match self-closing <mxCell ... /> or <mxCell ...>...</mxCell>, in document order.
|
||||
// The lazy [^>]*? tries "/>" first, so a self-closing cell never swallows
|
||||
// the following cells up to the next </mxCell>.
|
||||
const cellPattern = /<mxCell\b[^>]*?(?:\/>|>[\s\S]*?<\/mxCell>)/g
|
||||
const completeCells: Array<{ index: number; text: string }> = []
|
||||
|
||||
return (xml.match(cellPattern) || []).join("\n")
|
||||
// Match self-closing mxCell tags: <mxCell ... />
|
||||
// Also match mxCell with nested mxGeometry: <mxCell ...>...<mxGeometry .../></mxCell>
|
||||
const selfClosingPattern = /<mxCell\s+[^>]*\/>/g
|
||||
const nestedPattern = /<mxCell\s+[^>]*>[\s\S]*?<\/mxCell>/g
|
||||
|
||||
// Find all self-closing mxCell elements
|
||||
let match: RegExpExecArray | null
|
||||
while ((match = selfClosingPattern.exec(xml)) !== null) {
|
||||
completeCells.push({ index: match.index, text: match[0] })
|
||||
}
|
||||
|
||||
// Find all mxCell elements with nested content (like mxGeometry)
|
||||
while ((match = nestedPattern.exec(xml)) !== null) {
|
||||
completeCells.push({ index: match.index, text: match[0] })
|
||||
}
|
||||
|
||||
// Sort by position to maintain order
|
||||
completeCells.sort((a, b) => a.index - b.index)
|
||||
|
||||
// Remove duplicates (a self-closing match might overlap with nested match)
|
||||
const seen = new Set<number>()
|
||||
const uniqueCells = completeCells.filter((cell) => {
|
||||
if (seen.has(cell.index)) return false
|
||||
seen.add(cell.index)
|
||||
return true
|
||||
})
|
||||
|
||||
return uniqueCells.map((c) => c.text).join("\n")
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
@@ -474,31 +487,6 @@ export interface ApplyOperationsResult {
|
||||
errors: OperationError[]
|
||||
}
|
||||
|
||||
/**
|
||||
* draw.io wraps cells that have links, tooltips or custom data in
|
||||
* <object>/<UserObject>, and the wrapper carries the id instead of the mxCell.
|
||||
*/
|
||||
function getCellWrapper(cell: Element): Element | null {
|
||||
const parent = cell.parentElement
|
||||
return parent?.tagName === "object" || parent?.tagName === "UserObject"
|
||||
? parent
|
||||
: null
|
||||
}
|
||||
|
||||
/** Id of a cell, read from its wrapper when the mxCell has none */
|
||||
function getCellId(cell: Element): string | null {
|
||||
return (
|
||||
cell.getAttribute("id") ||
|
||||
getCellWrapper(cell)?.getAttribute("id") ||
|
||||
null
|
||||
)
|
||||
}
|
||||
|
||||
/** Element to replace or remove for a cell (the wrapper if there is one) */
|
||||
function getCellNode(cell: Element): Element {
|
||||
return getCellWrapper(cell) || cell
|
||||
}
|
||||
|
||||
/**
|
||||
* Apply diagram operations (update/add/delete) using ID-based lookup.
|
||||
* This replaces the text-matching approach with direct DOM manipulation.
|
||||
@@ -547,14 +535,12 @@ export function applyDiagramOperations(
|
||||
}
|
||||
}
|
||||
|
||||
// Build a map of cell IDs to elements (wrapper elements for wrapped cells)
|
||||
// Build a map of cell IDs to elements
|
||||
const cellMap = new Map<string, Element>()
|
||||
root.querySelectorAll("mxCell").forEach((cell) => {
|
||||
const id = getCellId(cell)
|
||||
if (id) cellMap.set(id, getCellNode(cell))
|
||||
const id = cell.getAttribute("id")
|
||||
if (id) cellMap.set(id, cell)
|
||||
})
|
||||
// Cells removed by delete operations in this batch
|
||||
const deletedIds = new Set<string>()
|
||||
|
||||
// Process each operation
|
||||
for (const op of operations) {
|
||||
@@ -594,7 +580,7 @@ export function applyDiagramOperations(
|
||||
}
|
||||
|
||||
// Validate ID matches
|
||||
const newCellId = getCellId(newCell)
|
||||
const newCellId = newCell.getAttribute("id")
|
||||
if (newCellId !== op.cell_id) {
|
||||
errors.push({
|
||||
type: "update",
|
||||
@@ -604,8 +590,8 @@ export function applyDiagramOperations(
|
||||
continue
|
||||
}
|
||||
|
||||
// Import and replace the node (with its wrapper, if any)
|
||||
const importedNode = doc.importNode(getCellNode(newCell), true)
|
||||
// Import and replace the node
|
||||
const importedNode = doc.importNode(newCell, true)
|
||||
existingCell.parentNode?.replaceChild(importedNode, existingCell)
|
||||
|
||||
// Update the map with the new element
|
||||
@@ -646,7 +632,7 @@ export function applyDiagramOperations(
|
||||
}
|
||||
|
||||
// Validate ID matches
|
||||
const newCellId = getCellId(newCell)
|
||||
const newCellId = newCell.getAttribute("id")
|
||||
if (newCellId !== op.cell_id) {
|
||||
errors.push({
|
||||
type: "add",
|
||||
@@ -656,8 +642,8 @@ export function applyDiagramOperations(
|
||||
continue
|
||||
}
|
||||
|
||||
// Import and append the node (with its wrapper, if any)
|
||||
const importedNode = doc.importNode(getCellNode(newCell), true)
|
||||
// Import and append the node
|
||||
const importedNode = doc.importNode(newCell, true)
|
||||
root.appendChild(importedNode)
|
||||
|
||||
// Add to map
|
||||
@@ -675,15 +661,8 @@ export function applyDiagramOperations(
|
||||
|
||||
const existingCell = cellMap.get(op.cell_id)
|
||||
if (!existingCell) {
|
||||
// Cells cascade-deleted earlier in this batch are skipped silently
|
||||
// (AI may redundantly list children/edges)
|
||||
if (!deletedIds.has(op.cell_id)) {
|
||||
errors.push({
|
||||
type: "delete",
|
||||
cellId: op.cell_id,
|
||||
message: `Cell with id="${op.cell_id}" not found`,
|
||||
})
|
||||
}
|
||||
// Cell not found - might have been cascade-deleted by a previous operation
|
||||
// Skip silently instead of erroring (AI may redundantly list children/edges)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -700,7 +679,7 @@ export function applyDiagramOperations(
|
||||
`mxCell[parent="${cellId}"]`,
|
||||
)
|
||||
children.forEach((child) => {
|
||||
const childId = getCellId(child)
|
||||
const childId = child.getAttribute("id")
|
||||
if (childId && childId !== "0" && childId !== "1") {
|
||||
collectDescendants(childId)
|
||||
}
|
||||
@@ -717,7 +696,7 @@ export function applyDiagramOperations(
|
||||
`mxCell[source="${cellId}"], mxCell[target="${cellId}"]`,
|
||||
)
|
||||
referencingEdges.forEach((edge) => {
|
||||
const edgeId = getCellId(edge)
|
||||
const edgeId = edge.getAttribute("id")
|
||||
// Protect root cells from being added via edge references
|
||||
if (edgeId && edgeId !== "0" && edgeId !== "1") {
|
||||
// Recurse to collect edge's children (like labels)
|
||||
@@ -739,7 +718,6 @@ export function applyDiagramOperations(
|
||||
if (cell) {
|
||||
cell.parentNode?.removeChild(cell)
|
||||
cellMap.delete(cellId)
|
||||
deletedIds.add(cellId)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -780,89 +758,24 @@ function checkDuplicateAttributes(xml: string): string | null {
|
||||
return null
|
||||
}
|
||||
|
||||
/** Matches one <diagram> page of a document (the last one may be unclosed) */
|
||||
const PAGE_PATTERN = /<diagram\b[\s\S]*?(?:<\/diagram>|$)/g
|
||||
|
||||
const ID_ATTR_PATTERN = /\bid\s*=\s*["']([^"']+)["']/gi
|
||||
|
||||
/**
|
||||
* Split XML into pages. Ids only need to be unique within a page: every
|
||||
* page of a multi-page document has its own root cells "0" and "1".
|
||||
*/
|
||||
function splitPages(xml: string): string[] {
|
||||
return xml.match(PAGE_PATTERN) || [xml]
|
||||
}
|
||||
|
||||
/** Ids that appear more than once, with their counts */
|
||||
function findDuplicateIds(xml: string): Map<string, number> {
|
||||
const ids = new Map<string, number>()
|
||||
for (const match of xml.matchAll(ID_ATTR_PATTERN)) {
|
||||
ids.set(match[1], (ids.get(match[1]) || 0) + 1)
|
||||
}
|
||||
return new Map(Array.from(ids).filter(([, count]) => count > 1))
|
||||
}
|
||||
|
||||
/** Check for duplicate IDs in XML (per page) */
|
||||
/** Check for duplicate IDs in XML */
|
||||
function checkDuplicateIds(xml: string): string | null {
|
||||
for (const page of splitPages(xml)) {
|
||||
const duplicateIds = Array.from(findDuplicateIds(page)).map(
|
||||
([id, count]) => `'${id}' (${count}x)`,
|
||||
)
|
||||
if (duplicateIds.length > 0) {
|
||||
return `Invalid XML: Found duplicate ID(s): ${duplicateIds.slice(0, 3).join(", ")}. All id attributes must be unique.`
|
||||
}
|
||||
const idPattern = /\bid\s*=\s*["']([^"']+)["']/gi
|
||||
const ids = new Map<string, number>()
|
||||
let idMatch
|
||||
while ((idMatch = idPattern.exec(xml)) !== null) {
|
||||
const id = idMatch[1]
|
||||
ids.set(id, (ids.get(id) || 0) + 1)
|
||||
}
|
||||
const duplicateIds = Array.from(ids.entries())
|
||||
.filter(([, count]) => count > 1)
|
||||
.map(([id, count]) => `'${id}' (${count}x)`)
|
||||
if (duplicateIds.length > 0) {
|
||||
return `Invalid XML: Found duplicate ID(s): ${duplicateIds.slice(0, 3).join(", ")}. All id attributes must be unique.`
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
/** Rename repeated ids in one page (keeps the first occurrence) */
|
||||
function renameDuplicateIds(xml: string): { xml: string; renamed: number } {
|
||||
const duplicateIds = findDuplicateIds(xml)
|
||||
if (duplicateIds.size === 0) return { xml, renamed: 0 }
|
||||
|
||||
const idCounters = new Map<string, number>()
|
||||
const renamedXml = xml.replace(ID_ATTR_PATTERN, (match, id) => {
|
||||
if (!duplicateIds.has(id)) return match
|
||||
|
||||
const count = idCounters.get(id) || 0
|
||||
idCounters.set(id, count + 1)
|
||||
|
||||
if (count === 0) return match // Keep first occurrence
|
||||
|
||||
// Rename subsequent occurrences (the id sits just before the closing quote)
|
||||
return `${match.slice(0, -id.length - 1)}${id}_dup${count}${match.slice(-1)}`
|
||||
})
|
||||
return { xml: renamedXml, renamed: duplicateIds.size }
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns a function telling whether a position is inside a quoted attribute
|
||||
* value. Positions must be queried in increasing order: the scan resumes where
|
||||
* it stopped instead of starting over, which keeps large documents fast.
|
||||
*/
|
||||
function createQuoteTracker(str: string): (pos: number) => boolean {
|
||||
let i = 0
|
||||
let inQuote = false
|
||||
let quoteChar = ""
|
||||
return (pos: number) => {
|
||||
for (; i < pos && i < str.length; i++) {
|
||||
const c = str[i]
|
||||
if (inQuote) {
|
||||
if (c === quoteChar) inQuote = false
|
||||
} else if (c === '"' || c === "'") {
|
||||
// Only quotes that follow "=" open an attribute value
|
||||
let j = i - 1
|
||||
while (j >= 0 && /\s/.test(str[j])) j--
|
||||
if (j >= 0 && str[j] === "=") {
|
||||
inQuote = true
|
||||
quoteChar = c
|
||||
}
|
||||
}
|
||||
}
|
||||
return inQuote
|
||||
}
|
||||
}
|
||||
|
||||
/** Check for tag mismatches using parsed tags */
|
||||
function checkTagMismatches(xml: string): string | null {
|
||||
const xmlWithoutComments = xml.replace(/<!--[\s\S]*?-->/g, "")
|
||||
@@ -1175,19 +1088,13 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
// 3b. Fix malformed attribute values where " is used as delimiter instead of actual quotes
|
||||
// Pattern: attr="value" should become attr="value" (the " was meant to be the quote delimiter)
|
||||
// This commonly happens with dashPattern="1 1;"
|
||||
// Matches inside another attribute value are kept: rich text labels like
|
||||
// value="<font color="#ff0000">..." are valid.
|
||||
const isInsideQuotesFor3b = createQuoteTracker(fixed)
|
||||
let malformedQuotesFixed = false
|
||||
fixed = fixed.replace(
|
||||
/(\s[a-zA-Z][a-zA-Z0-9_:-]*)="([^&]*?)"/g,
|
||||
(match: string, attr: string, value: string, offset: number) => {
|
||||
if (isInsideQuotesFor3b(offset)) return match
|
||||
malformedQuotesFixed = true
|
||||
return `${attr}="${value}"`
|
||||
},
|
||||
)
|
||||
if (malformedQuotesFixed) {
|
||||
const malformedQuotePattern = /(\s[a-zA-Z][a-zA-Z0-9_:-]*)="/
|
||||
if (malformedQuotePattern.test(fixed)) {
|
||||
// Replace =" with =" and trailing " before next attribute or tag end with "
|
||||
fixed = fixed.replace(
|
||||
/(\s[a-zA-Z][a-zA-Z0-9_:-]*)="([^&]*?)"/g,
|
||||
'$1="$2"',
|
||||
)
|
||||
fixes.push(
|
||||
'Fixed malformed attribute quotes (="..." to ="...")',
|
||||
)
|
||||
@@ -1201,11 +1108,9 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
}
|
||||
|
||||
// 3d. Fix missing space between attributes like vertex="1"parent="1"
|
||||
// Requires name=" right after the quote, so the opening quote of a value
|
||||
// such as style="rounded=1;..." is not mistaken for a closing one.
|
||||
const missingSpacePattern = /"([a-zA-Z_:][\w:.-]*=")/g
|
||||
const missingSpacePattern = /("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g
|
||||
if (missingSpacePattern.test(fixed)) {
|
||||
fixed = fixed.replace(missingSpacePattern, '" $1')
|
||||
fixed = fixed.replace(/("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g, "$1 $2")
|
||||
fixes.push("Added missing space between attributes")
|
||||
}
|
||||
|
||||
@@ -1335,13 +1240,32 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
"mxPoint",
|
||||
"Array",
|
||||
"Object",
|
||||
// Wrappers of cells with links, tooltips or custom data
|
||||
"object",
|
||||
"UserObject",
|
||||
"mxRectangle",
|
||||
])
|
||||
|
||||
const isInsideQuotesFor8c = createQuoteTracker(fixed)
|
||||
// Helper: Check if a position is inside a quoted attribute value
|
||||
// by counting unescaped quotes before that position
|
||||
const isInsideQuotes = (str: string, pos: number): boolean => {
|
||||
let inQuote = false
|
||||
let quoteChar = ""
|
||||
for (let i = 0; i < pos && i < str.length; i++) {
|
||||
const c = str[i]
|
||||
if (inQuote) {
|
||||
if (c === quoteChar) inQuote = false
|
||||
} else if (c === '"' || c === "'") {
|
||||
// Check if this quote is part of an attribute (preceded by =)
|
||||
// Look back for = sign
|
||||
let j = i - 1
|
||||
while (j >= 0 && /\s/.test(str[j])) j--
|
||||
if (j >= 0 && str[j] === "=") {
|
||||
inQuote = true
|
||||
quoteChar = c
|
||||
}
|
||||
}
|
||||
}
|
||||
return inQuote
|
||||
}
|
||||
|
||||
const foreignTagPattern = /<\/?([a-zA-Z][a-zA-Z0-9_]*)[^>]*>/g
|
||||
let foreignMatch
|
||||
const foreignTags = new Set<string>()
|
||||
@@ -1356,7 +1280,7 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
// Skip if this is a valid draw.io tag
|
||||
if (validDrawioTags.has(tagName)) continue
|
||||
// Skip if this tag is inside a quoted attribute value
|
||||
if (isInsideQuotesFor8c(foreignMatch.index)) continue
|
||||
if (isInsideQuotes(fixed, foreignMatch.index)) continue
|
||||
|
||||
foreignTags.add(tagName)
|
||||
foreignTagPositions.push({
|
||||
@@ -1428,11 +1352,10 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
>()
|
||||
// Match full tags to detect self-closing by checking if ends with />
|
||||
const fullTagPattern = /<(\/?[a-zA-Z][a-zA-Z0-9]*)[^>]*>/g
|
||||
const isInsideQuotesFor10b = createQuoteTracker(fixed)
|
||||
let tagCountMatch
|
||||
while ((tagCountMatch = fullTagPattern.exec(fixed)) !== null) {
|
||||
// Skip tags inside quoted attribute values (e.g., value="<b>Title</b>")
|
||||
if (isInsideQuotesFor10b(tagCountMatch.index)) continue
|
||||
if (isInsideQuotes(fixed, tagCountMatch.index)) continue
|
||||
|
||||
const fullMatch = tagCountMatch[0] // e.g., "<mxCell .../>" or "</mxCell>"
|
||||
const tagPart = tagCountMatch[1] // e.g., "mxCell" or "/mxCell"
|
||||
@@ -1522,112 +1445,125 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
// 11. Fix nested mxCell by flattening
|
||||
// Pattern A: <mxCell id="X">...<mxCell id="X">...</mxCell></mxCell> (duplicate ID)
|
||||
// Pattern B: <mxCell id="X">...<mxCell id="Y">...</mxCell></mxCell> (different ID - true nesting)
|
||||
// These passes work line by line and would break valid cells written on a
|
||||
// single line, so each one runs only when cells are really nested.
|
||||
if (checkNestedMxCells(fixed)) {
|
||||
const lines = fixed.split("\n")
|
||||
const newLines: string[] = []
|
||||
let nestedFixed = 0
|
||||
let extraClosingToRemove = 0
|
||||
const lines = fixed.split("\n")
|
||||
let newLines: string[] = []
|
||||
let nestedFixed = 0
|
||||
let extraClosingToRemove = 0
|
||||
|
||||
// First pass: fix duplicate ID nesting (same as before)
|
||||
for (let i = 0; i < lines.length; i++) {
|
||||
const line = lines[i]
|
||||
const nextLine = lines[i + 1]
|
||||
// First pass: fix duplicate ID nesting (same as before)
|
||||
for (let i = 0; i < lines.length; i++) {
|
||||
const line = lines[i]
|
||||
const nextLine = lines[i + 1]
|
||||
|
||||
// Check if current line and next line are both mxCell opening tags with same ID
|
||||
if (
|
||||
nextLine &&
|
||||
/<mxCell\s/.test(line) &&
|
||||
/<mxCell\s/.test(nextLine) &&
|
||||
!line.includes("/>") &&
|
||||
!nextLine.includes("/>")
|
||||
) {
|
||||
const id1 = line.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
|
||||
const id2 = nextLine.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
|
||||
// Check if current line and next line are both mxCell opening tags with same ID
|
||||
if (
|
||||
nextLine &&
|
||||
/<mxCell\s/.test(line) &&
|
||||
/<mxCell\s/.test(nextLine) &&
|
||||
!line.includes("/>") &&
|
||||
!nextLine.includes("/>")
|
||||
) {
|
||||
const id1 = line.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
|
||||
const id2 = nextLine.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
|
||||
|
||||
if (id1 && id1 === id2) {
|
||||
nestedFixed++
|
||||
extraClosingToRemove++ // Need to remove one </mxCell> later
|
||||
continue // Skip this duplicate opening line
|
||||
}
|
||||
if (id1 && id1 === id2) {
|
||||
nestedFixed++
|
||||
extraClosingToRemove++ // Need to remove one </mxCell> later
|
||||
continue // Skip this duplicate opening line
|
||||
}
|
||||
}
|
||||
|
||||
// Remove extra </mxCell> if we have pending removals
|
||||
if (extraClosingToRemove > 0 && /^\s*<\/mxCell>\s*$/.test(line)) {
|
||||
extraClosingToRemove--
|
||||
continue // Skip this closing tag
|
||||
// Remove extra </mxCell> if we have pending removals
|
||||
if (extraClosingToRemove > 0 && /^\s*<\/mxCell>\s*$/.test(line)) {
|
||||
extraClosingToRemove--
|
||||
continue // Skip this closing tag
|
||||
}
|
||||
|
||||
newLines.push(line)
|
||||
}
|
||||
|
||||
if (nestedFixed > 0) {
|
||||
fixed = newLines.join("\n")
|
||||
fixes.push(`Flattened ${nestedFixed} duplicate-ID nested mxCell(s)`)
|
||||
}
|
||||
|
||||
// Second pass: fix true nesting (different IDs)
|
||||
// Insert </mxCell> before nested child to close parent
|
||||
const lines2 = fixed.split("\n")
|
||||
newLines = []
|
||||
let trueNestedFixed = 0
|
||||
let cellDepth = 0
|
||||
let pendingCloseRemoval = 0
|
||||
|
||||
for (let i = 0; i < lines2.length; i++) {
|
||||
const line = lines2[i]
|
||||
const trimmed = line.trim()
|
||||
|
||||
// Track mxCell depth
|
||||
const isOpenCell = /<mxCell\s/.test(trimmed) && !trimmed.endsWith("/>")
|
||||
const isCloseCell = trimmed === "</mxCell>"
|
||||
|
||||
if (isOpenCell) {
|
||||
if (cellDepth > 0) {
|
||||
// Found nested cell - insert closing tag for parent before this line
|
||||
const indent = line.match(/^(\s*)/)?.[1] || ""
|
||||
newLines.push(indent + "</mxCell>")
|
||||
trueNestedFixed++
|
||||
pendingCloseRemoval++ // Need to remove one </mxCell> later
|
||||
}
|
||||
|
||||
cellDepth = 1 // Reset to 1 since we just opened a new cell
|
||||
newLines.push(line)
|
||||
} else if (isCloseCell) {
|
||||
if (pendingCloseRemoval > 0) {
|
||||
pendingCloseRemoval--
|
||||
// Skip this extra closing tag
|
||||
} else {
|
||||
cellDepth = Math.max(0, cellDepth - 1)
|
||||
newLines.push(line)
|
||||
}
|
||||
} else {
|
||||
newLines.push(line)
|
||||
}
|
||||
|
||||
if (nestedFixed > 0) {
|
||||
fixed = newLines.join("\n")
|
||||
fixes.push(`Flattened ${nestedFixed} duplicate-ID nested mxCell(s)`)
|
||||
}
|
||||
}
|
||||
|
||||
if (checkNestedMxCells(fixed)) {
|
||||
// Second pass: fix true nesting (different IDs)
|
||||
// Insert </mxCell> before nested child to close parent
|
||||
const lines2 = fixed.split("\n")
|
||||
const newLines: string[] = []
|
||||
let trueNestedFixed = 0
|
||||
let cellDepth = 0
|
||||
let pendingCloseRemoval = 0
|
||||
|
||||
for (let i = 0; i < lines2.length; i++) {
|
||||
const line = lines2[i]
|
||||
const trimmed = line.trim()
|
||||
|
||||
// Track mxCell depth
|
||||
const isOpenCell =
|
||||
/<mxCell\s/.test(trimmed) && !trimmed.endsWith("/>")
|
||||
const isCloseCell = trimmed === "</mxCell>"
|
||||
|
||||
if (isOpenCell) {
|
||||
if (cellDepth > 0) {
|
||||
// Found nested cell - insert closing tag for parent before this line
|
||||
const indent = line.match(/^(\s*)/)?.[1] || ""
|
||||
newLines.push(indent + "</mxCell>")
|
||||
trueNestedFixed++
|
||||
pendingCloseRemoval++ // Need to remove one </mxCell> later
|
||||
}
|
||||
cellDepth = 1 // Reset to 1 since we just opened a new cell
|
||||
newLines.push(line)
|
||||
} else if (isCloseCell) {
|
||||
if (pendingCloseRemoval > 0) {
|
||||
pendingCloseRemoval--
|
||||
// Skip this extra closing tag
|
||||
} else {
|
||||
cellDepth = Math.max(0, cellDepth - 1)
|
||||
newLines.push(line)
|
||||
}
|
||||
} else {
|
||||
newLines.push(line)
|
||||
}
|
||||
}
|
||||
|
||||
if (trueNestedFixed > 0) {
|
||||
fixed = newLines.join("\n")
|
||||
fixes.push(`Fixed ${trueNestedFixed} true nested mxCell(s)`)
|
||||
}
|
||||
if (trueNestedFixed > 0) {
|
||||
fixed = newLines.join("\n")
|
||||
fixes.push(`Fixed ${trueNestedFixed} true nested mxCell(s)`)
|
||||
}
|
||||
|
||||
// 12. Fix duplicate IDs by appending suffix, page by page (ids such as the
|
||||
// root cells "0" and "1" legitimately repeat across pages)
|
||||
let renamedIds = 0
|
||||
const renamePage = (page: string) => {
|
||||
const { xml: renamed, renamed: count } = renameDuplicateIds(page)
|
||||
renamedIds += count
|
||||
return renamed
|
||||
// 12. Fix duplicate IDs by appending suffix
|
||||
const seenIds = new Map<string, number>()
|
||||
const duplicateIds: string[] = []
|
||||
|
||||
// First pass: find duplicates
|
||||
const idPattern = /\bid\s*=\s*["']([^"']+)["']/gi
|
||||
let idMatch
|
||||
while ((idMatch = idPattern.exec(fixed)) !== null) {
|
||||
const id = idMatch[1]
|
||||
seenIds.set(id, (seenIds.get(id) || 0) + 1)
|
||||
}
|
||||
fixed = /<diagram\b/.test(fixed)
|
||||
? fixed.replace(PAGE_PATTERN, renamePage)
|
||||
: renamePage(fixed)
|
||||
if (renamedIds > 0) {
|
||||
fixes.push(`Renamed ${renamedIds} duplicate ID(s)`)
|
||||
|
||||
// Find which IDs are duplicated
|
||||
for (const [id, count] of seenIds) {
|
||||
if (count > 1) duplicateIds.push(id)
|
||||
}
|
||||
|
||||
// Second pass: rename duplicates (keep first occurrence, rename others)
|
||||
if (duplicateIds.length > 0) {
|
||||
const idCounters = new Map<string, number>()
|
||||
fixed = fixed.replace(/\bid\s*=\s*["']([^"']+)["']/gi, (match, id) => {
|
||||
if (!duplicateIds.includes(id)) return match
|
||||
|
||||
const count = idCounters.get(id) || 0
|
||||
idCounters.set(id, count + 1)
|
||||
|
||||
if (count === 0) return match // Keep first occurrence
|
||||
|
||||
// Rename subsequent occurrences
|
||||
const newId = `${id}_dup${count}`
|
||||
return match.replace(id, newId)
|
||||
})
|
||||
fixes.push(`Renamed ${duplicateIds.length} duplicate ID(s)`)
|
||||
}
|
||||
|
||||
// 9. Fix empty id attributes by generating unique IDs
|
||||
@@ -1737,11 +1673,6 @@ export function validateAndFixXml(xml: string): {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Decode an xmlsvg export (SVG data URL) into uncompressed diagram XML.
|
||||
* Only the first page is returned; for the full multi-page document use the
|
||||
* autosaved chartXML instead.
|
||||
*/
|
||||
export function extractDiagramXML(xml_svg_string: string): string {
|
||||
try {
|
||||
// 1. Parse the SVG string (using built-in DOMParser in a browser-like environment)
|
||||
|
||||
+1
-1
@@ -130,7 +130,7 @@
|
||||
"electron-builder": "^26.0.12",
|
||||
"esbuild": "^0.28.0",
|
||||
"eslint": "9.39.5",
|
||||
"eslint-config-next": "16.1.6",
|
||||
"eslint-config-next": "16.3.8",
|
||||
"husky": "^9.1.7",
|
||||
"jsdom": "^27.4.0",
|
||||
"lint-staged": "^16.2.7",
|
||||
|
||||
Generated
-19
@@ -12,7 +12,6 @@
|
||||
"@modelcontextprotocol/sdk": "^1.0.4",
|
||||
"linkedom": "^0.18.0",
|
||||
"open": "^11.0.0",
|
||||
"saxes": "^6.0.0",
|
||||
"zod": "^4.0.0"
|
||||
},
|
||||
"bin": {
|
||||
@@ -2835,18 +2834,6 @@
|
||||
"integrity": "sha512-YZo3K82SD7Riyi0E1EQPojLz7kpepnSQI9IyPbHHg1XXXevb5dJI7tpyN2ADxGcQbHG7vcyRHk0cbwqcQriUtg==",
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/saxes": {
|
||||
"version": "6.0.0",
|
||||
"resolved": "https://registry.npmjs.org/saxes/-/saxes-6.0.0.tgz",
|
||||
"integrity": "sha512-xAg7SOnEhrm5zI3puOOKyy1OMcMlIJZYNJY7xLBwSze0UjhPLnWfj2GF2EpT0jmzaJKIWKHLsaSSajf35bcYnA==",
|
||||
"license": "ISC",
|
||||
"dependencies": {
|
||||
"xmlchars": "^2.2.0"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">=v12.22.7"
|
||||
}
|
||||
},
|
||||
"node_modules/send": {
|
||||
"version": "1.2.1",
|
||||
"resolved": "https://registry.npmjs.org/send/-/send-1.2.1.tgz",
|
||||
@@ -3392,12 +3379,6 @@
|
||||
"url": "https://github.com/sponsors/sindresorhus"
|
||||
}
|
||||
},
|
||||
"node_modules/xmlchars": {
|
||||
"version": "2.2.0",
|
||||
"resolved": "https://registry.npmjs.org/xmlchars/-/xmlchars-2.2.0.tgz",
|
||||
"integrity": "sha512-JZnDKK8B0RCDw84FNdDAIpZK+JuJw+s7Lz8nksI7SIuU3UXJJslUthsi+uWBUYOwPFwW7W7PRLRfUKpxjtjFCw==",
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/zod": {
|
||||
"version": "4.6.5",
|
||||
"resolved": "https://registry.npmjs.org/zod/-/zod-4.6.5.tgz",
|
||||
|
||||
@@ -41,7 +41,6 @@
|
||||
"@modelcontextprotocol/sdk": "^1.0.4",
|
||||
"linkedom": "^0.18.0",
|
||||
"open": "^11.0.0",
|
||||
"saxes": "^6.0.0",
|
||||
"zod": "^4.0.0"
|
||||
},
|
||||
"devDependencies": {
|
||||
|
||||
@@ -7,8 +7,6 @@
|
||||
* first page is targeted (the "active page by convention" — see pages.ts).
|
||||
*/
|
||||
|
||||
import { getXmlSyntaxError } from "./dom.js"
|
||||
import { log } from "./logger.js"
|
||||
import { findPageElement, hasPageSelector, type PageSelector } from "./pages.js"
|
||||
|
||||
export interface DiagramOperation {
|
||||
@@ -28,18 +26,6 @@ export interface ApplyOperationsResult {
|
||||
errors: OperationError[]
|
||||
}
|
||||
|
||||
// Cells with links, tooltips or custom data are stored as
|
||||
// <UserObject id="..."><mxCell .../></UserObject> (or <object>): the id sits
|
||||
// on the wrapper, so the wrapper is treated as the cell.
|
||||
const CELL_SELECTOR = "mxCell, UserObject, object"
|
||||
|
||||
/** Read parent/source/target, which a wrapped cell keeps on its inner mxCell. */
|
||||
function cellAttr(cell: Element, name: string): string | null {
|
||||
const inner =
|
||||
cell.tagName === "mxCell" ? cell : cell.querySelector("mxCell")
|
||||
return inner?.getAttribute(name) ?? null
|
||||
}
|
||||
|
||||
/**
|
||||
* Apply diagram operations (update/add/delete) using ID-based lookup.
|
||||
*
|
||||
@@ -57,8 +43,12 @@ export function applyDiagramOperations(
|
||||
): ApplyOperationsResult {
|
||||
const errors: OperationError[] = []
|
||||
|
||||
// Check for syntax errors, then parse the XML
|
||||
const parseError = getXmlSyntaxError(xmlContent)
|
||||
// Parse the XML
|
||||
const parser = new DOMParser()
|
||||
const doc = parser.parseFromString(xmlContent, "text/xml")
|
||||
|
||||
// Check for parse errors
|
||||
const parseError = doc.querySelector("parsererror")
|
||||
if (parseError) {
|
||||
return {
|
||||
result: xmlContent,
|
||||
@@ -66,13 +56,11 @@ export function applyDiagramOperations(
|
||||
{
|
||||
type: "update",
|
||||
cellId: "",
|
||||
message: `XML parse error: ${parseError}`,
|
||||
message: `XML parse error: ${parseError.textContent}`,
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
const parser = new DOMParser()
|
||||
const doc = parser.parseFromString(xmlContent, "text/xml")
|
||||
|
||||
// Locate the <root> element to operate on.
|
||||
//
|
||||
@@ -144,12 +132,10 @@ export function applyDiagramOperations(
|
||||
|
||||
// Build a map of cell IDs to elements (scoped to the resolved page).
|
||||
const cellMap = new Map<string, Element>()
|
||||
root.querySelectorAll(CELL_SELECTOR).forEach((cell) => {
|
||||
root.querySelectorAll("mxCell").forEach((cell) => {
|
||||
const id = cell.getAttribute("id")
|
||||
if (id) cellMap.set(id, cell)
|
||||
})
|
||||
// Ids deleted so far in this batch; deleting one again is a no-op
|
||||
const deletedIds = new Set<string>()
|
||||
|
||||
// Process each operation
|
||||
for (const op of operations) {
|
||||
@@ -178,7 +164,7 @@ export function applyDiagramOperations(
|
||||
`<wrapper>${op.new_xml}</wrapper>`,
|
||||
"text/xml",
|
||||
)
|
||||
const newCell = newDoc.querySelector(CELL_SELECTOR)
|
||||
const newCell = newDoc.querySelector("mxCell")
|
||||
if (!newCell) {
|
||||
errors.push({
|
||||
type: "update",
|
||||
@@ -230,7 +216,7 @@ export function applyDiagramOperations(
|
||||
`<wrapper>${op.new_xml}</wrapper>`,
|
||||
"text/xml",
|
||||
)
|
||||
const newCell = newDoc.querySelector(CELL_SELECTOR)
|
||||
const newCell = newDoc.querySelector("mxCell")
|
||||
if (!newCell) {
|
||||
errors.push({
|
||||
type: "add",
|
||||
@@ -270,15 +256,8 @@ export function applyDiagramOperations(
|
||||
|
||||
const existingCell = cellMap.get(op.cell_id)
|
||||
if (!existingCell) {
|
||||
// Skip cells already cascade-deleted by a previous operation
|
||||
// (AI may redundantly list children/edges); warn otherwise
|
||||
if (!deletedIds.has(op.cell_id)) {
|
||||
errors.push({
|
||||
type: "delete",
|
||||
cellId: op.cell_id,
|
||||
message: `Cell with id="${op.cell_id}" not found`,
|
||||
})
|
||||
}
|
||||
// Cell not found - might have been cascade-deleted by a previous operation
|
||||
// Skip silently instead of erroring (AI may redundantly list children/edges)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -291,17 +270,17 @@ export function applyDiagramOperations(
|
||||
cellsToDelete.add(cellId)
|
||||
|
||||
// Find children (cells where parent === cellId)
|
||||
// cellMap only holds this page's cells, so other pages' cells
|
||||
// with the same parent id (notably "1") are never touched.
|
||||
for (const [childId, child] of cellMap) {
|
||||
if (
|
||||
childId !== "0" &&
|
||||
childId !== "1" &&
|
||||
cellAttr(child, "parent") === cellId
|
||||
) {
|
||||
// Scoped to `root` so other pages' cells with the same parent id
|
||||
// (notably "1") are never touched.
|
||||
const children = root!.querySelectorAll(
|
||||
`mxCell[parent="${cellId}"]`,
|
||||
)
|
||||
children.forEach((child) => {
|
||||
const childId = child.getAttribute("id")
|
||||
if (childId && childId !== "0" && childId !== "1") {
|
||||
collectDescendants(childId)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Collect the target cell and all its descendants
|
||||
@@ -310,23 +289,23 @@ export function applyDiagramOperations(
|
||||
// Find edges referencing any of the cells to be deleted
|
||||
// Also recursively collect children of those edges (e.g., edge labels)
|
||||
for (const cellId of cellsToDelete) {
|
||||
for (const [edgeId, edge] of cellMap) {
|
||||
const referencingEdges = root.querySelectorAll(
|
||||
`mxCell[source="${cellId}"], mxCell[target="${cellId}"]`,
|
||||
)
|
||||
referencingEdges.forEach((edge) => {
|
||||
const edgeId = edge.getAttribute("id")
|
||||
// Protect root cells from being added via edge references
|
||||
if (edgeId === "0" || edgeId === "1") continue
|
||||
if (
|
||||
cellAttr(edge, "source") === cellId ||
|
||||
cellAttr(edge, "target") === cellId
|
||||
) {
|
||||
if (edgeId && edgeId !== "0" && edgeId !== "1") {
|
||||
// Recurse to collect edge's children (like labels)
|
||||
collectDescendants(edgeId)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Log what will be deleted (stderr: stdout carries JSON-RPC)
|
||||
// Log what will be deleted
|
||||
if (cellsToDelete.size > 1) {
|
||||
log.debug(
|
||||
`Cascade delete "${op.cell_id}" → deleting ${cellsToDelete.size} cells: ${Array.from(cellsToDelete).join(", ")}`,
|
||||
console.log(
|
||||
`[applyDiagramOperations] Cascade delete "${op.cell_id}" → deleting ${cellsToDelete.size} cells: ${Array.from(cellsToDelete).join(", ")}`,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -337,7 +316,6 @@ export function applyDiagramOperations(
|
||||
cell.parentNode?.removeChild(cell)
|
||||
cellMap.delete(cellId)
|
||||
}
|
||||
deletedIds.add(cellId)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,89 +0,0 @@
|
||||
/**
|
||||
* DOM setup for Node.
|
||||
*
|
||||
* linkedom gives us a DOM with querySelector, but it is lenient: it never
|
||||
* reports syntax errors (no <parsererror>), and its serializer writes raw
|
||||
* newlines inside attribute values, which the browser reads back as spaces.
|
||||
* saxes, a strict XML parser, checks well-formedness the way draw.io's
|
||||
* DOMParser will, and serializeXml writes attribute values safely.
|
||||
*/
|
||||
import { DOMParser } from "linkedom"
|
||||
import { SaxesParser } from "saxes"
|
||||
|
||||
/**
|
||||
* Returns the first XML syntax error as "line:column: message", or null if
|
||||
* the XML is well-formed. Surrounding whitespace is ignored because every
|
||||
* caller trims before the XML reaches the browser.
|
||||
*/
|
||||
export function getXmlSyntaxError(xml: string): string | null {
|
||||
let error: string | null = null
|
||||
const parser = new SaxesParser()
|
||||
parser.on("error", (err) => {
|
||||
error ??= err.message
|
||||
})
|
||||
parser.write(xml.trim()).close()
|
||||
return error
|
||||
}
|
||||
|
||||
const ESCAPES: Record<string, string> = {
|
||||
"&": "&",
|
||||
"<": "<",
|
||||
">": ">",
|
||||
'"': """,
|
||||
"\t": "	",
|
||||
"\n": "
",
|
||||
"\r": "
",
|
||||
}
|
||||
const escapeChars = (text: string, chars: RegExp) =>
|
||||
text.replace(chars, (c) => ESCAPES[c])
|
||||
|
||||
/**
|
||||
* Serialize a linkedom node as XML. Attribute values escape tabs and line
|
||||
* breaks too, so multi-line labels (value="a
b") survive a round trip.
|
||||
*/
|
||||
export function serializeXml(node: Node): string {
|
||||
switch (node.nodeType) {
|
||||
case 9: {
|
||||
// Document
|
||||
const root = (node as Document).documentElement
|
||||
return root ? serializeXml(root) : ""
|
||||
}
|
||||
case 1: {
|
||||
// Element
|
||||
const el = node as Element
|
||||
let out = `<${el.tagName}`
|
||||
for (const attr of Array.from(el.attributes)) {
|
||||
out += ` ${attr.name}="${escapeChars(attr.value, /[&<>"\t\n\r]/g)}"`
|
||||
}
|
||||
if (el.childNodes.length === 0) return `${out}/>`
|
||||
out += ">"
|
||||
for (const child of Array.from(el.childNodes)) {
|
||||
out += serializeXml(child)
|
||||
}
|
||||
return `${out}</${el.tagName}>`
|
||||
}
|
||||
case 3:
|
||||
// Text
|
||||
return escapeChars(node.textContent ?? "", /[&<>]/g)
|
||||
case 4:
|
||||
// CDATA
|
||||
return `<![CDATA[${node.textContent ?? ""}]]>`
|
||||
case 8:
|
||||
// Comment
|
||||
return `<!--${node.textContent ?? ""}-->`
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
class XMLSerializerPolyfill {
|
||||
serializeToString(node: Node): string {
|
||||
return serializeXml(node)
|
||||
}
|
||||
}
|
||||
|
||||
/** Install the DOMParser and XMLSerializer globals the XML helpers use. */
|
||||
export function installDomPolyfill(): void {
|
||||
;(globalThis as any).DOMParser = DOMParser
|
||||
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
|
||||
}
|
||||
@@ -6,15 +6,7 @@
|
||||
import { log } from "./logger.js"
|
||||
|
||||
const MAX_HISTORY = 20
|
||||
|
||||
interface HistoryEntry {
|
||||
id: number // Stable across shifts of the circular buffer
|
||||
xml: string
|
||||
svg: string
|
||||
}
|
||||
|
||||
let nextEntryId = 0
|
||||
const historyStore = new Map<string, HistoryEntry[]>()
|
||||
const historyStore = new Map<string, Array<{ xml: string; svg: string }>>()
|
||||
|
||||
export function addHistory(sessionId: string, xml: string, svg = ""): number {
|
||||
let history = historyStore.get(sessionId)
|
||||
@@ -29,7 +21,7 @@ export function addHistory(sessionId: string, xml: string, svg = ""): number {
|
||||
return history.length - 1
|
||||
}
|
||||
|
||||
history.push({ id: nextEntryId++, xml, svg })
|
||||
history.push({ xml, svg })
|
||||
|
||||
// Circular buffer
|
||||
if (history.length > MAX_HISTORY) {
|
||||
@@ -40,16 +32,18 @@ export function addHistory(sessionId: string, xml: string, svg = ""): number {
|
||||
return history.length - 1
|
||||
}
|
||||
|
||||
export function getHistory(sessionId: string): HistoryEntry[] {
|
||||
export function getHistory(
|
||||
sessionId: string,
|
||||
): Array<{ xml: string; svg: string }> {
|
||||
return historyStore.get(sessionId) || []
|
||||
}
|
||||
|
||||
/** Look up an entry by its id; the array index shifts as old entries drop. */
|
||||
export function getHistoryEntry(
|
||||
sessionId: string,
|
||||
id: number,
|
||||
): HistoryEntry | undefined {
|
||||
return historyStore.get(sessionId)?.find((entry) => entry.id === id)
|
||||
index: number,
|
||||
): { xml: string; svg: string } | undefined {
|
||||
const history = historyStore.get(sessionId)
|
||||
return history?.[index]
|
||||
}
|
||||
|
||||
export function clearHistory(sessionId: string): void {
|
||||
|
||||
@@ -12,9 +12,7 @@ function readBody(
|
||||
res: http.ServerResponse,
|
||||
cb: (body: string) => void,
|
||||
): void {
|
||||
// Decode once at the end: a multi-byte UTF-8 character can be split
|
||||
// across two chunks.
|
||||
const chunks: Buffer[] = []
|
||||
let body = ""
|
||||
let size = 0
|
||||
req.on("data", (chunk: Buffer) => {
|
||||
size += chunk.length
|
||||
@@ -24,9 +22,9 @@ function readBody(
|
||||
req.destroy()
|
||||
return
|
||||
}
|
||||
chunks.push(chunk)
|
||||
body += chunk
|
||||
})
|
||||
req.on("end", () => cb(Buffer.concat(chunks).toString("utf8")))
|
||||
req.on("end", () => cb(body))
|
||||
}
|
||||
|
||||
import {
|
||||
@@ -64,11 +62,9 @@ function normalizeUrl(url: string): string {
|
||||
return url.replace(/\/$/, "")
|
||||
}
|
||||
|
||||
// Session ids look like "mcp-<base36 time>-<base36 random>" (start_session).
|
||||
// Only this charset is accepted, because ids are written into the page's
|
||||
// HTML and script and into the redirect Location header.
|
||||
function isValidSessionId(sessionId: string): boolean {
|
||||
return /^mcp-[a-z0-9-]{1,64}$/.test(sessionId)
|
||||
function isLikelyMcpSessionId(sessionId: string): boolean {
|
||||
// Keep this cheap and conservative to avoid creating state for arbitrary IDs.
|
||||
return sessionId.startsWith("mcp-") && sessionId.length <= 128
|
||||
}
|
||||
|
||||
// Find the most recent active session (for auto-redirect when no sessionId provided)
|
||||
@@ -84,7 +80,7 @@ function getMostRecentSessionId(): string | null {
|
||||
|
||||
function ensureSessionStateInitialized(sessionId: string): void {
|
||||
if (!sessionId) return
|
||||
if (!isValidSessionId(sessionId)) return
|
||||
if (!isLikelyMcpSessionId(sessionId)) return
|
||||
if (stateStore.has(sessionId)) return
|
||||
|
||||
setState(sessionId, DEFAULT_DIAGRAM_XML)
|
||||
@@ -93,11 +89,7 @@ function ensureSessionStateInitialized(sessionId: string): void {
|
||||
interface SessionState {
|
||||
xml: string
|
||||
version: number
|
||||
// Version of the last write the browser did not make itself (AI edit,
|
||||
// restore). A browser push based on an older version is rejected.
|
||||
serverVersion?: number
|
||||
lastUpdated: Date
|
||||
lastPolled?: number // Last browser poll; an open tab keeps the session alive
|
||||
svg?: string // Cached SVG from last browser save
|
||||
syncRequested?: number // Timestamp when sync requested, cleared when browser responds
|
||||
exportFormat?: "png" | "svg" // Set by MCP tool to request browser export
|
||||
@@ -116,20 +108,13 @@ export function getState(sessionId: string): SessionState | undefined {
|
||||
return stateStore.get(sessionId)
|
||||
}
|
||||
|
||||
export function setState(
|
||||
sessionId: string,
|
||||
xml: string,
|
||||
svg?: string,
|
||||
fromBrowser = false,
|
||||
): number {
|
||||
export function setState(sessionId: string, xml: string, svg?: string): number {
|
||||
const existing = stateStore.get(sessionId)
|
||||
const newVersion = (existing?.version || 0) + 1
|
||||
stateStore.set(sessionId, {
|
||||
xml,
|
||||
version: newVersion,
|
||||
serverVersion: fromBrowser ? existing?.serverVersion : newVersion,
|
||||
lastUpdated: new Date(),
|
||||
lastPolled: existing?.lastPolled,
|
||||
svg: svg || existing?.svg, // Preserve cached SVG if not provided
|
||||
syncRequested: undefined, // Clear sync request when browser pushes state
|
||||
exportFormat: existing?.exportFormat, // Preserve pending export request
|
||||
@@ -237,11 +222,7 @@ export function stopHttpServer(): void {
|
||||
function cleanupExpiredSessions(): void {
|
||||
const now = Date.now()
|
||||
for (const [sessionId, state] of stateStore) {
|
||||
const lastActive = Math.max(
|
||||
state.lastUpdated.getTime(),
|
||||
state.lastPolled ?? 0,
|
||||
)
|
||||
if (now - lastActive > SESSION_TTL) {
|
||||
if (now - state.lastUpdated.getTime() > SESSION_TTL) {
|
||||
stateStore.delete(sessionId)
|
||||
clearHistory(sessionId)
|
||||
log.info(`Cleaned up expired session: ${sessionId}`)
|
||||
@@ -264,48 +245,7 @@ function handleRequest(
|
||||
req: http.IncomingMessage,
|
||||
res: http.ServerResponse,
|
||||
): void {
|
||||
// A bad request must never take down the MCP process
|
||||
try {
|
||||
routeRequest(req, res)
|
||||
} catch (err) {
|
||||
log.error("HTTP request failed:", err)
|
||||
if (!res.headersSent) res.writeHead(500)
|
||||
res.end()
|
||||
}
|
||||
}
|
||||
|
||||
// Serve only requests addressed to localhost, sent by a localhost page or by
|
||||
// a non-browser client (no Origin header). This blocks DNS rebinding and
|
||||
// scripts on other websites.
|
||||
function isLocalRequest(req: http.IncomingMessage): boolean {
|
||||
const isLocalHost = (host: string) =>
|
||||
/^(localhost|127\.0\.0\.1)(:\d+)?$/.test(host)
|
||||
const origin = req.headers.origin
|
||||
return (
|
||||
isLocalHost(req.headers.host ?? "") &&
|
||||
(origin === undefined || isLocalHost(origin.replace(/^http:\/\//, "")))
|
||||
)
|
||||
}
|
||||
|
||||
function routeRequest(
|
||||
req: http.IncomingMessage,
|
||||
res: http.ServerResponse,
|
||||
): void {
|
||||
let url: URL
|
||||
try {
|
||||
url = new URL(req.url || "/", `http://localhost:${serverPort}`)
|
||||
} catch {
|
||||
// e.g. "//" is not a valid URL path
|
||||
res.writeHead(400)
|
||||
res.end("Bad Request")
|
||||
return
|
||||
}
|
||||
|
||||
if (!isLocalRequest(req)) {
|
||||
res.writeHead(403)
|
||||
res.end("Forbidden")
|
||||
return
|
||||
}
|
||||
const url = new URL(req.url || "/", `http://localhost:${serverPort}`)
|
||||
|
||||
const requestOrigin = req.headers.origin
|
||||
if (requestOrigin === `http://localhost:${serverPort}`) {
|
||||
@@ -322,19 +262,12 @@ function routeRequest(
|
||||
|
||||
if (url.pathname === "/" || url.pathname === "/index.html") {
|
||||
const sessionId = url.searchParams.get("mcp") || ""
|
||||
if (sessionId && !isValidSessionId(sessionId)) {
|
||||
res.writeHead(400)
|
||||
res.end("Invalid session id")
|
||||
return
|
||||
}
|
||||
|
||||
// Auto-redirect to most recent session if no sessionId provided
|
||||
if (!sessionId) {
|
||||
const recentSessionId = getMostRecentSessionId()
|
||||
if (recentSessionId) {
|
||||
res.writeHead(302, {
|
||||
Location: `/?mcp=${encodeURIComponent(recentSessionId)}`,
|
||||
})
|
||||
res.writeHead(302, { Location: `/?mcp=${recentSessionId}` })
|
||||
res.end()
|
||||
return
|
||||
}
|
||||
@@ -372,9 +305,6 @@ function handleStateApi(
|
||||
}
|
||||
ensureSessionStateInitialized(sessionId)
|
||||
const state = stateStore.get(sessionId)
|
||||
// Polling counts as activity, so a session stays alive while its
|
||||
// tab is open
|
||||
if (state) state.lastPolled = Date.now()
|
||||
res.writeHead(200, { "Content-Type": "application/json" })
|
||||
res.end(
|
||||
JSON.stringify({
|
||||
@@ -390,11 +320,9 @@ function handleStateApi(
|
||||
try {
|
||||
const data = JSON.parse(body)
|
||||
const { sessionId } = data
|
||||
if (!sessionId || !isValidSessionId(sessionId)) {
|
||||
if (!sessionId) {
|
||||
res.writeHead(400, { "Content-Type": "application/json" })
|
||||
res.end(
|
||||
JSON.stringify({ error: "valid sessionId required" }),
|
||||
)
|
||||
res.end(JSON.stringify({ error: "sessionId required" }))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -414,25 +342,7 @@ function handleStateApi(
|
||||
return
|
||||
}
|
||||
|
||||
// The browser edited a version older than the latest AI write
|
||||
// (it has not loaded that write yet). Keep the AI write; the
|
||||
// browser loads it on its next poll.
|
||||
const current = stateStore.get(sessionId)
|
||||
if (
|
||||
typeof data.baseVersion === "number" &&
|
||||
data.baseVersion < (current?.serverVersion ?? 0)
|
||||
) {
|
||||
res.writeHead(409, { "Content-Type": "application/json" })
|
||||
res.end(
|
||||
JSON.stringify({
|
||||
error: "Diagram changed on the server",
|
||||
version: current?.version,
|
||||
}),
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
const version = setState(sessionId, data.xml, data.svg, true)
|
||||
const version = setState(sessionId, data.xml, data.svg)
|
||||
res.writeHead(200, { "Content-Type": "application/json" })
|
||||
res.end(JSON.stringify({ success: true, version }))
|
||||
} catch {
|
||||
@@ -468,11 +378,7 @@ function handleHistoryApi(
|
||||
res.writeHead(200, { "Content-Type": "application/json" })
|
||||
res.end(
|
||||
JSON.stringify({
|
||||
entries: history.map((entry, i) => ({
|
||||
index: i,
|
||||
id: entry.id,
|
||||
svg: entry.svg,
|
||||
})),
|
||||
entries: history.map((entry, i) => ({ index: i, svg: entry.svg })),
|
||||
count: history.length,
|
||||
}),
|
||||
)
|
||||
@@ -490,14 +396,16 @@ function handleRestoreApi(
|
||||
|
||||
readBody(req, res, (body) => {
|
||||
try {
|
||||
const { sessionId, id } = JSON.parse(body)
|
||||
if (!sessionId || typeof id !== "number") {
|
||||
const { sessionId, index } = JSON.parse(body)
|
||||
if (!sessionId || index === undefined) {
|
||||
res.writeHead(400, { "Content-Type": "application/json" })
|
||||
res.end(JSON.stringify({ error: "sessionId and id required" }))
|
||||
res.end(
|
||||
JSON.stringify({ error: "sessionId and index required" }),
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
const entry = getHistoryEntry(sessionId, id)
|
||||
const entry = getHistoryEntry(sessionId, index)
|
||||
if (!entry) {
|
||||
res.writeHead(404, { "Content-Type": "application/json" })
|
||||
res.end(JSON.stringify({ error: "Entry not found" }))
|
||||
@@ -507,7 +415,7 @@ function handleRestoreApi(
|
||||
const newVersion = setState(sessionId, entry.xml)
|
||||
addHistory(sessionId, entry.xml, entry.svg)
|
||||
|
||||
log.info(`Restored session ${sessionId} to history entry ${id}`)
|
||||
log.info(`Restored session ${sessionId} to index ${index}`)
|
||||
|
||||
res.writeHead(200, { "Content-Type": "application/json" })
|
||||
res.end(JSON.stringify({ success: true, newVersion }))
|
||||
@@ -789,11 +697,10 @@ function getHtmlPage(sessionId: string): string {
|
||||
</div>
|
||||
</div>
|
||||
<script>
|
||||
const sessionId = ${JSON.stringify(sessionId).replace(/</g, "\\u003c")};
|
||||
const sessionId = "${sessionId}";
|
||||
const iframe = document.getElementById('drawio');
|
||||
let currentVersion = 0, isReady = false, pendingXml = null, lastXml = null;
|
||||
let pendingSvgExport = null;
|
||||
let pendingSvgBase = 0; // version the pending autosave was based on
|
||||
let pendingAiSvg = false;
|
||||
let pendingMcpExport = null; // 'png' or 'svg' when MCP requested export
|
||||
let projectionExportActive = false; // page-targeted export: showing a transient single-page projection
|
||||
@@ -811,29 +718,18 @@ function getHtmlPage(sessionId: string): string {
|
||||
// for a page-targeted export — otherwise we'd push the
|
||||
// transient projection back as the canonical session state.
|
||||
if (projectionExportActive) return;
|
||||
// Request SVG export, then push state with SVG. Remember the
|
||||
// version this edit is based on, so the server can reject it
|
||||
// if the AI wrote a newer version that is not loaded yet.
|
||||
// Request SVG export, then push state with SVG
|
||||
pendingSvgExport = msg.xml;
|
||||
pendingSvgBase = currentVersion;
|
||||
iframe.contentWindow.postMessage(JSON.stringify({ action: 'export', format: 'svg' }), '*');
|
||||
// Fallback if export doesn't respond
|
||||
setTimeout(() => { if (pendingSvgExport === msg.xml) { pushState(msg.xml, '', pendingSvgBase); pendingSvgExport = null; } }, 2000);
|
||||
} else if (msg.event === 'export' && msg.format === 'xml') {
|
||||
// Sync export requested by the server (get_diagram).
|
||||
// draw.io returns the XML in msg.xml, with no msg.data.
|
||||
if (pendingSyncExport && msg.xml) {
|
||||
pendingSyncExport = false;
|
||||
pushState(msg.xml, '');
|
||||
}
|
||||
setTimeout(() => { if (pendingSvgExport === msg.xml) { pushState(msg.xml, ''); pendingSvgExport = null; } }, 2000);
|
||||
} else if (msg.event === 'export' && msg.data) {
|
||||
// Handle MCP server export request (png/svg). fireExport tags
|
||||
// the request with mcpExport and draw.io echoes the request
|
||||
// back in msg.message, which tells it apart from autosave and
|
||||
// preview SVG exports.
|
||||
if (msg.message && msg.message.mcpExport) {
|
||||
// Handle MCP server export request (png/svg)
|
||||
// Verify the response matches the requested format to avoid capturing
|
||||
// unrelated exports (autosave SVG, sync XML)
|
||||
if (pendingMcpExport) {
|
||||
const d = msg.data;
|
||||
const isPng = pendingMcpExport === 'png' && d.startsWith('data:image/png');
|
||||
const isPng = pendingMcpExport === 'png' && (d.startsWith('data:image/png') || (typeof d === 'string' && d.length > 100 && !d.startsWith('<')));
|
||||
const isSvg = pendingMcpExport === 'svg' && (d.startsWith('data:image/svg') || d.startsWith('<svg'));
|
||||
if (isPng || isSvg) {
|
||||
pendingMcpExport = null;
|
||||
@@ -845,8 +741,8 @@ function getHtmlPage(sessionId: string): string {
|
||||
// Page-targeted export: restore the user's real
|
||||
// multi-page document now that we have the image.
|
||||
restoreFromProjection();
|
||||
return;
|
||||
}
|
||||
return;
|
||||
}
|
||||
// Handle file download export (PNG/SVG only, drawio uses lastXml directly)
|
||||
if (pendingDownload && (pendingDownload.format === 'png' || pendingDownload.format === 'svg')) {
|
||||
@@ -865,13 +761,19 @@ function getHtmlPage(sessionId: string): string {
|
||||
saveConfirmBtn.textContent = 'Save';
|
||||
return;
|
||||
}
|
||||
// Handle sync export (XML format) - server requested fresh state
|
||||
if (pendingSyncExport && !msg.data.startsWith('data:') && !msg.data.startsWith('<svg')) {
|
||||
pendingSyncExport = false;
|
||||
pushState(msg.data, '');
|
||||
return;
|
||||
}
|
||||
// Handle SVG export
|
||||
let svg = msg.data;
|
||||
if (!svg.startsWith('data:')) svg = 'data:image/svg+xml;base64,' + btoa(unescape(encodeURIComponent(svg)));
|
||||
if (pendingSvgExport) {
|
||||
const xml = pendingSvgExport;
|
||||
pendingSvgExport = null;
|
||||
pushState(xml, svg, pendingSvgBase);
|
||||
pushState(xml, svg);
|
||||
} else if (pendingAiSvg) {
|
||||
pendingAiSvg = false;
|
||||
fetch('/api/history-svg', {
|
||||
@@ -912,17 +814,15 @@ function getHtmlPage(sessionId: string): string {
|
||||
}
|
||||
}
|
||||
|
||||
async function pushState(xml, svg = '', baseVersion = currentVersion) {
|
||||
async function pushState(xml, svg = '') {
|
||||
if (!sessionId) return;
|
||||
try {
|
||||
const r = await fetch('/api/state', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ sessionId, xml, svg, baseVersion })
|
||||
body: JSON.stringify({ sessionId, xml, svg })
|
||||
});
|
||||
if (r.ok) { const d = await r.json(); currentVersion = d.version; lastXml = xml; }
|
||||
// 409: the AI wrote a newer version; load it now
|
||||
else if (r.status === 409) poll();
|
||||
} catch (e) { console.error('Push failed:', e); }
|
||||
}
|
||||
|
||||
@@ -930,22 +830,14 @@ function getHtmlPage(sessionId: string): string {
|
||||
|
||||
async function poll() {
|
||||
if (!sessionId) return;
|
||||
const knownVersion = currentVersion;
|
||||
try {
|
||||
const r = await fetch('/api/state?sessionId=' + encodeURIComponent(sessionId));
|
||||
if (!r.ok) return;
|
||||
const s = await r.json();
|
||||
// Handle sync request - server needs fresh state. Reset after a
|
||||
// while in case draw.io never answers, so later syncs still run.
|
||||
if (s.syncRequested && !pendingSyncExport && isReady) {
|
||||
// Handle sync request - server needs fresh state
|
||||
if (s.syncRequested && !pendingSyncExport) {
|
||||
pendingSyncExport = true;
|
||||
iframe.contentWindow.postMessage(JSON.stringify({ action: 'export', format: 'xml' }), '*');
|
||||
setTimeout(() => { pendingSyncExport = false; }, 5000);
|
||||
}
|
||||
// The server lost this session (e.g. it expired) and rebuilt it
|
||||
// with a blank diagram: push back what the browser shows.
|
||||
if (s.version < knownVersion && lastXml) {
|
||||
pushState(lastXml);
|
||||
}
|
||||
// Load new diagram from server (before export, so we export latest).
|
||||
// While a page-targeted projection is on screen, skip the reload
|
||||
@@ -970,10 +862,9 @@ function getHtmlPage(sessionId: string): string {
|
||||
if (s.exportFormat && !pendingMcpExport && isReady) {
|
||||
pendingMcpExport = s.exportFormat;
|
||||
const fireExport = () => {
|
||||
// mcpExport is echoed back in msg.message (see the handler)
|
||||
const exportOpts = pendingMcpExport === 'png'
|
||||
? { action: 'export', format: 'png', scale: 2, mcpExport: true }
|
||||
: { action: 'export', format: 'svg', mcpExport: true };
|
||||
? { action: 'export', format: 'png', scale: 2 }
|
||||
: { action: 'export', format: 'svg' };
|
||||
iframe.contentWindow.postMessage(JSON.stringify(exportOpts), '*');
|
||||
};
|
||||
if (s.exportXml) {
|
||||
@@ -1071,7 +962,7 @@ function getHtmlPage(sessionId: string): string {
|
||||
const historyEmpty = document.getElementById('history-empty');
|
||||
const restoreBtn = document.getElementById('restore-btn');
|
||||
const cancelBtn = document.getElementById('cancel-btn');
|
||||
let historyData = [], selectedId = null;
|
||||
let historyData = [], selectedIdx = null;
|
||||
|
||||
historyBtn.onclick = async () => {
|
||||
if (!sessionId) return;
|
||||
@@ -1086,7 +977,7 @@ function getHtmlPage(sessionId: string): string {
|
||||
historyModal.classList.add('open');
|
||||
};
|
||||
|
||||
cancelBtn.onclick = () => { historyModal.classList.remove('open'); selectedId = null; restoreBtn.disabled = true; };
|
||||
cancelBtn.onclick = () => { historyModal.classList.remove('open'); selectedIdx = null; restoreBtn.disabled = true; };
|
||||
historyModal.onclick = (e) => { if (e.target === historyModal) cancelBtn.onclick(); };
|
||||
|
||||
function renderHistory() {
|
||||
@@ -1098,30 +989,30 @@ function getHtmlPage(sessionId: string): string {
|
||||
historyGrid.style.display = 'grid';
|
||||
historyEmpty.style.display = 'none';
|
||||
historyGrid.innerHTML = historyData.map((e, i) => \`
|
||||
<div class="history-item" data-id="\${e.id}">
|
||||
<div class="history-item" data-idx="\${e.index}">
|
||||
<div class="thumb">\${e.svg ? \`<img src="\${e.svg}">\` : '#' + e.index}</div>
|
||||
<div class="label">#\${e.index}</div>
|
||||
</div>
|
||||
\`).join('');
|
||||
historyGrid.querySelectorAll('.history-item').forEach(item => {
|
||||
item.onclick = () => {
|
||||
const id = parseInt(item.dataset.id);
|
||||
if (selectedId === id) { selectedId = null; restoreBtn.disabled = true; }
|
||||
else { selectedId = id; restoreBtn.disabled = false; }
|
||||
historyGrid.querySelectorAll('.history-item').forEach(el => el.classList.toggle('selected', parseInt(el.dataset.id) === selectedId));
|
||||
const idx = parseInt(item.dataset.idx);
|
||||
if (selectedIdx === idx) { selectedIdx = null; restoreBtn.disabled = true; }
|
||||
else { selectedIdx = idx; restoreBtn.disabled = false; }
|
||||
historyGrid.querySelectorAll('.history-item').forEach(el => el.classList.toggle('selected', parseInt(el.dataset.idx) === selectedIdx));
|
||||
};
|
||||
});
|
||||
}
|
||||
|
||||
restoreBtn.onclick = async () => {
|
||||
if (selectedId === null) return;
|
||||
if (selectedIdx === null) return;
|
||||
restoreBtn.disabled = true;
|
||||
restoreBtn.textContent = 'Restoring...';
|
||||
try {
|
||||
const r = await fetch('/api/restore', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ sessionId, id: selectedId })
|
||||
body: JSON.stringify({ sessionId, index: selectedIdx })
|
||||
});
|
||||
if (r.ok) { cancelBtn.onclick(); await poll(); }
|
||||
else { alert('Restore failed'); }
|
||||
|
||||
@@ -18,6 +18,24 @@
|
||||
* surface.
|
||||
*/
|
||||
|
||||
// Setup DOM polyfill for Node.js (required for XML operations)
|
||||
import { DOMParser } from "linkedom"
|
||||
;(globalThis as any).DOMParser = DOMParser
|
||||
|
||||
// Create XMLSerializer polyfill using outerHTML
|
||||
class XMLSerializerPolyfill {
|
||||
serializeToString(node: any): string {
|
||||
if (node.outerHTML !== undefined) {
|
||||
return node.outerHTML
|
||||
}
|
||||
if (node.documentElement) {
|
||||
return node.documentElement.outerHTML
|
||||
}
|
||||
return ""
|
||||
}
|
||||
}
|
||||
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
|
||||
|
||||
import { createRequire } from "node:module"
|
||||
import { McpServer } from "@modelcontextprotocol/sdk/server/mcp.js"
|
||||
import { StdioServerTransport } from "@modelcontextprotocol/sdk/server/stdio.js"
|
||||
@@ -27,7 +45,6 @@ import {
|
||||
applyDiagramOperations,
|
||||
type DiagramOperation,
|
||||
} from "./diagram-operations.js"
|
||||
import { installDomPolyfill } from "./dom.js"
|
||||
import { checkEditGate } from "./edit-gate.js"
|
||||
import { addHistory } from "./history.js"
|
||||
import {
|
||||
@@ -55,9 +72,6 @@ import {
|
||||
} from "./pages.js"
|
||||
import { validateAndFixXml } from "./xml-validation.js"
|
||||
|
||||
// DOMParser/XMLSerializer globals for the XML helpers (Node has neither)
|
||||
installDomPolyfill()
|
||||
|
||||
// Server configuration
|
||||
const config = {
|
||||
port: parseInt(process.env.PORT || "6002", 10),
|
||||
@@ -894,47 +908,6 @@ server.registerTool(
|
||||
},
|
||||
)
|
||||
|
||||
// The browser bridge has one export slot per session, so export requests
|
||||
// run one at a time: a concurrent call waits for the previous one.
|
||||
let exportQueue: Promise<unknown> = Promise.resolve()
|
||||
|
||||
/**
|
||||
* Ask the browser to export (optionally via a page projection) and poll for
|
||||
* the resulting image data. Resolves to undefined on timeout.
|
||||
*/
|
||||
function exportViaBrowser(
|
||||
sessionId: string,
|
||||
format: "png" | "svg",
|
||||
projectionXml?: string,
|
||||
): Promise<string | undefined> {
|
||||
const run = exportQueue.then(async () => {
|
||||
requestExport(sessionId, format, projectionXml)
|
||||
|
||||
// A projection export does an extra load + render round-trip in the
|
||||
// browser, so give it a longer window. Re-read the live store entry
|
||||
// each tick: setState() (from a concurrent autosave or tool call)
|
||||
// replaces the Map entry with a new object, so a captured reference
|
||||
// would go stale and never observe the browser's exportData.
|
||||
const timeoutMs = projectionXml ? 15000 : 10000
|
||||
const start = Date.now()
|
||||
let exportData: string | undefined
|
||||
while (Date.now() - start < timeoutMs) {
|
||||
exportData = getState(sessionId)?.exportData
|
||||
if (exportData) break
|
||||
await new Promise((r) => setTimeout(r, 200))
|
||||
}
|
||||
const live = getState(sessionId)
|
||||
if (live) {
|
||||
live.exportData = undefined
|
||||
live.exportFormat = undefined
|
||||
live.exportXml = undefined
|
||||
}
|
||||
return exportData
|
||||
})
|
||||
exportQueue = run.catch(() => {})
|
||||
return run
|
||||
}
|
||||
|
||||
// Tool: export_diagram
|
||||
server.registerTool(
|
||||
"export_diagram",
|
||||
@@ -1106,12 +1079,34 @@ server.registerTool(
|
||||
projectionXml = projection.xml
|
||||
}
|
||||
|
||||
const exportData = await exportViaBrowser(
|
||||
// Ask the browser to export (optionally via a page projection) and
|
||||
// poll for the resulting image data.
|
||||
requestExport(
|
||||
currentSession.id,
|
||||
detectedFormat as "png" | "svg",
|
||||
projectionXml,
|
||||
)
|
||||
|
||||
// A projection export does an extra load + render round-trip in the
|
||||
// browser, so give it a longer window. Re-read the live store entry
|
||||
// each tick: setState() (from a concurrent autosave or tool call)
|
||||
// replaces the Map entry with a new object, so a captured reference
|
||||
// would go stale and never observe the browser's exportData.
|
||||
const timeoutMs = projectionXml ? 15000 : 10000
|
||||
const start = Date.now()
|
||||
let exportData: string | undefined
|
||||
while (Date.now() - start < timeoutMs) {
|
||||
exportData = getState(currentSession.id)?.exportData
|
||||
if (exportData) break
|
||||
await new Promise((r) => setTimeout(r, 200))
|
||||
}
|
||||
const live = getState(currentSession.id)
|
||||
if (live) {
|
||||
live.exportData = undefined
|
||||
live.exportFormat = undefined
|
||||
live.exportXml = undefined
|
||||
}
|
||||
|
||||
if (!exportData) {
|
||||
return {
|
||||
content: [
|
||||
@@ -1220,20 +1215,15 @@ async function loadMxfileForMutation(): Promise<
|
||||
doc,
|
||||
writeBack: (newDoc: Document) => {
|
||||
const newXml = serializeMxfile(newDoc)
|
||||
// The store may hold user edits the model has not seen yet.
|
||||
const sawLatest = checkEditGate(
|
||||
sessionRef.lastSeenXml,
|
||||
browserState?.xml ?? "",
|
||||
).ok
|
||||
// Save history before overwriting so the user can undo.
|
||||
addHistory(sessionRef.id, sessionRef.xml, browserState?.svg || "")
|
||||
sessionRef.xml = newXml
|
||||
sessionRef.version++
|
||||
setState(sessionRef.id, newXml)
|
||||
// The model just wrote this exact state. If it had seen the state
|
||||
// it built on, mark the result as seen so edit_diagram needs no
|
||||
// extra get_diagram; otherwise edit_diagram must ask for one.
|
||||
sessionRef.lastSeenXml = sawLatest ? newXml : ""
|
||||
// The model just wrote this exact state, so mark it as seen —
|
||||
// subsequent edit_diagram calls don't need a redundant
|
||||
// get_diagram round-trip.
|
||||
sessionRef.lastSeenXml = newXml
|
||||
addHistory(sessionRef.id, newXml, "")
|
||||
},
|
||||
}
|
||||
|
||||
@@ -9,7 +9,6 @@
|
||||
*/
|
||||
import { inflateRawSync } from "node:zlib"
|
||||
import { DOMParser } from "linkedom"
|
||||
import { getXmlSyntaxError } from "./dom.js"
|
||||
import {
|
||||
isMxFile,
|
||||
isMxGraphModel,
|
||||
@@ -83,7 +82,7 @@ export function parseDrawioFileContent(content: string): LoadResult {
|
||||
}
|
||||
const inner = new DOMParser().parseFromString(xml, "text/xml")
|
||||
if (
|
||||
getXmlSyntaxError(xml) ||
|
||||
inner.querySelector("parsererror") ||
|
||||
inner.documentElement?.tagName !== "mxGraphModel"
|
||||
) {
|
||||
return {
|
||||
|
||||
@@ -18,7 +18,6 @@
|
||||
*/
|
||||
|
||||
import { DOMParser } from "linkedom"
|
||||
import { getXmlSyntaxError } from "./dom.js"
|
||||
|
||||
export interface PageInfo {
|
||||
id: string
|
||||
@@ -111,8 +110,8 @@ export function normalizeToMxfile(
|
||||
*/
|
||||
export function parseMxfile(xml: string): Document | null {
|
||||
try {
|
||||
if (getXmlSyntaxError(xml)) return null
|
||||
const doc = new DOMParser().parseFromString(xml, "text/xml")
|
||||
if (doc.querySelector("parsererror")) return null
|
||||
if (doc.documentElement?.tagName !== "mxfile") return null
|
||||
return doc as unknown as Document
|
||||
} catch {
|
||||
@@ -259,12 +258,12 @@ export function addPageToDoc(
|
||||
}
|
||||
|
||||
const snippet = `<wrapper><diagram id="${escapeAttr(id)}" name="${escapeAttr(name)}">${inner}</diagram></wrapper>`
|
||||
if (getXmlSyntaxError(snippet)) {
|
||||
const tempDoc = new DOMParser().parseFromString(snippet, "text/xml")
|
||||
if (tempDoc.querySelector("parsererror")) {
|
||||
throw new Error(
|
||||
"Failed to parse new page xml — make sure it is a valid <mxGraphModel>",
|
||||
)
|
||||
}
|
||||
const tempDoc = new DOMParser().parseFromString(snippet, "text/xml")
|
||||
const newDiagram = tempDoc.querySelector("diagram")
|
||||
if (!newDiagram) {
|
||||
throw new Error("Failed to construct <diagram> element for new page")
|
||||
|
||||
@@ -3,8 +3,6 @@
|
||||
* Copied from lib/utils.ts to avoid cross-package imports
|
||||
*/
|
||||
|
||||
import { getXmlSyntaxError } from "./dom.js"
|
||||
|
||||
// ============================================================================
|
||||
// Constants
|
||||
// ============================================================================
|
||||
@@ -12,6 +10,9 @@ import { getXmlSyntaxError } from "./dom.js"
|
||||
/** Maximum XML size to process (1MB) - larger XMLs may cause performance issues */
|
||||
const MAX_XML_SIZE = 1_000_000
|
||||
|
||||
/** Maximum iterations for aggressive cell dropping to prevent infinite loops */
|
||||
const MAX_DROP_ITERATIONS = 10
|
||||
|
||||
/** Structural attributes that should not be duplicated in draw.io */
|
||||
const STRUCTURAL_ATTRS = [
|
||||
"edge",
|
||||
@@ -90,21 +91,6 @@ function parseXmlTags(xml: string): ParsedTag[] {
|
||||
return tags
|
||||
}
|
||||
|
||||
/** Rewrite every opening tag with fn, leaving text and closing tags as is. */
|
||||
function replaceInOpeningTags(
|
||||
xml: string,
|
||||
fn: (tag: string) => string,
|
||||
): string {
|
||||
let out = ""
|
||||
let last = 0
|
||||
for (const { tag, isClosing, startIndex, endIndex } of parseXmlTags(xml)) {
|
||||
if (isClosing) continue
|
||||
out += xml.slice(last, startIndex) + fn(tag)
|
||||
last = endIndex + 1
|
||||
}
|
||||
return out + xml.slice(last)
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Validation Helper Functions
|
||||
// ============================================================================
|
||||
@@ -142,7 +128,8 @@ function checkDuplicateAttributes(xml: string): string | null {
|
||||
* scope the cell-ID uniqueness check per <diagram>, and additionally check
|
||||
* that the <diagram> ids themselves are unique.
|
||||
*
|
||||
* The legacy regex-based check is kept as a fallback for non-mxfile inputs.
|
||||
* The legacy regex-based check is kept as a fallback for non-mxfile inputs
|
||||
* and for XML that won't DOM-parse.
|
||||
*/
|
||||
function checkDuplicateIds(xml: string): string | null {
|
||||
// The DOM-aware path only matters for <mxfile> wrappers; for legacy
|
||||
@@ -155,47 +142,51 @@ function checkDuplicateIds(xml: string): string | null {
|
||||
if (mightBeMxFile)
|
||||
try {
|
||||
const doc = new DOMParser().parseFromString(xml, "text/xml")
|
||||
const rootEl = doc.documentElement
|
||||
if (rootEl && rootEl.tagName === "mxfile") {
|
||||
const diagrams = doc.querySelectorAll("diagram")
|
||||
if (!doc.querySelector("parsererror")) {
|
||||
const rootEl = doc.documentElement
|
||||
if (rootEl && rootEl.tagName === "mxfile") {
|
||||
const diagrams = doc.querySelectorAll("diagram")
|
||||
|
||||
// 1) <diagram> ids must be unique across the file.
|
||||
const diagramIds = new Map<string, number>()
|
||||
diagrams.forEach((d) => {
|
||||
const id = d.getAttribute("id")
|
||||
if (id) diagramIds.set(id, (diagramIds.get(id) || 0) + 1)
|
||||
})
|
||||
const dupDiagrams = Array.from(diagramIds.entries())
|
||||
.filter(([, c]) => c > 1)
|
||||
.map(([id]) => `'${id}'`)
|
||||
if (dupDiagrams.length > 0) {
|
||||
return `Invalid XML: Found duplicate <diagram> id(s): ${dupDiagrams.slice(0, 3).join(", ")}. Each page must have a unique id.`
|
||||
}
|
||||
|
||||
// 2) Within each page, mxCell ids must be unique.
|
||||
for (let i = 0; i < diagrams.length; i++) {
|
||||
const diagram = diagrams[i]
|
||||
const pageId = diagram.getAttribute("id") || `(index ${i})`
|
||||
const cells = diagram.querySelectorAll("mxCell")
|
||||
const cellIds = new Map<string, number>()
|
||||
cells.forEach((c) => {
|
||||
const id = c.getAttribute("id")
|
||||
if (id) cellIds.set(id, (cellIds.get(id) || 0) + 1)
|
||||
// 1) <diagram> ids must be unique across the file.
|
||||
const diagramIds = new Map<string, number>()
|
||||
diagrams.forEach((d) => {
|
||||
const id = d.getAttribute("id")
|
||||
if (id)
|
||||
diagramIds.set(id, (diagramIds.get(id) || 0) + 1)
|
||||
})
|
||||
const dups = Array.from(cellIds.entries())
|
||||
const dupDiagrams = Array.from(diagramIds.entries())
|
||||
.filter(([, c]) => c > 1)
|
||||
.map(([id, count]) => `'${id}' (${count}x)`)
|
||||
if (dups.length > 0) {
|
||||
return `Invalid XML: Found duplicate cell ID(s) in page "${pageId}": ${dups.slice(0, 3).join(", ")}. All mxCell ids must be unique within a page.`
|
||||
.map(([id]) => `'${id}'`)
|
||||
if (dupDiagrams.length > 0) {
|
||||
return `Invalid XML: Found duplicate <diagram> id(s): ${dupDiagrams.slice(0, 3).join(", ")}. Each page must have a unique id.`
|
||||
}
|
||||
|
||||
// 2) Within each page, mxCell ids must be unique.
|
||||
for (let i = 0; i < diagrams.length; i++) {
|
||||
const diagram = diagrams[i]
|
||||
const pageId =
|
||||
diagram.getAttribute("id") || `(index ${i})`
|
||||
const cells = diagram.querySelectorAll("mxCell")
|
||||
const cellIds = new Map<string, number>()
|
||||
cells.forEach((c) => {
|
||||
const id = c.getAttribute("id")
|
||||
if (id) cellIds.set(id, (cellIds.get(id) || 0) + 1)
|
||||
})
|
||||
const dups = Array.from(cellIds.entries())
|
||||
.filter(([, c]) => c > 1)
|
||||
.map(([id, count]) => `'${id}' (${count}x)`)
|
||||
if (dups.length > 0) {
|
||||
return `Invalid XML: Found duplicate cell ID(s) in page "${pageId}": ${dups.slice(0, 3).join(", ")}. All mxCell ids must be unique within a page.`
|
||||
}
|
||||
}
|
||||
return null
|
||||
}
|
||||
return null
|
||||
}
|
||||
} catch {
|
||||
// fall through to regex
|
||||
}
|
||||
|
||||
// Legacy regex-based check for bare <mxGraphModel> inputs.
|
||||
// Legacy regex-based check for bare <mxGraphModel> and parse-error cases.
|
||||
const idPattern = /\bid\s*=\s*["']([^"']+)["']/gi
|
||||
const ids = new Map<string, number>()
|
||||
let idMatch
|
||||
@@ -324,11 +315,14 @@ export function validateMxCellStructure(xml: string): string | null {
|
||||
)
|
||||
}
|
||||
|
||||
// 0. DOM-based checks. Syntax errors are caught by the strict check at
|
||||
// the end: linkedom's DOMParser never reports them.
|
||||
// 0. First use DOM parser to catch syntax errors (most accurate)
|
||||
try {
|
||||
const parser = new DOMParser()
|
||||
const doc = parser.parseFromString(xml, "text/xml")
|
||||
const parseError = doc.querySelector("parsererror")
|
||||
if (parseError) {
|
||||
return `Invalid XML: The XML contains syntax errors (likely unescaped special characters like <, >, & in attribute values). Please escape special characters: use < for <, > for >, & for &, " for ". Regenerate the diagram with properly escaped values.`
|
||||
}
|
||||
|
||||
// DOM-based checks for nested mxCell
|
||||
const allCells = doc.querySelectorAll("mxCell")
|
||||
@@ -410,14 +404,6 @@ export function validateMxCellStructure(xml: string): string | null {
|
||||
return nestedCellError
|
||||
}
|
||||
|
||||
// 11. Strict XML syntax check, run last so the checks above can give
|
||||
// more specific messages. Catches what they miss, e.g. duplicate or
|
||||
// unquoted attributes, which make draw.io refuse to load the diagram.
|
||||
const syntaxError = getXmlSyntaxError(xml)
|
||||
if (syntaxError) {
|
||||
return `Invalid XML: syntax error at ${syntaxError} Escape special characters in attribute values (< for <, & for &, " for "), quote every attribute value, and do not repeat an attribute.`
|
||||
}
|
||||
|
||||
return null
|
||||
}
|
||||
|
||||
@@ -508,21 +494,13 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
}
|
||||
}
|
||||
|
||||
// 6. Fix malformed attribute quotes (name="value"). Quoted
|
||||
// values are matched first and kept, so " inside a rich-text
|
||||
// label like value="<font style="...">" is left alone.
|
||||
let quotesFixed = false
|
||||
fixed = replaceInOpeningTags(fixed, (tag) =>
|
||||
tag.replace(
|
||||
/("[^"]*"|'[^']*')|(\s[a-zA-Z][a-zA-Z0-9_:-]*)="([^&]*?)"/g,
|
||||
(match, quoted, name, value) => {
|
||||
if (quoted) return match
|
||||
quotesFixed = true
|
||||
return `${name}="${value}"`
|
||||
},
|
||||
),
|
||||
)
|
||||
if (quotesFixed) {
|
||||
// 6. Fix malformed attribute quotes
|
||||
const malformedQuotePattern = /(\s[a-zA-Z][a-zA-Z0-9_:-]*)="/
|
||||
if (malformedQuotePattern.test(fixed)) {
|
||||
fixed = fixed.replace(
|
||||
/(\s[a-zA-Z][a-zA-Z0-9_:-]*)="([^&]*?)"/g,
|
||||
'$1="$2"',
|
||||
)
|
||||
fixes.push("Fixed malformed attribute quotes")
|
||||
}
|
||||
|
||||
@@ -533,21 +511,10 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
fixes.push("Fixed malformed closing tags")
|
||||
}
|
||||
|
||||
// 8. Fix missing space between attributes (id="2"vertex="1"). Every
|
||||
// quoted value is consumed whole, so quotes always pair up within one
|
||||
// attribute.
|
||||
let spaceAdded = false
|
||||
fixed = replaceInOpeningTags(fixed, (tag) =>
|
||||
tag.replace(
|
||||
/("[^"]*"|'[^']*')([a-zA-Z_:])?/g,
|
||||
(match, quoted, next) => {
|
||||
if (!next) return match
|
||||
spaceAdded = true
|
||||
return `${quoted} ${next}`
|
||||
},
|
||||
),
|
||||
)
|
||||
if (spaceAdded) {
|
||||
// 8. Fix missing space between attributes
|
||||
const missingSpacePattern = /("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g
|
||||
if (missingSpacePattern.test(fixed)) {
|
||||
fixed = fixed.replace(/("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g, "$1 $2")
|
||||
fixes.push("Added missing space between attributes")
|
||||
}
|
||||
|
||||
@@ -665,9 +632,6 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
"Array",
|
||||
"Object",
|
||||
"mxRectangle",
|
||||
// Wrappers draw.io writes for cells with links, tooltips or data
|
||||
"UserObject",
|
||||
"object",
|
||||
])
|
||||
const foreignTagPattern = /<\/?([a-zA-Z][a-zA-Z0-9_]*)[^>]*>/g
|
||||
let foreignMatch
|
||||
@@ -832,10 +796,8 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
fixes.push(`Flattened ${nestedFixed} duplicate-ID nested mxCell(s)`)
|
||||
}
|
||||
|
||||
// 21. Fix true nested mxCell (different IDs). Runs only when the nesting
|
||||
// check finds real nesting, because this line-based rewrite can break
|
||||
// valid cells written over several lines.
|
||||
const lines2 = checkNestedMxCells(fixed) ? fixed.split("\n") : []
|
||||
// 21. Fix true nested mxCell (different IDs)
|
||||
const lines2 = fixed.split("\n")
|
||||
newLines = []
|
||||
let trueNestedFixed = 0
|
||||
let cellDepth = 0
|
||||
@@ -845,11 +807,7 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
const line = lines2[i]
|
||||
const trimmed = line.trim()
|
||||
|
||||
// A line holding a whole cell (<mxCell ...>...</mxCell>) opens nothing
|
||||
const isOpenCell =
|
||||
/<mxCell\s/.test(trimmed) &&
|
||||
!trimmed.endsWith("/>") &&
|
||||
!trimmed.endsWith("</mxCell>")
|
||||
const isOpenCell = /<mxCell\s/.test(trimmed) && !trimmed.endsWith("/>")
|
||||
const isCloseCell = trimmed === "</mxCell>"
|
||||
|
||||
if (isOpenCell) {
|
||||
@@ -902,11 +860,9 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
|
||||
if (duplicateIds.length > 0) {
|
||||
const idCounters = new Map<string, number>()
|
||||
// Rebuild from the captured parts so only the value changes (an id
|
||||
// like "d" or "i" also occurs in the attribute name itself)
|
||||
fixed = fixed.replace(
|
||||
/(\bid\s*=\s*["'])([^"']+)(["'])/gi,
|
||||
(match, before, id, after) => {
|
||||
/\bid\s*=\s*["']([^"']+)["']/gi,
|
||||
(match, id) => {
|
||||
if (!duplicateIds.includes(id)) return match
|
||||
|
||||
const count = idCounters.get(id) || 0
|
||||
@@ -914,7 +870,8 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
|
||||
if (count === 0) return match
|
||||
|
||||
return `${before}${id}_dup${count}${after}`
|
||||
const newId = `${id}_dup${count}`
|
||||
return match.replace(id, newId)
|
||||
},
|
||||
)
|
||||
fixes.push(`Renamed ${duplicateIds.length} duplicate ID(s)`)
|
||||
@@ -935,6 +892,49 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
|
||||
fixes.push(`Generated ${emptyIdCount} missing ID(s)`)
|
||||
}
|
||||
|
||||
// 24. Aggressive: drop broken mxCell elements
|
||||
if (typeof DOMParser !== "undefined") {
|
||||
let droppedCells = 0
|
||||
let maxIterations = MAX_DROP_ITERATIONS
|
||||
while (maxIterations-- > 0) {
|
||||
const parser = new DOMParser()
|
||||
const doc = parser.parseFromString(fixed, "text/xml")
|
||||
const parseError = doc.querySelector("parsererror")
|
||||
if (!parseError) break
|
||||
|
||||
const errText = parseError.textContent || ""
|
||||
const match = errText.match(/(\d+):\d+:/)
|
||||
if (!match) break
|
||||
|
||||
const errLine = parseInt(match[1], 10) - 1
|
||||
const lines = fixed.split("\n")
|
||||
|
||||
let cellStart = errLine
|
||||
let cellEnd = errLine
|
||||
|
||||
while (cellStart > 0 && !lines[cellStart].includes("<mxCell")) {
|
||||
cellStart--
|
||||
}
|
||||
|
||||
while (cellEnd < lines.length - 1) {
|
||||
if (
|
||||
lines[cellEnd].includes("</mxCell>") ||
|
||||
lines[cellEnd].trim().endsWith("/>")
|
||||
) {
|
||||
break
|
||||
}
|
||||
cellEnd++
|
||||
}
|
||||
|
||||
lines.splice(cellStart, cellEnd - cellStart + 1)
|
||||
fixed = lines.join("\n")
|
||||
droppedCells++
|
||||
}
|
||||
if (droppedCells > 0) {
|
||||
fixes.push(`Dropped ${droppedCells} unfixable mxCell element(s)`)
|
||||
}
|
||||
}
|
||||
|
||||
return { fixed, fixes }
|
||||
}
|
||||
|
||||
|
||||
@@ -1,85 +0,0 @@
|
||||
/**
|
||||
* Tests for edit_diagram operations on cells that draw.io wraps in
|
||||
* <UserObject> or <object> (cells with links, tooltips or custom data).
|
||||
* The id sits on the wrapper; the inner mxCell has none.
|
||||
*/
|
||||
|
||||
import { beforeAll, describe, expect, it, vi } from "vitest"
|
||||
import { installDomPolyfill } from "../src/dom.js"
|
||||
|
||||
beforeAll(() => {
|
||||
installDomPolyfill()
|
||||
})
|
||||
|
||||
import { applyDiagramOperations } from "../src/diagram-operations.js"
|
||||
|
||||
const DOC = `<mxfile><diagram id="p" name="Page-1"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><UserObject id="a" label="A" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><mxCell id="b" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><object id="e1" label="" tooltip="t"><mxCell edge="1" source="b" target="a" parent="1"><mxGeometry relative="1" as="geometry"/></mxCell></object><mxCell id="child" value="C" vertex="1" parent="a"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel></diagram></mxfile>`
|
||||
|
||||
describe("wrapped cells", () => {
|
||||
it("deletes a UserObject cell with its edges and children", () => {
|
||||
const { result, errors } = applyDiagramOperations(DOC, [
|
||||
{ operation: "delete", cell_id: "a" },
|
||||
])
|
||||
expect(errors).toEqual([])
|
||||
expect(result).not.toContain('id="a"')
|
||||
expect(result).not.toContain('id="e1"')
|
||||
expect(result).not.toContain('id="child"')
|
||||
expect(result).toContain('id="b"')
|
||||
})
|
||||
|
||||
it("cascades to a wrapped edge when deleting a plain cell", () => {
|
||||
const { result, errors } = applyDiagramOperations(DOC, [
|
||||
{ operation: "delete", cell_id: "b" },
|
||||
{ operation: "delete", cell_id: "e1" },
|
||||
])
|
||||
// e1 was already removed by the cascade, so no warning for it
|
||||
expect(errors).toEqual([])
|
||||
expect(result).not.toContain('id="e1"')
|
||||
expect(result).toContain('id="a"')
|
||||
})
|
||||
|
||||
it("warns when deleting a cell that does not exist", () => {
|
||||
const { errors } = applyDiagramOperations(DOC, [
|
||||
{ operation: "delete", cell_id: "missing" },
|
||||
])
|
||||
expect(errors).toHaveLength(1)
|
||||
expect(errors[0]).toMatchObject({ type: "delete", cellId: "missing" })
|
||||
})
|
||||
|
||||
it("updates a UserObject cell", () => {
|
||||
const { result, errors } = applyDiagramOperations(DOC, [
|
||||
{
|
||||
operation: "update",
|
||||
cell_id: "a",
|
||||
new_xml: `<UserObject id="a" label="A2" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject>`,
|
||||
},
|
||||
])
|
||||
expect(errors).toEqual([])
|
||||
expect(result).toContain('label="A2"')
|
||||
expect(result.match(/id="a"/g)).toHaveLength(1)
|
||||
})
|
||||
|
||||
it("refuses to add a cell whose id a UserObject already uses", () => {
|
||||
const { errors } = applyDiagramOperations(DOC, [
|
||||
{
|
||||
operation: "add",
|
||||
cell_id: "a",
|
||||
new_xml: `<mxCell id="a" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>`,
|
||||
},
|
||||
])
|
||||
expect(errors[0]?.message).toContain("already exists")
|
||||
})
|
||||
})
|
||||
|
||||
describe("cascade delete logging", () => {
|
||||
it("does not write cascade logs to stdout (the JSON-RPC channel)", () => {
|
||||
const plain = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="x" vertex="1" parent="1"/><mxCell id="y" vertex="1" parent="1"/><mxCell id="e" edge="1" source="x" target="y" parent="1"/></root></mxGraphModel>`
|
||||
const spy = vi.spyOn(console, "log").mockImplementation(() => {})
|
||||
const { result } = applyDiagramOperations(plain, [
|
||||
{ operation: "delete", cell_id: "x" },
|
||||
])
|
||||
expect(result).not.toContain('id="e"')
|
||||
expect(spy).not.toHaveBeenCalled()
|
||||
spy.mockRestore()
|
||||
})
|
||||
})
|
||||
@@ -9,11 +9,11 @@
|
||||
* reads as a user edit.
|
||||
*/
|
||||
|
||||
import { DOMParser } from "linkedom"
|
||||
import { beforeAll, describe, expect, it } from "vitest"
|
||||
import { installDomPolyfill } from "../src/dom.js"
|
||||
|
||||
beforeAll(() => {
|
||||
installDomPolyfill()
|
||||
;(globalThis as any).DOMParser = DOMParser
|
||||
})
|
||||
|
||||
import { checkEditGate, contentFingerprint } from "../src/edit-gate.js"
|
||||
|
||||
@@ -1,213 +0,0 @@
|
||||
/**
|
||||
* Tests for the embedded HTTP server (browser bridge).
|
||||
*
|
||||
* The server runs in-process on a random high port (never 6002, which is
|
||||
* also the default port of the Next.js dev server). Requests go through
|
||||
* node:http so tests can set raw paths and Host/Origin headers.
|
||||
*/
|
||||
|
||||
import http from "node:http"
|
||||
import { afterAll, beforeAll, describe, expect, it } from "vitest"
|
||||
import { addHistory } from "../src/history.js"
|
||||
import {
|
||||
getState,
|
||||
setState,
|
||||
shutdown,
|
||||
startHttpServer,
|
||||
} from "../src/http-server.js"
|
||||
|
||||
let port = 0
|
||||
|
||||
beforeAll(async () => {
|
||||
port = await startHttpServer(40000 + Math.floor(Math.random() * 10000))
|
||||
})
|
||||
|
||||
afterAll(() => {
|
||||
shutdown()
|
||||
})
|
||||
|
||||
interface Response {
|
||||
status: number
|
||||
headers: http.IncomingHttpHeaders
|
||||
body: string
|
||||
}
|
||||
|
||||
/** Send a request; `body` may be split into several writes. */
|
||||
function request(
|
||||
path: string,
|
||||
opts: {
|
||||
method?: string
|
||||
headers?: Record<string, string>
|
||||
body?: Buffer[]
|
||||
} = {},
|
||||
): Promise<Response> {
|
||||
return new Promise((resolve, reject) => {
|
||||
const req = http.request(
|
||||
{
|
||||
host: "127.0.0.1",
|
||||
port,
|
||||
path,
|
||||
method: opts.method ?? "GET",
|
||||
headers: { host: `localhost:${port}`, ...opts.headers },
|
||||
},
|
||||
(res) => {
|
||||
const chunks: Buffer[] = []
|
||||
res.on("data", (c: Buffer) => chunks.push(c))
|
||||
res.on("end", () =>
|
||||
resolve({
|
||||
status: res.statusCode ?? 0,
|
||||
headers: res.headers,
|
||||
body: Buffer.concat(chunks).toString("utf8"),
|
||||
}),
|
||||
)
|
||||
},
|
||||
)
|
||||
req.on("error", reject)
|
||||
const parts = opts.body ?? []
|
||||
// Pause between parts so the server reads them as separate chunks
|
||||
const writeNext = (i: number) => {
|
||||
if (i >= parts.length) return req.end()
|
||||
req.write(parts[i])
|
||||
setTimeout(() => writeNext(i + 1), 30)
|
||||
}
|
||||
writeNext(0)
|
||||
})
|
||||
}
|
||||
|
||||
const postJson = (path: string, data: unknown, headers = {}) =>
|
||||
request(path, {
|
||||
method: "POST",
|
||||
headers: { "content-type": "application/json", ...headers },
|
||||
body: [Buffer.from(JSON.stringify(data))],
|
||||
})
|
||||
|
||||
describe("session id in the page URL", () => {
|
||||
it("rejects a session id that could inject script", async () => {
|
||||
const res = await request(`/?mcp=${encodeURIComponent('";alert(1)//')}`)
|
||||
expect(res.status).toBe(400)
|
||||
expect(res.body).not.toContain("alert")
|
||||
})
|
||||
|
||||
it("writes a valid session id into the page script as a JSON string", async () => {
|
||||
const res = await request("/?mcp=mcp-test-page")
|
||||
expect(res.status).toBe(200)
|
||||
expect(res.body).toContain('const sessionId = "mcp-test-page";')
|
||||
})
|
||||
})
|
||||
|
||||
describe("requests that used to crash the process", () => {
|
||||
it("answers 400 for a path that is not a valid URL", async () => {
|
||||
const res = await request("//")
|
||||
expect(res.status).toBe(400)
|
||||
// The server is still alive
|
||||
expect((await request("/api/state?sessionId=mcp-alive")).status).toBe(
|
||||
200,
|
||||
)
|
||||
})
|
||||
|
||||
it("never creates sessions with ids unsafe for the Location header", async () => {
|
||||
const badId = "mcp-中"
|
||||
await request(`/api/state?sessionId=${encodeURIComponent(badId)}`)
|
||||
expect(getState(badId)).toBeUndefined()
|
||||
const post = await postJson("/api/state", {
|
||||
sessionId: badId,
|
||||
xml: "<mxfile/>",
|
||||
})
|
||||
expect(post.status).toBe(400)
|
||||
expect(getState(badId)).toBeUndefined()
|
||||
|
||||
const res = await request("/")
|
||||
expect([200, 302]).toContain(res.status)
|
||||
})
|
||||
})
|
||||
|
||||
describe("request origin checks", () => {
|
||||
it("refuses a foreign Host header (DNS rebinding)", async () => {
|
||||
const res = await request("/api/state?sessionId=mcp-alive", {
|
||||
headers: { host: `evil.example:${port}` },
|
||||
})
|
||||
expect(res.status).toBe(403)
|
||||
})
|
||||
|
||||
it("refuses writes from another website", async () => {
|
||||
const res = await postJson(
|
||||
"/api/state",
|
||||
{ sessionId: "mcp-csrf", xml: "<mxfile/>" },
|
||||
{ origin: "https://evil.example" },
|
||||
)
|
||||
expect(res.status).toBe(403)
|
||||
expect(getState("mcp-csrf")).toBeUndefined()
|
||||
})
|
||||
|
||||
it("accepts writes from the page itself", async () => {
|
||||
const res = await postJson(
|
||||
"/api/state",
|
||||
{ sessionId: "mcp-same-origin", xml: "<mxfile/>" },
|
||||
{ origin: `http://localhost:${port}` },
|
||||
)
|
||||
expect(res.status).toBe(200)
|
||||
})
|
||||
})
|
||||
|
||||
describe("POST /api/state", () => {
|
||||
it("decodes UTF-8 characters split across body chunks", async () => {
|
||||
const xml = `<mxfile>${"数据".repeat(30000)}</mxfile>`
|
||||
const body = Buffer.from(JSON.stringify({ sessionId: "mcp-utf8", xml }))
|
||||
// Cut inside a 3-byte character
|
||||
const cut = body.indexOf(Buffer.from("数")) + 1
|
||||
const res = await request("/api/state", {
|
||||
method: "POST",
|
||||
headers: { "content-type": "application/json" },
|
||||
body: [body.subarray(0, cut), body.subarray(cut)],
|
||||
})
|
||||
expect(res.status).toBe(200)
|
||||
expect(getState("mcp-utf8")?.xml).toBe(xml)
|
||||
})
|
||||
|
||||
it("rejects a browser push based on a version older than an AI write", async () => {
|
||||
const id = "mcp-conflict"
|
||||
setState(id, "<mxfile>user v1</mxfile>", undefined, true)
|
||||
const aiVersion = setState(id, "<mxfile>AI edit</mxfile>")
|
||||
|
||||
const stale = await postJson("/api/state", {
|
||||
sessionId: id,
|
||||
xml: "<mxfile>user edit on old version</mxfile>",
|
||||
baseVersion: aiVersion - 1,
|
||||
})
|
||||
expect(stale.status).toBe(409)
|
||||
expect(getState(id)?.xml).toBe("<mxfile>AI edit</mxfile>")
|
||||
|
||||
// Pushes based on the AI version are accepted, including a second
|
||||
// push sent before the first one's response updated the browser
|
||||
for (const xml of ["<mxfile>a</mxfile>", "<mxfile>b</mxfile>"]) {
|
||||
const ok = await postJson("/api/state", {
|
||||
sessionId: id,
|
||||
xml,
|
||||
baseVersion: aiVersion,
|
||||
})
|
||||
expect(ok.status).toBe(200)
|
||||
expect(getState(id)?.xml).toBe(xml)
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
describe("history restore", () => {
|
||||
it("restores the entry the user picked after older entries drop", async () => {
|
||||
const id = "mcp-history"
|
||||
setState(id, "<mxfile/>")
|
||||
for (let i = 0; i < 20; i++) addHistory(id, `<mxfile>${i}</mxfile>`)
|
||||
|
||||
const list = await request(`/api/history?sessionId=${id}`)
|
||||
const picked = JSON.parse(list.body).entries[5]
|
||||
|
||||
// A new AI edit shifts the buffer before the user clicks Restore
|
||||
addHistory(id, "<mxfile>new</mxfile>")
|
||||
|
||||
const res = await postJson("/api/restore", {
|
||||
sessionId: id,
|
||||
id: picked.id,
|
||||
})
|
||||
expect(res.status).toBe(200)
|
||||
expect(getState(id)?.xml).toBe("<mxfile>5</mxfile>")
|
||||
})
|
||||
})
|
||||
@@ -10,11 +10,18 @@
|
||||
import { deflateRawSync } from "node:zlib"
|
||||
import { DOMParser } from "linkedom"
|
||||
import { beforeAll, describe, expect, it } from "vitest"
|
||||
import { installDomPolyfill } from "../src/dom.js"
|
||||
|
||||
// Install the DOM polyfills exactly as index.ts does at runtime.
|
||||
beforeAll(() => {
|
||||
installDomPolyfill()
|
||||
;(globalThis as any).DOMParser = DOMParser
|
||||
class XMLSerializerPolyfill {
|
||||
serializeToString(node: any): string {
|
||||
if (node.outerHTML !== undefined) return node.outerHTML
|
||||
if (node.documentElement) return node.documentElement.outerHTML
|
||||
return ""
|
||||
}
|
||||
}
|
||||
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
|
||||
})
|
||||
|
||||
import {
|
||||
|
||||
@@ -15,13 +15,21 @@
|
||||
* (diagram-operations.ts) — i.e. the layers underneath the MCP tool surface.
|
||||
*/
|
||||
|
||||
import { DOMParser } from "linkedom"
|
||||
import { beforeAll, describe, expect, it } from "vitest"
|
||||
import { installDomPolyfill } from "../src/dom.js"
|
||||
|
||||
// Install the DOM polyfill exactly as index.ts does at runtime — the
|
||||
// helpers under test rely on it.
|
||||
beforeAll(() => {
|
||||
installDomPolyfill()
|
||||
;(globalThis as any).DOMParser = DOMParser
|
||||
class XMLSerializerPolyfill {
|
||||
serializeToString(node: any): string {
|
||||
if (node.outerHTML !== undefined) return node.outerHTML
|
||||
if (node.documentElement) return node.documentElement.outerHTML
|
||||
return ""
|
||||
}
|
||||
}
|
||||
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
|
||||
})
|
||||
|
||||
import { applyDiagramOperations } from "../src/diagram-operations.js"
|
||||
|
||||
@@ -1,186 +0,0 @@
|
||||
/**
|
||||
* Tests for XML syntax checking, autoFixXml and the XML serializer.
|
||||
*
|
||||
* linkedom (the DOM used in Node) parses leniently and never reports syntax
|
||||
* errors, so validation relies on the strict check in dom.ts. autoFixXml
|
||||
* runs on the whole document whenever any check fails, so its steps must
|
||||
* leave valid parts of the document untouched.
|
||||
*/
|
||||
|
||||
import { beforeAll, describe, expect, it } from "vitest"
|
||||
import { getXmlSyntaxError, installDomPolyfill } from "../src/dom.js"
|
||||
|
||||
beforeAll(() => {
|
||||
installDomPolyfill()
|
||||
})
|
||||
|
||||
import { addPageToDoc, parseMxfile, serializeMxfile } from "../src/pages.js"
|
||||
import { validateAndFixXml } from "../src/xml-validation.js"
|
||||
|
||||
/** Bare model with the root cells plus the given cells. */
|
||||
const model = (cells: string) =>
|
||||
`<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/>${cells}</root></mxGraphModel>`
|
||||
|
||||
// A bare & makes the first validation fail, which triggers autoFixXml on
|
||||
// the whole document.
|
||||
const BROKEN_CELL = `<mxCell id="9" value="R&D" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>`
|
||||
|
||||
describe("getXmlSyntaxError", () => {
|
||||
it("accepts well-formed XML", () => {
|
||||
expect(getXmlSyntaxError(model(""))).toBeNull()
|
||||
})
|
||||
|
||||
it.each([
|
||||
["duplicate attribute", `<a style="x" style="y"/>`],
|
||||
["unquoted attribute", `<a id=2/>`],
|
||||
["missing space between attributes", `<a id="2"vertex="1"/>`],
|
||||
["bare ampersand", `<a v="R&D"/>`],
|
||||
["unclosed tag", `<a><b></a>`],
|
||||
["plain text", `hello`],
|
||||
])("reports %s", (_name, xml) => {
|
||||
expect(getXmlSyntaxError(xml)).toMatch(/^\d+:\d+: /)
|
||||
})
|
||||
})
|
||||
|
||||
describe("validateAndFixXml", () => {
|
||||
it("rejects a duplicate style attribute", () => {
|
||||
const r = validateAndFixXml(
|
||||
model(
|
||||
`<mxCell id="2" style="a=1;" style="b=1;" vertex="1" parent="1"/>`,
|
||||
),
|
||||
)
|
||||
expect(r.valid).toBe(false)
|
||||
expect(r.error).toContain("duplicate attribute: style")
|
||||
})
|
||||
|
||||
it("rejects an unquoted attribute value", () => {
|
||||
const r = validateAndFixXml(
|
||||
model(`<mxCell id=2 vertex="1" parent="1"/>`),
|
||||
)
|
||||
expect(r.valid).toBe(false)
|
||||
})
|
||||
|
||||
it("keeps style values intact while fixing another cell", () => {
|
||||
const cells = `<mxCell id="2" style="shape=cylinder3;whiteSpace=wrap;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="3" style="edgeStyle=orthogonalEdgeStyle;" edge="1" parent="1" source="2" target="2"><mxGeometry relative="1" as="geometry"/></mxCell>`
|
||||
const r = validateAndFixXml(model(cells + BROKEN_CELL))
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixed).toContain('style="shape=cylinder3;whiteSpace=wrap;"')
|
||||
expect(r.fixed).toContain('style="edgeStyle=orthogonalEdgeStyle;"')
|
||||
expect(r.fixed).toContain('value="R&D"')
|
||||
})
|
||||
|
||||
it("keeps " inside rich-text labels", () => {
|
||||
const rich = `<mxCell id="4" value="<font style="color: red;">Hi</font>" style="html=1;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>`
|
||||
const r = validateAndFixXml(model(rich + BROKEN_CELL))
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixed).toContain(
|
||||
'value="<font style="color: red;">Hi</font>"',
|
||||
)
|
||||
expect(getXmlSyntaxError(r.fixed ?? "")).toBeNull()
|
||||
})
|
||||
|
||||
it("keeps UserObject and object wrappers", () => {
|
||||
const wrapped = `<UserObject id="u" label="L" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><object id="o" label="O"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></object>`
|
||||
const r = validateAndFixXml(model(wrapped + BROKEN_CELL))
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixed).toContain('<UserObject id="u"')
|
||||
expect(r.fixed).toContain('<object id="o"')
|
||||
})
|
||||
|
||||
it("leaves one-cell-per-line XML alone while fixing another cell", () => {
|
||||
const xml = [
|
||||
"<mxGraphModel>",
|
||||
"<root>",
|
||||
'<mxCell id="0"/>',
|
||||
'<mxCell id="1" parent="0"/>',
|
||||
BROKEN_CELL,
|
||||
'<mxCell id="3" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>',
|
||||
"</root>",
|
||||
"</mxGraphModel>",
|
||||
].join("\n")
|
||||
const r = validateAndFixXml(xml)
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixes).toEqual(["Escaped unescaped & characters"])
|
||||
})
|
||||
|
||||
it("leaves cells split over two lines alone while fixing another cell", () => {
|
||||
const xml = [
|
||||
"<mxGraphModel><root>",
|
||||
'<mxCell id="0"/><mxCell id="1" parent="0"/>',
|
||||
'<mxCell id="2" value="A" vertex="1" parent="1">',
|
||||
' <mxGeometry as="geometry"/></mxCell>',
|
||||
'<mxCell id="3" value="B" vertex="1" parent="1">',
|
||||
' <mxGeometry as="geometry"/></mxCell>',
|
||||
BROKEN_CELL,
|
||||
"</root></mxGraphModel>",
|
||||
].join("\n")
|
||||
const r = validateAndFixXml(xml)
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixes).toEqual(["Escaped unescaped & characters"])
|
||||
})
|
||||
|
||||
it("renames duplicate short ids without touching the attribute name", () => {
|
||||
const r = validateAndFixXml(
|
||||
model(
|
||||
`<mxCell id="d" vertex="1" parent="1"/><mxCell id="d" vertex="1" parent="1"/><mxCell id="i" vertex="1" parent="1"/><mxCell id="i" vertex="1" parent="1"/>`,
|
||||
),
|
||||
)
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixed).toContain('id="d_dup1"')
|
||||
expect(r.fixed).toContain('id="i_dup1"')
|
||||
})
|
||||
|
||||
it("adds a missing space between attributes", () => {
|
||||
const r = validateAndFixXml(
|
||||
model(
|
||||
`<mxCell id="2"value="a" style="x=1;" vertex="1" parent="1"/>`,
|
||||
),
|
||||
)
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixed).toContain('<mxCell id="2" value="a" style="x=1;"')
|
||||
})
|
||||
|
||||
it("fixes attribute values quoted with "", () => {
|
||||
const r = validateAndFixXml(
|
||||
model(
|
||||
`<mxCell id="2" value="Hello" vertex="1" parent="1"/>`,
|
||||
),
|
||||
)
|
||||
expect(r.valid).toBe(true)
|
||||
expect(r.fixed).toContain('value="Hello"')
|
||||
})
|
||||
})
|
||||
|
||||
describe("XML serializer and strict parsing in page helpers", () => {
|
||||
it("keeps line breaks and tabs in attribute values", () => {
|
||||
const xml = `<mxfile><diagram id="p" name="Page-1"><mxGraphModel><root><mxCell id="0"/><mxCell id="2" value="Multi-Head
Attention	x" vertex="1" parent="0"/></root></mxGraphModel></diagram></mxfile>`
|
||||
const out = serializeMxfile(parseMxfile(xml) as Document)
|
||||
expect(out).toContain('value="Multi-Head
Attention	x"')
|
||||
expect(out).not.toMatch(/value="[^"]*\n/)
|
||||
})
|
||||
|
||||
it("escapes special characters in attributes and text", () => {
|
||||
const xml = `<mxfile><diagram id="p" name="R&D">a < b<mxGraphModel><root><mxCell id="0" value="<b> & ""/></root></mxGraphModel></diagram></mxfile>`
|
||||
const out = serializeMxfile(parseMxfile(xml) as Document)
|
||||
expect(out).toBe(xml)
|
||||
})
|
||||
|
||||
it("parseMxfile returns null for malformed XML", () => {
|
||||
expect(
|
||||
parseMxfile(
|
||||
`<mxfile><diagram id="p" name="a" name="b"></diagram></mxfile>`,
|
||||
),
|
||||
).toBeNull()
|
||||
})
|
||||
|
||||
it("addPageToDoc rejects malformed page XML", () => {
|
||||
const doc = parseMxfile(
|
||||
`<mxfile><diagram id="p" name="Page-1">${model("")}</diagram></mxfile>`,
|
||||
) as Document
|
||||
expect(() =>
|
||||
addPageToDoc(doc, {
|
||||
xml: model(`<mxCell id=2 vertex="1" parent="1"/>`),
|
||||
}),
|
||||
).toThrow()
|
||||
})
|
||||
})
|
||||
@@ -48,12 +48,13 @@ export function proxy(request: NextRequest) {
|
||||
if (pathnameIsMissingLocale) {
|
||||
const locale = getLocale(request)
|
||||
|
||||
// Redirect to localized path. Cloning nextUrl keeps the basePath
|
||||
// (NEXT_PUBLIC_BASE_PATH) and query string, which
|
||||
// new URL("/...", request.url) would drop.
|
||||
const url = request.nextUrl.clone()
|
||||
url.pathname = `/${locale}${pathname}`
|
||||
return NextResponse.redirect(url)
|
||||
// Redirect to localized path
|
||||
return NextResponse.redirect(
|
||||
new URL(
|
||||
`/${locale}${pathname.startsWith("/") ? "" : "/"}${pathname}`,
|
||||
request.url,
|
||||
),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+67
-88
@@ -2,7 +2,7 @@
|
||||
|
||||
/**
|
||||
* Development script for running Electron with Next.js
|
||||
* 1. Reads the active preset's env vars (if any)
|
||||
* 1. Reads preset configuration (if exists)
|
||||
* 2. Starts Next.js dev server with preset env vars
|
||||
* 3. Waits for it to be ready
|
||||
* 4. Compiles Electron TypeScript
|
||||
@@ -47,39 +47,37 @@ function getUserDataPath() {
|
||||
}
|
||||
|
||||
/**
|
||||
* File where the Electron main process (in development) writes the active
|
||||
* preset's env vars, already decrypted and mapped to provider-specific keys
|
||||
* (see writeDevPresetEnv in electron/main/config-manager.ts)
|
||||
* Load preset configuration from config file
|
||||
*/
|
||||
const PRESET_ENV_FILE = "dev-preset-env.json"
|
||||
function loadPresetConfig() {
|
||||
const configPath = path.join(getUserDataPath(), "config-presets.json")
|
||||
|
||||
if (!existsSync(configPath)) {
|
||||
console.log("📋 No preset configuration found, using .env.local")
|
||||
return null
|
||||
}
|
||||
|
||||
/**
|
||||
* Read the active preset's env vars as JSON text (null if not available)
|
||||
*/
|
||||
function readPresetEnvFile() {
|
||||
try {
|
||||
const content = readFileSync(
|
||||
path.join(getUserDataPath(), PRESET_ENV_FILE),
|
||||
"utf-8",
|
||||
)
|
||||
JSON.parse(content) // Ignore a half-written file
|
||||
return content
|
||||
} catch {
|
||||
return null
|
||||
}
|
||||
}
|
||||
const content = readFileSync(configPath, "utf-8")
|
||||
const data = JSON.parse(content)
|
||||
|
||||
/**
|
||||
* Load the active preset's env vars
|
||||
*/
|
||||
function loadPresetEnv(content) {
|
||||
const env = content ? JSON.parse(content) : {}
|
||||
if (Object.keys(env).length === 0) {
|
||||
console.log("📋 No active preset, using .env.local")
|
||||
if (!data.currentPresetId) {
|
||||
console.log("📋 No active preset, using .env.local")
|
||||
return null
|
||||
}
|
||||
|
||||
const preset = data.presets.find((p) => p.id === data.currentPresetId)
|
||||
if (!preset) {
|
||||
console.log("📋 Active preset not found, using .env.local")
|
||||
return null
|
||||
}
|
||||
|
||||
console.log(`📋 Using preset: "${preset.name}"`)
|
||||
return preset.config
|
||||
} catch (error) {
|
||||
console.error("Failed to load preset config:", error.message)
|
||||
return null
|
||||
}
|
||||
console.log(`📋 Using preset env: ${Object.keys(env).join(", ")}`)
|
||||
return env
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -130,18 +128,6 @@ function runCommand(command, args, options = {}) {
|
||||
})
|
||||
}
|
||||
|
||||
/**
|
||||
* Kill a process started with shell: true. On Windows, kill() only ends the
|
||||
* cmd.exe wrapper and leaves next dev running, so kill the whole tree.
|
||||
*/
|
||||
function killProcess(proc) {
|
||||
if (process.platform === "win32" && proc.pid) {
|
||||
spawn("taskkill", ["/pid", String(proc.pid), "/T", "/F"])
|
||||
} else {
|
||||
proc.kill()
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Start Next.js dev server with preset environment
|
||||
*/
|
||||
@@ -178,8 +164,7 @@ async function main() {
|
||||
console.log("🚀 Starting Electron development environment...\n")
|
||||
|
||||
// Load preset configuration
|
||||
let presetEnvContent = readPresetEnvFile()
|
||||
const presetEnv = loadPresetEnv(presetEnvContent)
|
||||
const presetEnv = loadPresetConfig()
|
||||
|
||||
// Start Next.js dev server with preset env
|
||||
console.log("1. Starting Next.js development server...")
|
||||
@@ -191,7 +176,7 @@ async function main() {
|
||||
console.log("")
|
||||
} catch (err) {
|
||||
console.error("\n❌ Next.js server failed to start:", err.message)
|
||||
killProcess(nextProcess)
|
||||
nextProcess.kill()
|
||||
process.exit(1)
|
||||
}
|
||||
|
||||
@@ -201,7 +186,7 @@ async function main() {
|
||||
await runCommand("npm", ["run", "electron:compile"])
|
||||
} catch (err) {
|
||||
console.error("❌ Electron compilation failed:", err.message)
|
||||
killProcess(nextProcess)
|
||||
nextProcess.kill()
|
||||
process.exit(1)
|
||||
}
|
||||
|
||||
@@ -218,82 +203,76 @@ async function main() {
|
||||
},
|
||||
})
|
||||
|
||||
// Watch for preset env changes
|
||||
const userDataPath = getUserDataPath()
|
||||
// Watch for preset config changes
|
||||
const configPath = path.join(getUserDataPath(), "config-presets.json")
|
||||
let configWatcher = null
|
||||
let restartPending = false
|
||||
|
||||
function setupConfigWatcher() {
|
||||
if (!existsSync(userDataPath)) {
|
||||
if (!existsSync(path.dirname(configPath))) {
|
||||
// Directory doesn't exist yet, check again later
|
||||
setTimeout(setupConfigWatcher, 5000)
|
||||
return
|
||||
}
|
||||
|
||||
try {
|
||||
// Watch the directory, since the file may not exist yet
|
||||
configWatcher = watch(
|
||||
userDataPath,
|
||||
configPath,
|
||||
{ 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)
|
||||
async (eventType) => {
|
||||
if (eventType === "change" && !restartPending) {
|
||||
restartPending = true
|
||||
console.log(
|
||||
"✅ Next.js server restarted with new configuration\n",
|
||||
"\n🔄 Preset configuration changed, restarting Next.js server...",
|
||||
)
|
||||
} catch (err) {
|
||||
console.error(
|
||||
"❌ Failed to restart Next.js:",
|
||||
err.message,
|
||||
)
|
||||
}
|
||||
|
||||
restartPending = false
|
||||
// Kill current Next.js process
|
||||
nextProcess.kill()
|
||||
|
||||
// Wait a bit for process to die
|
||||
await new Promise((r) => setTimeout(r, 1000))
|
||||
|
||||
// Reload preset and restart
|
||||
const newPresetEnv = loadPresetConfig()
|
||||
nextProcess = startNextServer(newPresetEnv)
|
||||
|
||||
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
|
||||
}
|
||||
},
|
||||
)
|
||||
console.log("👀 Watching for preset configuration changes...")
|
||||
} catch (_err) {
|
||||
// Directory might not be ready yet, try again later
|
||||
// File might not exist yet, that's ok
|
||||
setTimeout(setupConfigWatcher, 5000)
|
||||
}
|
||||
}
|
||||
|
||||
// Start watching after a delay (user data directory might not exist yet)
|
||||
// Start watching after a delay (config file might not exist yet)
|
||||
setTimeout(setupConfigWatcher, 2000)
|
||||
|
||||
electronProcess.on("close", (code) => {
|
||||
console.log(`\nElectron exited with code ${code}`)
|
||||
if (configWatcher) configWatcher.close()
|
||||
killProcess(nextProcess)
|
||||
nextProcess.kill()
|
||||
process.exit(code || 0)
|
||||
})
|
||||
|
||||
electronProcess.on("error", (err) => {
|
||||
console.error("Electron error:", err)
|
||||
if (configWatcher) configWatcher.close()
|
||||
killProcess(nextProcess)
|
||||
nextProcess.kill()
|
||||
process.exit(1)
|
||||
})
|
||||
|
||||
@@ -301,8 +280,8 @@ async function main() {
|
||||
const cleanup = () => {
|
||||
console.log("\n🛑 Shutting down...")
|
||||
if (configWatcher) configWatcher.close()
|
||||
killProcess(electronProcess)
|
||||
killProcess(nextProcess)
|
||||
electronProcess.kill()
|
||||
nextProcess.kill()
|
||||
process.exit(0)
|
||||
}
|
||||
|
||||
|
||||
@@ -73,15 +73,6 @@ mkdirSync(targetDir, { recursive: true })
|
||||
console.log("Copying standalone directory...")
|
||||
copyDereferenced(standaloneDir, targetDir)
|
||||
|
||||
// Next.js copies the build machine's .env files into standalone; don't ship
|
||||
// them, they can hold the builder's API keys
|
||||
for (const entry of readdirSync(targetDir)) {
|
||||
if (entry.startsWith(".env")) {
|
||||
console.warn(`Removing ${entry} so it is not packaged into the app`)
|
||||
rmSync(join(targetDir, entry))
|
||||
}
|
||||
}
|
||||
|
||||
// Copy static files
|
||||
console.log("Copying static files...")
|
||||
const targetStaticDir = join(targetDir, ".next", "static")
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
// @vitest-environment node
|
||||
import fs from "fs"
|
||||
import os from "os"
|
||||
import path from "path"
|
||||
import { afterEach, beforeEach, describe, expect, it } from "vitest"
|
||||
import { GET, PUT } from "@/app/api/admin/providers/route"
|
||||
import { _resetForTests } from "@/lib/admin/settings"
|
||||
|
||||
let tmpDir: string
|
||||
|
||||
beforeEach(() => {
|
||||
tmpDir = fs.mkdtempSync(path.join(os.tmpdir(), "admin-providers-route-"))
|
||||
process.env.SETTINGS_FILE = path.join(tmpDir, "settings.json")
|
||||
process.env.ADMIN_PASSWORD = "pw"
|
||||
process.env.AI_MODELS_CONFIG_PATH = path.join(tmpDir, "none.json")
|
||||
_resetForTests()
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
_resetForTests()
|
||||
delete process.env.SETTINGS_FILE
|
||||
delete process.env.ADMIN_PASSWORD
|
||||
delete process.env.AI_MODELS_CONFIG_PATH
|
||||
delete process.env.AI_MODEL
|
||||
fs.rmSync(tmpDir, { recursive: true, force: true })
|
||||
})
|
||||
|
||||
const headers = { "x-admin-password": "pw" }
|
||||
|
||||
async function saveDefaultPanelProvider() {
|
||||
const res = await PUT(
|
||||
new Request("http://localhost/api/admin/providers", {
|
||||
method: "PUT",
|
||||
headers: { ...headers, "Content-Type": "application/json" },
|
||||
body: JSON.stringify({
|
||||
providers: [
|
||||
{
|
||||
id: "p1",
|
||||
provider: "openai",
|
||||
apiKey: "sk-test",
|
||||
models: ["gpt-panel"],
|
||||
isDefault: true,
|
||||
},
|
||||
],
|
||||
}),
|
||||
}),
|
||||
)
|
||||
expect(res.status).toBe(200)
|
||||
}
|
||||
|
||||
async function envHasDefaultModel(): Promise<boolean> {
|
||||
const res = await GET(
|
||||
new Request("http://localhost/api/admin/providers", { headers }),
|
||||
)
|
||||
return (await res.json()).envHasDefaultModel
|
||||
}
|
||||
|
||||
describe("envHasDefaultModel", () => {
|
||||
it("is true when .env sets AI_MODEL, even after a panel default", async () => {
|
||||
process.env.AI_MODEL = "gpt-env"
|
||||
expect(await envHasDefaultModel()).toBe(true)
|
||||
await saveDefaultPanelProvider()
|
||||
expect(await envHasDefaultModel()).toBe(true)
|
||||
})
|
||||
|
||||
it("ignores the AI_MODEL the panel default writes", async () => {
|
||||
await saveDefaultPanelProvider()
|
||||
expect(process.env.AI_MODEL).toBe("gpt-panel")
|
||||
expect(await envHasDefaultModel()).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -69,7 +69,7 @@ describe("deriveEnvUpdates", () => {
|
||||
expect(updates.ADMIN_OPENAI_API_KEY_2).toBe("sk-second")
|
||||
})
|
||||
|
||||
it("maps bedrock credentials to ADMIN_AWS_* env vars", () => {
|
||||
it("maps bedrock credentials to AWS env vars", () => {
|
||||
const updates = deriveEnvUpdates(
|
||||
[
|
||||
provider({
|
||||
@@ -83,26 +83,9 @@ describe("deriveEnvUpdates", () => {
|
||||
],
|
||||
[],
|
||||
)
|
||||
expect(updates.ADMIN_AWS_ACCESS_KEY_ID).toBe("AKIA123")
|
||||
expect(updates.ADMIN_AWS_SECRET_ACCESS_KEY).toBe("secret")
|
||||
expect(updates.ADMIN_AWS_REGION).toBe("us-west-2")
|
||||
// Standard AWS vars are left to the environment
|
||||
expect(updates.AWS_ACCESS_KEY_ID).toBeUndefined()
|
||||
})
|
||||
|
||||
it("clears AWS_* bedrock keys written by older versions", () => {
|
||||
const bedrock = provider({
|
||||
provider: "bedrock",
|
||||
apiKey: undefined,
|
||||
awsAccessKeyId: "AKIA123",
|
||||
awsSecretAccessKey: "secret",
|
||||
models: ["claude-x"],
|
||||
})
|
||||
const updates = deriveEnvUpdates([bedrock], [bedrock])
|
||||
expect(updates.AWS_ACCESS_KEY_ID).toBeNull()
|
||||
expect(updates.AWS_SECRET_ACCESS_KEY).toBeNull()
|
||||
expect(updates.AWS_REGION).toBeNull()
|
||||
expect(updates.ADMIN_AWS_ACCESS_KEY_ID).toBe("AKIA123")
|
||||
expect(updates.AWS_ACCESS_KEY_ID).toBe("AKIA123")
|
||||
expect(updates.AWS_SECRET_ACCESS_KEY).toBe("secret")
|
||||
expect(updates.AWS_REGION).toBe("us-west-2")
|
||||
})
|
||||
|
||||
it("clears keys owned by the previous list when providers are removed", () => {
|
||||
@@ -332,32 +315,6 @@ describe("validateAdminProviders", () => {
|
||||
expect(validateAdminProviders(list)).toMatch(/unique/)
|
||||
})
|
||||
|
||||
it("rejects names that differ only in case or punctuation", () => {
|
||||
const list = [
|
||||
provider({ id: "p1", name: "Open AI" }),
|
||||
provider({ id: "p2", name: "open-ai" }),
|
||||
]
|
||||
expect(validateAdminProviders(list)).toMatch(/unique/)
|
||||
})
|
||||
|
||||
it("rejects a case-only clash with an env-configured name", () => {
|
||||
expect(
|
||||
validateAdminProviders([provider({ name: "openai" })], {
|
||||
providers: [
|
||||
{ name: "OpenAI", provider: "openai", models: ["gpt-x"] },
|
||||
],
|
||||
}),
|
||||
).toMatch(/already defined/)
|
||||
})
|
||||
|
||||
it("accepts distinct CJK names", () => {
|
||||
const list = [
|
||||
provider({ id: "p1", provider: "deepseek", name: "主力" }),
|
||||
provider({ id: "p2", provider: "deepseek", name: "备用" }),
|
||||
]
|
||||
expect(validateAdminProviders(list)).toBeNull()
|
||||
})
|
||||
|
||||
it("rejects multiple defaults", () => {
|
||||
const list = [
|
||||
provider({ id: "p1", isDefault: true }),
|
||||
|
||||
@@ -1,87 +0,0 @@
|
||||
import { cleanup, fireEvent, render, screen } from "@testing-library/react"
|
||||
import type { ReactNode } from "react"
|
||||
import { afterEach, describe, expect, it, vi } from "vitest"
|
||||
import { SecretInput, SettingField } from "@/app/[lang]/admin/setting-field"
|
||||
import { DictionaryProvider } from "@/hooks/use-dictionary"
|
||||
import type { SettingDef } from "@/lib/admin/settings-registry"
|
||||
import type { Dictionary } from "@/lib/i18n/dictionaries"
|
||||
import en from "@/lib/i18n/dictionaries/en.json"
|
||||
|
||||
const STORED = { isSet: true as const, hint: "…abcd" }
|
||||
|
||||
function withDict(node: ReactNode) {
|
||||
return (
|
||||
<DictionaryProvider dictionary={en as unknown as Dictionary}>
|
||||
{node}
|
||||
</DictionaryProvider>
|
||||
)
|
||||
}
|
||||
|
||||
function typeInto(label: string, text: string) {
|
||||
fireEvent.change(screen.getByLabelText(label, { selector: "input" }), {
|
||||
target: { value: text },
|
||||
})
|
||||
}
|
||||
|
||||
afterEach(cleanup)
|
||||
|
||||
describe("SecretInput", () => {
|
||||
it("reverts to a key saved after mount instead of deleting it", () => {
|
||||
const onChange = vi.fn()
|
||||
const props = { id: "secret", keepOnEmpty: true, onChange }
|
||||
// New provider: nothing stored at mount, then saved
|
||||
const { rerender } = render(
|
||||
withDict(
|
||||
<>
|
||||
<label htmlFor="secret">secret</label>
|
||||
<SecretInput {...props} value={undefined} />
|
||||
</>,
|
||||
),
|
||||
)
|
||||
rerender(
|
||||
withDict(
|
||||
<>
|
||||
<label htmlFor="secret">secret</label>
|
||||
<SecretInput {...props} value={STORED} />
|
||||
</>,
|
||||
),
|
||||
)
|
||||
rerender(
|
||||
withDict(
|
||||
<>
|
||||
<label htmlFor="secret">secret</label>
|
||||
<SecretInput {...props} value="abc" />
|
||||
</>,
|
||||
),
|
||||
)
|
||||
|
||||
typeInto("secret", "")
|
||||
expect(onChange).toHaveBeenLastCalledWith(STORED)
|
||||
})
|
||||
})
|
||||
|
||||
describe("SettingField secret", () => {
|
||||
const def: SettingDef = {
|
||||
key: "LANGFUSE_SECRET_KEY",
|
||||
group: "observability",
|
||||
type: "secret",
|
||||
label: "Langfuse Secret Key",
|
||||
}
|
||||
|
||||
it("drops the pending change when a saved secret is typed over and cleared", () => {
|
||||
const onChange = vi.fn()
|
||||
const props = {
|
||||
def,
|
||||
state: { key: def.key, source: "file" as const, value: STORED },
|
||||
disabled: false,
|
||||
onChange,
|
||||
}
|
||||
const { rerender } = render(
|
||||
withDict(<SettingField {...props} pendingValue={undefined} />),
|
||||
)
|
||||
rerender(withDict(<SettingField {...props} pendingValue="a" />))
|
||||
|
||||
typeInto(en.admin.settings.LANGFUSE_SECRET_KEY.label, "")
|
||||
expect(onChange).toHaveBeenLastCalledWith(undefined)
|
||||
})
|
||||
})
|
||||
@@ -1,7 +1,7 @@
|
||||
import fs from "fs"
|
||||
import os from "os"
|
||||
import path from "path"
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
import { afterEach, beforeEach, describe, expect, it } from "vitest"
|
||||
import {
|
||||
_resetForTests,
|
||||
applyToEnv,
|
||||
@@ -100,24 +100,6 @@ describe("applyToEnv / saveSettings", () => {
|
||||
expect(process.env.TEST_ADMIN_VAR).toBeUndefined()
|
||||
})
|
||||
|
||||
it("a second module instance can remove a key the first one overlaid", async () => {
|
||||
// instrumentation.ts and API routes load separate copies in a build
|
||||
process.env.TEST_ADMIN_VAR = "from-env"
|
||||
fs.writeFileSync(
|
||||
process.env.SETTINGS_FILE!,
|
||||
JSON.stringify({ version: 1, values: { TEST_ADMIN_VAR: "abc" } }),
|
||||
)
|
||||
applyToEnv()
|
||||
expect(process.env.TEST_ADMIN_VAR).toBe("abc")
|
||||
|
||||
vi.resetModules()
|
||||
const second = await import("@/lib/admin/settings")
|
||||
expect(second.getValueSource("TEST_ADMIN_VAR")).toBe("file")
|
||||
second.saveSettings({ TEST_ADMIN_VAR: null })
|
||||
expect(process.env.TEST_ADMIN_VAR).toBe("from-env")
|
||||
expect(second.getEnvFallback("TEST_ADMIN_VAR")).toBe("from-env")
|
||||
})
|
||||
|
||||
it("persists across cache reset (file round-trip)", () => {
|
||||
saveSettings({ TEST_ADMIN_VAR: "persisted" })
|
||||
_resetForTests()
|
||||
|
||||
@@ -1,274 +0,0 @@
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
import {
|
||||
getAIModel,
|
||||
getValidationModel,
|
||||
usesServerCredentials,
|
||||
} from "@/lib/ai-providers"
|
||||
|
||||
const settings = vi.hoisted(() => ({ values: {} as Record<string, string> }))
|
||||
|
||||
vi.mock("@/lib/admin/settings", () => ({
|
||||
loadSettings: () => settings.values,
|
||||
}))
|
||||
|
||||
vi.mock("@ai-sdk/google-vertex", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createVertex: vi.fn(() => mockProviderFn) }
|
||||
})
|
||||
|
||||
vi.mock("@ai-sdk/openai", () => {
|
||||
const mockModel = { modelId: "test-model" }
|
||||
const mockProviderFn = vi.fn(() => mockModel) as any
|
||||
mockProviderFn.chat = vi.fn(() => mockModel)
|
||||
return {
|
||||
createOpenAI: vi.fn(() => mockProviderFn),
|
||||
openai: vi.fn(() => mockModel),
|
||||
}
|
||||
})
|
||||
|
||||
vi.mock("@ai-sdk/amazon-bedrock", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createAmazonBedrock: vi.fn(() => mockProviderFn) }
|
||||
})
|
||||
|
||||
vi.mock("@aws-sdk/credential-providers", () => ({
|
||||
fromNodeProviderChain: vi.fn(() => "node-chain"),
|
||||
}))
|
||||
|
||||
vi.mock("@openrouter/ai-sdk-provider", () => {
|
||||
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
|
||||
return { createOpenRouter: vi.fn(() => mockProviderFn) }
|
||||
})
|
||||
|
||||
const ENV_KEYS = [
|
||||
"GOOGLE_VERTEX_API_KEY",
|
||||
"GOOGLE_VERTEX_BASE_URL",
|
||||
"OPENAI_API_KEY",
|
||||
"OPENAI_BASE_URL",
|
||||
"OPENROUTER_API_KEY",
|
||||
"ADMIN_OPENAI_API_KEY",
|
||||
"ADMIN_OPENROUTER_API_KEY",
|
||||
"OLLAMA_API_KEY",
|
||||
"ADMIN_AWS_ACCESS_KEY_ID",
|
||||
"ADMIN_AWS_SECRET_ACCESS_KEY",
|
||||
"ADMIN_AWS_REGION",
|
||||
"AWS_REGION",
|
||||
"AI_PROVIDER",
|
||||
"AI_MODEL",
|
||||
"VALIDATION_MODEL",
|
||||
]
|
||||
const savedEnv: Record<string, string | undefined> = {}
|
||||
|
||||
beforeEach(() => {
|
||||
for (const key of ENV_KEYS) {
|
||||
savedEnv[key] = process.env[key]
|
||||
delete process.env[key]
|
||||
}
|
||||
settings.values = {}
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
for (const key of ENV_KEYS) {
|
||||
if (savedEnv[key] === undefined) delete process.env[key]
|
||||
else process.env[key] = savedEnv[key]
|
||||
}
|
||||
})
|
||||
|
||||
describe("Vertex AI key security", () => {
|
||||
it("never sends the server key to a client base URL", () => {
|
||||
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
|
||||
|
||||
// Any x-ai-api-key passes the outer guard; the branch must still refuse
|
||||
expect(() =>
|
||||
getAIModel({
|
||||
provider: "vertexai",
|
||||
apiKey: "x",
|
||||
baseUrl: "https://attacker.example",
|
||||
modelId: "gemini-2.5-flash",
|
||||
}),
|
||||
).toThrow("Vertex AI requires an API key")
|
||||
})
|
||||
|
||||
it("sends the client key to the client base URL", async () => {
|
||||
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
|
||||
const { createVertex } = await import("@ai-sdk/google-vertex")
|
||||
|
||||
getAIModel({
|
||||
provider: "vertexai",
|
||||
vertexApiKey: "client-key",
|
||||
baseUrl: "https://my-proxy.example",
|
||||
modelId: "gemini-2.5-flash",
|
||||
})
|
||||
|
||||
expect(createVertex).toHaveBeenCalledWith({
|
||||
apiKey: "client-key",
|
||||
baseURL: "https://my-proxy.example",
|
||||
})
|
||||
})
|
||||
|
||||
it("does not send the client key to the server's base URL", async () => {
|
||||
process.env.GOOGLE_VERTEX_BASE_URL = "https://server-proxy.internal"
|
||||
const { createVertex } = await import("@ai-sdk/google-vertex")
|
||||
|
||||
getAIModel({
|
||||
provider: "vertexai",
|
||||
vertexApiKey: "client-key",
|
||||
modelId: "gemini-2.5-flash",
|
||||
})
|
||||
|
||||
expect(createVertex).toHaveBeenCalledWith({ apiKey: "client-key" })
|
||||
})
|
||||
|
||||
it("still uses the server key and base URL without client overrides", async () => {
|
||||
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
|
||||
process.env.GOOGLE_VERTEX_BASE_URL = "https://server-proxy.internal"
|
||||
const { createVertex } = await import("@ai-sdk/google-vertex")
|
||||
|
||||
getAIModel({ provider: "vertexai", modelId: "gemini-2.5-flash" })
|
||||
|
||||
expect(createVertex).toHaveBeenCalledWith({
|
||||
apiKey: "server-vertex-key",
|
||||
baseURL: "https://server-proxy.internal",
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe("Bedrock admin panel credentials", () => {
|
||||
it("uses the ADMIN_AWS_* keys when the client sends none", async () => {
|
||||
process.env.ADMIN_AWS_ACCESS_KEY_ID = "panel-id"
|
||||
process.env.ADMIN_AWS_SECRET_ACCESS_KEY = "panel-secret"
|
||||
process.env.ADMIN_AWS_REGION = "eu-west-1"
|
||||
process.env.AWS_REGION = "us-east-1"
|
||||
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
|
||||
|
||||
getAIModel({ provider: "bedrock", modelId: "amazon.nova-lite-v1:0" })
|
||||
|
||||
expect(createAmazonBedrock).toHaveBeenCalledWith({
|
||||
region: "eu-west-1",
|
||||
accessKeyId: "panel-id",
|
||||
secretAccessKey: "panel-secret",
|
||||
})
|
||||
})
|
||||
|
||||
it("prefers the client's keys and region", async () => {
|
||||
process.env.ADMIN_AWS_ACCESS_KEY_ID = "panel-id"
|
||||
process.env.ADMIN_AWS_SECRET_ACCESS_KEY = "panel-secret"
|
||||
process.env.ADMIN_AWS_REGION = "eu-west-1"
|
||||
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
|
||||
|
||||
getAIModel({
|
||||
provider: "bedrock",
|
||||
modelId: "amazon.nova-lite-v1:0",
|
||||
awsAccessKeyId: "client-id",
|
||||
awsSecretAccessKey: "client-secret",
|
||||
awsRegion: "ap-northeast-1",
|
||||
})
|
||||
|
||||
expect(createAmazonBedrock).toHaveBeenCalledWith({
|
||||
region: "ap-northeast-1",
|
||||
accessKeyId: "client-id",
|
||||
secretAccessKey: "client-secret",
|
||||
})
|
||||
})
|
||||
|
||||
it("falls back to the default AWS credential chain", async () => {
|
||||
process.env.AWS_REGION = "us-east-1"
|
||||
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
|
||||
|
||||
getAIModel({ provider: "bedrock", modelId: "amazon.nova-lite-v1:0" })
|
||||
|
||||
expect(createAmazonBedrock).toHaveBeenCalledWith({
|
||||
region: "us-east-1",
|
||||
credentialProvider: "node-chain",
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe("usesServerCredentials", () => {
|
||||
it("is true when no key comes with the request", () => {
|
||||
expect(usesServerCredentials("openai", {})).toBe(true)
|
||||
expect(usesServerCredentials("openai", { apiKey: "k" })).toBe(false)
|
||||
})
|
||||
|
||||
it("looks at the credential each provider actually uses", () => {
|
||||
// A stray x-ai-api-key does not replace the IAM role or Vertex key
|
||||
expect(usesServerCredentials("bedrock", { apiKey: "x" })).toBe(true)
|
||||
expect(
|
||||
usesServerCredentials("bedrock", {
|
||||
awsAccessKeyId: "id",
|
||||
awsSecretAccessKey: "secret",
|
||||
}),
|
||||
).toBe(false)
|
||||
expect(usesServerCredentials("vertexai", { apiKey: "x" })).toBe(true)
|
||||
expect(usesServerCredentials("vertexai", { vertexApiKey: "k" })).toBe(
|
||||
false,
|
||||
)
|
||||
})
|
||||
|
||||
it("treats keyless EdgeOne and local Ollama as free", () => {
|
||||
expect(usesServerCredentials("edgeone", {})).toBe(false)
|
||||
expect(usesServerCredentials("ollama", {})).toBe(false)
|
||||
expect(
|
||||
usesServerCredentials("ollama", {
|
||||
baseUrl: "http://localhost:11434",
|
||||
}),
|
||||
).toBe(false)
|
||||
|
||||
process.env.OLLAMA_API_KEY = "server-ollama-key"
|
||||
expect(usesServerCredentials("ollama", {})).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
describe("server model apiKeyEnv", () => {
|
||||
it("uses the custom env var on the official OpenAI endpoint", async () => {
|
||||
process.env.ADMIN_OPENAI_API_KEY = "panel-key"
|
||||
const { createOpenAI, openai } = await import("@ai-sdk/openai")
|
||||
|
||||
getAIModel({
|
||||
provider: "openai",
|
||||
modelId: "gpt-4o",
|
||||
apiKeyEnv: "ADMIN_OPENAI_API_KEY",
|
||||
})
|
||||
|
||||
// The default instance would read OPENAI_API_KEY instead
|
||||
expect(openai).not.toHaveBeenCalled()
|
||||
expect(createOpenAI).toHaveBeenCalledWith({ apiKey: "panel-key" })
|
||||
})
|
||||
})
|
||||
|
||||
describe("getValidationModel", () => {
|
||||
it("uses the admin panel default's ADMIN_ key", async () => {
|
||||
settings.values = {
|
||||
ADMIN_PROVIDERS: JSON.stringify([
|
||||
{
|
||||
id: "p1",
|
||||
provider: "openrouter",
|
||||
name: "My OpenRouter",
|
||||
apiKey: "panel-key",
|
||||
models: ["openai/gpt-4o"],
|
||||
isDefault: true,
|
||||
},
|
||||
]),
|
||||
}
|
||||
// What deriveEnvUpdates writes for that panel config
|
||||
process.env.AI_PROVIDER = "openrouter"
|
||||
process.env.AI_MODEL = "openai/gpt-4o"
|
||||
process.env.ADMIN_OPENROUTER_API_KEY = "panel-key"
|
||||
const { createOpenRouter } = await import("@openrouter/ai-sdk-provider")
|
||||
|
||||
expect(() => getValidationModel()).not.toThrow()
|
||||
expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "panel-key" })
|
||||
})
|
||||
|
||||
it("uses the standard env vars without a panel default", async () => {
|
||||
process.env.AI_PROVIDER = "openrouter"
|
||||
process.env.AI_MODEL = "openai/gpt-4o"
|
||||
process.env.OPENROUTER_API_KEY = "env-key"
|
||||
const { createOpenRouter } = await import("@openrouter/ai-sdk-provider")
|
||||
|
||||
getValidationModel()
|
||||
|
||||
expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "env-key" })
|
||||
})
|
||||
})
|
||||
@@ -1,156 +0,0 @@
|
||||
// @vitest-environment node
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
import { POST as parseUrl } from "@/app/api/parse-url/route"
|
||||
import { POST as validateDiagram } from "@/app/api/validate-diagram/route"
|
||||
import { POST as validateModel } from "@/app/api/validate-model/route"
|
||||
import { POST as verifyAccessCode } from "@/app/api/verify-access-code/route"
|
||||
import { checkAccessCode } from "@/lib/access-code"
|
||||
|
||||
// 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,
|
||||
}))
|
||||
|
||||
function post(path: string, body: unknown, accessCode?: string): Request {
|
||||
return new Request(`http://localhost${path}`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
...(accessCode ? { "x-access-code": accessCode } : {}),
|
||||
},
|
||||
body: JSON.stringify(body),
|
||||
})
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
process.env.ACCESS_CODE_LIST = "secret, other"
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
delete process.env.ACCESS_CODE_LIST
|
||||
delete process.env.ALLOW_PRIVATE_URLS
|
||||
vi.unstubAllGlobals()
|
||||
})
|
||||
|
||||
describe("checkAccessCode", () => {
|
||||
it("passes when no access codes are configured", () => {
|
||||
delete process.env.ACCESS_CODE_LIST
|
||||
expect(checkAccessCode(post("/x", {}))).toBeNull()
|
||||
})
|
||||
|
||||
it("rejects a missing or wrong code and accepts a listed one", () => {
|
||||
expect(checkAccessCode(post("/x", {}))?.status).toBe(401)
|
||||
expect(checkAccessCode(post("/x", {}, "nope"))?.status).toBe(401)
|
||||
expect(checkAccessCode(post("/x", {}, "other"))).toBeNull()
|
||||
})
|
||||
})
|
||||
|
||||
describe("routes that spend server resources require the access code", () => {
|
||||
it("parse-url", async () => {
|
||||
const res = await parseUrl(
|
||||
post("/api/parse-url", { url: "https://example.com" }),
|
||||
)
|
||||
expect(res.status).toBe(401)
|
||||
})
|
||||
|
||||
it("validate-diagram", async () => {
|
||||
const res = await validateDiagram(
|
||||
post("/api/validate-diagram", {
|
||||
imageData: "data:image/png;base64,AAAA",
|
||||
}),
|
||||
)
|
||||
expect(res.status).toBe(401)
|
||||
})
|
||||
|
||||
it("validate-model", async () => {
|
||||
const res = await validateModel(
|
||||
post("/api/validate-model", {
|
||||
provider: "openai",
|
||||
apiKey: "sk",
|
||||
modelId: "m",
|
||||
}),
|
||||
)
|
||||
expect(res.status).toBe(401)
|
||||
})
|
||||
|
||||
it("verify-access-code", async () => {
|
||||
const bad = await verifyAccessCode(post("/api/verify-access-code", {}))
|
||||
expect(bad.status).toBe(401)
|
||||
expect((await bad.json()).valid).toBe(false)
|
||||
|
||||
const good = await verifyAccessCode(
|
||||
post("/api/verify-access-code", {}, "secret"),
|
||||
)
|
||||
expect((await good.json()).valid).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
describe("size limits", () => {
|
||||
it("parse-url stops reading a body over the download limit", async () => {
|
||||
const chunk = new Uint8Array(1024 * 1024)
|
||||
let sent = 0
|
||||
const body = new ReadableStream<Uint8Array>({
|
||||
pull(controller) {
|
||||
sent++
|
||||
controller.enqueue(chunk)
|
||||
},
|
||||
})
|
||||
vi.stubGlobal(
|
||||
"fetch",
|
||||
vi.fn(
|
||||
async () =>
|
||||
new Response(body, {
|
||||
headers: { "content-type": "text/html" },
|
||||
}),
|
||||
),
|
||||
)
|
||||
|
||||
const res = await parseUrl(
|
||||
post("/api/parse-url", { url: "https://example.com" }, "secret"),
|
||||
)
|
||||
expect(res.status).toBe(413)
|
||||
expect(sent).toBeLessThan(10)
|
||||
})
|
||||
|
||||
it("validate-diagram rejects oversized image data", async () => {
|
||||
const imageData = `data:image/png;base64,${"A".repeat(6 * 1024 * 1024)}`
|
||||
const res = await validateDiagram(
|
||||
post("/api/validate-diagram", { imageData }, "secret"),
|
||||
)
|
||||
expect(res.status).toBe(413)
|
||||
})
|
||||
})
|
||||
|
||||
describe("validate-model redirects", () => {
|
||||
it("refuses redirects when private URLs are blocked", async () => {
|
||||
process.env.ALLOW_PRIVATE_URLS = "false"
|
||||
vi.stubGlobal(
|
||||
"fetch",
|
||||
vi.fn(
|
||||
async () =>
|
||||
new Response(null, {
|
||||
status: 302,
|
||||
headers: { location: "http://169.254.169.254/" },
|
||||
}),
|
||||
),
|
||||
)
|
||||
|
||||
const res = await validateModel(
|
||||
post(
|
||||
"/api/validate-model",
|
||||
{
|
||||
provider: "openai",
|
||||
apiKey: "sk",
|
||||
modelId: "m",
|
||||
baseUrl: "https://attacker.example/v1",
|
||||
},
|
||||
"secret",
|
||||
),
|
||||
)
|
||||
const data = await res.json()
|
||||
expect(data.valid).toBe(false)
|
||||
expect(data.error).toMatch(/Redirects are not allowed/)
|
||||
expect(fetch).toHaveBeenCalledTimes(1)
|
||||
})
|
||||
})
|
||||
@@ -14,35 +14,12 @@ describe("findCachedResponse", () => {
|
||||
expect(result?.xml).toContain("Transformer Architecture")
|
||||
})
|
||||
|
||||
it("returns cached response for exact match with the example file", () => {
|
||||
const result = findCachedResponse(
|
||||
"Replicate this in aws style",
|
||||
true,
|
||||
"architecture.png",
|
||||
)
|
||||
it("returns cached response for exact match with image", () => {
|
||||
const result = findCachedResponse("Replicate this in aws style", true)
|
||||
expect(result).toBeDefined()
|
||||
expect(result?.xml).toContain("AWS")
|
||||
})
|
||||
|
||||
it("returns undefined when the user attached their own file", () => {
|
||||
expect(
|
||||
findCachedResponse("Replicate this flowchart.", true, "mine.png"),
|
||||
).toBeUndefined()
|
||||
expect(
|
||||
findCachedResponse(
|
||||
"Summarize this paper as a diagram",
|
||||
true,
|
||||
"thesis.pdf",
|
||||
),
|
||||
).toBeUndefined()
|
||||
})
|
||||
|
||||
it("returns undefined for file examples when the file name is unknown", () => {
|
||||
// The server only knows whether a file is attached, not which one
|
||||
const result = findCachedResponse("Replicate this in aws style", true)
|
||||
expect(result).toBeUndefined()
|
||||
})
|
||||
|
||||
it("returns undefined for non-matching prompt", () => {
|
||||
const result = findCachedResponse(
|
||||
"random prompt that doesn't exist",
|
||||
|
||||
@@ -1,11 +1,6 @@
|
||||
// @vitest-environment node
|
||||
|
||||
import { convertToModelMessages } from "ai"
|
||||
import { jsonrepair } from "jsonrepair"
|
||||
import { describe, expect, it } from "vitest"
|
||||
import {
|
||||
dropInvalidToolCalls,
|
||||
fixToolInputJson,
|
||||
isMinimalDiagram,
|
||||
replaceHistoricalToolInputs,
|
||||
validateFileParts,
|
||||
@@ -70,29 +65,6 @@ describe("validateFileParts", () => {
|
||||
expect(result.valid).toBe(false)
|
||||
expect(result.error).toContain("exceeds")
|
||||
})
|
||||
|
||||
it("rejects file URLs the server would have to download", () => {
|
||||
for (const url of [
|
||||
"http://10.0.0.5/secret.png",
|
||||
"https://example.com/a.png",
|
||||
undefined,
|
||||
]) {
|
||||
const messages = [{ role: "user", parts: [{ type: "file", url }] }]
|
||||
expect(validateFileParts(messages).valid).toBe(false)
|
||||
}
|
||||
})
|
||||
|
||||
it("checks files in earlier messages too", () => {
|
||||
const messages = [
|
||||
{
|
||||
role: "user",
|
||||
parts: [{ type: "file", url: "http://169.254.169.254/x" }],
|
||||
},
|
||||
{ role: "assistant", parts: [{ type: "text", text: "ok" }] },
|
||||
{ role: "user", parts: [{ type: "text", text: "hello" }] },
|
||||
]
|
||||
expect(validateFileParts(messages).valid).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe("isMinimalDiagram", () => {
|
||||
@@ -111,18 +83,6 @@ describe("isMinimalDiagram", () => {
|
||||
const xml = ' <mxCell id="0"/> <mxCell id="1" parent="0"/> '
|
||||
expect(isMinimalDiagram(xml)).toBe(true)
|
||||
})
|
||||
|
||||
it("returns false for a shape drawn in draw.io with a random id", () => {
|
||||
const xml =
|
||||
'<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="xY3kQ9-1" value="" style="rounded=0;" vertex="1" parent="1"><mxGeometry x="10" y="10" width="120" height="60" as="geometry"/></mxCell></root></mxGraphModel>'
|
||||
expect(isMinimalDiagram(xml)).toBe(false)
|
||||
})
|
||||
|
||||
it("does not mistake ids that start with 0 or 1 for root cells", () => {
|
||||
const xml =
|
||||
'<mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="10"/>'
|
||||
expect(isMinimalDiagram(xml)).toBe(false)
|
||||
})
|
||||
})
|
||||
|
||||
describe("replaceHistoricalToolInputs", () => {
|
||||
@@ -164,7 +124,7 @@ describe("replaceHistoricalToolInputs", () => {
|
||||
)
|
||||
})
|
||||
|
||||
it("leaves tool calls with invalid inputs for dropInvalidToolCalls", () => {
|
||||
it("removes tool calls with invalid inputs", () => {
|
||||
const messages = [
|
||||
{
|
||||
role: "assistant",
|
||||
@@ -183,7 +143,7 @@ describe("replaceHistoricalToolInputs", () => {
|
||||
},
|
||||
]
|
||||
const result = replaceHistoricalToolInputs(messages)
|
||||
expect(result[0].content).toEqual(messages[0].content)
|
||||
expect(result[0].content).toHaveLength(0)
|
||||
})
|
||||
|
||||
it("preserves non-assistant messages", () => {
|
||||
@@ -209,123 +169,3 @@ describe("replaceHistoricalToolInputs", () => {
|
||||
expect(result[0].content[0].input).toEqual({ foo: "bar" })
|
||||
})
|
||||
})
|
||||
|
||||
describe("dropInvalidToolCalls", () => {
|
||||
it("drops an invalid tool-call together with its tool-result", () => {
|
||||
const messages = [
|
||||
{ role: "user", content: [{ type: "text", text: "draw" }] },
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{
|
||||
type: "tool-call",
|
||||
toolCallId: "call-1",
|
||||
toolName: "display_diagram",
|
||||
input: undefined,
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
role: "tool",
|
||||
content: [
|
||||
{
|
||||
type: "tool-result",
|
||||
toolCallId: "call-1",
|
||||
toolName: "display_diagram",
|
||||
output: { type: "error-text", value: "Stopped" },
|
||||
},
|
||||
],
|
||||
},
|
||||
{ role: "user", content: [{ type: "text", text: "again" }] },
|
||||
]
|
||||
const result = dropInvalidToolCalls(messages)
|
||||
expect(result.map((m) => m.role)).toEqual(["user", "user"])
|
||||
})
|
||||
|
||||
it("keeps valid calls and results in the same messages", () => {
|
||||
const messages = [
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{ type: "text", text: "Here you go" },
|
||||
{
|
||||
type: "tool-call",
|
||||
toolCallId: "bad",
|
||||
toolName: "edit_diagram",
|
||||
input: "{broken",
|
||||
},
|
||||
{
|
||||
type: "tool-call",
|
||||
toolCallId: "good",
|
||||
toolName: "display_diagram",
|
||||
input: { xml: "<mxCell/>" },
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
role: "tool",
|
||||
content: [
|
||||
{ type: "tool-result", toolCallId: "bad", output: {} },
|
||||
{ type: "tool-result", toolCallId: "good", output: {} },
|
||||
],
|
||||
},
|
||||
]
|
||||
const result = dropInvalidToolCalls(messages)
|
||||
expect(result[0].content.map((p: any) => p.toolCallId)).toEqual([
|
||||
undefined,
|
||||
"good",
|
||||
])
|
||||
expect(result[1].content.map((p: any) => p.toolCallId)).toEqual([
|
||||
"good",
|
||||
])
|
||||
})
|
||||
|
||||
it("cleans up a tool call the user stopped before its input arrived", async () => {
|
||||
// handleStop turns a still-streaming call into output-error with no input
|
||||
const modelMessages = await convertToModelMessages([
|
||||
{ role: "user", parts: [{ type: "text", text: "draw" }] },
|
||||
{
|
||||
role: "assistant",
|
||||
parts: [
|
||||
{
|
||||
type: "tool-display_diagram",
|
||||
toolCallId: "call-1",
|
||||
state: "output-error",
|
||||
input: undefined,
|
||||
errorText: "Stopped by user",
|
||||
} as any,
|
||||
],
|
||||
},
|
||||
{ role: "user", parts: [{ type: "text", text: "again" }] },
|
||||
])
|
||||
const result = dropInvalidToolCalls(modelMessages)
|
||||
expect(result.map((m) => m.role)).toEqual(["user", "user"])
|
||||
})
|
||||
|
||||
it("leaves messages with string content alone", () => {
|
||||
const messages = [{ role: "system", content: "You are..." }]
|
||||
expect(dropInvalidToolCalls(messages)).toEqual(messages)
|
||||
})
|
||||
})
|
||||
|
||||
describe("fixToolInputJson", () => {
|
||||
it("fixes an attribute whose closing quote alone is escaped", () => {
|
||||
const input =
|
||||
'{"xml": "<mxCell id=\\"2\\" vertex=\\"1\\"><mxGeometry x=\\"10\\" y="-20\\" as=\\"geometry\\"/></mxCell>"}'
|
||||
const parsed = JSON.parse(jsonrepair(fixToolInputJson(input)))
|
||||
expect(parsed.xml).toContain('y="-20"')
|
||||
expect(parsed.xml).toContain('id="2"')
|
||||
})
|
||||
|
||||
it("fixes = used instead of : after a JSON key", () => {
|
||||
const input = '{"xml"= "<mxCell id=\\"2\\"/>"}'
|
||||
const parsed = JSON.parse(jsonrepair(fixToolInputJson(input)))
|
||||
expect(parsed.xml).toBe('<mxCell id="2"/>')
|
||||
})
|
||||
|
||||
it("leaves well-formed input unchanged", () => {
|
||||
const input =
|
||||
'{"operations": [{"operation": "add", "cell_id": "a", "new_xml": "<mxCell id=\\"a\\" value=\\"x=1\\"/>"}]}'
|
||||
expect(fixToolInputJson(input)).toBe(input)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
// @vitest-environment node
|
||||
import { describe, expect, it } from "vitest"
|
||||
import { onRequest } from "@/edge-functions/api/edgeai/chat/completions"
|
||||
|
||||
function request(headers: Record<string, string>): Request {
|
||||
return new Request("http://localhost/api/edgeai/chat/completions", {
|
||||
method: "POST",
|
||||
headers,
|
||||
// Non-streaming requests return a mock reply without calling AI
|
||||
body: JSON.stringify({ messages: [{ role: "user", content: "hi" }] }),
|
||||
})
|
||||
}
|
||||
|
||||
const json = { "Content-Type": "application/json" }
|
||||
|
||||
describe("EdgeOne chat completions function", () => {
|
||||
it("sends no CORS headers", async () => {
|
||||
const res = await onRequest({ request: request(json), env: {} })
|
||||
expect(res.status).toBe(200)
|
||||
expect(res.headers.get("access-control-allow-origin")).toBeNull()
|
||||
})
|
||||
|
||||
it("rejects non-JSON requests", async () => {
|
||||
const res = await onRequest({
|
||||
request: request({ "Content-Type": "text/plain" }),
|
||||
env: {},
|
||||
})
|
||||
expect(res.status).toBe(400)
|
||||
})
|
||||
|
||||
it("checks the access code when ACCESS_CODE_LIST is set", async () => {
|
||||
const env = { ACCESS_CODE_LIST: "secret" }
|
||||
const missing = await onRequest({ request: request(json), env })
|
||||
expect(missing.status).toBe(401)
|
||||
|
||||
const ok = await onRequest({
|
||||
request: request({ ...json, "x-access-code": "secret" }),
|
||||
env,
|
||||
})
|
||||
expect(ok.status).toBe(200)
|
||||
})
|
||||
|
||||
it("lets requests through when env is unavailable", async () => {
|
||||
const res = await onRequest({ request: request(json) })
|
||||
expect(res.status).toBe(200)
|
||||
})
|
||||
})
|
||||
@@ -3,7 +3,6 @@ import {
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
parseOutputTokenLimit,
|
||||
resolveMaxOutputTokens,
|
||||
retryOutputTokens,
|
||||
withOutputTokenLimitFallback,
|
||||
} from "@/lib/output-token-limit"
|
||||
|
||||
@@ -99,204 +98,34 @@ describe("parseOutputTokenLimit", () => {
|
||||
}
|
||||
expect(parseOutputTokenLimit(error)).toBeNull()
|
||||
})
|
||||
|
||||
it("reads the ceiling from a Volcengine Ark rejection", () => {
|
||||
const error = {
|
||||
message:
|
||||
"The parameter `max_tokens` specified in the request are not valid: integer above maximum value, expected a value <= 32768, but got 64000 instead.",
|
||||
statusCode: 400,
|
||||
}
|
||||
expect(parseOutputTokenLimit(error)).toBe(32768)
|
||||
// Same message JSON-escaped in the response body
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message: "Bad request",
|
||||
responseBody:
|
||||
'{"error":{"message":"The parameter `max_tokens` specified in the request are not valid: integer above maximum value, expected a value \\u003c= 16384, but got 64000 instead."}}',
|
||||
}),
|
||||
).toBe(16384)
|
||||
})
|
||||
|
||||
it("reads the ceiling from a DashScope rejection", () => {
|
||||
const error = {
|
||||
message:
|
||||
"<400> InternalError.Algo.InvalidParameter: Range of max_tokens should be [1, 8192]",
|
||||
}
|
||||
expect(parseOutputTokenLimit(error)).toBe(8192)
|
||||
})
|
||||
|
||||
it("subtracts the input in SGLang and vLLM context rejections", () => {
|
||||
// SGLang
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message:
|
||||
"Requested token count exceeds the model's maximum context length of 32768 tokens. You requested a total of 70000 tokens: 6000 tokens from the input messages and 64000 tokens for the completion.",
|
||||
}),
|
||||
).toBe(32768 - 6000 - 1024)
|
||||
// vLLM, older wording
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message:
|
||||
"This model's maximum context length is 32768 tokens. However, you requested 70000 tokens (6000 in the messages, 64000 in the completion).",
|
||||
}),
|
||||
).toBe(32768 - 6000 - 1024)
|
||||
// vLLM, newer wording
|
||||
expect(
|
||||
parseOutputTokenLimit({
|
||||
message:
|
||||
"This model's maximum context length is 32768 tokens and your request has 6000 input tokens (64000 > 32768 - 6000).",
|
||||
}),
|
||||
).toBe(32768 - 6000 - 1024)
|
||||
})
|
||||
})
|
||||
|
||||
describe("retryOutputTokens", () => {
|
||||
const bedrockLimit = Object.assign(
|
||||
new Error(
|
||||
"The maximum tokens you requested exceeds the model limit of 64000.",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
|
||||
it("leaves room for the Bedrock thinking budget the provider adds", () => {
|
||||
// 64000 + 12000 thinking was sent, so the ceiling of 64000 is below it
|
||||
expect(
|
||||
retryOutputTokens(bedrockLimit, {
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
bedrock: {
|
||||
reasoningConfig: {
|
||||
type: "enabled",
|
||||
budgetTokens: 12000,
|
||||
},
|
||||
},
|
||||
},
|
||||
}),
|
||||
).toBe(52000)
|
||||
})
|
||||
|
||||
it("leaves room for the Anthropic thinking budget the provider adds", () => {
|
||||
const error = Object.assign(
|
||||
new Error(
|
||||
"max_tokens: 76000 > 64000, which is the maximum allowed number of output tokens",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
expect(
|
||||
retryOutputTokens(error, {
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
anthropic: {
|
||||
thinking: { type: "enabled", budgetTokens: 12000 },
|
||||
},
|
||||
},
|
||||
}),
|
||||
).toBe(52000)
|
||||
})
|
||||
|
||||
it("does not retry when the ceiling covers what was sent", () => {
|
||||
expect(
|
||||
retryOutputTokens(bedrockLimit, { maxOutputTokens: 64000 }),
|
||||
).toBeNull()
|
||||
})
|
||||
|
||||
it("does not retry when the thinking budget leaves no usable room", () => {
|
||||
const error = Object.assign(new Error("model limit of 16000"), {
|
||||
statusCode: 400,
|
||||
})
|
||||
expect(
|
||||
retryOutputTokens(error, {
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
bedrock: {
|
||||
reasoningConfig: {
|
||||
type: "enabled",
|
||||
budgetTokens: 15500,
|
||||
},
|
||||
},
|
||||
},
|
||||
}),
|
||||
).toBeNull()
|
||||
})
|
||||
|
||||
it("falls back to 16000 when the budget is named but no number can be read", () => {
|
||||
const error = Object.assign(
|
||||
new Error("max_tokens (64000) exceeds the limit for this model"),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBe(16000)
|
||||
// Nothing to gain when the request was already that small
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 16000 })).toBeNull()
|
||||
})
|
||||
|
||||
it("does not fall back for errors that do not name the budget", () => {
|
||||
const error = Object.assign(new Error("temperature must be <= 2"), {
|
||||
statusCode: 400,
|
||||
})
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBeNull()
|
||||
// Not a bad request, even though it names the budget
|
||||
const auth = Object.assign(new Error("max_tokens: invalid API key"), {
|
||||
statusCode: 401,
|
||||
})
|
||||
expect(retryOutputTokens(auth, { maxOutputTokens: 64000 })).toBeNull()
|
||||
})
|
||||
|
||||
it("does not fall back when a ceiling was found but is too small", () => {
|
||||
const error = Object.assign(
|
||||
new Error(
|
||||
"max_tokens: 64000 > 512, which is the maximum allowed number of output tokens",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
)
|
||||
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBeNull()
|
||||
})
|
||||
})
|
||||
|
||||
describe("resolveMaxOutputTokens", () => {
|
||||
it("uses a valid header value", () => {
|
||||
expect(resolveMaxOutputTokens("32000", false)).toBe(32000)
|
||||
expect(resolveMaxOutputTokens("32000", true)).toBe(32000)
|
||||
expect(resolveMaxOutputTokens("32000")).toBe(32000)
|
||||
})
|
||||
|
||||
it("falls back to the default for missing or bogus values", () => {
|
||||
for (const value of [null, "", "abc", "0", "-5", "1.5", "640000"]) {
|
||||
// "640000" is above the sanity ceiling, e.g. an extra zero
|
||||
expect(resolveMaxOutputTokens(value, false)).toBe(
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
}
|
||||
expect(resolveMaxOutputTokens(null)).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("abc")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("0")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("-5")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
expect(resolveMaxOutputTokens("1.5")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
// Above the sanity ceiling, e.g. an extra zero
|
||||
expect(resolveMaxOutputTokens("640000")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
})
|
||||
|
||||
it("uses the env value when no header is sent, and validates it too", () => {
|
||||
const original = process.env.MAX_OUTPUT_TOKENS
|
||||
try {
|
||||
process.env.MAX_OUTPUT_TOKENS = "24000"
|
||||
expect(resolveMaxOutputTokens(null, true)).toBe(24000)
|
||||
// A lower header still wins
|
||||
expect(resolveMaxOutputTokens("8000", true)).toBe(8000)
|
||||
expect(resolveMaxOutputTokens(null)).toBe(24000)
|
||||
// Header still wins
|
||||
expect(resolveMaxOutputTokens("8000")).toBe(8000)
|
||||
|
||||
process.env.MAX_OUTPUT_TOKENS = "-1"
|
||||
expect(resolveMaxOutputTokens(null, true)).toBe(
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
} finally {
|
||||
if (original === undefined) delete process.env.MAX_OUTPUT_TOKENS
|
||||
else process.env.MAX_OUTPUT_TOKENS = original
|
||||
}
|
||||
})
|
||||
|
||||
it("lets the header raise the budget only on the user's own credentials", () => {
|
||||
const original = process.env.MAX_OUTPUT_TOKENS
|
||||
try {
|
||||
process.env.MAX_OUTPUT_TOKENS = "16000"
|
||||
expect(resolveMaxOutputTokens("200000", true)).toBe(16000)
|
||||
expect(resolveMaxOutputTokens("200000", false)).toBe(200000)
|
||||
|
||||
// Without MAX_OUTPUT_TOKENS the default is the cap
|
||||
delete process.env.MAX_OUTPUT_TOKENS
|
||||
expect(resolveMaxOutputTokens("100000", true)).toBe(
|
||||
DEFAULT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
expect(resolveMaxOutputTokens(null)).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
|
||||
} finally {
|
||||
if (original === undefined) delete process.env.MAX_OUTPUT_TOKENS
|
||||
else process.env.MAX_OUTPUT_TOKENS = original
|
||||
@@ -349,34 +178,6 @@ describe("withOutputTokenLimitFallback", () => {
|
||||
expect(calls.map((c) => c.maxOutputTokens)).toEqual([64000, 4096])
|
||||
})
|
||||
|
||||
it("subtracts the thinking budget from the retry", async () => {
|
||||
const [model, calls] = fakeModel([
|
||||
() =>
|
||||
Promise.reject(
|
||||
Object.assign(
|
||||
new Error(
|
||||
"The maximum tokens you requested exceeds the model limit of 64000.",
|
||||
),
|
||||
{ statusCode: 400 },
|
||||
),
|
||||
),
|
||||
() => Promise.resolve(STREAM_OK),
|
||||
])
|
||||
|
||||
const wrapped = withOutputTokenLimitFallback(model)
|
||||
await wrapped.doStream({
|
||||
prompt: [],
|
||||
maxOutputTokens: 64000,
|
||||
providerOptions: {
|
||||
bedrock: {
|
||||
reasoningConfig: { type: "enabled", budgetTokens: 12000 },
|
||||
},
|
||||
},
|
||||
} as any)
|
||||
|
||||
expect(calls.map((c) => c.maxOutputTokens)).toEqual([64000, 52000])
|
||||
})
|
||||
|
||||
it("does not retry an error it cannot attribute to the budget", async () => {
|
||||
const [model, calls] = fakeModel([
|
||||
() =>
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
import { describe, expect, it } from "vitest"
|
||||
import { isTextFile } from "@/lib/pdf-utils"
|
||||
|
||||
describe("isTextFile", () => {
|
||||
it("treats SVG files as text so their markup is sent to the model", () => {
|
||||
const svg = new File(["<svg/>"], "diagram.svg", {
|
||||
type: "image/svg+xml",
|
||||
})
|
||||
expect(isTextFile(svg)).toBe(true)
|
||||
})
|
||||
|
||||
it("does not treat raster images as text", () => {
|
||||
const png = new File(["x"], "photo.png", { type: "image/png" })
|
||||
expect(isTextFile(png)).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -4,7 +4,6 @@ import {
|
||||
loadFlattenedServerModels,
|
||||
type ServerModelsConfig,
|
||||
ServerModelsConfigSchema,
|
||||
slugify,
|
||||
} from "@/lib/server-model-config"
|
||||
|
||||
const ORIGINAL_ENV = { ...process.env }
|
||||
@@ -234,50 +233,3 @@ describe("loadFlattenedServerModels", () => {
|
||||
expect(models[0].apiKeyEnv).toEqual(["OPENAI_KEY_1", "OPENAI_KEY_2"])
|
||||
})
|
||||
})
|
||||
|
||||
describe("slugify", () => {
|
||||
it("keeps ASCII names readable", () => {
|
||||
expect(slugify("OpenAI Production")).toBe("openai-production")
|
||||
})
|
||||
|
||||
it("gives distinct ASCII slugs to distinct CJK names", () => {
|
||||
const slugs = ["主力", "备用", "DeepSeek 官方", "DeepSeek 备用"].map(
|
||||
slugify,
|
||||
)
|
||||
expect(new Set(slugs).size).toBe(4)
|
||||
for (const slug of slugs) expect(slug).toMatch(/^[a-z0-9-]+$/)
|
||||
})
|
||||
})
|
||||
|
||||
describe("loadFlattenedServerModels id collisions", () => {
|
||||
it("drops a model whose id repeats an earlier provider's", async () => {
|
||||
const config: ServerModelsConfig = {
|
||||
providers: [
|
||||
{ name: "OpenAI", provider: "openai", models: ["gpt-4o"] },
|
||||
{
|
||||
name: "openai",
|
||||
provider: "openai",
|
||||
models: ["gpt-4o"],
|
||||
apiKeyEnv: "OTHER_KEY",
|
||||
},
|
||||
{
|
||||
name: "主力",
|
||||
provider: "deepseek",
|
||||
models: ["deepseek-chat"],
|
||||
},
|
||||
{
|
||||
name: "备用",
|
||||
provider: "deepseek",
|
||||
models: ["deepseek-chat"],
|
||||
},
|
||||
],
|
||||
}
|
||||
process.env.AI_MODELS_CONFIG = JSON.stringify(config)
|
||||
|
||||
const models = await loadFlattenedServerModels()
|
||||
const ids = models.map((m) => m.id)
|
||||
expect(new Set(ids).size).toBe(ids.length)
|
||||
expect(ids).toHaveLength(3)
|
||||
expect(models[0].apiKeyEnv).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,25 +0,0 @@
|
||||
import { afterEach, describe, expect, it, vi } from "vitest"
|
||||
import { STORAGE_KEYS } from "@/lib/storage"
|
||||
import { extractUrlContent } from "@/lib/url-utils"
|
||||
|
||||
describe("extractUrlContent", () => {
|
||||
afterEach(() => {
|
||||
vi.unstubAllGlobals()
|
||||
localStorage.clear()
|
||||
})
|
||||
|
||||
it("sends the saved access code with the request", async () => {
|
||||
localStorage.setItem(STORAGE_KEYS.accessCode, "secret")
|
||||
const body = { title: "T", content: "body", charCount: 4 }
|
||||
const fetchMock = vi
|
||||
.fn()
|
||||
.mockResolvedValue(new Response(JSON.stringify(body)))
|
||||
vi.stubGlobal("fetch", fetchMock)
|
||||
|
||||
const data = await extractUrlContent("https://example.com")
|
||||
|
||||
expect(data.content).toBe("body")
|
||||
const headers = fetchMock.mock.calls[0][1].headers
|
||||
expect(headers["x-access-code"]).toBe("secret")
|
||||
})
|
||||
})
|
||||
@@ -1,102 +0,0 @@
|
||||
import { act, renderHook } from "@testing-library/react"
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest"
|
||||
import { extractPdfText, extractTextFileContent } from "@/lib/pdf-utils"
|
||||
import { useFileProcessor } from "@/lib/use-file-processor"
|
||||
|
||||
vi.mock("sonner", () => ({ toast: { error: vi.fn() } }))
|
||||
vi.mock("@/lib/pdf-utils", async (importOriginal) => ({
|
||||
...(await importOriginal<typeof import("@/lib/pdf-utils")>()),
|
||||
extractPdfText: vi.fn(),
|
||||
extractTextFileContent: vi.fn(),
|
||||
}))
|
||||
|
||||
// A promise we can resolve from the test, to control extraction timing
|
||||
function deferred<T>() {
|
||||
let resolve!: (value: T) => void
|
||||
const promise = new Promise<T>((r) => {
|
||||
resolve = r
|
||||
})
|
||||
return { promise, resolve }
|
||||
}
|
||||
|
||||
const pdfFile = () =>
|
||||
new File(["%PDF"], "slow.pdf", { type: "application/pdf" })
|
||||
const textFile = () => new File(["notes"], "notes.txt", { type: "text/plain" })
|
||||
|
||||
describe("useFileProcessor", () => {
|
||||
beforeEach(() => {
|
||||
vi.mocked(extractPdfText).mockReset()
|
||||
vi.mocked(extractTextFileContent).mockReset()
|
||||
})
|
||||
|
||||
it("marks queued files as extracting before the first one finishes", async () => {
|
||||
const pdf = deferred<string>()
|
||||
vi.mocked(extractPdfText).mockReturnValue(pdf.promise)
|
||||
vi.mocked(extractTextFileContent).mockResolvedValue("notes")
|
||||
const a = pdfFile()
|
||||
const b = textFile()
|
||||
const { result } = renderHook(() => useFileProcessor())
|
||||
|
||||
let done!: Promise<void>
|
||||
act(() => {
|
||||
done = result.current.handleFileChange([a, b])
|
||||
})
|
||||
|
||||
expect(result.current.pdfData.get(a)?.isExtracting).toBe(true)
|
||||
expect(result.current.pdfData.get(b)?.isExtracting).toBe(true)
|
||||
|
||||
await act(async () => {
|
||||
pdf.resolve("pdf text")
|
||||
await done
|
||||
})
|
||||
expect(result.current.pdfData.get(a)?.text).toBe("pdf text")
|
||||
expect(result.current.pdfData.get(b)?.text).toBe("notes")
|
||||
})
|
||||
|
||||
it("keeps text of a file added while an earlier file is extracting", async () => {
|
||||
const pdf = deferred<string>()
|
||||
vi.mocked(extractPdfText).mockReturnValue(pdf.promise)
|
||||
vi.mocked(extractTextFileContent).mockResolvedValue("notes")
|
||||
const a = pdfFile()
|
||||
const b = textFile()
|
||||
const { result } = renderHook(() => useFileProcessor())
|
||||
|
||||
let first!: Promise<void>
|
||||
act(() => {
|
||||
first = result.current.handleFileChange([a])
|
||||
})
|
||||
await act(async () => {
|
||||
await result.current.handleFileChange([a, b])
|
||||
})
|
||||
expect(result.current.pdfData.get(b)?.text).toBe("notes")
|
||||
|
||||
await act(async () => {
|
||||
pdf.resolve("pdf text")
|
||||
await first
|
||||
})
|
||||
expect(result.current.pdfData.get(a)?.text).toBe("pdf text")
|
||||
expect(result.current.pdfData.get(b)?.text).toBe("notes")
|
||||
})
|
||||
|
||||
it("does not bring back a file removed while extracting", async () => {
|
||||
const pdf = deferred<string>()
|
||||
vi.mocked(extractPdfText).mockReturnValue(pdf.promise)
|
||||
const a = pdfFile()
|
||||
const { result } = renderHook(() => useFileProcessor())
|
||||
|
||||
let first!: Promise<void>
|
||||
act(() => {
|
||||
first = result.current.handleFileChange([a])
|
||||
})
|
||||
await act(async () => {
|
||||
await result.current.handleFileChange([])
|
||||
})
|
||||
|
||||
await act(async () => {
|
||||
pdf.resolve("pdf text")
|
||||
await first
|
||||
})
|
||||
expect(result.current.pdfData.has(a)).toBe(false)
|
||||
expect(result.current.files).toEqual([])
|
||||
})
|
||||
})
|
||||
@@ -1,133 +0,0 @@
|
||||
import { act, cleanup, renderHook, waitFor } from "@testing-library/react"
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
|
||||
import { useModelConfig } from "@/hooks/use-model-config"
|
||||
import type { FlattenedServerModel } from "@/lib/server-model-config"
|
||||
import { STORAGE_KEYS } from "@/lib/storage"
|
||||
import type { MultiModelConfig } from "@/lib/types/model-config"
|
||||
|
||||
const SERVER_MODELS: FlattenedServerModel[] = [
|
||||
{
|
||||
id: "server:openai-main:gpt-4o-mini",
|
||||
modelId: "gpt-4o-mini",
|
||||
provider: "openai",
|
||||
providerLabel: "OpenAI Main",
|
||||
isDefault: false,
|
||||
},
|
||||
{
|
||||
id: "server:openai-main:gpt-4o",
|
||||
modelId: "gpt-4o",
|
||||
provider: "openai",
|
||||
providerLabel: "OpenAI Main",
|
||||
isDefault: true,
|
||||
},
|
||||
]
|
||||
|
||||
const USER_CONFIG: MultiModelConfig = {
|
||||
version: 1,
|
||||
providers: [
|
||||
{
|
||||
id: "p1",
|
||||
provider: "openai",
|
||||
apiKey: "sk-test",
|
||||
models: [{ id: "m1", modelId: "gpt-4o" }],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
function storeConfig(config: MultiModelConfig) {
|
||||
localStorage.setItem(STORAGE_KEYS.modelConfigs, JSON.stringify(config))
|
||||
}
|
||||
|
||||
async function renderLoaded() {
|
||||
const hook = renderHook(() => useModelConfig())
|
||||
await waitFor(() => expect(hook.result.current.isLoaded).toBe(true))
|
||||
return hook
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
localStorage.clear()
|
||||
vi.stubGlobal(
|
||||
"fetch",
|
||||
vi.fn(async () => ({
|
||||
ok: true,
|
||||
json: async () => ({ models: SERVER_MODELS }),
|
||||
})),
|
||||
)
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
cleanup()
|
||||
vi.unstubAllGlobals()
|
||||
})
|
||||
|
||||
describe("useModelConfig server model selection", () => {
|
||||
it("replaces a saved server model that no longer exists", async () => {
|
||||
storeConfig({
|
||||
...USER_CONFIG,
|
||||
selectedModelId: "server:openai-production:gpt-4o",
|
||||
})
|
||||
const { result } = await renderLoaded()
|
||||
expect(result.current.selectedModelId).toBe("server:openai-main:gpt-4o")
|
||||
})
|
||||
|
||||
it("keeps a saved server model that still exists", async () => {
|
||||
storeConfig({
|
||||
...USER_CONFIG,
|
||||
selectedModelId: "server:openai-main:gpt-4o-mini",
|
||||
})
|
||||
const { result } = await renderLoaded()
|
||||
expect(result.current.selectedModelId).toBe(
|
||||
"server:openai-main:gpt-4o-mini",
|
||||
)
|
||||
})
|
||||
|
||||
it("keeps a selected user model", async () => {
|
||||
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
|
||||
const { result } = await renderLoaded()
|
||||
expect(result.current.selectedModelId).toBe("m1")
|
||||
})
|
||||
|
||||
it("falls back to the default server model when the selected model is deleted", async () => {
|
||||
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
|
||||
const { result } = await renderLoaded()
|
||||
act(() => result.current.deleteModel("p1", "m1"))
|
||||
expect(result.current.selectedModelId).toBe("server:openai-main:gpt-4o")
|
||||
})
|
||||
|
||||
it("falls back to the default server model when the selected provider is deleted", async () => {
|
||||
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
|
||||
const { result } = await renderLoaded()
|
||||
act(() => result.current.deleteProvider("p1"))
|
||||
expect(result.current.selectedModelId).toBe("server:openai-main:gpt-4o")
|
||||
})
|
||||
})
|
||||
|
||||
describe("useModelConfig across tabs", () => {
|
||||
it("reloads the config when another tab saves it", async () => {
|
||||
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
|
||||
const { result } = await renderLoaded()
|
||||
|
||||
const fromOtherTab: MultiModelConfig = {
|
||||
...USER_CONFIG,
|
||||
providers: [
|
||||
...USER_CONFIG.providers,
|
||||
{
|
||||
id: "p2",
|
||||
provider: "anthropic",
|
||||
apiKey: "sk-ant",
|
||||
models: [{ id: "m2", modelId: "claude-sonnet-4-5" }],
|
||||
},
|
||||
],
|
||||
selectedModelId: "m2",
|
||||
}
|
||||
act(() => {
|
||||
storeConfig(fromOtherTab)
|
||||
window.dispatchEvent(
|
||||
new StorageEvent("storage", { key: STORAGE_KEYS.modelConfigs }),
|
||||
)
|
||||
})
|
||||
|
||||
expect(result.current.selectedModelId).toBe("m2")
|
||||
expect(result.current.config.providers).toHaveLength(2)
|
||||
})
|
||||
})
|
||||
+1
-185
@@ -1,14 +1,5 @@
|
||||
import { describe, expect, it } from "vitest"
|
||||
import {
|
||||
applyDiagramOperations,
|
||||
autoFixXml,
|
||||
cn,
|
||||
extractCompleteMxCells,
|
||||
isMxCellXmlComplete,
|
||||
validateAndFixXml,
|
||||
validateMxCellStructure,
|
||||
wrapWithMxFile,
|
||||
} from "@/lib/utils"
|
||||
import { cn, isMxCellXmlComplete, wrapWithMxFile } from "@/lib/utils"
|
||||
|
||||
describe("isMxCellXmlComplete", () => {
|
||||
it("returns false for empty/null input", () => {
|
||||
@@ -42,31 +33,6 @@ describe("isMxCellXmlComplete", () => {
|
||||
expect(isMxCellXmlComplete(xml)).toBe(false)
|
||||
})
|
||||
|
||||
it("returns false when output stops after a child of an open mxCell", () => {
|
||||
const xml = `<mxCell id="2" value="A" vertex="1" parent="1">
|
||||
<mxGeometry x="0" y="0" width="80" height="40" as="geometry"/>
|
||||
</mxCell>
|
||||
<mxCell id="3" value="B" vertex="1" parent="1">
|
||||
<mxGeometry x="100" y="0" width="80" height="40" as="geometry"/>`
|
||||
expect(isMxCellXmlComplete(xml)).toBe(false)
|
||||
})
|
||||
|
||||
it("returns false when output stops after </mxGeometry> of an open mxCell", () => {
|
||||
const xml = `<mxCell id="e1" edge="1" parent="1" source="2" target="3">
|
||||
<mxGeometry relative="1" as="geometry">
|
||||
<mxPoint x="10" y="10" as="sourcePoint"/>
|
||||
</mxGeometry>`
|
||||
expect(isMxCellXmlComplete(xml)).toBe(false)
|
||||
})
|
||||
|
||||
it("returns true for a self-closing last mxCell with > in its value", () => {
|
||||
const xml = `<mxCell id="2" value="A" vertex="1" parent="1">
|
||||
<mxGeometry as="geometry"/>
|
||||
</mxCell>
|
||||
<mxCell id="3" value="A -> B" vertex="1" parent="1"/></root>`
|
||||
expect(isMxCellXmlComplete(xml)).toBe(true)
|
||||
})
|
||||
|
||||
it("returns true for multiple complete mxCells", () => {
|
||||
const xml = `<mxCell id="2" value="A" vertex="1" parent="1"/>
|
||||
<mxCell id="3" value="B" vertex="1" parent="1"/>`
|
||||
@@ -118,153 +84,3 @@ describe("cn (class name utility)", () => {
|
||||
expect(cn("text-red-500", "text-blue-500")).toBe("text-blue-500")
|
||||
})
|
||||
})
|
||||
|
||||
describe("extractCompleteMxCells", () => {
|
||||
it("keeps the cell right after self-closing root cells", () => {
|
||||
const xml = `<mxfile><diagram id="p1"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="2" value="A" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="3" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel></diagram></mxfile>`
|
||||
const ids = [
|
||||
...extractCompleteMxCells(xml).matchAll(/<mxCell id="([^"]+)"/g),
|
||||
].map((m) => m[1])
|
||||
expect(ids).toEqual(["0", "1", "2", "3"])
|
||||
})
|
||||
|
||||
it("drops an incomplete trailing cell", () => {
|
||||
const xml = `<mxCell id="2" vertex="1" parent="1"/><mxCell id="3" vertex="1" parent="1"><mxGeometry as="geometry"/>`
|
||||
expect(extractCompleteMxCells(xml)).toBe(
|
||||
'<mxCell id="2" vertex="1" parent="1"/>',
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
const page = (id: string, cells: string) =>
|
||||
`<diagram name="${id}" id="${id}"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/>${cells}</root></mxGraphModel></diagram>`
|
||||
|
||||
describe("duplicate ids in multi-page documents", () => {
|
||||
const shape = (id: string, value = "Box") =>
|
||||
`<mxCell id="${id}" value="${value}" vertex="1" parent="1"><mxGeometry x="0" y="0" width="80" height="40" as="geometry"/></mxCell>`
|
||||
|
||||
it("accepts the same ids on different pages", () => {
|
||||
const xml = `<mxfile>${page("p1", shape("2"))}${page("p2", shape("2"))}</mxfile>`
|
||||
expect(validateMxCellStructure(xml)).toBeNull()
|
||||
})
|
||||
|
||||
it("still reports duplicate ids within one page", () => {
|
||||
const xml = `<mxfile>${page("p1", shape("2") + shape("2"))}${page("p2", "")}</mxfile>`
|
||||
expect(validateMxCellStructure(xml)).toContain("duplicate ID")
|
||||
})
|
||||
|
||||
it("does not rename the root cells of other pages when fixing", () => {
|
||||
const xml = `<mxfile>${page("p1", shape("2", "R&D"))}${page("p2", shape("3"))}</mxfile>`
|
||||
const result = validateAndFixXml(xml)
|
||||
expect(result.valid).toBe(true)
|
||||
expect(result.fixed).not.toContain("_dup")
|
||||
expect(result.fixed).toContain("R&D")
|
||||
})
|
||||
|
||||
it("renames a duplicate id within a page", () => {
|
||||
const xml = `<mxfile>${page("p1", shape("d") + shape("d"))}</mxfile>`
|
||||
const { fixed } = autoFixXml(xml)
|
||||
expect(fixed).toContain('<mxCell id="d" ')
|
||||
expect(fixed).toContain('<mxCell id="d_dup1" ')
|
||||
})
|
||||
})
|
||||
|
||||
describe("autoFixXml", () => {
|
||||
it("does not insert a space at the start of style values", () => {
|
||||
const xml = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="2" value="R&D" style="rounded=1;whiteSpace=wrap;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel>`
|
||||
const { fixed } = autoFixXml(xml)
|
||||
expect(fixed).toContain('style="rounded=1;whiteSpace=wrap;"')
|
||||
})
|
||||
|
||||
it("adds a missing space between attributes", () => {
|
||||
const xml = `<mxCell id="2" vertex="1"parent="1"/>`
|
||||
expect(autoFixXml(xml).fixed).toContain('vertex="1" parent="1"')
|
||||
})
|
||||
|
||||
it("keeps " inside rich text labels", () => {
|
||||
const label = "<font color="#ff0000">Hello</font>"
|
||||
const xml = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="2" value="${label}" style="html=1;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="3" value="Q&A" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel>`
|
||||
const result = validateAndFixXml(xml)
|
||||
expect(result.valid).toBe(true)
|
||||
expect(result.fixed).toContain(`value="${label}"`)
|
||||
})
|
||||
|
||||
it("fixes an attribute delimited by "", () => {
|
||||
const xml = `<mxCell id="2" dashPattern="1 1;" vertex="1" parent="1"/>`
|
||||
expect(autoFixXml(xml).fixed).toContain('dashPattern="1 1;"')
|
||||
})
|
||||
|
||||
it("keeps cells written on one line next to multi-line cells", () => {
|
||||
const xml = `<mxGraphModel><root>
|
||||
<mxCell id="0"/>
|
||||
<mxCell id="1" parent="0"/>
|
||||
<mxCell id="2" value="Q&A" vertex="1" parent="1">
|
||||
<mxGeometry x="0" y="0" width="80" height="40" as="geometry"/>
|
||||
</mxCell>
|
||||
<mxCell id="e1" edge="1" parent="1" source="2" target="3"><mxGeometry relative="1" as="geometry"/></mxCell>
|
||||
<mxCell id="3" value="B" vertex="1" parent="1">
|
||||
<mxGeometry x="200" y="0" width="80" height="40" as="geometry"/>
|
||||
</mxCell>
|
||||
</root></mxGraphModel>`
|
||||
const result = validateAndFixXml(xml)
|
||||
expect(result.valid).toBe(true)
|
||||
for (const id of ["2", "e1", "3"]) {
|
||||
expect(result.fixed).toContain(`<mxCell id="${id}"`)
|
||||
}
|
||||
})
|
||||
|
||||
it("keeps object and UserObject wrappers", () => {
|
||||
const xml = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><UserObject id="2" label="Docs" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><object id="3" label="A&B" owner="me"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></object></root></mxGraphModel>`
|
||||
const result = validateAndFixXml(xml)
|
||||
expect(result.valid).toBe(true)
|
||||
expect(result.fixed).toContain('<UserObject id="2"')
|
||||
expect(result.fixed).toContain('<object id="3"')
|
||||
})
|
||||
})
|
||||
|
||||
describe("applyDiagramOperations with wrapped cells", () => {
|
||||
const xml = `<mxfile><diagram id="p1"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><UserObject id="5" label="Docs" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><mxCell id="6" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="e1" edge="1" parent="1" source="5" target="6"><mxGeometry relative="1" as="geometry"/></mxCell></root></mxGraphModel></diagram></mxfile>`
|
||||
|
||||
it("deletes a wrapped cell and its edges", () => {
|
||||
const { result, errors } = applyDiagramOperations(xml, [
|
||||
{ operation: "delete", cell_id: "5" },
|
||||
{ operation: "delete", cell_id: "e1" },
|
||||
])
|
||||
expect(errors).toEqual([])
|
||||
expect(result).not.toContain("UserObject")
|
||||
expect(result).not.toContain('id="e1"')
|
||||
expect(result).toContain('id="6"')
|
||||
})
|
||||
|
||||
it("rejects adding a cell with the id of a wrapped cell", () => {
|
||||
const { errors } = applyDiagramOperations(xml, [
|
||||
{
|
||||
operation: "add",
|
||||
cell_id: "5",
|
||||
new_xml: '<mxCell id="5" vertex="1" parent="1"/>',
|
||||
},
|
||||
])
|
||||
expect(errors[0]?.message).toContain("already exists")
|
||||
})
|
||||
|
||||
it("updates a wrapped cell", () => {
|
||||
const { result, errors } = applyDiagramOperations(xml, [
|
||||
{
|
||||
operation: "update",
|
||||
cell_id: "5",
|
||||
new_xml:
|
||||
'<UserObject id="5" label="New" link="https://example.org"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject>',
|
||||
},
|
||||
])
|
||||
expect(errors).toEqual([])
|
||||
expect(result).toContain('label="New"')
|
||||
expect(result).not.toContain('label="Docs"')
|
||||
})
|
||||
|
||||
it("reports deleting a cell that does not exist", () => {
|
||||
const { errors } = applyDiagramOperations(xml, [
|
||||
{ operation: "delete", cell_id: "missing" },
|
||||
])
|
||||
expect(errors[0]?.message).toContain("not found")
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user