diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..e2d6602 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,51 @@ +name: CI + +# Fast feedback for every push and pull request: typecheck the whole tree +# (including the renderer), run the full electron-vite build, then the unit +# test suite. Heavy per-OS packaging stays in build.yml. +on: + push: + branches: [master] + pull_request: + +permissions: + contents: read + +concurrency: + group: ci-${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + verify: + name: Typecheck, build and test + runs-on: ubuntu-latest + timeout-minutes: 20 + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-node@v4 + with: + node-version: 22 + cache: npm + + - name: Install dependencies + run: npm ci + + # Full pipeline: `tsc -p tsconfig.json` typechecks CLI, Electron main and + # renderer sources, then electron-vite bundles main/preload/renderer. + - name: Build (tsc + electron-vite) + run: npm run build + + # Compile before testing: dist/ is gitignored, and `node --test` exits 0 + # on an unmatched glob, so testing before building silently runs nothing. + - name: Unit tests + run: npm test + + # `doctor` intentionally fails on a fresh workspace (it requires a fully + # configured install with SecScore MCP enabled), so the smoke test only + # covers init + a read-only config command. + - name: Smoke test the CLI + run: | + SMOKE_WORKSPACE="$(mktemp -d)/workspace" + SECTL_WORKSPACE="$SMOKE_WORKSPACE" node dist/index.js init + SECTL_WORKSPACE="$SMOKE_WORKSPACE" node dist/index.js mcp list diff --git a/.gitignore b/.gitignore index 5af685b..d826e1e 100644 --- a/.gitignore +++ b/.gitignore @@ -21,5 +21,9 @@ __pycache__/ .claude-tmp-* secagent-ce-icon-generated.png +# Windows-style temp fixtures the test suite creates when run on Linux +# (literal-backslash names such as "\tmp\secagent-...") +\\tmp* + AGENTS.local.md \ No newline at end of file diff --git a/README.md b/README.md index 68bc36e..07d2c6d 100644 --- a/README.md +++ b/README.md @@ -1,22 +1,83 @@ -# SecAgent CLI +# SecAgent -## Workspace selection +SecAgent 是一个把自然语言转换为工具调用的桌面 Agent:支持多模型提供商、语音输入、本地工具、MCP 服务与插件。本仓库包含桌面端(Electron + React)与 CLI 两种使用方式。 -The default workspace is `~/SecAgentWorkspace`. Set `SECTL_WORKSPACE` before starting -the CLI or desktop app to use another directory. An explicit `--workspace` argument -overrides the environment variable for CLI commands. +## 数据目录 + +SecAgent 的所有数据(配置、密钥、会话、日志)集中存放在**一个**按平台约定的工作区目录: + +| 平台 | 默认工作区 | +|---|---| +| Windows | `%APPDATA%\SecAgent\workspace` | +| macOS | `~/Library/Application Support/SecAgent/workspace` | +| Linux | `$XDG_CONFIG_HOME/SecAgent/workspace`(未设置时为 `~/.config/SecAgent/workspace`) | + +首次启动时,旧版遗留在 `~/SecAgentWorkspace` 的数据会**自动迁移**到上述目录(跨盘符时降级为复制,旧目录改名为 `SecAgentWorkspace.migrated` 备查)。 + +需要使用其他目录时,设置 `SECTL_WORKSPACE` 环境变量;CLI 命令可用 `--workspace` 参数覆盖。工作区内布局: -```powershell -$env:SECTL_WORKSPACE = "D:\Temp\SecAgentTest" -npm run build:cli -node dist/index.js init -node dist/index.js sessions list ``` +<工作区>/ +├── secagent.yaml # 配置(模型提供商、MCP、语音、更新等) +├── .env # API 密钥(不要提交、不要手写变量名,见下文) +├── sessions/ # 会话历史(session.json + runtime.jsonl) +├── logs/ # 运行日志 +└── skills/ # SKILL.md 技能文件 +``` + +## 模型提供商 -To use the setting for the desktop development app, keep the environment variable in -the same PowerShell session and run `npm run build` followed by `npm run start`. +桌面端“设置 → 模型提供商”中添加提供商:填名称、Base URL、粘贴 API Key 即可。**不需要手写环境变量名**——保存时按提供商名自动生成(如 `SECAGENT_DEEPSEEK_API_KEY`),密钥只写入工作区 `.env`,绝不进入 `secagent.yaml`。同一个预设添加两次(例如两个账号)会自动加编号后缀,避免密钥互相覆盖。 -## CLI 调试 Agent +所有模型选择处(主页模型菜单、默认模型、唤醒模型)均**按提供商分组显示**,不同提供商下的同名模型不会再混淆。多个 Google 提供商(官方 key + 中转)的模型会全部列出。 + +YAML 手写示例(与设置界面等价): + +```yaml +agent: + providers: + - id: deepseek + name: DeepSeek + provider: openai-compatible + apiKeyEnv: SECAGENT_DEEPSEEK_API_KEY # .env 中的变量名(自动生成) + baseUrl: https://api.deepseek.com/v1 + endpoint: /chat/completions + maxTokens: 16384 + models: + - id: deepseek-chat + name: DeepSeek V3 +``` + +## 模型稳定性(重试 / 备用切换 / 冷却) + +面向阿里云百炼等“赠送资源包”平台设计——资源包用尽时无需手动换模型: + +- **自动重试**:同一模型对网络类错误重试一次; +- **多轮 fallback**:当前模型失败后按顺序切换到其他已配置模型,全部失败才报错;切换链覆盖每一个已启用模型; +- **失败记忆与冷却**:配额耗尽 / 鉴权失败的模型进入冷却期(普通错误 5 分钟起指数退避;资源包类 60 分钟),期间被跳过,成功一次即自动恢复。状态持久化在工作区 `.model-health.json`,重启后仍然生效; +- 语音识别(ASR)链同样支持:第三方 → 官方 → 本地逐级回退,连续失败的提供方短暂禁用。 + +以上行为均可在“设置 → 系统 → 模型稳定性”中分开关控制。 + +## 敏感操作确认(Codex 风格) + +模型请求执行删除文件、格式化磁盘、强制推送、写系统注册表、写工作区外路径、下载并执行等操作时,会弹窗要求确认: + +- **拒绝**:工具调用被拦截,模型收到说明并改用其他方式; +- **允许一次**:仅本次放行; +- **总是允许此类**:按「工具 + 命令头」签名记忆(如 `bash|rm`),不同命令不会误放行;签名保存在 `secagent.yaml` 的 `guard.approved`。 + +CLI 模式下通过终端 `y/N` 确认;非交互环境(管道 / CI)默认拒绝。5 分钟无响应自动拒绝。总开关位于“设置 → 系统 → 安全与检测”。 + +## 幻觉检测 + +最终回答会经过轻量启发式检测,命中时在回答下方显示提醒条(仅提醒、不拦截): + +- 工具调用**失败**后回答却声称“已成功完成”; +- 回答陷入重复循环(小模型过载时的常见模式); +- 引用了本轮从未产生的材料(“如上表所示”但没有任何工具产出)。 + +## CLI CLI 的每次 `run` 都会持久化为一个会话,并默认实时打印模型思考片段、工具调用、工具返回结果和最终回答。模型请求失败时会保存错误消息并返回非零退出码。 @@ -25,8 +86,8 @@ cd SecAgent npm install npm run build:cli -# 初始化工作区,并在 .env 中填写模型密钥 -node dist/index.js init --workspace ./demo-workspace +# 初始化工作区(默认使用上文平台目录;也可用 --workspace 指定) +node dist/index.js init # 执行单条消息;命令结束时会打印 [session] <会话 ID> node dist/index.js run "查询李明当前积分" --workspace ./demo-workspace @@ -41,101 +102,25 @@ node dist/index.js run "把刚才的结果总结一下" --session <会话 ID> -- node dist/index.js chat --session <会话 ID> --workspace ./demo-workspace ``` -交互式 `chat` 中输入 `:history` 查看当前会话,输入 `:use <会话 ID>` 切换会话,输入 `exit` 退出。需要完整的模型请求/响应原始事件时,加上 `--verbose`;普通模式已经会打印思考和工具过程。 - -会话文件位于工作区的 `sessions/<会话 ID>/session.json`,运行时事件位于同目录的 `runtime.jsonl`,因此 CLI 和桌面端可以共享历史会话。 - -SecAgent:把自然语言转换为工具调用。 - -```bash -npm install -npm run build -node dist/index.js init --workspace ./demo-workspace -node dist/index.js run "给高一三班的李明加 2 分" --workspace ./demo-workspace -``` +交互式 `chat` 中输入 `:history` 查看当前会话,输入 `:use <会话 ID>` 切换会话,输入 `exit` 退出。需要完整的模型请求/响应原始事件时,加上 `--verbose`。 CLI 直接调用 SecScore 的 HTTP MCP(默认 `http://127.0.0.1:3901/mcp`),支持查学生、真实写入、审计和撤销。 -## 云端中文语音输入 +## 语音输入(多提供方 + 自动回退) -桌面端麦克风按钮统一通过 SecAgent 官方服务的 WebSocket 接口进行云端识别。使用前需要登录官方服务并配置 `SECTL_OFFICIAL_API_URL` 和 `SECTL_OFFICIAL_TOKEN`;音频不会在本地使用 `sherpa-onnx` 模型处理。 - -主界面输入框支持鼠标或触摸长按 0.7 秒说话,松开后一次性识别并插入输入框;向左侧“拖动至此取消”区域松开可取消。也可以点击麦克风按钮开始,再在录音条上松开完成识别。 +语音识别(ASR)被抽象为独立的提供方层(`src/asr/`),支持三种后端并按链自动回退: -## 模型配置 +| 顺序 | 提供方 | 说明 | +|---|---|---| +| 1 | 第三方云端 | 任意 OpenAI 兼容 `/audio/transcriptions` 端点,内置小米 MiMo ASR / SiliconFlow SenseVoice / Groq Whisper 预设 | +| 2 | 官方云端 | SECTL 官方服务 WebSocket(需登录),仅在位于回退链中时启用 | +| 3 | 本地离线 | 随应用打包的 sherpa-onnx 流式模型,无需网络 | -`secagent init` 会在工作区创建 `.env`。将密钥填入其中,密钥不会写进 `secagent.yaml`: +设置 → 语音识别中可选择“自动”(默认,按上表顺序回退)或固定某一后端,并支持一键“测试识别服务连通性”。API Key 同样不需要手写环境变量名,保存时自动写入 `.env`。 -```dotenv -OPENAI_API_KEY=... -ANTHROPIC_API_KEY=... -GEMINI_API_KEY=... -``` - -桌面端打开“设置”后,可在“协议”中选择“Google Gemini”,粘贴从 Google AI Studio 获取的 API key 并保存。程序会使用 Gemini 原生 API;key 会保存到工作区 `.env`,不会写入 `secagent.yaml`。 - -在 `secagent.yaml` 的 `agent` 区块选择协议、模型和端点: - -```yaml -# OpenAI Responses 协议 -agent: - provider: openai-responses - model: gpt-5 - apiKeyEnv: OPENAI_API_KEY - baseUrl: https://api.openai.com/v1 - endpoint: /responses - maxTokens: 16384 - -# OpenAI 或任何兼容 Chat Completions 的服务 -agent: - provider: openai-compatible - model: gpt-5 - apiKeyEnv: OPENAI_API_KEY - baseUrl: https://api.openai.com/v1 - endpoint: /chat/completions - maxTokens: 16384 - - # 可选:配置多个模型后,可在桌面端输入框右侧切换。 - models: - - id: gpt-5 - name: GPT-5 - provider: openai-compatible - model: gpt-5 - apiKeyEnv: OPENAI_API_KEY - baseUrl: https://api.openai.com/v1 - endpoint: /chat/completions - maxTokens: 16384 - - id: claude - name: Claude Sonnet - provider: anthropic - model: claude-sonnet-4-20250514 - apiKeyEnv: ANTHROPIC_API_KEY - baseUrl: https://api.anthropic.com - endpoint: /v1/messages - anthropicVersion: "2023-06-01" - maxTokens: 16384 - -# Anthropic Messages API 或其兼容端点 -# agent: -# provider: anthropic -# model: claude-sonnet-4-20250514 -# apiKeyEnv: ANTHROPIC_API_KEY -# baseUrl: https://api.anthropic.com -# endpoint: /v1/messages -# anthropicVersion: "2023-06-01" -# maxTokens: 16384 - -# Google Gemini 原生 API -# agent: -# provider: google -# model: gemini-2.5-flash -# apiKeyEnv: GEMINI_API_KEY -# baseUrl: https://generativelanguage.googleapis.com/v1beta -# endpoint: "" -# maxTokens: 16384 -``` +主界面输入框支持鼠标或触摸长按 0.7 秒说话,松开后一次性识别并插入输入框;向左侧“拖动至此取消”区域松开可取消。也可以点击麦克风按钮开始,再在录音条上松开完成识别。 -桌面端输入框右侧的模型菜单可以分别选择模型和推理强度(不思考、低、中、高)。OpenAI Responses 会将其映射到 `reasoning.effort`;Anthropic 和 Gemini 会映射到各自的 thinking 配置。Responses、Anthropic thinking 和 Gemini thought summary 的流式内容会按时间顺序显示在工具执行过程内,最终答案仍单独显示。 +## 工具与技能 模型可直接调用所有已发现的 MCP 工具,以及 Pi 风格的 `look_at`、`read`、`write`、`edit`、`bash` 五个本地工具;`look_at` 会读取工作区或本地路径中的图片并以多模态内容返回给模型,每次调用仍会写入本地审计。 @@ -151,10 +136,23 @@ description: 处理学生查询、积分加减分和撤销。 隐藏 MCP 工具的声明方式、通用调用入口,以及 Skill/MCP 开发者约定见 [`docs/skill-mcp-convention.md`](docs/skill-mcp-convention.md)。 +## 开发与 CI + +```bash +npm install +npm run build # tsc 全量类型检查 + electron-vite 打包(main/preload/renderer) +npm test # node --test dist/**/*.test.js +``` + +GitHub Actions: + +- **CI**(`.github/workflows/ci.yml`):每次 push / PR 运行——类型检查、完整构建、单元测试、CLI 冒烟(`init` + `doctor`); +- **Build**(`.github/workflows/build.yml`):三平台打包并上传构建产物。 + ## 更新检查与诊断日志 SecAgent 会优先读取签名的 `updates.json` 通道清单;清单暂不可用时回退到 GitHub Releases API,并依次尝试代理和直连。安装包下载后必须通过 SHA-256 校验。 -更新设置页面中的“打开日志目录”可直接打开工作区的 `logs` 目录;“导出诊断日志”会生成脱敏 ZIP,适合提交故障信息。日志通常位于 `%USERPROFILE%\SecAgentWorkspace\logs`。 +更新设置页面中的“打开日志目录”可直接打开工作区的 `logs` 目录;“导出诊断日志”会生成脱敏 ZIP,适合提交故障信息。日志位于上文“数据目录”表格中对应平台的 `<工作区>/logs`。 发布者如需启用签名清单,应将与 `src/update-public-key.ts` 匹配的私钥配置为 GitHub Actions Secret:`SECAGENT_UPDATE_PRIVATE_KEY`。私钥不得提交到仓库。 diff --git a/electron.vite.config.ts b/electron.vite.config.ts index 01a0a2d..f31fa50 100644 --- a/electron.vite.config.ts +++ b/electron.vite.config.ts @@ -1,11 +1,19 @@ -import { defineConfig } from "electron-vite"; +import { defineConfig, externalizeDepsPlugin } from "electron-vite"; import react from "@vitejs/plugin-react"; export default defineConfig({ - main: { build: { rollupOptions: { input: "src/electron/main.ts" } } }, + // Keep native/WASM-heavy dependencies as runtime requires: bundling + // sherpa-onnx (Emscripten loader + .wasm + onnx models) into the main or + // preload bundle breaks its on-disk lookups, which is exactly what made the + // local speech model fail to load in packaged builds. + main: { + plugins: [externalizeDepsPlugin()], + build: { rollupOptions: { input: "src/electron/main.ts" } } + }, // Electron runs sandboxed preload scripts as CommonJS. A `.cjs` filename is // important because this package otherwise opts into ESM via `type: module`. preload: { + plugins: [externalizeDepsPlugin()], build: { rollupOptions: { input: "src/electron/preload.ts", diff --git a/package.json b/package.json index 48f2061..b625c82 100644 --- a/package.json +++ b/package.json @@ -26,6 +26,7 @@ }, "dependencies": { "@andresaya/edge-tts": "^1.8.0", + "ws": "^8.21.1", "@opencode-ai/models": "^0.0.35", "@react-three/drei": "^10.7.8", "@react-three/fiber": "^9.7.0", diff --git a/scripts/fetch-sensevoice-pack.mjs b/scripts/fetch-sensevoice-pack.mjs new file mode 100644 index 0000000..6cf9431 --- /dev/null +++ b/scripts/fetch-sensevoice-pack.mjs @@ -0,0 +1,68 @@ +#!/usr/bin/env node +/** + * 本地模型升级附加包安装器 — SenseVoice (sherpa-onnx int8, ~230MB). + * + * Downloads the official sherpa-onnx SenseVoice pack and unpacks it to + * `models/sense-voice/` next to the resources the app searches: + * - dev: /models/sense-voice + * - installed desktop app: /models/sense-voice (pass --out ) + * + * Usage: + * node scripts/fetch-sensevoice-pack.mjs # dev (repo/models) + * node scripts/fetch-sensevoice-pack.mjs --out "C:\Program Files\SecAgent\resources" + * + * Source: https://github.com/k2-fsa/sherpa-onnx/releases (asr-models) + */ +import fs from "node:fs"; +import path from "node:path"; +import os from "node:os"; +import { spawn } from "node:child_process"; +import { pipeline } from "node:stream/promises"; +import { Readable } from "node:stream"; + +const URL_BASE = "https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-sense-voice-zh-en-ja-ko-yue-int8.tar.bz2"; +const FILES = ["model.int8.onnx", "tokens.txt"]; + +const args = process.argv.slice(2); +const outIdx = args.indexOf("--out"); +const targetRoot = outIdx >= 0 && args[outIdx + 1] ? args[outIdx + 1] : path.resolve(process.cwd(), "models"); +const packDir = path.join(targetRoot, "sense-voice"); + +function have(files) { return files.every((f) => fs.existsSync(path.join(packDir, f))); } + +if (have(FILES)) { + console.log(`✓ SenseVoice 附加包已存在:${packDir}(如需重装请先删除该目录)`); + process.exit(0); +} + +fs.mkdirSync(packDir, { recursive: true }); +const tmp = fs.mkdtempSync(path.join(os.tmpdir(), "sensevoice-")); +const archive = path.join(tmp, "pack.tar.bz2"); + +console.log("↓ 下载 SenseVoice int8 附加包(约 230 MB,只需一次)…"); +console.log(" " + URL_BASE); +const response = await fetch(URL_BASE, { redirect: "follow" }); +if (!response.ok || !response.body) { console.error(`✗ 下载失败:HTTP ${response.status}`); process.exit(1); } +let seen = 0; +const total = Number(response.headers.get("content-length") || 0); +const tracked = new Readable().wrap(response.body); +tracked.on("data", (chunk) => { + seen += chunk.length; + if (total) process.stdout.write(`\r ${(seen / 1048576).toFixed(1)} / ${(total / 1048576).toFixed(1)} MB`); +}); +await pipeline(tracked, fs.createWriteStream(archive)); +process.stdout.write("\n"); + +console.log("→ 解压(tar xjf)…"); +const child = spawn("tar", ["-xjf", archive, "-C", tmp, "--strip-components", "1"], { stdio: "inherit" }); +await new Promise((resolve, reject) => child.on("close", (code) => (code === 0 ? resolve() : reject(new Error(`tar 退出码 ${code}`))))); + +for (const file of FILES) { + const from = path.join(tmp, file); + const to = path.join(packDir, file); + fs.copyFileSync(from, to); +} +fs.rmSync(tmp, { recursive: true, force: true }); +const size = FILES.map((f) => fs.statSync(path.join(packDir, f)).size).reduce((a, b) => a + b, 0); +console.log(`✓ 安装完成:${packDir}(${(size / 1048576).toFixed(1)} MB)`); +console.log(" 在 SecAgent 设置 → 语音识别 中选择「本地增强(SenseVoice 附加包)」即可启用。"); diff --git a/src/asr/bailian-config.ts b/src/asr/bailian-config.ts new file mode 100644 index 0000000..1789d68 --- /dev/null +++ b/src/asr/bailian-config.ts @@ -0,0 +1,61 @@ +/** + * Shared 阿里云百炼 ASR configuration resolution. + * + * Values entered in the settings UI win; the workspace `.env` (already loaded + * into `process.env` by the config layer) can supply the API key plus the + * endpoints/models for headless setups. Documented defaults from + * `BAILIAN_DEFAULTS` fill whatever is left empty. + */ +import { BAILIAN_DEFAULTS, type BailianAsrSettings } from "./settings.js"; + +export interface ResolvedBailianConfig { + settings?: BailianAsrSettings; + apiKey: string; + /** OpenAI-compatible base URL for channel A (`chat/completions`). */ + baseUrl: string; + /** Realtime WebSocket endpoint for channel B. */ + wsUrl: string; + /** Non-streaming model name. */ + model: string; + /** Realtime streaming model name. */ + streamModel: string; + /** Optional language hint; absent means auto detect. */ + language?: string; + enableItn: boolean; +} + +function text(value: unknown): string { + return typeof value === "string" ? value.trim() : ""; +} + +function envFlag(value: string | undefined): boolean | undefined { + const normalized = text(value).toLowerCase(); + if (!normalized) return undefined; + return normalized === "1" || normalized === "true" || normalized === "yes" || normalized === "on"; +} + +/** Resolve 百炼 settings/env/defaults; `null` when no API key is available. */ +export function resolveBailianConfig( + settings: BailianAsrSettings | undefined, + getApiKey: (envName: string) => string | undefined, + env: Record = process.env +): ResolvedBailianConfig | null { + const apiKey = text(getApiKey(text(settings?.apiKeyEnv) || BAILIAN_DEFAULTS.apiKeyEnv)); + if (!apiKey) return null; + const language = text(settings?.language) || text(env.BAILIAN_LANGUAGE); + return { + ...(settings ? { settings } : {}), + apiKey, + baseUrl: (text(settings?.baseUrl) || text(env.BAILIAN_BASE_URL) || BAILIAN_DEFAULTS.baseUrl).replace(/\/+$/, ""), + wsUrl: (text(settings?.wsUrl) || text(env.BAILIAN_WS_URL) || BAILIAN_DEFAULTS.wsUrl).replace(/\/+$/, ""), + model: text(settings?.model) || text(env.BAILIAN_ASR_MODEL) || BAILIAN_DEFAULTS.model, + streamModel: text(settings?.streamModel) || text(env.BAILIAN_STREAM_MODEL) || BAILIAN_DEFAULTS.streamModel, + ...(language ? { language } : {}), + enableItn: settings?.enableItn ?? envFlag(env.BAILIAN_ENABLE_ITN) ?? false + }; +} + +/** Display name for logs and connectivity probes. */ +export function bailianDisplayName(settings: BailianAsrSettings | undefined, fallback: string): string { + return text(settings?.name) || fallback; +} \ No newline at end of file diff --git a/src/asr/bailian-http.ts b/src/asr/bailian-http.ts new file mode 100644 index 0000000..42a0e83 --- /dev/null +++ b/src/asr/bailian-http.ts @@ -0,0 +1,232 @@ +/** + * 阿里云百炼 ASR — channel A (non-streaming whole-utterance recognition). + * + * 百炼 does not implement OpenAI's `/audio/transcriptions` (it answers 404); + * its ASR runs on the OpenAI-compatible `chat/completions` endpoint with an + * `input_audio` payload carrying a `data:audio/wav;base64,…` Data URL. + * Requests are intentionally non-streaming (`stream: false`); plain HTTP + * cannot stream, so the session flushes buffered audio as `partial` text + * roughly every three seconds and as the `final` result on stop. + * + * Docs: https://help.aliyun.com/zh/model-studio/qwen-asr-api-reference + */ +import type { AsrEventSink, AsrProvider, AsrSession, AsrTestResult } from "./types.js"; +import { encodeWav, mergeSamples, ASR_SAMPLE_RATE } from "./wav.js"; +import type { BailianAsrSettings } from "./settings.js"; +import { bailianDisplayName, resolveBailianConfig } from "./bailian-config.js"; + +const PARTIAL_FLUSH_MS = 3_000; +const MIN_CHUNK_MS = 900; +const CONNECT_TIMEOUT_MS = 12_000; +/** Documented cap: the Base64 Data URL must stay within 10MB. */ +const MAX_DATA_URL_LENGTH = 10 * 1024 * 1024; +/** Exponential backoff for 429 responses: 3 retries after the first attempt. */ +const RATE_LIMIT_RETRIES = 3; +const RATE_LIMIT_BASE_DELAY_MS = 500; + +export interface BailianHttpAsrOptions { + /** Current 百炼 settings (re-read on each start). */ + getSettings: () => BailianAsrSettings | undefined; + /** Resolves the API key for an env var name (usually process.env). */ + getApiKey: (envName: string) => string | undefined; + fetchImpl?: typeof fetch; + /** Base delay for the 429 exponential backoff (tests shorten this). */ + rateLimitBaseDelayMs?: number; + log?: (message: string) => void; +} + +interface BailianRequest { + baseUrl: string; + apiKey: string; + model: string; + language?: string; + enableItn: boolean; +} + +interface ChatCompletionResponse { + choices?: Array<{ message?: { content?: string | Array<{ text?: string }> } }>; + error?: { message?: string } | string; +} + +function base64Length(bytes: number): number { + return Math.ceil(bytes / 3) * 4; +} + +/** + * Encode PCM as a WAV Data URL, trimming the oldest audio when the Base64 + * payload would exceed the documented 10MB input limit. + */ +function toWavDataUrl(samples: Float32Array, log?: (message: string) => void): string { + let wav = encodeWav(samples); + if (base64Length(wav.length) > MAX_DATA_URL_LENGTH) { + const maxWavBytes = Math.floor(MAX_DATA_URL_LENGTH / 4) * 3; + const maxSamples = Math.max(0, Math.floor((maxWavBytes - 44) / 2)); + wav = encodeWav(samples.subarray(samples.length - maxSamples)); + log?.(`[asr:bailian] 音频超过 Base64 10MB 上限,已截断至 ${(maxSamples / ASR_SAMPLE_RATE).toFixed(1)}s`); + } + return `data:audio/wav;base64,${Buffer.from(wav).toString("base64")}`; +} + +function statusMessage(status: number, body: string): string { + if (status === 401 || status === 403) return "API Key 无效或额度已耗尽(百炼“用完即停”):请检查密钥与免费额度,或在设置中切换识别服务"; + if (status === 404) return "接口不存在(404):请检查 Base URL 是否包含 /compatible-mode/v1"; + if (status === 400) return `请求被拒绝(400):音频超过 10MB 或格式不支持${body.slice(0, 160) ? `:${body.slice(0, 160)}` : ""}`; + if (status === 429) return "请求过于频繁(429 限流),已退避重试仍失败,请稍后重试"; + if (status >= 500) return `百炼服务端错误(${status}):${body.slice(0, 160)}`; + return `请求失败(${status}):${body.slice(0, 200)}`; +} + +/** `choices[0].message.content` is the transcript (string, or text parts). */ +function extractText(payload: ChatCompletionResponse): string { + const content = payload.choices?.[0]?.message?.content; + if (typeof content === "string") return content.trim(); + if (Array.isArray(content)) return content.map((part) => part?.text || "").join("").trim(); + return ""; +} + +const delay = (ms: number): Promise => new Promise((resolve) => setTimeout(resolve, ms)); + +export class BailianHttpAsrProvider implements AsrProvider { + readonly id = "bailian"; + readonly label = "阿里云百炼(chat/completions)"; + private readonly options: BailianHttpAsrOptions; + + constructor(options: BailianHttpAsrOptions) { + this.options = options; + } + + private displayName(): string { + return bailianDisplayName(this.options.getSettings(), this.label); + } + + isConfigured(): boolean { + return resolveBailianConfig(this.options.getSettings(), this.options.getApiKey) !== null; + } + + private async postOnce(request: BailianRequest, body: string, timeoutMs: number): Promise<{ ok: true; text: string } | { ok: false; message: string; rateLimited?: boolean }> { + const fetchImpl = this.options.fetchImpl || fetch; + const controller = new AbortController(); + const timer = setTimeout(() => controller.abort(), timeoutMs); + try { + const response = await fetchImpl(`${request.baseUrl}/chat/completions`, { + method: "POST", + headers: { Authorization: `Bearer ${request.apiKey}`, "Content-Type": "application/json" }, + body, + signal: controller.signal + }); + const raw = await response.text(); + if (!response.ok) return { ok: false, message: statusMessage(response.status, raw), ...(response.status === 429 ? { rateLimited: true } : {}) }; + let payload: ChatCompletionResponse; + try { payload = JSON.parse(raw) as ChatCompletionResponse; } + catch { return { ok: false, message: "服务返回了无法解析的内容(确认 Base URL 为百炼 /compatible-mode/v1)" }; } + if (payload.error) return { ok: false, message: typeof payload.error === "string" ? payload.error : payload.error.message || "百炼语音识别返回错误" }; + return { ok: true, text: extractText(payload) }; + } catch (error) { + if (error instanceof Error && error.name === "AbortError") return { ok: false, message: `连接超时(${timeoutMs / 1000}s):无法访问 ${request.baseUrl}` }; + return { ok: false, message: `无法连接 ${request.baseUrl}:${error instanceof Error ? error.message : String(error)}` }; + } finally { + clearTimeout(timer); + } + } + + private async transcribe(request: BailianRequest, samples: Float32Array, timeoutMs: number): Promise<{ ok: true; text: string } | { ok: false; message: string }> { + const body = JSON.stringify({ + model: request.model, + messages: [{ role: "user", content: [{ type: "input_audio", input_audio: { data: toWavDataUrl(samples, this.options.log) } }] }], + // 百炼 ASR 非流式场景固定 stream=false;asr_options 是百炼扩展参数,放 body 顶层。 + stream: false, + asr_options: { ...(request.language ? { language: request.language } : {}), enable_itn: request.enableItn } + }); + for (let attempt = 0; ; attempt += 1) { + const result = await this.postOnce(request, body, timeoutMs); + if (result.ok) return result; + if (!result.rateLimited || attempt >= RATE_LIMIT_RETRIES) return { ok: false, message: result.message }; + const backoff = (this.options.rateLimitBaseDelayMs ?? RATE_LIMIT_BASE_DELAY_MS) * 2 ** attempt; + this.options.log?.(`[asr:bailian] 429 限流,${backoff}ms 后重试(第 ${attempt + 1}/${RATE_LIMIT_RETRIES} 次)`); + await delay(backoff); + } + } + + private requestFor(config: { baseUrl: string; apiKey: string; model: string; language?: string; enableItn: boolean }): BailianRequest { + const request: BailianRequest = { baseUrl: config.baseUrl, apiKey: config.apiKey, model: config.model, enableItn: config.enableItn }; + if (config.language) request.language = config.language; + return request; + } + + async test(): Promise { + const config = resolveBailianConfig(this.options.getSettings(), this.options.getApiKey); + if (!config) return { ok: false, message: `环境变量 ${this.options.getSettings()?.apiKeyEnv || "BAILIAN_API_KEY"} 中没有 API Key(可在下方填写并保存到工作区 .env)` }; + // A 250ms near-silent probe is enough to validate auth, endpoint and model. + const silence = new Float32Array(Math.round(ASR_SAMPLE_RATE * 0.25)); + const result = await this.transcribe(this.requestFor(config), silence, CONNECT_TIMEOUT_MS); + if (!result.ok) return { ok: false, message: result.message }; + return { ok: true, message: `${this.displayName()} 连接成功${result.text ? `(识别:${result.text.slice(0, 40)})` : ""}` }; + } + + async start(sink: AsrEventSink): Promise { + const config = resolveBailianConfig(this.options.getSettings(), this.options.getApiKey); + if (!config) { + const settings = this.options.getSettings(); + const envName = settings?.apiKeyEnv?.trim() || "BAILIAN_API_KEY"; + throw new Error(`阿里云百炼语音识别未配置:请在设置中填写 API Key(保存到工作区 .env 的 ${envName})`); + } + // Snapshot the request template: settings may be replaced by a concurrent save. + const request = this.requestFor(config); + const provider = this; + let buffer: Float32Array[] = []; + let bufferedMs = 0; + let inFlight = false; + let cancelled = false; + let stopped = false; + const flushTimer = setInterval(() => { void maybeFlush(false); }, PARTIAL_FLUSH_MS); + + async function maybeFlush(final: boolean): Promise { + if (cancelled || inFlight) return; + if (!final && (stopped || bufferedMs < MIN_CHUNK_MS)) return; + const samples = mergeSamples(buffer); + buffer = []; + bufferedMs = 0; + if (!samples.length) return; + inFlight = true; + const result = await provider.transcribe(request, samples, CONNECT_TIMEOUT_MS); + inFlight = false; + if (cancelled) return; + if (!result.ok) { + provider.options.log?.(`[asr:bailian] ${result.message}`); + sink({ type: "error", message: result.message }); + return; + } + if (result.text) sink({ type: final ? "final" : "partial", text: result.text, provider: "bailian" }); + else if (final) sink({ type: "final", text: "", provider: "bailian" }); + } + + return { + providerId: this.id, + push: (samples: Float32Array): void => { + if (cancelled || stopped) return; + buffer.push(samples); + bufferedMs += (samples.length / ASR_SAMPLE_RATE) * 1_000; + }, + stop: async (): Promise => { + if (stopped || cancelled) return; + stopped = true; + clearInterval(flushTimer); + try { + // Wait for an in-flight partial flush so the final request is ordered. + const deadline = Date.now() + CONNECT_TIMEOUT_MS; + while (inFlight && Date.now() < deadline) await new Promise((resolve) => setTimeout(resolve, 50)); + await maybeFlush(true); + } finally { + clearInterval(flushTimer); + sink({ type: "stopped" }); + } + }, + cancel: (): void => { + cancelled = true; + clearInterval(flushTimer); + buffer = []; + bufferedMs = 0; + } + }; + } +} \ No newline at end of file diff --git a/src/asr/bailian-ws.ts b/src/asr/bailian-ws.ts new file mode 100644 index 0000000..942d3c6 --- /dev/null +++ b/src/asr/bailian-ws.ts @@ -0,0 +1,414 @@ +/** + * 阿里云百炼 ASR — channel B (realtime streaming, "type-as-you-speak"). + * + * DashScope WebSocket duplex protocol on `wss://…/api-ws/v1/inference`: + * 1. send `run-task` (JSON text frame), wait for `task-started` + * 2. stream raw PCM binary frames (16 kHz / 16-bit / mono, ~100ms / 3.2KB + * each, never above 16KB and never a burst of sub-1KB frames) + * 3. each frame may yield `result-generated` → `payload.output.sentence.text` + * 4. send `finish-task` and close after `task-finished`; `task-failed` + * surfaces `header.error_message`. + * An unexpected drop is reconnected once (the already-received audio is + * replayed); if that fails the session reports an error so the manager can + * fall back to channel A (`bailian`) or the local model. + * + * Docs: https://help.aliyun.com/zh/model-studio/fun-asr-realtime-websocket-api + */ +import type { AsrEventSink, AsrProvider, AsrSession, AsrTestResult } from "./types.js"; +import { ASR_SAMPLE_RATE, encodePcm16, mergeSamples } from "./wav.js"; +import type { AsrNoiseSettings, BailianAsrSettings } from "./settings.js"; +import { bailianDisplayName, resolveBailianConfig, type ResolvedBailianConfig } from "./bailian-config.js"; +// `ws`(而非全局 WebSocket):DashScope 在握手阶段校验 Authorization 头, +// Node 全局 WebSocket 不支持自定义请求头,必须用 ws 包。 +import Ws from "ws"; + +const CONNECT_TIMEOUT_MS = 8_000; +const FINISH_TIMEOUT_MS = 8_000; +/** ~100ms of 16 kHz 16-bit mono PCM (3.2KB per frame, well within the documented 16KB cap). */ +const FRAME_BYTES = 3_200; +/** Cap audio queued while (re)connecting (~3s at 16 kHz). */ +const MAX_PENDING_BYTES = 3 * ASR_SAMPLE_RATE * 2; +/** Cap the replay buffer used to restart an utterance after a drop (~30s). */ +const MAX_REPLAY_SAMPLES = 30 * ASR_SAMPLE_RATE; +/** Documented behaviour: reconnect once when the socket drops. */ +const MAX_RECONNECTS = 1; + +const WS_OPEN = 1; + +type WsLike = Pick & { + onopen: ((event: Event) => void) | null; + onclose: ((event: CloseEvent) => void) | null; + onerror: ((event: Event) => void) | null; + onmessage: ((event: { data?: unknown }) => void) | null; +}; + +/** Node's WebSocket accepts handshake headers through a non-standard option bag. */ +export type BailianWsConstructor = new (url: string, options?: { headers?: Record }) => WebSocket; + +export interface BailianWsAsrOptions { + /** Current 百炼 settings (re-read on each start). */ + getSettings: () => BailianAsrSettings | undefined; + /** Resolves the API key for an env var name (usually process.env). */ + getApiKey: (envName: string) => string | undefined; + /** 嘈杂环境参数(speech_noise_threshold / vad_model / 即时热词)。 */ + getNoise?: () => AsrNoiseSettings | undefined; + WebSocketCtor?: BailianWsConstructor; + log?: (message: string) => void; + /** Injectable UUID source (tests only). */ + uuid?: () => string; +} + +/** Split Float32 PCM into ≤16KB binary frames of ~100ms. */ +class PcmFramer { + private pending = new Float32Array(0); + + push(samples: Float32Array): ArrayBuffer[] { + if (!samples.length) return []; + const merged = new Float32Array(this.pending.length + samples.length); + merged.set(this.pending, 0); + merged.set(samples, this.pending.length); + const perFrame = FRAME_BYTES / 2; + const frames: ArrayBuffer[] = []; + let offset = 0; + while (merged.length - offset >= perFrame) { + frames.push(encodePcm16(merged.subarray(offset, offset + perFrame)).buffer as ArrayBuffer); + offset += perFrame; + } + this.pending = merged.slice(offset); + return frames; + } + + /** Final tail (may be shorter than 100ms); only used when the utterance ends. */ + flush(): ArrayBuffer | undefined { + if (!this.pending.length) return undefined; + const frame = encodePcm16(this.pending).buffer as ArrayBuffer; + this.pending = new Float32Array(0); + return frame; + } + + /** Drop the un-flushed tail (the replay rebuilds frames from the full history). */ + reset(): void { + this.pending = new Float32Array(0); + } +} + +/** + * One utterance over the duplex protocol. `ready` resolves after task-started + * so the manager only reports the provider as usable once audio can flow. + */ +class BailianWsSession implements AsrSession { + readonly providerId = "bailian-ws"; + private readonly sink: AsrEventSink; + private readonly config: ResolvedBailianConfig; + private readonly options: BailianWsAsrOptions; + private readonly taskId: string; + private readonly framer = new PcmFramer(); + /** Every sample received this utterance, kept for a single replay. */ + private readonly history: Float32Array[] = []; + private socket: WsLike | undefined; + private started = false; + private finished = false; + private failed = false; + private cancelled = false; + private stopRequested = false; + private reconnects = 0; + private pendingBytes = 0; + private readonly pending: ArrayBuffer[] = []; + private historySamples = 0; + private reconnectPromise: Promise | undefined; + + constructor(config: ResolvedBailianConfig, sink: AsrEventSink, options: BailianWsAsrOptions, taskId: string) { + this.config = config; + this.sink = sink; + this.options = options; + this.taskId = taskId; + } + + /** Establish the connection and start the task; retries once on failure. */ + async ready(): Promise { + try { + await this.connect(); + } catch (error) { + if (this.reconnects >= MAX_RECONNECTS || this.cancelled || this.stopRequested) throw error; + this.reconnects += 1; + this.options.log?.("[asr:bailian-ws] connect failed, retrying once"); + await this.connect(); + } + } + + private connect(): Promise { + const Ctor = this.options.WebSocketCtor || (Ws as unknown as BailianWsConstructor); + if (typeof Ctor === "undefined") return Promise.reject(new Error("当前环境不支持 WebSocket")); + return new Promise((resolve, reject) => { + let socket: WsLike; + try { + socket = new Ctor(this.config.wsUrl, { headers: { Authorization: `Bearer ${this.config.apiKey}`, "user-agent": "SecAgent" } }) as WsLike; + } catch (error) { + reject(new Error(`百炼实时识别连接创建失败:${error instanceof Error ? error.message : String(error)}`)); + return; + } + socket.binaryType = "arraybuffer"; + this.socket = socket; + this.started = false; + let settled = false; + const timer = setTimeout(() => { + if (settled) return; + settled = true; + try { socket.close(); } catch { /* may already be closed */ } + reject(new Error(`百炼实时识别连接超时(${CONNECT_TIMEOUT_MS / 1000}s),请检查网络,或改用非流式/本地识别`)); + }, CONNECT_TIMEOUT_MS); + const succeed = (): void => { + if (settled) return; + settled = true; + clearTimeout(timer); + resolve(); + }; + const abort = (message: string): void => { + if (settled) return; + settled = true; + clearTimeout(timer); + if (this.socket === socket) this.socket = undefined; + try { socket.close(); } catch { /* may already be closed */ } + reject(new Error(message)); + }; + + socket.onopen = () => { + this.options.log?.("[asr:bailian-ws] websocket opened"); + try { + // 嘈杂环境参数(官方文档): + // - speech_noise_threshold [-1,1]:-1 方向更不易漏音(教室场景建议 -0.2 ~ -0.6) + // - vad_model:far_field_meeting_16k 远场(希沃顶部麦克风)/ near_meeting_16k 近场 + // - vocabulary:即时热词(权重 50 = 超级热词) + const noise = this.options.getNoise?.(); + const parameters: Record = { format: "pcm", sample_rate: ASR_SAMPLE_RATE }; + if (typeof noise?.speechNoiseThreshold === "number") parameters.speech_noise_threshold = noise.speechNoiseThreshold; + if (noise?.vadModel) parameters.vad_model = noise.vadModel; + if (noise?.hotwords?.length) parameters.vocabulary = Object.fromEntries(noise.hotwords.map((word) => [word, 50])); + if (this.config.language) parameters.language_hints = [this.config.language]; + socket.send(JSON.stringify({ + header: { action: "run-task", task_id: this.taskId, streaming: "duplex" }, + payload: { + task_group: "audio", + task: "asr", + function: "recognition", + model: this.config.streamModel, + parameters, + input: {} + } + })); + } catch (error) { + abort(`百炼 run-task 发送失败:${error instanceof Error ? error.message : String(error)}`); + } + }; + + socket.onmessage = (event) => { + // ws 包的文本帧给 Buffer,全局 WebSocket 给 string——统一成 string。 + const raw = event.data; + const text = typeof raw === "string" ? raw : Buffer.isBuffer(raw) ? raw.toString("utf8") : typeof ArrayBuffer !== "undefined" && raw instanceof ArrayBuffer ? new TextDecoder().decode(raw) : undefined; + if (!text) return; // binary audio is never sent back by the server + let message: { header?: { event?: string; error_code?: string; error_message?: string }; payload?: { message?: string; output?: { sentence?: { text?: string; sentence_end?: boolean; heartbeat?: boolean } } } }; + try { message = JSON.parse(text) as typeof message; } catch { return; } + const header = message.header || {}; + if (header.event === "task-started") { + this.started = true; + this.options.log?.("[asr:bailian-ws] task started"); + for (const pcm of this.pending.splice(0)) { try { socket.send(pcm); } catch { /* socket may close mid-send */ } } + this.pendingBytes = 0; + succeed(); + return; + } + if (header.event === "result-generated") { + const sentence = message.payload?.output?.sentence; + if (sentence?.heartbeat) return; // keep-alive packet carries no transcript + const text = (sentence?.text || "").trim(); + if (!text) return; + this.sink({ type: sentence?.sentence_end === false ? "partial" : "final", text, provider: this.providerId }); + return; + } + if (header.event === "task-finished") { + this.finished = true; + return; + } + if (header.event === "task-failed") { + this.failed = true; + const detail = header.error_message || message.payload?.message || header.error_code || "未知错误"; + abort(`百炼实时识别任务失败:${detail}`); + this.sink({ type: "error", message: `百炼实时识别失败:${detail}` }); + } + }; + + socket.onerror = (event: Event) => { + const errorEvent = event as ErrorEvent; + this.options.log?.(`[asr:bailian-ws] websocket error state=${socket.readyState} message=${errorEvent.message || ""}`); + abort("百炼实时识别连接失败(网络波动或密钥/地址有误),已尝试改用其他识别通道"); + }; + + socket.onclose = (event: CloseEvent) => { + this.options.log?.(`[asr:bailian-ws] websocket closed code=${event.code}`); + if (this.socket === socket) this.socket = undefined; + if (!settled) { + // Handshake dropped before task-started: surface a retryable failure. + abort(`百炼实时识别连接已断开(code=${event.code}):请检查 API Key(无效或额度用尽)、WS URL 与网络`); + return; + } + // Mid-stream drop: restart the task once and replay the buffered audio. + this.scheduleReconnect(event.code); + }; + }); + } + + private pushHistory(samples: Float32Array): void { + this.history.push(samples); + this.historySamples += samples.length; + while (this.historySamples > MAX_REPLAY_SAMPLES && this.history.length > 1) { + const dropped = this.history.shift(); + if (dropped) this.historySamples -= dropped.length; + } + } + + private queueFrame(frame: ArrayBuffer): void { + if (this.pendingBytes + frame.byteLength > MAX_PENDING_BYTES) { + const dropped = this.pending.shift(); + if (dropped) this.pendingBytes -= dropped.byteLength; + } + this.pending.push(frame); + this.pendingBytes += frame.byteLength; + } + + /** Unexpected close after task-started: one reconnect attempt, then error. */ + private scheduleReconnect(code: number): void { + if (this.cancelled || this.stopRequested || this.finished || this.failed) return; + if (this.reconnects >= MAX_RECONNECTS) { + this.failed = true; + this.sink({ type: "error", message: `百炼实时识别连接中断(code=${code}),自动重连失败:请重试,或改用非流式/本地识别` }); + return; + } + this.options.log?.("[asr:bailian-ws] connection dropped mid-stream, reconnecting once"); + this.reconnectPromise = this.reconnect(code); + } + + private async reconnect(code: number): Promise { + try { + this.reconnects += 1; + this.prepareReplay(); + await this.connect(); + this.options.log?.("[asr:bailian-ws] reconnected, buffered audio replayed"); + } catch (error) { + this.failed = true; + const message = error instanceof Error ? error.message : String(error); + this.sink({ type: "error", message: `百炼实时识别连接中断(code=${code}),重连失败:${message}` }); + } finally { + this.reconnectPromise = undefined; + } + } + + /** Rebuild the already-received audio as ≤16KB frames for the new task. */ + private prepareReplay(): void { + this.framer.reset(); + this.pending.length = 0; + this.pendingBytes = 0; + const samples = mergeSamples(this.history); + if (!samples.length) return; + for (const frame of new PcmFramer().push(samples)) { + this.pending.push(frame); + this.pendingBytes += frame.byteLength; + } + this.options.log?.(`[asr:bailian-ws] replaying ${(samples.length / ASR_SAMPLE_RATE).toFixed(1)}s of audio`); + } + + push(samples: Float32Array): void { + if (this.cancelled || this.stopRequested) return; + this.pushHistory(samples); + const frames = this.framer.push(samples); + const socket = this.socket; + if (!socket || socket.readyState !== WS_OPEN || !this.started) { + // Queue framed audio while connecting (or reconnecting). + for (const frame of frames) this.queueFrame(frame); + return; + } + for (const frame of frames) { + try { socket.send(frame); } catch { /* socket may close between the state check and send */ } + } + } + + async stop(): Promise { + if (this.cancelled || this.stopRequested) return; + this.stopRequested = true; + // A drop may be mid-reconnect: finish the recovery attempt before closing. + if (this.reconnectPromise) await this.reconnectPromise; + const socket = this.socket; + if (socket && socket.readyState === WS_OPEN && this.started) { + const tail = this.framer.flush(); + try { + if (tail) socket.send(tail); + socket.send(JSON.stringify({ + header: { action: "finish-task", task_id: this.taskId, streaming: "duplex" }, + payload: { input: {} } + })); + } catch { /* socket may already be closing */ } + const deadline = Date.now() + FINISH_TIMEOUT_MS; + while (!this.finished && !this.failed && Date.now() < deadline) await new Promise((resolve) => setTimeout(resolve, 50)); + } + if (!this.finished && !this.failed) this.sink({ type: "error", message: "百炼实时识别未在预期时间内返回最终结果" }); + this.close("finished"); + this.sink({ type: "stopped" }); + } + + cancel(): void { + this.cancelled = true; + this.stopRequested = true; + this.close("cancelled"); + } + + private close(reason: string): void { + const socket = this.socket; + this.socket = undefined; + try { socket?.close(1000, reason); } catch { /* may already be closed */ } + } +} + +export class BailianWsAsrProvider implements AsrProvider { + readonly id = "bailian-ws"; + readonly label = "阿里云百炼(WebSocket 流式)"; + private readonly options: BailianWsAsrOptions; + + constructor(options: BailianWsAsrOptions) { + this.options = options; + } + + private displayName(): string { + return bailianDisplayName(this.options.getSettings(), this.label); + } + + isConfigured(): boolean { + return resolveBailianConfig(this.options.getSettings(), this.options.getApiKey) !== null; + } + + private uuid(): string { + const injected = this.options.uuid?.(); + if (injected) return injected; + try { return crypto.randomUUID(); } catch { /* runtime without randomUUID */ } + return `task-${Date.now()}-${Math.floor(Math.random() * 0xffffff).toString(16)}`; + } + + async test(): Promise { + const config = resolveBailianConfig(this.options.getSettings(), this.options.getApiKey); + if (!config) return { ok: false, message: `环境变量 ${this.options.getSettings()?.apiKeyEnv || "BAILIAN_API_KEY"} 中没有 API Key(可在下方填写并保存到工作区 .env)` }; + const endpoint = config.wsUrl.replace(/^wss?:\/\//, ""); + return { ok: true, message: `${this.displayName()} 已配置(实时连接在开始说话时建立:${endpoint})` }; + } + + async start(sink: AsrEventSink): Promise { + const config = resolveBailianConfig(this.options.getSettings(), this.options.getApiKey); + if (!config) throw new Error("阿里云百炼实时语音识别未配置:请在设置中填写 API Key(保存到工作区 .env 的 BAILIAN_API_KEY)"); + this.options.log?.(`[asr:bailian-ws] connecting ${config.wsUrl.replace(/^wss?:\/\//, "")} model=${config.streamModel}`); + const session = new BailianWsSession(config, sink, this.options, this.uuid()); + try { + await session.ready(); + } catch (error) { + session.cancel(); + throw error; + } + return session; + } +} \ No newline at end of file diff --git a/src/asr/bailian.test.ts b/src/asr/bailian.test.ts new file mode 100644 index 0000000..726b5c7 --- /dev/null +++ b/src/asr/bailian.test.ts @@ -0,0 +1,558 @@ +/** + * 阿里云百炼 ASR 测试 — 通道 A(chat/completions)与通道 B(实时 WebSocket)。 + * + * 断言严格对齐官方协议(禁止凭记忆编造): + * A: POST {baseUrl}/chat/completions(绝不触碰 /audio/transcriptions), + * `stream:false`、`asr_options` 在 body 顶层、音频为 data:audio/wav;base64, Data URL; + * B: run-task(streaming=duplex)+ ≤16KB 的 ~100ms PCM 二进制帧 + finish-task。 + * + * Docs: + * https://help.aliyun.com/zh/model-studio/qwen-asr-api-reference + * https://help.aliyun.com/zh/model-studio/fun-asr-realtime-websocket-api + */ +import test from "node:test"; +import assert from "node:assert/strict"; +import { BailianHttpAsrProvider } from "./bailian-http.js"; +import { BailianWsAsrProvider, type BailianWsConstructor } from "./bailian-ws.js"; +import { OpenAiHttpAsrProvider } from "./openai-http.js"; +import { resolveBailianConfig } from "./bailian-config.js"; +import { AsrManager } from "./manager.js"; +import type { AsrEvent, AsrEventSink, AsrProvider, AsrSession } from "./types.js"; +import type { BailianAsrSettings } from "./settings.js"; + +const BAILIAN_SETTINGS: BailianAsrSettings = { + name: "百炼测试", + apiKeyEnv: "BAILIAN_API_KEY", + baseUrl: "https://ws-test123.cn-beijing.maas.aliyuncs.com/compatible-mode/v1", + wsUrl: "wss://ws-test123.cn-beijing.maas.aliyuncs.com/api-ws/v1/inference", + model: "qwen3-asr-flash", + streamModel: "qwen-audio-3.1-asr-flash-streaming", + language: "zh", + enableItn: false +}; + +/* ------------------------------------------------------------------ */ +/* Channel A — non-streaming chat/completions */ +/* ------------------------------------------------------------------ */ + +interface CapturedRequest { + url: string; + auth: string | null; + body: Record; +} + +function fakeFetch(responses: Array<{ status: number; body: string }>): { fetchImpl: typeof fetch; requests: CapturedRequest[] } { + const requests: CapturedRequest[] = []; + let call = 0; + const fetchImpl = (async (input: RequestInfo | URL, init?: RequestInit) => { + requests.push({ + url: String(input), + auth: (init?.headers as Record | undefined)?.Authorization ?? null, + body: JSON.parse(typeof init?.body === "string" ? init.body : "{}") as Record + }); + const response = responses[Math.min(call, responses.length - 1)]; + call += 1; + return new Response(response.body, { status: response.status, headers: { "Content-Type": "application/json" } }); + }) as unknown as typeof fetch; + return { fetchImpl, requests }; +} + +function httpProvider( + settings: BailianAsrSettings | undefined, + apiKey: string, + responses: Array<{ status: number; body: string }> +): { provider: BailianHttpAsrProvider; requests: CapturedRequest[] } { + const { fetchImpl, requests } = fakeFetch(responses); + return { + provider: new BailianHttpAsrProvider({ getSettings: () => settings, getApiKey: () => apiKey || undefined, fetchImpl, rateLimitBaseDelayMs: 1 }), + requests + }; +} + +function audioDataUrl(body: Record): string { + const messages = body.messages as Array<{ role: string; content: Array<{ type: string; input_audio: { data: string } }> }>; + assert.equal(messages[0].role, "user"); + assert.equal(messages[0].content[0].type, "input_audio"); + return messages[0].content[0].input_audio.data; +} + +const CHAT_REPLY = JSON.stringify({ choices: [{ message: { content: "你好世界" }, finish_reason: "stop" }], usage: {} }); + +test("config resolution prefers settings, then .env, then documented defaults", () => { + const env = { + BAILIAN_API_KEY: "sk-env", + BAILIAN_BASE_URL: "https://env.example.com/compatible-mode/v1", + BAILIAN_WS_URL: "wss://env.example.com/api-ws/v1/inference", + BAILIAN_ASR_MODEL: "env-model", + BAILIAN_STREAM_MODEL: "env-stream-model", + BAILIAN_LANGUAGE: "en", + BAILIAN_ENABLE_ITN: "true" + }; + const getApiKey = (name: string): string | undefined => env[name as keyof typeof env]; + + assert.equal(resolveBailianConfig(undefined, () => undefined, {}), null); + + const fromEnv = resolveBailianConfig(undefined, getApiKey, env); + assert.ok(fromEnv); + assert.equal(fromEnv.apiKey, "sk-env"); + assert.equal(fromEnv.baseUrl, "https://env.example.com/compatible-mode/v1"); + assert.equal(fromEnv.wsUrl, "wss://env.example.com/api-ws/v1/inference"); + assert.equal(fromEnv.model, "env-model"); + assert.equal(fromEnv.streamModel, "env-stream-model"); + assert.equal(fromEnv.language, "en"); + assert.equal(fromEnv.enableItn, true); + + const fromSettings = resolveBailianConfig(BAILIAN_SETTINGS, getApiKey, env); + assert.ok(fromSettings); + assert.equal(fromSettings.baseUrl, BAILIAN_SETTINGS.baseUrl); + assert.equal(fromSettings.model, "qwen3-asr-flash"); + assert.equal(fromSettings.language, "zh"); + + // 空 env 回落到文档默认值(含公共域名回退)。 + const defaults = resolveBailianConfig(undefined, () => "sk-test", {}); + assert.ok(defaults); + assert.equal(defaults.baseUrl, "https://dashscope.aliyuncs.com/compatible-mode/v1"); + assert.equal(defaults.wsUrl, "wss://dashscope.aliyuncs.com/api-ws/v1/inference"); + assert.equal(defaults.model, "qwen3-asr-flash"); + assert.equal(defaults.streamModel, "qwen-audio-3.1-asr-flash-streaming"); + assert.equal(defaults.enableItn, false); + assert.equal("language" in defaults, false); +}); + +test("channel A posts a non-streaming chat/completions request with a WAV Data URL and asr_options", async () => { + const { provider, requests } = httpProvider(BAILIAN_SETTINGS, "sk-test", [{ status: 200, body: CHAT_REPLY }]); + const result = await provider.test(); + + assert.equal(requests.length, 1); + const request = requests[0]; + assert.equal(request.url, `${BAILIAN_SETTINGS.baseUrl}/chat/completions`); + // 明确禁止:百炼不支持 OpenAI Whisper 范式的转写端点。 + assert.equal(request.url.includes("/audio/transcriptions"), false); + assert.equal(request.auth, "Bearer sk-test"); + assert.equal(request.body.model, "qwen3-asr-flash"); + // 非流式场景固定 stream=false。 + assert.equal(request.body.stream, false); + // asr_options 是百炼扩展参数,直连 HTTP 时放 body 顶层。 + assert.deepEqual(request.body.asr_options, { language: "zh", enable_itn: false }); + assert.equal("asr_options" in ((request.body.messages as Array<{ content: Array> }>)[0].content[0]), false); + + const dataUrl = audioDataUrl(request.body); + assert.ok(dataUrl.startsWith("data:audio/wav;base64,")); + const wav = Buffer.from(dataUrl.slice("data:audio/wav;base64,".length), "base64"); + assert.equal(wav.subarray(0, 4).toString("ascii"), "RIFF"); + assert.equal(wav.subarray(8, 12).toString("ascii"), "WAVE"); + assert.equal(wav.readUInt32LE(24), 16_000); // sample rate + // 250ms probe → 44-byte header + 16000Hz * 2B * 0.25s。 + assert.equal(wav.length, 44 + 8_000); + + assert.equal(result.ok, true); + assert.match(result.message, /你好世界/); +}); + +test("channel A omits the language hint when it is not configured (auto detect)", async () => { + const { provider, requests } = httpProvider({ ...BAILIAN_SETTINGS, language: undefined }, "sk-test", [{ status: 200, body: CHAT_REPLY }]); + const previous = process.env.BAILIAN_LANGUAGE; + process.env.BAILIAN_LANGUAGE = ""; + try { + await provider.test(); + } finally { + if (previous === undefined) delete process.env.BAILIAN_LANGUAGE; + else process.env.BAILIAN_LANGUAGE = previous; + } + const options = requests[0].body.asr_options as Record; + assert.deepEqual(options, { enable_itn: false }); + assert.equal("language" in options, false); +}); + +test("a 1-second utterance is sent as one Data URL and yields the final transcript", async () => { + const { provider, requests } = httpProvider(BAILIAN_SETTINGS, "sk-test", [{ status: 200, body: CHAT_REPLY }]); + const events: AsrEvent[] = []; + const session = await provider.start((event) => events.push(event)); + session.push(new Float32Array(16_000)); // exactly 1 second + await session.stop(); + + assert.equal(requests.length, 1); + const dataUrl = audioDataUrl(requests[0].body); + assert.ok(dataUrl.startsWith("data:audio/wav;base64,")); + const wav = Buffer.from(dataUrl.slice("data:audio/wav;base64,".length), "base64"); + assert.equal(wav.length, 44 + 16_000 * 2); + assert.equal(wav.readUInt32LE(40), 16_000 * 2); // data chunk size + + const final = events.find((event) => event.type === "final"); + assert.equal(final && final.type === "final" ? final.text : "", "你好世界"); + assert.equal(events[events.length - 1].type, "stopped"); +}); + +test("channel A start() rejects with the .env variable name when no key is stored", async () => { + const { provider } = httpProvider(BAILIAN_SETTINGS, "", []); + assert.equal(provider.isConfigured(), false); + await assert.rejects(() => provider.start(() => {}), /BAILIAN_API_KEY/); +}); + +test("401/403 explains the invalid key or exhausted quota and points at other providers", async () => { + const { provider } = httpProvider(BAILIAN_SETTINGS, "sk-bad", [{ status: 401, body: JSON.stringify({ error: { message: "invalid api key" } }) }]); + const result = await provider.test(); + assert.equal(result.ok, false); + assert.match(result.message, /API Key 无效|额度/); + + const events: AsrEvent[] = []; + const session = await provider.start((event) => events.push(event)); + session.push(new Float32Array(16_000)); + await session.stop(); + const error = events.find((event) => event.type === "error"); + assert.ok(error && error.type === "error"); + if (error.type === "error") assert.match(error.message, /切换识别服务|API Key 无效/); +}); + +test("404 points at the missing /compatible-mode/v1 suffix", async () => { + const { provider } = httpProvider({ ...BAILIAN_SETTINGS, baseUrl: "https://ws-test123.cn-beijing.maas.aliyuncs.com" }, "sk-test", [{ status: 404, body: "" }]); + const result = await provider.test(); + assert.equal(result.ok, false); + assert.match(result.message, /compatible-mode\/v1/); +}); + +test("429 retries with exponential backoff and succeeds within the retry budget", async () => { + const { provider, requests } = httpProvider(BAILIAN_SETTINGS, "sk-test", [ + { status: 429, body: "" }, + { status: 429, body: "" }, + { status: 200, body: CHAT_REPLY } + ]); + const result = await provider.test(); + assert.equal(result.ok, true); + assert.equal(requests.length, 3); +}); + +test("429 exhausts three retries and reports the rate limit to the sink", async () => { + const { provider, requests } = httpProvider(BAILIAN_SETTINGS, "sk-test", [{ status: 429, body: "" }]); + const events: AsrEvent[] = []; + const session = await provider.start((event) => events.push(event)); + session.push(new Float32Array(16_000)); + await session.stop(); + assert.equal(requests.length, 4); // initial attempt + 3 retries + const error = events.find((event) => event.type === "error"); + assert.ok(error && error.type === "error"); + if (error.type === "error") assert.match(error.message, /429|限流/); +}); + +/* ------------------------------------------------------------------ */ +/* Channel B — realtime WebSocket (duplex) */ +/* ------------------------------------------------------------------ */ + +interface RunTaskMessage { + header: { action: string; task_id: string; streaming: string }; + payload: { + task_group: string; + task: string; + function: string; + model: string; + parameters: Record; + input: Record; + }; +} + +class FakeSocket { + static instances: FakeSocket[] = []; + static onOpen: (socket: FakeSocket) => void = (socket) => { + socket.readyState = 1; + socket.onopen?.(new Event("open")); + }; + static onRunTask: (socket: FakeSocket) => void = (socket) => { + setTimeout(() => socket.emit({ header: { event: "task-started", task_id: "test-task-id" }, payload: {} }), 0); + }; + + static reset(): void { + FakeSocket.instances = []; + FakeSocket.onOpen = (socket) => { + socket.readyState = 1; + socket.onopen?.(new Event("open")); + }; + FakeSocket.onRunTask = (socket) => { + setTimeout(() => socket.emit({ header: { event: "task-started", task_id: "test-task-id" }, payload: {} }), 0); + }; + } + + readyState = 0; + binaryType = "blob"; + onopen: ((event: Event) => void) | null = null; + onclose: ((event: CloseEvent) => void) | null = null; + onerror: ((event: Event) => void) | null = null; + onmessage: ((event: MessageEvent) => void) | null = null; + readonly url: string; + readonly headers: Record | undefined; + readonly sent: Array = []; + closed = false; + + constructor(url: string, options?: { headers?: Record }) { + this.url = url; + this.headers = options?.headers; + FakeSocket.instances.push(this); + setTimeout(() => { if (!this.closed) FakeSocket.onOpen(this); }, 0); + } + + send(data: string | ArrayBuffer): void { + if (this.closed) throw new Error("socket is closed"); + this.sent.push(data); + if (typeof data !== "string") return; + if (data.includes("\"run-task\"")) FakeSocket.onRunTask(this); + else if (data.includes("\"finish-task\"")) { + setTimeout(() => this.emit({ header: { event: "task-finished", task_id: "test-task-id" }, payload: {} }), 0); + } + } + + close(code?: number, reason?: string): void { + void reason; + void code; + this.closed = true; + this.readyState = 3; + } + + emit(message: unknown): void { + this.onmessage?.({ data: JSON.stringify(message) } as MessageEvent); + } + + /** Simulate an unexpected drop (network jitter, 401 handshake……). */ + drop(code = 1006): void { + this.closed = true; + this.readyState = 3; + this.onclose?.({ code } as CloseEvent); + } + + binaryFrames(): ArrayBuffer[] { + return this.sent.filter((item): item is ArrayBuffer => typeof item !== "string"); + } + + textMessages(): string[] { + return this.sent.filter((item): item is string => typeof item === "string"); + } +} + +function wsProvider(settings: BailianAsrSettings | undefined = BAILIAN_SETTINGS, apiKey = "sk-test"): BailianWsAsrProvider { + return new BailianWsAsrProvider({ + getSettings: () => settings, + getApiKey: () => apiKey || undefined, + WebSocketCtor: FakeSocket as unknown as BailianWsConstructor, + uuid: () => "test-task-id" + }); +} + +async function waitFor(condition: () => boolean, timeoutMs = 2_000): Promise { + const deadline = Date.now() + timeoutMs; + while (!condition() && Date.now() < deadline) await new Promise((resolve) => setTimeout(resolve, 5)); + assert.ok(condition(), "condition not met before the timeout"); +} + +test("channel B runs the documented duplex handshake and streams ~100ms PCM frames", async () => { + FakeSocket.reset(); + const events: AsrEvent[] = []; + const session = await wsProvider().start((event) => events.push(event)); + const socket = FakeSocket.instances.at(-1); + assert.ok(socket); + + assert.equal(socket.url, BAILIAN_SETTINGS.wsUrl); + assert.equal(socket.url.includes("/audio/transcriptions"), false); + assert.equal(socket.headers?.Authorization, "Bearer sk-test"); + assert.ok(socket.headers?.["user-agent"]); + + const runTask = JSON.parse(socket.textMessages()[0]) as RunTaskMessage; + assert.equal(runTask.header.action, "run-task"); + assert.equal(runTask.header.task_id, "test-task-id"); + assert.equal(runTask.header.streaming, "duplex"); // 文档固定值(非 "out") + assert.equal(runTask.payload.task_group, "audio"); + assert.equal(runTask.payload.task, "asr"); + assert.equal(runTask.payload.function, "recognition"); + assert.equal(runTask.payload.model, "qwen-audio-3.1-asr-flash-streaming"); + assert.deepEqual(runTask.payload.parameters, { format: "pcm", sample_rate: 16_000, language_hints: ["zh"] }); + + session.push(new Float32Array(16_000)); // 1s → ten 3.2KB frames + const frames = socket.binaryFrames(); + assert.equal(frames.length, 10); + for (const frame of frames) { + assert.equal(frame.byteLength, 3_200); + assert.ok(frame.byteLength <= 16 * 1024, "frame exceeds the 16KB cap"); + assert.ok(frame.byteLength >= 1024, "frame below the 1KB floor was sent mid-stream"); + } + + // 中间结果 → partial,最终结果 → final,心跳包忽略。 + socket.emit({ header: { event: "result-generated", task_id: "test-task-id" }, payload: { output: { sentence: { text: "你好", sentence_end: false } } } }); + socket.emit({ header: { event: "result-generated", task_id: "test-task-id" }, payload: { output: { sentence: { text: "你好世界", sentence_end: true, begin_time: 0, end_time: 900 } } } }); + socket.emit({ header: { event: "result-generated", task_id: "test-task-id" }, payload: { output: { sentence: { text: "心跳", heartbeat: true, sentence_id: 0 } } } }); + + await session.stop(); + const types = events.map((event) => event.type); + assert.deepEqual(types, ["partial", "final", "stopped"]); + assert.equal(events[0].type === "partial" ? events[0].text : "", "你好"); + assert.equal(events[1].type === "final" ? events[1].text : "", "你好世界"); + assert.equal(events.some((event) => event.type !== "stopped" && "text" in event && event.text === "心跳"), false); + + const finishTask = JSON.parse(socket.textMessages().at(-1) || "{}") as { header: { action: string; task_id: string } }; + assert.equal(finishTask.header.action, "finish-task"); + assert.equal(finishTask.header.task_id, "test-task-id"); + assert.equal(socket.binaryFrames().length, 10); // 无多余尾帧 +}); + +test("channel B buffers sub-100ms tails until stop instead of bursting small frames", async () => { + FakeSocket.reset(); + const session = await wsProvider().start(() => {}); + const socket = FakeSocket.instances.at(-1); + assert.ok(socket); + + session.push(new Float32Array(1_600)); // 100ms → 1 full frame + session.push(new Float32Array(400)); // 25ms → buffered + assert.equal(socket.binaryFrames().length, 1); + + await session.stop(); + const frames = socket.binaryFrames(); + assert.equal(frames.length, 2); + assert.equal(frames[0].byteLength, 3_200); + assert.equal(frames[1].byteLength, 800); // 尾帧仅在结束 utterance 时单独发送 + assert.match(socket.textMessages().at(-1) || "", /finish-task/); +}); + +test("channel B surfaces task-failed with the server error message", async () => { + FakeSocket.reset(); + FakeSocket.onRunTask = (socket) => { + setTimeout(() => socket.emit({ + header: { event: "task-failed", task_id: "test-task-id", error_code: "CLIENT_ERROR", error_message: "request timeout after 23 seconds." }, + payload: {} + }), 0); + }; + try { + const events: AsrEvent[] = []; + await assert.rejects(() => wsProvider().start((event) => events.push(event)), /request timeout after 23 seconds/); + const error = events.find((event) => event.type === "error"); + assert.ok(error && error.type === "error"); + } finally { + FakeSocket.reset(); + } +}); + +test("channel B handshake failures recommend checking key and WS URL", async () => { + FakeSocket.reset(); + FakeSocket.onOpen = (socket) => { socket.drop(1006); }; + try { + await assert.rejects(() => wsProvider().start(() => {}), /API Key|连接已断开/); + } finally { + FakeSocket.reset(); + } +}); + +test("channel B reconnects once after a mid-stream drop and replays the buffered audio", async () => { + FakeSocket.reset(); + const events: AsrEvent[] = []; + const session = await wsProvider().start((event) => events.push(event)); + const first = FakeSocket.instances.at(-1); + assert.ok(first); + + session.push(new Float32Array(8_000)); // 0.5s → 5 frames + assert.equal(first.binaryFrames().length, 5); + + first.drop(1006); + session.push(new Float32Array(8_000)); // queued while reconnecting + await waitFor(() => FakeSocket.instances.length === 2 && FakeSocket.instances[1].binaryFrames().length === 10); + + const second = FakeSocket.instances[1]; + const replay = second.binaryFrames(); + assert.equal(replay.length, 10); // 1s of audio (0.5s replayed + 0.5s queued) as ≤16KB frames + for (const frame of replay) assert.ok(frame.byteLength <= 16 * 1024); + + second.emit({ header: { event: "result-generated", task_id: "test-task-id" }, payload: { output: { sentence: { text: "重连成功", sentence_end: true } } } }); + await session.stop(); + assert.equal(events.some((event) => event.type === "final" && event.text === "重连成功"), true); + assert.equal(events[events.length - 1].type, "stopped"); + assert.equal(events.some((event) => event.type === "error"), false); +}); + +test("channel B gives up after a second drop so the manager can fall back", async () => { + FakeSocket.reset(); + const events: AsrEvent[] = []; + const session = await wsProvider().start((event) => events.push(event)); + const first = FakeSocket.instances.at(-1); + assert.ok(first); + session.push(new Float32Array(3_200)); + first.drop(1006); + await waitFor(() => FakeSocket.instances.length === 2 && FakeSocket.instances[1].binaryFrames().length > 0); + + FakeSocket.instances[1].drop(1006); + await waitFor(() => events.some((event) => event.type === "error")); + const error = events.find((event) => event.type === "error"); + assert.ok(error && error.type === "error"); + if (error.type === "error") assert.match(error.message, /重连失败|改用非流式/); + session.cancel(); +}); + +test("channel B cancel closes the socket without emitting results", async () => { + FakeSocket.reset(); + const events: AsrEvent[] = []; + const session = await wsProvider().start((event) => events.push(event)); + const socket = FakeSocket.instances.at(-1); + assert.ok(socket); + session.push(new Float32Array(16_000)); + session.cancel(); + assert.equal(socket.closed, true); + session.push(new Float32Array(16_000)); // ignored after cancel + await session.stop(); + assert.equal(events.some((event) => event.type === "final" || event.type === "partial" || event.type === "stopped"), false); +}); + +test("channel B test() reports the endpoint and requires an API key", async () => { + const ready = await wsProvider().test(); + assert.equal(ready.ok, true); + assert.match(ready.message, /api-ws\/v1\/inference/); + const missing = await wsProvider(BAILIAN_SETTINGS, "").test(); + assert.equal(missing.ok, false); + assert.match(missing.message, /BAILIAN_API_KEY/); +}); + +/* ------------------------------------------------------------------ */ +/* Orchestration — new provider kinds and regressions */ +/* ------------------------------------------------------------------ */ + +function fakeProvider(id: string, configured = true): AsrProvider { + return { + id, + label: id, + isConfigured: () => configured, + start: async (sink: AsrEventSink): Promise => ({ + providerId: id, + push: () => {}, + stop: async () => { sink({ type: "stopped" }); }, + cancel: () => {} + }) + }; +} + +test("manager chains the 百炼 channels and leaves existing kinds untouched", () => { + FakeSocket.reset(); + let kind: "auto" | "bailian" | "bailian-ws" = "auto"; + const manager = new AsrManager({ getProviderKind: () => kind }); + manager.register(fakeProvider("openai")); + manager.register(fakeProvider("official")); + manager.register(new BailianHttpAsrProvider({ getSettings: () => BAILIAN_SETTINGS, getApiKey: () => "sk-test" })); + manager.register(new BailianWsAsrProvider({ getSettings: () => BAILIAN_SETTINGS, getApiKey: () => "sk-test", WebSocketCtor: FakeSocket as unknown as BailianWsConstructor })); + manager.register(fakeProvider("local")); + + assert.deepEqual(manager.chain(), ["openai", "official", "local"]); + kind = "bailian"; + assert.deepEqual(manager.chain(), ["bailian", "local"]); + kind = "bailian-ws"; + assert.deepEqual(manager.chain(), ["bailian-ws", "bailian", "local"]); + + // 未配置 Key 时百炼通道被跳过,本地模型仍然可用。 + const unconfigured = new AsrManager({ getProviderKind: () => "bailian-ws" }); + unconfigured.register(new BailianHttpAsrProvider({ getSettings: () => undefined, getApiKey: () => undefined })); + unconfigured.register(new BailianWsAsrProvider({ getSettings: () => undefined, getApiKey: () => undefined })); + unconfigured.register(fakeProvider("local")); + assert.deepEqual(unconfigured.chain(), ["local"]); +}); + +test("third-party OpenAI-compatible presets are unchanged (SiliconFlow regression)", async () => { + const captured = fakeFetch([{ status: 200, body: JSON.stringify({ text: "回归通过" }) }]); + const provider = new OpenAiHttpAsrProvider({ + getSettings: () => ({ name: "SiliconFlow SenseVoice", baseUrl: "https://api.siliconflow.cn/v1", apiKeyEnv: "SILICONFLOW_API_KEY", model: "FunAudioLLM/SenseVoiceSmall" }), + getApiKey: () => "sk-sf", + fetchImpl: captured.fetchImpl + }); + const result = await provider.test(); + assert.equal(result.ok, true); + assert.equal(captured.requests[0].url, "https://api.siliconflow.cn/v1/audio/transcriptions"); +}); \ No newline at end of file diff --git a/src/asr/local-sensevoice.ts b/src/asr/local-sensevoice.ts new file mode 100644 index 0000000..efee22f --- /dev/null +++ b/src/asr/local-sensevoice.ts @@ -0,0 +1,148 @@ +/** + * Local ASR upgrade pack — SenseVoice (offline, sherpa-onnx). + * + * The bundled streaming zipformer is tiny and fast, but its accuracy drops in + * noisy classrooms. SenseVoice-small is the optional "本地增强包": a larger + * offline model (中/英/日/韩/粤) that recognises the whole utterance after the + * user stops talking. It is NOT bundled — `scripts/fetch-sensevoice-pack.mjs` + * downloads it (~230 MB int8) into `models/sense-voice/`. + * + * sherpa-onnx Node API (official examples): + * const recognizer = new sherpa.OnlineRecognizer… — offline variant: + * recognizer = new sherpa.OfflineRecognizer({ modelConfig: { senseVoice: + * { model, useInverseTextNormalization }, tokens, numThreads, provider } }) + * stream = recognizer.createStream(); stream.acceptWaveform(16000, samples); + * recognizer.decode(stream); recognizer.getResult(stream).text + */ +import fs from "node:fs"; +import path from "node:path"; +import type { AsrEventSink, AsrProvider, AsrSession, AsrTestResult } from "./types.js"; +import { loadSherpaOnnx } from "./sherpa-loader.js"; +import { mergeSamples } from "./wav.js"; + +const PACK_DIR = "sense-voice"; +const MODEL_FILES = ["model.int8.onnx", "tokens.txt"] as const; + +export interface LocalSenseVoiceOptions { + /** Extra directories that may contain `models/sense-voice`. */ + extraRoots?: string[]; + language?: string; // "" (auto) | zh | en | ja | ko | yue + numThreads?: number; + /** Partial decode cadence in ms; 0 disables interim results. Default 3500. */ + partialIntervalMs?: number; + log?(message: string): void; +} + +function electronResourcesPath(): string | undefined { + return (process as NodeJS.Process & { resourcesPath?: string }).resourcesPath; +} + +/** Directories that may hold the optional SenseVoice pack. */ +function candidateRoots(options: LocalSenseVoiceOptions): string[] { + const resourcesPath = electronResourcesPath(); + return [ + ...(process.env.SECAGENT_ASR_MODELS_ROOT ? [process.env.SECAGENT_ASR_MODELS_ROOT] : []), + ...(resourcesPath ? [resourcesPath] : []), + ...(options.extraRoots || []), + process.cwd() + ]; +} + +export function resolveSenseVoicePack(options: LocalSenseVoiceOptions = {}): string | undefined { + for (const root of candidateRoots(options)) { + const dir = path.join(root, "models", PACK_DIR); + if (MODEL_FILES.every((file) => fs.existsSync(path.join(dir, file)))) return dir; + } + return undefined; +} + +interface SenseVoiceRecognizerLike { + createStream(): { acceptWaveform(sampleRate: number, samples: Float32Array): void }; + decode(stream: unknown): void; + getResult(stream: unknown): { text: string }; +} + +export class LocalSenseVoiceProvider implements AsrProvider { + readonly id = "local-pro"; + readonly label = "本地增强(SenseVoice 附加包)"; + private readonly options: LocalSenseVoiceOptions; + + constructor(options: LocalSenseVoiceOptions = {}) { this.options = options; } + + isConfigured(): boolean { return Boolean(resolveSenseVoicePack(this.options)); } + + displayName(): string { return this.options.language ? `${this.label} · ${this.options.language}` : this.label; } + + async test(): Promise { + const dir = resolveSenseVoicePack(this.options); + if (!dir) return { ok: false, message: "未安装附加包:在设备上运行 scripts/fetch-sensevoice-pack.mjs 下载(约 230 MB,含 model.int8.onnx 与 tokens.txt)" }; + return { ok: true, message: `SenseVoice 附加包已就绪(${dir};整句离线识别,嘈杂环境准确率显著高于默认小模型)` }; + } + + async start(sink: AsrEventSink): Promise { + const dir = resolveSenseVoicePack(this.options); + if (!dir) throw new Error("SenseVoice 附加包未安装:请先运行 scripts/fetch-sensevoice-pack.mjs"); + const sherpa = await loadSherpaOnnx(); + // sherpa-onnx 的类型声明未覆盖 OfflineRecognizer(SenseVoice),运行时存在。 + const OfflineRecognizer = (sherpa as unknown as { OfflineRecognizer: new (config: unknown) => unknown }).OfflineRecognizer; + const recognizer = new OfflineRecognizer({ + modelConfig: { + senseVoice: { model: path.join(dir, "model.int8.onnx"), useInverseTextNormalization: 1 }, + tokens: path.join(dir, "tokens.txt"), + numThreads: this.options.numThreads ?? 2, + debug: 0, + provider: "cpu" + } + }) as unknown as SenseVoiceRecognizerLike; + this.options.log?.("[asr:local-pro] sensevoice ready"); + sink({ type: "ready", provider: "local-pro" }); + return new SenseVoiceSession(recognizer, sink, this.options); + } +} + +class SenseVoiceSession implements AsrSession { + readonly providerId = "local-pro"; + private readonly recognizer: SenseVoiceRecognizerLike; + private readonly sink: AsrEventSink; + private readonly options: LocalSenseVoiceOptions; + private readonly buffers: Float32Array[] = []; + private timer: NodeJS.Timeout | undefined; + private stopped = false; + + constructor(recognizer: SenseVoiceRecognizerLike, sink: AsrEventSink, options: LocalSenseVoiceOptions) { + this.recognizer = recognizer; + this.sink = sink; + this.options = options; + const interval = options.partialIntervalMs ?? 3500; + if (interval > 0) this.timer = setInterval(() => this.decode(false), interval); + } + + push(samples: Float32Array): void { if (!this.stopped) this.buffers.push(samples); } + + private decode(final: boolean): void { + const merged = mergeSamples(this.buffers); + if (!merged.length) return; + if (!final && merged.length < 16000) return; // wait for ≥1s of audio for interim decodes + try { + const stream = this.recognizer.createStream(); + stream.acceptWaveform(16000, merged); + this.recognizer.decode(stream); + const text = (this.recognizer.getResult(stream).text || "").trim(); + if (text) this.sink({ type: final ? "final" : "partial", text, provider: "local-pro" }); + } catch (error) { + this.sink({ type: "log", message: `[asr:local-pro] decode 失败:${error instanceof Error ? error.message : String(error)}` }); + } + } + + async stop(): Promise { + if (this.stopped) return; + this.stopped = true; + if (this.timer) clearInterval(this.timer); + this.decode(true); + } + + cancel(): void { + this.stopped = true; + if (this.timer) clearInterval(this.timer); + } +} diff --git a/src/asr/local-sherpa.ts b/src/asr/local-sherpa.ts new file mode 100644 index 0000000..14e194f --- /dev/null +++ b/src/asr/local-sherpa.ts @@ -0,0 +1,148 @@ +/** + * Local offline ASR backed by the bundled sherpa-onnx streaming zipformer. + * + * The recognizer is created lazily on first use and kept for the app lifetime; + * a load failure surfaces as a rejected `start()` so the manager can fall back + * to a cloud provider instead of leaving speech dead. + */ +import fs from "node:fs"; +import path from "node:path"; +import type { AsrEventSink, AsrProvider, AsrSession } from "./types.js"; +import { loadSherpaOnnx } from "./sherpa-loader.js"; + +const RECOGNIZER_DIR = "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23"; + +type Recognizer = ReturnType; +type RecognizerStream = ReturnType; + +export interface LocalAsrOptions { + /** Additional directories to search for the bundled `models/` folder. */ + extraRoots?: string[]; + log?: (message: string) => void; +} + +/** Electron exposes `process.resourcesPath`; plain Node does not. */ +function electronResourcesPath(): string | undefined { + return (process as NodeJS.Process & { resourcesPath?: string }).resourcesPath; +} + +function resolveModelRoot(options: LocalAsrOptions): string { + const resourcesPath = electronResourcesPath(); + const candidates = [ + ...(options.extraRoots || []), + ...(resourcesPath ? [resourcesPath] : []), + process.cwd(), + ...(process.env.SECAGENT_ASR_MODELS_ROOT ? [process.env.SECAGENT_ASR_MODELS_ROOT] : []) + ]; + for (const root of candidates) { + const candidate = path.join(root, "models", RECOGNIZER_DIR); + if (fs.existsSync(candidate)) return candidate; + } + const searched = candidates.map((root) => path.join(root, "models")).join("、"); + throw new Error(`找不到本地语音模型 ${RECOGNIZER_DIR},已搜索:${searched || "(无候选目录)"}`); +} + +export class LocalSherpaAsrProvider implements AsrProvider { + readonly id = "local"; + readonly label = "本地离线识别(sherpa-onnx)"; + + private recognizer: Recognizer | undefined; + private recognizerError: string | undefined; + private readonly options: LocalAsrOptions; + + constructor(options: LocalAsrOptions = {}) { + this.options = options; + } + + isConfigured(): boolean { + return true; // bundled with the app; actual load errors surface on start + } + + private async ensureRecognizer(): Promise { + if (this.recognizerError) throw new Error(this.recognizerError); + if (this.recognizer) return this.recognizer; + try { + const modelRoot = resolveModelRoot(this.options); + this.options.log?.(`[asr:local] loading model from ${modelRoot}`); + const sherpa = await loadSherpaOnnx(); + this.recognizer = sherpa.createOnlineRecognizer({ + featConfig: { sampleRate: 16_000, featureDim: 80 }, + modelConfig: { + transducer: { + encoder: path.join(modelRoot, "encoder-epoch-99-avg-1.int8.onnx"), + decoder: path.join(modelRoot, "decoder-epoch-99-avg-1.onnx"), + joiner: path.join(modelRoot, "joiner-epoch-99-avg-1.onnx") + }, + tokens: path.join(modelRoot, "tokens.txt"), + provider: "cpu", + numThreads: 1 + }, + enableEndpoint: 1, + rule1MinTrailingSilence: 2.4, + rule2MinTrailingSilence: 1.2, + rule3MinUtteranceLength: 20 + }); + return this.recognizer; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + this.recognizerError = message; + throw error; + } + } + + async start(sink: AsrEventSink): Promise { + const recognizer = await this.ensureRecognizer(); + let stream: RecognizerStream = recognizer.createStream(); + let stopped = false; + + const decode = (): void => { + while (recognizer.isReady(stream)) recognizer.decode(stream); + }; + + return { + providerId: this.id, + push: (samples: Float32Array): void => { + if (stopped) return; + stream.acceptWaveform(16_000, samples); + decode(); + if (recognizer.isEndpoint(stream)) { + const text = (recognizer.getResult(stream).text || "").trim(); + if (text) sink({ type: "final", text, provider: this.id }); + recognizer.reset(stream); + } else { + const text = (recognizer.getResult(stream).text || "").trim(); + if (text) sink({ type: "partial", text, provider: this.id }); + } + }, + stop: async (): Promise => { + if (stopped) return; + stopped = true; + try { + stream.inputFinished(); + decode(); + const text = (recognizer.getResult(stream).text || "").trim(); + if (text) sink({ type: "final", text, provider: this.id }); + } catch (error) { + sink({ type: "error", message: error instanceof Error ? error.message : String(error) }); + } finally { + // A stream that has seen inputFinished() cannot be reused. + try { stream.free(); } catch { /* already freed */ } + stream = recognizer.createStream(); + stopped = false; + sink({ type: "stopped" }); + } + }, + cancel: (): void => { + stopped = true; + try { stream.free(); } catch { /* already freed */ } + stream = recognizer.createStream(); + stopped = false; + } + }; + } +} + +// Keep the exported helper name used by older callers (wake window diagnostics). +export function isLocalAsrAvailable(provider: LocalSherpaAsrProvider): boolean { + return provider.isConfigured(); +} diff --git a/src/asr/manager.test.ts b/src/asr/manager.test.ts new file mode 100644 index 0000000..5eb87b2 --- /dev/null +++ b/src/asr/manager.test.ts @@ -0,0 +1,75 @@ +import test from "node:test"; +import assert from "node:assert/strict"; +import { AsrManager } from "./manager.js"; +import type { AsrEventSink, AsrProvider, AsrSession } from "./types.js"; +import type { AsrProviderKind } from "./settings.js"; + +function fakeProvider(id: string, options: { configured?: boolean; failStart?: boolean } = {}): AsrProvider { + return { + id, + label: `provider ${id}`, + isConfigured: () => options.configured !== false, + start: async (sink: AsrEventSink): Promise => { + if (options.failStart) throw new Error(`${id} cannot start`); + return { + providerId: id, + push: () => {}, + stop: async () => { sink({ type: "stopped" }); }, + cancel: () => {} + }; + } + }; +} + +test("auto chain prefers third-party, then official, then local", () => { + let kind: AsrProviderKind | undefined = "auto"; + const manager = new AsrManager({ getProviderKind: () => kind }); + manager.register(fakeProvider("openai")); + manager.register(fakeProvider("official")); + manager.register(fakeProvider("local")); + assert.deepEqual(manager.chain(), ["openai", "official", "local"]); + kind = "local"; + assert.deepEqual(manager.chain(), ["local"]); + kind = "official"; + assert.deepEqual(manager.chain(), ["official", "local"]); +}); + +test("unconfigured providers are skipped unless they are the local fallback", () => { + const manager = new AsrManager({ getProviderKind: () => "auto" }); + manager.register(fakeProvider("openai", { configured: false })); + manager.register(fakeProvider("official", { configured: false })); + manager.register(fakeProvider("local")); + assert.deepEqual(manager.chain(), ["local"]); +}); + +test("start falls back when the preferred provider rejects", async () => { + const manager = new AsrManager({ getProviderKind: () => "openai" }); + manager.register(fakeProvider("openai", { failStart: true })); + const local = fakeProvider("local"); + manager.register(local); + const started = await manager.start(() => {}); + assert.equal(started.providerId, "local"); + assert.deepEqual(started.fallbacks, ["openai"]); + manager.cancel(); +}); + +test("start resolves with the first working provider and records no fallbacks", async () => { + const manager = new AsrManager({ getProviderKind: () => "auto" }); + manager.register(fakeProvider("openai")); + manager.register(fakeProvider("local")); + const started = await manager.start(() => {}); + assert.equal(started.providerId, "openai"); + assert.deepEqual(started.fallbacks, []); + manager.cancel(); +}); + +test("start rejects when every provider fails", async () => { + const manager = new AsrManager({ getProviderKind: () => "auto" }); + manager.register(fakeProvider("openai", { failStart: true })); + manager.register(fakeProvider("official", { failStart: true })); + manager.register(fakeProvider("local", { failStart: true })); + await assert.rejects(() => manager.start(() => {}), (error: unknown) => { + const message = error instanceof Error ? error.message : String(error); + return message.includes("openai:") && message.includes("official:") && message.includes("local:"); + }); +}); diff --git a/src/asr/manager.ts b/src/asr/manager.ts new file mode 100644 index 0000000..5ff5e66 --- /dev/null +++ b/src/asr/manager.ts @@ -0,0 +1,190 @@ +/** + * ASR orchestration: picks providers by settings, starts an utterance with + * automatic fallback, and keeps at most one active session at a time. + * + * Fallback chain: + * `auto` third-party (explicit user config) → official relay → local + * `official` official relay → local + * `openai` third-party → local + * `local` local only + * `bailian` 百炼 chat/completions → local + * `bailian-ws` 百炼 realtime WebSocket → 百炼 chat/completions → local + */ +import type { AsrEvent, AsrEventSink, AsrProvider, AsrSession } from "./types.js"; +import type { AsrProviderKind } from "./settings.js"; + +export interface AsrManagerOptions { + /** Reads the live provider preference (`auto` when absent). */ + getProviderKind: () => AsrProviderKind | undefined; + /** Reads the user's fully custom ordered fallback chain, if configured. */ + getCustomChain?: () => string[] | undefined; + log?: (message: string) => void; +} + +export interface StartedAsr { + session: AsrSession; + providerId: string; + /** Ordered provider ids that were tried before one started. */ + fallbacks: string[]; +} + +export class AsrManager { + private readonly providers = new Map(); + private active: AsrSession | undefined; + private readonly options: AsrManagerOptions; + + constructor(options: AsrManagerOptions) { + this.options = options; + } + + register(provider: AsrProvider): this { + this.providers.set(provider.id, provider); + return this; + } + + getProvider(id: string): AsrProvider | undefined { + return this.providers.get(id); + } + + listProviders(): AsrProvider[] { + return [...this.providers.values()]; + } + + /** Resolve the fallback chain for the configured provider kind. */ + resolveChain(): AsrProvider[] { + // 用户自定义链优先:完全按用户给的顺序,只过滤未配置项(local 始终保留在末尾兜底)。 + const custom = this.options.getCustomChain?.(); + if (custom && custom.length) { + const picked = custom + .map((id) => this.providers.get(id)) + .filter((provider): provider is AsrProvider => Boolean(provider)) + .filter((provider) => provider.isConfigured() || provider.id === "local"); + if (!picked.some((provider) => provider.id === "local")) { + const local = this.providers.get("local"); + if (local && (local.isConfigured() || true)) picked.push(local); + } + if (picked.length) return picked; + } + const kind = this.options.getProviderKind() || "auto"; + const chainFor: Record = { + auto: ["openai", "official", "local"], + official: ["official", "local"], + openai: ["openai", "local"], + local: ["local"], + bailian: ["bailian", "local"], + "bailian-ws": ["bailian-ws", "bailian", "local"], + mimo: ["mimo", "local"], + "local-pro": ["local-pro", "local"] + }; + return chainFor[kind] + .map((id) => this.providers.get(id)) + .filter((provider): provider is AsrProvider => Boolean(provider)) + .filter((provider) => provider.isConfigured() || provider.id === "local"); + } + + /** Provider ids the current settings would try, in order (diagnostics). */ + chain(): string[] { + return this.resolveChain().map((provider) => provider.id); + } + + get activeProviderId(): string | undefined { + return this.active?.providerId; + } + + /** + * Consecutive-start failures per provider (session-scoped). After 3 in a + * row a provider is skipped for a cool-off window (5 minutes), mirroring the + * model-level resilience cooldown. Reset on any successful start. + */ + private startFailures = new Map(); + private isAsrCoolingDown(id: string): boolean { + const entry = this.startFailures.get(id); + return Boolean(entry && entry.until > Date.now()); + } + + /** Start an utterance, falling back down the chain when a provider cannot start. */ + async start(sink: AsrEventSink): Promise { + if (this.active) await this.cancel(); + const chain = this.resolveChain(); + if (!chain.length) throw new Error("没有可用的语音识别服务:请登录官方服务、配置第三方识别,或安装本地模型"); + // Keep at least one option even when everything is cooling down. + const ready = chain.filter((provider) => !this.isAsrCoolingDown(provider.id)); + const ordered = ready.length ? [...ready, ...chain.filter((provider) => this.isAsrCoolingDown(provider.id))] : chain; + const failures: Array<{ id: string; message: string }> = []; + for (const provider of ordered) { + try { + const session = await provider.start(sink); + this.active = session; + this.startFailures.delete(provider.id); + this.options.log?.(`[asr] session started provider=${provider.id} chain=${ordered.map((item) => item.id).join(">")}`); + sink({ type: "ready", provider: provider.id }); + return { session, providerId: provider.id, fallbacks: failures.map((failure) => failure.id) }; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + this.options.log?.(`[asr] provider ${provider.id} failed to start: ${message}`); + const entry = this.startFailures.get(provider.id) || { count: 0, until: 0 }; + entry.count += 1; + if (entry.count >= 3) entry.until = Date.now() + 5 * 60_000; + this.startFailures.set(provider.id, entry); + failures.push({ id: provider.id, message }); + } + } + throw new Error(failures.map((failure) => `${failure.id}:${failure.message}`).join(";")); + } + + /** Push audio into the active session (if any). */ + push(samples: Float32Array): void { + this.active?.push(samples); + } + + /** Finish the active utterance. */ + async stop(): Promise { + const session = this.active; + this.active = undefined; + if (!session) return; + try { await session.stop(); } catch { /* stop errors surface as error events */ } + } + + /** Abort the active utterance without a final result. */ + async cancel(): Promise { + const session = this.active; + this.active = undefined; + if (!session) return; + try { session.cancel(); } catch { /* cancel must never throw */ } + } + + /** Probe every configured provider (used by the settings page). */ + async test(kind: AsrProviderKind): Promise> { + const ids: Record = { + auto: ["openai", "official", "local"], + official: ["official"], + openai: ["openai"], + local: ["local"], + bailian: ["bailian"], + "bailian-ws": ["bailian-ws"], + mimo: ["mimo"], + "local-pro": ["local-pro"] + }; + const results: Array<{ id: string; label: string; ok: boolean; message: string }> = []; + for (const id of ids[kind]) { + const provider = this.providers.get(id); + if (!provider) continue; + if (!provider.test) { + results.push({ id, label: provider.label, ok: provider.isConfigured(), message: provider.isConfigured() ? "已配置" : "未配置" }); + continue; + } + try { + const result = await provider.test(); + results.push({ id, label: provider.label, ok: result.ok, message: result.message }); + } catch (error) { + results.push({ id, label: provider.label, ok: false, message: error instanceof Error ? error.message : String(error) }); + } + } + return results; + } +} + +/** Convenience wrapper matching the legacy `startSpeech` result shape. */ +export function isRemoteAsrEvent(event: AsrEvent): boolean { + return event.type === "ready" || event.type === "partial" || event.type === "final"; +} diff --git a/src/asr/mimo-http.ts b/src/asr/mimo-http.ts new file mode 100644 index 0000000..5fde5e6 --- /dev/null +++ b/src/asr/mimo-http.ts @@ -0,0 +1,189 @@ +/** + * MiMo ASR — OpenAI-compatible chat/completions with input_audio. + * + * Official docs: https://mimo.mi.com/docs/zh-CN/api/audio/Speech-Recognition + * POST https://api.xiaomimimo.com/v1/chat/completions + * Auth (one of): `api-key: $MIMO_API_KEY` or `Authorization: Bearer ` + * Body: { model, messages: [{ role: "user", content: [{ type: "input_audio", + * input_audio: { data: "data:audio/wav;base64,…", format: "wav" | "mp3" } }] }], + * asr_options: { language: "auto" | "zh" | "en" } } + * Audio: mp3 or wav only (we always send WAV); one-shot per request. + * Response: choices[0].message.content holds the transcription. + * + * The session buffers audio and flushes every ~3s as a partial, mirroring the + * OpenAI HTTP provider, then sends the whole utterance on stop for the final + * text — MiMo is not a streaming API. + */ +import type { AsrEventSink, AsrProvider, AsrSession, AsrTestResult } from "./types.js"; +import { encodeWav } from "./wav.js"; + +export interface MimoAsrSettings { + name?: string; + apiKeyEnv?: string; + baseUrl?: string; + model?: string; + language?: string; +} + +export const MIMO_ASR_DEFAULTS = { + apiKeyEnv: "MIMO_API_KEY", + baseUrl: "https://api.xiaomimimo.com/v1", + model: "MiMo-V2.5-ASR", + language: "auto" +} as const; + +export interface MimoAsrProviderOptions { + getSettings(): MimoAsrSettings | undefined; + getApiKey(name: string): string | undefined; + log?(message: string): void; + fetchImpl?: typeof fetch; + uuid?(): string; +} + +function text(value: string | undefined): string | undefined { + const trimmed = (value ?? "").trim(); + return trimmed ? trimmed : undefined; +} + +export interface ResolvedMimoAsrConfig { + apiKey: string; + baseUrl: string; + model: string; + language: string; +} + +export function resolveMimoAsrConfig(settings: MimoAsrSettings | undefined, getApiKey: (name: string) => string | undefined, env: NodeJS.ProcessEnv = process.env): ResolvedMimoAsrConfig | null { + const apiKey = (text(settings?.apiKeyEnv) && getApiKey(text(settings!.apiKeyEnv)!)) || text(env.MIMO_API_KEY); + if (!apiKey) return null; + return { + apiKey, + baseUrl: (text(settings?.baseUrl) || text(env.MIMO_ASR_BASE_URL) || MIMO_ASR_DEFAULTS.baseUrl).replace(/\/+$/, ""), + model: text(settings?.model) || text(env.MIMO_ASR_MODEL) || MIMO_ASR_DEFAULTS.model, + language: text(settings?.language) || text(env.MIMO_ASR_LANGUAGE) || MIMO_ASR_DEFAULTS.language + }; +} + +interface MimoAsrSessionDeps { + config: ResolvedMimoAsrConfig; + sink: AsrEventSink; + fetchImpl: typeof fetch; + log(message: string): void; +} + +export class MimoAsrSession implements AsrSession { + readonly providerId = "mimo"; + private readonly deps: MimoAsrSessionDeps; + private buffer: Float32Array = new Float32Array(0); + private stopped = false; + private flushing: Promise = Promise.resolve(); + + constructor(deps: MimoAsrSessionDeps) { this.deps = deps; } + + push(samples: Float32Array): void { + if (this.stopped || samples.length === 0) return; + const merged = new Float32Array(this.buffer.length + samples.length); + merged.set(this.buffer, 0); + merged.set(samples, this.buffer.length); + this.buffer = merged; + if (this.buffer.length >= 48000) void this.flush("partial"); + } + + private async transcribe(samples: Float32Array, timeoutMs = 30000): Promise { + const { config, fetchImpl } = this.deps; + const wav = encodeWav(samples); + const dataUrl = `data:audio/wav;base64,${Buffer.from(wav).toString("base64")}`; + const body = { + model: config.model, + messages: [{ + role: "user", + content: [{ type: "input_audio", input_audio: { data: dataUrl, format: "wav" } }] + }], + asr_options: { language: config.language }, + stream: false + }; + const controller = new AbortController(); + const timer = setTimeout(() => controller.abort(), timeoutMs); + try { + const response = await fetchImpl(`${config.baseUrl}/chat/completions`, { + method: "POST", + headers: { "Content-Type": "application/json", "api-key": config.apiKey, Authorization: `Bearer ${config.apiKey}` }, + body: JSON.stringify(body), + signal: controller.signal + }); + if (!response.ok) { + const detail = await response.text().catch(() => ""); + throw new Error(`MiMo ASR HTTP ${response.status}${detail ? `:${detail.slice(0, 300)}` : ""}`); + } + const json = (await response.json()) as { choices?: Array<{ message?: { content?: string } }>; error?: { message?: string } }; + if (json.error?.message) throw new Error(`MiMo ASR:${json.error.message}`); + return json.choices?.[0]?.message?.content?.trim() || null; + } finally { + clearTimeout(timer); + } + } + + private async flush(reason: "partial" | "final"): Promise { + if (this.buffer.length === 0) return; + const samples = this.buffer; + if (reason === "final") this.buffer = new Float32Array(0); + this.flushing = this.flushing.then(async () => { + try { + const transcript = await this.transcribe(samples); + if (!transcript) return; + this.deps.sink({ type: reason === "final" ? "final" : "partial", text: transcript, provider: "mimo" }); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + this.deps.log(`[asr:mimo] ${message}`); + if (reason === "final") this.deps.sink({ type: "error", message }); + } + }); + await this.flushing; + } + + async stop(): Promise { + if (this.stopped) return; + this.stopped = true; + await this.flush("final"); + } + + cancel(): void { + this.stopped = true; + this.buffer = new Float32Array(0); + this.deps.sink({ type: "stopped" }); + } +} + +export class MimoHttpAsrProvider implements AsrProvider { + readonly id = "mimo"; + readonly label = "小米 MiMo(chat/completions)"; + private readonly options: MimoAsrProviderOptions; + + constructor(options: MimoAsrProviderOptions) { this.options = options; } + + isConfigured(): boolean { + return resolveMimoAsrConfig(this.options.getSettings(), this.options.getApiKey) !== null; + } + + displayName(): string { + return this.options.getSettings()?.name?.trim() || "小米 MiMo"; + } + + async test(): Promise { + const config = resolveMimoAsrConfig(this.options.getSettings(), this.options.getApiKey); + if (!config) return { ok: false, message: `环境变量 ${this.options.getSettings()?.apiKeyEnv || "MIMO_API_KEY"} 中没有 API Key(可在下方填写并保存到工作区 .env)` }; + return { ok: true, message: `${this.displayName()} 已配置(${config.baseUrl}/chat/completions,模型 ${config.model})` }; + } + + async start(sink: AsrEventSink): Promise { + const config = resolveMimoAsrConfig(this.options.getSettings(), this.options.getApiKey); + if (!config) throw new Error("小米 MiMo 语音识别未配置:请在设置中填写 API Key(保存到工作区 .env 的 MIMO_API_KEY)"); + this.options.log?.(`[asr:mimo] ready ${config.baseUrl} model=${config.model} language=${config.language}`); + sink({ type: "ready", provider: "mimo" }); + return new MimoAsrSession({ + config, + sink, + fetchImpl: this.options.fetchImpl || fetch, + log: (message) => this.options.log?.(message) + }); + } +} diff --git a/src/asr/openai-http.test.ts b/src/asr/openai-http.test.ts new file mode 100644 index 0000000..971ffe8 --- /dev/null +++ b/src/asr/openai-http.test.ts @@ -0,0 +1,92 @@ +import test from "node:test"; +import assert from "node:assert/strict"; +import { OpenAiHttpAsrProvider } from "./openai-http.js"; +import type { AsrEvent } from "./types.js"; +import type { OpenAiAsrSettings } from "./settings.js"; + +function setup(settings: OpenAiAsrSettings | undefined, apiKey: string, responses: Array<{ status: number; body: string }>): { provider: OpenAiHttpAsrProvider; requests: Array<{ url: string; auth: string | null; body: FormData }> } { + const requests: Array<{ url: string; auth: string | null; body: FormData }> = []; + let call = 0; + const fetchImpl = (async (input: RequestInfo | URL, init?: RequestInit) => { + const url = String(input); + const auth = init?.headers ? (init.headers as Record).Authorization ?? null : null; + requests.push({ url, auth, body: init?.body as FormData }); + const response = responses[Math.min(call, responses.length - 1)]; + call += 1; + return new Response(response.body, { status: response.status, headers: { "Content-Type": "application/json" } }); + }) as typeof fetch; + const provider = new OpenAiHttpAsrProvider({ + getSettings: () => settings, + getApiKey: () => apiKey, + fetchImpl + }); + return { provider, requests }; +} + +const validSettings: OpenAiAsrSettings = { name: "小米 MiMo ASR", baseUrl: "https://token-plan-cn.xiaomimimo.com/v1", apiKeyEnv: "MIMO_API_KEY", model: "MiMo-ASR" }; + +test("isConfigured requires settings, key and endpoint fields", () => { + const { provider } = setup(undefined, "", []); + assert.equal(provider.isConfigured(), false); + const missingKey = setup(validSettings, "", []); + assert.equal(missingKey.provider.isConfigured(), false); + const ready = setup(validSettings, "sk-test", []); + assert.equal(ready.provider.isConfigured(), true); +}); + +test("test() posts a WAV to the OpenAI-compatible transcriptions endpoint", async () => { + const { provider, requests } = setup(validSettings, "sk-test", [{ status: 200, body: JSON.stringify({ text: "你好" }) }]); + const result = await provider.test(); + assert.equal(result.ok, true); + assert.match(result.message, /MiMo-ASR|连接成功/); + assert.equal(requests.length, 1); + assert.equal(requests[0].url, "https://token-plan-cn.xiaomimimo.com/v1/audio/transcriptions"); + assert.equal(requests[0].auth, "Bearer sk-test"); + assert.equal(requests[0].body.get("model"), "MiMo-ASR"); + const file = requests[0].body.get("file"); + assert.ok(file instanceof File); + assert.equal((file as File).type, "audio/wav"); +}); + +test("test() reports auth failures in plain language", async () => { + const { provider } = setup(validSettings, "sk-bad", [{ status: 401, body: JSON.stringify({ error: { message: "bad key" } }) }]); + const result = await provider.test(); + assert.equal(result.ok, false); + assert.match(result.message, /API Key 无效或无权限/); +}); + +test("test() explains a wrong base URL", async () => { + const { provider } = setup(validSettings, "sk-test", [{ status: 404, body: "" }]); + const result = await provider.test(); + assert.equal(result.ok, false); + assert.match(result.message, /接口不存在/); +}); + +test("start() rejects with guidance when the key is missing", async () => { + const { provider } = setup(validSettings, "", []); + await assert.rejects(() => provider.start(() => {}), /MIMO_API_KEY/); +}); + +test("a session emits final text on stop, then a stopped event", async () => { + const { provider } = setup(validSettings, "sk-test", [{ status: 200, body: JSON.stringify({ text: "你好世界" }) }]); + const events: AsrEvent[] = []; + const session = await provider.start((event) => events.push(event)); + session.push(new Float32Array(16_000)); // 1 second + await session.stop(); + const types = events.map((event) => event.type); + // The provider emits final + stopped; the manager adds "ready" itself. + assert.ok(types.includes("final")); + assert.equal(types[types.length - 1], "stopped"); + const final = events.find((event) => event.type === "final"); + assert.equal(final && final.type === "final" ? final.text : "", "你好世界"); +}); + +test("cancel drops buffered audio without emitting results", async () => { + const { provider } = setup(validSettings, "sk-test", [{ status: 200, body: JSON.stringify({ text: "你好" }) }]); + const events: AsrEvent[] = []; + const session = await provider.start((event) => events.push(event)); + session.push(new Float32Array(16_000)); + session.cancel(); + await session.stop(); + assert.equal(events.some((event) => event.type === "final" || event.type === "partial"), false); +}); diff --git a/src/asr/openai-http.ts b/src/asr/openai-http.ts new file mode 100644 index 0000000..05adae3 --- /dev/null +++ b/src/asr/openai-http.ts @@ -0,0 +1,179 @@ +/** + * Third-party cloud ASR through an OpenAI-compatible `/audio/transcriptions` + * endpoint — works with 小米 MiMo ASR、SiliconFlow SenseVoice、Groq Whisper and + * any other provider that implements the same multipart protocol. + * + * Plain HTTP cannot stream, so the session chunks the utterance: buffered + * audio is flushed as a WAV roughly every three seconds and surfaced as + * `partial` text, and the stop() flush emits the `final` result. + */ +import type { AsrEventSink, AsrProvider, AsrSession, AsrTestResult } from "./types.js"; +import { encodeWav, mergeSamples, ASR_SAMPLE_RATE } from "./wav.js"; +import type { OpenAiAsrSettings } from "./settings.js"; + +const PARTIAL_FLUSH_MS = 3_000; +const MIN_CHUNK_MS = 900; +const CONNECT_TIMEOUT_MS = 12_000; + +export interface OpenAiAsrOptions { + /** Current provider settings (re-read on each start). */ + getSettings: () => OpenAiAsrSettings | undefined; + /** Resolves the API key for an env var name (usually process.env). */ + getApiKey: (envName: string) => string | undefined; + fetchImpl?: typeof fetch; + log?: (message: string) => void; +} + +interface TranscriptionResponse { text?: string; error?: { message?: string } | string } + +function statusMessage(status: number, body: string): string { + if (status === 401 || status === 403) return "API Key 无效或无权限(检查密钥是否正确、是否有语音模型权限)"; + if (status === 404) return "接口不存在:请检查 Base URL 是否为 OpenAI 兼容地址(一般以 /v1 结尾)"; + if (status === 422 || status === 400) return `请求被拒绝(${status}):${body.slice(0, 200) || "请检查模型名称"} `; + if (status === 429) return "请求过于频繁(429 限流),请稍后重试"; + if (status >= 500) return `服务端错误(${status})`; + return `请求失败(${status}):${body.slice(0, 200)}`; +} + +export class OpenAiHttpAsrProvider implements AsrProvider { + readonly id = "openai"; + readonly label = "第三方云端语音识别"; + private readonly options: OpenAiAsrOptions; + + constructor(options: OpenAiAsrOptions) { + this.options = options; + } + + private config(): { settings: OpenAiAsrSettings; apiKey: string } | null { + const settings = this.options.getSettings(); + if (!settings || !settings.baseUrl?.trim() || !settings.model?.trim() || !settings.apiKeyEnv?.trim()) return null; + const apiKey = (this.options.getApiKey(settings.apiKeyEnv) || "").trim(); + if (!apiKey) return null; + return { settings, apiKey }; + } + + private displayName(): string { + return this.options.getSettings()?.name?.trim() || this.label; + } + + isConfigured(): boolean { + return this.config() !== null; + } + + private async transcribe(request: { baseUrl: string; apiKey: string; model: string; language?: string }, samples: Float32Array, timeoutMs: number): Promise<{ ok: true; text: string } | { ok: false; message: string }> { + const fetchImpl = this.options.fetchImpl || fetch; + const form = new FormData(); + const wav = encodeWav(samples); + // Copy into a plain ArrayBuffer-backed Blob part: TS 5.7 rejects the + // `ArrayBufferLike` union in Blob constructors. + form.append("file", new Blob([wav.slice().buffer as ArrayBuffer], { type: "audio/wav" }), "speech.wav"); + form.append("model", request.model); + if (request.language) form.append("language", request.language); + form.append("response_format", "json"); + const url = `${request.baseUrl.replace(/\/+$/, "")}/audio/transcriptions`; + const controller = new AbortController(); + const timer = setTimeout(() => controller.abort(), timeoutMs); + try { + const response = await fetchImpl(url, { + method: "POST", + headers: { Authorization: `Bearer ${request.apiKey}` }, + body: form, + signal: controller.signal + }); + const body = await response.text(); + if (!response.ok) return { ok: false, message: statusMessage(response.status, body) }; + let payload: TranscriptionResponse; + try { payload = JSON.parse(body) as TranscriptionResponse; } + catch { return { ok: false, message: "服务返回了无法解析的内容(确认接口为 OpenAI 兼容的 /audio/transcriptions)" }; } + if (payload.error) return { ok: false, message: typeof payload.error === "string" ? payload.error : payload.error.message || "第三方语音识别返回错误" }; + return { ok: true, text: (payload.text || "").trim() }; + } catch (error) { + if (error instanceof Error && error.name === "AbortError") return { ok: false, message: `连接超时(${timeoutMs / 1000}s):无法访问 ${request.baseUrl}` }; + return { ok: false, message: `无法连接 ${request.baseUrl}:${error instanceof Error ? error.message : String(error)}` }; + } finally { + clearTimeout(timer); + } + } + + async test(): Promise { + const config = this.config(); + if (!config) { + const settings = this.options.getSettings(); + if (!settings?.baseUrl || !settings.model || !settings.apiKeyEnv) return { ok: false, message: "请先填写 Base URL、模型名称和 API Key 环境变量名" }; + return { ok: false, message: `环境变量 ${settings.apiKeyEnv} 中没有 API Key(保存设置后填写密钥再测试)` }; + } + // A 250ms near-silent probe is enough to validate auth, endpoint and model. + const silence = new Float32Array(Math.round(ASR_SAMPLE_RATE * 0.25)); + const result = await this.transcribe({ baseUrl: config.settings.baseUrl, apiKey: config.apiKey, model: config.settings.model, language: config.settings.language }, silence, CONNECT_TIMEOUT_MS); + if (!result.ok) return { ok: false, message: result.message }; + return { ok: true, message: `${this.displayName()} 连接成功${result.text ? `(识别:${result.text.slice(0, 40)})` : ""}` }; + } + + async start(sink: AsrEventSink): Promise { + const config = this.config(); + if (!config) { + const settings = this.options.getSettings(); + if (!settings || !settings.baseUrl?.trim() || !settings.model?.trim() || !settings.apiKeyEnv?.trim()) throw new Error("第三方语音识别未配置:请在设置中填写 Base URL、模型和 API Key"); + throw new Error(`环境变量 ${settings.apiKeyEnv} 缺少 API Key,无法使用第三方语音识别`); + } + // Snapshot the request template: the settings object may be replaced by a + // concurrent save while this utterance is running. + const request = { baseUrl: config.settings.baseUrl, apiKey: config.apiKey, model: config.settings.model, language: config.settings.language }; + const provider = this; + let buffer: Float32Array[] = []; + let bufferedMs = 0; + let inFlight = false; + let cancelled = false; + let stopped = false; + const flushTimer = setInterval(() => { void maybeFlush(false); }, PARTIAL_FLUSH_MS); + + async function maybeFlush(final: boolean): Promise { + if (cancelled || inFlight) return; + if (!final && (stopped || bufferedMs < MIN_CHUNK_MS)) return; + const samples = mergeSamples(buffer); + buffer = []; + bufferedMs = 0; + if (!samples.length) return; + inFlight = true; + const result = await provider.transcribe(request, samples, CONNECT_TIMEOUT_MS); + inFlight = false; + if (cancelled) return; + if (!result.ok) { + provider.options.log?.(`[asr:openai] ${result.message}`); + sink({ type: "error", message: result.message }); + return; + } + if (result.text) sink({ type: final ? "final" : "partial", text: result.text, provider: "openai" }); + else if (final) sink({ type: "final", text: "", provider: "openai" }); + } + + return { + providerId: this.id, + push: (samples: Float32Array): void => { + if (cancelled || stopped) return; + buffer.push(samples); + bufferedMs += (samples.length / ASR_SAMPLE_RATE) * 1_000; + }, + stop: async (): Promise => { + if (stopped || cancelled) return; + stopped = true; + clearInterval(flushTimer); + try { + // Wait for an in-flight partial flush so the final request is ordered. + const deadline = Date.now() + CONNECT_TIMEOUT_MS; + while (inFlight && Date.now() < deadline) await new Promise((resolve) => setTimeout(resolve, 50)); + await maybeFlush(true); + } finally { + clearInterval(flushTimer); + sink({ type: "stopped" }); + } + }, + cancel: (): void => { + cancelled = true; + clearInterval(flushTimer); + buffer = []; + bufferedMs = 0; + } + }; + } +} diff --git a/src/asr/relay.ts b/src/asr/relay.ts new file mode 100644 index 0000000..f7d5a8a --- /dev/null +++ b/src/asr/relay.ts @@ -0,0 +1,156 @@ +/** + * Cloud ASR via the official SECTL relay (`/asr/ws` WebSocket). + * + * Requires a signed-in session (`SECTL_OFFICIAL_TOKEN` + API URL). Connection + * failures reject `start()` so the manager can fall back instead of leaving + * the user without speech input. + */ +import type { AsrEvent, AsrEventSink, AsrProvider, AsrSession } from "./types.js"; + +const CONNECT_TIMEOUT_MS = 8_000; +/** Cap buffered audio while the socket is still connecting (~1s at 16 kHz). */ +const MAX_PENDING_CHUNKS = 32; + +export interface RelayAsrOptions { + getToken?: () => string; + getApiBaseUrl?: () => string; + log?: (message: string) => void; + WebSocketCtor?: typeof WebSocket; +} + +export class RelayAsrProvider implements AsrProvider { + readonly id = "official"; + readonly label = "官方云端语音识别"; + private readonly options: RelayAsrOptions; + /** Current socket, when a session started through this provider. */ + private socket: WebSocket | undefined; + + constructor(options: RelayAsrOptions = {}) { + this.options = options; + } + + private endpoint(): string | null { + const token = this.options.getToken?.() || process.env.SECTL_OFFICIAL_TOKEN || ""; + const baseUrl = (this.options.getApiBaseUrl?.() || process.env.SECTL_OFFICIAL_API_URL || "").replace(/\/$/, ""); + if (!token || !baseUrl) return null; + const wsBase = baseUrl.replace(/^https:/, "wss:").replace(/^http:/, "ws:"); + return `${wsBase}/asr/ws?token=${encodeURIComponent(token)}`; + } + + isConfigured(): boolean { + return this.endpoint() !== null; + } + + async test(): Promise<{ ok: boolean; message: string }> { + const url = this.endpoint(); + if (!url) return { ok: false, message: "尚未登录官方服务(缺少 SECTL_OFFICIAL_TOKEN),无法使用官方云端语音识别" }; + return { ok: true, message: "官方云端语音识别已配置(连接在开始说话时建立)" }; + } + + async start(sink: AsrEventSink): Promise { + const url = this.endpoint(); + if (!url) throw new Error("官方云端语音识别未配置:请先登录官方服务"); + const WebSocketCtor = this.options.WebSocketCtor || WebSocket; + if (typeof WebSocketCtor === "undefined") throw new Error("当前环境不支持 WebSocket"); + + let logTarget = ""; + try { + const parsed = new URL(url); + logTarget = `${parsed.protocol}//${parsed.host}${parsed.pathname}`; + } catch { /* keep the placeholder */ } + this.options.log?.(`[asr:official] connecting to ${logTarget}`); + + return await new Promise((resolve, reject) => { + let socket: WebSocket; + try { + socket = new WebSocketCtor(url); + } catch (error) { + reject(new Error(`云端语音识别连接创建失败:${error instanceof Error ? error.message : String(error)}`)); + return; + } + socket.binaryType = "arraybuffer"; + const pending: ArrayBuffer[] = []; + let settled = false; + const timeout = setTimeout(() => { + if (settled) return; + settled = true; + try { socket.close(); } catch { /* may already be closed */ } + reject(new Error(`云端语音识别连接超时(${CONNECT_TIMEOUT_MS / 1000}s)`)); + }, CONNECT_TIMEOUT_MS); + const finish = (error?: Error, session?: AsrSession): void => { + if (settled) return; + settled = true; + clearTimeout(timeout); + if (error || !session) reject(error || new Error("云端语音识别连接失败")); + else resolve(session); + }; + + socket.onopen = () => { + this.options.log?.("[asr:official] websocket opened"); + // The relay must receive the start control message before binary + // audio. Sending buffered audio first makes the relay discard it. + try { socket.send(JSON.stringify({ type: "start" })); } catch { /* socket may close during startup */ } + for (const pcm of pending.splice(0)) socket.send(pcm); + if (this.socket && this.socket !== socket) { try { this.socket.close(1000, "superseded"); } catch { /* ignore */ } } + this.socket = socket; + // Set once the user finishes the utterance; the server-initiated close + // that follows "Done" is then expected, not an error. + let finished = false; + finish(undefined, { + providerId: this.id, + push: (samples: Float32Array): void => { + const pcm = samples.buffer.slice(samples.byteOffset, samples.byteOffset + samples.byteLength) as ArrayBuffer; + if (socket.readyState === WebSocket.OPEN) { + try { socket.send(pcm); } catch { /* socket may close between the state check and send */ } + return; + } + if (socket.readyState === WebSocket.CONNECTING) { + if (pending.length >= MAX_PENDING_CHUNKS) pending.shift(); + pending.push(pcm); + } + }, + stop: async (): Promise => { + finished = true; + if (socket.readyState === WebSocket.OPEN) { + try { socket.send("Done"); } catch { /* socket may already be closing */ } + } else if (socket.readyState === WebSocket.CONNECTING) socket.close(); + }, + cancel: (): void => { + finished = true; + try { socket.close(1000, "cancelled"); } catch { /* socket may already be closed */ } + if (this.socket === socket) this.socket = undefined; + } + }); + socket.onclose = (event: CloseEvent): void => onclose(event, () => finished); + }; + socket.onmessage = (event: MessageEvent) => { + try { + sink(typeof event.data === "string" ? JSON.parse(event.data) as AsrEvent : (event.data as AsrEvent)); + } catch { + sink({ type: "log", message: String(event.data ?? "") }); + } + }; + socket.onerror = (event: Event) => { + const errorEvent = event as ErrorEvent; + const error = errorEvent.error as { message?: string; code?: string } | undefined; + this.options.log?.(`[asr:official] websocket error state=${socket.readyState} message=${errorEvent.message || error?.message || ""}`); + const message = "云端语音识别连接失败"; + if (settled) sink({ type: "error", message }); + else finish(new Error(message)); + }; + const onclose = (event: CloseEvent, isFinished: () => boolean): void => { + this.options.log?.(`[asr:official] websocket closed code=${event.code}`); + if (this.socket === socket) this.socket = undefined; + if (!settled) { + finish(new Error(`云端语音识别连接已断开(code=${event.code}),请检查网络或改用第三方/本地识别`)); + return; + } + // A close after a normal stop/cancel is expected; only a mid-utterance + // drop should surface as an error so the UI does not hang. + if (isFinished()) return; + sink({ type: "error", message: `云端语音识别连接断开(code=${event.code})` }); + }; + socket.onclose = (event: CloseEvent): void => onclose(event, () => false); + }); + } +} diff --git a/src/asr/settings.test.ts b/src/asr/settings.test.ts new file mode 100644 index 0000000..8f450df --- /dev/null +++ b/src/asr/settings.test.ts @@ -0,0 +1,45 @@ +import test from "node:test"; +import assert from "node:assert/strict"; +import { ASR_OPENAI_PRESETS, MIMO_ASR_DEFAULTS, isOpenAiAsrConfigured, normalizeSpeechSettings, type SpeechAsrSettings } from "./settings.js"; + +test("MiMo ASR ships dedicated chat/completions defaults (not /audio/transcriptions)", () => { + // 小米 MiMo 走专用 chat/completions + input_audio 协议(mimo-http.ts), + // 不在 ASR_OPENAI_PRESETS 里,但必须提供开箱默认端点。 + assert.equal(ASR_OPENAI_PRESETS.some((preset) => preset.id === "mimo"), false); + assert.match(MIMO_ASR_DEFAULTS.baseUrl, /^https:\/\/.+\/v1$/); + assert.ok(MIMO_ASR_DEFAULTS.model); + assert.ok(MIMO_ASR_DEFAULTS.apiKeyEnv); +}); + +test("presets are unique by id and base URL", () => { + const ids = ASR_OPENAI_PRESETS.map((preset) => preset.id); + const urls = ASR_OPENAI_PRESETS.map((preset) => preset.baseUrl); + assert.equal(new Set(ids).size, ids.length); + assert.equal(new Set(urls).size, urls.length); +}); + +test("normalizeSpeechSettings accepts undefined and garbage", () => { + // MiMo 默认块始终带出(有官方默认端点,设置页开箱即用)。 + const defaults = { betterRecognition: false, provider: "auto", mimo: { apiKeyEnv: "MIMO_API_KEY", baseUrl: "https://api.xiaomimimo.com/v1", model: "mimo-v2.5-asr" } }; + assert.deepEqual(normalizeSpeechSettings(undefined), defaults); + assert.deepEqual(normalizeSpeechSettings("nonsense"), defaults); +}); + +test("normalizeSpeechSettings keeps a valid provider and trims endpoint fields", () => { + const normalized = normalizeSpeechSettings({ provider: "openai", openai: { baseUrl: " https://api.example.com/v1/ ", apiKeyEnv: "MIMO_API_KEY", model: " MiMo-ASR " } }); + assert.equal(normalized.provider, "openai"); + assert.equal(normalized.openai?.baseUrl, "https://api.example.com/v1"); + assert.equal(normalized.openai?.model, "MiMo-ASR"); +}); + +test("normalizeSpeechSettings rejects malformed env var names", () => { + const normalized = normalizeSpeechSettings({ openai: { baseUrl: "https://x.example.com/v1", apiKeyEnv: "not a name!", model: "m" } }); + assert.equal(normalized.openai?.apiKeyEnv, ""); +}); + +test("isOpenAiAsrConfigured requires endpoint, model and key name", () => { + const base: SpeechAsrSettings = { provider: "openai", openai: { baseUrl: "https://x/v1", apiKeyEnv: "K", model: "m" } }; + assert.equal(isOpenAiAsrConfigured(base), true); + assert.equal(isOpenAiAsrConfigured({ ...base, openai: { ...base.openai!, model: "" } }), false); + assert.equal(isOpenAiAsrConfigured(undefined), false); +}); diff --git a/src/asr/settings.ts b/src/asr/settings.ts new file mode 100644 index 0000000..f2c3be5 --- /dev/null +++ b/src/asr/settings.ts @@ -0,0 +1,266 @@ +/** Settings-facing ASR configuration shared between the config layer and UI. */ + +/** + * Which speech-to-text backend to use. `auto` follows the fallback chain. + * + * `bailian` is 阿里云百炼's OpenAI-compatible `chat/completions` ASR channel + * (non-streaming utterances), `bailian-ws` its realtime WebSocket channel. + */ +export type AsrProviderKind = "auto" | "official" | "openai" | "local" | "bailian" | "bailian-ws" | "mimo" | "local-pro"; + +export interface OpenAiAsrSettings { + /** Optional display name (e.g. 小米 MiMo ASR). */ + name?: string; + /** OpenAI-compatible base URL, e.g. `https://token-plan-cn.xiaomimimo.com/v1`. */ + baseUrl: string; + /** Env var name that holds the API key inside the workspace `.env`. */ + apiKeyEnv: string; + /** Model name posted to `/audio/transcriptions`. */ + model: string; + /** Optional ISO language hint (`zh`, `en`…). */ + language?: string; +} + +/** + * 阿里云百炼 ASR settings. Channel A talks to `POST {baseUrl}/chat/completions` + * with an `input_audio` Data URL; channel B talks to the realtime WebSocket at + * `wsUrl`. Both share one DashScope API key (env-isolated like other providers). + * Documented defaults are only applied when a field is left empty. + */ +export interface BailianAsrSettings { + /** Optional display name. */ + name?: string; + /** Env var name that holds the API key inside the workspace `.env`. */ + apiKeyEnv: string; + /** OpenAI-compatible base URL, e.g. `https://…/compatible-mode/v1`. */ + baseUrl: string; + /** Realtime endpoint, e.g. `wss://…/api-ws/v1/inference`. */ + wsUrl: string; + /** Non-streaming model name (`qwen3-asr-flash`). */ + model: string; + /** Realtime streaming model name (`qwen-audio-3.1-asr-flash-streaming`). */ + streamModel: string; + /** Optional single language hint; empty means auto detect (recommended). */ + language?: string; + /** ITN (digits normalization); 百炼 defaults to false. */ + enableItn?: boolean; +} + +/** Documented 百炼 defaults (DASHSCOPE public domain fallback). */ +export const BAILIAN_DEFAULTS = { + apiKeyEnv: "BAILIAN_API_KEY", + baseUrl: "https://dashscope.aliyuncs.com/compatible-mode/v1", + wsUrl: "wss://dashscope.aliyuncs.com/api-ws/v1/inference", + model: "qwen3-asr-flash", + streamModel: "qwen-audio-3.1-asr-flash-streaming" +} as const; + +/** 小米 MiMo ASR — dedicated chat/completions + input_audio protocol. */ +export interface MimoAsrSettings { + name?: string; + baseUrl: string; + apiKeyEnv: string; + model: string; + /** "auto" | "zh" | "en" — maps to asr_options.language. */ + language?: string; +} + +export const MIMO_ASR_DEFAULTS = { + apiKeyEnv: "MIMO_API_KEY", + baseUrl: "https://api.xiaomimimo.com/v1", + model: "mimo-v2.5-asr", + language: "auto" +} as const; + +/** 嘈杂环境优化(教室/希沃一体机:麦克风在屏幕顶部,远场+高噪声)。 */ +export interface AsrNoiseSettings { + /** classroom: 远场 VAD + 高噪声容忍 + 拾音增强(默认推荐)。 */ + profile?: "standard" | "classroom" | "custom"; + /** 百炼 speech_noise_threshold [-1,1]:越接近 -1 越不容易漏掉语音。 */ + speechNoiseThreshold?: number; + /** 百炼 VAD 模型(qwen-audio 系)。 */ + vadModel?: "near_meeting_16k" | "far_field_meeting_16k"; + /** 即时热词(百炼 vocabulary;权重 50 为超级热词)。 */ + hotwords?: string[]; +} + +/** 输入/输出音频设备选择("auto" 自动检测 / "default" 系统默认 / 设备 ID)。 */ +export interface AsrAudioDeviceSettings { + input?: string; + output?: string; +} + +export interface SpeechAsrSettings { + betterRecognition?: boolean; + provider?: AsrProviderKind; + openai?: OpenAiAsrSettings; + bailian?: BailianAsrSettings; + mimo?: MimoAsrSettings; + /** 完全自定义的回退链(按顺序尝试,例如 ["bailian-ws","mimo","local-pro","local"])。 */ + chain?: string[]; + noise?: AsrNoiseSettings; + audio?: AsrAudioDeviceSettings; +} + +/** Provider dropdown entries for the 第三方云端 panel's 百炼 channels. */ +export interface AsrBailianPreset { + id: string; + provider: "bailian" | "bailian-ws"; + label: string; + note: string; +} + +export const ASR_BAILIAN_PRESETS: readonly AsrBailianPreset[] = [ + { + id: "bailian", + provider: "bailian", + label: "阿里云百炼(chat/completions)", + note: "百炼 OpenAI 兼容模式:POST /chat/completions + input_audio(非流式,整句返回)。Base URL 需含 /compatible-mode/v1。" + }, + { + id: "bailian-ws", + provider: "bailian-ws", + label: "阿里云百炼(WebSocket 流式)", + note: "百炼实时语音识别:wss /api-ws/v1/inference,边说边出字;失败自动回退非流式通道与本地模型。" + } +]; + +export interface AsrOpenAiPreset { + id: string; + label: string; + baseUrl: string; + model: string; + apiKeyEnv: string; + note?: string; +} + +/** + * Presets for OpenAI-compatible third-party ASR endpoints. Every field stays + * editable in settings, so regional variants or renamed models keep working. + */ +export const ASR_OPENAI_PRESETS: readonly AsrOpenAiPreset[] = [ + // 小米 MiMo 已移到这里之外:它不是 /audio/transcriptions,而是专用 + // chat/completions + input_audio 协议(见 mimo-http.ts 与官方文档)。 + { + id: "siliconflow", + label: "SiliconFlow SenseVoice", + baseUrl: "https://api.siliconflow.cn/v1", + model: "FunAudioLLM/SenseVoiceSmall", + apiKeyEnv: "SILICONFLOW_API_KEY", + note: "SiliconFlow 语音识别,OpenAI 兼容接口。" + }, + { + id: "groq", + label: "Groq Whisper", + baseUrl: "https://api.groq.com/openai/v1", + model: "whisper-large-v3", + apiKeyEnv: "GROQ_API_KEY", + note: "Groq Whisper,OpenAI 兼容接口。" + }, + { + id: "custom", + label: "自定义 OpenAI 兼容", + baseUrl: "", + model: "", + apiKeyEnv: "CUSTOM_ASR_API_KEY", + note: "任何兼容 /v1/audio/transcriptions 的服务。" + } +]; + +export function findAsrPreset(id: string | undefined): AsrOpenAiPreset | undefined { + return ASR_OPENAI_PRESETS.find((preset) => preset.id === (id || "custom")); +} + +export function isAsrProviderKind(value: unknown): value is AsrProviderKind { + return value === "auto" || value === "official" || value === "openai" || value === "local" || value === "bailian" || value === "bailian-ws" || value === "mimo" || value === "local-pro"; +} + +/** Trim a URL-ish field and drop trailing slashes (mirrors the `openai` block). */ +function normalizeUrl(value: unknown): string { + return typeof value === "string" ? value.trim().replace(/\/+$/, "") : ""; +} + +/** Normalize raw (YAML/UI) ASR settings; always returns a defined object. */ +export function normalizeSpeechSettings(raw: unknown): SpeechAsrSettings { + const source = raw && typeof raw === "object" ? raw as Record : {}; + const betterRecognition = source.betterRecognition === true; + const provider = isAsrProviderKind(source.provider) ? source.provider : "auto"; + const openaiRaw = source.openai && typeof source.openai === "object" ? source.openai as Record : {}; + const openai: OpenAiAsrSettings = { + ...(typeof openaiRaw.name === "string" && openaiRaw.name.trim() ? { name: openaiRaw.name.trim() } : {}), + baseUrl: typeof openaiRaw.baseUrl === "string" ? openaiRaw.baseUrl.trim().replace(/\/+$/, "") : "", + apiKeyEnv: typeof openaiRaw.apiKeyEnv === "string" && /^[A-Za-z_][A-Za-z0-9_]*$/.test(openaiRaw.apiKeyEnv) ? openaiRaw.apiKeyEnv : "", + model: typeof openaiRaw.model === "string" ? openaiRaw.model.trim() : "", + ...(typeof openaiRaw.language === "string" && openaiRaw.language.trim() ? { language: openaiRaw.language.trim() } : {}) + }; + // Only emit the third-party block when it carries a usable endpoint, so an + // untouched config stays `{ betterRecognition, provider }` without an empty + // `openai:` mapping in the YAML. + const hasOpenAi = Boolean(openai.baseUrl || openai.model); + const bailianRaw = source.bailian && typeof source.bailian === "object" ? source.bailian as Record : {}; + const language = typeof bailianRaw.language === "string" ? bailianRaw.language.trim() : ""; + const bailian: BailianAsrSettings = { + ...(typeof bailianRaw.name === "string" && bailianRaw.name.trim() ? { name: bailianRaw.name.trim() } : {}), + apiKeyEnv: typeof bailianRaw.apiKeyEnv === "string" && /^[A-Za-z_][A-Za-z0-9_]*$/.test(bailianRaw.apiKeyEnv) ? bailianRaw.apiKeyEnv : "", + baseUrl: normalizeUrl(bailianRaw.baseUrl), + wsUrl: normalizeUrl(bailianRaw.wsUrl), + model: typeof bailianRaw.model === "string" ? bailianRaw.model.trim() : "", + streamModel: typeof bailianRaw.streamModel === "string" ? bailianRaw.streamModel.trim() : "", + ...(language ? { language } : {}), + ...(bailianRaw.enableItn === true ? { enableItn: true } : {}) + }; + const hasBailian = Boolean(bailian.baseUrl || bailian.wsUrl || bailian.model || bailian.streamModel); + + // MiMo(专用 chat/completions + input_audio 协议)。 + const mimoRaw = source.mimo && typeof source.mimo === "object" ? source.mimo as Record : {}; + const mimo: MimoAsrSettings = { + ...(typeof mimoRaw.name === "string" && mimoRaw.name.trim() ? { name: mimoRaw.name.trim() } : {}), + baseUrl: typeof mimoRaw.baseUrl === "string" && mimoRaw.baseUrl.trim() ? normalizeUrl(mimoRaw.baseUrl) : MIMO_ASR_DEFAULTS.baseUrl, + apiKeyEnv: typeof mimoRaw.apiKeyEnv === "string" && /^[A-Za-z_][A-Za-z0-9_]*$/.test(mimoRaw.apiKeyEnv) ? mimoRaw.apiKeyEnv : MIMO_ASR_DEFAULTS.apiKeyEnv, + model: typeof mimoRaw.model === "string" && mimoRaw.model.trim() ? mimoRaw.model.trim() : MIMO_ASR_DEFAULTS.model, + ...(typeof mimoRaw.language === "string" && mimoRaw.language.trim() ? { language: mimoRaw.language.trim() } : {}) + }; + const hasMimo = Boolean(mimo.baseUrl && mimo.model && mimo.apiKeyEnv) || mimoRaw.apiKeyEnv !== undefined || mimoRaw.baseUrl !== undefined; + + // 自定义回退链:白名单内的 provider id,按用户顺序保留。 + const chain = Array.isArray(source.chain) + ? source.chain.filter((id): id is string => typeof id === "string" && isAsrProviderKind(id) && id !== "auto") + : undefined; + + // 嘈杂环境参数。 + const noiseRaw = source.noise && typeof source.noise === "object" ? source.noise as Record : {}; + const noise: AsrNoiseSettings | undefined = (source.noise && typeof source.noise === "object") ? { + profile: noiseRaw.profile === "classroom" || noiseRaw.profile === "custom" ? noiseRaw.profile : "standard", + ...(typeof noiseRaw.speechNoiseThreshold === "number" && Number.isFinite(noiseRaw.speechNoiseThreshold) + ? { speechNoiseThreshold: Math.min(1, Math.max(-1, noiseRaw.speechNoiseThreshold)) } + : {}), + ...(noiseRaw.vadModel === "near_meeting_16k" || noiseRaw.vadModel === "far_field_meeting_16k" ? { vadModel: noiseRaw.vadModel } : {}), + ...(Array.isArray(noiseRaw.hotwords) + ? { hotwords: noiseRaw.hotwords.filter((word): word is string => typeof word === "string" && word.trim().length > 0).slice(0, 50).map((word) => word.trim()) } + : {}) + } : undefined; + + // 输入/输出音频设备。 + const audioRaw = source.audio && typeof source.audio === "object" ? source.audio as Record : {}; + const audio: AsrAudioDeviceSettings | undefined = (source.audio && typeof source.audio === "object") ? { + ...(typeof audioRaw.input === "string" && audioRaw.input ? { input: audioRaw.input } : {}), + ...(typeof audioRaw.output === "string" && audioRaw.output ? { output: audioRaw.output } : {}) + } : undefined; + + return { + betterRecognition, + provider, + ...(hasOpenAi ? { openai } : {}), + ...(hasBailian ? { bailian } : {}), + ...(hasMimo ? { mimo } : {}), + ...(chain && chain.length ? { chain } : {}), + ...(noise ? { noise } : {}), + ...(audio ? { audio } : {}) + }; +} + +/** An OpenAI-compatible provider is usable when endpoint, model and key name exist. */ +export function isOpenAiAsrConfigured(settings: SpeechAsrSettings | undefined): boolean { + const openai = settings?.openai; + return Boolean(openai && openai.baseUrl && openai.model && openai.apiKeyEnv); +} diff --git a/src/asr/sherpa-loader.ts b/src/asr/sherpa-loader.ts new file mode 100644 index 0000000..1426068 --- /dev/null +++ b/src/asr/sherpa-loader.ts @@ -0,0 +1,62 @@ +/** + * Resilient loader for the local sherpa-onnx speech engine. + * + * `sherpa-onnx` is a CommonJS package wrapping an Emscripten WASM runtime; its + * on-disk layout (wasm + model files) must stay intact at require time. Bundling + * it into the main-process bundle breaks the wasm lookup, so the loader always + * resolves the package at runtime — via `createRequire` in ESM hosts, the host + * `require` in CJS bundles, or dynamic import as a last resort — and reports a + * readable error instead of crashing the app. + */ +import { createRequire } from "node:module"; + +type SherpaModule = typeof import("sherpa-onnx"); + +const MODULE_ID = "sherpa-onnx"; +let cached: SherpaModule | undefined; + +function describe(error: unknown): string { + return error instanceof Error ? error.message : String(error); +} + +export async function loadSherpaOnnx(): Promise { + if (cached) return cached; + const failures: string[] = []; + + // 1) createRequire anchored at this file: works under plain `node` (CLI) and + // in the Electron main bundle, because node_modules is still resolvable + // from the emitted file location. + try { + const requireAtHere = createRequire(import.meta.url); + cached = requireAtHere(MODULE_ID) as SherpaModule; + return cached; + } catch (error) { + failures.push(`createRequire: ${describe(error)}`); + } + + // 2) Host-provided require: electron-vite emits a CJS bundle where a real + // `require` exists on the module scope. The dynamic member access keeps + // bundlers from rewriting it. + const hostRequire = (globalThis as { require?: (id: string) => unknown }).require; + if (typeof hostRequire === "function") { + try { + cached = hostRequire(MODULE_ID) as SherpaModule; + return cached; + } catch (error) { + failures.push(`require: ${describe(error)}`); + } + } + + // 3) Dynamic import: last resort for ESM-only hosts. + try { + cached = (await import(/* @vite-ignore */ MODULE_ID)) as SherpaModule; + return cached; + } catch (error) { + failures.push(`import: ${describe(error)}`); + } + + throw new Error( + `本地语音引擎 sherpa-onnx 加载失败(${failures.join(";")})。` + + "已安装的桌面版会把引擎放在 resources 目录;开发模式请先执行 npm install。" + ); +} diff --git a/src/asr/types.ts b/src/asr/types.ts new file mode 100644 index 0000000..9cb5a07 --- /dev/null +++ b/src/asr/types.ts @@ -0,0 +1,49 @@ +/** + * Provider-agnostic ASR (speech-to-text) contracts. + * + * The core layer knows nothing about Electron: providers receive Float32 PCM + * chunks at 16 kHz and emit events through a sink. Electron-specific wiring + * (windows, IPC, telemetry) lives in `src/electron/`. + */ + +export type AsrEvent = + | { type: "ready"; provider?: string } + | { type: "partial"; text: string; provider?: string } + | { type: "final"; text: string; provider?: string } + | { type: "log"; message: string } + | { type: "stopped" } + | { type: "error"; message: string }; + +export type AsrEventSink = (event: AsrEvent) => void; + +export interface AsrSession { + readonly providerId: string; + /** Feed 16 kHz mono Float32 samples. */ + push(samples: Float32Array): void; + /** Finish the utterance; resolves after the final text has been emitted. */ + stop(): Promise; + /** Abort without emitting a final result. */ + cancel(): void; +} + +export interface AsrTestResult { + ok: boolean; + message: string; +} + +export interface AsrProvider { + /** Stable identifier used in settings and logs (`official`, `openai`, `local`). */ + readonly id: string; + /** Human-readable label for logs and diagnostics. */ + readonly label: string; + /** Whether the provider has everything it needs (config present, logged in…). */ + isConfigured(): boolean; + /** + * Start an utterance. The returned promise resolves once the provider is + * ready to accept audio and rejects when it cannot start, so the manager can + * fall back to the next provider in the chain. + */ + start(sink: AsrEventSink): Promise; + /** Optional connectivity probe used by the settings page. */ + test?(): Promise; +} diff --git a/src/asr/voice-wake.ts b/src/asr/voice-wake.ts new file mode 100644 index 0000000..6b5b517 --- /dev/null +++ b/src/asr/voice-wake.ts @@ -0,0 +1,125 @@ +/** + * Voice wake ("小泽同学") keyword spotting with the bundled sherpa-onnx KWS model. + * + * Loading is lazy and async: a missing or broken engine now surfaces as a + * rejected promise with a readable message instead of an import-time crash. + */ +import fs from "node:fs"; +import path from "node:path"; +import { pinyin } from "pinyin-pro"; +import { loadSherpaOnnx } from "./sherpa-loader.js"; + +const KWS_DIR = "sherpa-onnx-kws-zipformer-zh-en-3M-2025-12-20"; + +type Kws = ReturnType; +type KwsStream = ReturnType; + +export interface VoiceWakeOptions { + extraRoots?: string[]; + log?: (message: string) => void; +} + +export function keywordTokens(phrase: string): string { + const syllables = pinyin(phrase.replace(/\s+/g, ""), { toneType: "symbol", type: "array" }) as string[]; + const initials = ["zh", "ch", "sh", "b", "p", "m", "f", "d", "t", "n", "l", "g", "k", "h", "j", "q", "x", "r", "z", "c", "s", "y", "w"]; + return syllables.map((syllable) => { + const initial = initials.find((candidate) => syllable.startsWith(candidate)) || ""; + return `${initial} ${syllable.slice(initial.length)}`; + }).join(" "); +} + +export class VoiceWakeEngine { + private kws: Kws | undefined; + private stream: KwsStream | undefined; + private detected: (() => void) | undefined; + private startedAt = 0; + private audioFrames = 0; + private awaitingFirstAudio = false; + private lastHeartbeatAt = 0; + private readonly options: VoiceWakeOptions; + + constructor(options: VoiceWakeOptions = {}) { + this.options = options; + } + + get active(): boolean { + return Boolean(this.kws && this.stream); + } + + async start(phrase: string, onDetected: () => void): Promise { + this.detected = onDetected; + if (this.kws) { + this.options.log?.(`[voice-wake] local KWS already active phrase=${phrase}`); + return; + } + const resourcesPath = (process as NodeJS.Process & { resourcesPath?: string }).resourcesPath; + const candidates = [ + ...(this.options.extraRoots || []), + ...(resourcesPath ? [resourcesPath] : []), + process.cwd() + ]; + const root = candidates.map((candidate) => path.join(candidate, "models", KWS_DIR)).find((candidate) => fs.existsSync(candidate)); + if (!root) throw new Error(`找不到语音唤醒模型 ${KWS_DIR},已搜索:${candidates.map((candidate) => path.join(candidate, "models")).join("、")}`); + const sherpa = await loadSherpaOnnx(); + this.kws = sherpa.createKws({ + featConfig: { samplingRate: 16_000, featureDim: 80 }, + modelConfig: { + transducer: { + encoder: path.join(root, "encoder-epoch-13-avg-2-chunk-16-left-64.int8.onnx"), + decoder: path.join(root, "decoder-epoch-13-avg-2-chunk-16-left-64.onnx"), + joiner: path.join(root, "joiner-epoch-13-avg-2-chunk-16-left-64.int8.onnx") + }, + tokens: path.join(root, "tokens.txt"), + provider: "cpu", + numThreads: 1, + modelingUnit: "ppinyin" + }, + maxActivePaths: 4, + numTrailingBlanks: 1, + keywordsScore: 1.5, + keywordsThreshold: 0.55, + keywords: `${keywordTokens(phrase)} @${phrase}` + }); + this.stream = this.kws.createStream(); + this.startedAt = Date.now(); + this.audioFrames = 0; + this.awaitingFirstAudio = true; + this.lastHeartbeatAt = this.startedAt; + this.options.log?.(`[voice-wake] local KWS ready phrase=${phrase}`); + } + + feed(samples: Float32Array): void { + const kws = this.kws; + const stream = this.stream; + if (!kws || !stream) return; + const now = Date.now(); + this.audioFrames += 1; + if (this.awaitingFirstAudio) { + this.awaitingFirstAudio = false; + this.options.log?.(`[voice-wake] local KWS received first audio elapsed=${now - this.startedAt}ms`); + } else if (now - this.lastHeartbeatAt >= 15_000) { + this.lastHeartbeatAt = now; + this.options.log?.(`[voice-wake] local KWS audio heartbeat frames=${this.audioFrames} elapsed=${now - this.startedAt}ms`); + } + stream.acceptWaveform(16_000, samples); + while (kws.isReady(stream)) kws.decode(stream); + const result = kws.getResult(stream); + if (result.keyword) { + this.options.log?.(`[voice-wake] local KWS detected keyword=${result.keyword} frames=${this.audioFrames}`); + // Reset before invoking the callback. The callback may stop the engine + // and release the KWS instance immediately. + kws.reset(stream); + this.detected?.(); + } + } + + stop(): void { + this.detected = undefined; + this.stream = undefined; + try { this.kws?.free(); } catch { /* already freed */ } + this.kws = undefined; + this.startedAt = 0; + this.audioFrames = 0; + this.awaitingFirstAudio = false; + } +} diff --git a/src/asr/wav.test.ts b/src/asr/wav.test.ts new file mode 100644 index 0000000..014c97c --- /dev/null +++ b/src/asr/wav.test.ts @@ -0,0 +1,40 @@ +import test from "node:test"; +import assert from "node:assert/strict"; +import { encodeWav, mergeSamples, ASR_SAMPLE_RATE } from "./wav.js"; + +test("encodeWav writes a canonical 44-byte RIFF header", () => { + const samples = new Float32Array([0, 0.5, -0.5, 1, -1]); + const wav = encodeWav(samples); + assert.equal(wav.byteLength, 44 + samples.length * 2); + const chunk = (offset: number, length: number): string => Array.from(wav.slice(offset, offset + length)).map((byte) => String.fromCharCode(byte)).join(""); + assert.equal(chunk(0, 4), "RIFF"); + assert.equal(chunk(8, 4), "WAVE"); + assert.equal(chunk(12, 4), "fmt "); + assert.equal(chunk(36, 4), "data"); + const view = new DataView(wav.buffer, wav.byteOffset, wav.byteLength); + assert.equal(view.getUint32(24, true), ASR_SAMPLE_RATE); + assert.equal(view.getUint16(22, true), 1); // mono + assert.equal(view.getUint16(34, true), 16); // bits per sample +}); + +test("encodeWav clamps out-of-range samples to int16 bounds", () => { + const wav = encodeWav(new Float32Array([2, -2, NaN])); + const view = new DataView(wav.buffer, wav.byteOffset, wav.byteLength); + assert.equal(view.getInt16(44, true), 0x7fff); + assert.equal(view.getInt16(46, true), -0x8000); + assert.equal(view.getInt16(48, true), 0); // NaN quantizes to silence +}); + +test("mergeSamples concatenates without mutating inputs", () => { + const a = new Float32Array([1, 2]); + const b = new Float32Array([3]); + const merged = mergeSamples([a, b]); + assert.deepEqual([...merged], [1, 2, 3]); + assert.deepEqual([...a], [1, 2]); + assert.equal(merged.byteOffset, 0); +}); + +test("mergeSamples handles an empty list", () => { + const merged = mergeSamples([]); + assert.equal(merged.length, 0); +}); diff --git a/src/asr/wav.ts b/src/asr/wav.ts new file mode 100644 index 0000000..496e239 --- /dev/null +++ b/src/asr/wav.ts @@ -0,0 +1,58 @@ +/** Encode 16 kHz mono Float32 PCM as a 16-bit little-endian WAV file. */ + +export const ASR_SAMPLE_RATE = 16_000; + +export function encodeWav(samples: Float32Array, sampleRate = ASR_SAMPLE_RATE): Uint8Array { + const dataBytes = samples.length * 2; + const buffer = new ArrayBuffer(44 + dataBytes); + const view = new DataView(buffer); + const writeAscii = (offset: number, text: string): void => { + for (let index = 0; index < text.length; index += 1) view.setUint8(offset + index, text.charCodeAt(index)); + }; + writeAscii(0, "RIFF"); + view.setUint32(4, 36 + dataBytes, true); + writeAscii(8, "WAVE"); + writeAscii(12, "fmt "); + view.setUint32(16, 16, true); // PCM chunk size + view.setUint16(20, 1, true); // PCM format + view.setUint16(22, 1, true); // mono + view.setUint32(24, sampleRate, true); + view.setUint32(28, sampleRate * 2, true); // byte rate + view.setUint16(32, 2, true); // block align + view.setUint16(34, 16, true); // bits per sample + writeAscii(36, "data"); + view.setUint32(40, dataBytes, true); + let offset = 44; + for (const sample of samples) { + // Clamp and dither-free quantization to signed 16-bit. + const clamped = Math.max(-1, Math.min(1, Number.isFinite(sample) ? sample : 0)); + view.setInt16(offset, clamped < 0 ? clamped * 0x8000 : clamped * 0x7fff, true); + offset += 2; + } + return new Uint8Array(buffer); +} + +/** Encode Float32 samples as 16-bit little-endian mono PCM (raw stream frames). */ +export function encodePcm16(samples: Float32Array): Uint8Array { + const bytes = new Uint8Array(samples.length * 2); + const view = new DataView(bytes.buffer); + for (let index = 0; index < samples.length; index += 1) { + const sample = samples[index]; + const clamped = Math.max(-1, Math.min(1, Number.isFinite(sample) ? sample : 0)); + view.setInt16(index * 2, clamped < 0 ? clamped * 0x8000 : clamped * 0x7fff, true); + } + return bytes; +} + +/** Concatenate Float32 chunks without mutating the inputs. */ +export function mergeSamples(chunks: readonly Float32Array[]): Float32Array { + let length = 0; + for (const chunk of chunks) length += chunk.length; + const merged = new Float32Array(length); + let offset = 0; + for (const chunk of chunks) { + merged.set(chunk, offset); + offset += chunk.length; + } + return merged; +} diff --git a/src/classisland.test.ts b/src/classisland.test.ts index 0706885..c4bad0a 100644 --- a/src/classisland.test.ts +++ b/src/classisland.test.ts @@ -38,12 +38,12 @@ function classIslandHealthResponse(input: string | URL): Response | undefined { : undefined; } -test("ClassIsland versions enforce the 2.1.1.0 minimum", () => { - assert.equal(compareClassIslandVersions("2.1.1.0", "2.1.1.0"), 0); - assert.equal(compareClassIslandVersions("2.1.1.1", "2.1.1.0") > 0, true); - assert.equal(compareClassIslandVersions("2.1.0.9", "2.1.1.0") < 0, true); - assert.equal(isCompatibleClassIslandVersion("2.1.1.0"), true); - assert.equal(isCompatibleClassIslandVersion("2.1.0.9"), false); +test("ClassIsland versions enforce the 2.0.0.0 minimum", () => { + assert.equal(compareClassIslandVersions("2.0.0.0", "2.0.0.0"), 0); + assert.equal(compareClassIslandVersions("2.0.0.1", "2.0.0.0") > 0, true); + assert.equal(compareClassIslandVersions("1.9.0.9", "2.0.0.0") < 0, true); + assert.equal(isCompatibleClassIslandVersion("2.0.0.0"), true); + assert.equal(isCompatibleClassIslandVersion("1.9.0.9"), false); assert.equal(isCompatibleClassIslandVersion(undefined), false); }); @@ -74,8 +74,8 @@ test("discovers multiple ClassIsland versions and marks old versions incompatibl const paths = ["C:\\Portable\\ClassIsland.exe", "D:\\Old\\ClassIsland.exe", "C:\\Program Files\\ClassIsland\\ClassIsland.exe"]; const versions: Record = { [paths[0]]: "2.1.1.0", - [paths[1]]: "2.0.4.0", - [paths[2]]: "2.1.1.0" + [paths[1]]: "1.9.0.0", + [paths[2]]: "2.0.4.0" }; const found = await discoverClassIslandInstallations({ platform: "win32", @@ -91,7 +91,9 @@ test("discovers multiple ClassIsland versions and marks old versions incompatibl assert.equal(found.find((item) => item.executablePath === paths[0])?.isRunning, true); assert.deepEqual(found.find((item) => item.executablePath === paths[0])?.launchArgs, ["--quiet"]); assert.equal(found.find((item) => item.executablePath === paths[1])?.compatible, false); - assert.match(found.find((item) => item.executablePath === paths[1])?.reason || "", /2\.1\.1\.0/); + assert.match(found.find((item) => item.executablePath === paths[1])?.reason || "", /2\.0\.0\.0/); + // 2.0.x releases (e.g. the 2.0.4.0 the issue reporter runs) are now supported. + assert.equal(found.find((item) => item.executablePath === paths[2])?.compatible, true); }); test("maps the running ClassIsland.Desktop process back to its launcher and closes the real instance", async () => { @@ -289,7 +291,7 @@ test("does not write a package when ClassIsland is too old or the digest is inva try { const exe = path.win32.join(root, "ClassIsland.exe"); fs.writeFileSync(exe, "test executable"); - const oldInstaller = new ClassIslandInstaller({ platform: "win32", executablePaths: [exe], runningProcesses: [], versionOf: () => "2.1.0.9", exists: (candidate) => fs.existsSync(candidate) }); + const oldInstaller = new ClassIslandInstaller({ platform: "win32", executablePaths: [exe], runningProcesses: [], versionOf: () => "1.9.0.0", exists: (candidate) => fs.existsSync(candidate) }); const [oldTarget] = await oldInstaller.detect(); const [oldResult] = await oldInstaller.install([oldTarget.id]); assert.equal(oldResult.action, "skipped"); diff --git a/src/classisland.ts b/src/classisland.ts index 15d4875..0ad6660 100644 --- a/src/classisland.ts +++ b/src/classisland.ts @@ -10,7 +10,12 @@ import { DEFAULT_MARKETPLACE_PROXY_URL, describeDownloadAttempt, marketplaceRequ export const CLASSISLAND_PLUGIN_REPOSITORY = "SECTL/ClassIsland-SecAgent-Plugin"; export const CLASSISLAND_PLUGIN_ID = "classisland.secagent"; export const CLASSISLAND_PLUGIN_ASSET_NAME = "ClassIsland.SecAgent.Plugin.cipx"; -export const MIN_CLASSISLAND_VERSION = "2.1.1.0"; +// The ClassIsland-side plugin (manifest.yml) declares apiVersion: 2.0.0.0, and +// ClassIsland 2.0.0.0 already enforces that plugins target API version +// >= 2.0.0.0 (PluginService.InitializePlugins) and lists only those in the +// marketplace (PluginMarketService). Every host API the plugin uses exists +// unchanged in 2.0.0.0, so 2.0.0.0 is the real lower bound. +export const MIN_CLASSISLAND_VERSION = "2.0.0.0"; export const CLASSISLAND_RELEASE_API_URL = `https://api.github.com/repos/${CLASSISLAND_PLUGIN_REPOSITORY}/releases/latest`; const CLASSISLAND_RELEASE_PAGE_URL = `https://github.com/${CLASSISLAND_PLUGIN_REPOSITORY}/releases/latest`; diff --git a/src/config.test.ts b/src/config.test.ts index 51246a2..118e43b 100644 --- a/src/config.test.ts +++ b/src/config.test.ts @@ -4,8 +4,9 @@ import os from "node:os"; import path from "node:path"; import test from "node:test"; import YAML from "yaml"; -import { configPath, initializeWorkspace, isOnboardingComplete, loadConfig, markOnboardingComplete, oobeProgressPath, readOobeProgress, readSettings, saveOobeProgress, saveSettings, type OobeProgress } from "./config.js"; +import { configPath, initializeWorkspace, isOnboardingComplete, loadConfig, markOnboardingComplete, OFFICIAL_VISION_MODEL, oobeProgressPath, readOobeProgress, readSettings, resolveVisionAgentConfig, saveOobeProgress, saveSettings, type OobeProgress } from "./config.js"; import { SYSTEM_PROMPT } from "./system-prompt.js"; +import type { SecAgentConfig } from "./types.js"; function temporaryWorkspace(): string { return fs.mkdtempSync(path.join(os.tmpdir(), "secagent-oobe-")); @@ -121,3 +122,87 @@ test("defaults and persists Windows update preferences", () => { fs.rmSync(workspace, { recursive: true, force: true }); } }); + +test("persists the vision model id with the settings", () => { + const workspace = temporaryWorkspace(); + try { + initializeWorkspace(workspace); + const settings = readSettings(workspace); + const saved = saveSettings(workspace, { ...settings, visionModelId: "custom-provider:vision" }); + assert.equal(saved.visionModelId, "custom-provider:vision"); + assert.equal(readSettings(workspace).visionModelId, "custom-provider:vision"); + // Clearing the value removes the persisted key. + const cleared = saveSettings(workspace, { ...settings, visionModelId: undefined }); + assert.equal(cleared.visionModelId, undefined); + assert.equal(readSettings(workspace).visionModelId, undefined); + } finally { + fs.rmSync(workspace, { recursive: true, force: true }); + } +}); + +function multiModelConfig(overrides: Partial = {}): SecAgentConfig { + return { + version: 1, + workspace: "/tmp/secagent-test-ws", + agent: { + provider: "openai-compatible", + model: "main", + apiKeyEnv: "MAIN_KEY", + baseUrl: "https://main.test/v1", + endpoint: "/chat/completions", + maxTokens: 100, + systemPrompt: "main", + models: [ + { id: "main", provider: "openai-compatible", model: "main", apiKeyEnv: "MAIN_KEY", baseUrl: "https://main.test/v1", endpoint: "/chat/completions", maxTokens: 100 }, + { id: "vision", provider: "openai-compatible", model: "vision-model", apiKeyEnv: "VISION_KEY", baseUrl: "https://vision.test/v1", endpoint: "/chat/completions", maxTokens: 100 }, + { id: "sectl-official:deepseek-v4-flash", provider: "openai-responses", model: "deepseek-v4-flash", apiKeyEnv: "SECTL_OFFICIAL_TOKEN", baseUrl: "https://relay.test/v1", endpoint: "/responses", maxTokens: 100 } + ] + }, + mcp: { servers: {} }, + ...overrides + } as SecAgentConfig; +} + +test("resolveVisionAgentConfig honors an explicit vision model id without mutating the source config", () => { + const config = multiModelConfig({ defaults: { modelId: "main", visionModelId: "vision" } }); + const vision = resolveVisionAgentConfig(config); + assert.ok(vision); + assert.equal(vision.agent.model, "vision-model"); + assert.equal(vision.agent.apiKeyEnv, "VISION_KEY"); + assert.equal(vision.agent.baseUrl, "https://vision.test/v1"); + // The caller's config keeps its own model. + assert.equal(config.agent.model, "main"); +}); + +test("resolveVisionAgentConfig returns undefined for a stale vision model id", () => { + const config = multiModelConfig({ defaults: { visionModelId: "does-not-exist" } }); + assert.equal(resolveVisionAgentConfig(config), undefined); +}); + +test("resolveVisionAgentConfig falls back to the official virtual-vision model in official mode", () => { + const previous = process.env.SECTL_OFFICIAL_TOKEN; + process.env.SECTL_OFFICIAL_TOKEN = "test-token"; + try { + const config = multiModelConfig({ defaults: { customModelMode: false } }); + const vision = resolveVisionAgentConfig(config); + assert.ok(vision); + assert.equal(vision.agent.model, OFFICIAL_VISION_MODEL); + assert.equal(vision.agent.apiKeyEnv, "SECTL_OFFICIAL_TOKEN"); + assert.equal(vision.agent.endpoint, "/responses"); + } finally { + if (previous === undefined) delete process.env.SECTL_OFFICIAL_TOKEN; + else process.env.SECTL_OFFICIAL_TOKEN = previous; + } +}); + +test("resolveVisionAgentConfig does not fall back in custom model mode", () => { + const previous = process.env.SECTL_OFFICIAL_TOKEN; + process.env.SECTL_OFFICIAL_TOKEN = "test-token"; + try { + const config = multiModelConfig({ defaults: { customModelMode: true } }); + assert.equal(resolveVisionAgentConfig(config), undefined); + } finally { + if (previous === undefined) delete process.env.SECTL_OFFICIAL_TOKEN; + else process.env.SECTL_OFFICIAL_TOKEN = previous; + } +}); diff --git a/src/config.ts b/src/config.ts index 4a989a1..dda3d89 100644 --- a/src/config.ts +++ b/src/config.ts @@ -3,12 +3,22 @@ import path from "node:path"; import YAML from "yaml"; import { expandPath } from "./paths.js"; import type { McpServerConfig, ModelProfile, ProviderConfig, ReasoningEffort, SecAgentConfig, TelemetrySettings, UpdatePreferences } from "./types.js"; +import { normalizeSpeechSettings, type BailianAsrSettings, type MimoAsrSettings, type OpenAiAsrSettings, type SpeechAsrSettings } from "./asr/settings.js"; import type { GoogleModelInfo } from "./google-models.js"; import { DEFAULT_WAKE_HOTKEY, normalizeWakeHotkey } from "./wake-hotkey.js"; +import { normalizeResilienceSettings } from "./resilience.js"; +import { normalizeToolGuardSettings } from "./tool-guard.js"; import { SYSTEM_PROMPT } from "./system-prompt.js"; export const DEFAULT_GOOGLE_MODEL = "gemini-2.5-flash"; export const DEFAULT_MAX_TOKENS = 16_384; +/** + * Client-facing virtual vision model served by the official relay. The relay routes this + * model id to a vision-capable upstream, so the client does not need to know which concrete + * model is behind it. Only the client part lives in this repository; the relay contract is + * documented in the README/settings help text. + */ +export const OFFICIAL_VISION_MODEL = "virtual-vision"; const ONBOARDING_MARKER = ".oobe-complete"; const OOBE_PROGRESS_FILE = ".oobe-progress.json"; const LEGACY_AGENT_MODEL_FIELDS = ["provider", "model", "apiKeyEnv", "baseUrl", "endpoint", "anthropicVersion", "maxTokens"] as const; @@ -23,6 +33,49 @@ export const PROJECT_ENV_FILE = BUNDLED_ENV_FILES.find((file) => fs.existsSync(f if (fs.existsSync(PROJECT_ENV_FILE)) loadEnvFile(PROJECT_ENV_FILE, "project"); export const DEFAULT_TTS_VOICE = "zh-CN-XiaoxiaoNeural"; export const DEFAULT_TTS_RATE = "+0%"; + +/** TTS provider kinds allowed in the YAML block (mirror of tts/types.ts). */ +const TTS_KINDS = new Set(["edge", "windows", "mimo", "bailian"]); + +/** + * Normalize the `tts:` YAML block into a full TtsSettings, keeping provider + * and the ordered fallback chain plus per-provider sub-blocks intact. + */ +function normalizeTtsBlock(raw: unknown): import("./tts/types.js").TtsSettings { + const source = raw && typeof raw === "object" ? raw as Record : {}; + const text = (value: unknown, fallback: string): string => (typeof value === "string" && value.trim() ? value.trim() : fallback); + const provider = typeof source.provider === "string" && TTS_KINDS.has(source.provider) ? source.provider as "edge" | "windows" | "mimo" | "bailian" : "edge"; + const chain = Array.isArray(source.chain) + ? (source.chain.filter((kind): kind is "edge" | "windows" | "mimo" | "bailian" => typeof kind === "string" && TTS_KINDS.has(kind)) as Array<"edge" | "windows" | "mimo" | "bailian">) + : (["edge", "windows"] as Array<"edge" | "windows" | "mimo" | "bailian">); + const sub = (key: string): Record => source[key] && typeof source[key] === "object" ? source[key] as Record : {}; + const windows = sub("windows"); + const mimo = sub("mimo"); + const bailian = sub("bailian"); + const has = (block: Record): boolean => Object.keys(block).length > 0; + return { + provider, + chain: chain.length ? chain : [provider], + voice: text(source.voice, DEFAULT_TTS_VOICE), + rate: text(source.rate, DEFAULT_TTS_RATE), + ...(has(windows) ? { windows: { ...(typeof windows.voice === "string" && windows.voice.trim() ? { voice: windows.voice.trim() } : {}) } } : {}), + ...(has(mimo) ? { mimo: { + ...(typeof mimo.apiKeyEnv === "string" && /^[A-Za-z_][A-Za-z0-9_]*$/.test(mimo.apiKeyEnv) ? { apiKeyEnv: mimo.apiKeyEnv } : {}), + ...(typeof mimo.baseUrl === "string" && mimo.baseUrl.trim() ? { baseUrl: mimo.baseUrl.trim().replace(/\/+$/, "") } : {}), + ...(typeof mimo.model === "string" && mimo.model.trim() ? { model: mimo.model.trim() } : {}), + ...(typeof mimo.voice === "string" && mimo.voice.trim() ? { voice: mimo.voice.trim() } : {}), + ...(typeof mimo.format === "string" && mimo.format.trim() ? { format: mimo.format.trim() } : {}), + ...(typeof mimo.voiceDescription === "string" && mimo.voiceDescription.trim() ? { voiceDescription: mimo.voiceDescription.trim() } : {}) + } } : {}), + ...(has(bailian) ? { bailian: { + ...(typeof bailian.apiKeyEnv === "string" && /^[A-Za-z_][A-Za-z0-9_]*$/.test(bailian.apiKeyEnv) ? { apiKeyEnv: bailian.apiKeyEnv } : {}), + ...(typeof bailian.baseUrl === "string" && bailian.baseUrl.trim() ? { baseUrl: bailian.baseUrl.trim().replace(/\/+$/, "") } : {}), + ...(typeof bailian.model === "string" && bailian.model.trim() ? { model: bailian.model.trim() } : {}), + ...(typeof bailian.voice === "string" && bailian.voice.trim() ? { voice: bailian.voice.trim() } : {}), + ...(typeof bailian.format === "string" && bailian.format.trim() ? { format: bailian.format.trim() } : {}) + } } : {}) + }; +} export const DEFAULT_WAKE_PHRASE = "小泽同学"; export const DEFAULT_UPDATE_PREFERENCES: UpdatePreferences = { channel: "stable", autoCheck: true, autoDownload: true, autoInstallOnQuit: true }; // Installers for managed/education deployments can opt out before the first @@ -44,7 +97,7 @@ const template = (workspace: string): SecAgentConfig => ({ maxTokens: DEFAULT_MAX_TOKENS }] } as SecAgentConfig["agent"], - tts: { voice: DEFAULT_TTS_VOICE, rate: DEFAULT_TTS_RATE }, + tts: { provider: "edge" as const, chain: ["edge", "windows"], voice: DEFAULT_TTS_VOICE, rate: DEFAULT_TTS_RATE }, wake: { hotkey: DEFAULT_WAKE_HOTKEY, voiceEnabled: false, voicePhrase: DEFAULT_WAKE_PHRASE }, updates: { ...DEFAULT_UPDATE_PREFERENCES }, telemetry: { ...DEFAULT_TELEMETRY_SETTINGS }, @@ -167,18 +220,31 @@ export function normalizeAndValidate(raw: SecAgentConfig, workspace: string): Se // Multi-model configuration is canonical. Populate the legacy top-level fields in memory // so the runtime can keep using one normalized AgentConfig shape. if (raw?.agent?.providers?.length) { - raw.agent.models = raw.agent.providers.flatMap((provider) => provider.models.map((model) => ({ - id: `${provider.id}:${model.id}`, - name: model.name || model.id, - enabled: model.enabled, - provider: provider.provider, - model: model.id, - apiKeyEnv: provider.apiKeyEnv, - baseUrl: provider.baseUrl, - endpoint: provider.endpoint, - anthropicVersion: provider.anthropicVersion, - maxTokens: provider.maxTokens - }))); + raw.agent.models = raw.agent.providers.flatMap((provider) => { + const seenModelIds = new Set(); + // Drop duplicate model ids within one provider instead of failing the + // whole save: duplicates used to abort normalizeAndValidate with a bare + // "id 重复" error, which made every later autosave fail too until the + // YAML was fixed by hand. + const uniqueModels = provider.models.filter((model) => { + if (!model?.id || seenModelIds.has(model.id)) return false; + seenModelIds.add(model.id); + return true; + }); + return uniqueModels.map((model) => ({ + id: `${provider.id}:${model.id}`, + name: model.name || model.id, + enabled: model.enabled, + provider: provider.provider, + model: model.id, + apiKeyEnv: provider.apiKeyEnv, + baseUrl: provider.baseUrl, + endpoint: provider.endpoint, + anthropicVersion: provider.anthropicVersion, + maxTokens: provider.maxTokens, + providerName: provider.name || provider.id + })); + }); } if (raw?.agent?.models?.length) { const first = raw.agent.models[0]; @@ -207,7 +273,7 @@ export function normalizeAndValidate(raw: SecAgentConfig, workspace: string): Se } // 系统提示词写死在源码 system-prompt.ts 中,忽略工作区 YAML 里的 agent.systemPrompt。 raw.agent.systemPrompt = SYSTEM_PROMPT; - raw.tts = { voice: raw.tts?.voice || DEFAULT_TTS_VOICE, rate: raw.tts?.rate || DEFAULT_TTS_RATE }; + raw.tts = normalizeTtsBlock(raw.tts); raw.updates = { channel: raw.updates?.channel === "preview" ? "preview" : DEFAULT_UPDATE_PREFERENCES.channel, autoCheck: raw.updates?.autoCheck !== false, @@ -226,11 +292,16 @@ export function normalizeAndValidate(raw: SecAgentConfig, workspace: string): Se if (errors.length) throw new Error(`配置校验失败:${errors.join(";")}`); raw.agent.baseUrl = raw.agent.baseUrl.replace(/\/$/, ""); raw.agent.maxTokens = raw.agent.maxTokens || DEFAULT_MAX_TOKENS; + // Keep the speech/ASR block canonical (no UI-only extras like raw API keys). + raw.speech = normalizeSpeechSettings(raw.speech); + raw.resilience = normalizeResilienceSettings(raw.resilience); + raw.guard = normalizeToolGuardSettings(raw.guard); + raw.hallucination = { enabled: raw.hallucination?.enabled !== false }; for (const model of raw.agent.models ?? []) validateModelProfile(model, errors); if (raw.agent.models?.length) { const ids = new Set(); for (const model of raw.agent.models) { - if (ids.has(model.id)) errors.push(`agent.models.id 重复:${model.id}`); + if (ids.has(model.id)) errors.push(`agent.models.id 重复:${model.id}${model.providerName ? `(提供商「${model.providerName}」)` : ""}`); ids.add(model.id); model.name = model.name?.trim() || model.model; model.baseUrl = model.baseUrl.replace(/\/$/, ""); @@ -256,6 +327,8 @@ export interface ModelOption { name: string; model: string; provider: SecAgentConfig["agent"]["provider"]; + /** Display name of the provider this model belongs to, used to group the model pickers. */ + providerLabel?: string; } function commaValues(value: string | undefined): string[] { @@ -265,26 +338,27 @@ function commaValues(value: string | undefined): string[] { export function configuredModels(config: SecAgentConfig, googleModels: GoogleModelInfo[] = []): ModelOption[] { const profiles = config.agent.models?.length ? config.agent.models : [{ id: "default", name: config.agent.model, model: config.agent.model, provider: config.agent.provider, apiKeyEnv: config.agent.apiKeyEnv, baseUrl: config.agent.baseUrl } as ModelProfile]; const options: ModelOption[] = []; - let googleSeen = false; for (const profile of profiles) { if (profile.enabled === false) continue; - if (profile.provider === "google" && googleSeen) continue; - if (profile.provider === "google") googleSeen = true; + const providerLabel = profile.providerName || profile.id; const configuredNames = commaValues(profile.name); const configuredModelNames = commaValues(profile.model); + // Every provider gets its own entry now. Previously only the first Google + // profile was expanded (`googleSeen`) and later Google providers — e.g. an + // official key plus a relay — silently lost all of their models. if (profile.provider !== "google" || !googleModels.length) { const modelNames = configuredModelNames.length ? configuredModelNames : [""]; - modelNames.forEach((modelName, index) => options.push({ id: index ? `${profile.id}#${index}` : profile.id, name: configuredNames[index] || configuredNames[0] || modelName || "Google Gemini(自动选择)", model: modelName, provider: profile.provider })); + modelNames.forEach((modelName, index) => options.push({ id: index ? `${profile.id}#${index}` : profile.id, name: configuredNames[index] || configuredNames[0] || modelName || "Google Gemini(自动选择)", model: modelName, provider: profile.provider, providerLabel })); continue; } if (configuredModelNames.length) { - configuredModelNames.forEach((modelName, index) => options.push({ id: index ? `${profile.id}#${index}` : profile.id, name: configuredNames[index] || configuredNames[0] || modelName, model: modelName, provider: profile.provider })); + configuredModelNames.forEach((modelName, index) => options.push({ id: index ? `${profile.id}#${index}` : profile.id, name: configuredNames[index] || configuredNames[0] || modelName, model: modelName, provider: profile.provider, providerLabel })); continue; } for (const model of googleModels) { const modelName = model.name?.replace(/^models\//, ""); if (!modelName) continue; - options.push({ id: `google:${profile.id}:${modelName}`, name: model.displayName || modelName, model: modelName, provider: "google" }); + options.push({ id: `google:${profile.id}:${modelName}`, name: model.displayName || modelName, model: modelName, provider: "google", providerLabel }); } } return options; @@ -303,23 +377,67 @@ export function useConfiguredModel(config: SecAgentConfig, id?: string): void { config.agent = { ...config.agent, ...selected, model: dynamicModel || selectedModels[profileIndex] || selectedModels[0] || DEFAULT_GOOGLE_MODEL, maxTokens: selected.maxTokens || config.agent.maxTokens, systemPrompt: config.agent.systemPrompt, models: config.agent.models }; } +/** + * Return a copy of the config with `agent` resolved to the given model id, without + * mutating the caller's config. Unlike `useConfiguredModel`, this is safe to call for + * a secondary (e.g. vision) model while the main session keeps using its own model. + */ +export function resolveModelConfig(config: SecAgentConfig, modelId: string): SecAgentConfig { + const next = { ...config, agent: { ...config.agent } }; + useConfiguredModel(next, modelId); + if (!next.agent.models?.length) throw new Error(`未找到配置模型:${modelId}`); + return next; +} + +/** + * Resolve the dedicated image-recognition model config, if any. + * 1. An explicitly configured `defaults.visionModelId` wins. + * 2. In official (non-custom) mode, fall back to the relay's virtual-vision model so + * the feature works without any manual setup. + * Returns undefined when no vision model is configured or the id is stale — the vision + * tool is then simply not exposed to the main agent. + */ +export function resolveVisionAgentConfig(config: SecAgentConfig): SecAgentConfig | undefined { + const id = config.defaults?.visionModelId; + if (id) { + try { return resolveModelConfig(config, id); } + catch { return undefined; } + } + if (config.defaults?.customModelMode === false + && process.env.SECTL_OFFICIAL_TOKEN + && config.agent.models?.some((model) => model.id.startsWith("sectl-official:"))) { + try { return resolveModelConfig(config, `official:sectl-official:${OFFICIAL_VISION_MODEL}`); } + catch { return undefined; } + } + return undefined; +} + export interface SettingsPayload { providers: Array; /** Compatibility field for older IPC callers; the settings UI uses providers. */ models: Array; - tts: { voice: string; rate: string }; + tts: import("./tts/types.js").TtsSettings & { mimo?: import("./tts/types.js").MimoTtsSettings & { apiKey?: string; apiKeyConfigured?: boolean }; bailian?: import("./tts/types.js").BailianTtsSettings & { apiKey?: string; apiKeyConfigured?: boolean } }; wake: { hotkey: string; modelId?: string; voiceEnabled?: boolean; voicePhrase?: string }; - speech: { betterRecognition?: boolean }; + /** + * Speech-to-text settings; `openai.apiKey`/`bailian.apiKey` and their + * `apiKeyConfigured` flags are UI-only extras (keys live in the workspace .env). + */ + speech: SpeechAsrSettings & { openai?: OpenAiAsrSettings & { apiKey?: string; apiKeyConfigured?: boolean }; bailian?: BailianAsrSettings & { apiKey?: string; apiKeyConfigured?: boolean }; mimo?: MimoAsrSettings & { apiKey?: string; apiKeyConfigured?: boolean } }; updates: UpdatePreferences; telemetry: TelemetrySettings; mcp: { servers: Record }; defaultModelId?: string; defaultReasoningEffort?: ReasoningEffort; + /** Model used by the dedicated image-recognition tool when the main model has no vision input. */ + visionModelId?: string; autostart?: boolean; /** On by default: an autostart launch stays in the tray instead of opening the main window. */ autostartHidden?: boolean; /** Off by default: custom providers are ignored and the official service (login) is required. */ customModelMode?: boolean; + resilience?: import("./resilience.js").ResilienceSettings; + guard?: import("./tool-guard.js").ToolGuardSettings; + hallucinationEnabled?: boolean; } export function readSettings(workspaceInput: string): SettingsPayload { @@ -338,7 +456,8 @@ export function readSettings(workspaceInput: string): SettingsPayload { maxTokens: config.agent.maxTokens }]; const providers = config.agent.providers?.length ? config.agent.providers : groupLegacyModels(configured); - return { providers: providers.map((provider) => ({ ...provider, apiKeyConfigured: Boolean(process.env[provider.apiKeyEnv]) })), models: configured.map((model) => ({ ...model, apiKeyConfigured: Boolean(process.env[model.apiKeyEnv]) })), tts: { voice: config.tts?.voice || DEFAULT_TTS_VOICE, rate: config.tts?.rate || DEFAULT_TTS_RATE }, wake: { hotkey: config.wake?.hotkey || DEFAULT_WAKE_HOTKEY, ...(config.wake?.modelId ? { modelId: config.wake.modelId } : {}), voiceEnabled: config.wake?.voiceEnabled === true, voicePhrase: config.wake?.voicePhrase || DEFAULT_WAKE_PHRASE }, speech: { betterRecognition: config.speech?.betterRecognition === true }, updates: { ...(config.updates || DEFAULT_UPDATE_PREFERENCES) }, telemetry: { enabled: config.telemetry?.enabled !== false }, mcp: config.mcp, defaultModelId: config.defaults?.modelId, defaultReasoningEffort: config.defaults?.reasoningEffort, autostart: config.defaults?.autostart === true, autostartHidden: config.defaults?.autostartHidden !== false, customModelMode: config.defaults?.customModelMode ?? false }; + const speech = normalizeSpeechSettings(config.speech); + return { providers: providers.map((provider) => ({ ...provider, apiKeyConfigured: Boolean(process.env[provider.apiKeyEnv]) })), models: configured.map((model) => ({ ...model, apiKeyConfigured: Boolean(process.env[model.apiKeyEnv]) })), tts: { ...normalizeTtsBlock(config.tts), ...(config.tts?.mimo ? { mimo: { ...config.tts.mimo, apiKeyConfigured: Boolean(config.tts.mimo.apiKeyEnv && process.env[config.tts.mimo.apiKeyEnv]) } } : {}), ...(config.tts?.bailian ? { bailian: { ...config.tts.bailian, apiKeyConfigured: Boolean(process.env[config.tts.bailian.apiKeyEnv || "BAILIAN_API_KEY"]) } } : {}) }, wake: { hotkey: config.wake?.hotkey || DEFAULT_WAKE_HOTKEY, ...(config.wake?.modelId ? { modelId: config.wake.modelId } : {}), voiceEnabled: config.wake?.voiceEnabled === true, voicePhrase: config.wake?.voicePhrase || DEFAULT_WAKE_PHRASE }, speech: { ...speech, ...(speech.openai ? { openai: { ...speech.openai, apiKeyConfigured: Boolean(speech.openai.apiKeyEnv && process.env[speech.openai.apiKeyEnv]) } } : {}), ...(speech.bailian ? { bailian: { ...speech.bailian, apiKeyConfigured: Boolean(process.env[speech.bailian.apiKeyEnv || "BAILIAN_API_KEY"]) } } : {}), ...(speech.mimo ? { mimo: { ...speech.mimo, apiKeyConfigured: Boolean(process.env[speech.mimo.apiKeyEnv || "MIMO_API_KEY"]) } } : {}) }, updates: { ...(config.updates || DEFAULT_UPDATE_PREFERENCES) }, telemetry: { enabled: config.telemetry?.enabled !== false }, mcp: config.mcp, defaultModelId: config.defaults?.modelId, defaultReasoningEffort: config.defaults?.reasoningEffort, visionModelId: config.defaults?.visionModelId, autostart: config.defaults?.autostart === true, autostartHidden: config.defaults?.autostartHidden !== false, customModelMode: config.defaults?.customModelMode ?? false, resilience: normalizeResilienceSettings(config.resilience), guard: normalizeToolGuardSettings(config.guard), hallucinationEnabled: config.hallucination?.enabled !== false }; } function groupLegacyModels(models: ModelProfile[]): ProviderConfig[] { @@ -353,6 +472,25 @@ function groupLegacyModels(models: ModelProfile[]): ProviderConfig[] { return [...groups.values()]; } +/** + * Derive a stable, filesystem-safe env-var name for a provider from its name + * (or baseUrl host when the name has no ASCII letters, e.g. Chinese-only). + * Users never see or type this name — it only lives in the workspace .env. + */ +function deriveEnvName(name: string, baseUrl: string): string { + const slug = (source: string) => source.replace(/[^A-Za-z0-9]+/g, "_").replace(/^_+|_+$/g, "").toUpperCase().slice(0, 32); + const fromName = slug(name); + if (fromName) return `SECAGENT_${fromName}_API_KEY`; + try { + const host = slug(new URL(baseUrl).hostname.replace(/\./g, "_")); + if (host) return `SECAGENT_${host}_API_KEY`; + } catch { /* invalid/empty baseUrl — fall through */ } + return "SECAGENT_CUSTOM_API_KEY"; +} + +/** Legacy placeholder the old "new provider" form used to plant into the config. */ +const LEGACY_DEFAULT_PROVIDER_ENVS = new Set(["CUSTOM_API_KEY"]); + export function saveSettings(workspaceInput: string, payload: SettingsPayload): SettingsPayload { const workspace = expandPath(workspaceInput); const file = configPath(workspace); @@ -360,19 +498,85 @@ export function saveSettings(workspaceInput: string, payload: SettingsPayload): const inputProviders: Array = Array.isArray(payload?.providers) && payload.providers.length ? payload.providers : groupLegacyModels(payload?.models || []); if (!inputProviders.length) throw new Error("至少需要配置一个提供商"); if (!payload.mcp?.servers || typeof payload.mcp.servers !== "object") throw new Error("MCP 服务配置无效"); + // The env-var name used to be a manual text field in the settings UI, which + // forced users to invent a valid identifier before an API key could be + // saved. Auto-derive one instead whenever it is missing or still the legacy + // placeholder; explicitly configured names (yaml, presets) are preserved. + const takenEnvs = new Set(inputProviders.map((provider) => provider.apiKeyEnv).filter(Boolean)); + for (const provider of inputProviders) { + const current = provider.apiKeyEnv?.trim() || ""; + if (current && !LEGACY_DEFAULT_PROVIDER_ENVS.has(current)) continue; + let generated = deriveEnvName(provider.name || "", provider.baseUrl || ""); + if (takenEnvs.has(generated) && current !== generated) { + let suffix = 2; + while (takenEnvs.has(`${generated}_${suffix}`)) suffix++; + generated = `${generated}_${suffix}`; + } + takenEnvs.add(generated); + provider.apiKeyEnv = generated; + } const providers = inputProviders.map(({ apiKey, apiKeyConfigured: _apiKeyConfigured, ...provider }) => { if (typeof apiKey === "string" && apiKey.trim()) writeWorkspaceEnv(workspace, provider.apiKeyEnv, apiKey.trim()); return provider; }); + // Two providers sharing one env var used to silently overwrite each other's + // API key (last write wins in .env), which looked like "I fixed the key but + // the other provider broke". Refuse ambiguous saves with an explicit error. + const envOwners = new Map(); + for (const provider of providers) { + const owner = envOwners.get(provider.apiKeyEnv); + if (owner && owner !== provider.name) throw new Error(`提供商「${owner}」和「${provider.name}」使用了相同的环境变量 ${provider.apiKeyEnv},请为其中一个改用独立变量名,否则 API Key 会互相覆盖`); + envOwners.set(provider.apiKeyEnv, provider.name); + } const models = providers.flatMap((provider) => provider.models.map((model) => ({ id: `${provider.id}:${model.id}`, name: model.name || model.id, enabled: model.enabled, provider: provider.provider, model: model.id, apiKeyEnv: provider.apiKeyEnv, baseUrl: provider.baseUrl, endpoint: provider.endpoint, anthropicVersion: provider.anthropicVersion, maxTokens: provider.maxTokens }))); - const nextTts = { voice: payload.tts?.voice || DEFAULT_TTS_VOICE, rate: payload.tts?.rate || DEFAULT_TTS_RATE }; + // TTS: keep the full provider/fallback-chain block (voice/rate stay shared). + const nextTts = normalizeTtsBlock(payload.tts); + // TTS API keys follow the same env-var convention as model providers. + const inputTtsMimo = payload.tts?.mimo; + if (inputTtsMimo && typeof (inputTtsMimo as { apiKey?: string }).apiKey === "string" && (inputTtsMimo as { apiKey?: string }).apiKey!.trim()) { + const envName = /^[A-Za-z_][A-Za-z0-9_]*$/.test(inputTtsMimo.apiKeyEnv || "") ? inputTtsMimo.apiKeyEnv! : "MIMO_TTS_API_KEY"; + inputTtsMimo.apiKeyEnv = envName; + writeWorkspaceEnv(workspace, envName, (inputTtsMimo as { apiKey?: string }).apiKey!.trim()); + } + const inputTtsBailian = payload.tts?.bailian; + if (inputTtsBailian && typeof (inputTtsBailian as { apiKey?: string }).apiKey === "string" && (inputTtsBailian as { apiKey?: string }).apiKey!.trim()) { + const envName = /^[A-Za-z_][A-Za-z0-9_]*$/.test(inputTtsBailian.apiKeyEnv || "") ? inputTtsBailian.apiKeyEnv! : "BAILIAN_TTS_API_KEY"; + inputTtsBailian.apiKeyEnv = envName; + writeWorkspaceEnv(workspace, envName, (inputTtsBailian as { apiKey?: string }).apiKey!.trim()); + } const nextWake = { hotkey: normalizeWakeHotkey(payload.wake?.hotkey || DEFAULT_WAKE_HOTKEY), ...(payload.wake?.modelId ? { modelId: payload.wake.modelId } : {}), voiceEnabled: payload.wake?.voiceEnabled === true, voicePhrase: payload.wake?.voicePhrase?.trim() || DEFAULT_WAKE_PHRASE }; const canonicalAgent = { ...(raw.agent as unknown as Record), providers, models } as SecAgentConfig["agent"]; for (const field of LEGACY_AGENT_MODEL_FIELDS) delete (canonicalAgent as unknown as Record)[field]; // 系统提示词写死在源码中,保存时从工作区配置文件里移除该键。 delete (canonicalAgent as { systemPrompt?: unknown }).systemPrompt; const candidateAgent = { ...canonicalAgent, models: models.map((model) => ({ ...model })) } as SecAgentConfig["agent"]; - const nextSpeech = { betterRecognition: payload.speech?.betterRecognition === true }; + // Third-party ASR keys follow the same env-var convention as model providers. + // The env-var name is no longer a visible form field; auto-assign a stable + // default whenever the payload arrives without a valid one. + const inputOpenAi = payload.speech?.openai; + if (inputOpenAi && typeof inputOpenAi.apiKey === "string" && inputOpenAi.apiKey.trim()) { + let envName = (inputOpenAi.apiKeyEnv || "").trim(); + if (!/^[A-Za-z_][A-Za-z0-9_]*$/.test(envName)) envName = "SECAGENT_ASR_KEY"; + inputOpenAi.apiKeyEnv = envName; + writeWorkspaceEnv(workspace, envName, inputOpenAi.apiKey.trim()); + } + // 百炼 uses the documented BAILIAN_API_KEY name so headless .env setups and + // the settings UI share one variable. + const inputBailian = payload.speech?.bailian; + if (inputBailian && typeof inputBailian.apiKey === "string" && inputBailian.apiKey.trim()) { + let envName = (inputBailian.apiKeyEnv || "").trim(); + if (!/^[A-Za-z_][A-Za-z0-9_]*$/.test(envName)) envName = "BAILIAN_API_KEY"; + inputBailian.apiKeyEnv = envName; + writeWorkspaceEnv(workspace, envName, inputBailian.apiKey.trim()); + } + // MiMo ASR key: same convention, MIMO_API_KEY by default. + const inputMimo = payload.speech?.mimo; + if (inputMimo && typeof (inputMimo as { apiKey?: string }).apiKey === "string" && (inputMimo as { apiKey?: string }).apiKey!.trim()) { + const envName = /^[A-Za-z_][A-Za-z0-9_]*$/.test(inputMimo.apiKeyEnv || "") ? inputMimo.apiKeyEnv! : "MIMO_API_KEY"; + inputMimo.apiKeyEnv = envName; + writeWorkspaceEnv(workspace, envName, (inputMimo as { apiKey?: string }).apiKey!.trim()); + } + const nextSpeech = normalizeSpeechSettings(payload.speech); const currentUpdates = raw.updates || DEFAULT_UPDATE_PREFERENCES; const nextUpdates: UpdatePreferences = { channel: payload.updates?.channel === "preview" ? "preview" : payload.updates?.channel === "stable" ? "stable" : currentUpdates.channel, autoCheck: payload.updates ? payload.updates.autoCheck !== false : currentUpdates.autoCheck, autoDownload: payload.updates ? payload.updates.autoDownload !== false : currentUpdates.autoDownload, autoInstallOnQuit: payload.updates ? payload.updates.autoInstallOnQuit !== false : currentUpdates.autoInstallOnQuit }; const nextTelemetry: TelemetrySettings = { enabled: payload.telemetry?.enabled !== false }; @@ -388,7 +592,10 @@ export function saveSettings(workspaceInput: string, payload: SettingsPayload): raw.updates = nextUpdates; raw.telemetry = nextTelemetry; raw.mcp = payload.mcp; - raw.defaults = { modelId: payload.defaultModelId || undefined, reasoningEffort: payload.defaultReasoningEffort || undefined, customModelMode: Boolean(payload.customModelMode), autostart: payload.autostart === true, autostartHidden: payload.autostartHidden !== false }; + raw.defaults = { modelId: payload.defaultModelId || undefined, reasoningEffort: payload.defaultReasoningEffort || undefined, customModelMode: Boolean(payload.customModelMode), autostart: payload.autostart === true, autostartHidden: payload.autostartHidden !== false, visionModelId: payload.visionModelId || undefined }; + raw.resilience = normalizeResilienceSettings(payload.resilience); + raw.guard = normalizeToolGuardSettings(payload.guard); + raw.hallucination = { enabled: payload.hallucinationEnabled !== false }; delete (raw as SecAgentConfig & { policy?: unknown }).policy; fs.writeFileSync(file, YAML.stringify(raw), "utf8"); return readSettings(workspace); diff --git a/src/electron/main.ts b/src/electron/main.ts index 1a95bde..1292cd6 100644 --- a/src/electron/main.ts +++ b/src/electron/main.ts @@ -8,18 +8,20 @@ import crypto from "node:crypto"; import fs from "node:fs"; import os from "node:os"; import path from "node:path"; +import YAML from "yaml"; import { pathToFileURL } from "node:url"; -import { DEFAULT_WORKSPACE } from "../paths.js"; -import { configuredModels, configPath, DEFAULT_TELEMETRY_SETTINGS, initializeWorkspace, isOnboardingComplete, loadConfig, markOnboardingComplete, readOobeProgress, readSettings, saveOobeProgress, saveSettings, useConfiguredModel, writeWorkspaceEnv, type OobeProgress, type SettingsPayload } from "../config.js"; +import { DEFAULT_WORKSPACE, migrateLegacyWorkspace } from "../paths.js"; +import { configuredModels, configPath, DEFAULT_TELEMETRY_SETTINGS, initializeWorkspace, isOnboardingComplete, loadConfig, markOnboardingComplete, OFFICIAL_VISION_MODEL, readOobeProgress, readSettings, resolveVisionAgentConfig, saveOobeProgress, saveSettings, useConfiguredModel, writeWorkspaceEnv, type OobeProgress, type SettingsPayload } from "../config.js"; import { loadEnabledSkills } from "../skills.js"; import { AuditStore } from "../audit.js"; import { SecAgentRuntime, type TraceEvent } from "../runtime.js"; import type { ConversationMessage } from "../model-provider.js"; import { SessionStore, type AssistantActivity, type SessionData, type ToolCallRecord } from "../session-store.js"; -import { cancelSpeech, sendSpeechAudio, sendVoiceWakeAudio, startSpeech, startVoiceWake, stopSpeech, stopVoiceWake } from "./speech.js"; +import { cancelSpeech, configureSpeech, sendSpeechAudio, sendVoiceWakeAudio, speechChain, startSpeech, startVoiceWake, stopSpeech, stopVoiceWake, testSpeech } from "./speech.js"; +import { runSectlOAuthFlow, type SectlOAuthResult } from "./oauth.js"; import type { ChatAttachment, ReasoningEffort, UpdateState } from "../types.js"; -import { listGoogleModels } from "../google-models.js"; -import { synthesizeSpeech } from "./tts.js"; +import { listGoogleModels, type GoogleModelInfo } from "../google-models.js"; +import { synthesizeSpeech, testTts, ttsChain, listWindowsVoices, configureTts } from "./tts.js"; import { PluginManager, type SvgPreviewRequest } from "../plugin-manager.js"; import { MarketplaceClient, type MarketplaceVersion } from "../marketplace.js"; import { detectCompanionApps } from "../companion-apps.js"; @@ -37,6 +39,16 @@ import { WindowsUpdateManager } from "./update-manager.js"; import { diagnosticLogDirectory, exportDiagnosticLogs } from "./diagnostic-logs.js"; import { TelemetryClient, hashIdentifier, normalizeMessage, sanitizeStack, type TelemetryFailure } from "../telemetry.js"; +// One-time move of the legacy `~/SecAgentWorkspace` tree into the +// platform-standard data directory. Must run before anything reads the +// workspace, so it lives at module scope ahead of the first loadConfig call. +try { + const migratedTo = migrateLegacyWorkspace(); + if (migratedTo) console.info(`[paths] 已将旧工作区迁移到 ${migratedTo}`); +} catch (error) { + console.warn("[paths] 旧工作区迁移失败,将继续使用新路径", error); +} + const SENTRY_DSN = process.env.SENTRY_DSN?.trim() || ""; function readInitialTelemetryEnabled(): boolean { if (!fs.existsSync(configPath(DEFAULT_WORKSPACE))) return DEFAULT_TELEMETRY_SETTINGS.enabled; @@ -835,8 +847,12 @@ const OFFICIAL_TIER_IDS = ["virtual-fast", "virtual-standard", "virtual-deep"] a ipcMain.handle("models:list", async () => { const { config } = loadConfig(DEFAULT_WORKSPACE); - const googleProfile = config.agent.models?.find((model) => model.provider === "google"); - const googleModels = googleProfile ? await listGoogleModels(process.env[googleProfile.apiKeyEnv] || "", googleProfile.baseUrl).catch(() => []) : []; + // Pull the live catalog for every Google provider (official key, relays, ...) + // instead of only the first one — the rest used to lose all their models. + const googleProfiles = (config.agent.models || []).filter((model) => model.provider === "google"); + const googleModels = googleProfiles.length + ? (await Promise.all(googleProfiles.map((profile) => listGoogleModels(process.env[profile.apiKeyEnv] || "", profile.baseUrl).catch(() => [] as GoogleModelInfo[])))).flat() + : []; const options = configuredModels(config, googleModels).filter((option) => option.id !== "sectl-official" && !option.id.startsWith("sectl-official:")); const customModelMode = Boolean(config.defaults?.customModelMode); const token = process.env.SECTL_OFFICIAL_TOKEN; @@ -861,8 +877,11 @@ ipcMain.handle("models:list", async () => { // 自定义模型模式开启:只加入后台允许的官方真实模型与本地自定义模型。 return [...visibleRemote, ...options]; } - // 关闭:官方档位模式 —— 下拉只有快速/标准/深度三个虚拟档位,看不到具体模型。 - return visibleRemote.filter((model) => (OFFICIAL_TIER_IDS as readonly string[]).includes(model.model)); + // 关闭:官方档位模式 —— 下拉只有快速/标准/深度三个虚拟档位,看不到具体模型; + // 另外提供一个识图虚拟模型(virtual-vision),它只作为识图工具的后端模型, + // 不作为主 Agent 模型出现在前端下拉中(前端按 vision 标记过滤)。 + return visibleRemote.filter((model) => (OFFICIAL_TIER_IDS as readonly string[]).includes(model.model) || model.model === OFFICIAL_VISION_MODEL) + .map((model) => ({ ...model, vision: model.model === OFFICIAL_VISION_MODEL })); } catch { return customModelMode ? options : []; } }); ipcMain.handle("providers:list", async () => { @@ -976,111 +995,21 @@ ipcMain.handle("official:login", async (_event, email: string, password: string) const providers = current.providers.some((provider) => provider.id === "sectl-official") ? current.providers : [...current.providers, officialProvider(baseUrl)]; return saveSettings(DEFAULT_WORKSPACE, { ...current, providers }); }); -const PUBLIC_IP_ENDPOINTS = [ - "https://api.ipify.org?format=json", - "https://httpbin.org/ip", - "https://api64.ipify.org?format=json" -]; - -async function resolvePublicIpv4(): Promise { - for (const endpoint of PUBLIC_IP_ENDPOINTS) { - try { - const response = await fetch(endpoint, { signal: AbortSignal.timeout(5_000) }); - if (!response.ok) continue; - const payload = await response.json().catch(() => ({})) as { ip?: unknown; origin?: unknown }; - const candidate = String(payload.ip ?? payload.origin ?? "").split(",")[0].trim(); - if (isIPv4(candidate)) return candidate; - } catch { - // Try the next public-IP provider. - } - } - throw new Error("无法获取本机公网 IPv4,请检查网络连接后重试"); -} -async function runSectlOAuthLogin(): Promise<{ accessToken: string; userId?: string; email?: string; name?: string }> { +async function runSectlOAuthLogin(): Promise { loadConfig(DEFAULT_WORKSPACE); - const relayUrl = (process.env.SECTL_OFFICIAL_API_URL || "").replace(/\/$/, ""); - const oauthUrl = (process.env.SECTL_OAUTH_API_URL || "https://appwrite.sectl.cn").replace(/\/$/, ""); - const oauthWebUrl = (process.env.SECTL_OAUTH_WEB_URL || "https://sectl.cn").replace(/\/$/, ""); - const clientId = process.env.SECTL_OFFICIAL_CLIENT_ID || ""; - const port = Number(process.env.SECTL_OAUTH_CALLBACK_PORT || 49152); - if (!relayUrl) throw new Error("请在 SecAgent .env 配置 SECTL_OFFICIAL_API_URL"); - if (!clientId) throw new Error("请在 SecAgent .env 配置 SECTL_OFFICIAL_CLIENT_ID"); - if (!Number.isInteger(port) || port < 49152 || port > 65535) throw new Error("SECTL_OAUTH_CALLBACK_PORT 必须是 49152-65535 的固定端口"); - const redirectUri = `http://127.0.0.1:${port}/oauth/callback`; - const state = crypto.randomBytes(24).toString("base64url"); - const verifier = crypto.randomBytes(48).toString("base64url"); - const challenge = crypto.createHash("sha256").update(verifier).digest("base64url"); - const authorize = new URL(`${oauthWebUrl}/oauth/authorize`); - authorize.search = new URLSearchParams({ client_id: clientId, redirect_uri: redirectUri, response_type: "code", scope: "user:read", state, code_challenge: challenge, code_challenge_method: "S256" }).toString(); - const callback = await new Promise<{ code: string }>((resolve, reject) => { - const server = createServer((request, response) => { - const url = new URL(request.url || "/", `http://127.0.0.1:${port}`); - if (url.pathname !== "/oauth/callback") { response.writeHead(404); response.end("Not found"); return; } - if (url.searchParams.get("state") !== state) { response.writeHead(400); response.end("Invalid state"); reject(new Error("OAuth state validation failed")); server.close(); return; } - const error = url.searchParams.get("error"); - if (error) { response.writeHead(400, { "Content-Type": "text/html; charset=utf-8" }); response.end("

Login failed. You can close this page.

"); reject(new Error(url.searchParams.get("error_description") || error)); server.close(); return; } - const code = url.searchParams.get("code"); - if (!code) { response.writeHead(400); response.end("Missing code"); reject(new Error("OAuth callback missing code")); server.close(); return; } - response.writeHead(200, { "Content-Type": "text/html; charset=utf-8" }); response.end("

Login successful. You can close this page.

"); resolve({ code }); server.close(); - }); - server.on("error", (error) => reject(new Error(`无法监听 OAuth 回调端口 ${port}: ${error.message}`))); - server.listen(port, "127.0.0.1", () => { void shell.openExternal(authorize.toString()); }); - setTimeout(() => { server.close(); reject(new Error("OAuth 登录超时,请重试")); }, 5 * 60 * 1000).unref(); - }); - const ipAddress = await resolvePublicIpv4(); - const response = await fetch(`${oauthUrl}/api/oauth/token`, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ grant_type: "authorization_code", code: callback.code, client_id: clientId, redirect_uri: redirectUri, code_verifier: verifier, device_uuid: crypto.randomUUID(), ip_address: ipAddress }) }); - const payload = await response.json().catch(() => ({})) as { access_token?: string; error_description?: string }; - if (!response.ok || !payload.access_token) throw new Error(payload.error_description || "SECTL OAuth 换取令牌失败"); - const relayResponse = await fetch(`${relayUrl}/auth/oauth`, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ access_token: payload.access_token, client_id: clientId, platform_id: process.env.SECTL_OFFICIAL_PLATFORM_ID || clientId }) }); - const relayPayload = await relayResponse.json().catch(() => ({})) as { access_token?: string; user?: { id?: string; email?: string; name?: string }; detail?: string }; - if (!relayResponse.ok || !relayPayload.access_token) throw new Error(relayPayload.detail || "官方服务 OAuth 登录失败"); - return { accessToken: relayPayload.access_token, userId: relayPayload.user?.id, email: relayPayload.user?.email, name: relayPayload.user?.name }; + return runSectlOAuthFlow(); } ipcMain.handle("sectl:oauth-login", () => runSectlOAuthLogin()); ipcMain.handle("official:oauth-login", async () => { loadConfig(DEFAULT_WORKSPACE); const relayUrl = (process.env.SECTL_OFFICIAL_API_URL || "").replace(/\/$/, ""); - const oauthUrl = (process.env.SECTL_OAUTH_API_URL || "https://appwrite.sectl.cn").replace(/\/$/, ""); - const oauthWebUrl = (process.env.SECTL_OAUTH_WEB_URL || "https://sectl.cn").replace(/\/$/, ""); - const clientId = process.env.SECTL_OFFICIAL_CLIENT_ID || ""; - const port = Number(process.env.SECTL_OAUTH_CALLBACK_PORT || 49152); - if (!relayUrl) throw new Error("请在 SecAgent 代码目录 .env 配置 SECTL_OFFICIAL_API_URL"); - if (!clientId) throw new Error("请在 SecAgent 代码目录 .env 配置 SECTL_OFFICIAL_CLIENT_ID"); - if (!Number.isInteger(port) || port < 49152 || port > 65535) throw new Error("SECTL_OAUTH_CALLBACK_PORT 必须是 49152-65535 的固定端口"); - const redirectUri = `http://127.0.0.1:${port}/oauth/callback`; - const state = crypto.randomBytes(24).toString("base64url"); - const verifier = crypto.randomBytes(48).toString("base64url"); - const challenge = crypto.createHash("sha256").update(verifier).digest("base64url"); - const authorize = new URL(`${oauthWebUrl}/oauth/authorize`); - authorize.search = new URLSearchParams({ client_id: clientId, redirect_uri: redirectUri, response_type: "code", scope: "user:read", state, code_challenge: challenge, code_challenge_method: "S256" }).toString(); - const callback = await new Promise<{ code: string }>((resolve, reject) => { - const server = createServer((request, response) => { - const url = new URL(request.url || "/", `http://127.0.0.1:${port}`); - if (url.pathname !== "/oauth/callback") { response.writeHead(404); response.end("Not found"); return; } - if (url.searchParams.get("state") !== state) { response.writeHead(400); response.end("Invalid state"); reject(new Error("OAuth state 校验失败")); server.close(); return; } - const error = url.searchParams.get("error"); - if (error) { response.writeHead(400, { "Content-Type": "text/html; charset=utf-8" }); response.end("

登录未完成,请返回 SecAgent 重试。

"); reject(new Error(url.searchParams.get("error_description") || error)); server.close(); return; } - const code = url.searchParams.get("code"); - if (!code) { response.writeHead(400); response.end("Missing code"); reject(new Error("OAuth 回调缺少 code")); server.close(); return; } - response.writeHead(200, { "Content-Type": "text/html; charset=utf-8" }); response.end("

SecAgent 登录成功,可以关闭此页面。

"); resolve({ code }); server.close(); - }); - server.on("error", (error) => reject(new Error(`无法监听 OAuth 回调端口 ${port}: ${error.message}`))); - server.listen(port, "127.0.0.1", () => { void shell.openExternal(authorize.toString()); }); - setTimeout(() => { server.close(); reject(new Error("OAuth 登录超时,请重试")); }, 5 * 60 * 1000).unref(); - }); - const ipAddress = await resolvePublicIpv4(); - const tokenResponse = await fetch(`${oauthUrl}/api/oauth/token`, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ grant_type: "authorization_code", code: callback.code, client_id: clientId, redirect_uri: redirectUri, code_verifier: verifier, device_uuid: crypto.randomUUID(), ip_address: ipAddress }) }); - const tokenPayload = await tokenResponse.json().catch(() => ({})) as { access_token?: string; error_description?: string }; - if (!tokenResponse.ok || !tokenPayload.access_token) throw new Error(tokenPayload.error_description || "SECTL OAuth 换取令牌失败"); - const relayResponse = await fetch(`${relayUrl}/auth/oauth`, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ access_token: tokenPayload.access_token, client_id: clientId, platform_id: process.env.SECTL_OFFICIAL_PLATFORM_ID || clientId }) }); - const relayPayload = await relayResponse.json().catch(() => ({})) as { access_token?: string; user?: { id?: string; email?: string; name?: string }; detail?: string }; - if (!relayResponse.ok || !relayPayload.access_token) throw new Error(relayPayload.detail || "官方服务登录失败"); - writeWorkspaceEnv(DEFAULT_WORKSPACE, "SECTL_OFFICIAL_TOKEN", relayPayload.access_token); + const result = await runSectlOAuthFlow(); + writeWorkspaceEnv(DEFAULT_WORKSPACE, "SECTL_OFFICIAL_TOKEN", result.accessToken); writeWorkspaceEnv(DEFAULT_WORKSPACE, "SECTL_OFFICIAL_SECTL_TOKEN", ""); - writeWorkspaceEnv(DEFAULT_WORKSPACE, "SECTL_OFFICIAL_USER_ID", relayPayload.user?.id || ""); - writeWorkspaceEnv(DEFAULT_WORKSPACE, "SECTL_OFFICIAL_EMAIL", relayPayload.user?.email || "SECTL 用户"); + writeWorkspaceEnv(DEFAULT_WORKSPACE, "SECTL_OFFICIAL_USER_ID", result.userId || ""); + writeWorkspaceEnv(DEFAULT_WORKSPACE, "SECTL_OFFICIAL_EMAIL", result.email || "SECTL 用户"); const current = readSettings(DEFAULT_WORKSPACE); const providers = current.providers.some((provider) => provider.id === "sectl-official") ? current.providers : [...current.providers, officialProvider(relayUrl)]; return saveSettings(DEFAULT_WORKSPACE, { ...current, providers }); @@ -1327,6 +1256,49 @@ ipcMain.handle("oobe:complete", (event) => { if (windowRef && !windowRef.isDestroyed()) { windowRef.show(); windowRef.focus(); } return { ok: true }; }); +const pendingToolConfirmations = new Map void; timer: NodeJS.Timeout }>(); +let toolConfirmationSeq = 0; + +/** Pause the agent until the user approves a sensitive tool call (Codex-style). */ +function confirmSensitiveToolCall(sessionId: string, confirmation: { tool: string; arguments: Record; reason: string }): Promise { + return new Promise((resolve) => { + const confirmationId = `tool-confirm-${++toolConfirmationSeq}`; + const timer = setTimeout(() => { + pendingToolConfirmations.delete(confirmationId); + logMain("tool.confirm.timeout", { confirmationId }); + resolve(false); + }, 5 * 60_000); + pendingToolConfirmations.set(confirmationId, { resolve, timer }); + logMain("tool.confirm.request", { confirmationId, sessionId, tool: confirmation.tool }); + sendToAppWindows("runtime:tool-confirmation", { confirmationId, sessionId, ...confirmation }); + }); +} + +ipcMain.handle("runtime:tool-confirmation-reply", (_event, payload: { confirmationId: string; approved: boolean; always?: boolean; signature?: string }) => { + const pending = pendingToolConfirmations.get(payload.confirmationId); + if (!pending) return { ok: false, error: "确认请求已过期" }; + pendingToolConfirmations.delete(payload.confirmationId); + clearTimeout(pending.timer); + if (payload.approved && payload.always && payload.signature) appendGuardApproval(payload.signature); + logMain("tool.confirm.reply", { confirmationId: payload.confirmationId, approved: payload.approved, always: Boolean(payload.always) }); + pending.resolve(payload.approved); + return { ok: true }; +}); + +/** Persist a "不再提示" approval straight into the yaml without a full settings rewrite. */ +function appendGuardApproval(signature: string): void { + try { + const file = configPath(DEFAULT_WORKSPACE); + const raw = YAML.parse(fs.readFileSync(file, "utf8")) as { guard?: { approved?: string[] } }; + const approved = new Set(raw?.guard?.approved || []); + approved.add(signature); + raw.guard = { ...(raw.guard || {}), approved: [...approved] }; + fs.writeFileSync(file, YAML.stringify(raw), "utf8"); + } catch (error) { + logMain("tool.confirm.persist.failed", { error: error instanceof Error ? error.message : String(error) }); + } +} + ipcMain.handle("settings:save", (_event, payload: SettingsPayload) => { const customModelMode = Boolean(payload?.customModelMode); let providers = Array.isArray(payload?.providers) ? payload.providers : []; @@ -1371,6 +1343,9 @@ ipcMain.handle("settings:save", (_event, payload: SettingsPayload) => { sentryTelemetryEnabled = saved.telemetry.enabled; initializeSentry(); telemetry?.setEnabled(saved.telemetry.enabled); + // Apply the new speech-recognition preference (provider chain) immediately. + configureSpeech(saved.speech); + configureTts(saved.tts); sendToAppWindows("settings:changed", saved); updateManager?.setPreferences(saved.updates); closeVoiceWakeWindow(); @@ -1391,27 +1366,32 @@ ipcMain.on("wake:context", (_event, payload: unknown) => { }); ipcMain.handle("wake:close", () => { closeWakeWindow(); return { ok: true }; }); ipcMain.on("wake:interactive", (_event, interactive: boolean) => { if (wakeWindow && !wakeWindow.isDestroyed()) wakeWindow.setIgnoreMouseEvents(!interactive); }); -ipcMain.handle("speech:start", (event) => { +ipcMain.handle("speech:start", async (event) => { const target = wakeWindow?.webContents.id === event.sender.id ? wakeWindow : windowRef; logMain("speech.start", { window: target === wakeWindow ? "wake" : "main" }); try { - const result = startSpeech(target); - logMain("speech.start.ready", { window: target === wakeWindow ? "wake" : "main", remote: result.remote }); + const result = await startSpeech(target); + logMain("speech.start.ready", { window: target === wakeWindow ? "wake" : "main", provider: result.provider, fallbacks: result.fallbacks }); return result; + } catch (error) { + recordTelemetryFailure({ type: "speech.failed", error, context: { phase: "start" } }); + throw error; } - catch (error) { recordTelemetryFailure({ type: "speech.failed", error, context: { phase: "start" } }); throw error; } }); -ipcMain.handle("speech:stop", () => { logMain("speech.stop"); stopSpeech(); return { ok: true }; }); +ipcMain.handle("speech:stop", () => { logMain("speech.stop"); void stopSpeech(); return { ok: true }; }); ipcMain.handle("speech:cancel", () => { logMain("speech.cancel"); cancelSpeech(); return { ok: true }; }); +ipcMain.handle("speech:chain", () => speechChain()); +ipcMain.handle("speech:test", (_event, kind: unknown) => { + const scope = kind === "official" || kind === "openai" || kind === "local" || kind === "bailian" || kind === "bailian-ws" ? kind : "auto"; + return testSpeech(scope); +}); ipcMain.handle("voice-wake:start", (event, phrase: string) => { - try { - return startVoiceWake(voiceWakeWindow?.webContents.id === event.sender.id ? voiceWakeWindow : undefined, phrase, () => { - // Keep the hidden microphone window alive so the listener can be resumed - // after the one-shot wake overlay closes. - stopVoiceWake(); - void openWakeWindow().catch((error) => logMain("wake.open.failed", { error: String(error), reason: "voice" })); - }); - } catch (error) { recordTelemetryFailure({ type: "speech.failed", error, context: { phase: "voice-wake-start" } }); throw error; } + return startVoiceWake(voiceWakeWindow?.webContents.id === event.sender.id ? voiceWakeWindow : undefined, phrase, () => { + // Keep the hidden microphone window alive so the listener can be resumed + // after the one-shot wake overlay closes. + stopVoiceWake(); + void openWakeWindow().catch((error) => logMain("wake.open.failed", { error: String(error), reason: "voice" })); + }).catch((error) => { recordTelemetryFailure({ type: "speech.failed", error, context: { phase: "voice-wake-start" } }); throw error; }); }); ipcMain.handle("voice-wake:stop", () => { stopVoiceWake(); return { ok: true }; }); ipcMain.on("voice-wake:log", (_event, payload: unknown) => { @@ -1441,6 +1421,19 @@ ipcMain.handle("tts:synthesize", async (_event, text: string) => { } }); ipcMain.on("wake:tts-log", (_event, payload: unknown) => logMain("wake.tts.playback", payload)); +// TTS diagnostics: connectivity probe, active fallback chain, installed SAPI voices. +ipcMain.handle("tts:test", (_event, kind?: string) => testTts(kind as never)); +ipcMain.handle("tts:chain", () => ttsChain()); +ipcMain.handle("tts:voices", () => listWindowsVoices()); +// Fetch an OpenAI-compatible provider's model catalogue (GET {base}/models), +// e.g. https://api.xiaomimimo.com/v1/models — feeds the settings dropdowns. +ipcMain.handle("models:fetch", async (_event, request: { baseUrl?: string; apiKey?: string; apiKeyEnv?: string }) => { + const apiKey = (request.apiKey && request.apiKey.trim()) || (request.apiKeyEnv ? process.env[request.apiKeyEnv] || "" : ""); + if (!request.baseUrl?.trim()) return { ok: false, message: "请填写 Base URL(例如 https://api.xiaomimimo.com/v1)", models: [] }; + if (!apiKey) return { ok: false, message: "缺少 API Key(先保存到工作区 .env 或在输入框填写)", models: [] }; + const { fetchProviderModels } = await import("../models/fetch-models.js"); + return fetchProviderModels({ baseUrl: request.baseUrl, apiKey, timeoutMs: 15_000 }); +}); ipcMain.on("speech:audio", (_event, samples: Float32Array) => sendSpeechAudio(samples)); ipcMain.on("voice-wake:audio", (_event, samples: Float32Array) => sendVoiceWakeAudio(samples)); ipcMain.handle("sessions:stop", (_event, id: string) => { @@ -1529,7 +1522,9 @@ ipcMain.handle("sessions:send", async (_event, id: string, text: string, modelId const runtimeConfig = isWakeRequest ? { ...config, agent: { ...config.agent, systemPrompt: `${config.agent.systemPrompt}\n\n## 快速唤起输出协议\n${QUICK_WAKE_OUTPUT_PROMPT}` } } : config; - runtime = new SecAgentRuntime(runtimeConfig, audit, skills, trace, pluginManager); + runtime = new SecAgentRuntime(runtimeConfig, audit, skills, trace, pluginManager, { confirmToolCall: (confirmation) => confirmSensitiveToolCall(id, confirmation) }); + const visionConfig = resolveVisionAgentConfig(config); + if (visionConfig) logMain("session.vision-model", { model: visionConfig.agent.model, provider: visionConfig.agent.provider, baseUrl: visionConfig.agent.baseUrl }); const previousReadSkillNames = before.messages.flatMap((message) => message.toolCalls || []).filter((call) => call.name === "secagent__read_skill" || call.name === "read_skill").map((call) => typeof (call.arguments as { name?: unknown })?.name === "string" ? (call.arguments as { name: string }).name : ""); const result = await runtime.run(historyInput(before, text), selectedReasoningEffort, conversationInput(before, text, attachments), abortController.signal, { previousAutoLoadedSkills: before.autoLoadedSkills, previousReadSkillNames, preRule }); if (result.autoLoadedSkills?.length) { @@ -1538,7 +1533,12 @@ ipcMain.handle("sessions:send", async (_event, id: string, text: string, modelId // Reuse the store's normal persistence path without adding a visible message. sessionStore.setAutoLoadedSkills(id, current.autoLoadedSkills); } - sessionStore.appendMessage(id, "assistant", result.message, toolCalls, activities); + // Surface resilience + hallucination findings inline so they survive the + // session history and stay visible without extra UI plumbing. + let finalMessage = result.message; + if ("usedFallbackModels" in result && result.usedFallbackModels?.length) finalMessage += `\n\n> ⚙️ 模型稳定性:已自动切换备用模型(${result.usedFallbackModels.join(" → ")}),原模型暂时不可用。`; + if ("hallucination" in result && result.hallucination?.signals.length) finalMessage += `\n\n> ⚠️ 幻觉检测提醒(仅提示,不代表一定有错):\n${result.hallucination.signals.map((signal) => `> - ${signal.detail}`).join("\n")}\n> 建议人工核对以上要点。`; + sessionStore.appendMessage(id, "assistant", finalMessage, toolCalls, activities); const title = await titlePromise; if (title) sessionStore.setTitle(id, title); trace({ stage: "assistant.response", data: { text: result.message } }); @@ -1622,6 +1622,9 @@ app.whenReady().then(async () => { async function startApplication(): Promise { const needsOnboarding = !fs.existsSync(configPath(DEFAULT_WORKSPACE)) || !isOnboardingComplete(DEFAULT_WORKSPACE); const initialSettings = readSettings(DEFAULT_WORKSPACE); + // Apply the persisted ASR preference before any speech session can start. + configureSpeech(initialSettings.speech); + configureTts(initialSettings.tts); pluginManager = new PluginManager(DEFAULT_WORKSPACE, { getSession: async () => { loadConfig(DEFAULT_WORKSPACE); diff --git a/src/electron/oauth.ts b/src/electron/oauth.ts new file mode 100644 index 0000000..954aa2b --- /dev/null +++ b/src/electron/oauth.ts @@ -0,0 +1,85 @@ +/** + * Shared SECTL OAuth (PKCE) login flow. + * + * Extracted from the two near-identical copies that used to live in + * `main.ts` (the plugin-manager login and the settings-page login). Both + * callers now run the same flow; the settings page additionally persists the + * returned tokens to the workspace. + */ +import { createServer } from "node:http"; +import { isIPv4 } from "node:net"; +import crypto from "node:crypto"; +import { shell } from "electron"; + +const PUBLIC_IP_ENDPOINTS = [ + "https://api.ipify.org?format=json", + "https://httpbin.org/ip", + "https://api64.ipify.org?format=json" +]; + +export interface SectlOAuthResult { + accessToken: string; + userId?: string; + email?: string; + name?: string; +} + +async function resolvePublicIpv4(): Promise { + for (const endpoint of PUBLIC_IP_ENDPOINTS) { + try { + const response = await fetch(endpoint, { signal: AbortSignal.timeout(5_000) }); + if (!response.ok) continue; + const payload = await response.json().catch(() => ({})) as { ip?: unknown; origin?: unknown }; + const candidate = String(payload.ip ?? payload.origin ?? "").split(",")[0].trim(); + if (isIPv4(candidate)) return candidate; + } catch { + // Try the next public-IP provider. + } + } + throw new Error("无法获取本机公网 IPv4,请检查网络连接后重试"); +} + +/** + * Run the full OAuth authorization-code + PKCE flow against the SECTL relay + * and return the relay session tokens. Rejects with a readable error when any + * step fails. + */ +export async function runSectlOAuthFlow(): Promise { + const relayUrl = (process.env.SECTL_OFFICIAL_API_URL || "").replace(/\/$/, ""); + const oauthUrl = (process.env.SECTL_OAUTH_API_URL || "https://appwrite.sectl.cn").replace(/\/$/, ""); + const oauthWebUrl = (process.env.SECTL_OAUTH_WEB_URL || "https://sectl.cn").replace(/\/$/, ""); + const clientId = process.env.SECTL_OFFICIAL_CLIENT_ID || ""; + const port = Number(process.env.SECTL_OAUTH_CALLBACK_PORT || 49152); + if (!relayUrl) throw new Error("请在 SecAgent .env 配置 SECTL_OFFICIAL_API_URL"); + if (!clientId) throw new Error("请在 SecAgent .env 配置 SECTL_OFFICIAL_CLIENT_ID"); + if (!Number.isInteger(port) || port < 49152 || port > 65535) throw new Error("SECTL_OAUTH_CALLBACK_PORT 必须是 49152-65535 的固定端口"); + const redirectUri = `http://127.0.0.1:${port}/oauth/callback`; + const state = crypto.randomBytes(24).toString("base64url"); + const verifier = crypto.randomBytes(48).toString("base64url"); + const challenge = crypto.createHash("sha256").update(verifier).digest("base64url"); + const authorize = new URL(`${oauthWebUrl}/oauth/authorize`); + authorize.search = new URLSearchParams({ client_id: clientId, redirect_uri: redirectUri, response_type: "code", scope: "user:read", state, code_challenge: challenge, code_challenge_method: "S256" }).toString(); + const callback = await new Promise<{ code: string }>((resolve, reject) => { + const server = createServer((request, response) => { + const url = new URL(request.url || "/", `http://127.0.0.1:${port}`); + if (url.pathname !== "/oauth/callback") { response.writeHead(404); response.end("Not found"); return; } + if (url.searchParams.get("state") !== state) { response.writeHead(400); response.end("Invalid state"); reject(new Error("OAuth state 校验失败")); server.close(); return; } + const error = url.searchParams.get("error"); + if (error) { response.writeHead(400, { "Content-Type": "text/html; charset=utf-8" }); response.end("

登录未完成,请返回 SecAgent 重试。

"); reject(new Error(url.searchParams.get("error_description") || error)); server.close(); return; } + const code = url.searchParams.get("code"); + if (!code) { response.writeHead(400); response.end("Missing code"); reject(new Error("OAuth 回调缺少 code")); server.close(); return; } + response.writeHead(200, { "Content-Type": "text/html; charset=utf-8" }); response.end("

SecAgent 登录成功,可以关闭此页面。

"); resolve({ code }); server.close(); + }); + server.on("error", (error) => reject(new Error(`无法监听 OAuth 回调端口 ${port}: ${error.message}`))); + server.listen(port, "127.0.0.1", () => { void shell.openExternal(authorize.toString()); }); + setTimeout(() => { server.close(); reject(new Error("OAuth 登录超时,请重试")); }, 5 * 60 * 1000).unref(); + }); + const ipAddress = await resolvePublicIpv4(); + const tokenResponse = await fetch(`${oauthUrl}/api/oauth/token`, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ grant_type: "authorization_code", code: callback.code, client_id: clientId, redirect_uri: redirectUri, code_verifier: verifier, device_uuid: crypto.randomUUID(), ip_address: ipAddress }) }); + const tokenPayload = await tokenResponse.json().catch(() => ({})) as { access_token?: string; error_description?: string }; + if (!tokenResponse.ok || !tokenPayload.access_token) throw new Error(tokenPayload.error_description || "SECTL OAuth 换取令牌失败"); + const relayResponse = await fetch(`${relayUrl}/auth/oauth`, { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ access_token: tokenPayload.access_token, client_id: clientId, platform_id: process.env.SECTL_OFFICIAL_PLATFORM_ID || clientId }) }); + const relayPayload = await relayResponse.json().catch(() => ({})) as { access_token?: string; user?: { id?: string; email?: string; name?: string }; detail?: string }; + if (!relayResponse.ok || !relayPayload.access_token) throw new Error(relayPayload.detail || "官方服务 OAuth 登录失败"); + return { accessToken: relayPayload.access_token, userId: relayPayload.user?.id, email: relayPayload.user?.email, name: relayPayload.user?.name }; +} diff --git a/src/electron/preload.ts b/src/electron/preload.ts index 973acdf..f0b87a8 100644 --- a/src/electron/preload.ts +++ b/src/electron/preload.ts @@ -81,11 +81,17 @@ contextBridge.exposeInMainWorld("secagent", { sendSpeechAudio: (samples: Float32Array) => ipcRenderer.send("speech:audio", samples), stopSpeech: () => ipcRenderer.invoke("speech:stop"), cancelSpeech: () => ipcRenderer.invoke("speech:cancel"), + testSpeech: (kind?: string) => ipcRenderer.invoke("speech:test", kind), + speechChain: () => ipcRenderer.invoke("speech:chain"), startVoiceWake: (phrase: string) => ipcRenderer.invoke("voice-wake:start", phrase), sendVoiceWakeAudio: (samples: Float32Array) => ipcRenderer.send("voice-wake:audio", samples), stopVoiceWake: () => ipcRenderer.invoke("voice-wake:stop"), logVoiceWake: (event: unknown) => ipcRenderer.send("voice-wake:log", event), synthesizeSpeech: (text: string) => ipcRenderer.invoke("tts:synthesize", text), + testTts: (kind?: string) => ipcRenderer.invoke("tts:test", kind), + ttsChain: () => ipcRenderer.invoke("tts:chain"), + listWindowsVoices: () => ipcRenderer.invoke("tts:voices"), + fetchRemoteModels: (request: { baseUrl: string; apiKey?: string; apiKeyEnv?: string }) => ipcRenderer.invoke("models:fetch", request), logWakeTts: (event: unknown) => ipcRenderer.send("wake:tts-log", event), logSpeech: (event: unknown) => ipcRenderer.send("speech:log", event), setWakeContext: (context: unknown) => ipcRenderer.send("wake:context", context), @@ -120,5 +126,11 @@ contextBridge.exposeInMainWorld("secagent", { const wrapped = (_event: Electron.IpcRendererEvent, payload: unknown) => listener(payload); ipcRenderer.on("plugins:changed", wrapped); return () => ipcRenderer.removeListener("plugins:changed", wrapped); + }, + respondToolConfirmation: (payload: { confirmationId: string; approved: boolean; always?: boolean; signature?: string }) => ipcRenderer.invoke("runtime:tool-confirmation-reply", payload), + onToolConfirmation: (listener: (payload: unknown) => void) => { + const wrapped = (_event: Electron.IpcRendererEvent, payload: unknown) => listener(payload); + ipcRenderer.on("runtime:tool-confirmation", wrapped); + return () => ipcRenderer.removeListener("runtime:tool-confirmation", wrapped); } }); diff --git a/src/electron/speech.ts b/src/electron/speech.ts index ee508d0..2c41a90 100644 --- a/src/electron/speech.ts +++ b/src/electron/speech.ts @@ -1,328 +1,141 @@ -import fs from "node:fs"; -import path from "node:path"; -import { app, BrowserWindow } from "electron"; -import { createKws, createOnlineRecognizer } from "sherpa-onnx"; -import { pinyin } from "pinyin-pro"; - -const modelName = "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23"; -/** Local offline ASR via the bundled WASM build; the model stays loaded for the app lifetime. */ -let localRecognizer: ReturnType | undefined; -let localStream: ReturnType["createStream"]> | undefined; -let remoteSocket: WebSocket | undefined; -let speechWindow: BrowserWindow | undefined; -let pendingRemoteAudio: ArrayBuffer[] = []; -/** "remote" = backend relay /asr/ws; "local" = bundled sherpa-onnx WASM recognizer; "idle" = none. */ -let mode: "remote" | "local" | "idle" = "idle"; -let voiceWakeKws: ReturnType | undefined; -let voiceWakeStream: ReturnType["createStream"]> | undefined; -let voiceWakePhrase = ""; -let voiceWakeDetected: (() => void) | undefined; -let voiceWakeStartedAt = 0; -let voiceWakeAudioFrames = 0; -let voiceWakeAwaitingFirstAudio = false; -let voiceWakeLastHeartbeatAt = 0; -let voiceWakeLastInactiveAudioLogAt = 0; -let remoteSocketSequence = 0; - -function projectPath(...parts: string[]): string { - const candidates = [ - path.join(process.resourcesPath, ...parts), - path.join(process.cwd(), ...parts), - path.join(app.getAppPath(), ...parts), - path.join(__dirname, "../../", ...parts) - ]; - const found = candidates.find((candidate) => fs.existsSync(candidate)); - if (!found) throw new Error(`找不到语音资源:${parts.join("/")}`); - return found; +/** + * Electron glue for speech recognition and voice wake. + * + * All recognition logic lives in the provider-agnostic `src/asr/` layer; this + * module owns the process-wide `AsrManager`, routes events to the requesting + * window over the `speech:event` channel, and keeps the exported surface that + * `main.ts` consumes. + */ +import { BrowserWindow, app } from "electron"; +import { AsrManager, type StartedAsr } from "../asr/manager.js"; +import type { AsrEvent } from "../asr/types.js"; +import type { AsrProviderKind } from "../asr/settings.js"; +import { LocalSherpaAsrProvider } from "../asr/local-sherpa.js"; +import { RelayAsrProvider } from "../asr/relay.js"; +import { OpenAiHttpAsrProvider } from "../asr/openai-http.js"; +import { BailianHttpAsrProvider } from "../asr/bailian-http.js"; +import { BailianWsAsrProvider } from "../asr/bailian-ws.js"; +import { MimoHttpAsrProvider } from "../asr/mimo-http.js"; +import { LocalSenseVoiceProvider } from "../asr/local-sensevoice.js"; +import { VoiceWakeEngine } from "../asr/voice-wake.js"; +import type { SpeechAsrSettings } from "../asr/settings.js"; + +let speechSettings: SpeechAsrSettings | undefined; + +/** Update the live ASR preference (called after settings load/save). */ +export function configureSpeech(settings: SpeechAsrSettings | undefined): void { + speechSettings = settings; } -function send(window: BrowserWindow | undefined, payload: unknown): void { +function send(window: BrowserWindow | undefined, event: AsrEvent): void { if (!window || window.isDestroyed() || window.webContents.isDestroyed()) return; - try { window.webContents.send("speech:event", payload); } catch { /* Window may close during an async callback. */ } + try { window.webContents.send("speech:event", event); } catch { /* Window may close during an async callback. */ } } -/** Remote ASR endpoint on the official relay. */ -function remoteAsrUrl(): string | null { - const token = process.env.SECTL_OFFICIAL_TOKEN || ""; - const baseUrl = (process.env.SECTL_OFFICIAL_API_URL || "").replace(/\/$/, ""); - if (!token || !baseUrl) return null; - const wsBase = baseUrl.replace(/^https:/, "wss:").replace(/^http:/, "ws:"); - return `${wsBase}/asr/ws?token=${encodeURIComponent(token)}`; -} +const appModelRoots = (): string[] => { + // Packaged layout keeps models/ next to the asar archive; dev runs from the + // repository root. Also cover out/ bundles two levels deep. + const roots: string[] = []; + try { roots.push(app.getAppPath()); } catch { /* not under Electron */ } + if (typeof process.resourcesPath === "string") roots.push(process.resourcesPath); + roots.push(process.cwd(), __dirname); + return [...new Set(roots)]; +}; + +const log = (message: string): void => console.info(message); + +const manager = new AsrManager({ + getProviderKind: () => speechSettings?.provider, + // 完全自定义回退链(用户在设置里排好的顺序,例:A → C → 本地)。 + getCustomChain: () => speechSettings?.chain, + log +}); +manager.register(new RelayAsrProvider({ + getToken: () => process.env.SECTL_OFFICIAL_TOKEN || "", + getApiBaseUrl: () => process.env.SECTL_OFFICIAL_API_URL || "", + log +})); +manager.register(new OpenAiHttpAsrProvider({ + getSettings: () => speechSettings?.openai, + getApiKey: (envName) => process.env[envName] || "", + log +})); +manager.register(new BailianHttpAsrProvider({ + getSettings: () => speechSettings?.bailian, + getApiKey: (envName) => process.env[envName] || "", + log +})); +manager.register(new BailianWsAsrProvider({ + getSettings: () => speechSettings?.bailian, + getApiKey: (envName) => process.env[envName] || "", + getNoise: () => speechSettings?.noise, + log +})); +manager.register(new MimoHttpAsrProvider({ + getSettings: () => speechSettings?.mimo, + getApiKey: (envName) => process.env[envName] || "", + log +})); +manager.register(new LocalSenseVoiceProvider({ extraRoots: appModelRoots(), log })); +manager.register(new LocalSherpaAsrProvider({ extraRoots: appModelRoots(), log })); + +const voiceWake = new VoiceWakeEngine({ extraRoots: appModelRoots(), log }); -function remoteAsrLogTarget(url: string): string { - try { - const parsed = new URL(url); - return `${parsed.protocol}//${parsed.host}${parsed.pathname}`; - } catch { - return ""; - } -} +let speechWindow: BrowserWindow | undefined; +let currentSession: StartedAsr | undefined; -/** Load the bundled streaming transducer once and give each utterance a fresh stream. */ -function startLocalRecognizer(): void { - if (localRecognizer) { - // A stream that has seen inputFinished() cannot accept more audio. - localStream?.free(); - localStream = localRecognizer.createStream(); - return; +/** Start a recognition utterance for `window`; falls back down the provider chain. */ +export async function startSpeech(window: BrowserWindow | undefined): Promise<{ ok: true; provider: string; remote: boolean; fallbacks: string[] }> { + speechWindow = window; + if (currentSession) { + // The chat window and the wake overlay share one utterance; re-route events. + send(speechWindow, { type: "ready", provider: currentSession.providerId }); + return { ok: true, provider: currentSession.providerId, remote: currentSession.providerId !== "local", fallbacks: [] }; } - const model = projectPath("models", modelName); - localRecognizer = createOnlineRecognizer({ - featConfig: { sampleRate: 16000, featureDim: 80 }, - modelConfig: { - transducer: { - encoder: path.join(model, "encoder-epoch-99-avg-1.int8.onnx"), - decoder: path.join(model, "decoder-epoch-99-avg-1.onnx"), - joiner: path.join(model, "joiner-epoch-99-avg-1.int8.onnx") - }, - tokens: path.join(model, "tokens.txt"), provider: "cpu", numThreads: 1 - }, - enableEndpoint: 1, - rule1MinTrailingSilence: 2.4, - rule2MinTrailingSilence: 1.2, - rule3MinUtteranceLength: 20 - }); - localStream = localRecognizer.createStream(); - console.info(`[speech] local WASM recognizer ready model=${modelName}`); + currentSession = await manager.start((event) => send(speechWindow, event)); + return { ok: true, provider: currentSession.providerId, remote: currentSession.providerId !== "local", fallbacks: currentSession.fallbacks }; } -function keywordTokens(phrase: string): string { - const syllables = pinyin(phrase.replace(/\s+/g, ""), { toneType: "symbol", type: "array" }) as string[]; - const initials = ["zh", "ch", "sh", "b", "p", "m", "f", "d", "t", "n", "l", "g", "k", "h", "j", "q", "x", "r", "z", "c", "s", "y", "w"]; - return syllables.map((syllable) => { - const initial = initials.find((candidate) => syllable.startsWith(candidate)) || ""; - return `${initial} ${syllable.slice(initial.length)}`; - }).join(" "); +export function sendSpeechAudio(samples: Float32Array): void { + if (currentSession) currentSession.session.push(samples); } -export function startVoiceWake(window: BrowserWindow | undefined, phrase: string, onDetected: () => void): { ok: true } { - void window; - voiceWakePhrase = phrase.trim(); - voiceWakeDetected = onDetected; - if (voiceWakeKws) { - console.info(`[voice-wake] local KWS already active phrase=${voiceWakePhrase}`); - return { ok: true }; - } - const model = projectPath("models", "sherpa-onnx-kws-zipformer-zh-en-3M-2025-12-20"); - voiceWakeKws = createKws({ - featConfig: { samplingRate: 16000, featureDim: 80 }, - modelConfig: { - transducer: { - encoder: path.join(model, "encoder-epoch-13-avg-2-chunk-16-left-64.int8.onnx"), - decoder: path.join(model, "decoder-epoch-13-avg-2-chunk-16-left-64.onnx"), - joiner: path.join(model, "joiner-epoch-13-avg-2-chunk-16-left-64.int8.onnx") - }, - tokens: path.join(model, "tokens.txt"), provider: "cpu", numThreads: 1, modelingUnit: "ppinyin" - }, - maxActivePaths: 4, numTrailingBlanks: 1, keywordsScore: 1.5, keywordsThreshold: 0.55, - keywords: `${keywordTokens(voiceWakePhrase)} @${voiceWakePhrase}` - }); - voiceWakeStream = voiceWakeKws.createStream(); - voiceWakeStartedAt = Date.now(); - voiceWakeAudioFrames = 0; - voiceWakeAwaitingFirstAudio = true; - voiceWakeLastHeartbeatAt = voiceWakeStartedAt; - console.info(`[voice-wake] local KWS ready phrase=${voiceWakePhrase}`); - return { ok: true }; +export async function stopSpeech(): Promise { + const active = currentSession; + if (!active) return; + currentSession = undefined; + try { await active.session.stop(); } catch (error) { send(speechWindow, { type: "error", message: error instanceof Error ? error.message : String(error) }); } + speechWindow = undefined; } -export function sendVoiceWakeAudio(samples: Float32Array): void { - const kws = voiceWakeKws; - const stream = voiceWakeStream; - const now = Date.now(); - if (!kws || !stream) { - if (now - voiceWakeLastInactiveAudioLogAt >= 15000) { - voiceWakeLastInactiveAudioLogAt = now; - console.info(`[voice-wake] audio ignored kws=${Boolean(kws)} stream=${Boolean(stream)}`); - } - return; - } - voiceWakeAudioFrames += 1; - if (voiceWakeAwaitingFirstAudio) { - voiceWakeAwaitingFirstAudio = false; - console.info(`[voice-wake] local KWS received first audio elapsed=${now - voiceWakeStartedAt}ms`); - } else if (now - voiceWakeLastHeartbeatAt >= 15000) { - voiceWakeLastHeartbeatAt = now; - console.info(`[voice-wake] local KWS audio heartbeat frames=${voiceWakeAudioFrames} elapsed=${now - voiceWakeStartedAt}ms`); - } - stream.acceptWaveform(16000, samples); - while (kws.isReady(stream)) kws.decode(stream); - const result = kws.getResult(stream); - if (result.keyword) { - console.info(`[voice-wake] local KWS detected keyword=${result.keyword} frames=${voiceWakeAudioFrames}`); - // Reset before invoking the callback. The callback closes the hidden - // window and releases the KWS instance immediately. - kws.reset(stream); - const detected = voiceWakeDetected; - detected?.(); - } +export function cancelSpeech(): void { + const active = currentSession; + if (!active) return; + currentSession = undefined; + try { active.session.cancel(); } catch { /* best effort */ } + speechWindow = undefined; } -export function stopVoiceWake(): void { - if (voiceWakeKws || voiceWakeStream) console.info(`[voice-wake] local KWS stopping frames=${voiceWakeAudioFrames} activeFor=${voiceWakeStartedAt ? Date.now() - voiceWakeStartedAt : 0}ms`); - voiceWakeDetected = undefined; - voiceWakeStream = undefined; - voiceWakeKws?.free(); - voiceWakeKws = undefined; - voiceWakeStartedAt = 0; - voiceWakeAudioFrames = 0; - voiceWakeAwaitingFirstAudio = false; +/** Connectivity probe for the settings page; `kind` scopes which providers run. */ +export function testSpeech(kind: AsrProviderKind): ReturnType { + return manager.test(kind); } -export function startSpeech(window: BrowserWindow | undefined): { ok: true; remote: boolean } { - speechWindow = window; - if (mode === "local" && localRecognizer) return { ok: true, remote: false }; - if (remoteSocket && (remoteSocket.readyState === WebSocket.OPEN || remoteSocket.readyState === WebSocket.CONNECTING)) { - // The main chat window and the wake overlay share one ASR connection. If the - // other window started it first, route subsequent events to the latest caller. - if (remoteSocket.readyState === WebSocket.OPEN) send(speechWindow, { type: "ready" }); - return { ok: true, remote: true }; - } - - const url = remoteAsrUrl(); - if (url && typeof WebSocket !== "undefined") { - const connectionId = ++remoteSocketSequence; - const connectedAt = Date.now(); - console.info(`[speech] connecting id=${connectionId} to ${remoteAsrLogTarget(url)}`); - try { - const socket = new WebSocket(url); - remoteSocket = socket; - mode = "remote"; - socket.binaryType = "arraybuffer"; - socket.onopen = () => { - console.info(`[speech] ASR WebSocket opened id=${connectionId} elapsed=${Date.now() - connectedAt}ms`); - if (remoteSocket === socket && socket.readyState === WebSocket.OPEN) { - // The relay must receive the start control message before binary - // audio. Sending the buffered audio first makes the relay discard - // everything spoken while the WebSocket was connecting. - try { socket.send(JSON.stringify({ type: "start" })); } catch { /* Socket may close during startup. */ } - for (const pcm of pendingRemoteAudio) socket.send(pcm); - pendingRemoteAudio = []; - } - send(speechWindow, { type: "ready" }); - }; - socket.onmessage = (event) => { - try { send(speechWindow, typeof event.data === "string" ? JSON.parse(event.data) : event.data); } - catch { send(speechWindow, { type: "log", message: String(event.data ?? "") }); } - }; - socket.onerror = (event) => { - const errorEvent = event as ErrorEvent; - const error = errorEvent.error as { message?: string; code?: string } | undefined; - console.error("[speech] ASR WebSocket error", { - id: connectionId, - elapsed: Date.now() - connectedAt, - readyState: socket.readyState, - readyStateName: ["CONNECTING", "OPEN", "CLOSING", "CLOSED"][socket.readyState] || "UNKNOWN", - message: errorEvent.message || error?.message || "", - errorCode: error?.code || "", - eventType: errorEvent.type - }); - if (mode === "remote") send(speechWindow, { type: "error", message: "云端语音识别连接失败" }); - }; - socket.onclose = (event) => { - console.warn("[speech] ASR WebSocket closed", { - id: connectionId, - code: event.code, - reason: event.reason || "", - wasClean: event.wasClean, - elapsed: Date.now() - connectedAt, - hadCurrentSocket: remoteSocket === socket - }); - const wasRemote = mode === "remote"; - if (remoteSocket === socket) { remoteSocket = undefined; mode = "idle"; } - pendingRemoteAudio = []; - if (wasRemote) send(speechWindow, { type: "stopped" }); - }; - return { ok: true, remote: true }; - } catch { - remoteSocket = undefined; - mode = "idle"; - } - } - - mode = "local"; - try { - startLocalRecognizer(); - } catch (error) { - mode = "idle"; - send(speechWindow, { type: "error", message: error instanceof Error ? error.message : String(error) }); - return { ok: true, remote: false }; - } - send(speechWindow, { type: "ready" }); - return { ok: true, remote: false }; +/** Provider chain that `auto` would try right now, for diagnostics in the UI. */ +export function speechChain(): string[] { + return manager.chain(); } -export function sendSpeechAudio(samples: Float32Array): void { - const pcm = samples.buffer.slice(samples.byteOffset, samples.byteOffset + samples.byteLength) as ArrayBuffer; - if (remoteSocket?.readyState === WebSocket.OPEN) { - try { remoteSocket.send(pcm); } catch { /* Socket may close between the state check and send. */ } - return; - } - if (remoteSocket?.readyState === WebSocket.CONNECTING) { - if (pendingRemoteAudio.length >= 32) pendingRemoteAudio.shift(); - pendingRemoteAudio.push(pcm); - return; - } - const recognizer = localRecognizer; - const stream = localStream; - if (!recognizer || !stream) return; - try { - stream.acceptWaveform(16000, samples); - while (recognizer.isReady(stream)) recognizer.decode(stream); - const text = (recognizer.getResult(stream).text || "").trim(); - if (text) send(speechWindow, { type: "partial", text }); - if (recognizer.isEndpoint(stream)) { - // The endpoint fires with the segment's full text; reset starts a new segment. - recognizer.reset(stream); - if (text) send(speechWindow, { type: "final", text }); - } - } catch (error) { - send(speechWindow, { type: "error", message: error instanceof Error ? error.message : String(error) }); - } +export async function startVoiceWake(window: BrowserWindow | undefined, phrase: string, onDetected: () => void): Promise<{ ok: true }> { + void window; + await voiceWake.start(phrase.trim(), onDetected); + return { ok: true }; } -export function stopSpeech(): void { - if (remoteSocket) { - pendingRemoteAudio = []; - if (remoteSocket.readyState === WebSocket.OPEN) { - try { remoteSocket.send("Done"); } catch { /* Socket may already be closing. */ } - } else if (remoteSocket.readyState === WebSocket.CONNECTING) remoteSocket.close(); - return; - } - const recognizer = localRecognizer; - const stream = localStream; - if (!recognizer || !stream) return; - try { - // Flush the tail of the utterance so the last segment is not lost. - stream.inputFinished(); - while (recognizer.isReady(stream)) recognizer.decode(stream); - const text = (recognizer.getResult(stream).text || "").trim(); - if (text) send(speechWindow, { type: "final", text }); - } catch (error) { - send(speechWindow, { type: "error", message: error instanceof Error ? error.message : String(error) }); - } - // A stream that has seen inputFinished() cannot be reused. - stream.free(); - localStream = recognizer.createStream(); - mode = "idle"; - send(speechWindow, { type: "stopped" }); - speechWindow = undefined; +export function stopVoiceWake(): void { + voiceWake.stop(); } -/** Abort the current utterance without waiting for recognition or enhancement. */ -export function cancelSpeech(): void { - pendingRemoteAudio = []; - if (remoteSocket) { - const socket = remoteSocket; - remoteSocket = undefined; - mode = "idle"; - try { socket.close(1000, "cancelled"); } catch { /* Socket may already be closed. */ } - return; - } - if (localRecognizer) { - try { localStream?.free(); } catch { /* Stream may already be freed. */ } - localStream = undefined; - mode = "idle"; - } +export function sendVoiceWakeAudio(samples: Float32Array): void { + voiceWake.feed(samples); } diff --git a/src/electron/tts.ts b/src/electron/tts.ts index ee4e180..ccd5d61 100644 --- a/src/electron/tts.ts +++ b/src/electron/tts.ts @@ -1,27 +1,83 @@ -import { EdgeTTS } from "@andresaya/edge-tts"; -import { DEFAULT_TTS_RATE, DEFAULT_TTS_VOICE } from "../config.js"; +/** + * Electron glue for text-to-speech. + * + * The provider implementations and the fallback-chain manager live in + * `src/tts/`; this module keeps the process-wide TtsManager, feeds it live + * settings, and exposes the small surface `main.ts` consumes: + * - configureTts(settings) after every settings load/save + * - synthesizeSpeech(text) speak one utterance (tries chain in order) + * - testTts / ttsChain diagnostics for the settings page + * - listWindowsVoices() installed SAPI voices for the picker + */ +import { execFile } from "node:child_process"; +import { promisify } from "node:util"; +import { EdgeTtsProvider, WindowsSapiTtsProvider, MimoTtsProvider, BailianTtsProvider } from "../tts/providers.js"; +import { TtsManager } from "../tts/manager.js"; +import type { TtsProviderKind, TtsSettings } from "../tts/types.js"; +import { DEFAULT_TTS_VOICE, DEFAULT_TTS_RATE } from "../config.js"; -function escapeXml(value: string): string { - return value.replace(/&/g, "&").replace(//g, ">").replace(/"/g, """).replace(/'/g, "'"); +const execFileAsync = promisify(execFile); + +let ttsSettings: TtsSettings | undefined; +/** Cache keyed on the settings object identity — configureTts() swaps the object. */ +let managerCache: { ref: TtsSettings | undefined; manager: TtsManager } | undefined; + +const log = (message: string): void => console.info(message); + +/** Update the live TTS preference (called after settings load/save). */ +export function configureTts(settings: TtsSettings | undefined): void { + ttsSettings = settings; + managerCache = undefined; // rebuilt lazily with the latest settings } -/** Generate one short MP3 chunk so the renderer can start playback immediately. */ -export async function synthesizeSpeech(text: string, options: { voice?: string; rate?: string } = {}): Promise { +function ensureManager(): TtsManager { + if (managerCache && managerCache.ref === ttsSettings) return managerCache.manager; + const liveSettings = (): TtsSettings => ttsSettings || { provider: "edge", voice: DEFAULT_TTS_VOICE, rate: DEFAULT_TTS_RATE }; + const manager = new TtsManager({ getSettings: liveSettings, log }); + // Edge/Windows take no constructor options — voice/rate ride on synthesize options. + manager.register(new EdgeTtsProvider()); + manager.register(new WindowsSapiTtsProvider()); + manager.register(new MimoTtsProvider({ + getSettings: () => ttsSettings?.mimo, + getApiKey: (envName) => process.env[envName] || "", + log + })); + manager.register(new BailianTtsProvider({ + getSettings: () => ttsSettings?.bailian, + getApiKey: (envName) => process.env[envName] || "", + log + })); + managerCache = { ref: ttsSettings, manager }; + return manager; +} + +/** Speak `text` through the primary provider, falling back down the chain. */ +export async function synthesizeSpeech(text: string, settingsSnapshot?: TtsSettings): Promise { + if (settingsSnapshot) configureTts(settingsSnapshot); const clean = text.replace(/\s+/g, " ").trim(); if (!clean) return Buffer.alloc(0); - let lastError: unknown; - // Edge TTS occasionally resets the TLS socket before the WebSocket - // handshake completes. A fresh client per attempt avoids reusing that - // broken connection and makes short-lived network hiccups transparent. - for (let attempt = 0; attempt < 3; attempt += 1) { - try { - const client = new EdgeTTS(); - await client.synthesize(escapeXml(clean), options.voice || DEFAULT_TTS_VOICE, { rate: options.rate || DEFAULT_TTS_RATE, outputFormat: "audio-24khz-48kbitrate-mono-mp3" }); - return client.toBuffer(); - } catch (error) { - lastError = error; - if (attempt < 2) await new Promise((resolve) => setTimeout(resolve, 250 * 2 ** attempt)); - } + const chunk = await ensureManager().synthesize(clean, { voice: ttsSettings?.voice, rate: ttsSettings?.rate }); + return chunk.data; +} + +/** Connectivity probe for the settings page. */ +export function testTts(kind?: TtsProviderKind): ReturnType { + return ensureManager().test(kind); +} + +/** Provider chain that would be tried right now, for diagnostics. */ +export function ttsChain(): TtsProviderKind[] { + return ensureManager().chain(); +} + +/** Installed Windows SAPI voices ("Name|Culture" per line) for the picker. */ +export async function listWindowsVoices(): Promise { + if (process.platform !== "win32") return []; + try { + const script = "Add-Type -AssemblyName System.Speech; (New-Object System.Speech.Synthesis.SpeechSynthesizer).GetInstalledVoices() | ForEach-Object { $_.VoiceInfo.Name + '|' + $_.VoiceInfo.Culture }"; + const { stdout } = await execFileAsync("powershell.exe", ["-NoProfile", "-NonInteractive", "-Command", script], { timeout: 15_000 }); + return stdout.split(/\r?\n/).map((line) => line.trim()).filter(Boolean); + } catch { + return []; } - throw lastError instanceof Error ? lastError : new Error(String(lastError)); } diff --git a/src/hallucination.test.ts b/src/hallucination.test.ts new file mode 100644 index 0000000..a7f81ea --- /dev/null +++ b/src/hallucination.test.ts @@ -0,0 +1,30 @@ +import assert from "node:assert/strict"; +import test from "node:test"; +import { detectHallucination } from "./hallucination.js"; + +test("flags answers claiming success after failed tool calls", () => { + const report = detectHallucination("操作已经成功完成,积分已加 5 分。", { toolCalls: [{ name: "secagent__secscore", ok: false }], runCompleted: true }); + assert.ok(report.signals.some((signal) => signal.id === "claims_success_after_tool_failure")); +}); + +test("does not flag honest failure reports", () => { + const report = detectHallucination("很抱歉,积分添加失败:数据库不可用。", { toolCalls: [{ name: "secagent__secscore", ok: false }], runCompleted: true }); + assert.equal(report.signals.some((signal) => signal.id === "claims_success_after_tool_failure"), false); +}); + +test("flags repetition loops", () => { + const line = "这个答案就是不断重复同样的一句话没有新信息"; + const text = Array(40).fill(line).join("\n"); + const report = detectHallucination(text, { toolCalls: [], runCompleted: true }); + assert.ok(report.signals.some((signal) => signal.id === "repetition_loop")); +}); + +test("flags fabricated references to never-produced evidence", () => { + const report = detectHallucination("如上图所示,数据呈上升趋势。", { toolCalls: [], runCompleted: true }); + assert.ok(report.signals.some((signal) => signal.id === "fabricated_tool_reference")); +}); + +test("clean answers stay clean", () => { + const report = detectHallucination("根据你提供的配置文件,发现默认模型未设置。建议在设置页选择一个默认模型。", { toolCalls: [{ name: "secagent__read_file", ok: true }], runCompleted: true }); + assert.equal(report.score, 0); +}); diff --git a/src/hallucination.ts b/src/hallucination.ts new file mode 100644 index 0000000..5fff1a3 --- /dev/null +++ b/src/hallucination.ts @@ -0,0 +1,96 @@ +/** + * Lightweight hallucination signals for final answers. + * + * Heuristics only — nothing here blocks an answer. Findings are surfaced as a + * warning strip in the UI and as trace events, so the user can double-check + * claims the model makes after tools failed, or answers degenerated into + * repetition loops (a common failure mode of smaller models under load). + */ +export interface HallucinationSignal { + /** Stable id for UI i18n/lookup. */ + id: "claims_success_after_tool_failure" | "repetition_loop" | "contradicts_empty_tools" | "fabricated_tool_reference"; + detail: string; +} + +export interface HallucinationReport { + /** 0 = clean; each signal adds 1. Not a probability. */ + score: number; + signals: HallucinationSignal[]; +} + +export interface TurnEvidence { + /** Tool calls from this run with their outcome, in call order. */ + toolCalls: Array<{ name: string; ok: boolean }>; + /** Whether the run completed without the agent loop erroring out. */ + runCompleted: boolean; +} + +const REPETITION_WINDOW = 24; + +function normalizedLines(text: string): string[] { + return text.split(/\n+/).map((line) => line.trim().replace(/\s+/g, " ")).filter((line) => line.length >= 8); +} + +/** + * Detect n-gram repetition loops. Answers where a single chunk of ~8 words + * covers most of the text, repeated many times, are almost always generation + * loops rather than intentional emphasis. + */ +function detectRepetitionLoop(text: string): HallucinationSignal | undefined { + const lines = normalizedLines(text); + if (lines.length < REPETITION_WINDOW) return undefined; + const chunks = lines.map((line) => { + const words = line.split(" "); + return words.length <= 8 ? words.join(" ") : words.slice(0, 8).join(" "); + }); + const counts = new Map(); + for (const chunk of chunks) counts.set(chunk, (counts.get(chunk) || 0) + 1); + let maxCount = 0; + let topChunk = ""; + for (const [chunk, count] of counts) { + if (count > maxCount) { maxCount = count; topChunk = chunk; } + } + if (maxCount >= REPETITION_WINDOW && maxCount / chunks.length >= 0.4) { + return { id: "repetition_loop", detail: `回答疑似陷入循环重复(「${topChunk.slice(0, 40)}…」出现 ${maxCount} 次,占正文 ${(100 * maxCount / chunks.length).toFixed(0)}%)。` }; + } + return undefined; +} + +const SUCCESS_CLAIM_PATTERN = /(?:已(?:经)?(?:成功|完成|执行)|操作已(?:成功)?|successfully (?:completed|done)|done\.)/i; +const FAILURE_ACK_PATTERN = /(?:失败|未能|无法|没有成功|出错|error|failed)/i; +const TOOL_MENTION = /(?:工具|调用|secagent__|secscore|plugin)/i; + +/** The model asserts an operation succeeded while the very tools it called failed. */ +function detectSuccessClaimAfterFailure(evidence: TurnEvidence, text: string): HallucinationSignal | undefined { + const failures = evidence.toolCalls.filter((call) => !call.ok); + if (!failures.length) return undefined; + const head = text.slice(0, 600); + const claimsSuccess = SUCCESS_CLAIM_PATTERN.test(head) && !FAILURE_ACK_PATTERN.test(head); + const mentionsTool = TOOL_MENTION.test(text) || /(?:操作|积分|写入|保存|修改|创建)/.test(head); + if (claimsSuccess && mentionsTool) { + const names = [...new Set(failures.map((call) => call.name))].slice(0, 3).join("、"); + return { id: "claims_success_after_tool_failure", detail: `本轮工具 ${names} 实际执行失败,但回答开头声称操作已成功。请核实后再采信。` }; + } + return undefined; +} + +const FABRICATED_REFERENCE = /(?:如上(?:方|图|文)所示|从(?:上述|以上)(?:结果|截图|表格)可见|(?:as shown|as mentioned) (?:above|in the table))/i; + +/** Refers to evidence (tables/screenshots/results) that were never produced this turn. */ +function detectFabricatedReference(evidence: TurnEvidence, text: string): HallucinationSignal | undefined { + if (evidence.toolCalls.length > 0) return undefined; + const match = text.match(FABRICATED_REFERENCE); + if (match) return { id: "fabricated_tool_reference", detail: `回答引用了不存在的材料(「${match[0]}」),但本轮没有任何工具产生数据。` }; + return undefined; +} + +export function detectHallucination(finalText: string, evidence: TurnEvidence): HallucinationReport { + const signals: HallucinationSignal[] = []; + const repetition = detectRepetitionLoop(finalText); + if (repetition) signals.push(repetition); + const successClaim = detectSuccessClaimAfterFailure(evidence, finalText); + if (successClaim) signals.push(successClaim); + const fabricated = detectFabricatedReference(evidence, finalText); + if (fabricated) signals.push(fabricated); + return { score: signals.length, signals }; +} diff --git a/src/index.ts b/src/index.ts index 6e73be2..cab8137 100644 --- a/src/index.ts +++ b/src/index.ts @@ -2,7 +2,7 @@ import readline from "node:readline/promises"; import { stdin as input, stdout as output } from "node:process"; import { initializeWorkspace, loadConfig, normalizeAndValidate, useConfiguredModel } from "./config.js"; -import { DEFAULT_WORKSPACE, expandPath } from "./paths.js"; +import { DEFAULT_WORKSPACE, expandPath, migrateLegacyWorkspace } from "./paths.js"; import { loadEnabledSkills } from "./skills.js"; import { AuditStore } from "./audit.js"; import { SecAgentRuntime, type RunResult, type TraceEvent } from "./runtime.js"; @@ -160,7 +160,19 @@ async function openRuntime(workspace: string, modelId: string | undefined, trace const plugins = new PluginManager(workspace); await plugins.initialize(); const skills = [...loadEnabledSkills(config), ...plugins.getSkills()]; - return { runtime: new SecAgentRuntime(config, audit, skills, trace, plugins), audit, plugins, config }; + // Interactive sensitive-tool confirmation for the CLI: default-deny when + // stdin is not a TTY (piped/CI runs) so nothing dangerous executes unattended. + const confirmToolCall = async (confirmation: { tool: string; reason: string }): Promise => { + if (!process.stdin.isTTY) return false; + process.stdout.write(`\n⚠ 敏感操作确认(${confirmation.tool}):${confirmation.reason}\n允许执行?[y/N] `); + const reply = await new Promise((resolve) => { + const onData = (chunk: Buffer) => { process.stdin.removeListener("data", onData); resolve(chunk.toString("utf8")); }; + process.stdin.once("data", onData); + setTimeout(() => { process.stdin.removeListener("data", onData); resolve(""); }, 60_000).unref(); + }); + return /^y(es)?$/i.test(reply.trim()); + }; + return { runtime: new SecAgentRuntime(config, audit, skills, trace, plugins, { confirmToolCall }), audit, plugins, config }; } async function closeRuntime(handle: RuntimeHandle | undefined): Promise { @@ -296,6 +308,12 @@ async function main(): Promise { if (!command || ["-h", "--help", "help"].includes(command)) return void console.log(usage()); const options = parseOptions(args); const { workspace, positionals } = options; + // Move a legacy ~/SecAgentWorkspace into the platform data directory before + // any command touches the default workspace. + try { + const migratedTo = workspace === DEFAULT_WORKSPACE ? migrateLegacyWorkspace() : undefined; + if (migratedTo) console.log(`已将旧工作区迁移到 ${migratedTo}`); + } catch { /* 迁移失败不阻塞命令 */ } if (command === "init") { initializeWorkspace(workspace); diff --git a/src/model-provider.test.ts b/src/model-provider.test.ts index d9c59a2..43e4558 100644 --- a/src/model-provider.test.ts +++ b/src/model-provider.test.ts @@ -224,3 +224,35 @@ test("persisted tool calls and results are restored to the next OpenAI request", delete process.env.TEST_MODEL_KEY; } }); +test("a single-turn sub-agent encodes user attachments as OpenAI image parts", async () => { + const originalFetch = globalThis.fetch; + process.env.TEST_MODEL_KEY = "test-key"; + let requestBody: Record | undefined; + globalThis.fetch = async (_url, init) => { + requestBody = JSON.parse(String(init?.body || "{}")) as Record; + return response('data: {"choices":[{"delta":{"content":"red"}}]}\n\ndata: [DONE]\n\n'); + }; + try { + // Tool-less vision sub-agent: allowEmptyTools and no runtime prompts. + const agent = new ModelToolAgent(config(), [], undefined, undefined, false, true); + const conversation: ConversationMessage[] = [ + { + role: "user", + content: "图中是什么颜色", + attachments: [{ id: "a1", name: "test.png", mimeType: "image/png", dataUrl: "data:image/png;base64,iVBORw0KGgo=", size: 6 }] + } + ]; + const result = await agent.run("图中是什么颜色", [], async () => { + throw new Error("识图模型不应调用工具"); + }, "low", conversation); + assert.equal(result, "red"); + const messages = requestBody?.messages as Array>; + const content = messages[1]?.content as Array>; + assert.ok(Array.isArray(content)); + assert.deepEqual(content[0], { type: "text", text: "图中是什么颜色" }); + assert.deepEqual(content[1], { type: "image_url", image_url: { url: "data:image/png;base64,iVBORw0KGgo=" } }); + } finally { + globalThis.fetch = originalFetch; + delete process.env.TEST_MODEL_KEY; + } +}); \ No newline at end of file diff --git a/src/model-provider.ts b/src/model-provider.ts index d2eec5e..3b1dea5 100644 --- a/src/model-provider.ts +++ b/src/model-provider.ts @@ -451,7 +451,7 @@ export class ModelToolAgent { } if (type === "response.failed") { const failed = event.response as { error?: { message?: string; code?: string } } | undefined; - throw new Error(failed?.error?.message || "妯″瀷璇锋眰澶辫触"); + throw new Error(failed?.error?.message || "模型请求失败"); } }, () => ({ output: [{ type: "message", content: answer || undefined }, ...[...calls.values()].map((call) => ({ type: "function_call", call_id: call.callId, name: call.name, arguments: call.arguments }))] }), signal); const functionCalls = [...calls.values()].filter((call) => call.name && call.callId); diff --git a/src/models/fetch-models.ts b/src/models/fetch-models.ts new file mode 100644 index 0000000..015a77d --- /dev/null +++ b/src/models/fetch-models.ts @@ -0,0 +1,59 @@ +/** + * Fetch the model catalogue of an OpenAI-compatible endpoint — + * `GET {baseUrl}/models` — so the settings UI can offer live lists instead of + * hand-typed model ids. + * + * Verified sources: + * MiMo: GET https://api.xiaomimimo.com/v1/models → { object:"list", data:[{id,…}] } + * (auth: `api-key: $MIMO_API_KEY` or `Authorization: Bearer`) + * OpenAI-compatible relays: GET {base}/models with `Authorization: Bearer`. + * + * `baseUrl` must already include the version segment (`…/v1`), exactly like + * the provider base URLs used for chat/completions. + */ +export interface RemoteModel { + id: string; + ownedBy?: string; +} + +export interface RemoteModelsResult { + ok: boolean; + message: string; + models: RemoteModel[]; +} + +export async function fetchProviderModels(request: { + baseUrl: string; + apiKey: string; + fetchImpl?: typeof fetch; + timeoutMs?: number; +}): Promise { + const base = (request.baseUrl || "").trim().replace(/\/+$/, ""); + if (!base) return { ok: false, message: "Base URL 为空", models: [] }; + if (!request.apiKey) return { ok: false, message: "API Key 未配置(请先填写并保存)", models: [] }; + const fetchImpl = request.fetchImpl || fetch; + const controller = new AbortController(); + const timer = setTimeout(() => controller.abort(), request.timeoutMs ?? 20000); + try { + const response = await fetchImpl(`${base}/models`, { + method: "GET", + headers: { Authorization: `Bearer ${request.apiKey}`, "api-key": request.apiKey }, + signal: controller.signal + }); + if (!response.ok) { + const detail = await response.text().catch(() => ""); + return { ok: false, message: `HTTP ${response.status}${detail ? `:${detail.slice(0, 200)}` : ""}`, models: [] }; + } + const json = (await response.json()) as { data?: Array<{ id?: string; owned_by?: string }> }; + const models = Array.isArray(json.data) + ? json.data.filter((entry) => typeof entry?.id === "string" && entry.id).map((entry) => ({ id: entry.id as string, ownedBy: entry.owned_by })) + : []; + models.sort((a, b) => a.id.localeCompare(b.id)); + return { ok: true, message: `获取到 ${models.length} 个模型`, models }; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + return { ok: false, message: `请求失败:${message.includes("aborted") ? "超时" : message}`, models: [] }; + } finally { + clearTimeout(timer); + } +} diff --git a/src/paths.test.ts b/src/paths.test.ts index 3e941a3..68ff9f7 100644 --- a/src/paths.test.ts +++ b/src/paths.test.ts @@ -1,11 +1,31 @@ import assert from "node:assert/strict"; +import fs from "node:fs"; import os from "node:os"; import path from "node:path"; import test from "node:test"; -import { WORKSPACE_ENV, resolveDefaultWorkspace } from "./paths.js"; +import { LEGACY_WORKSPACE, WORKSPACE_ENV, defaultWorkspaceRoot, migrateLegacyWorkspace, resolveDefaultWorkspace } from "./paths.js"; -test("uses the default SecAgent workspace when no override is set", () => { - assert.equal(resolveDefaultWorkspace({}), path.join(os.homedir(), "SecAgentWorkspace")); +test("default workspace follows the platform data directory", () => { + const resolved = resolveDefaultWorkspace({}); + assert.equal(resolved, defaultWorkspaceRoot({})); + if (process.platform === "win32") { + assert.match(resolved, /[\\/]SecAgent[\\/]workspace$/); + assert.notEqual(resolved, path.join(os.homedir(), "SecAgentWorkspace")); + } else if (process.platform === "darwin") { + assert.equal(resolved, path.join(os.homedir(), "Library", "Application Support", "SecAgent", "workspace")); + } else { + assert.equal(resolved, path.join(os.homedir(), ".config", "SecAgent", "workspace")); + } +}); + +test("honours XDG_CONFIG_HOME on linux-style environments", () => { + const xdg = path.join(os.tmpdir(), "secagent-xdg-home"); + assert.equal(defaultWorkspaceRoot({ XDG_CONFIG_HOME: xdg, APPDATA: path.join(os.tmpdir(), "appdata") }, "linux"), path.join(xdg, "SecAgent", "workspace")); +}); + +test("honours APPDATA on windows-style environments", () => { + const appData = path.join(os.tmpdir(), "appdata"); + assert.equal(defaultWorkspaceRoot({ APPDATA: appData }, "win32"), path.join(appData, "SecAgent", "workspace")); }); test("resolves SECTL_WORKSPACE as an absolute workspace path", () => { @@ -14,5 +34,30 @@ test("resolves SECTL_WORKSPACE as an absolute workspace path", () => { }); test("ignores an empty SECTL_WORKSPACE value", () => { - assert.equal(resolveDefaultWorkspace({ [WORKSPACE_ENV]: " " }), path.join(os.homedir(), "SecAgentWorkspace")); + assert.equal(resolveDefaultWorkspace({ [WORKSPACE_ENV]: " " }), defaultWorkspaceRoot({})); +}); + +test("migrateLegacyWorkspace moves the legacy home directory once", () => { + const target = defaultWorkspaceRoot({}); + if (path.resolve(LEGACY_WORKSPACE) === path.resolve(target)) return; // already colocated + const legacyExisted = fs.existsSync(LEGACY_WORKSPACE); + const targetExisted = fs.existsSync(target); + if (!legacyExisted && !targetExisted) { + fs.mkdirSync(LEGACY_WORKSPACE, { recursive: true }); + fs.writeFileSync(path.join(LEGACY_WORKSPACE, "marker.txt"), "data", "utf8"); + } + const first = migrateLegacyWorkspace({}); + if (legacyExisted || targetExisted) { + assert.equal(first, undefined); // nothing to do or target already present + return; + } + assert.equal(first, target); + assert.ok(fs.existsSync(path.join(target, "marker.txt"))); + assert.ok(!fs.existsSync(LEGACY_WORKSPACE)); + assert.equal(migrateLegacyWorkspace({}), undefined); // idempotent + fs.rmSync(target, { recursive: true, force: true }); +}); + +test("migrateLegacyWorkspace respects an explicit workspace override", () => { + assert.equal(migrateLegacyWorkspace({ [WORKSPACE_ENV]: "/tmp/secagent-override" }), undefined); }); diff --git a/src/paths.ts b/src/paths.ts index cdf92eb..9f2600d 100644 --- a/src/paths.ts +++ b/src/paths.ts @@ -1,12 +1,43 @@ +import fs from "node:fs"; import os from "node:os"; import path from "node:path"; export const WORKSPACE_ENV = "SECTL_WORKSPACE"; +/** + * Platform-appropriate per-user data root for SecAgent. + * + * Historically the workspace defaulted to `~/SecAgentWorkspace`, which litters + * the home directory on every platform (and on Windows lands directly in + * `C:\Users\`). The new default follows each OS convention and matches + * Electron's own `app.getPath("userData")` root, so there is exactly one + * directory to look at: + * + * - Windows: %APPDATA%\SecAgent\workspace + * - macOS: ~/Library/Application Support/SecAgent/workspace + * - Linux: $XDG_CONFIG_HOME/SecAgent/workspace (~/.config/SecAgent/workspace) + * + * Electron's userData directory (Roaming on Windows) then holds both the + * Chromium caches and the SecAgent workspace side by side, instead of + * spreading state across home, Roaming and the install directory. + */ +export function defaultWorkspaceRoot(env: NodeJS.ProcessEnv = process.env, platform: NodeJS.Platform = process.platform): string { + if (platform === "win32") { + const appData = env.APPDATA?.trim(); + if (appData) return path.join(appData, "SecAgent", "workspace"); + return path.join(os.homedir(), "AppData", "Roaming", "SecAgent", "workspace"); + } + if (platform === "darwin") { + return path.join(os.homedir(), "Library", "Application Support", "SecAgent", "workspace"); + } + const xdg = env.XDG_CONFIG_HOME?.trim(); + return path.join(xdg && path.isAbsolute(xdg) ? xdg : path.join(os.homedir(), ".config"), "SecAgent", "workspace"); +} + /** Resolve the default workspace, allowing the host process to override it. */ export function resolveDefaultWorkspace(env: NodeJS.ProcessEnv = process.env): string { const configured = env[WORKSPACE_ENV]?.trim(); - return configured ? expandPath(configured) : path.join(os.homedir(), "SecAgentWorkspace"); + return configured ? expandPath(configured) : defaultWorkspaceRoot(env); } export const DEFAULT_WORKSPACE = resolveDefaultWorkspace(); @@ -17,3 +48,34 @@ export function expandPath(input: string, base = process.cwd()): string { : input; return path.resolve(base, expanded); } + +/** The pre-migration default workspace location. */ +export const LEGACY_WORKSPACE = path.join(os.homedir(), "SecAgentWorkspace"); + +/** + * Move a legacy `~/SecAgentWorkspace` to the platform-appropriate default the + * first time the app runs after the change. No-op when the override env var is + * set, the legacy directory is gone, or the target already exists. + */ +export function migrateLegacyWorkspace(env: NodeJS.ProcessEnv = process.env): string | undefined { + if (env[WORKSPACE_ENV]?.trim()) return undefined; + const target = defaultWorkspaceRoot(env); + if (path.resolve(LEGACY_WORKSPACE) === path.resolve(target)) return undefined; + if (!fs.existsSync(LEGACY_WORKSPACE) || fs.existsSync(target)) return undefined; + try { + fs.mkdirSync(path.dirname(target), { recursive: true }); + fs.renameSync(LEGACY_WORKSPACE, target); + return target; + } catch { + // Cross-device rename (e.g. home on another drive): copy then retire the old tree. + try { + fs.cpSync(LEGACY_WORKSPACE, target, { recursive: true }); + const retired = `${LEGACY_WORKSPACE}.migrated`; + fs.rmSync(retired, { recursive: true, force: true }); + fs.renameSync(LEGACY_WORKSPACE, retired); + return target; + } catch { + return undefined; + } + } +} diff --git a/src/pi-tools.ts b/src/pi-tools.ts index 41ef92a..eaff795 100644 --- a/src/pi-tools.ts +++ b/src/pi-tools.ts @@ -7,6 +7,36 @@ import type { ToolImageContent } from "./tool-content.js"; const execAsync = promisify(exec); +/** Supported local image formats shared by `look_at` and the vision sub-model tool. */ +const IMAGE_MEDIA_TYPES: Record = { ".png": "image/png", ".jpg": "image/jpeg", ".jpeg": "image/jpeg", ".webp": "image/webp", ".gif": "image/gif" }; +export const MAX_IMAGE_BYTES = 12 * 1024 * 1024; + +export interface ReadImageResult { + filePath: string; + name: string; + mimeType: string; + base64: string; +} + +export function resolveWorkspacePath(workspace: string, filePath: string): string { + return path.isAbsolute(filePath) ? filePath : path.resolve(workspace, filePath); +} + +/** + * Shared image read + validation used both by the `look_at` Pi tool (returns the image to a + * vision-capable main model) and by `secagent__look_at_image` (feeds the image to a dedicated + * vision sub-model and returns text). + */ +export async function readImageFile(workspace: string, filePath: string): Promise { + const resolved = resolveWorkspacePath(workspace, filePath); + const mediaType = IMAGE_MEDIA_TYPES[path.extname(resolved).toLowerCase()]; + if (!mediaType) throw new Error("仅支持 png、jpg、jpeg、webp、gif 图片"); + const stat = await fs.stat(resolved); + if (!stat.isFile()) throw new Error("path 不是文件"); + if (stat.size > MAX_IMAGE_BYTES) throw new Error("图片不能超过 12 MB"); + return { filePath: resolved, name: path.basename(resolved), mimeType: mediaType, base64: (await fs.readFile(resolved)).toString("base64") }; +} + export const piTools: AgentTool[] = [ { key: "look_at", description: "查看本地图片并把图片内容直接提供给模型。path 可使用绝对路径或相对于工作区的路径;仅支持 png、jpg、jpeg、webp、gif 图片。需要理解图片内容时必须调用此工具,不要只读取图片文件的二进制内容。", inputSchema: { type: "object", additionalProperties: false, required: ["path"], properties: { path: { type: "string", description: "图片的绝对路径或相对于工作区的路径" } } } }, { key: "read", description: "读取文件内容。path 可使用绝对路径或相对于工作区的路径。", inputSchema: { type: "object", additionalProperties: false, required: ["path"], properties: { path: { type: "string" }, offset: { type: "integer", minimum: 1 }, limit: { type: "integer", minimum: 1 } } } }, @@ -20,14 +50,8 @@ function resolvePath(workspace: string, filePath: string): string { return path. export async function callPiTool(workspace: string, key: string, args: Record): Promise { if (key === "look_at") { if (typeof args.path !== "string" || !args.path.trim()) throw new Error("look_at 需要非空 path"); - const filePath = resolvePath(workspace, args.path); - const mediaTypes: Record = { ".png": "image/png", ".jpg": "image/jpeg", ".jpeg": "image/jpeg", ".webp": "image/webp", ".gif": "image/gif" }; - const mediaType = mediaTypes[path.extname(filePath).toLowerCase()]; - if (!mediaType) throw new Error("look_at 仅支持 png、jpg、jpeg、webp、gif 图片"); - const stat = await fs.stat(filePath); - if (!stat.isFile()) throw new Error("look_at 的 path 不是文件"); - if (stat.size > 12 * 1024 * 1024) throw new Error("look_at 图片不能超过 12 MB"); - const result: ToolImageContent = { type: "image", data: (await fs.readFile(filePath)).toString("base64"), mimeType: mediaType, name: path.basename(filePath), path: filePath }; + const image = await readImageFile(workspace, args.path); + const result: ToolImageContent = { type: "image", data: image.base64, mimeType: image.mimeType, name: image.name, path: image.filePath }; return result; } if (key === "read") { diff --git a/src/renderer/src/App.tsx b/src/renderer/src/App.tsx index f517dca..85a346c 100644 --- a/src/renderer/src/App.tsx +++ b/src/renderer/src/App.tsx @@ -11,9 +11,10 @@ import { WorkspaceFileStrip } from "./components/WorkspaceFileStrip.js"; import { stripWorkspaceFilesMarkup } from "../../workspace-file-contract.js"; import { reasoningEffortLabels, traceLabel } from "./constants.js"; import type { TraceEvent } from "./constants.js"; -import { isOfficialModel, isOfficialTierModel, reasoningEffortsForModel } from "./utils.js"; +import { isOfficialModel, isOfficialTierModel, isOfficialVisionModel, reasoningEffortsForModel } from "./utils.js"; import { officialTiers, tierDefaultId } from "./constants.js"; import { buildQuotedUserMessage, parseQuotedUserMessage, webSearchUrl } from "../../quoted-message.js"; +import { DaySeparator, DeleteButton, ErrorStateCard, MatrixOrb, MessageActions, ScrollProgress, ThoughtLine, VoicePill, daySeparatorLabel } from "./components/ui/Bits.js"; function selectionInElement(element: HTMLElement): string { const selection = window.getSelection(); @@ -83,6 +84,8 @@ export function App() { const [quotedText, setQuotedText] = useState(""); const [speakingMessageId, setSpeakingMessageId] = useState(null); const [readingStatus, setReadingStatus] = useState<"loading" | "playing" | null>(null); + // Codex-style sensitive tool-call confirmation, pending user decision. + const [toolConfirmation, setToolConfirmation] = useState<{ confirmationId: string; tool: string; arguments: Record; reason: string } | null>(null); const [trace, setTrace] = useState([]); const messagesRef = useRef(null); const answerContentRef = useRef(null); @@ -91,7 +94,25 @@ export function App() { const answerScrollLockTimer = useRef(undefined); const answerStartScrollPending = useRef(false); const modelMenuEnd = useRef(null); - const orderedModels = useMemo(() => [...models.filter(isOfficialModel), ...models.filter((model) => !isOfficialModel(model))], [models]); + // Official tiers first, then custom models clustered by provider so the + // submenu can render a labelled group header per provider. The relay's + // virtual-vision model is a vision-tool backend only, never a main model. + const orderedModels = useMemo(() => { + const visible = models.filter((model) => !isOfficialVisionModel(model)); + const official = visible.filter(isOfficialModel); + const custom = visible.filter((model) => !isOfficialModel(model)); + const clustered: ModelOption[] = []; + const byProvider = new Map(); + for (const model of custom) { + const group = model.providerLabel || "自定义模型"; + const bucket = byProvider.get(group); + if (bucket) bucket.push(model); + else { byProvider.set(group, [model]); } + } + for (const bucket of byProvider.values()) clustered.push(...bucket); + return [...official, ...clustered]; + }, [models]); + const modelGroupLabel = (model: ModelOption): string => isOfficialModel(model) ? "官方服务" : (model.providerLabel || "自定义模型"); const selectedModel = models.find((model) => model.id === selectedModelId); const reasoningEfforts = useMemo(() => reasoningEffortsForModel(selectedModel), [selectedModel]); useEffect(() => { @@ -106,6 +127,8 @@ export function App() { const fileInputRef = useRef(null); const audioRef = useRef<{ context: AudioContext; stream: MediaStream; source: MediaStreamAudioSourceNode; processor: ScriptProcessorNode } | undefined>(undefined); const recordingRef = useRef(false); + /** 选定的输入音频设备 ID(settings.speech.audio.input;undefined = 系统默认)。 */ + const audioInputRef = useRef(undefined); const voicePendingSendRef = useRef(undefined); const speechInputId = useRef(0); const speechSessionRef = useRef<{ @@ -180,6 +203,8 @@ export function App() { ]); const customMode = Boolean(savedSettings.customModelMode); setCustomModelMode(customMode); + const savedInput = (savedSettings as { speech?: { audio?: { input?: string } } }).speech?.audio?.input; + audioInputRef.current = savedInput && savedInput !== "auto" && savedInput !== "default" ? savedInput : undefined; const defaultReasoning = (savedSettings.defaultReasoningEffort || "high") as ReasoningEffort; setDefaultEffort(defaultReasoning); setReasoningEffort(defaultReasoning); @@ -187,9 +212,10 @@ export function App() { setSession(active); requestAnimationFrame(() => textareaRef.current?.focus()); const configured = await modelsPromise; - const preferred = configured.find((model) => model.id === savedSettings.defaultModelId) - || configured.find((model) => isOfficialTierModel(model) && model.model === tierDefaultId) - || configured[0]; + const mainModels = configured.filter((model) => !isOfficialVisionModel(model)); + const preferred = mainModels.find((model) => model.id === savedSettings.defaultModelId) + || mainModels.find((model) => isOfficialTierModel(model) && model.model === tierDefaultId) + || mainModels[0]; setSelectedModelId(preferred?.id || ""); })(); }, [bridge]); @@ -197,6 +223,10 @@ export function App() { useEffect(() => { if (!bridge) return; return bridge.onSettingsChanged(() => { + void bridge.getSettings().then((current) => { + const input = (current as { speech?: { audio?: { input?: string } } } | null)?.speech?.audio?.input; + audioInputRef.current = input && input !== "auto" && input !== "default" ? input : undefined; + }).catch(() => undefined); void bridge.listModels().then((models) => { setModels(models); void bridge.getSettings().then((settings) => { @@ -204,10 +234,10 @@ export function App() { setCustomModelMode(customMode); setDefaultEffort((settings.defaultReasoningEffort || "high") as ReasoningEffort); setSelectedModelId((current) => { - if (models.some((model) => model.id === (settings.defaultModelId || current))) return settings.defaultModelId || current; + if (models.some((model) => !isOfficialVisionModel(model) && model.id === (settings.defaultModelId || current))) return settings.defaultModelId || current; const tier = models.find((model) => isOfficialTierModel(model) && model.model === tierDefaultId); if (tier) return tier.id; - return models[0]?.id || ""; + return models.find((model) => !isOfficialVisionModel(model))?.id || ""; }); setReasoningEffort(settings.defaultReasoningEffort || "high"); }); @@ -275,6 +305,19 @@ export function App() { }); }, [bridge]); + // Sensitive tool calls pause here until the user approves, rejects, or the + // 5-minute main-process timeout fires. + useEffect(() => { + if (!bridge) return; + return bridge.onToolConfirmation((payload) => setToolConfirmation(payload)); + }, [bridge]); + + const resolveToolConfirmation = (approved: boolean, always = false) => { + if (!toolConfirmation) return; + void bridge?.respondToolConfirmation({ confirmationId: toolConfirmation.confirmationId, approved, always, signature: always ? `${toolConfirmation.tool}|${String(toolConfirmation.arguments.command ?? "").trim().split(/\s+/)[0] || "*"}`.toLowerCase() : undefined }); + setToolConfirmation(null); + }; + useEffect(() => { const closeOnOutsideClick = (event: PointerEvent) => { if (!modelMenuEnd.current?.contains(event.target as Node)) { setModelMenuOpen(false); setModelSubmenu(null); } @@ -577,7 +620,7 @@ export function App() { setSpeechMode(null); return false; } - const stream = await navigator.mediaDevices.getUserMedia({ audio: { channelCount: 1, echoCancellation: true, noiseSuppression: true, autoGainControl: true } }); + const stream = await navigator.mediaDevices.getUserMedia({ audio: { channelCount: 1, echoCancellation: true, noiseSuppression: true, autoGainControl: true, ...(audioInputRef.current ? { deviceId: { exact: audioInputRef.current } } : {}) } }); if (speechSession.stopRequested) { stream.getTracks().forEach((track) => track.stop()); if (speechSessionRef.current === speechSession) speechSessionRef.current = undefined; @@ -802,7 +845,7 @@ export function App() { }, [recording, speechProcessing]); if (!bridge) { - return

SecAgent 桌面桥接未加载

请退出应用后重新运行 pnpm dev。若仍出现此提示,请检查 Electron 的 preload 启动日志。

; + return
; } return
@@ -827,7 +870,7 @@ export function App() { {sessions.map((item) =>
- + void deleteSession(item.id)} />
)} @@ -848,21 +891,37 @@ export function App() {
- {session?.messages.length === 0 &&

开始一个课堂操作

例如:查询张三积分,或给张三加 2 分。

} - {session?.messages.map((message) => { + + {session?.messages.length === 0 &&

开始一个课堂操作

例如:查询张三积分,或给张三加 2 分。

} + {session?.messages.map((message, index) => { + const previous = index > 0 ? session.messages[index - 1] : undefined; + const dayLabel = daySeparatorLabel(message.createdAt); const activities = message.activities?.length ? message.activities : message.toolCalls?.length ? message.toolCalls.map((call) => ({ kind: "tool" as const, ...call })) : message.id === latestAssistantId ? traceActivities : []; const reading = speakingMessageId === message.id; const visibleContent = message.role === "assistant" ? stripWorkspaceFilesMarkup(message.content) : message.content; - return
{message.role === "user" ? "教师" : "SecAgent"} · {new Date(message.createdAt).toLocaleTimeString()}
{message.role === "assistant" && }{message.role === "user" && message.attachments?.length ? : null}
{message.role === "user" ? "你" : SecAgent}
{ event.preventDefault(); const selection = selectionInElement(event.currentTarget); setMessageMenu({ x: Math.min(event.clientX, window.innerWidth - 180), y: Math.min(event.clientY, window.innerHeight - 176), messageId: message.id, text: message.content, role: message.role, selection }); }}>{message.role === "assistant" ? {visibleContent} : }
{message.role === "assistant" && }{reading && (readingStatus === "loading" ? : )}
; + return {(!previous || daySeparatorLabel(previous.createdAt) !== dayLabel) && }
{message.role === "user" ? "教师" : "SecAgent"} · {new Date(message.createdAt).toLocaleTimeString()}
{message.role === "assistant" && }{message.role === "user" && message.attachments?.length ? : null}
{message.role === "user" ? "你" : SecAgent}
{ event.preventDefault(); const selection = selectionInElement(event.currentTarget); setMessageMenu({ x: Math.min(event.clientX, window.innerWidth - 180), y: Math.min(event.clientY, window.innerHeight - 176), messageId: message.id, text: message.content, role: message.role, selection }); }}>{message.role === "assistant" ? {visibleContent} : }
{message.role === "assistant" && }{reading && (readingStatus === "loading" ? : )}
void copyText(message.content)} />
; })} - {sending && !finishing &&
SecAgent · 正在生成
SecAgent
{streamingOutput ? {stripWorkspaceFilesMarkup(streamingOutput)} : "正在调用模型与工具…"}
} + {sending && !finishing &&
SecAgent · 正在生成
SecAgent
{streamingOutput ? {stripWorkspaceFilesMarkup(streamingOutput)} : "正在调用模型与工具…"}
}
-
{ if ((event.target as Element).closest('.icon-button img[src="/image-icon.svg"]')) fileInputRef.current?.click(); }} onPaste={handlePaste} onDragEnter={(event) => { if (event.dataTransfer.types.includes("Files")) { event.preventDefault(); setComposerDragging(true); } }} onDragOver={(event) => { if (event.dataTransfer.types.includes("Files")) event.preventDefault(); }} onDragLeave={(event) => { if (!event.currentTarget.contains(event.relatedTarget as Node)) setComposerDragging(false); }} onDrop={handleDrop}> { void addImageFiles(event.target.files || []); event.target.value = ""; }} />{attachments.length > 0 &&
setAttachments((current) => current.filter((attachment) => attachment.id !== id))} />
}{quotedText &&
引用

{quotedText}

}{attachmentError &&
{attachmentError}
}{speechStatus && !recording && !speechProcessing &&
{speechStatus}
} + {toolConfirmation &&
+
+

模型请求执行敏感操作

+

{toolConfirmation.reason}

+
{toolConfirmation.tool}
{JSON.stringify(toolConfirmation.arguments, null, 2).slice(0, 2000)}
+

允许后该操作将在本机执行。如不信任此请求请拒绝;拒绝后模型会收到拦截说明并尝试其他方式。

+
+ + + +
+
+
} + { if ((event.target as Element).closest('.icon-button img[src="/image-icon.svg"]')) fileInputRef.current?.click(); }} onPaste={handlePaste} onDragEnter={(event) => { if (event.dataTransfer.types.includes("Files")) { event.preventDefault(); setComposerDragging(true); } }} onDragOver={(event) => { if (event.dataTransfer.types.includes("Files")) event.preventDefault(); }} onDragLeave={(event) => { if (!event.currentTarget.contains(event.relatedTarget as Node)) setComposerDragging(false); }} onDrop={handleDrop}> { void addImageFiles(event.target.files || []); event.target.value = ""; }} />{attachments.length > 0 &&
setAttachments((current) => current.filter((attachment) => attachment.id !== id))} />
}{quotedText &&
引用

{quotedText}

}{attachmentError &&
{attachmentError}
}{(recording || speechProcessing) && speechMode === "streaming" ? : speechStatus && !recording && !speechProcessing &&
{speechStatus}
} {speechMode === "hold" && (recording || speechProcessing) ?
{!speechProcessing &&
拖到这里取消松开取消识别
拖到这里转文字松开写入输入框
}
{speechProcessing ? speechStatus || "正在识别…" : voiceDropZone === "edit" ? "松开写入输入框" : voiceDropZone === "cancel" ? "松开取消" : "松开直接发送"}
-
: <>
+
: <>