diff --git a/.github/CONTRIBUTING.md b/.github/CONTRIBUTING.md index edcb1184..83f65401 100644 --- a/.github/CONTRIBUTING.md +++ b/.github/CONTRIBUTING.md @@ -47,6 +47,10 @@ To run tests with UI mode: npx playwright test --ui ``` +## Before You Start + +For **significant changes** (new features, architecture changes, large refactors, etc.), please **open an issue first** to discuss your proposal before writing code. This helps avoid wasted effort and ensures alignment with the project direction. Small bug fixes and minor improvements can go straight to a PR. + ## Pull Requests 1. Create a feature branch @@ -57,6 +61,21 @@ npx playwright test --ui CI will run the full test suite on your PR. +## Using AI Tools + +AI-assisted contributions are welcome. But please **review the output before opening a PR**: + +1. **Review the code** — understand what was generated, don't just commit blindly +2. **Write a PR description** — explain what changed and why +3. **Rebase on latest `main`** — AI tools often work on stale branches, run `git rebase origin/main` before pushing +4. **Clean up artifacts** — remove IDE configs (`.idea/`, `.kiro/`), env files, scratch notes, and throwaway test scripts that AI tools leave behind + +## Code Review + +This project uses GitHub Copilot for automated code review. If you receive review comments from Copilot on your PR: +- **Valid suggestions**: Please address them in your code. +- **Invalid or irrelevant suggestions**: Feel free to click "Resolve" to dismiss them. + ## Issues Include steps to reproduce, expected vs actual behavior, and AI provider used. diff --git a/.github/renovate.json b/.github/renovate.json index 437ea5d6..6d7eb58b 100644 --- a/.github/renovate.json +++ b/.github/renovate.json @@ -33,6 +33,11 @@ "matchPackagePatterns": ["@ai-sdk/*", "ai", "next"], "groupName": "Core framework packages", "automerge": false + }, + { + "matchPackageNames": ["@biomejs/biome"], + "groupName": "Biome", + "automerge": false } ], "vulnerabilityAlerts": { diff --git a/.github/workflows/auto-format.yml b/.github/workflows/auto-format.yml index 03a1ca96..ada9eba8 100644 --- a/.github/workflows/auto-format.yml +++ b/.github/workflows/auto-format.yml @@ -23,7 +23,9 @@ jobs: node-version: '24' - name: Run Biome format - run: npx @biomejs/biome@latest check --write --no-errors-on-unmatched . + # Pin to the version in package.json so CI matches local/pre-commit + # (npx @latest drifts — e.g. 2.5.0 broke this job on unrelated PRs). + run: npx @biomejs/biome@2.5.7 check --write --no-errors-on-unmatched . - name: Check for changes id: changes diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d2eade46..0566d3f4 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -40,5 +40,3 @@ jobs: - name: Build run: npm run build - - name: Security audit - run: npm audit --audit-level=high --omit=dev diff --git a/.github/workflows/docker-build.yml b/.github/workflows/docker-build.yml index 59d90b63..66ac4e7d 100644 --- a/.github/workflows/docker-build.yml +++ b/.github/workflows/docker-build.yml @@ -58,6 +58,8 @@ jobs: with: context: . push: ${{ github.event_name != 'pull_request' }} + provenance: mode=max + sbom: true tags: ${{ steps.meta.outputs.tags }} labels: ${{ steps.meta.outputs.labels }} cache-from: type=gha @@ -89,4 +91,3 @@ jobs: docker pull ghcr.io/${REPO_LOWER}:latest docker tag ghcr.io/${REPO_LOWER}:latest ${{ secrets.AWS_ACCOUNT_ID }}.dkr.ecr.ap-northeast-1.amazonaws.com/next-ai-draw-io:latest docker push ${{ secrets.AWS_ACCOUNT_ID }}.dkr.ecr.ap-northeast-1.amazonaws.com/next-ai-draw-io:latest - diff --git a/.github/workflows/electron-release.yml b/.github/workflows/electron-release.yml index 9cb15ac6..9e8ef191 100644 --- a/.github/workflows/electron-release.yml +++ b/.github/workflows/electron-release.yml @@ -34,6 +34,15 @@ jobs: node-version: 24 cache: "npm" + - name: Download draw.io static files for offline use + run: | + rm -rf public/drawio + git clone --depth 1 https://github.com/jgraph/drawio.git /tmp/drawio + mkdir -p public/drawio + cp -r /tmp/drawio/src/main/webapp/* public/drawio/ + rm -rf public/drawio/WEB-INF + rm -rf public/drawio/META-INF + - name: Install dependencies run: npm install @@ -57,6 +66,16 @@ jobs: node-version: 24 cache: "npm" + - name: Download draw.io static files for offline use + shell: bash + run: | + rm -rf public/drawio + git clone --depth 1 https://github.com/jgraph/drawio.git /tmp/drawio + mkdir -p public/drawio + cp -r /tmp/drawio/src/main/webapp/* public/drawio/ + rm -rf public/drawio/WEB-INF + rm -rf public/drawio/META-INF + - name: Install dependencies run: npm install @@ -80,7 +99,7 @@ jobs: api-token: ${{ secrets.SIGNPATH_API_TOKEN }} organization-id: '880a211d-2cd3-4e7b-8d04-3d1f8eb39df5' project-slug: 'next-ai-draw-io' - signing-policy-slug: 'test-signing' + signing-policy-slug: 'release-signing' artifact-configuration-slug: 'windows-exe' github-artifact-id: ${{ steps.upload-unsigned.outputs.artifact-id }} wait-for-completion: true diff --git a/.github/workflows/publish-mcp.yml b/.github/workflows/publish-mcp.yml new file mode 100644 index 00000000..6b816c45 --- /dev/null +++ b/.github/workflows/publish-mcp.yml @@ -0,0 +1,71 @@ +name: Publish MCP Server + +# Publishes @next-ai-drawio/mcp-server to npm via OIDC trusted publishing +# (no token, no OTP). Triggers when packages/mcp-server changes on main; +# skips silently if the package.json version is already on npm — so a +# release is just "bump the version in a PR and merge". +on: + push: + branches: + - main + paths: + - "packages/mcp-server/**" + workflow_dispatch: + +permissions: + contents: read + id-token: write # OIDC token for npm trusted publishing + +concurrency: + group: publish-mcp + cancel-in-progress: false + +jobs: + publish: + runs-on: ubuntu-latest + defaults: + run: + working-directory: packages/mcp-server + steps: + - name: Checkout + uses: actions/checkout@v6 + + - name: Setup Node.js + uses: actions/setup-node@v6 + with: + node-version: 24 + cache: "npm" + cache-dependency-path: packages/mcp-server/package-lock.json + registry-url: "https://registry.npmjs.org" + + # Trusted publishing requires npm >= 11.5.1 + - name: Update npm + run: npm install -g npm@latest + + - name: Check if version is already published + id: version + run: | + LOCAL=$(node -p "require('./package.json').version") + if npm view "@next-ai-drawio/mcp-server@${LOCAL}" version >/dev/null 2>&1; then + echo "Version ${LOCAL} already on npm - nothing to publish" + echo "publish=false" >> "$GITHUB_OUTPUT" + else + echo "Version ${LOCAL} not on npm - publishing" + echo "publish=true" >> "$GITHUB_OUTPUT" + fi + + - name: Install dependencies + if: steps.version.outputs.publish == 'true' + run: npm ci + + - name: Test + if: steps.version.outputs.publish == 'true' + run: npm test + + - name: Build and check package contents + if: steps.version.outputs.publish == 'true' + run: npm run build && npm run check-package + + - name: Publish to npm + if: steps.version.outputs.publish == 'true' + run: npm publish diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 52428f18..bc8376c2 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -28,6 +28,20 @@ jobs: - name: Run unit tests run: npm run test -- --run + # The MCP server package ships its own vitest because its DOM polyfill + # (linkedom) needs `environment: node`, while the root vitest uses jsdom + # for the Next.js app. Install + run its tests separately so CI catches + # multi-page mxfile regressions. + - name: Install MCP server dependencies + run: npm --prefix packages/mcp-server ci + + - name: Run MCP server unit tests + run: npm --prefix packages/mcp-server test + + # Tests run from src/, so check the built npm package separately + - name: Build MCP server and check package contents + run: npm --prefix packages/mcp-server run build && npm --prefix packages/mcp-server run check-package + e2e: name: E2E Tests runs-on: ubuntu-latest diff --git a/.gitignore b/.gitignore index ef02da91..0e8010a0 100644 --- a/.gitignore +++ b/.gitignore @@ -56,6 +56,8 @@ push-via-ec2.sh /dist-electron/ /release/ /electron-standalone/ +# Draw.io static files (downloaded during CI build) +public/drawio/ *.dmg *.exe *.AppImage @@ -68,4 +70,12 @@ CLAUDE.md # edgeone .edgeone -opencode.json \ No newline at end of file +opencode.json +ai-models.json + +# local backups +*.bak +.gstack/ + +# admin panel settings (contains secrets) +data/ diff --git a/Dockerfile b/Dockerfile index a4af9806..037b48c1 100644 --- a/Dockerfile +++ b/Dockerfile @@ -9,6 +9,7 @@ WORKDIR /app COPY package.json package-lock.json* ./ # Install dependencies +ARG ELECTRON_SKIP_BINARY_DOWNLOAD=1 RUN npm install # Stage 2: Build application @@ -34,6 +35,11 @@ ENV NEXT_PUBLIC_SHOW_ABOUT_AND_NOTICE=${NEXT_PUBLIC_SHOW_ABOUT_AND_NOTICE} ARG NEXT_PUBLIC_BASE_PATH="" ENV NEXT_PUBLIC_BASE_PATH=${NEXT_PUBLIC_BASE_PATH} +# Control sponsorship and self-hosting messaging in quota notifications. +# Set NEXT_PUBLIC_SELFHOSTED="true" in self-hosted deployments to hide sponsorship/self-host links and related text in quota popups. +ARG NEXT_PUBLIC_SELFHOSTED="" +ENV NEXT_PUBLIC_SELFHOSTED="${NEXT_PUBLIC_SELFHOSTED}" + # Build Next.js application (standalone mode) RUN npm run build @@ -55,6 +61,9 @@ COPY --from=builder /app/public ./public COPY --from=builder --chown=nextjs:nodejs /app/.next/standalone ./ COPY --from=builder --chown=nextjs:nodejs /app/.next/static ./.next/static +# Writable dir for admin panel settings (data/settings.json) +RUN mkdir -p /app/data && chown nextjs:nodejs /app/data + USER nextjs EXPOSE 3000 diff --git a/README.md b/README.md index 7db249aa..5c94015a 100644 --- a/README.md +++ b/README.md @@ -19,7 +19,18 @@ English | [中文](./docs/cn/README_CN.md) | [日本語](./docs/ja/README_JA.md) A Next.js web application that integrates AI capabilities with draw.io diagrams. Create, modify, and enhance diagrams through natural language commands and AI-assisted visualization. -> Note: Thanks to [ByteDance Doubao](https://console.volcengine.com/ark/region:ark+cn-beijing/overview?briefPage=0&briefType=introduce&type=new&utm_campaign=doubao&utm_content=aidrawio&utm_medium=github&utm_source=coopensrc&utm_term=project) sponsorship, the demo site now uses the powerful K2-thinking model! +> Note: Thanks to [ByteDance Doubao](https://www.volcengine.com/activity/codingplan?ac=MMAP8JTTCAQ2&rc=Z9Z3LDTJ&utm_campaign=drawio&utm_content=drawio&utm_medium=devrel&utm_source=OWO&utm_term=drawio) sponsorship, the demo site now uses the powerful glm-4.7 model! + +

+ + + + Atlas Cloud + + +

+ +> 🎁 Thanks to **[Atlas Cloud](https://www.atlascloud.ai/?utm_source=github&utm_medium=link&utm_campaign=next-ai-draw-io)** for sponsoring next-ai-draw-io. Its OpenAI-compatible API gives diagram workflows one provider connection for DeepSeek, Qwen, GLM, Kimi, MiniMax, and more. Budget-friendly access is available through the [Coding Plan](https://www.atlascloud.ai/console/coding-plan). https://github.com/user-attachments/assets/9d60a3e8-4a1c-4b5e-acbb-26af2d3eabd1 @@ -31,7 +42,7 @@ https://github.com/user-attachments/assets/9d60a3e8-4a1c-4b5e-acbb-26af2d3eabd1 - [Table of Contents](#table-of-contents) - [Examples](#examples) - [Features](#features) - - [MCP Server (Preview)](#mcp-server-preview) + - [MCP Server](#mcp-server) - [Claude Code CLI](#claude-code-cli) - [Getting Started](#getting-started) - [Try it Online](#try-it-online) @@ -43,6 +54,8 @@ https://github.com/user-attachments/assets/9d60a3e8-4a1c-4b5e-acbb-26af2d3eabd1 - [Deploy on Vercel](#deploy-on-vercel) - [Deploy on Cloudflare Workers](#deploy-on-cloudflare-workers) - [Multi-Provider Support](#multi-provider-support) + - [Server-Side Multi-Model Configuration](#server-side-multi-model-configuration) + - [Admin Panel](#admin-panel) - [How It Works](#how-it-works) - [Support \& Contact](#support--contact) - [FAQ](#faq) @@ -63,24 +76,24 @@ Here are some example prompts and their generated diagrams: - GCP architecture diagram
-

Prompt: Generate a GCP architecture diagram with **GCP icons**. In this diagram, users connect to a frontend hosted on an instance.

- GCP Architecture Diagram + RAG Technique Diagram
+

Prompt: Generate a RAG architecture diagram for **chat application**. Use connected diagram for data ingestion

+ RAG Architecture Diagram - AWS architecture diagram
-

Prompt: Generate a AWS architecture diagram with **AWS icons**. In this diagram, users connect to a frontend hosted on an instance.

- AWS Architecture Diagram + Authentication using React and AWS
+

Prompt: Generate authentication process using React with **AWS**. Use Serverless architecture.

+ Authentication Architecture Diagram - Azure architecture diagram
-

Prompt: Generate a Azure architecture diagram with **Azure icons**. In this diagram, users connect to a frontend hosted on an instance.

- Azure Architecture Diagram + Open Innovation
+

Prompt: Create visualization of Henry Chesbrough's Open Innovation model.

+ Open Innovation Diagram - Cat sketch prompt
+ Cat sketch

Prompt: Draw a cute cat for me.

Cat Drawing @@ -99,9 +112,7 @@ Here are some example prompts and their generated diagrams: - **Cloud Architecture Diagram Support**: Specialized support for generating cloud architecture diagrams (AWS, GCP, Azure) - **Animated Connectors**: Create dynamic and animated connectors between diagram elements for better visualization -## MCP Server (Preview) - -> **Preview Feature**: This feature is experimental and may not be stable. +## MCP Server Use Next AI Draw.io with AI agents like Claude Desktop, Cursor, and VS Code via MCP (Model Context Protocol). @@ -127,6 +138,13 @@ Then ask Claude to create diagrams: The diagram appears in your browser in real-time! +The MCP server includes most of the web app's drawing features: + +- The same drawing rules and shape libraries (AWS, Azure, GCP, Kubernetes and more) +- A screenshot tool, so the AI can check the rendered diagram and fix it +- Version history, multi-page diagrams, and download as `.drawio`, `.png`, `.svg`, or `.drawio.svg` +- Auto-save to `~/.next-ai-drawio/`, so you can continue a diagram after a restart + See the [MCP Server README](./packages/mcp-server/README.md) for VS Code, Cursor, and other client configurations. ## Getting Started @@ -202,25 +220,38 @@ See the [Next.js deployment documentation](https://nextjs.org/docs/app/building- ## Multi-Provider Support -- [ByteDance Doubao](https://console.volcengine.com/ark/region:ark+cn-beijing/overview?briefPage=0&briefType=introduce&type=new&utm_campaign=doubao&utm_content=aidrawio&utm_medium=github&utm_source=coopensrc&utm_term=project) +- [ByteDance Doubao](https://www.volcengine.com/activity/codingplan?ac=MMAP8JTTCAQ2&rc=Z9Z3LDTJ&utm_campaign=drawio&utm_content=drawio&utm_medium=devrel&utm_source=OWO&utm_term=drawio) - AWS Bedrock (default) - OpenAI - Anthropic - Google AI +- Google Vertex AI - Azure OpenAI - Ollama - OpenRouter +- AIHubMix - DeepSeek - SiliconFlow - ModelScope - SGLang - Vercel AI Gateway +- [Atlas Cloud](https://www.atlascloud.ai/?utm_source=github&utm_medium=link&utm_campaign=next-ai-draw-io) All providers except AWS Bedrock and OpenRouter support custom endpoints. 📖 **[Detailed Provider Configuration Guide](./docs/en/ai-providers.md)** - See setup instructions for each provider. +### Server-Side Multi-Model Configuration + +Administrators can configure multiple server-side models that are available to all users without requiring personal API keys. Configure via `AI_MODELS_CONFIG` environment variable (JSON string) or `ai-models.json` file. For a single-provider quick setup, list comma-separated model IDs in `AI_MODEL`. + +### Admin Panel + +Set the `ADMIN_PASSWORD` environment variable and visit `/admin` to manage server settings (models, access codes, features, observability, quota) from a web panel instead of hand-editing `.env`. + +📖 **[Admin Panel Guide](./docs/en/admin-panel.md)** — setup, precedence rules, and notes. + **Model Requirements**: This task requires strong model capabilities for generating long-form text with strict formatting constraints (draw.io XML). Recommended models include Claude Sonnet 4.5, GPT-5.1, Gemini 3 Pro, and DeepSeek V3.2/R1. Note that the `claude` series has been trained on draw.io diagrams with cloud architecture logos like AWS, Azure, GCP. So if you want to create cloud architecture diagrams, this is the best choice. @@ -239,7 +270,9 @@ Diagrams are represented as XML that can be rendered in draw.io. The AI processe ## Support & Contact -**Special thanks to [ByteDance Doubao](https://console.volcengine.com/ark/region:ark+cn-beijing/overview?briefPage=0&briefType=introduce&type=new&utm_campaign=doubao&utm_content=aidrawio&utm_medium=github&utm_source=coopensrc&utm_term=project) for sponsoring the API token usage of the demo site!** Register on the ARK platform to get 500K free tokens for all models! +**Special thanks to [ByteDance Doubao](https://www.volcengine.com/activity/codingplan?ac=MMAP8JTTCAQ2&rc=Z9Z3LDTJ&utm_campaign=drawio&utm_content=drawio&utm_medium=devrel&utm_source=OWO&utm_term=drawio) for sponsoring the API token usage of the demo site!** Register on the ARK platform to get 500K free tokens for all models! + +**Special thanks to [Atlas Cloud](https://www.atlascloud.ai/?utm_source=github&utm_medium=link&utm_campaign=next-ai-draw-io) for sponsoring next-ai-draw-io and supporting its multi-provider ecosystem!** Try its OpenAI-compatible LLM API through the [Atlas Cloud Coding Plan](https://www.atlascloud.ai/console/coding-plan). If you find this project useful, please consider [sponsoring](https://github.com/sponsors/DayuanJiang) to help me host the live demo site! diff --git a/app/[lang]/about/cn/page.tsx b/app/[lang]/about/cn/page.tsx index e88a8add..bacd18e9 100644 --- a/app/[lang]/about/cn/page.tsx +++ b/app/[lang]/about/cn/page.tsx @@ -1,7 +1,7 @@ import type { Metadata } from "next" -import Image from "next/image" import Link from "next/link" import { FaGithub } from "react-icons/fa" +import Image from "@/components/image-with-basepath" export const metadata: Metadata = { title: "关于 - Next AI Draw.io", @@ -78,7 +78,7 @@ export default function AboutCN() {

好消息!感谢{" "} 的慷慨赞助,演示站点现已接入强大的{" "} - K2-thinking + glm-4.7 {" "} 模型,图表生成效果更佳!点击链接注册即可领取{" "} @@ -97,6 +97,23 @@ export default function AboutCN() {

+ {/* Invite Poster */} +
+ + 火山引擎方舟 Coding Plan + +
+ {/* Bring Your Own Key */}

@@ -158,92 +175,106 @@ export default function AboutCN() {

- {/* Animated Transformer */} + {/* ResNet50 Architecture */}

- 动画Transformer连接器 + ResNet50模型架构动画

- 提示词: 给我一个带有 - 动画连接器的Transformer架构图。 + Prompt: Give me an{" "} + animated architecture diagram + of the ResNet50 model.

- 带动画连接器的Transformer架构 +
+ ResNet50模型架构图 +
- {/* Cloud Architecture Grid */} + {/* Diagram Grid */}

- GCP架构图 + RAG技术图

- 提示词: 使用 - GCP图标 - 生成一个GCP架构图。用户连接到托管在实例上的前端。 + Prompt: Generate a RAG + architecture diagram for{" "} + chat application. Use + connected diagram for data ingestion

- GCP架构图 +
+ RAG架构图 +

- AWS架构图 + React和AWS认证流程

- 提示词: 使用 - AWS图标 - 生成一个AWS架构图。用户连接到托管在实例上的前端。 + Prompt: Generate + authentication process using React with{" "} + AWS. Use Serverless + architecture.

- AWS架构图 +
+ 认证架构图 +

- Azure架构图 + 敏捷Scrum流程

- 提示词: 使用 - Azure图标 - 生成一个Azure架构图。用户连接到托管在实例上的前端。 + Prompt: Generate agile + scrum workflow diagram for software + development team.

- Azure架构图 +
+ 敏捷Scrum流程图 +

- 猫咪素描 + 开放式创新

- 提示词:{" "} - 给我画一只可爱的猫。 + Prompt: Create + visualization of Henry Chesbrough's + Open Innovation model.

- 猫咪绘图 +
+ 开放式创新图 +
@@ -277,7 +308,7 @@ export default function AboutCN() {

@@ -199,7 +220,6 @@ export default function Home() { {/* Chat Panel */} p.isDefault && p.models.length > 0, + ) + return { + writable: isSettingsWritable(), + providers: maskAdminProviders(adminProviders), + envProviders: + envConfig?.providers.map((p) => ({ + name: p.name, + provider: p.provider, + 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"), + } +} + +export async function GET(req: Request) { + const authError = checkAdminAuth(req) + if (authError) return authError + return Response.json(await payload()) +} + +export async function PUT(req: Request) { + const authError = checkAdminAuth(req) + if (authError) return authError + + if (!isSettingsWritable()) { + return Response.json( + { + error: "Settings file is not writable on this deployment. Configure via environment variables instead.", + }, + { status: 503 }, + ) + } + + let body: unknown + try { + body = await req.json() + } catch { + return Response.json({ error: "Invalid JSON body" }, { status: 400 }) + } + + const parsed = AdminProvidersSchema.safeParse( + (body as { providers?: unknown })?.providers, + ) + if (!parsed.success) { + return Response.json( + { + error: `Invalid providers: ${parsed.error.issues[0]?.message ?? "schema mismatch"}`, + }, + { status: 400 }, + ) + } + + const stored = loadAdminProviders() + const merged = mergeSecrets(parsed.data, stored) + + const envConfig = await loadEnvServerModelsConfig() + const validationError = validateAdminProviders(merged, envConfig) + if (validationError) { + return Response.json({ error: validationError }, { status: 400 }) + } + + saveSettings(deriveEnvUpdates(merged, stored)) + + return Response.json(await payload()) +} diff --git a/app/api/admin/settings/route.ts b/app/api/admin/settings/route.ts new file mode 100644 index 00000000..d4d09755 --- /dev/null +++ b/app/api/admin/settings/route.ts @@ -0,0 +1,126 @@ +import { checkAdminAuth, maskSecret } from "@/lib/admin/auth" +import { + getEnvFallback, + getValueSource, + isSettingsWritable, + loadSettings, + saveSettings, +} from "@/lib/admin/settings" +import { + SETTINGS_BY_KEY, + SETTINGS_REGISTRY, + type SettingDef, +} from "@/lib/admin/settings-registry" + +export const runtime = "nodejs" +export const dynamic = "force-dynamic" + +function serializeSettings() { + const fileValues = loadSettings() + return SETTINGS_REGISTRY.map((def) => { + const source = getValueSource(def.key) + const raw = + source === "file" + ? fileValues[def.key] + : (getEnvFallback(def.key) ?? null) + const value = def.type === "secret" && raw ? maskSecret(raw) : raw + return { key: def.key, source, value } + }) +} + +export async function GET(req: Request) { + const authError = checkAdminAuth(req) + if (authError) return authError + + return Response.json({ + writable: isSettingsWritable(), + settings: serializeSettings(), + }) +} + +function validateValue(def: SettingDef, value: string): string | null { + switch (def.type) { + case "number": { + const num = Number(value) + if (!Number.isFinite(num)) return "Must be a number" + if (def.min !== undefined && num < def.min) + return `Must be at least ${def.min}` + if (def.max !== undefined && num > def.max) + return `Must be at most ${def.max}` + return null + } + case "boolean": + return value === "true" || value === "false" + ? null + : 'Must be "true" or "false"' + case "enum": + return def.options?.includes(value) + ? null + : `Must be one of: ${def.options?.join(", ")}` + default: + return null + } +} + +export async function PUT(req: Request) { + const authError = checkAdminAuth(req) + if (authError) return authError + + if (!isSettingsWritable()) { + return Response.json( + { + error: "Settings file is not writable on this deployment. Configure via environment variables instead.", + }, + { status: 503 }, + ) + } + + let body: { values?: Record } + try { + body = await req.json() + } catch { + return Response.json({ error: "Invalid JSON body" }, { status: 400 }) + } + if (!body.values || typeof body.values !== "object") { + return Response.json( + { error: "Body must contain a values object" }, + { status: 400 }, + ) + } + + const updates: Record = {} + const errors: Record = {} + + for (const [key, value] of Object.entries(body.values)) { + const def = SETTINGS_BY_KEY.get(key) + if (!def) { + errors[key] = "Unknown setting" + continue + } + if (value === null || value === "") { + updates[key] = null + continue + } + if (typeof value !== "string") { + errors[key] = "Value must be a string" + continue + } + const error = validateValue(def, value) + if (error) { + errors[key] = error + continue + } + updates[key] = value + } + + if (Object.keys(errors).length > 0) { + return Response.json({ errors }, { status: 400 }) + } + + saveSettings(updates) + + return Response.json({ + writable: true, + settings: serializeSettings(), + }) +} diff --git a/app/api/admin/test-model/route.ts b/app/api/admin/test-model/route.ts new file mode 100644 index 00000000..51999ca8 --- /dev/null +++ b/app/api/admin/test-model/route.ts @@ -0,0 +1,84 @@ +import { POST as validateModel } from "@/app/api/validate-model/route" +import { checkAdminAuth } from "@/lib/admin/auth" +import { + AdminProviderSchema, + loadAdminProviders, + mergeSecrets, +} from "@/lib/admin/providers" +import { globalBaseUrl } from "@/lib/ai-providers" + +export const runtime = "nodejs" +export const dynamic = "force-dynamic" + +// Test a model with the client's CURRENT provider state (which may be +// unsaved). Secret fields arrive either as plaintext (newly typed) or as +// masked {isSet} markers, which are resolved against settings.json — so +// testing works both before and after saving. +export async function POST(req: Request) { + const authError = checkAdminAuth(req) + if (authError) return authError + + let body: { provider?: unknown; modelId?: string } + try { + body = await req.json() + } catch { + return Response.json({ error: "Invalid JSON body" }, { status: 400 }) + } + + const parsed = AdminProviderSchema.safeParse(body.provider) + if (!parsed.success || !body.modelId) { + return Response.json( + { valid: false, error: "Invalid provider or model" }, + { status: 400 }, + ) + } + + // SECURITY: a stored secret is only resolved from an {isSet} marker if + // the endpoint it would be sent to (provider + baseUrl) still matches + // the stored entry. Otherwise a tampered baseUrl could exfiltrate the + // stored key to an arbitrary host. Mismatches must re-supply plaintext. + const stored = loadAdminProviders().find((p) => p.id === parsed.data.id) + const sameEndpoint = + stored && + stored.provider === parsed.data.provider && + (stored.baseUrl ?? "") === (parsed.data.baseUrl ?? "") && + (stored.awsRegion ?? "") === (parsed.data.awsRegion ?? "") + const [resolved] = mergeSecrets( + [parsed.data], + sameEndpoint && stored ? [stored] : [], + ) + + const serverUrl = globalBaseUrl(resolved.provider) + return validateModel( + new Request(new URL("/api/validate-model", req.url), { + method: "POST", + headers: { + "Content-Type": "application/json", + // Checked again there, in place of an access code + "x-admin-password": req.headers.get("x-admin-password") || "", + // The EdgeOne function checks the access code and Pages + // cookies, and its URL is built from the page's origin + "x-access-code": req.headers.get("x-access-code") || "", + cookie: req.headers.get("cookie") || "", + ...(req.headers.get("origin") && { + origin: req.headers.get("origin") as string, + }), + }, + body: JSON.stringify({ + provider: resolved.provider, + apiKey: resolved.apiKey, + // Without a URL of its own, chat sends the entry's key to + // the server's

_BASE_URL: test that endpoint, not + // another one. It is the server's own, which chat uses + // without the checks for a URL a user typed. + baseUrl: resolved.baseUrl || serverUrl, + ...(!resolved.baseUrl && serverUrl && { serverBaseUrl: true }), + modelId: body.modelId, + awsAccessKeyId: resolved.awsAccessKeyId, + awsSecretAccessKey: resolved.awsSecretAccessKey, + awsRegion: resolved.awsRegion, + vertexApiKey: resolved.vertexApiKey, + }), + }), + ) +} diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts index 872b8702..c5ea47a2 100644 --- a/app/api/chat/route.ts +++ b/app/api/chat/route.ts @@ -4,40 +4,66 @@ import { createUIMessageStream, createUIMessageStreamResponse, InvalidToolInputError, - LoadAPIKeyError, stepCountIs, streamText, } from "ai" -import fs from "fs/promises" import { jsonrepair } from "jsonrepair" import path from "path" import { z } from "zod" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" import { + CACHE_POINT, + edgeOneEndpoint, getAIModel, - supportsImageInput, + getServerProvider, + SINGLE_SYSTEM_PROVIDERS, supportsPromptCaching, + usesServerCredentials, + usesServerEndpoint, } from "@/lib/ai-providers" import { findCachedResponse } from "@/lib/cached-responses" import { - isMinimalDiagram, + dropInvalidToolCalls, + fixToolInputJson, replaceHistoricalToolInputs, validateFileParts, } from "@/lib/chat-helpers" +import { withDeprecatedParamsFallback } from "@/lib/deprecated-params" import { checkAndIncrementRequest, isQuotaEnabled, recordTokenUsage, } from "@/lib/dynamo-quota-manager" import { + endTrace, getTelemetryConfig, setTraceInput, setTraceOutput, wrapWithObserve, } from "@/lib/langfuse" +import { classifyLLMError, streamErrorText } from "@/lib/llm-errors" +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 { getSystemPrompt } from "@/lib/system-prompts" +import { normalizeBaseUrl } from "@/lib/types/model-config" import { getUserIdFromRequest } from "@/lib/user-id" +import { hasCells } from "@/packages/mcp-server/src/pages.ts" +import { + getShapeLibrary, + SHAPE_LIBRARY_LIST, +} from "@/packages/mcp-server/src/shape-library.ts" +import { SWIMLANE_EXAMPLE } from "@/packages/mcp-server/src/xml-examples.ts" -export const maxDuration = 120 +// No explicit cap: a reasoning model can spend minutes planning before it emits +// the tool call, so take whatever the host allows. Vercel's own default is 300s, +// which is also where Node's response-body timeout on the upstream stream lands. // Helper function to create cached stream response function createCachedStreamResponse(xml: string): Response { @@ -69,26 +95,25 @@ function createCachedStreamResponse(xml: string): Response { return createUIMessageStreamResponse({ stream }) } -// Inner handler function -async function handleChatRequest(req: Request): Promise { - // Check for access code - const accessCodes = - process.env.ACCESS_CODE_LIST?.split(",") - .map((code) => code.trim()) - .filter(Boolean) || [] - if (accessCodes.length > 0) { - const accessCodeHeader = req.headers.get("x-access-code") - if (!accessCodeHeader || !accessCodes.includes(accessCodeHeader)) { - return Response.json( - { - error: "Invalid or missing access code. Please configure it in Settings.", - }, - { status: 401 }, - ) - } - } +// Responses streamed from the model, whose trace streamText's callbacks end +const modelStreamResponses = new WeakSet() - const { messages, xml, previousXml, sessionId } = await req.json() +// Inner handler function +const DEBUG_LLM_PAYLOAD = process.env.DEBUG_LLM_PAYLOAD === "true" + +async function handleChatRequest(req: Request): Promise { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite + // Check for access code + const accessDenied = checkAccessCode(req) + if (accessDenied) return accessDenied + + const body = await req.json() + const { messages, xml, previousXml, sessionId } = body + const customSystemMessage = + typeof body.customSystemMessage === "string" + ? body.customSystemMessage.slice(0, 5000) + : "" // Get user ID for Langfuse tracking and quota const userId = getUserIdFromRequest(req) @@ -114,14 +139,165 @@ async function handleChatRequest(req: Request): Promise { userId: userId, }) - // === SERVER-SIDE QUOTA CHECK START === - // Quota is opt-in: only enabled when DYNAMODB_QUOTA_TABLE env var is set - const hasOwnApiKey = !!( - req.headers.get("x-ai-provider") && req.headers.get("x-ai-api-key") + // === FILE VALIDATION START === + const fileValidation = validateFileParts(messages) + if (!fileValidation.valid) { + return Response.json({ error: fileValidation.error }, { status: 400 }) + } + // === FILE VALIDATION END === + + // === CACHE CHECK START === + const isFirstMessage = messages.length === 1 + const isEmptyDiagram = !xml || !hasCells(xml) + + if (isFirstMessage && isEmptyDiagram) { + const lastMessage = messages[0] + const textPart = lastMessage.parts?.find((p: any) => p.type === "text") + const filePart = lastMessage.parts?.find((p: any) => p.type === "file") + + const cached = findCachedResponse(textPart?.text || "", !!filePart) + + if (cached) { + return createCachedStreamResponse(cached.xml) + } + } + // === CACHE CHECK END === + + // Read client AI provider overrides from headers + const provider = req.headers.get("x-ai-provider") + let baseUrl = req.headers.get("x-ai-base-url") + const selectedModelId = req.headers.get("x-selected-model-id") + + // Check if this is a server model with custom env var names + let serverModelConfig: { + apiKeyEnv?: string | string[] + baseUrlEnv?: string + provider?: string + } = {} + let serverModel: FlattenedServerModel | null = null + if (selectedModelId?.startsWith("server:")) { + serverModel = await findServerModelById(selectedModelId) + console.log( + `[Server Model Lookup] ID: ${selectedModelId}, Found: ${!!serverModel}, Provider: ${serverModel?.provider}`, + ) + if (serverModel) { + serverModelConfig = { + apiKeyEnv: serverModel.apiKeyEnv, + baseUrlEnv: serverModel.baseUrlEnv, + // Use actual provider from config (client header may have incorrect value due to ID format change) + provider: serverModel.provider, + } + } + } + + // A server model's provider comes from its config: for one set up in + // the admin panel the header holds the provider name's slug. Without + // either, the server's own AI_PROVIDER. + const isEdgeOne = + (serverModelConfig.provider || provider || getServerProvider()) === + "edgeone" + + // EdgeOne is this deployment's own function, whatever URL the request + // names: another host would get the user's EdgeOne cookies, and the + // quota counts it. Absolute, as the SDK needs. + if (isEdgeOne) baseUrl = edgeOneEndpoint(req) + + // Same rule as validate-model: with ALLOW_PRIVATE_URLS=false a request may + // not point the server at a private or internal address + if (baseUrl && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) { + return Response.json( + { error: "Private or internal base URLs are not allowed." }, + { status: 400 }, + ) + } + + // Get cookie header for EdgeOne authentication (eo_token, eo_time) + const cookieHeader = req.headers.get("cookie") + + const clientOverrides = { + // Server model provider takes precedence over client header; EdgeOne + // named only in AI_PROVIDER is named here, for its own base URL + provider: + serverModelConfig.provider || + provider || + (isEdgeOne ? "edgeone" : null), + 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"), + // AWS Bedrock credentials + awsAccessKeyId: req.headers.get("x-aws-access-key-id"), + awsSecretAccessKey: req.headers.get("x-aws-secret-access-key"), + awsRegion: req.headers.get("x-aws-region"), + awsSessionToken: req.headers.get("x-aws-session-token"), + // Server model custom env var names + ...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 + ...(isEdgeOne && { + headers: { + ...(cookieHeader && { cookie: cookieHeader }), + "x-access-code": req.headers.get("x-access-code") || "", + }, + }), + } + + // Read minimal style preference from header + const minimalStyle = req.headers.get("x-minimal-style") === "true" + + console.log( + `[Client Overrides] provider: ${clientOverrides.provider}, modelId: ${clientOverrides.modelId}`, ) - // Skip quota check if: quota disabled, user has own API key, or is anonymous - if (isQuotaEnabled() && !hasOwnApiKey && userId !== "anonymous") { + // Get AI model with optional client overrides + const { + model: baseModel, + providerOptions, + modelId, + 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 + // on AI_PROVIDER. With their own key, users can run any model. + const onServerCredentials = usesServerCredentials( + resolvedProvider, + clientOverrides, + ) + const envModels = + process.env.AI_MODEL?.split(",").map((m) => m.trim()) || [] + const offeredInEnv = + envModels.includes(modelId) && resolvedProvider === getServerProvider() + if (onServerCredentials && !serverModel && !offeredInEnv) { + return Response.json( + { + error: `Model "${modelId}" is not available on this server. Add your own API key in Settings to use it.`, + }, + { status: 400 }, + ) + } + + // === SERVER-SIDE QUOTA CHECK START === + // Quota is opt-in (DYNAMODB_QUOTA_TABLE) and counts what runs on the + // server's keys, or on the server's own endpoints: EdgeOne, its keyless + // Ollama, and anything at a private address (the server's network, + // which ignores a dummy key header). Bedrock and EdgeOne never use the + // base URL header. In the desktop app every endpoint is the user's. + const clientBaseUrl = normalizeBaseUrl( + req.headers.get("x-ai-base-url") ?? "", + ) + const onServerEndpoint = await usesServerEndpoint( + resolvedProvider, + clientBaseUrl, + clientOverrides.apiKey, + ) + const countsQuota = + isQuotaEnabled() && + (onServerCredentials || onServerEndpoint) && + userId !== "anonymous" + if (countsQuota) { const quotaCheck = await checkAndIncrementRequest(userId, { requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10, tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000, @@ -141,67 +317,20 @@ async function handleChatRequest(req: Request): Promise { } // === SERVER-SIDE QUOTA CHECK END === - // === FILE VALIDATION START === - const fileValidation = validateFileParts(messages) - if (!fileValidation.valid) { - return Response.json({ error: fileValidation.error }, { status: 400 }) - } - // === FILE VALIDATION END === + // Retry once if the provider rejects the requested budget, or (newer + // Claude models) the sampling or thinking settings + const model = withOutputTokenLimitFallback( + withDeprecatedParamsFallback(baseModel), + ) - // === CACHE CHECK START === - const isFirstMessage = messages.length === 1 - const isEmptyDiagram = !xml || xml.trim() === "" || isMinimalDiagram(xml) - - if (isFirstMessage && isEmptyDiagram) { - const lastMessage = messages[0] - const textPart = lastMessage.parts?.find((p: any) => p.type === "text") - const filePart = lastMessage.parts?.find((p: any) => p.type === "file") - - const cached = findCachedResponse(textPart?.text || "", !!filePart) - - if (cached) { - return createCachedStreamResponse(cached.xml) - } - } - // === CACHE CHECK END === - - // Read client AI provider overrides from headers - const provider = req.headers.get("x-ai-provider") - let baseUrl = req.headers.get("x-ai-base-url") - - // For EdgeOne provider, construct full URL from request origin - // because createOpenAI needs absolute URL, not relative path - if (provider === "edgeone" && !baseUrl) { - const origin = req.headers.get("origin") || new URL(req.url).origin - baseUrl = `${origin}/api/edgeai` - } - - // Get cookie header for EdgeOne authentication (eo_token, eo_time) - const cookieHeader = req.headers.get("cookie") - - const clientOverrides = { - provider, - baseUrl, - apiKey: req.headers.get("x-ai-api-key"), - 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"), - awsRegion: req.headers.get("x-aws-region"), - awsSessionToken: req.headers.get("x-aws-session-token"), - // Pass cookies for EdgeOne Pages authentication - ...(provider === "edgeone" && - cookieHeader && { - headers: { cookie: cookieHeader }, - }), - } - - // Read minimal style preference from header - const minimalStyle = req.headers.get("x-minimal-style") === "true" - - // Get AI model with optional client overrides - const { model, providerOptions, headers, modelId } = - getAIModel(clientOverrides) + // The user setting can raise the budget only on their own key (in the + // desktop app every key is the user's); on the server's keys or own + // endpoints it can only lower it + const maxOutputTokens = resolveMaxOutputTokens( + req.headers.get("x-max-output-tokens"), + onServerCredentials || onServerEndpoint, + ) + console.log(`[maxOutputTokens] ${maxOutputTokens}`) // Check if model supports prompt caching const shouldCache = supportsPromptCaching(modelId) @@ -211,22 +340,19 @@ async function handleChatRequest(req: Request): Promise { // Get the appropriate system prompt based on model (extended for Opus/Haiku 4.5) const systemMessage = getSystemPrompt(modelId, minimalStyle) + const finalSystemMessage = customSystemMessage + ? `${systemMessage}\n\n## Custom Instructions\n${customSystemMessage}` + : systemMessage // Extract file parts (images) from the last user message const fileParts = lastUserMessage?.parts?.filter((part: any) => part.type === "file") || [] - // Check if user is sending images to a model that doesn't support them - // AI SDK silently drops unsupported parts, so we need to catch this early - if (fileParts.length > 0 && !supportsImageInput(modelId)) { - return Response.json( - { - error: `The model "${modelId}" does not support image input. Please use a vision-capable model (e.g., GPT-4o, Claude, Gemini) or remove the image.`, - }, - { status: 400 }, - ) - } + // Note: we used to pre-emptively reject images for models we guessed were + // text-only (by name matching). That heuristic misfired on newer models + // (see issue #874), so we now let the request through and surface the real + // provider error if the model genuinely can't accept images. // User input only - XML is now in a separate cached system message const formattedUserInput = `User input: @@ -234,39 +360,46 @@ async function handleChatRequest(req: Request): Promise { ${userInputText} """` - // Convert UIMessages to ModelMessages and add system message - const modelMessages = await convertToModelMessages(messages) - - // DEBUG: Log incoming messages structure - console.log("[route.ts] Incoming messages count:", messages.length) - messages.forEach((msg: any, idx: number) => { - console.log( - `[route.ts] Message ${idx} role:`, - msg.role, - "parts count:", - msg.parts?.length, - ) - if (msg.parts) { - msg.parts.forEach((part: any, partIdx: number) => { - if ( - part.type === "tool-invocation" || - part.type === "tool-result" - ) { - console.log(`[route.ts] Part ${partIdx}:`, { - type: part.type, - toolName: part.toolName, - hasInput: !!part.input, - inputType: typeof part.input, - inputKeys: - part.input && typeof part.input === "object" - ? Object.keys(part.input) - : null, - }) - } - }) - } + // Convert UIMessages to ModelMessages and add system message. A tool + // call that never got its result (the user stopped while it ran) is + // left out: the SDK would refuse this and every later request of the + // chat (MissingToolResultsError) + const modelMessages = await convertToModelMessages(messages, { + ignoreIncompleteToolCalls: true, }) + // DEBUG_LLM_PAYLOAD=true logs the incoming message structure + if (DEBUG_LLM_PAYLOAD) { + console.log("[route.ts] Incoming messages count:", messages.length) + messages.forEach((msg: any, idx: number) => { + console.log( + `[route.ts] Message ${idx} role:`, + msg.role, + "parts count:", + msg.parts?.length, + ) + if (msg.parts) { + msg.parts.forEach((part: any, partIdx: number) => { + if ( + part.type === "tool-invocation" || + part.type === "tool-result" + ) { + console.log(`[route.ts] Part ${partIdx}:`, { + type: part.type, + toolName: part.toolName, + hasInput: !!part.input, + inputType: typeof part.input, + inputKeys: + part.input && typeof part.input === "object" + ? Object.keys(part.input) + : null, + }) + } + }) + } + }) + } + // Replace historical tool call XML with placeholders to reduce tokens // Disabled by default - some models (e.g. minimax) copy placeholders instead of generating XML const enableHistoryReplace = @@ -283,61 +416,43 @@ ${userInputText} ) // Filter out tool-calls with invalid inputs (from failed repair or interrupted streaming) - // Bedrock API rejects messages where toolUse.input is not a valid JSON object - enhancedMessages = enhancedMessages - .map((msg: any) => { - if (msg.role !== "assistant" || !Array.isArray(msg.content)) { - return msg - } - const filteredContent = msg.content.filter((part: any) => { - if (part.type === "tool-call") { - // Check if input is a valid object (not null, undefined, or empty) - if ( - !part.input || - typeof part.input !== "object" || - Object.keys(part.input).length === 0 - ) { - console.warn( - `[route.ts] Filtering out tool-call with invalid input:`, - { toolName: part.toolName, input: part.input }, - ) - return false - } - } - return true - }) - return { ...msg, content: filteredContent } - }) - .filter((msg: any) => msg.content && msg.content.length > 0) + // and their results. Bedrock API rejects messages where toolUse.input is not a valid + // JSON object, and every provider rejects a tool result whose call is gone. + enhancedMessages = dropInvalidToolCalls(enhancedMessages) - // DEBUG: Log modelMessages structure (what's being sent to AI) - console.log("[route.ts] Model messages count:", enhancedMessages.length) - enhancedMessages.forEach((msg: any, idx: number) => { - console.log( - `[route.ts] ModelMsg ${idx} role:`, - msg.role, - "content count:", - msg.content?.length, - ) - if (msg.content) { - msg.content.forEach((part: any, partIdx: number) => { - if (part.type === "tool-call" || part.type === "tool-result") { - console.log(`[route.ts] Content ${partIdx}:`, { - type: part.type, - toolName: part.toolName, - hasInput: !!part.input, - inputType: typeof part.input, - inputValue: - part.input === undefined - ? "undefined" - : part.input === null - ? "null" - : "object", - }) - } - }) - } - }) + // DEBUG_LLM_PAYLOAD=true logs what is sent to the model + if (DEBUG_LLM_PAYLOAD) { + console.log("[route.ts] Model messages count:", enhancedMessages.length) + enhancedMessages.forEach((msg: any, idx: number) => { + console.log( + `[route.ts] ModelMsg ${idx} role:`, + msg.role, + "content count:", + msg.content?.length, + ) + if (msg.content) { + msg.content.forEach((part: any, partIdx: number) => { + if ( + part.type === "tool-call" || + part.type === "tool-result" + ) { + console.log(`[route.ts] Content ${partIdx}:`, { + type: part.type, + toolName: part.toolName, + hasInput: !!part.input, + inputType: typeof part.input, + inputValue: + part.input === undefined + ? "undefined" + : part.input === null + ? "null" + : "object", + }) + } + }) + } + }) + } // Update the last message with user input only (XML moved to separate cached system message) if (enhancedMessages.length >= 1) { @@ -353,7 +468,7 @@ ${userInputText} contentParts.push({ type: "image", image: filePart.url, - mimeType: filePart.mediaType, + mediaType: filePart.mediaType, }) } @@ -373,9 +488,7 @@ ${userInputText} if (enhancedMessages[i].role === "assistant") { enhancedMessages[i] = { ...enhancedMessages[i], - providerOptions: { - bedrock: { cachePoint: { type: "default" } }, - }, + providerOptions: CACHE_POINT, } break // Only cache the last assistant message } @@ -383,40 +496,75 @@ ${userInputText} } // System messages with multiple cache breakpoints for optimal caching: - // - Breakpoint 1: Static instructions (~1500 tokens) - rarely changes + // - Breakpoint 1: System instructions + custom instructions - changes when user updates custom system message // - Breakpoint 2: Current XML context - changes per diagram, but constant within a conversation turn - // This allows: if only user message changes, both system caches are reused - // if XML changes, instruction cache is still reused - const systemMessages = [ - // Cache breakpoint 1: Instructions (rarely change) - { - role: "system" as const, - content: systemMessage, - ...(shouldCache && { - providerOptions: { - bedrock: { cachePoint: { type: "default" } }, - }, - }), - }, - // Cache breakpoint 2: Previous and Current diagram XML context - { - role: "system" as const, - content: `${previousXml ? `Previous diagram XML (before user's last message):\n"""xml\n${previousXml}\n"""\n\n` : ""}Current diagram XML (AUTHORITATIVE - the source of truth):\n"""xml\n${xml || ""}\n"""\n\nIMPORTANT: 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!`, - ...(shouldCache && { - providerOptions: { - bedrock: { cachePoint: { type: "default" } }, - }, - }), - }, - ] + // Some providers (e.g. MiniMax) don't support multiple system messages + // Merge them into a single system message for compatibility + // Also merge for OpenAI-compatible providers with custom base URLs (e.g. vLLM, LMStudio) + // because open-source model chat templates (Qwen, Llama, etc.) typically reject multiple system messages + const isCustomOpenAIEndpoint = + resolvedProvider === "openai" && + !!( + baseUrl || + process.env.OPENAI_BASE_URL || + (serverModelConfig.baseUrlEnv && + process.env[serverModelConfig.baseUrlEnv]) + ) + const isSingleSystemProvider = + SINGLE_SYSTEM_PROVIDERS.has(resolvedProvider) || isCustomOpenAIEndpoint + + const xmlContext = `${ + previousXml + ? `Previous diagram XML (before user's last message): +"""xml +${previousXml} +""" + +` + : "" + }Current diagram XML (AUTHORITATIVE - the source of truth): +"""xml +${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.` + + const systemMessages = isSingleSystemProvider + ? [ + { + role: "system" as const, + content: `${finalSystemMessage}\n\n${xmlContext}`, + }, + ] + : [ + // Cache breakpoint 1: Instructions (+ optional custom instructions) + { + role: "system" as const, + content: finalSystemMessage, + ...(shouldCache && { providerOptions: CACHE_POINT }), + }, + // Cache breakpoint 2: Previous and Current diagram XML context + { + role: "system" as const, + content: xmlContext, + ...(shouldCache && { providerOptions: CACHE_POINT }), + }, + ] const allMessages = [...systemMessages, ...enhancedMessages] + // Set by onAbort, which records the finished steps' tokens itself + let stopped = false const result = streamText({ model, - ...(process.env.MAX_OUTPUT_TOKENS && { - maxOutputTokens: parseInt(process.env.MAX_OUTPUT_TOKENS, 10), - }), + // The system messages carry cache points, so they go in messages. + // A client's own system messages have string content and were + // dropped by the empty-content filter above. + allowSystemInMessages: true, + abortSignal: req.signal, + // Must be sent: unset means the provider's own default, and Bedrock's is + // 4096, enough for a small diagram, so larger ones were cut off mid-attribute. + maxOutputTokens, stopWhen: stepCountIs(5), // Repair truncated tool calls when maxOutputTokens is reached mid-JSON experimental_repairToolCall: async ({ toolCall, error }) => { @@ -434,16 +582,11 @@ ${userInputText} error.name === "AI_InvalidToolInputError" ) { try { - // Pre-process to fix common LLM JSON errors that jsonrepair can't handle - let inputToRepair = toolCall.input - if (typeof inputToRepair === "string") { - // Fix `:=` instead of `: ` (LLM sometimes generates this) - inputToRepair = inputToRepair.replace(/:=/g, ": ") - // Fix `= "` instead of `: "` - inputToRepair = inputToRepair.replace(/=\s*"/g, ': "') - } - // Use jsonrepair to fix truncated JSON - const repairedInput = jsonrepair(inputToRepair) + // Pre-process to fix common LLM JSON errors that jsonrepair can't handle, + // then use jsonrepair to fix truncated JSON + const repairedInput = jsonrepair( + fixToolInputJson(toolCall.input), + ) console.log( `[repairToolCall] Repaired truncated JSON for tool: ${toolCall.toolName}`, ) @@ -453,26 +596,8 @@ ${userInputText} `[repairToolCall] Failed to repair JSON for tool: ${toolCall.toolName}`, repairError, ) - // Return a placeholder input to avoid API errors in multi-step - // The tool will fail gracefully on client side - if (toolCall.toolName === "edit_diagram") { - return { - ...toolCall, - input: { - operations: [], - _error: "JSON repair failed - no operations to apply", - }, - } - } - if (toolCall.toolName === "display_diagram") { - return { - ...toolCall, - input: { - xml: "", - _error: "JSON repair failed - empty diagram", - }, - } - } + // Keep the original error, so the model and the client see why + // the input was rejected and the model can retry the call return null } } @@ -481,7 +606,6 @@ ${userInputText} }, messages: allMessages, ...(providerOptions && { providerOptions }), // This now includes all reasoning configs - ...(headers && { headers }), // Langfuse telemetry config (returns undefined if not configured) ...(getTelemetryConfig({ sessionId: validSessionId, userId }) && { experimental_telemetry: getTelemetryConfig({ @@ -495,21 +619,36 @@ ${userInputText} // Record token usage for server-side quota tracking (if enabled) // Use totalUsage (cumulative across all steps) instead of usage (final step only) - // Include all 4 token types: input, output, cache read, cache write - if ( - isQuotaEnabled() && - !hasOwnApiKey && - userId !== "anonymous" && - totalUsage - ) { + // inputTokens already includes cache reads and writes in AI SDK 6 + if (countsQuota && totalUsage && !stopped) { const totalTokens = (totalUsage.inputTokens || 0) + - (totalUsage.outputTokens || 0) + - (totalUsage.cachedInputTokens || 0) + - (totalUsage.inputTokenDetails?.cacheWriteTokens || 0) + (totalUsage.outputTokens || 0) recordTokenUsage(userId, totalTokens) } }, + // onFinish is skipped when the stream fails or is aborted, so end the trace here + onError: ({ error }) => { + console.error(error) // what AI SDK does without an onError + endTrace() + }, + onAbort: ({ steps }) => { + stopped = true + endTrace() + // Stopped (or disconnected) after some steps finished: their + // tokens were used, or stopping every request after a costly + // first step would get around the token limits + if (countsQuota) { + const tokens = steps.reduce( + (sum, step) => + sum + + (step.usage.inputTokens || 0) + + (step.usage.outputTokens || 0), + 0, + ) + if (tokens > 0) recordTokenUsage(userId, tokens) + } + }, tools: { // Client-side tool that will be executed on the client display_diagram: { @@ -524,21 +663,7 @@ VALIDATION RULES (XML will be rejected if violated): 6. Escape special chars in values: < > & " Example (generate ONLY this - no wrapper tags): - - - - - - - - - - - - - - - +${SWIMLANE_EXAMPLE} Notes: - For AWS diagrams, use **AWS 2025 icons**. @@ -616,14 +741,7 @@ Example: If previous output ended with ' streamErrorText(error, onServerCredentials), messageMetadata: ({ part }) => { if (part.type === "finish") { const usage = (part as any).totalUsage @@ -695,63 +784,28 @@ 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 +// Errors before the stream starts, as JSON the chat panel reads function handleError(error: unknown): Response { console.error("Error in chat route:", error) const isDev = process.env.NODE_ENV === "development" - - // Check for specific AI SDK error types - if (APICallError.isInstance(error)) { - return Response.json( - { - error: error.message, - ...(isDev && { - details: error.responseBody, - stack: error.stack, - }), - }, - { status: error.statusCode || 500 }, - ) - } - - if (LoadAPIKeyError.isInstance(error)) { - return Response.json( - { - error: "Authentication failed. Please check your API key.", - ...(isDev && { - stack: error.stack, - }), - }, - { status: 401 }, - ) - } - - // Fallback for other errors with safety filter - const message = - error instanceof Error ? error.message : "An unexpected error occurred" - const status = (error as any)?.statusCode || (error as any)?.status || 500 - - // Prevent leaking API keys, tokens, or other sensitive data - const lowerMessage = message.toLowerCase() - const safeMessage = - lowerMessage.includes("key") || - lowerMessage.includes("token") || - lowerMessage.includes("sig") || - lowerMessage.includes("signature") || - lowerMessage.includes("secret") || - lowerMessage.includes("password") || - lowerMessage.includes("credential") - ? "Authentication failed. Please check your credentials." - : message + const classified = classifyLLMError(error) + const status = + (error as { statusCode?: number })?.statusCode || + (error as { status?: number })?.status || + (classified.code === "invalid_api_key" ? 401 : 500) return Response.json( { - error: safeMessage, + ...classified, ...(isDev && { - details: message, + details: APICallError.isInstance(error) + ? error.responseBody + : undefined, stack: error instanceof Error ? error.stack : undefined, }), }, @@ -761,11 +815,16 @@ function handleError(error: unknown): Response { // Wrap handler with error handling async function safeHandler(req: Request): Promise { + let response: Response try { - return await handleChatRequest(req) + response = await handleChatRequest(req) } catch (error) { - return handleError(error) + response = handleError(error) } + // Early returns, cache hits and errors never reach streamText's callbacks, + // so their Langfuse trace has to be ended here + if (!modelStreamResponses.has(response)) endTrace() + return response } // Wrap with Langfuse observe (if configured) diff --git a/app/api/log-save/route.ts b/app/api/log-save/route.ts index fc73fb2b..eb30e0fe 100644 --- a/app/api/log-save/route.ts +++ b/app/api/log-save/route.ts @@ -4,7 +4,7 @@ import { getLangfuseClient } from "@/lib/langfuse" const saveSchema = z.object({ filename: z.string().min(1).max(255), - format: z.enum(["drawio", "png", "svg"]), + format: z.enum(["drawio", "png", "svg", "xmlsvg"]), sessionId: z.string().min(1).max(200).optional(), }) diff --git a/app/api/parse-url/route.ts b/app/api/parse-url/route.ts index f5278e65..33a15c4b 100644 --- a/app/api/parse-url/route.ts +++ b/app/api/parse-url/route.ts @@ -1,62 +1,46 @@ -import { extract } from "@extractus/article-extractor" +import { extractFromHtml } from "@extractus/article-extractor" import { NextResponse } from "next/server" import TurndownService from "turndown" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" +import { readLimitedBody } from "@/lib/read-limited-body" +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)" -// SSRF protection - block private/internal addresses -function isPrivateUrl(urlString: string): boolean { +// Detect the page's charset so non-UTF-8 pages (Shift_JIS/GBK/EUC/Big5, common +// on CJK sites) are decoded correctly. Response.text() always assumes UTF-8 and +// would produce mojibake; the article-extractor library does the same detection +// when it fetches the page itself, which we no longer rely on. +function detectCharset( + contentType: string | null, + buffer: ArrayBuffer, +): string { + // 1. HTTP Content-Type header charset (most authoritative). + const headerCharset = contentType?.match(/charset=([^;]+)/i)?.[1]?.trim() + // 2. / in the first bytes of the document. + const head = new TextDecoder("utf-8").decode(buffer.slice(0, 4096)) + const metaCharset = + head.match(/]+charset=["']?\s*([\w-]+)/i)?.[1] || + head.match(/]+content=["'][^"']*charset=([\w-]+)/i)?.[1] + const charset = (headerCharset || metaCharset || "utf-8").toLowerCase() + // TextDecoder throws on unknown encoding labels; fall back to UTF-8. try { - const url = new URL(urlString) - const hostname = url.hostname.toLowerCase() - - // Block localhost - if ( - hostname === "localhost" || - hostname === "127.0.0.1" || - hostname === "::1" - ) { - return true - } - - // Block AWS/cloud metadata endpoints - if ( - hostname === "169.254.169.254" || - hostname === "metadata.google.internal" - ) { - return true - } - - // Check for private IPv4 ranges - const ipv4Match = hostname.match( - /^(\d{1,3})\.(\d{1,3})\.(\d{1,3})\.(\d{1,3})$/, - ) - if (ipv4Match) { - const [, a, b] = ipv4Match.map(Number) - if (a === 10) return true // 10.0.0.0/8 - if (a === 172 && b >= 16 && b <= 31) return true // 172.16.0.0/12 - if (a === 192 && b === 168) return true // 192.168.0.0/16 - if (a === 169 && b === 254) return true // 169.254.0.0/16 (link-local) - if (a === 127) return true // 127.0.0.0/8 (loopback) - } - - // Block common internal hostnames - if ( - hostname.endsWith(".local") || - hostname.endsWith(".internal") || - hostname.endsWith(".localhost") - ) { - return true - } - - return false + new TextDecoder(charset) + return charset } catch { - return true // Invalid URL - block it + return "utf-8" } } export async function POST(req: Request) { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite + const accessError = checkAccessCode(req) + if (accessError) return accessError + try { const { url } = await req.json() @@ -77,28 +61,61 @@ export async function POST(req: Request) { ) } - // SSRF protection - if (isPrivateUrl(url)) { + // SSRF protection: parse-url has no use case for fetching internal + // hosts, so private URLs are always rejected. ALLOW_PRIVATE_URLS only + // governs LLM provider baseUrl overrides (validate-model, chat). + if (await isPrivateUrl(url)) { return NextResponse.json( { error: "Cannot access private/internal URLs" }, { status: 400 }, ) } - - // Extract article content with timeout to avoid tying up server resources + // Fetch the page ourselves so we control redirect handling. The + // article-extractor library follows redirects internally and ignores a + // `redirect` option, which would let a public URL 302 to an internal + // host and bypass the SSRF check above. `redirect: "error"` rejects any + // redirect outright. const controller = new AbortController() const timeoutId = setTimeout(() => { controller.abort() }, EXTRACT_TIMEOUT_MS) - let article + let html: string try { - article = await extract(url, undefined, { - headers: { - "User-Agent": "Mozilla/5.0 (compatible; NextAIDrawio/1.0)", - }, + const response = await fetch(url, { + headers: { "User-Agent": USER_AGENT }, + redirect: "error", signal: controller.signal, }) + + const contentType = response.headers.get("content-type") + if (contentType?.includes("application/pdf")) { + return NextResponse.json( + { + error: "PDF URLs are not supported. Please download and upload the PDF file directly", + }, + { status: 422 }, + ) + } + + if (!response.ok) { + return NextResponse.json( + { error: "Could not fetch URL content" }, + { status: 400 }, + ) + } + + const buffer = await readLimitedBody(response, MAX_RESPONSE_BYTES) + if (!buffer) { + return NextResponse.json( + { + error: `Page exceeds the ${MAX_RESPONSE_BYTES / 1024 / 1024} MB download limit`, + }, + { status: 413 }, + ) + } + const charset = detectCharset(contentType, buffer) + html = new TextDecoder(charset).decode(buffer) } catch (err: any) { if (err?.name === "AbortError") { return NextResponse.json( @@ -106,9 +123,26 @@ export async function POST(req: Request) { { status: 504 }, ) } - throw err + // Redirects are rejected with a TypeError ("failed to fetch" / + // "unexpected redirect") when redirect: "error" is set. + return NextResponse.json( + { error: "Could not fetch URL content" }, + { status: 400 }, + ) } finally { clearTimeout(timeoutId) + // Ends a download left unread (too large, PDF, error status); + // a body already read is not affected + controller.abort() + } + + // extractFromHtml throws (not returns null) on empty/non-HTML bodies, + // so map any parse error to the same 400 as the no-content case. + let article: Awaited> + try { + article = await extractFromHtml(html, url) + } catch { + article = null } if (!article || !article.content) { diff --git a/app/api/provider-models/route.ts b/app/api/provider-models/route.ts new file mode 100644 index 00000000..6b362e73 --- /dev/null +++ b/app/api/provider-models/route.ts @@ -0,0 +1,79 @@ +import { NextResponse } from "next/server" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" +import { classifyLLMError } from "@/lib/llm-errors" +import { + canListModels, + listProviderModels, + ModelListError, +} from "@/lib/provider-models" +import { + allowPrivateUrls, + isPrivateUrl, + RedirectRefusedError, + redirectGuardedFetch, +} from "@/lib/ssrf-protection" +import type { ProviderName } from "@/lib/types/model-config" + +export const runtime = "nodejs" + +// Public lists need no key +const NO_KEY_NEEDED = new Set([ + "ollama", + "openrouter", + "aihubmix", +]) + +/** + * The models a provider offers, for the "Fetch models" button in model + * settings. Answers { models: null } for providers that cannot list them, + * so the dialog keeps its suggested models. + */ +export async function POST(req: Request) { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite + // Sends requests to a URL the client chose, so require the access code + const accessError = checkAccessCode(req) + if (accessError) return accessError + + const { provider, apiKey, baseUrl } = (await req.json()) as { + provider: ProviderName + apiKey?: string + baseUrl?: string + } + if (!canListModels(provider)) { + return NextResponse.json({ models: null }) + } + // SECURITY: Block SSRF attacks via custom baseUrl + if (baseUrl && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) { + return NextResponse.json({ error: "Invalid base URL" }, { status: 400 }) + } + if (!apiKey && !NO_KEY_NEEDED.has(provider)) { + return NextResponse.json( + { error: "API key is required" }, + { status: 400 }, + ) + } + + try { + const models = await listProviderModels( + provider, + { apiKey, baseUrl }, + (baseUrl && redirectGuardedFetch()) || fetch, + ) + return NextResponse.json({ models }) + } catch (error) { + console.warn("[provider-models] Listing failed:", error) + // Only our own explanations go back: the URL may be an internal + // address, whose answer or host names must not reach the caller. + // The Gateway SDK wraps them, keeping ours as the cause. + const isOwn = (e: unknown): e is Error => + e instanceof ModelListError || e instanceof RedirectRefusedError + const cause = (error as { cause?: unknown })?.cause + const own = isOwn(error) ? error : isOwn(cause) ? cause : null + const { code } = classifyLLMError(own ?? error) + return NextResponse.json({ + code, + error: own?.message ?? "The model list request failed.", + }) + } +} diff --git a/app/api/server-models/route.ts b/app/api/server-models/route.ts new file mode 100644 index 00000000..49ea12bb --- /dev/null +++ b/app/api/server-models/route.ts @@ -0,0 +1,14 @@ +import { NextResponse } from "next/server" +import { loadFlattenedServerModels } from "@/lib/server-model-config" + +// Use dynamic rendering to read AI_MODEL/AI_PROVIDER env vars at runtime +// This ensures Docker users can set these values when starting containers +export const dynamic = "force-dynamic" + +export async function GET() { + const models = await loadFlattenedServerModels() + return NextResponse.json({ + models, + hasConfig: models.length > 0, + }) +} diff --git a/app/api/validate-diagram/route.ts b/app/api/validate-diagram/route.ts new file mode 100644 index 00000000..3fba6628 --- /dev/null +++ b/app/api/validate-diagram/route.ts @@ -0,0 +1,184 @@ +/** + * API endpoint for VLM-based diagram validation. + * Accepts a PNG image and streams validation results using useObject-compatible format. + */ + +import { Output, streamText } from "ai" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" +import { getValidationModel } from "@/lib/ai-providers" +import { + checkAndIncrementRequest, + isQuotaEnabled, + recordTokenUsage, +} from "@/lib/dynamo-quota-manager" +import { getUserIdFromRequest } from "@/lib/user-id" +import { VALIDATION_SYSTEM_PROMPT } from "@/lib/validation-prompts" +import { + type ValidationResult, + ValidationResultSchema, +} from "@/lib/validation-schema" + +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 +} + +// Default valid result for disabled/error cases +const DEFAULT_VALID_RESULT: ValidationResult = { + valid: true, + issues: [], + suggestions: [], +} + +/** A fixed result in the text format useObject reads */ +function createStreamingResponse(result: ValidationResult): Response { + return new Response(JSON.stringify(result), { + headers: { "Content-Type": "text/plain; charset=utf-8" }, + }) +} + +export async function POST(req: Request): Promise { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite + // 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" + if (!enableValidation) { + return createStreamingResponse(DEFAULT_VALID_RESULT) + } + + const body: ValidateDiagramRequest = await req.json() + const { imageData, sessionId } = body + + if (!imageData) { + return Response.json( + { error: "Missing imageData" }, + { status: 400 }, + ) + } + + // Validate image data format + if ( + !imageData.startsWith("data:image/png;base64,") && + !imageData.startsWith("data:image/") + ) { + return Response.json( + { error: "Invalid image data format" }, + { status: 400 }, + ) + } + + if (imageData.length > MAX_IMAGE_DATA_LENGTH) { + return Response.json( + { error: "Image data too large" }, + { status: 413 }, + ) + } + + // It runs the server's vision model: with the quota on, the daily + // and per-minute token limits apply, and its tokens are counted. Not + // the request limit, which is for chats: the day's last chat still + // gets its check, and a check does not count as a chat. + const userId = getUserIdFromRequest(req) + const countsQuota = isQuotaEnabled() && userId !== "anonymous" + if (countsQuota) { + const quotaCheck = await checkAndIncrementRequest( + userId, + { + requests: 0, + tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000, + tpm: Number(process.env.TPM_LIMIT) || 20000, + }, + 0, + ) + if (!quotaCheck.allowed) { + return Response.json( + { + error: quotaCheck.error, + type: quotaCheck.type, + used: quotaCheck.used, + limit: quotaCheck.limit, + }, + { status: 429 }, + ) + } + } + + // Get the validation model + let model + try { + model = getValidationModel() + } catch (error) { + console.warn( + "[validate-diagram] Validation model not available:", + error, + ) + // Return valid if no vision model is configured + return createStreamingResponse(DEFAULT_VALID_RESULT) + } + + // Parse timeout with validation (minimum 1000ms, default 10000ms) + const timeout = + Math.max( + 1000, + parseInt(process.env.VALIDATION_TIMEOUT || "10000", 10), + ) || 10000 + + // Stream the VLM response for useObject consumption + const result = streamText({ + model, + output: Output.object({ schema: ValidationResultSchema }), + system: VALIDATION_SYSTEM_PROMPT, + messages: [ + { + role: "user", + content: [ + { + type: "image", + image: imageData, + }, + { + type: "text", + text: "Please analyze this diagram for visual quality issues.", + }, + ], + }, + ], + maxOutputTokens: 1024, + abortSignal: AbortSignal.timeout(timeout), + onFinish: ({ output, totalUsage }) => { + if (countsQuota && totalUsage) { + recordTokenUsage( + userId, + (totalUsage.inputTokens || 0) + + (totalUsage.outputTokens || 0), + ) + } + if (sessionId && output) { + console.log( + `[validate-diagram] Session ${sessionId}: valid=${output.valid}, issues=${output.issues?.length ?? 0}`, + ) + } + }, + }) + + return result.toTextStreamResponse() + } catch (error) { + // Log with session context if available + const errorMessage = + error instanceof Error ? error.message : String(error) + console.error("[validate-diagram] Error:", errorMessage) + + // On error, return valid to not block the user + return createStreamingResponse(DEFAULT_VALID_RESULT) + } +} diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts index b8b258e6..1109a976 100644 --- a/app/api/validate-model/route.ts +++ b/app/api/validate-model/route.ts @@ -1,78 +1,28 @@ -import { createAmazonBedrock } from "@ai-sdk/amazon-bedrock" -import { createAnthropic } from "@ai-sdk/anthropic" -import { createDeepSeek, deepseek } from "@ai-sdk/deepseek" -import { createGateway } from "@ai-sdk/gateway" -import { createGoogleGenerativeAI } from "@ai-sdk/google" -import { createOpenAI } from "@ai-sdk/openai" -import { createOpenRouter } from "@openrouter/ai-sdk-provider" -import { generateText } from "ai" +import { streamText, tool } from "ai" import { NextResponse } from "next/server" -import { createOllama } from "ollama-ai-provider-v2" +import { z } from "zod" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" +import { checkAdminAuth } from "@/lib/admin/auth" +import { + edgeOneEndpoint, + getAIModel, + globalBaseUrl, + usesServerCredentials, + usesServerEndpoint, +} from "@/lib/ai-providers" +import { + checkAndIncrementRequest, + isQuotaEnabled, +} from "@/lib/dynamo-quota-manager" +import { classifyLLMError } from "@/lib/llm-errors" +import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection" +import { normalizeBaseUrl, type ProviderName } from "@/lib/types/model-config" +import { getUserIdFromRequest } from "@/lib/user-id" export const runtime = "nodejs" -/** - * SECURITY: Check if URL points to private/internal network (SSRF protection) - * Blocks: localhost, private IPs, link-local, AWS metadata service - */ -function isPrivateUrl(urlString: string): boolean { - try { - const url = new URL(urlString) - const hostname = url.hostname.toLowerCase() - - // Block localhost - if ( - hostname === "localhost" || - hostname === "127.0.0.1" || - hostname === "::1" - ) { - return true - } - - // Block AWS/cloud metadata endpoints - if ( - hostname === "169.254.169.254" || - hostname === "metadata.google.internal" - ) { - return true - } - - // Check for private IPv4 ranges - const ipv4Match = hostname.match( - /^(\d{1,3})\.(\d{1,3})\.(\d{1,3})\.(\d{1,3})$/, - ) - if (ipv4Match) { - const [, a, b] = ipv4Match.map(Number) - // 10.0.0.0/8 - if (a === 10) return true - // 172.16.0.0/12 - if (a === 172 && b >= 16 && b <= 31) return true - // 192.168.0.0/16 - if (a === 192 && b === 168) return true - // 169.254.0.0/16 (link-local) - if (a === 169 && b === 254) return true - // 127.0.0.0/8 (loopback) - if (a === 127) return true - } - - // Block common internal hostnames - if ( - hostname.endsWith(".local") || - hostname.endsWith(".internal") || - hostname.endsWith(".localhost") - ) { - return true - } - - return false - } catch { - // Invalid URL - block it - return true - } -} - interface ValidateRequest { - provider: string + provider: ProviderName apiKey: string baseUrl?: string modelId: string @@ -80,19 +30,44 @@ interface ValidateRequest { awsAccessKeyId?: string awsSecretAccessKey?: string awsRegion?: string + awsSessionToken?: string + // Vertex AI specific + vertexApiKey?: string // Express Mode API key + // Set by the admin panel's Test: baseUrl is the server's

_BASE_URL + serverBaseUrl?: boolean } +const TEST_TIMEOUT_MS = 15_000 + +// Drawing works through tool calls, so the test asks for one +const PING_TOOL = tool({ + description: "Report that the connection works.", + inputSchema: z.object({}), +}) + +const NO_TOOL_CALL_WARNING = + "Connected, but the model answered without calling a tool. It may not support tool calls, which drawing needs." + export async function POST(req: Request) { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite + // Lets the server send requests to arbitrary URLs, so require the access + // code, or the admin password (the admin panel's Test button) + const accessError = checkAccessCode(req) + if (accessError && checkAdminAuth(req)) return accessError + try { const body: ValidateRequest = await req.json() const { provider, apiKey, - baseUrl, modelId, awsAccessKeyId, awsSecretAccessKey, awsRegion, + awsSessionToken, + // Note: Express Mode only needs vertexApiKey + vertexApiKey, } = body if (!provider || !modelId) { @@ -101,9 +76,26 @@ export async function POST(req: Request) { { status: 400 }, ) } + // EdgeOne is this site's own function, as in the chat; the admin + // panel's Test sends no URL, and a relative one cannot be fetched + const baseUrl = + provider === "edgeone" ? edgeOneEndpoint(req) : body.baseUrl + // The admin panel's Test of an entry without a URL sends the + // server's own

_BASE_URL, which chat uses as it is: not a URL a + // user chose, so no private-address or redirect rules + const serverUrl = + body.serverBaseUrl === true && + !!baseUrl && + baseUrl === globalBaseUrl(provider) && + !checkAdminAuth(req) // SECURITY: Block SSRF attacks via custom baseUrl - if (baseUrl && isPrivateUrl(baseUrl)) { + if ( + baseUrl && + !serverUrl && + !allowPrivateUrls() && + (await isPrivateUrl(baseUrl)) + ) { return NextResponse.json( { valid: false, error: "Invalid base URL" }, { status: 400 }, @@ -121,278 +113,133 @@ export async function POST(req: Request) { { status: 400 }, ) } + } else if (provider === "vertexai") { + if (!vertexApiKey) { + return NextResponse.json( + { + valid: false, + error: "Vertex AI API key is required for Express Mode", + }, + { status: 400 }, + ) + } } else if (provider !== "ollama" && provider !== "edgeone" && !apiKey) { return NextResponse.json( { valid: false, error: "API key is required" }, { status: 400 }, ) } - - let model: any - - switch (provider) { - case "openai": { - const openai = createOpenAI({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - }) - model = openai.chat(modelId) - break - } - - case "anthropic": { - const anthropic = createAnthropic({ - apiKey, - baseURL: baseUrl || "https://api.anthropic.com/v1", - }) - model = anthropic(modelId) - break - } - - case "google": { - const google = createGoogleGenerativeAI({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - }) - model = google(modelId) - break - } - - case "azure": { - const azure = createOpenAI({ - apiKey, - baseURL: baseUrl, - }) - model = azure.chat(modelId) - break - } - - case "bedrock": { - const bedrock = createAmazonBedrock({ - accessKeyId: awsAccessKeyId, - secretAccessKey: awsSecretAccessKey, - region: awsRegion, - }) - model = bedrock(modelId) - break - } - - case "openrouter": { - const openrouter = createOpenRouter({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - }) - model = openrouter(modelId) - break - } - - case "deepseek": { - if (baseUrl || apiKey) { - const ds = createDeepSeek({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - }) - model = ds(modelId) - } else { - model = deepseek(modelId) - } - break - } - - case "siliconflow": { - const sf = createOpenAI({ - apiKey, - baseURL: baseUrl || "https://api.siliconflow.cn/v1", - }) - model = sf.chat(modelId) - break - } - - case "ollama": { - const ollama = createOllama({ - baseURL: baseUrl || "http://localhost:11434", - }) - model = ollama(modelId) - break - } - - case "gateway": { - const gw = createGateway({ - apiKey, - ...(baseUrl && { baseURL: baseUrl }), - }) - model = gw(modelId) - break - } - - case "edgeone": { - // EdgeOne uses OpenAI-compatible API via Edge Functions - // 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", - headers: { - cookie: cookieHeader, - }, - }) - model = edgeone.chat(modelId) - break - } - - case "sglang": { - // SGLang is OpenAI-compatible - const sglang = createOpenAI({ - apiKey: apiKey || "not-needed", - baseURL: baseUrl || "http://127.0.0.1:8000/v1", - }) - model = sglang.chat(modelId) - break - } - - case "doubao": { - // ByteDance Doubao: use DeepSeek for DeepSeek/Kimi models, OpenAI for others - const doubaoBaseUrl = - baseUrl || "https://ark.cn-beijing.volces.com/api/v3" - const lowerModelId = modelId.toLowerCase() - if ( - lowerModelId.includes("deepseek") || - lowerModelId.includes("kimi") - ) { - const doubao = createDeepSeek({ - apiKey, - baseURL: doubaoBaseUrl, - }) - model = doubao(modelId) - } else { - const doubao = createOpenAI({ - apiKey, - baseURL: doubaoBaseUrl, - }) - model = doubao.chat(modelId) - } - break - } - - case "modelscope": { - const baseURL = - baseUrl || "https://api-inference.modelscope.cn/v1" - const startTime = Date.now() - - try { - // Initiate a streaming request (required for QwQ-32B and certain Qwen3 models) - const response = await fetch( - `${baseURL}/chat/completions`, - { - method: "POST", - headers: { - "Content-Type": "application/json", - Authorization: `Bearer ${apiKey}`, - }, - body: JSON.stringify({ - model: modelId, - messages: [ - { role: "user", content: "Say 'OK'" }, - ], - max_tokens: 20, - stream: true, - enable_thinking: false, - }), - }, - ) - - if (!response.ok) { - const errorText = await response.text() - throw new Error( - `ModelScope API error (${response.status}): ${errorText}`, - ) - } - - const contentType = - response.headers.get("content-type") || "" - const isValidStreamingResponse = - response.status === 200 && - (contentType.includes("text/event-stream") || - contentType.includes("application/json")) - - if (!isValidStreamingResponse) { - throw new Error( - `Unexpected response format: ${contentType}`, - ) - } - - const responseTime = Date.now() - startTime - - if (response.body) { - response.body.cancel().catch(() => { - /* Ignore cancellation errors */ - }) - } - - return NextResponse.json({ - valid: true, - responseTime, - note: "ModelScope model validated (using streaming API)", - }) - } catch (error) { - console.error( - "[validate-model] ModelScope validation failed:", - error, - ) - throw error - } - } - - default: - return NextResponse.json( - { valid: false, error: `Unknown provider: ${provider}` }, - { status: 400 }, - ) + // The Test button checks the user's own provider. On the server's + // keys (Ollama Cloud without a key or URL) anyone could run any model. + if ( + usesServerCredentials(provider, { + apiKey, + baseUrl, + awsAccessKeyId, + awsSecretAccessKey, + vertexApiKey, + }) + ) { + return NextResponse.json( + { valid: false, error: "API key is required" }, + { status: 400 }, + ) } - // Make a minimal test request - const startTime = Date.now() - await generateText({ - model, - prompt: "Say 'OK'", - maxOutputTokens: 20, + // On the deployment's own endpoints a Test runs a model as a chat + // does, so with the quota on it counts as a chat request (an + // admin's Test of the server's URL does not) + const userId = getUserIdFromRequest(req) + if ( + isQuotaEnabled() && + !serverUrl && + userId !== "anonymous" && + (await usesServerEndpoint( + provider, + normalizeBaseUrl(body.baseUrl ?? ""), + apiKey, + )) + ) { + const quotaCheck = await checkAndIncrementRequest(userId, { + requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10, + tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000, + tpm: Number(process.env.TPM_LIMIT) || 20000, + }) + if (!quotaCheck.allowed) { + return NextResponse.json( + { valid: false, error: quotaCheck.error }, + { status: 429 }, + ) + } + } + + // The same model the chat would use. A client base URL makes it + // refuse redirects to internal hosts. + const { model } = getAIModel({ + provider, + modelId, + apiKey, + baseUrl, + trustedBaseUrl: serverUrl, + awsAccessKeyId, + awsSecretAccessKey, + awsRegion, + // Temporary AWS credentials need it, as in the chat + awsSessionToken, + vertexApiKey, + // EdgeOne checks the Pages cookies and the access code + ...(provider === "edgeone" && { + headers: { + cookie: req.headers.get("cookie") || "", + "x-access-code": req.headers.get("x-access-code") || "", + }, + }), }) + + // Streaming, like the chat (some models only stream). Stop at the + // first tool call; a reasoning model that runs out of tokens first + // proves the connection but not tool support. + const startTime = Date.now() + const result = streamText({ + model, + prompt: "Call the ping tool.", + tools: { ping: PING_TOOL }, + maxOutputTokens: 1024, + maxRetries: 0, + abortSignal: AbortSignal.timeout(TEST_TIMEOUT_MS), + }) + let calledTool = false + let finishReason: string | undefined + for await (const part of result.fullStream) { + if (part.type === "error") throw part.error + // The timeout ends the stream with an abort part, not an error + if (part.type === "abort") { + const timeout = new Error( + `The model did not answer within ${TEST_TIMEOUT_MS / 1000} s.`, + ) + timeout.name = "TimeoutError" + throw timeout + } + if (part.type === "tool-call") { + calledTool = true + break + } + if (part.type === "finish") finishReason = part.finishReason + } const responseTime = Date.now() - startTime return NextResponse.json({ valid: true, responseTime, + ...(!calledTool && + finishReason !== "length" && { warning: NO_TOOL_CALL_WARNING }), }) } catch (error) { console.error("[validate-model] Error:", error) - let errorMessage = "Validation failed" - if (error instanceof Error) { - // Extract meaningful error message - if ( - error.message.includes("401") || - error.message.includes("Unauthorized") - ) { - errorMessage = "Invalid API key" - } else if ( - error.message.includes("404") || - error.message.includes("not found") - ) { - errorMessage = "Model not found" - } else if ( - error.message.includes("429") || - error.message.includes("rate limit") - ) { - errorMessage = "Rate limited - try again later" - } else if (error.message.includes("ECONNREFUSED")) { - errorMessage = "Cannot connect to server" - } else { - errorMessage = error.message.slice(0, 100) - } - } - + const { code, message } = classifyLLMError(error) return NextResponse.json( - { valid: false, error: errorMessage }, + { valid: false, code, error: message }, { status: 200 }, // Return 200 so client can read error message ) } diff --git a/app/api/verify-access-code/route.ts b/app/api/verify-access-code/route.ts index d69f59d1..55cbc94b 100644 --- a/app/api/verify-access-code/route.ts +++ b/app/api/verify-access-code/route.ts @@ -1,29 +1,9 @@ +import { checkAccessCode } from "@/lib/access-code" + export async function POST(req: Request) { - 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) { + if (checkAccessCode(req)) { return Response.json( - { valid: false, message: "Access code is required" }, - { status: 401 }, - ) - } - - if (!accessCodes.includes(accessCodeHeader)) { - return Response.json( - { valid: false, message: "Invalid access code" }, + { valid: false, message: "Invalid or missing access code" }, { status: 401 }, ) } diff --git a/biome.json b/biome.json index 32874167..bf56b8b0 100644 --- a/biome.json +++ b/biome.json @@ -1,12 +1,18 @@ { - "$schema": "https://biomejs.dev/schemas/2.3.10/schema.json", + "$schema": "https://biomejs.dev/schemas/2.4.14/schema.json", "vcs": { "enabled": true, "clientKind": "git", "useIgnoreFile": true }, "files": { - "ignoreUnknown": false + "ignoreUnknown": false, + "includes": [ + "**", + "!public", + "!packages/mcp-server/src/preview", + "!lib/model-catalog.json" + ] }, "formatter": { "enabled": true, diff --git a/components/ai-elements/model-selector.tsx b/components/ai-elements/model-selector.tsx index 1b71cb70..7164f44b 100644 --- a/components/ai-elements/model-selector.tsx +++ b/components/ai-elements/model-selector.tsx @@ -1,5 +1,6 @@ import { Cloud } from "lucide-react" -import type { ComponentProps, ReactNode } from "react" +import type { ComponentProps, ElementRef, ReactNode } from "react" +import { useEffect, useRef, useState } from "react" import { Command, CommandDialog, @@ -69,20 +70,62 @@ export type ModelSelectorListProps = ComponentProps export const ModelSelectorList = ({ className, ...props -}: ModelSelectorListProps) => ( -

- - {/* Bottom shadow indicator for scrollable content */} -
-
-) +}: ModelSelectorListProps) => { + const listRef = useRef>(null) + const [showShadow, setShowShadow] = useState(false) + + useEffect(() => { + const listElement = listRef.current + if (!listElement) return + + const checkScroll = () => { + const { scrollTop, scrollHeight, clientHeight } = listElement + // Show shadow if there is more content below + // Using a small threshold to handle fractional pixel rendering + setShowShadow( + scrollHeight > Math.ceil(scrollTop + clientHeight) + 1, + ) + } + + // Initial check + checkScroll() + + // Event listeners + listElement.addEventListener("scroll", checkScroll) + window.addEventListener("resize", checkScroll) + + // Observe content changes (e.g. async loading of items) + const observer = new MutationObserver(checkScroll) + observer.observe(listElement, { childList: true, subtree: true }) + + return () => { + listElement.removeEventListener("scroll", checkScroll) + window.removeEventListener("resize", checkScroll) + observer.disconnect() + } + }, []) + + return ( +
+ + {/* Bottom shadow indicator for scrollable content */} +
+
+ ) +} export type ModelSelectorEmptyProps = ComponentProps @@ -169,3 +212,27 @@ export const ModelSelectorName = ({ }: ModelSelectorNameProps) => ( ) + +export type ModelSelectorSectionHeaderProps = { + icon: ReactNode + label: string + className?: string +} + +export const ModelSelectorSectionHeader = ({ + icon, + label, + className, +}: ModelSelectorSectionHeaderProps) => ( +
+ + {label} +
+) diff --git a/components/chat-example-panel.tsx b/components/chat-example-panel.tsx index a74f42cc..4e721292 100644 --- a/components/chat-example-panel.tsx +++ b/components/chat-example-panel.tsx @@ -141,9 +141,6 @@ export default function ExamplePanel({ {dict.examples.mcpServer} - - {dict.examples.preview} -

{dict.examples.mcpDescription} diff --git a/components/chat-input.tsx b/components/chat-input.tsx index 6848dc58..0f375925 100644 --- a/components/chat-input.tsx +++ b/components/chat-input.tsx @@ -1,17 +1,28 @@ "use client" import { + BookmarkPlus, Download, History, Image as ImageIcon, Link, - Loader2, Send, + Square, } from "lucide-react" import type React from "react" -import { useCallback, useEffect, useRef, useState } from "react" +import { + type Dispatch, + forwardRef, + type SetStateAction, + useCallback, + useEffect, + useImperativeHandle, + useRef, + useState, +} from "react" import { toast } from "sonner" import { ButtonWithTooltip } from "@/components/button-with-tooltip" +import { TemplateCreateDialog } from "@/components/chat/TemplateCreateDialog" import { ErrorToast } from "@/components/error-toast" import { HistoryDialog } from "@/components/history-dialog" import { ModelSelector } from "@/components/model-selector" @@ -27,13 +38,25 @@ import { isPdfFile, isTextFile } from "@/lib/pdf-utils" import { STORAGE_KEYS } from "@/lib/storage" import type { FlattenedModel } from "@/lib/types/model-config" import { extractUrlContent, type UrlData } from "@/lib/url-utils" +import { isRealDiagram } from "@/lib/utils" 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 file.type.startsWith("image/") || isPdfFile(file) || isTextFile(file) + return ( + SUPPORTED_IMAGE_TYPES.includes(file.type) || + isPdfFile(file) || + isTextFile(file) + ) } function formatFileSize(bytes: number): string { @@ -137,11 +160,16 @@ function showValidationErrors(errors: string[], dict: any) { } } +export interface ChatInputRef { + focus: () => void +} + interface ChatInputProps { input: string status: "submitted" | "streaming" | "ready" | "error" onSubmit: (e: React.FormEvent) => void onChange: (e: React.ChangeEvent) => void + onStop?: () => void files?: File[] onFileChange?: (files: File[]) => void pdfData?: Map< @@ -149,7 +177,7 @@ interface ChatInputProps { { text: string; charCount: number; isExtracting: boolean } > urlData?: Map - onUrlChange?: (data: Map) => void + onUrlChange?: Dispatch>> sessionId?: string error?: Error | null @@ -157,122 +185,230 @@ interface ChatInputProps { models?: FlattenedModel[] selectedModelId?: string onModelSelect?: (modelId: string | undefined) => void - showUnvalidatedModels?: boolean onConfigureModels?: () => void + showUnvalidatedModels?: boolean + // Focus control props + shouldFocus?: boolean + onFocused?: () => void } -export function ChatInput({ - input, - status, - onSubmit, - onChange, - files = [], - onFileChange = () => {}, - pdfData = new Map(), - urlData, - onUrlChange, - sessionId, - error = null, - models = [], - selectedModelId, - onModelSelect = () => {}, - showUnvalidatedModels = false, - onConfigureModels = () => {}, -}: ChatInputProps) { - const dict = useDictionary() - const { - diagramHistory, - saveDiagramToFile, - showSaveDialog, - setShowSaveDialog, - } = useDiagram() +export const ChatInput = forwardRef( + function ChatInput( + { + input, + status, + onSubmit, + onChange, + onStop, + files = [], + onFileChange = () => {}, + pdfData = new Map(), + urlData, + onUrlChange, + sessionId, + error = null, + models = [], + selectedModelId, + onModelSelect = () => {}, + onConfigureModels, + showUnvalidatedModels = false, + shouldFocus = false, + onFocused, + }, + ref, + ) { + const dict = useDictionary() + const { + chartXML, + diagramHistory, + saveDiagramToFile, + showSaveDialog, + setShowSaveDialog, + } = useDiagram() - const textareaRef = useRef(null) - const fileInputRef = useRef(null) - const [isDragging, setIsDragging] = useState(false) - const [showHistory, setShowHistory] = useState(false) - const [showUrlDialog, setShowUrlDialog] = useState(false) - const [isExtractingUrl, setIsExtractingUrl] = useState(false) - const [sendShortcut, setSendShortcut] = useState("ctrl-enter") - // Allow retry when there's an error (even if status is still "streaming" or "submitted") - const isDisabled = - (status === "streaming" || status === "submitted") && !error + const textareaRef = useRef(null) + const fileInputRef = useRef(null) + const [isDragging, setIsDragging] = useState(false) - const adjustTextareaHeight = useCallback(() => { - const textarea = textareaRef.current - if (textarea) { - textarea.style.height = "auto" - textarea.style.height = `${Math.min(textarea.scrollHeight, 200)}px` - } - }, []) - // Handle programmatic input changes (e.g., setInput("") after form submission) - useEffect(() => { - adjustTextareaHeight() - }, [input, adjustTextareaHeight]) + // Expose focus method via ref + useImperativeHandle(ref, () => ({ + focus: () => { + textareaRef.current?.focus() + }, + })) - // Load send shortcut preference from localStorage and listen for changes - useEffect(() => { - const stored = localStorage.getItem(STORAGE_KEYS.sendShortcut) - if (stored) setSendShortcut(stored) + // Focus the textarea when shouldFocus becomes true + // Use setTimeout to ensure focus happens after drawio iframe settles + useEffect(() => { + if (shouldFocus) { + const timer = setTimeout(() => { + textareaRef.current?.focus() + onFocused?.() + }, 150) + return () => clearTimeout(timer) + } + }, [shouldFocus, onFocused]) - const handleChange = (e: CustomEvent) => - setSendShortcut(e.detail) - window.addEventListener( - "sendShortcutChange", - handleChange as EventListener, - ) - return () => - window.removeEventListener( + const [showHistory, setShowHistory] = useState(false) + const [showUrlDialog, setShowUrlDialog] = useState(false) + const [showSaveAsTemplate, setShowSaveAsTemplate] = useState(false) + const [isExtractingUrl, setIsExtractingUrl] = useState(false) + const [sendShortcut, setSendShortcut] = useState("ctrl-enter") + // 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 + if (textarea) { + textarea.style.height = "auto" + textarea.style.height = `${Math.min(textarea.scrollHeight, 200)}px` + } + }, []) + // Handle programmatic input changes (e.g., setInput("") after form submission) + useEffect(() => { + adjustTextareaHeight() + }, [input, adjustTextareaHeight]) + + // Load send shortcut preference from localStorage and listen for changes + useEffect(() => { + const stored = localStorage.getItem(STORAGE_KEYS.sendShortcut) + if (stored) setSendShortcut(stored) + + const handleChange = (e: CustomEvent) => + setSendShortcut(e.detail) + window.addEventListener( "sendShortcutChange", handleChange as EventListener, ) - }, []) + return () => + window.removeEventListener( + "sendShortcutChange", + handleChange as EventListener, + ) + }, []) - const handleChange = (e: React.ChangeEvent) => { - onChange(e) - adjustTextareaHeight() - } + const handleChange = (e: React.ChangeEvent) => { + onChange(e) + adjustTextareaHeight() + } - const handleKeyDown = (e: React.KeyboardEvent) => { - const shouldSend = - sendShortcut === "enter" - ? e.key === "Enter" && !e.shiftKey && !e.ctrlKey && !e.metaKey - : (e.metaKey || e.ctrlKey) && e.key === "Enter" + const handleKeyDown = (e: React.KeyboardEvent) => { + // Enter that confirms an IME candidate must not send the message + if (e.nativeEvent.isComposing || e.keyCode === 229) return - if (shouldSend) { - e.preventDefault() - const form = e.currentTarget.closest("form") - if (form && input.trim() && !isDisabled) { - form.requestSubmit() + const shouldSend = + sendShortcut === "enter" + ? e.key === "Enter" && + !e.shiftKey && + !e.ctrlKey && + !e.metaKey + : (e.metaKey || e.ctrlKey) && e.key === "Enter" + + if (shouldSend) { + e.preventDefault() + const form = e.currentTarget.closest("form") + if ( + form && + input.trim() && + !isDisabled && + !isExtractingAttachments + ) { + form.requestSubmit() + } } } - } - const handlePaste = async (e: React.ClipboardEvent) => { - if (isDisabled) return + const handlePaste = async (e: React.ClipboardEvent) => { + if (isDisabled) return - const items = e.clipboardData.items - const imageItems = Array.from(items).filter((item) => - item.type.startsWith("image/"), - ) + const items = e.clipboardData.items + const imageItems = Array.from(items).filter((item) => + item.type.startsWith("image/"), + ) - if (imageItems.length > 0) { - const imageFiles = ( - await Promise.all( - imageItems.map(async (item, index) => { - const file = item.getAsFile() - if (!file) return null - return new File( - [file], - `pasted-image-${Date.now()}-${index}.${file.type.split("/")[1]}`, - { type: file.type }, - ) - }), + if (imageItems.length > 0) { + const imageFiles = ( + await Promise.all( + imageItems.map(async (item, index) => { + const file = item.getAsFile() + if (!file) return null + return new File( + [file], + `pasted-image-${Date.now()}-${index}.${file.type.split("/")[1]}`, + { type: file.type }, + ) + }), + ) + ).filter((f): f is File => f !== null) + + const { validFiles, errors } = validateFiles( + imageFiles, + files.length, + dict, ) - ).filter((f): f is File => f !== null) + showValidationErrors(errors, dict) + if (validFiles.length > 0) { + onFileChange([...files, ...validFiles]) + } + } + } + const handleFileChange = (e: React.ChangeEvent) => { + const newFiles = Array.from(e.target.files || []) const { validFiles, errors } = validateFiles( - imageFiles, + newFiles, + files.length, + dict, + ) + showValidationErrors(errors, dict) + if (validFiles.length > 0) { + onFileChange([...files, ...validFiles]) + } + + if (fileInputRef.current) { + fileInputRef.current.value = "" + } + } + + const handleRemoveFile = (fileToRemove: File) => { + onFileChange(files.filter((file) => file !== fileToRemove)) + if (fileInputRef.current) { + fileInputRef.current.value = "" + } + } + + const triggerFileInput = () => { + fileInputRef.current?.click() + } + + const handleDragOver = (e: React.DragEvent) => { + e.preventDefault() + e.stopPropagation() + setIsDragging(true) + } + + const handleDragLeave = (e: React.DragEvent) => { + e.preventDefault() + e.stopPropagation() + setIsDragging(false) + } + + const handleDrop = (e: React.DragEvent) => { + e.preventDefault() + e.stopPropagation() + setIsDragging(false) + + if (isDisabled) return + + // Let validateFiles show a toast for unsupported types + const { validFiles, errors } = validateFiles( + Array.from(e.dataTransfer.files), files.length, dict, ) @@ -281,278 +417,253 @@ export function ChatInput({ onFileChange([...files, ...validFiles]) } } - } - const handleFileChange = (e: React.ChangeEvent) => { - const newFiles = Array.from(e.target.files || []) - const { validFiles, errors } = validateFiles( - newFiles, - files.length, - dict, - ) - showValidationErrors(errors, dict) - if (validFiles.length > 0) { - onFileChange([...files, ...validFiles]) + const handleUrlExtract = async (url: string) => { + if (!onUrlChange) return + + 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 data = await extractUrlContent(url) + + // Skip if the URL was removed while extracting + onUrlChange((prev) => + prev.has(url) ? new Map(prev).set(url, data) : prev, + ) + + 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 + }) + showErrorToast( + + {error instanceof Error + ? error.message + : "Failed to extract URL content"} + , + ) + } finally { + setIsExtractingUrl(false) + } } - if (fileInputRef.current) { - fileInputRef.current.value = "" - } - } - - const handleRemoveFile = (fileToRemove: File) => { - onFileChange(files.filter((file) => file !== fileToRemove)) - if (fileInputRef.current) { - fileInputRef.current.value = "" - } - } - - const triggerFileInput = () => { - fileInputRef.current?.click() - } - - const handleDragOver = (e: React.DragEvent) => { - e.preventDefault() - e.stopPropagation() - setIsDragging(true) - } - - const handleDragLeave = (e: React.DragEvent) => { - e.preventDefault() - e.stopPropagation() - setIsDragging(false) - } - - const handleDrop = (e: React.DragEvent) => { - e.preventDefault() - e.stopPropagation() - setIsDragging(false) - - if (isDisabled) return - - const droppedFiles = e.dataTransfer.files - const supportedFiles = Array.from(droppedFiles).filter((file) => - isValidFileType(file), - ) - - const { validFiles, errors } = validateFiles( - supportedFiles, - files.length, - dict, - ) - showValidationErrors(errors, dict) - if (validFiles.length > 0) { - onFileChange([...files, ...validFiles]) - } - } - - const handleUrlExtract = async (url: string) => { - if (!onUrlChange) return - - setIsExtractingUrl(true) - - try { - const existing = urlData - ? new Map(urlData) - : new Map() - existing.set(url, { - url, - title: url, - content: "", - charCount: 0, - isExtracting: true, - }) - onUrlChange(existing) - - const data = await extractUrlContent(url) - - const newUrlData = new Map(existing) - newUrlData.set(url, data) - onUrlChange(newUrlData) - - setShowUrlDialog(false) - } catch (error) { - // Remove the URL from the data map on error - const newUrlData = urlData - ? new Map(urlData) - : new Map() - newUrlData.delete(url) - onUrlChange(newUrlData) - showErrorToast( - - {error instanceof Error - ? error.message - : "Failed to extract URL content"} - , - ) - } finally { - setIsExtractingUrl(false) - } - } - - return ( -

- {/* File & URL previews */} - {(files.length > 0 || (urlData && urlData.size > 0)) && ( -
- { - const next = new Map(urlData) - next.delete(url) - onUrlChange(next) - } - : undefined - } + return ( + + {/* File & URL previews */} + {(files.length > 0 || (urlData && urlData.size > 0)) && ( +
+ + onUrlChange((prev) => { + const next = new Map(prev) + next.delete(url) + return next + }) + : undefined + } + /> +
+ )} +
+