Compare commits

...
Author SHA1 Message Date
yxxandClaude Code dd5739bb43 fix(backend): FTS 阈值过滤处理 span=0 退化场景
全部命中 bm25 相同时归一化皆为 1.0,阈值超过 1.0 应无命中,
与旧 normalize_scores 语义对齐。

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-03 23:57:00 +08:00
yxx 58a8d64fdb Merge remote-tracking branch 'origin/main' into feat/knowledge-retrieval-core
# Conflicts:
#	README.md
#	backend/app/routes.py
#	docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md
#	docs/development/Knowledge与Retrieval-Core开发说明.md
2026-09-03 23:45:52 +08:00
yxxandClaude Code e3f8d76243 fix(backend): 落实 PR #12 评审意见
- P1 事件循环让出:run_rag 在样本边界 await asyncio.sleep(0),运行中取消/进度/SSE 可及时调度
- P2 SSE 终止事件:历史回放期间识别终止事件并结束流,try/finally 保证订阅清理
- P2 FTS 截断:fts 走数据库侧精确分页与计数,阈值经 bm25 截止值换算,不再受 5000 条固定截断
- P2 仅块标注:expected_block_ids 从块反查所属笔记,避免合法样本被判零分

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-03 23:37:09 +08:00
Kronecker 0a4ca1284c Merge pull request 'feat(mcp): 完善独立 MCP 服务器配置中心与连接生命周期管理' (#13) from feat/mcp-server-registry into main
Reviewed-on: #13
2026-09-03 22:45:55 +08:00
admin f8e499df5d fix(mcp): 修复配置导入、凭据管理与协议边界
修复生命周期锁阻塞事件循环、旧连接回调误停新连接及超时契约不一致。

补齐 MCP JSON 兼容导入、密钥拆分与失败重试,修复 Header 大小写草稿丢失,迁移大小写敏感的环境变量凭据。

在 SSE 行拼接前限制缓冲大小,增加并发、迁移和流式输入回归测试;忽略本机 MCP 数据及 server.json/servers.json。

验证:后端 185 项、前端 54 项测试通过,前端生产构建、相关文件 Ruff 与暂存差异检查通过。
2026-09-03 22:38:35 +08:00
yxxandClaude Code 5df2ec5ecd docs: 补齐 PR #11 评审要求的文档同步
- 技术栈说明实施状态:RAG Benchmark 标记为已完成、Agent Benchmark 暂缓
- 新增 Benchmark 开发说明,并登记到文档索引
- README 回归基线更新为后端 157 / 前端 29

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-03 22:35:45 +08:00
yxxandClaude Code bc7c934375 fix(backend): 落实 PR #11 第二轮评审意见
- Benchmark 容量淘汰只删终态 run,满容量且全活动时返回 BENCHMARK_CAPACITY_EXCEEDED
- 创建 run 前校验索引兼容性(BENCHMARK_INDEX_INCOMPATIBLE)
- 取消 run 补发 RunCancelled 终止事件;失败分支脱敏(BENCHMARK_RUN_FAILED)
- 失败样本计入汇总分母,报告输出 total/successful/failed/failure_rate
- load_dataset 按文件名隔离无关损坏文件,顶层非对象拒绝
- FTS score_threshold 先于计数/分页,total 与 items 一致
- Benchmark SSE 支持 Last-Event-ID 游标
- 移除 Agent Benchmark 501 占位接口
- 同步第二阶段接口契约与开发说明文档

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-03 22:20:11 +08:00
admin 3001d6089a 添加MCP客户端超时配置和连接管理改进
添加了MCP客户端的超时配置功能,包括启动超时和工具调用超时参数。
改进了HTTP客户端和标准IO客户端的超时处理机制,确保请求在指定时间内完成或取消。
增加了对MCP服务器数量的限制,防止配置过多服务器导致系统不稳定。
增强了错误处理机制,当连接异常时能够正确清理资源并移除桥接主机。
添加了对大型MCP消息的大小验证,防止过大的请求导致系统问题。
优化了密钥更改后的处理流程,确保在修改密钥时停用服务器并要求重新测试。
2026-09-03 16:16:58 +08:00
admin 9f2c46ab39 feat(mcp): complete remote transports and configuration workflow 2026-09-03 15:25:41 +08:00
admin d4ffbdcadd feat(mcp): add standalone server registry
Implement C.1 stdio MCP server CRUD, encrypted environment secrets, command digest approval, connection tests, lifecycle recovery, and dynamic tool registration. Add the standalone frontend configuration center, contracts, regression tests, and development documentation.
2026-09-03 14:49:49 +08:00
admin 0c392eebe4 feat(frontend): complete MCP plugin management UI 2026-09-03 14:22:29 +08:00
yxxandClaude Code d591a79942 fix(backend): 落实 PR #9 评审意见
- 检索调优参数(rrf_k/rerank/rerank_candidates/score_threshold)透传到引擎实际执行
- Recall 去重,避免同一 Note 多 Block 重复导致 Recall 超 1
- RAG 运行改为后台异步执行:创建即 queued + 202,支持取消与 SSE 实时事件
- 数据集元数据校验,坏文件隔离跳过;citation_required 语义修正
- modes 空/重复校验;配置快照记录模型版本与索引元信息

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-02 23:20:34 +08:00
admin 0bbff0090e docs: 同步第二阶段进度与扩展复盘 2026-09-02 20:22:06 +08:00
Kronecker c7294982f6 Merge pull request 'feat(extension): 实现 Plugin Command 与 Settings Contribution' (#10) from feat/plugin-command-settings into main
Reviewed-on: #10
2026-09-02 20:05:42 +08:00
admin 1218b5cd71 fix(extension): 收紧插件命令运行时契约 2026-09-02 20:03:39 +08:00
admin 09c1bdac21 fix(extension): 向MCP命令传递插件设置 2026-09-02 18:23:32 +08:00
admin b16230ac4c fix(extension): 按Schema资源作用域校验引用 2026-09-02 15:47:54 +08:00
admin ec88795b11 fix(extension): 保护插件凭据引用与删除事务 2026-09-02 15:31:39 +08:00
admin cff38158f6 fix(extension): 完成MCP命令目标并收紧Schema边界 2026-09-02 14:52:03 +08:00
admin c1bac00d12 fix(extension): 收紧插件密钥与命令清单边界 2026-09-02 14:08:47 +08:00
admin 0a9cad1c76 feat(extension): 实现插件命令与设置贡献 2026-09-02 12:52:40 +08:00
yxxandClaude Code 5cb7e6afae chore(backend): 移除误提交的验收笔记
验收笔记此前被误纳入 benchmark 提交,现摘除跟踪,文件保留在本地磁盘。

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-01 23:35:11 +08:00
yxx f74b751568 feat: benchmark功能开发完成 2026-09-01 23:30:11 +08:00
yxx 0e7bb20a5d feat: 完成benchmark后端功能 2026-09-01 23:19:28 +08:00
Kronecker d99f93b617 Merge pull request 'feat(extension): 接入 stdio MCP Bridge 与 Plugin Host' (#8) from feat/mcp-plugin-host into main
Reviewed-on: #8
2026-09-01 22:10:56 +08:00
admin 605bfc1c1a fix(extension): 强化 MCP 参数与生产运行门禁 2026-09-01 21:25:05 +08:00
admin 199dd25c3e fix(extension): 完善 MCP 参数与运行安全边界 2026-09-01 16:06:23 +08:00
admin 20920b6845 fix(extension): 修复 MCP Host 资源与协议边界 2026-09-01 12:11:30 +08:00
admin a8e0fe8ff3 docs(extension): 补充阶段C MCP开发说明 2026-09-01 11:32:47 +08:00
admin 4894029a0f feat(extension): 接入 stdio MCP Plugin Host 2026-09-01 11:32:05 +08:00
Kronecker 8463f70fc9 Merge pull request 'feat(agent): 持久化 Agent Trace 并支持 SSE 断点恢复' (#7) from feat/agent-trace-persistence into main
Reviewed-on: #7
2026-09-01 10:36:43 +08:00
admin 3a6643ca6f fix(agent): 保留持久化Run完整内容 2026-09-01 10:05:58 +08:00
admin d1458ee9fb docs: 重组文档目录并补充CI/CD细则 2026-09-01 09:55:40 +08:00
admin 1f840bf48f docs(agent): 记录Trace问题与修复方案 2026-09-01 00:49:30 +08:00
admin a820d14656 feat(agent): 持久化Trace并支持SSE恢复 2026-09-01 00:39:37 +08:00
admin 8553d6f3c3 feat(workspace): 接入真实Vault数据链路 2026-08-31 21:38:56 +08:00
admin 2cb47c08f1 docs: 清理接口文档行尾格式 2026-08-31 20:27:23 +08:00
admin f899b50930 docs(api): 规划第二阶段统一接口契约 2026-08-31 20:26:53 +08:00
admin b5e1d5c06b docs: 更新第二阶段技术栈基线 2026-08-31 20:08:14 +08:00
admin 9c38a2f6af docx:添加第二阶段团队分工表
添加详细的第二阶段开发计划文档,包括:

- 阶段目标和总体分工安排
- 各成员具体职责和任务分配(范涵宇、杨星萱、吉海燕)
- 技术实现方案和架构设计
- 跨模块协作关系和接口定义
- 优先级划分(P0/P1/P2)和验收标准
- 项目演示Demo规划和完成定义
2026-08-31 19:53:44 +08:00
admin 0b77b3c08a merge: 补充前后端代码注释与TODO约定 2026-08-30 23:03:33 +08:00
admin 17188357b8 docs: 建立代码注释与TODO维护约定 2026-08-30 23:00:11 +08:00
admin da4dde2951 chore(frontend): 补充状态与服务边界注释 2026-08-30 22:59:58 +08:00
admin 35dc1ddefb chore(backend): 补充核心流程注释与待办 2026-08-30 22:59:45 +08:00
Kronecker cd7b47116b Merge pull request 'feat(frontend): 统一页面视觉并完善 GitHub 代码主题' (#6) from feat/frontend-visual-polish into main
Reviewed-on: #6
2026-08-30 22:39:15 +08:00
admin a03dcafd5b docs(frontend): 补充 Shiki 预览与选择器回归说明 2026-08-30 22:31:51 +08:00
admin cba8f3f324 fix(frontend): 修复 Shiki 主题选择器并添加真实预览 2026-08-30 22:31:21 +08:00
admin c1008ea08d docs(frontend): 记录 GitHub 代码主题配置 2026-08-30 20:42:00 +08:00
admin d85362ab53 feat(frontend): 添加 GitHub 代码块主题设置 2026-08-30 20:41:42 +08:00
admin 71047aea17 chore(frontend): 将暂定品牌名统一为 NotesAgent 2026-08-30 20:25:51 +08:00
admin dd99a7a6f5 fix(frontend): 提升 Markdown 表格与列表对比度 2026-08-30 20:23:22 +08:00
admin afacee0ffd docs(frontend): 记录视觉优化与动效约束 2026-08-30 20:14:22 +08:00
admin 51ed010e54 feat(frontend): 统一页面视觉与轻量动效 2026-08-30 20:13:29 +08:00
admin d0dc358938 docs: 添加第一阶段测试验证手册 2026-08-30 15:20:12 +08:00
Kronecker 8d78af018b Merge pull request 'docs: 同步当前工程实现与验证基线' (#5) from fix/frontend-review-findings into main
Reviewed-on: #5
2026-08-30 15:16:17 +08:00
114 changed files with 13868 additions and 849 deletions
+6
View File
@@ -14,6 +14,12 @@ backend/.env
# 运行期生成的 SQLite 索引(vault 下的 Markdown 测试数据需提交) # 运行期生成的 SQLite 索引(vault 下的 Markdown 测试数据需提交)
backend/data/*.db* backend/data/*.db*
backend/data/credentials/ backend/data/credentials/
# 阶段验收笔记(验收用,不提交)
backend/data/vault/验收/
# 本机 MCP 配置、授权状态及服务器工作目录不得提交。
backend/data/mcp/
server.json
servers.json
# Editors and operating systems # Editors and operating systems
.idea/ .idea/
+20 -20
View File
@@ -2,7 +2,7 @@
> 本文件用于团队开发期间快速配置环境和启动项目,不是正式的项目 README。 > 本文件用于团队开发期间快速配置环境和启动项目,不是正式的项目 README。
> 当前基线:2026-08-30。第一阶段 Web 联调版的前端页面、Knowledge/Retrieval Core、AI/Agent Core、Extension Core、Provider 预设与本地加密凭据链路均已实现Tauri Host、Stronghold、真实桌面文件系统和 Sync Server 尚未接入。 > 当前基线:2026-09-03。第一阶段 Web 联调前后端已经完成;第二阶段已完成 Workspace 去 Mock、Agent Trace 持久化与 SSE 恢复、stdio MCP Bridge、隔离 Plugin Host、Plugin Command/Settings,以及独立 MCP Server 配置中心 C.1stdio、Streamable HTTP 与旧 SSE 兼容)。真实音频、Provider 协议增强、Benchmark、导出、主题包、Trace 可视化、Mermaid 与函数图像仍在后续开发Tauri Host、Stronghold、原生多 Vault 文件系统和 Sync Server 尚未接入。
## 当前目录 ## 当前目录
@@ -10,7 +10,7 @@
NotesAgent/ NotesAgent/
├── frontend/ Vue 3 + TypeScript + Vite 前端 ├── frontend/ Vue 3 + TypeScript + Vite 前端
├── backend/ FastAPI + Pydantic 后端 ├── backend/ FastAPI + Pydantic 后端
├── docs/ 分工与技术栈说明 ├── docs/ 架构、契约、开发说明、协作规范与问题复盘
└── server sync/ 云同步服务预留目录,当前未实现 └── server sync/ 云同步服务预留目录,当前未实现
``` ```
@@ -36,7 +36,7 @@ python --version
uv --version uv --version
``` ```
当前 Web 联调不需要 Rust 和 Tauri。开始桌面端集成后,再按照 `docs/AI笔记软件技术栈说明-团队版-v2.2.md` 安装 Rust Toolchain 与 Tauri CLI。 当前 Web 联调不需要 Rust 和 Tauri。开始桌面端集成后,再按照 `docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md` 安装 Rust Toolchain 与 Tauri CLI。
## 首次初始化 ## 首次初始化
@@ -118,7 +118,7 @@ cd frontend
pnpm test pnpm test
``` ```
当前回归基线为后端 71 项测试、前端 14 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 当前回归基线为后端 218 项测试、前端 32 项测试,且 TypeScript 类型检查和生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。
构建产物位于 `frontend/dist`,该目录不提交到 Git。 构建产物位于 `frontend/dist`,该目录不提交到 Git。
@@ -126,19 +126,18 @@ pnpm test
| 文档 | 用途 | | 文档 | 用途 |
| --- | --- | | --- | --- |
| [技术栈说明](docs/AI笔记软件技术栈说明-团队版-v2.2.md) | 目标架构、当前实施边界与模块依赖 | | [文档总索引](docs/README.md) | 文档分类、阅读顺序和维护规则 |
| [第一阶段分工表](docs/第一阶段分工表.md) | 成员职责、协作关系与当前交付状态 | | [技术栈说明](docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md) | 目标架构、第二阶段技术边界与模块依赖 |
| [后端接口契约](docs/后端接口契约-开发版.md) | HTTP/SSE 接口、错误和当前实现状态 | | [第二阶段分工表](docs/architecture/第二阶段团队分工表.md) | 第二阶段人员职责、任务顺序、协作关系与验收项 |
| [AI Core 与 Agent Core](docs/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core | | [后端接口契约](docs/contracts/后端接口契约-开发版.md) | HTTP/SSE 接口、错误和当前实现状态 |
| [Knowledge 与 Retrieval Core](docs/Knowledge与Retrieval-Core开发说明.md) | Block、索引、混合检索和 Citation | | [第二阶段接口契约](docs/contracts/第二阶段接口契约-开发版.md) | 第二阶段公共 DTO、计划接口、SSE、错误码与联调顺序 |
| [模型提供商与模型发现](docs/模型提供商与模型发现开发说明.md) | Provider 预设、模型发现和凭据边界 | | [AI Core 与 Agent Core](docs/development/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core |
| [前端页面需求](docs/前端页面需求说明-开发版.md) | 页面、交互、状态与验收基线 | | [MCP Bridge 与 Plugin Host](docs/development/MCP-Bridge与Plugin-Host开发说明.md) | stdio MCP、隔离进程、Tool 映射、状态与错误边界 |
| [前端实现说明](docs/前端壳子与接口层开发说明.md) | 当前前端目录、Service、SSE 和运行边界 | | [Plugin Command 与 Settings](docs/development/Plugin-Command与Settings开发说明.md) | Command Registry、Settings Schema、Secret 引用与联调边界 |
| [前端写作体验](docs/前端写作体验优化开发说明.md) | Milkdown、CodeMirror、格式栏和 Shiki | | [Plugin Command 与 Settings 复盘](docs/retrospectives/Plugin-Command与Settings问题与修复复盘.md) | 阶段 D 连续审阅发现的安全、事务、Schema 与运行时契约问题 |
| [Git 使用细则](docs/Git使用细则-团队开发版.md) | 分支、提交、PR、Review 与合并流程 | | [Git 使用细则](docs/guides/Git使用细则-团队开发版.md) | 分支、提交、PR、Review 与合并流程 |
| [后端审阅复盘](docs/后端全面审阅问题与修复复盘.md) | 后端问题原因、后果与修复方案 | | [CI/CD 细则](docs/guides/CI-CD细则-团队开发版.md) | Gitea 流水线、质量门禁、产物、发布与回滚规则 |
| [Knowledge/Retrieval 复盘](docs/Knowledge与Retrieval-Core问题与修复复盘.md) | 检索与事务问题复盘 | | [Agent Trace 复盘](docs/retrospectives/Agent-Core第二阶段问题与修复复盘.md) | Agent 持久化、SSE 恢复、事件契约与脱敏问题复盘 |
| [前端审阅复盘](docs/前端合并审阅问题与修复复盘.md) | 前端工程、契约和交互问题复盘 |
## 日常开发注意事项 ## 日常开发注意事项
@@ -148,6 +147,7 @@ pnpm test
- API 默认监听 `127.0.0.1:8000`,前端默认监听 `127.0.0.1:5173` - API 默认监听 `127.0.0.1:8000`,前端默认监听 `127.0.0.1:5173`
- 后端附件目录默认是 `backend/data/attachments`,可通过 `APP_ATTACHMENTS_PATH` 覆盖;该目录由桌面 Host 管理。 - 后端附件目录默认是 `backend/data/attachments`,可通过 `APP_ATTACHMENTS_PATH` 覆盖;该目录由桌面 Host 管理。
- 跨模块接口发生变化时,需要同步更新前后端类型和 `docs` 中的接口说明。 - 跨模块接口发生变化时,需要同步更新前后端类型和 `docs` 中的接口说明。
- 当前前后端接口清单`docs/后端接口契约-开发版.md`OpenAPI `/openapi.json` 为准。 - 当前已实现接口`docs/contracts/后端接口契约-开发版.md`第二阶段规划接口见 `docs/contracts/第二阶段接口契约-开发版.md`;已实现能力`/openapi.json` 为准。
- 前端页面、交互、状态管理和第一阶段验收要求见 `docs/前端页面需求说明-开发版.md` - 前端页面、交互、状态管理及当前阶段后续页面需求见 `docs/contracts/前端页面需求说明-开发版.md`
- 分支、提交、Pull Request、Review 和冲突处理规范见 `docs/Git使用细则-团队开发版.md` - 分支、提交、Pull Request、Review 和冲突处理规范见 `docs/guides/Git使用细则-团队开发版.md`
- CI 检查、产物、发布和回滚规范见 `docs/guides/CI-CD细则-团队开发版.md`
+5 -5
View File
@@ -2,7 +2,7 @@
FastAPI + Pydantic 的本地 AI Core / Agent Core。项目使用 uv 管理依赖和虚拟环境。 FastAPI + Pydantic 的本地 AI Core / Agent Core。项目使用 uv 管理依赖和虚拟环境。
当前实现包含 Knowledge/Retrieval、Chat、Agent Runtime、Tool/Permission、Skill/Plugin、Provider Adapter、任务、索引和开发阶段凭据加密存储。Provider 支持 Mock、OpenAI Chat/OpenAI-Compatible 与 OllamaOpenAI Responses、Anthropic Messages、MCP 独立 Host 和真实语音模型仍属于后续阶段。 当前实现包含 Knowledge/Retrieval、Chat、Agent Runtime、Tool/Permission、Skill/Plugin、stdio MCP Host、Plugin Command/Settings、Provider Adapter、任务、索引和开发阶段凭据加密存储。Provider 支持 Mock、OpenAI Chat/OpenAI-Compatible 与 OllamaOpenAI Responses、Anthropic Messages、操作系统级 Plugin 沙箱和真实语音模型仍属于后续阶段。
```powershell ```powershell
uv sync uv sync
@@ -23,10 +23,10 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000
uv run pytest uv run pytest
``` ```
当前基线为 71 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY``DEEPSEEK_API_KEY``AINOTE_CREDENTIAL_<ID>` 注入;不要把真实密钥写入仓库。 当前基线为 136 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY``DEEPSEEK_API_KEY``AINOTE_CREDENTIAL_<ID>` 注入;不要把真实密钥写入仓库。`plugin.*` 是 Plugin Settings 的保留凭据命名空间,通用 Provider 凭据接口不能读写。
团队接口清单见 `../docs/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。 团队接口清单见 `../docs/contracts/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。
AI Core 与 Agent Core 的模块边界、Mock Provider 和 Tool Calling 调试方式见 `../docs/AI-Core与Agent-Core开发说明.md` AI Core 与 Agent Core 的模块边界、Mock Provider 和 Tool Calling 调试方式见 `../docs/development/AI-Core与Agent-Core开发说明.md`
Knowledge Core 与 Retrieval Core 的模块边界、数据模型、接口与检索流程见 `../docs/Knowledge与Retrieval-Core开发说明.md` Knowledge Core 与 Retrieval Core 的模块边界、数据模型、接口与检索流程见 `../docs/development/Knowledge与Retrieval-Core开发说明.md`
+11
View File
@@ -1,3 +1,5 @@
"""Agent 工具权限策略与一次性确认票据。"""
import asyncio import asyncio
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum from enum import Enum
@@ -51,6 +53,7 @@ class PermissionPolicy:
def mode_for(self, permission: str | None) -> PermissionMode: def mode_for(self, permission: str | None) -> PermissionMode:
if permission is None: if permission is None:
return PermissionMode.allow return PermissionMode.allow
# 未登记权限一律拒绝,防止扩展通过拼写错误或新权限绕过策略。
return self._rules.get(permission, PermissionMode.deny) return self._rules.get(permission, PermissionMode.deny)
@@ -63,6 +66,8 @@ class PermissionTicket:
class PermissionManager: class PermissionManager:
"""管理当前进程内的确认请求与会话级授权。"""
def __init__(self, policy: PermissionPolicy) -> None: def __init__(self, policy: PermissionPolicy) -> None:
self.policy = policy self.policy = policy
self._pending: dict[tuple[str, str], PermissionTicket] = {} self._pending: dict[tuple[str, str], PermissionTicket] = {}
@@ -94,10 +99,16 @@ class PermissionManager:
if ticket is None or ticket.future.done(): if ticket is None or ticket.future.done():
return False return False
if decision == "allow_session": if decision == "allow_session":
# 会话授权只存在于进程内,应用重启后按默认策略重新确认。
self._session_grants.add(ticket.permission) self._session_grants.add(ticket.permission)
ticket.future.set_result(decision) ticket.future.set_result(decision)
return True return True
def get_ticket(self, run_id: str, request_id: str) -> PermissionTicket | None:
"""只读返回待确认票据,供 Trace 记录权限类型;不暴露 Future 给接口层。"""
return self._pending.get((run_id, request_id))
def cancel_run(self, run_id: str) -> None: def cancel_run(self, run_id: str) -> None:
for key, ticket in list(self._pending.items()): for key, ticket in list(self._pending.items()):
if ticket.run_id == run_id: if ticket.run_id == run_id:
+187 -25
View File
@@ -1,3 +1,5 @@
"""Agent 运行时:负责模型轮次、工具调用、权限确认与事件发布。"""
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
@@ -5,17 +7,20 @@ import json
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime, timezone from datetime import datetime, timezone
from time import perf_counter
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from uuid import uuid4 from uuid import uuid4
from app.agent.permissions import PermissionManager, PermissionMode from app.agent.permissions import PermissionManager, PermissionMode
from app.agent.tools import ToolExecutionContext, ToolNotFoundError, ToolRegistry from app.agent.tools import ToolExecutionContext, ToolNotFoundError, ToolRegistry
from app.agent.trace_repository import AgentTraceRepository, sanitize_trace_value
from app.contracts import ( from app.contracts import (
AgentEvent, AgentEvent,
AgentEventType, AgentEventType,
AgentRun, AgentRun,
AgentRunCreateRequest, AgentRunCreateRequest,
AgentRunStatus, AgentRunStatus,
AgentTraceResponse,
Citation, Citation,
Message, Message,
MessageRole, MessageRole,
@@ -50,6 +55,8 @@ MAX_TOOL_CALLS_PER_TURN = 50
@dataclass(slots=True) @dataclass(slots=True)
class RunRecord: class RunRecord:
"""单次运行的可变上下文,仅由 AgentRuntime 持有。"""
run: AgentRun run: AgentRun
request: AgentRunCreateRequest request: AgentRunCreateRequest
skill_config: AgentConfiguration | None = None skill_config: AgentConfiguration | None = None
@@ -57,20 +64,25 @@ class RunRecord:
events: list[AgentEvent] = field(default_factory=list) events: list[AgentEvent] = field(default_factory=list)
subscribers: set[asyncio.Queue[AgentEvent]] = field(default_factory=set) subscribers: set[asyncio.Queue[AgentEvent]] = field(default_factory=set)
task: asyncio.Task[None] | None = None task: asyncio.Task[None] | None = None
next_sequence: int = 0
class AgentRuntime: class AgentRuntime:
"""进程内 Agent 编排器;对外返回深拷贝,避免调用方修改运行状态。"""
def __init__( def __init__(
self, self,
providers: ProviderRegistry, providers: ProviderRegistry,
tools: ToolRegistry, tools: ToolRegistry,
permissions: PermissionManager, permissions: PermissionManager,
skills: SkillRuntime | None = None, skills: SkillRuntime | None = None,
trace_repository: AgentTraceRepository | None = None,
) -> None: ) -> None:
self.providers = providers self.providers = providers
self.tools = tools self.tools = tools
self.permissions = permissions self.permissions = permissions
self.skills = skills self.skills = skills
self.trace_repository = trace_repository or AgentTraceRepository()
self._records: dict[str, RunRecord] = {} self._records: dict[str, RunRecord] = {}
async def create_run(self, request: AgentRunCreateRequest) -> AgentRun: async def create_run(self, request: AgentRunCreateRequest) -> AgentRun:
@@ -98,6 +110,7 @@ class AgentRuntime:
) )
allowed_tools = list(request.allowed_tools) allowed_tools = list(request.allowed_tools)
if skill_config is not None: if skill_config is not None:
# 同时指定 Skill 与工具白名单时取交集,避免 Skill 扩大调用权限。
allowed_tools = ( allowed_tools = (
[name for name in skill_config.allowed_tools if name in allowed_tools] [name for name in skill_config.allowed_tools if name in allowed_tools]
if allowed_tools if allowed_tools
@@ -109,22 +122,38 @@ class AgentRuntime:
skill_config=skill_config, skill_config=skill_config,
allowed_tools=allowed_tools, allowed_tools=allowed_tools,
) )
self.trace_repository.create_run(
run,
request,
self._config_snapshot(record),
)
self._records[run.run_id] = record self._records[run.run_id] = record
record.task = asyncio.create_task(self._execute(record), name=run.run_id) record.task = asyncio.create_task(self._execute(record), name=run.run_id)
return run.model_copy(deep=True) return run.model_copy(deep=True)
def get_run(self, run_id: str) -> AgentRun: def get_run(self, run_id: str) -> AgentRun:
return self._get_record(run_id).run.model_copy(deep=True) record = self._records.get(run_id)
if record is not None:
return record.run.model_copy(deep=True)
run = self.trace_repository.recover_interrupted(run_id)
if run is None:
raise AgentRunNotFoundError(run_id)
return run.model_copy(deep=True)
def list_runs(self, limit: int, offset: int) -> tuple[list[AgentRun], int]: def list_runs(self, limit: int, offset: int) -> tuple[list[AgentRun], int]:
records = sorted( items, total = self.trace_repository.list_runs(limit=limit, offset=offset)
self._records.values(), key=lambda item: item.run.created_at, reverse=True recovered = [
) self.trace_repository.recover_interrupted(item.run_id) or item
items = [item.run.model_copy(deep=True) for item in records[offset : offset + limit]] if item.run_id not in self._records
return items, len(records) else self._records[item.run_id].run.model_copy(deep=True)
for item in items
]
return recovered, total
async def cancel(self, run_id: str) -> AgentRun: async def cancel(self, run_id: str) -> AgentRun:
record = self._get_record(run_id) record = self._records.get(run_id)
if record is None:
return self.get_run(run_id)
if record.run.status in TERMINAL_STATUSES: if record.run.status in TERMINAL_STATUSES:
return record.run.model_copy(deep=True) return record.run.model_copy(deep=True)
record.run.cancelled = True record.run.cancelled = True
@@ -137,21 +166,53 @@ class AgentRuntime:
return record.run.model_copy(deep=True) return record.run.model_copy(deep=True)
def resolve_permission(self, run_id: str, request_id: str, decision: str) -> bool: def resolve_permission(self, run_id: str, request_id: str, decision: str) -> bool:
self._get_record(run_id) record = self._records.get(run_id)
return self.permissions.resolve(run_id, request_id, decision) if record is None:
return False
ticket = self.permissions.get_ticket(run_id, request_id)
resolved = self.permissions.resolve(run_id, request_id, decision)
if resolved:
self._publish(
record,
AgentEventType.permission_resolved,
{
"request_id": request_id,
"permission": ticket.permission if ticket else None,
"decision": decision,
},
)
return resolved
async def events(self, run_id: str) -> AsyncIterator[AgentEvent]: async def events(
record = self._get_record(run_id) self, run_id: str, *, after_sequence: int = -1
) -> AsyncIterator[AgentEvent]:
record = self._records.get(run_id)
run = self.get_run(run_id)
if record is None:
for event in self.trace_repository.list_events(
run_id, after_sequence=after_sequence
):
yield event
return
# 先注册订阅再读持久化历史;同一事件循环内没有 await,不会丢失交界事件。
queue: asyncio.Queue[AgentEvent] = asyncio.Queue() queue: asyncio.Queue[AgentEvent] = asyncio.Queue()
record.subscribers.add(queue) record.subscribers.add(queue)
history = [event.model_copy(deep=True) for event in record.events] history = self.trace_repository.list_events(
run_id, after_sequence=after_sequence
)
last_sequence = after_sequence
try: try:
for event in history: for event in history:
last_sequence = event.sequence
yield event yield event
if record.run.status in TERMINAL_STATUSES: if run.status in TERMINAL_STATUSES:
return return
while True: while True:
event = await queue.get() event = await queue.get()
if event.sequence <= last_sequence:
continue
last_sequence = event.sequence
yield event.model_copy(deep=True) yield event.model_copy(deep=True)
if event.event in { if event.event in {
AgentEventType.run_completed, AgentEventType.run_completed,
@@ -163,7 +224,9 @@ class AgentRuntime:
record.subscribers.discard(queue) record.subscribers.discard(queue)
async def wait(self, run_id: str) -> AgentRun: async def wait(self, run_id: str) -> AgentRun:
record = self._get_record(run_id) record = self._records.get(run_id)
if record is None:
return self.get_run(run_id)
if record.task: if record.task:
try: try:
await asyncio.shield(record.task) await asyncio.shield(record.task)
@@ -171,6 +234,17 @@ class AgentRuntime:
pass pass
return record.run.model_copy(deep=True) return record.run.model_copy(deep=True)
def get_trace(
self, run_id: str, *, after_sequence: int, limit: int
) -> AgentTraceResponse:
self.get_run(run_id)
trace = self.trace_repository.get_trace(
run_id, after_sequence=after_sequence, limit=limit
)
if trace is None:
raise AgentRunNotFoundError(run_id)
return trace
async def _execute(self, record: RunRecord) -> None: async def _execute(self, record: RunRecord) -> None:
try: try:
async with asyncio.timeout(record.request.run_timeout_seconds): async with asyncio.timeout(record.request.run_timeout_seconds):
@@ -201,6 +275,19 @@ class AgentRuntime:
for step in range(1, record.request.max_steps + 1): for step in range(1, record.request.max_steps + 1):
record.run.current_step = step record.run.current_step = step
record.run.updated_at = datetime.now(timezone.utc) record.run.updated_at = datetime.now(timezone.utc)
model_call_id = f"model_call_{uuid4().hex}"
started_at = perf_counter()
self._publish(
record,
AgentEventType.model_call_started,
{
"model_call_id": model_call_id,
"step": step,
"provider_id": record.request.provider_id,
"model": record.request.model,
},
)
try:
turn = await provider.complete( turn = await provider.complete(
ModelRequest( ModelRequest(
provider_id=record.request.provider_id, provider_id=record.request.provider_id,
@@ -211,6 +298,29 @@ class AgentRuntime:
metadata=self._request_metadata(record), metadata=self._request_metadata(record),
) )
) )
except Exception as exc:
self._publish(
record,
AgentEventType.model_call_failed,
{
"model_call_id": model_call_id,
"duration_ms": int((perf_counter() - started_at) * 1000),
"error_code": getattr(exc, "code", type(exc).__name__),
},
)
raise
self._publish(
record,
AgentEventType.model_call_completed,
{
"model_call_id": model_call_id,
"duration_ms": int((perf_counter() - started_at) * 1000),
"finish_reason": "tool_calls" if turn.tool_calls else "stop",
"input_tokens": turn.input_tokens,
"output_tokens": turn.output_tokens,
"tool_call_count": len(turn.tool_calls),
},
)
record.run.token_usage += turn.input_tokens + turn.output_tokens record.run.token_usage += turn.input_tokens + turn.output_tokens
self._publish( self._publish(
record, record,
@@ -243,11 +353,12 @@ class AgentRuntime:
messages.append( messages.append(
Message(role=MessageRole.assistant, content=turn.text or "", tool_calls=calls) Message(role=MessageRole.assistant, content=turn.text or "", tool_calls=calls)
) )
# 工具可以并发执行,但结果按模型原始调用顺序写回上下文,保证轮次可复现。
semaphore = asyncio.Semaphore(record.request.max_concurrent_tools) semaphore = asyncio.Semaphore(record.request.max_concurrent_tools)
async def execute(call: ToolCall) -> ToolResult: async def execute(call: ToolCall) -> ToolResult:
async with semaphore: async with semaphore:
return await self._execute_tool(record, call) return await self._execute_tool(record, call, model_call_id)
results = await asyncio.gather(*(execute(call) for call in calls)) results = await asyncio.gather(*(execute(call) for call in calls))
for call, result in zip(calls, results): for call, result in zip(calls, results):
@@ -280,8 +391,13 @@ class AgentRuntime:
self._fail(record, "MAX_STEPS_EXCEEDED", "Agent reached its maximum step count.") self._fail(record, "MAX_STEPS_EXCEEDED", "Agent reached its maximum step count.")
async def _execute_tool(self, record: RunRecord, call: ToolCall) -> ToolResult: async def _execute_tool(
self._publish(record, AgentEventType.tool_call, call.model_dump(mode="json")) self, record: RunRecord, call: ToolCall, parent_model_call_id: str
) -> ToolResult:
started_at = perf_counter()
call_data = call.model_dump(mode="json")
call_data["parent_model_call_id"] = parent_model_call_id
self._publish(record, AgentEventType.tool_call, call_data)
try: try:
registered = self.tools.get(call.name) registered = self.tools.get(call.name)
except ToolNotFoundError: except ToolNotFoundError:
@@ -295,7 +411,9 @@ class AgentRuntime:
error_code="TOOL_NOT_ALLOWED", error_code="TOOL_NOT_ALLOWED",
error_message="Tool is not included in allowed_tools.", error_message="Tool is not included in allowed_tools.",
) )
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json")) self._publish_tool_result(
record, result, parent_model_call_id, started_at
)
return result return result
permission = registered.definition.permission if registered else None permission = registered.definition.permission if registered else None
@@ -307,12 +425,15 @@ class AgentRuntime:
error_code="NETWORK_NOT_ALLOWED", error_code="NETWORK_NOT_ALLOWED",
error_message="Agent run does not allow network tools.", error_message="Agent run does not allow network tools.",
) )
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json")) self._publish_tool_result(
record, result, parent_model_call_id, started_at
)
return result return result
mode = self.permissions.mode_for(permission) mode = self.permissions.mode_for(permission)
if mode == PermissionMode.deny: if mode == PermissionMode.deny:
result = self._permission_denied(call) result = self._permission_denied(call)
elif mode == PermissionMode.confirm and permission: elif mode == PermissionMode.confirm and permission:
# 运行状态必须在等待期间可见,前端才能展示并处理权限确认卡片。
ticket = self.permissions.create_ticket(record.run.run_id, permission) ticket = self.permissions.create_ticket(record.run.run_id, permission)
record.run.status = AgentRunStatus.waiting_permission record.run.status = AgentRunStatus.waiting_permission
self._publish( self._publish(
@@ -337,11 +458,13 @@ class AgentRuntime:
error_code="PERMISSION_TIMEOUT", error_code="PERMISSION_TIMEOUT",
error_message="Tool permission confirmation timed out.", error_message="Tool permission confirmation timed out.",
) )
self._publish( self._publish_tool_result(
record, AgentEventType.tool_result, result.model_dump(mode="json") record, result, parent_model_call_id, started_at
) )
return result return result
record.run.status = AgentRunStatus.running record.run.status = AgentRunStatus.running
record.run.updated_at = datetime.now(timezone.utc)
self.trace_repository.save_run(record.run)
result = ( result = (
await self._invoke_tool(record, call) await self._invoke_tool(record, call)
if decision in {"allow_once", "allow_session"} if decision in {"allow_once", "allow_session"}
@@ -350,13 +473,31 @@ class AgentRuntime:
else: else:
result = await self._invoke_tool(record, call) result = await self._invoke_tool(record, call)
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json")) self._publish_tool_result(record, result, parent_model_call_id, started_at)
return result return result
def _publish_tool_result(
self,
record: RunRecord,
result: ToolResult,
parent_model_call_id: str,
started_at: float,
) -> None:
data = result.model_dump(mode="json")
data["parent_model_call_id"] = parent_model_call_id
data["duration_ms"] = int((perf_counter() - started_at) * 1000)
self._publish(record, AgentEventType.tool_result, data)
async def _invoke_tool(self, record: RunRecord, call: ToolCall) -> ToolResult: async def _invoke_tool(self, record: RunRecord, call: ToolCall) -> ToolResult:
try: try:
return await asyncio.wait_for( return await asyncio.wait_for(
self.tools.execute(call, ToolExecutionContext(run_id=record.run.run_id)), self.tools.execute(
call,
ToolExecutionContext(
run_id=record.run.run_id,
tool_call_id=call.tool_call_id,
),
),
timeout=record.request.tool_timeout_seconds, timeout=record.request.tool_timeout_seconds,
) )
except TimeoutError: except TimeoutError:
@@ -400,14 +541,19 @@ class AgentRuntime:
def _publish( def _publish(
self, record: RunRecord, event_type: AgentEventType, data: dict[str, object] self, record: RunRecord, event_type: AgentEventType, data: dict[str, object]
) -> None: ) -> None:
sanitized = sanitize_trace_value(data)
assert isinstance(sanitized, dict)
event = AgentEvent( event = AgentEvent(
event=event_type, event=event_type,
run_id=record.run.run_id, run_id=record.run.run_id,
sequence=len(record.events), sequence=record.next_sequence,
data=data, data=sanitized,
timestamp=datetime.now(timezone.utc), timestamp=datetime.now(timezone.utc),
) )
record.next_sequence += 1
record.events.append(event) record.events.append(event)
self.trace_repository.append_event(record.run, event)
# 内存只保留实时订阅窗口;完整审计轨迹由 SQLite 保存。
if len(record.events) > MAX_EVENTS_PER_RUN: if len(record.events) > MAX_EVENTS_PER_RUN:
del record.events[: len(record.events) - MAX_EVENTS_PER_RUN] del record.events[: len(record.events) - MAX_EVENTS_PER_RUN]
for queue in record.subscribers: for queue in record.subscribers:
@@ -421,6 +567,21 @@ class AgentRuntime:
metadata["retrieval"] = record.skill_config.retrieval.model_dump(mode="json") metadata["retrieval"] = record.skill_config.retrieval.model_dump(mode="json")
return metadata return metadata
def _config_snapshot(self, record: RunRecord) -> dict[str, object]:
provider = self.providers.get(record.request.provider_id).config
return {
"provider_id": record.request.provider_id,
"provider_type": provider.provider_type.value,
"model": record.request.model,
"capabilities": [item.value for item in provider.capabilities],
"skill_id": record.request.skill_id,
"allowed_tools": list(record.allowed_tools),
"max_steps": record.request.max_steps,
"token_budget": record.request.token_budget,
"allow_network": record.request.allow_network,
"metadata": record.request.metadata,
}
def _collect_citations(self, record: RunRecord, result: ToolResult) -> None: def _collect_citations(self, record: RunRecord, result: ToolResult) -> None:
if not result.success or not isinstance(result.output, dict): if not result.success or not isinstance(result.output, dict):
return return
@@ -448,6 +609,7 @@ class AgentRuntime:
raise AgentRunNotFoundError(run_id) from exc raise AgentRunNotFoundError(run_id) from exc
def _prune_records(self) -> None: def _prune_records(self) -> None:
# 只清理终态记录,绝不为了容量取消仍在执行或等待授权的任务。
overflow = len(self._records) - MAX_RUN_RECORDS + 1 overflow = len(self._records) - MAX_RUN_RECORDS + 1
if overflow <= 0: if overflow <= 0:
return return
+35 -1
View File
@@ -1,4 +1,7 @@
"""Agent 工具注册与执行边界。"""
import inspect import inspect
import threading
from dataclasses import dataclass from dataclasses import dataclass
from time import perf_counter from time import perf_counter
from typing import Any, Awaitable, Callable from typing import Any, Awaitable, Callable
@@ -8,6 +11,7 @@ from jsonschema import Draft202012Validator
from jsonschema.exceptions import ValidationError as JsonSchemaValidationError from jsonschema.exceptions import ValidationError as JsonSchemaValidationError
from app.contracts import ToolCall, ToolDefinition, ToolResult from app.contracts import ToolCall, ToolDefinition, ToolResult
from app.schema_security import reject_external_schema_references
ToolExecutor = Callable[[BaseModel, "ToolExecutionContext"], Any | Awaitable[Any]] ToolExecutor = Callable[[BaseModel, "ToolExecutionContext"], Any | Awaitable[Any]]
@@ -15,6 +19,7 @@ ToolExecutor = Callable[[BaseModel, "ToolExecutionContext"], Any | Awaitable[Any
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class ToolExecutionContext: class ToolExecutionContext:
run_id: str run_id: str
tool_call_id: str | None = None
@dataclass(slots=True) @dataclass(slots=True)
@@ -28,9 +33,21 @@ class ToolNotFoundError(LookupError):
pass pass
class ToolExecutionError(RuntimeError):
"""Executor 可预期失败,保留领域错误码而不是折叠成通用异常。"""
def __init__(self, code: str, message: str) -> None:
super().__init__(message)
self.code = code
self.message = message
class ToolRegistry: class ToolRegistry:
"""统一校验工具入参并隔离执行异常,避免单个工具击穿 Agent 主循环。"""
def __init__(self) -> None: def __init__(self) -> None:
self._tools: dict[str, RegisteredTool] = {} self._tools: dict[str, RegisteredTool] = {}
self._lock = threading.RLock()
def register( def register(
self, self,
@@ -38,6 +55,9 @@ class ToolRegistry:
arguments_model: type[BaseModel], arguments_model: type[BaseModel],
executor: ToolExecutor, executor: ToolExecutor,
) -> None: ) -> None:
Draft202012Validator.check_schema(definition.parameters)
reject_external_schema_references(definition.parameters)
with self._lock:
if definition.name in self._tools: if definition.name in self._tools:
raise ValueError(f"Tool already registered: {definition.name}") raise ValueError(f"Tool already registered: {definition.name}")
self._tools[definition.name] = RegisteredTool( self._tools[definition.name] = RegisteredTool(
@@ -47,12 +67,15 @@ class ToolRegistry:
) )
def unregister(self, name: str) -> None: def unregister(self, name: str) -> None:
with self._lock:
self._tools.pop(name, None) self._tools.pop(name, None)
def contains(self, name: str) -> bool: def contains(self, name: str) -> bool:
with self._lock:
return name in self._tools return name in self._tools
def get(self, name: str) -> RegisteredTool: def get(self, name: str) -> RegisteredTool:
with self._lock:
try: try:
return self._tools[name] return self._tools[name]
except KeyError as exc: except KeyError as exc:
@@ -60,6 +83,7 @@ class ToolRegistry:
def definitions(self, allowed: list[str] | None = None) -> list[ToolDefinition]: def definitions(self, allowed: list[str] | None = None) -> list[ToolDefinition]:
names = set(allowed) if allowed is not None else None names = set(allowed) if allowed is not None else None
with self._lock:
return [ return [
item.definition.model_copy(deep=True) item.definition.model_copy(deep=True)
for name, item in self._tools.items() for name, item in self._tools.items()
@@ -80,6 +104,7 @@ class ToolRegistry:
) )
try: try:
# JSON Schema 约束模型可见的协议,Pydantic 再完成运行时类型转换。
Draft202012Validator(registered.definition.parameters).validate(call.arguments) Draft202012Validator(registered.definition.parameters).validate(call.arguments)
arguments = registered.arguments_model.model_validate(call.arguments) arguments = registered.arguments_model.model_validate(call.arguments)
except (ValidationError, JsonSchemaValidationError) as exc: except (ValidationError, JsonSchemaValidationError) as exc:
@@ -103,7 +128,16 @@ class ToolRegistry:
output=output, output=output,
duration_ms=round((perf_counter() - started) * 1000), duration_ms=round((perf_counter() - started) * 1000),
) )
except Exception as exc: # Tool failures are isolated from the Agent loop. except ToolExecutionError as exc:
return ToolResult(
tool_call_id=call.tool_call_id,
name=call.name,
success=False,
error_code=exc.code,
error_message=exc.message,
duration_ms=round((perf_counter() - started) * 1000),
)
except Exception as exc: # 工具失败转换成结构化结果,由模型决定是否降级或重试。
return ToolResult( return ToolResult(
tool_call_id=call.tool_call_id, tool_call_id=call.tool_call_id,
name=call.name, name=call.name,
+372
View File
@@ -0,0 +1,372 @@
"""Agent Run/Event 持久化与 Trace 查询。
SQLite 中的事件是 SSE、前端 Trace 和 Benchmark 的共同事实来源。写入前统一脱敏和
限长,避免 Secret 或无限大的 Tool Result 进入审计数据。
"""
from __future__ import annotations
import json
import re
from datetime import datetime, timezone
from typing import Any
from app.contracts import (
AgentEvent,
AgentEventType,
AgentRun,
AgentRunCreateRequest,
AgentRunStatus,
AgentTraceResponse,
AgentTraceSummary,
)
from app.database.db import connect, transaction
MAX_TRACE_STRING = 4_096
MAX_TRACE_COLLECTION = 100
MAX_TRACE_DEPTH = 8
_SECRET_KEYS = {
"api_key",
"apikey",
"authorization",
"access_token",
"refresh_token",
"client_secret",
"password",
"secret",
"token",
}
_SECRET_KEY_SUFFIXES = ("_api_key", "_password", "_secret")
_TERMINAL_VALUES = {
AgentRunStatus.completed.value,
AgentRunStatus.failed.value,
AgentRunStatus.cancelled.value,
}
_BEARER_PATTERN = re.compile(r"(?i)\bBearer\s+[^\s,;]+")
_API_KEY_PATTERN = re.compile(r"\bsk-[A-Za-z0-9_-]{8,}\b")
def sanitize_trace_value(
value: Any, *, depth: int = 0, apply_limits: bool = True
) -> Any:
"""递归净化持久化数据;可按审计用途限制体积,Secret 始终脱敏。"""
if apply_limits and depth >= MAX_TRACE_DEPTH:
return "[MAX_DEPTH]"
if isinstance(value, dict):
sanitized: dict[str, Any] = {}
for index, (key, item) in enumerate(value.items()):
if apply_limits and index >= MAX_TRACE_COLLECTION:
sanitized["__truncated__"] = True
break
normalized = str(key).casefold().replace("-", "_")
sanitized[str(key)] = (
"[REDACTED]"
if normalized in _SECRET_KEYS
or normalized.endswith(_SECRET_KEY_SUFFIXES)
else sanitize_trace_value(
item, depth=depth + 1, apply_limits=apply_limits
)
)
return sanitized
if isinstance(value, (list, tuple)):
source_items = value[:MAX_TRACE_COLLECTION] if apply_limits else value
items = [
sanitize_trace_value(
item, depth=depth + 1, apply_limits=apply_limits
)
for item in source_items
]
if apply_limits and len(value) > MAX_TRACE_COLLECTION:
items.append("[TRUNCATED]")
return items
if isinstance(value, str):
value = _BEARER_PATTERN.sub("Bearer [REDACTED]", value)
value = _API_KEY_PATTERN.sub("[REDACTED]", value)
if apply_limits and len(value) > MAX_TRACE_STRING:
return f"{value[:MAX_TRACE_STRING]}...[TRUNCATED]"
return value
if value is None or isinstance(value, (str, int, float, bool)):
return value
return sanitize_trace_value(
str(value), depth=depth + 1, apply_limits=apply_limits
)
class AgentTraceRepository:
def create_run(
self,
run: AgentRun,
request: AgentRunCreateRequest,
config_snapshot: dict[str, Any],
) -> None:
conn = connect()
try:
with transaction(conn):
conn.execute(
"""
INSERT INTO agent_runs(
run_id, status, run_json, request_json, config_snapshot_json,
created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?)
""",
(
run.run_id,
run.status.value,
self._serialize_run(run),
json.dumps(
sanitize_trace_value(request.model_dump(mode="json")),
ensure_ascii=False,
),
json.dumps(
sanitize_trace_value(config_snapshot), ensure_ascii=False
),
run.created_at.isoformat(),
run.updated_at.isoformat(),
),
)
finally:
conn.close()
def save_run(self, run: AgentRun) -> None:
conn = connect()
try:
with transaction(conn):
self._update_run(conn, run)
finally:
conn.close()
def append_event(self, run: AgentRun, event: AgentEvent) -> None:
"""在同一事务中保存最新 Run 和事件;复写同一序号时保持幂等。"""
conn = connect()
try:
with transaction(conn):
self._update_run(conn, run)
conn.execute(
"""
INSERT INTO agent_events(run_id, sequence, event, data_json, timestamp)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(run_id, sequence) DO NOTHING
""",
(
event.run_id,
event.sequence,
event.event.value,
json.dumps(event.data, ensure_ascii=False),
event.timestamp.isoformat(),
),
)
finally:
conn.close()
def get_run(self, run_id: str) -> AgentRun | None:
conn = connect()
try:
row = conn.execute(
"SELECT run_json FROM agent_runs WHERE run_id = ?", (run_id,)
).fetchone()
return AgentRun.model_validate_json(row["run_json"]) if row else None
finally:
conn.close()
def list_runs(self, limit: int, offset: int) -> tuple[list[AgentRun], int]:
conn = connect()
try:
total = int(conn.execute("SELECT COUNT(*) FROM agent_runs").fetchone()[0])
rows = conn.execute(
"""
SELECT run_json FROM agent_runs
ORDER BY created_at DESC LIMIT ? OFFSET ?
""",
(limit, offset),
).fetchall()
return [AgentRun.model_validate_json(row["run_json"]) for row in rows], total
finally:
conn.close()
def list_events(
self, run_id: str, *, after_sequence: int = -1, limit: int | None = None
) -> list[AgentEvent]:
conn = connect()
try:
sql = """
SELECT event, sequence, data_json, timestamp
FROM agent_events
WHERE run_id = ? AND sequence > ?
ORDER BY sequence
"""
params: tuple[Any, ...] = (run_id, after_sequence)
if limit is not None:
sql += " LIMIT ?"
params += (limit,)
return [self._event_from_row(run_id, row) for row in conn.execute(sql, params)]
finally:
conn.close()
def get_trace(
self, run_id: str, *, after_sequence: int, limit: int
) -> AgentTraceResponse | None:
conn = connect()
try:
row = conn.execute(
"""
SELECT run_json, config_snapshot_json
FROM agent_runs WHERE run_id = ?
""",
(run_id,),
).fetchone()
if row is None:
return None
run = AgentRun.model_validate_json(row["run_json"])
event_rows = conn.execute(
"""
SELECT event, sequence, data_json, timestamp
FROM agent_events
WHERE run_id = ? AND sequence > ?
ORDER BY sequence LIMIT ?
""",
(run_id, after_sequence, limit + 1),
).fetchall()
has_more = len(event_rows) > limit
items = [
self._event_from_row(run_id, item) for item in event_rows[:limit]
]
counts = {
item["event"]: int(item["count"])
for item in conn.execute(
"""
SELECT event, COUNT(*) AS count
FROM agent_events WHERE run_id = ? GROUP BY event
""",
(run_id,),
)
}
tool_errors = int(
conn.execute(
"""
SELECT COUNT(*) FROM agent_events
WHERE run_id = ? AND event = 'ToolResult'
AND json_extract(data_json, '$.success') = 0
""",
(run_id,),
).fetchone()[0]
)
errors = (
counts.get(AgentEventType.run_failed.value, 0)
+ counts.get(AgentEventType.model_call_failed.value, 0)
+ tool_errors
)
duration_ms = max(
0, int((run.updated_at - run.created_at).total_seconds() * 1000)
)
return AgentTraceResponse(
run_id=run_id,
status=run.status,
items=items,
next_sequence=items[-1].sequence if items else after_sequence,
has_more=has_more,
summary=AgentTraceSummary(
model_calls=counts.get(AgentEventType.model_call_started.value, 0),
tool_calls=counts.get(AgentEventType.tool_call.value, 0),
duration_ms=duration_ms,
token_usage=run.token_usage,
errors=errors,
),
config_snapshot=json.loads(row["config_snapshot_json"]),
)
finally:
conn.close()
def recover_interrupted(self, run_id: str) -> AgentRun | None:
"""把上个进程遗留的非终态 Run 收束为失败,并追加可回放终止事件。"""
conn = connect()
try:
with transaction(conn):
row = conn.execute(
"SELECT run_json, status FROM agent_runs WHERE run_id = ?", (run_id,)
).fetchone()
if row is None:
return None
run = AgentRun.model_validate_json(row["run_json"])
if row["status"] in _TERMINAL_VALUES:
return run
run.status = AgentRunStatus.failed
run.error_code = "AGENT_PROCESS_RESTARTED"
run.error_message = "Agent process restarted before the run completed."
run.updated_at = datetime.now(timezone.utc)
next_sequence = int(
conn.execute(
"""
SELECT COALESCE(MAX(sequence), -1) + 1
FROM agent_events WHERE run_id = ?
""",
(run_id,),
).fetchone()[0]
)
event = AgentEvent(
event=AgentEventType.run_failed,
run_id=run_id,
sequence=next_sequence,
data={
"code": run.error_code,
"message": run.error_message,
},
timestamp=run.updated_at,
)
self._update_run(conn, run)
conn.execute(
"""
INSERT INTO agent_events(run_id, sequence, event, data_json, timestamp)
VALUES (?, ?, ?, ?, ?)
""",
(
run_id,
next_sequence,
event.event.value,
json.dumps(event.data, ensure_ascii=False),
event.timestamp.isoformat(),
),
)
return run
finally:
conn.close()
@staticmethod
def _update_run(conn, run: AgentRun) -> None:
cursor = conn.execute(
"""
UPDATE agent_runs
SET status = ?, run_json = ?, updated_at = ?
WHERE run_id = ?
""",
(
run.status.value,
AgentTraceRepository._serialize_run(run),
run.updated_at.isoformat(),
run.run_id,
),
)
if cursor.rowcount != 1:
raise LookupError(run.run_id)
@staticmethod
def _event_from_row(run_id: str, row) -> AgentEvent:
return AgentEvent(
event=AgentEventType(row["event"]),
run_id=run_id,
sequence=int(row["sequence"]),
data=json.loads(row["data_json"]),
timestamp=datetime.fromisoformat(row["timestamp"]),
)
@staticmethod
def _serialize_run(run: AgentRun) -> str:
# Run 是重启后 GET/list 的完整事实;只做 Secret 脱敏,不套用 Trace 摘要限长。
return json.dumps(
sanitize_trace_value(
run.model_dump(mode="json"), apply_limits=False
),
ensure_ascii=False,
)
+8
View File
@@ -0,0 +1,8 @@
"""Benchmark 服务:RAG / Agent 数据集注册、指标计算与运行管理。
模块划分:
- metrics.py 纯函数指标(Hit@K / Recall@K / MRR / CitationHit / 分位数)
- datasets.py 受控目录的 Dataset 注册与校验
- rag.py RAG Benchmark Runner(调用 retrieval.engine.search
- service.py 运行注册表、配置快照与报告组装
"""
+198
View File
@@ -0,0 +1,198 @@
"""Benchmark Dataset 注册:从受控目录加载 JSON 数据集并校验。
Dataset 只能来自配置目录(settings.benchmark_datasets_path),API 不接受调用方提交
任意文件路径。目录不存在或为空时按「无数据集」处理,不报错。
"""
from __future__ import annotations
import hashlib
import json
from dataclasses import dataclass, field
from pathlib import Path
from pydantic import BaseModel, Field, ValidationError
from app.config import get_settings
from app.contracts import (
BenchmarkDatasetInfo,
BenchmarkKind,
RAGDatasetCase,
)
from app.errors import ApiError
@dataclass
class RAGDataset:
"""内存中的 RAG 数据集:元信息 + 已校验的 Case 列表 + 内容哈希。"""
dataset_id: str
kind: BenchmarkKind
version: str
description: str
cases: list[RAGDatasetCase] = field(default_factory=list)
content_hash: str = ""
class _DatasetMeta(BaseModel):
"""Dataset 元数据的最小校验模型。
list_datasets 用它逐文件校验元信息字段结构,把「合法 JSON 但字段类型错误」
(如 cases: 42)这类损坏文件隔离掉,而不是让 len() 抛 TypeError 拖垮整个列表。
"""
dataset_id: str = Field(min_length=1)
kind: str = ""
version: str = ""
description: str = ""
cases: list = Field(default_factory=list)
def _datasets_dir() -> Path:
return get_settings().benchmark_datasets_path
def _dataset_files() -> list[Path]:
directory = _datasets_dir()
if not directory.is_dir():
return []
return sorted(directory.glob("*.json"))
def _content_hash(raw: bytes) -> str:
return "sha256:" + hashlib.sha256(raw).hexdigest()
def _read_json(path: Path) -> tuple[dict, bytes]:
"""读取并解析 JSON 文件,返回 (dict, 原始字节);非法 JSON 抛 BENCHMARK_DATASET_INVALID。"""
try:
raw_bytes = path.read_bytes()
return json.loads(raw_bytes.decode("utf-8")), raw_bytes
except (json.JSONDecodeError, OSError, UnicodeDecodeError) as exc:
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
f"Dataset file is not valid JSON: {path.name}",
{"path": str(path)},
) from exc
def _dataset_from_raw(raw: dict, raw_bytes: bytes, kind: BenchmarkKind) -> RAGDataset:
"""把单个数据集 JSON 解析为 RAGDataset,非法结构抛 BENCHMARK_DATASET_INVALID。"""
dataset_id = raw.get("dataset_id")
if not isinstance(dataset_id, str) or not dataset_id:
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
"Dataset must declare a non-empty string 'dataset_id'.",
{},
)
file_kind = raw.get("kind", kind.value)
if file_kind != kind.value:
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
f"Dataset kind mismatch: expected '{kind.value}', got '{file_kind}'.",
{"dataset_id": dataset_id},
)
raw_cases = raw.get("cases")
if not isinstance(raw_cases, list) or not raw_cases:
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
"Dataset 'cases' must be a non-empty list.",
{"dataset_id": dataset_id},
)
cases: list[RAGDatasetCase] = []
for index, case in enumerate(raw_cases):
try:
parsed = RAGDatasetCase.model_validate(case)
except ValidationError as exc:
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
f"Dataset case #{index} is invalid.",
{"dataset_id": dataset_id, "case_index": index, "errors": exc.errors()},
) from exc
# 每个 Case 至少要声明一个期望 id,否则无法计算命中/召回
if not parsed.expected_note_ids and not parsed.expected_block_ids:
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
f"Dataset case '{parsed.case_id}' must declare expected_note_ids or expected_block_ids.",
{"dataset_id": dataset_id, "case_id": parsed.case_id},
)
# citation_required=true 时必须声明 expected_block_ids,否则无法计算 Citation Hit Rate
if parsed.citation_required and not parsed.expected_block_ids:
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
f"Dataset case '{parsed.case_id}' requires expected_block_ids when citation_required is true.",
{"dataset_id": dataset_id, "case_id": parsed.case_id},
)
cases.append(parsed)
return RAGDataset(
dataset_id=dataset_id,
kind=kind,
version=str(raw.get("version", "")),
description=str(raw.get("description", "")),
cases=cases,
content_hash=_content_hash(raw_bytes),
)
def list_datasets(kind: BenchmarkKind) -> list[BenchmarkDatasetInfo]:
"""枚举受控目录下指定 kind 的数据集元信息(不含 Case 内容)。
逐文件用 _DatasetMeta 校验元信息字段结构,单个损坏文件隔离跳过而非整体失败,
保证列表接口健壮;损坏细节由 load_dataset 抛出。
"""
infos: list[BenchmarkDatasetInfo] = []
for path in _dataset_files():
try:
raw, raw_bytes = _read_json(path)
meta = _DatasetMeta.model_validate(raw)
except (ApiError, ValidationError):
continue
if meta.kind not in ("", kind.value):
continue
infos.append(
BenchmarkDatasetInfo(
dataset_id=meta.dataset_id,
kind=kind,
version=meta.version,
description=meta.description,
case_count=len(meta.cases),
content_hash=_content_hash(raw_bytes),
)
)
return infos
def load_dataset(dataset_id: str, kind: BenchmarkKind) -> RAGDataset:
"""按文件名加载并校验数据集;找不到抛 BENCHMARK_DATASET_NOT_FOUND。
只读取与请求 dataset_id 同名的文件({dataset_id}.json),无关文件的损坏(JSON 语法
错误、UTF-8 解码错误、顶层非对象)不会阻断目标数据集加载;只有目标文件本身损坏
才抛 BENCHMARK_DATASET_INVALID。按现有文件 stem 精确匹配,不拼接调用方传入的路径。
"""
for path in _dataset_files():
if path.stem != dataset_id:
continue
raw, raw_bytes = _read_json(path)
if not isinstance(raw, dict):
raise ApiError(
422,
"BENCHMARK_DATASET_INVALID",
"Dataset top-level must be a JSON object.",
{"dataset_id": dataset_id, "path": path.name},
)
return _dataset_from_raw(raw, raw_bytes, kind)
raise ApiError(
404,
"BENCHMARK_DATASET_NOT_FOUND",
f"Benchmark dataset does not exist: {dataset_id}",
{"dataset_id": dataset_id, "kind": kind.value},
)
+58
View File
@@ -0,0 +1,58 @@
"""Benchmark 指标纯函数。
所有指标只依赖「按相关性降序的 retrieved id 列表」和「期望 id 集合」,不接触任何
外部状态,便于单元测试与未来 Agent Benchmark 复用。retrieved 顺序越靠前越相关。
"""
from __future__ import annotations
def hit_at_k(retrieved: list[str], expected: set[str], k: int) -> bool:
"""前 k 个结果里是否命中任意期望 id(用于 Hit@1 / Hit@5)。"""
return any(item in expected for item in retrieved[:k])
def recall_at_k(retrieved: list[str], expected: set[str], k: int) -> float:
"""前 k 个结果召回的期望 id 占比;期望为空时视为 0。
结果先去重:检索结果是 Block 级,同一 Note 可能经多个 Block 重复出现,
直接逐项计数会把同一 Note 算多次、导致 Recall 超过 1。
"""
if not expected:
return 0.0
return len(set(retrieved[:k]) & expected) / len(expected)
def reciprocal_rank(retrieved: list[str], expected: set[str]) -> float:
"""首个命中的倒数排名;未命中返回 0。rank 从 1 开始。"""
for rank, item in enumerate(retrieved, start=1):
if item in expected:
return 1.0 / rank
return 0.0
def citation_hit(retrieved_block_ids: list[str], expected: set[str]) -> bool:
"""首条结果的 block_id 是否为期望引用块(Citation Hit Rate 的逐 Case 判据)。"""
if not retrieved_block_ids or not expected:
return False
return retrieved_block_ids[0] in expected
def mean(values: list[float]) -> float:
return sum(values) / len(values) if values else 0.0
def percentile(values: list[float], p: float) -> float:
"""线性插值分位数(p ∈ [0, 100]),用于 P50 / P95 延迟。空列表返回 0。"""
if not values:
return 0.0
ordered = sorted(values)
if len(ordered) == 1:
return ordered[0]
rank = (len(ordered) - 1) * (p / 100.0)
lo = int(rank)
hi = lo + 1
if hi >= len(ordered):
return ordered[-1]
frac = rank - lo
return ordered[lo] + (ordered[hi] - ordered[lo]) * frac
+158
View File
@@ -0,0 +1,158 @@
"""RAG Benchmark Runner:调用检索引擎对数据集逐 Case 求值并聚合指标。
只读操作,直接复用 app.retrieval.engine 的 search(),不旁路检索链路。指标按
(mode, case, repeat) 逐样本计算,再按 mode 聚合;失败样本按零分计入质量指标分母,
避免把执行失败误判为检索质量(同时保留 total/successful/failed/failure_rate)。
"""
from __future__ import annotations
import asyncio
import logging
import time
from collections.abc import Callable
from app import repository
from app.benchmarks import metrics as m
from app.benchmarks.datasets import RAGDataset
from app.contracts import (
RAGCaseResult,
RAGDatasetCase,
RAGMetrics,
RAGRunRequest,
SearchMode,
SearchRequest,
)
from app.retrieval.engine import engine
logger = logging.getLogger(__name__)
class BenchmarkCancelled(Exception):
"""运行在 Case 之间被取消时抛出,用于中断后台执行并标记 cancelled。"""
async def run_rag(
dataset: RAGDataset,
request: RAGRunRequest,
on_case: Callable[[RAGCaseResult, int, int], None] | None = None,
should_cancel: Callable[[], bool] | None = None,
) -> tuple[dict[str, RAGMetrics], list[RAGCaseResult]]:
"""执行 RAG Benchmark,返回 (按 mode 聚合的指标, 全部逐样本结果)。
on_case 在每个样本求值完成后回调 (result, done, total),供上层更新进度与事件。
should_cancel 在每个样本开始前被检查;返回 True 时抛出 BenchmarkCancelled 中断运行。
"""
total = len(request.modes) * len(dataset.cases) * request.repeat
done = 0
results: list[RAGCaseResult] = []
for mode in request.modes:
for case in dataset.cases:
expected_notes = _expected_notes(case)
for repeat in range(request.repeat):
# 让出事件循环:使运行中取消、SSE 进度与并发 API 请求能及时得到调度
await asyncio.sleep(0)
if should_cancel is not None and should_cancel():
raise BenchmarkCancelled()
result = await _evaluate_one(case, mode, request, repeat, expected_notes)
results.append(result)
done += 1
if on_case is not None:
on_case(result, done, total)
metrics_by_mode = {mode.value: _aggregate(results, mode) for mode in request.modes}
return metrics_by_mode, results
def _expected_notes(case: RAGDatasetCase) -> set[str]:
"""返回笔记级期望 id;仅标注块 ID 时从块反查所属笔记,避免把标注缺失误判为检索失败。"""
if case.expected_note_ids:
return set(case.expected_note_ids)
return {hit.note_id for hit in repository.get_block_hits(case.expected_block_ids)}
async def _evaluate_one(
case: RAGDatasetCase,
mode: SearchMode,
request: RAGRunRequest,
repeat: int,
expected_notes: set[str],
) -> RAGCaseResult:
search_request = SearchRequest(
query=case.query,
mode=mode,
limit=request.retrieval.top_k,
include_snippet=False,
rrf_k=request.retrieval.rrf_k,
rerank=request.retrieval.rerank,
rerank_candidates=request.retrieval.rerank_candidates,
score_threshold=request.retrieval.score_threshold,
)
start = time.perf_counter()
try:
response = await engine.search(search_request)
latency_ms = (time.perf_counter() - start) * 1000.0
except Exception as exc: # 单个样本失败不中断整个 Benchmark
# 详细异常只进日志,公开响应只带项目错误码与安全消息,避免泄露路径/SQL 等敏感信息
logger.warning(
"RAG case evaluation failed: case=%s mode=%s", case.case_id, mode.value,
exc_info=exc,
)
return RAGCaseResult(
case_id=case.case_id,
mode=mode,
repeat=repeat,
latency_ms=(time.perf_counter() - start) * 1000.0,
citation_applicable=case.citation_required,
error="RAG case evaluation failed.",
error_code="BENCHMARK_CASE_EVALUATION_FAILED",
)
retrieved_note_ids = [item.note_id for item in response.items]
retrieved_block_ids = [item.block_id for item in response.items]
expected_blocks = set(case.expected_block_ids)
k = request.retrieval.top_k
return RAGCaseResult(
case_id=case.case_id,
mode=mode,
repeat=repeat,
latency_ms=latency_ms,
retrieved_note_ids=retrieved_note_ids,
retrieved_block_ids=retrieved_block_ids,
hit_at_1=m.hit_at_k(retrieved_note_ids, expected_notes, 1),
hit_at_5=m.hit_at_k(retrieved_note_ids, expected_notes, 5),
recall=m.recall_at_k(retrieved_note_ids, expected_notes, k),
reciprocal_rank=m.reciprocal_rank(retrieved_note_ids, expected_notes),
citation_hit=m.citation_hit(retrieved_block_ids, expected_blocks),
citation_applicable=case.citation_required,
)
def _aggregate(cases: list[RAGCaseResult], mode: SearchMode) -> RAGMetrics:
samples = [c for c in cases if c.mode == mode]
total = len(samples)
failed = sum(1 for c in samples if c.error is not None)
successful = total - failed
if total == 0:
return RAGMetrics()
# 延迟只统计成功样本;失败样本按零分计入质量指标分母,避免汇总虚高
latencies = [c.latency_ms for c in samples if c.error is None]
citation_samples = [c for c in samples if c.citation_applicable]
return RAGMetrics(
hit_at_1=m.mean([1.0 if (c.error is None and c.hit_at_1) else 0.0 for c in samples]),
hit_at_5=m.mean([1.0 if (c.error is None and c.hit_at_5) else 0.0 for c in samples]),
recall_at_k=m.mean([c.recall if c.error is None else 0.0 for c in samples]),
mrr=m.mean([c.reciprocal_rank if c.error is None else 0.0 for c in samples]),
citation_hit_rate=m.mean(
[1.0 if (c.error is None and c.citation_hit) else 0.0 for c in citation_samples]
),
p50_latency_ms=m.percentile(latencies, 50.0),
p95_latency_ms=m.percentile(latencies, 95.0),
total_cases=total,
successful_cases=successful,
failed_cases=failed,
failure_rate=failed / total,
)
+348
View File
@@ -0,0 +1,348 @@
"""Benchmark 服务:运行注册表、配置快照与报告组装。
RAG Benchmark 采用「创建即返回 queued、后台 Task 异步执行」的模式(与 index_service
的 rebuild 一致):POST 创建后立即返回 202 queued 的 BenchmarkRun,由受管 asyncio.Task
在后台逐 Case 求值,进度与事件实时写入内存注册表,供 SSE 订阅。运行记录、事件与报告
暂存内存(_runs/_events/_reports),不持久化到 SQLite;后续接入异步任务队列时再落库。
"""
from __future__ import annotations
import asyncio
import logging
import sys
from datetime import datetime, timezone
from uuid import uuid4
from app import repository
from app.benchmarks import datasets
from app.benchmarks.datasets import RAGDataset
from app.benchmarks.rag import BenchmarkCancelled, run_rag
from app.config import get_settings
from app.contracts import (
BenchmarkEvent,
BenchmarkEventType,
BenchmarkKind,
BenchmarkReport,
BenchmarkRun,
BenchmarkStatus,
RAGCaseResult,
RAGMetrics,
RAGRunRequest,
SearchMode,
)
from app.errors import ApiError
from app.retrieval.engine import engine
logger = logging.getLogger(__name__)
_runs: dict[str, BenchmarkRun] = {}
_events: dict[str, list[BenchmarkEvent]] = {}
_reports: dict[str, BenchmarkReport] = {}
_tasks: dict[str, asyncio.Task] = {}
_subscribers: dict[str, list[asyncio.Queue[BenchmarkEvent]]] = {}
_cancel_flags: dict[str, asyncio.Event] = {}
MAX_RUNS = 100
def _now() -> datetime:
return datetime.now(timezone.utc)
def _forget(run_id: str) -> None:
"""移除一条 run 的全部内存态;仅在 run 处于终态时调用,避免打断活动任务。"""
_runs.pop(run_id, None)
_events.pop(run_id, None)
_reports.pop(run_id, None)
_tasks.pop(run_id, None)
_subscribers.pop(run_id, None)
_cancel_flags.pop(run_id, None)
def _evict_terminal() -> bool:
"""超过容量时淘汰最旧的终态 run;全部为活动 run 无法淘汰时返回 False。
绝不能删除仍在运行(queued/running)的 run:那会连带移除其 _cancel_flags 与
_subscribers,使后台 Task 访问时抛出 KeyError。
"""
terminal = (BenchmarkStatus.completed, BenchmarkStatus.failed, BenchmarkStatus.cancelled)
while len(_runs) >= MAX_RUNS:
victim = next(
(rid for rid, run in _runs.items() if run.status in terminal), None
)
if victim is None:
return False
_forget(victim)
return True
def _config_snapshot(request: RAGRunRequest, dataset: RAGDataset) -> dict:
"""记录运行时的模型 / 索引 / 环境信息,保证报告可解释、可复现。"""
settings = get_settings()
return {
"dataset_id": dataset.dataset_id,
"dataset_hash": dataset.content_hash,
"dataset_version": dataset.version,
"modes": [m.value for m in request.modes],
"retrieval": request.retrieval.model_dump(),
"repeat": request.repeat,
"embedding": {
"model_id": engine.embedding.model_id,
"version": engine.embedding.version,
"dim": engine.embedding.dim,
},
"reranker": {
"model_id": engine.reranker.model_id,
"version": engine.reranker.version,
},
"index_meta": repository.get_index_meta(),
"app": {"version": settings.version, "environment": settings.environment},
"python": sys.version.split()[0],
"metadata": request.metadata,
}
async def _validate_index_compatibility(request: RAGRunRequest) -> None:
"""创建 RAG Run 前校验索引已建立且与当前 Embedding 模型/维度兼容。
空索引或不兼容索引会让所有模式得到全 0 指标,把环境/索引错误误判为检索质量差,
故在创建时即拒绝,返回 BENCHMARK_INDEX_INCOMPATIBLE。
"""
stats = repository.stats()
meta = repository.get_index_meta()
needs_vector = any(m in (SearchMode.vector, SearchMode.hybrid) for m in request.modes)
reasons: list[str] = []
if stats["blocks"] == 0:
reasons.append("index is empty (no indexed blocks; run /api/index/rebuild first)")
if needs_vector:
if meta.get("embedding_model") != engine.embedding.model_id:
reasons.append(
f"embedding model mismatch: index={meta.get('embedding_model')!r}, "
f"engine={engine.embedding.model_id!r}"
)
if meta.get("embedding_dim") != str(engine.embedding.dim):
reasons.append(
f"embedding dimension mismatch: index={meta.get('embedding_dim')!r}, "
f"engine={engine.embedding.dim}"
)
if await engine.vector_store.count() == 0:
reasons.append("vector index is empty")
if reasons:
raise ApiError(
409,
"BENCHMARK_INDEX_INCOMPATIBLE",
"Benchmark index is not built or is incompatible with the current retrieval engine.",
{"reasons": reasons},
)
async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun:
"""创建一次 RAG Benchmark,立即返回 queued 的 BenchmarkRun,由后台 Task 执行。"""
dataset = datasets.load_dataset(request.dataset_id, BenchmarkKind.rag)
await _validate_index_compatibility(request)
# 容量检查:先淘汰终态 run 腾空间;满容量且全为活动 run 时拒绝创建
if not _evict_terminal():
raise ApiError(
429,
"BENCHMARK_CAPACITY_EXCEEDED",
"Benchmark run capacity exceeded; wait for active runs to finish.",
{},
)
run_id = "benchmark_" + uuid4().hex[:12]
snapshot = _config_snapshot(request, dataset)
run = BenchmarkRun(
run_id=run_id,
kind=BenchmarkKind.rag,
dataset_id=dataset.dataset_id,
dataset_hash=dataset.content_hash,
status=BenchmarkStatus.queued,
progress=0.0,
config_snapshot=snapshot,
created_at=_now(),
)
_runs[run_id] = run
_events[run_id] = []
_subscribers[run_id] = []
_cancel_flags[run_id] = asyncio.Event()
_tasks[run_id] = asyncio.create_task(_execute_rag(run_id, request, dataset, snapshot))
return run
async def _execute_rag(
run_id: str, request: RAGRunRequest, dataset: RAGDataset, snapshot: dict
) -> None:
"""后台执行 RAG Benchmark,实时更新进度/事件,结束后写入报告并关闭订阅。"""
cancel_event = _cancel_flags[run_id]
def emit(event_type: BenchmarkEventType, data: dict) -> None:
sequence = len(_events[run_id])
event = BenchmarkEvent(
event=event_type, run_id=run_id, sequence=sequence, data=data, timestamp=_now()
)
_events[run_id].append(event)
for queue in _subscribers.get(run_id, []):
queue.put_nowait(event)
def finish() -> None:
_subscribers.pop(run_id, None)
_cancel_flags.pop(run_id, None)
_runs[run_id] = _runs[run_id].model_copy(
update={"status": BenchmarkStatus.running, "started_at": _now()}
)
emit(
BenchmarkEventType.run_started,
{"dataset_id": dataset.dataset_id, "modes": [m.value for m in request.modes]},
)
total = len(request.modes) * len(dataset.cases) * request.repeat
def on_case(result: RAGCaseResult, done: int, _total: int) -> None:
progress = done / total if total else 1.0
_runs[run_id] = _runs[run_id].model_copy(update={"progress": progress})
emit(BenchmarkEventType.case_completed, result.model_dump(mode="json"))
try:
metrics_by_mode, results = await run_rag(
dataset,
request,
on_case=on_case,
should_cancel=cancel_event.is_set,
)
except BenchmarkCancelled:
_runs[run_id] = _runs[run_id].model_copy(
update={
"status": BenchmarkStatus.cancelled,
"progress": 1.0,
"completed_at": _now(),
}
)
emit(BenchmarkEventType.run_cancelled, {"status": BenchmarkStatus.cancelled.value})
_reports[run_id] = BenchmarkReport(
run_id=run_id,
kind=BenchmarkKind.rag,
dataset_id=dataset.dataset_id,
dataset_hash=dataset.content_hash,
status=BenchmarkStatus.cancelled,
config_snapshot=snapshot,
)
finish()
return
except Exception as exc: # 单次运行失败不拖垮服务,记录错误后结束
# 详细异常只进日志,公开响应仅带项目错误码与安全消息,避免泄露路径/SQL 等敏感信息
logger.exception("Benchmark run failed: run_id=%s", run_id)
_runs[run_id] = _runs[run_id].model_copy(
update={
"status": BenchmarkStatus.failed,
"progress": 1.0,
"error": "Benchmark run failed.",
"error_code": "BENCHMARK_RUN_FAILED",
"completed_at": _now(),
}
)
emit(
BenchmarkEventType.run_failed,
{"error": "Benchmark run failed.", "error_code": "BENCHMARK_RUN_FAILED"},
)
_reports[run_id] = BenchmarkReport(
run_id=run_id,
kind=BenchmarkKind.rag,
dataset_id=dataset.dataset_id,
dataset_hash=dataset.content_hash,
status=BenchmarkStatus.failed,
config_snapshot=snapshot,
error="Benchmark run failed.",
error_code="BENCHMARK_RUN_FAILED",
)
finish()
return
metrics = {mode: m.model_dump() for mode, m in metrics_by_mode.items()}
_runs[run_id] = _runs[run_id].model_copy(
update={
"status": BenchmarkStatus.completed,
"progress": 1.0,
"metrics": metrics,
"completed_at": _now(),
}
)
emit(BenchmarkEventType.run_completed, {"metrics": metrics})
_reports[run_id] = BenchmarkReport(
run_id=run_id,
kind=BenchmarkKind.rag,
dataset_id=dataset.dataset_id,
dataset_hash=dataset.content_hash,
status=BenchmarkStatus.completed,
config_snapshot=snapshot,
metrics=metrics,
cases=results,
)
finish()
def list_runs(
kind: BenchmarkKind | None = None,
status: BenchmarkStatus | None = None,
limit: int = 50,
offset: int = 0,
) -> tuple[list[BenchmarkRun], int]:
runs = list(_runs.values())
if kind is not None:
runs = [r for r in runs if r.kind == kind]
if status is not None:
runs = [r for r in runs if r.status == status]
runs.sort(key=lambda r: r.created_at, reverse=True)
total = len(runs)
return runs[offset : offset + limit], total
def get_run(run_id: str) -> BenchmarkRun | None:
return _runs.get(run_id)
def get_report(run_id: str) -> BenchmarkReport | None:
return _reports.get(run_id)
def get_events(run_id: str) -> list[BenchmarkEvent]:
return _events.get(run_id, [])
def cancel_run(run_id: str) -> BenchmarkRun | None:
"""取消运行:对 queued/running 设置取消标志,后台 Task 在 Case 边界检查后置为 cancelled。"""
run = _runs.get(run_id)
if run is None:
return None
if run.status in (BenchmarkStatus.queued, BenchmarkStatus.running):
_cancel_flags[run_id].set()
return run
def subscribe(run_id: str) -> asyncio.Queue[BenchmarkEvent] | None:
"""订阅运行事件流;运行已结束(completed/failed/cancelled)时返回 None。"""
run = _runs.get(run_id)
if run is None or run.status in (
BenchmarkStatus.completed,
BenchmarkStatus.failed,
BenchmarkStatus.cancelled,
):
return None
queue: asyncio.Queue[BenchmarkEvent] = asyncio.Queue()
_subscribers.setdefault(run_id, []).append(queue)
return queue
def unsubscribe(run_id: str, queue: asyncio.Queue[BenchmarkEvent]) -> None:
subscribers = _subscribers.get(run_id)
if subscribers and queue in subscribers:
subscribers.remove(queue)
async def wait_for_run(run_id: str) -> BenchmarkRun:
"""等待后台任务结束(测试/轮询用);无任务时直接返回当前状态。"""
task = _tasks.get(run_id)
if task is not None:
await task
return _runs.get(run_id)
+4
View File
@@ -24,6 +24,7 @@ class Settings:
db_path: Path db_path: Path
vault_path: Path vault_path: Path
attachments_path: Path attachments_path: Path
benchmark_datasets_path: Path
@lru_cache @lru_cache
@@ -41,4 +42,7 @@ def get_settings() -> Settings:
attachments_path=Path( attachments_path=Path(
os.getenv("APP_ATTACHMENTS_PATH", str(data_dir / "attachments")) os.getenv("APP_ATTACHMENTS_PATH", str(data_dir / "attachments"))
), ),
benchmark_datasets_path=Path(
os.getenv("APP_BENCHMARK_DATASETS_PATH", str(data_dir / "benchmarks"))
),
) )
+20 -2
View File
@@ -3,8 +3,9 @@ from dataclasses import dataclass
from app.agent import AgentRuntime, PermissionManager, PermissionPolicy, ToolRegistry from app.agent import AgentRuntime, PermissionManager, PermissionPolicy, ToolRegistry
from app.agent.builtin_tools import register_builtin_tools from app.agent.builtin_tools import register_builtin_tools
from app.contracts import ModelCapability, ProviderConfig, ProviderType from app.contracts import ModelCapability, ProviderConfig, ProviderType
from app.config import BACKEND_DIR from app.config import BACKEND_DIR, get_settings
from app.extensions import PluginRuntime, SkillRuntime from app.extensions import PluginRuntime, SkillRuntime
from app.extensions.mcp_registry import McpServerRegistry
from app.providers import MockProvider, ProviderFactory, ProviderRegistry from app.providers import MockProvider, ProviderFactory, ProviderRegistry
from app.providers.credentials import ( from app.providers.credentials import (
ChainedCredentialResolver, ChainedCredentialResolver,
@@ -22,10 +23,12 @@ class ApplicationContainer:
permissions: PermissionManager permissions: PermissionManager
skills: SkillRuntime skills: SkillRuntime
plugins: PluginRuntime plugins: PluginRuntime
mcp_servers: McpServerRegistry
agent: AgentRuntime agent: AgentRuntime
def build_container() -> ApplicationContainer: def build_container() -> ApplicationContainer:
settings = get_settings()
credentials = EncryptedCredentialStore() credentials = EncryptedCredentialStore()
provider_factory = ProviderFactory( provider_factory = ProviderFactory(
ChainedCredentialResolver(credentials, EnvironmentCredentialResolver()) ChainedCredentialResolver(credentials, EnvironmentCredentialResolver())
@@ -50,10 +53,24 @@ def build_container() -> ApplicationContainer:
tools = ToolRegistry() tools = ToolRegistry()
register_builtin_tools(tools) register_builtin_tools(tools)
plugins = PluginRuntime(tools) plugins = PluginRuntime(
tools,
credentials=credentials,
# 当前 Python Host 尚无 OS 沙箱。生产构建必须保持关闭,直到
# Tauri/Rust Host 能签发绑定命令摘要的可信启动许可。
allow_unsandboxed_mcp=settings.environment == "development",
)
plugins.install(BACKEND_DIR / "extensions" / "plugins" / "text-tools") plugins.install(BACKEND_DIR / "extensions" / "plugins" / "text-tools")
plugins.enable("text-tools") plugins.enable("text-tools")
mcp_servers = McpServerRegistry(
tools,
credentials,
settings.data_dir,
allow_process_launch=settings.environment == "development",
)
mcp_servers.restore_enabled()
skills = SkillRuntime(tools) skills = SkillRuntime(tools)
skills.install(BACKEND_DIR / "extensions" / "skills" / "knowledge-assistant") skills.install(BACKEND_DIR / "extensions" / "skills" / "knowledge-assistant")
skills.enable("knowledge-assistant") skills.enable("knowledge-assistant")
@@ -74,6 +91,7 @@ def build_container() -> ApplicationContainer:
permissions=permissions, permissions=permissions,
skills=skills, skills=skills,
plugins=plugins, plugins=plugins,
mcp_servers=mcp_servers,
agent=agent, agent=agent,
) )
+508 -3
View File
@@ -1,8 +1,8 @@
from datetime import datetime from datetime import datetime
from enum import Enum from enum import Enum
from typing import Any, Literal from typing import Annotated, Any, Literal
from pydantic import BaseModel, ConfigDict, Field, SecretStr from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
class Contract(BaseModel): class Contract(BaseModel):
@@ -31,6 +31,48 @@ class OperationResponse(Contract):
message: str | None = None message: str | None = None
# Workspace boundary (single configured Vault in Web development mode)
class WorkspaceInfo(Contract):
vault_id: str = "default"
name: str
path: str
file_count: int = 0
indexed_note_count: int = 0
requires_refresh: bool = False
class WorkspaceEntry(Contract):
entry_id: str
name: str
path: str
type: Literal["file", "folder"]
note_id: str | None = None
children: list["WorkspaceEntry"] = Field(default_factory=list)
class WorkspaceSnapshot(Contract):
workspace: WorkspaceInfo
items: list[WorkspaceEntry] = Field(default_factory=list)
class WorkspaceOpenRequest(Contract):
path: str | None = None
class FolderCreateRequest(Contract):
parent: str = ""
name: str = Field(min_length=1)
class FolderRenameRequest(Contract):
path: str
new_name: str = Field(min_length=1)
class FolderDeleteRequest(Contract):
path: str
# Notes and retrieval # Notes and retrieval
class NoteBlock(Contract): class NoteBlock(Contract):
block_id: str block_id: str
@@ -79,6 +121,10 @@ class NoteMoveRequest(Contract):
folder: str folder: str
class NoteRenameRequest(Contract):
file_name: str = Field(min_length=1)
class SearchMode(str, Enum): class SearchMode(str, Enum):
fts = "fts" fts = "fts"
vector = "vector" vector = "vector"
@@ -98,6 +144,12 @@ class SearchRequest(Contract):
limit: int = Field(default=20, ge=1, le=100) limit: int = Field(default=20, ge=1, le=100)
offset: int = Field(default=0, ge=0) offset: int = Field(default=0, ge=0)
include_snippet: bool = True include_snippet: bool = True
# 检索调优参数(Benchmark 与 Skill 共用):控制 RRF / 精排 / 候选池 / 分数阈值。
# rerank_candidates=None 表示对全部候选精排(保留原有行为),Benchmark 传显式值。
rrf_k: int = Field(default=60, ge=1)
rerank: bool = True
rerank_candidates: int | None = Field(default=None, ge=1)
score_threshold: float = Field(default=0.0, ge=0.0)
class Citation(Contract): class Citation(Contract):
@@ -153,7 +205,7 @@ class ToolDefinition(Contract):
description: str description: str
parameters: dict[str, Any] = Field(default_factory=dict) parameters: dict[str, Any] = Field(default_factory=dict)
permission: str | None = None permission: str | None = None
source: Literal["builtin", "plugin"] = "builtin" source: Literal["builtin", "plugin", "mcp_server"] = "builtin"
class ToolCall(Contract): class ToolCall(Contract):
@@ -283,6 +335,10 @@ class AgentEventType(str, Enum):
permission_required = "PermissionRequired" permission_required = "PermissionRequired"
usage = "Usage" usage = "Usage"
citation = "Citation" citation = "Citation"
model_call_started = "ModelCallStarted"
model_call_completed = "ModelCallCompleted"
model_call_failed = "ModelCallFailed"
permission_resolved = "PermissionResolved"
run_completed = "RunCompleted" run_completed = "RunCompleted"
run_failed = "RunFailed" run_failed = "RunFailed"
run_cancelled = "RunCancelled" run_cancelled = "RunCancelled"
@@ -296,6 +352,24 @@ class AgentEvent(Contract):
timestamp: datetime timestamp: datetime
class AgentTraceSummary(Contract):
model_calls: int = 0
tool_calls: int = 0
duration_ms: int = 0
token_usage: int = 0
errors: int = 0
class AgentTraceResponse(Contract):
run_id: str
status: AgentRunStatus
items: list[AgentEvent] = Field(default_factory=list)
next_sequence: int
has_more: bool = False
summary: AgentTraceSummary = Field(default_factory=AgentTraceSummary)
config_snapshot: dict[str, Any] = Field(default_factory=dict)
class PermissionDecisionRequest(Contract): class PermissionDecisionRequest(Contract):
decision: Literal["allow_once", "allow_session", "deny"] decision: Literal["allow_once", "allow_session", "deny"]
@@ -349,6 +423,10 @@ class ExtensionInstallRequest(Contract):
class PluginBackend(Contract): class PluginBackend(Contract):
type: Literal["mcp", "internal_rpc", "none"] = "none" type: Literal["mcp", "internal_rpc", "none"] = "none"
transport: Literal["stdio", "http", "none"] = "none" transport: Literal["stdio", "http", "none"] = "none"
command: str | None = None
args: list[str] = Field(default_factory=list)
startup_timeout_seconds: int = Field(default=10, ge=1, le=60)
tool_timeout_seconds: int = Field(default=30, ge=1, le=600)
class PluginContribution(Contract): class PluginContribution(Contract):
@@ -392,6 +470,285 @@ class PluginListResponse(Contract):
items: list[Plugin] = Field(default_factory=list) items: list[Plugin] = Field(default_factory=list)
class PluginHostState(str, Enum):
stopped = "stopped"
starting = "starting"
ready = "ready"
unhealthy = "unhealthy"
error = "error"
class PluginHostStatus(Contract):
plugin_id: str
backend_type: Literal["mcp", "internal_rpc", "none"]
transport: Literal["stdio", "http", "none"]
status: PluginHostState
tools_count: int = 0
started_at: datetime | None = None
last_seen_at: datetime | None = None
protocol_version: str | None = None
server_name: str | None = None
server_version: str | None = None
error: str | None = None
# Independent user-managed MCP Server Registry. This is deliberately separate
# from Plugin manifests: a server can contribute tools without being a Plugin.
class McpServerTransport(str, Enum):
stdio = "stdio"
streamable_http = "streamable_http"
sse = "sse"
class McpServerConfig(Contract):
name: str = Field(min_length=1, max_length=80)
transport: McpServerTransport = McpServerTransport.stdio
command: str | None = Field(default=None, max_length=1024)
args: list[str] = Field(default_factory=list, max_length=64)
url: str | None = Field(default=None, max_length=4096)
headers: dict[str, str] = Field(default_factory=dict)
environment: dict[str, str] = Field(default_factory=dict)
secret_environment_keys: list[str] = Field(default_factory=list)
secret_header_keys: list[str] = Field(default_factory=list)
permissions: list[str] = Field(default_factory=list)
startup_timeout_seconds: float = Field(default=15, ge=1, le=120)
tool_timeout_seconds: float = Field(default=30, ge=1, le=300)
class McpServerCreateRequest(McpServerConfig):
pass
class McpServerUpdateRequest(McpServerConfig):
version: int = Field(ge=1)
class McpServerSecretWriteRequest(Contract):
secret: SecretStr = Field(min_length=1, max_length=32768)
class McpServerSecretStatus(Contract):
key: str
configured: bool
class McpServerTrustRequest(Contract):
command_digest: str = Field(min_length=64, max_length=64)
class McpServerStatus(Contract):
enabled: bool = False
status: PluginHostState = PluginHostState.stopped
tools_count: int = 0
protocol_version: str | None = None
remote_server_name: str | None = None
remote_server_version: str | None = None
error: str | None = None
last_tested_at: datetime | None = None
last_test_succeeded: bool | None = None
class McpServer(McpServerStatus):
server_id: str
version: int
name: str
transport: McpServerTransport
command: str | None = None
args: list[str] = Field(default_factory=list)
url: str | None = None
headers: dict[str, str] = Field(default_factory=dict)
environment: dict[str, str] = Field(default_factory=dict)
permissions: list[str] = Field(default_factory=list)
startup_timeout_seconds: float
tool_timeout_seconds: float
secret_environment: dict[str, bool] = Field(default_factory=dict)
secret_headers: dict[str, bool] = Field(default_factory=dict)
trusted: bool = False
command_digest: str
command_summary: str
class McpServerListResponse(Contract):
items: list[McpServer] = Field(default_factory=list)
class McpToolSummary(Contract):
name: str
remote_name: str
description: str
permission: str | None = None
class McpToolSummaryListResponse(Contract):
items: list[McpToolSummary] = Field(default_factory=list)
class PluginCommandLocation(str, Enum):
command_palette = "command_palette"
context_menu = "context_menu"
toolbar = "toolbar"
class PluginCommand(Contract):
command_id: str
plugin_id: str
title: str
description: str = ""
icon: str | None = None
locations: list[PluginCommandLocation] = Field(default_factory=list)
when: list[str] = Field(default_factory=list)
parameters: dict[str, Any] = Field(default_factory=dict)
enabled: bool = True
class PluginCommandListResponse(Contract):
items: list[PluginCommand] = Field(default_factory=list)
class PluginCommandContext(Contract):
vault_id: str | None = None
note_id: str | None = None
file_path: str | None = None
selection: str | None = None
class PluginCommandExecuteRequest(Contract):
arguments: dict[str, Any] = Field(default_factory=dict)
context: PluginCommandContext = Field(default_factory=PluginCommandContext)
class PluginNotificationEffectPayload(Contract):
level: Literal["info", "success", "warning", "error"] = "info"
message: str = Field(min_length=1, max_length=4096)
class PluginNavigateEffectPayload(Contract):
route: Literal[
"vault-entry",
"workspace",
"search",
"chat",
"agent",
"tasks",
"skills",
"plugins",
"themes",
"settings",
]
class PluginRefreshEffectPayload(Contract):
scope: Literal["workspace", "commands", "settings", "plugins"]
class PluginJobEffectPayload(Contract):
job_id: str = Field(
min_length=1,
max_length=128,
pattern=r"^[A-Za-z0-9][A-Za-z0-9._:-]*$",
)
class PluginNoEffectPayload(Contract):
pass
class PluginNoEffect(Contract):
type: Literal["none"] = "none"
payload: PluginNoEffectPayload = Field(default_factory=PluginNoEffectPayload)
class PluginNotificationEffect(Contract):
type: Literal["notification"] = "notification"
payload: PluginNotificationEffectPayload
class PluginNavigateEffect(Contract):
type: Literal["navigate"] = "navigate"
payload: PluginNavigateEffectPayload
class PluginRefreshEffect(Contract):
type: Literal["refresh"] = "refresh"
payload: PluginRefreshEffectPayload
class PluginJobEffect(Contract):
type: Literal["job"] = "job"
payload: PluginJobEffectPayload
PluginCommandEffect = Annotated[
PluginNoEffect
| PluginNotificationEffect
| PluginNavigateEffect
| PluginRefreshEffect
| PluginJobEffect,
Field(discriminator="type"),
]
PLUGIN_COMMAND_EFFECT_TYPES = (
PluginNoEffect,
PluginNotificationEffect,
PluginNavigateEffect,
PluginRefreshEffect,
PluginJobEffect,
)
class PluginCommandResult(Contract):
command_id: str
status: Literal["completed"] = "completed"
effect: PluginCommandEffect = Field(default_factory=PluginNoEffect)
class PluginSettingType(str, Enum):
string = "string"
number = "number"
boolean = "boolean"
select = "select"
secret = "secret"
class PluginSettingField(Contract):
key: str
label: str
description: str = ""
type: PluginSettingType
required: bool = False
default: Any | None = None
minimum: float | None = None
maximum: float | None = None
options: list[str] = Field(default_factory=list)
class PluginSecretState(Contract):
configured: bool = False
class PluginSettingsSchema(Contract):
plugin_id: str
schema_version: int = Field(ge=1)
fields: list[PluginSettingField] = Field(default_factory=list)
values: dict[str, Any] = Field(default_factory=dict)
secrets: dict[str, PluginSecretState] = Field(default_factory=dict)
class PluginSettingsUpdateRequest(Contract):
schema_version: int = Field(ge=1)
values: dict[str, Any] = Field(default_factory=dict)
class PluginSecretWriteRequest(Contract):
secret: SecretStr
class PluginSecretStatus(Contract):
plugin_id: str
key: str
configured: bool
class PluginPermissionGrantRequest(Contract): class PluginPermissionGrantRequest(Contract):
permissions: list[str] = Field(default_factory=list) permissions: list[str] = Field(default_factory=list)
@@ -558,3 +915,151 @@ class IndexJob(Contract):
status: Literal["queued", "running", "completed", "failed"] status: Literal["queued", "running", "completed", "failed"]
scope: Literal["all", "notes", "vectors"] scope: Literal["all", "notes", "vectors"]
created_at: datetime created_at: datetime
# Benchmark
class BenchmarkKind(str, Enum):
rag = "rag"
agent = "agent"
class BenchmarkStatus(str, Enum):
queued = "queued"
running = "running"
completed = "completed"
failed = "failed"
cancelled = "cancelled"
class RAGDatasetCase(Contract):
case_id: str
query: str = Field(min_length=1)
expected_note_ids: list[str] = Field(default_factory=list)
expected_block_ids: list[str] = Field(default_factory=list)
citation_required: bool = False
tags: list[str] = Field(default_factory=list)
class RAGRetrievalConfig(Contract):
"""RAG Benchmark 的检索参数。top_k 映射到 SearchRequest.limit
其余参数透传到 SearchRequest,由检索引擎实际执行。"""
top_k: int = Field(default=10, ge=1, le=100)
rrf_k: int = Field(default=60, ge=1)
rerank: bool = True
rerank_candidates: int = Field(default=20, ge=1)
score_threshold: float = Field(default=0.0, ge=0.0)
class RAGRunRequest(Contract):
dataset_id: str = Field(min_length=1)
modes: list[SearchMode] = Field(
default_factory=lambda: [SearchMode.fts, SearchMode.vector, SearchMode.hybrid],
min_length=1,
)
retrieval: RAGRetrievalConfig = Field(default_factory=RAGRetrievalConfig)
repeat: int = Field(default=1, ge=1, le=10)
metadata: dict[str, Any] = Field(default_factory=dict)
@field_validator("modes")
@classmethod
def _no_duplicate_modes(cls, value: list[SearchMode]) -> list[SearchMode]:
if len(value) != len(set(value)):
raise ValueError("modes must not contain duplicates")
return value
class RAGMetrics(Contract):
hit_at_1: float = 0.0
hit_at_5: float = 0.0
recall_at_k: float = 0.0
mrr: float = 0.0
citation_hit_rate: float = 0.0
p50_latency_ms: float = 0.0
p95_latency_ms: float = 0.0
# 样本构成:失败样本按零分计入质量指标,汇总不虚高;报告据此可知实际分母
total_cases: int = 0
successful_cases: int = 0
failed_cases: int = 0
failure_rate: float = 0.0
class BenchmarkDatasetInfo(Contract):
dataset_id: str
kind: BenchmarkKind
version: str
description: str = ""
case_count: int
content_hash: str
class BenchmarkDatasetListResponse(Contract):
items: list[BenchmarkDatasetInfo] = Field(default_factory=list)
class BenchmarkRun(Contract):
run_id: str
kind: BenchmarkKind
dataset_id: str
dataset_hash: str
status: BenchmarkStatus
progress: float | None = None
metrics: dict[str, Any] | None = None
config_snapshot: dict[str, Any] = Field(default_factory=dict)
error: str | None = None
error_code: str | None = None
created_at: datetime
started_at: datetime | None = None
completed_at: datetime | None = None
class BenchmarkRunListResponse(Contract):
items: list[BenchmarkRun] = Field(default_factory=list)
page: PageMeta = Field(default_factory=PageMeta)
class BenchmarkEventType(str, Enum):
run_started = "RunStarted"
case_completed = "CaseCompleted"
run_completed = "RunCompleted"
run_failed = "RunFailed"
run_cancelled = "RunCancelled"
class BenchmarkEvent(Contract):
event: BenchmarkEventType
run_id: str
sequence: int
data: dict[str, Any] = Field(default_factory=dict)
timestamp: datetime
class RAGCaseResult(Contract):
case_id: str
mode: SearchMode
repeat: int
latency_ms: float
retrieved_note_ids: list[str] = Field(default_factory=list)
retrieved_block_ids: list[str] = Field(default_factory=list)
hit_at_1: bool = False
hit_at_5: bool = False
recall: float = 0.0
reciprocal_rank: float = 0.0
citation_hit: bool = False
# 该 Case 是否声明了 expected_block_ids(决定是否计入 citation_hit_rate 分母)
citation_applicable: bool = False
error: str | None = None
error_code: str | None = None
class BenchmarkReport(Contract):
run_id: str
kind: BenchmarkKind
dataset_id: str
dataset_hash: str
status: BenchmarkStatus
config_snapshot: dict[str, Any] = Field(default_factory=dict)
metrics: dict[str, Any] = Field(default_factory=dict)
cases: list[RAGCaseResult] = Field(default_factory=list)
error: str | None = None
error_code: str | None = None
+27
View File
@@ -69,6 +69,33 @@ MIGRATIONS: list[str] = [
); );
CREATE INDEX IF NOT EXISTS idx_tasks_status_due ON tasks(status, due_at); CREATE INDEX IF NOT EXISTS idx_tasks_status_due ON tasks(status, due_at);
""", """,
# v3: 第二阶段 Agent TraceRun 与事件事实持久化,供 SSE 恢复和 Benchmark 复用。
"""
CREATE TABLE IF NOT EXISTS agent_runs (
run_id TEXT PRIMARY KEY,
status TEXT NOT NULL,
run_json TEXT NOT NULL,
request_json TEXT NOT NULL,
config_snapshot_json TEXT NOT NULL DEFAULT '{}',
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_agent_runs_created
ON agent_runs(created_at DESC);
CREATE INDEX IF NOT EXISTS idx_agent_runs_status
ON agent_runs(status, updated_at DESC);
CREATE TABLE IF NOT EXISTS agent_events (
run_id TEXT NOT NULL REFERENCES agent_runs(run_id) ON DELETE CASCADE,
sequence INTEGER NOT NULL,
event TEXT NOT NULL,
data_json TEXT NOT NULL DEFAULT '{}',
timestamp TEXT NOT NULL,
PRIMARY KEY (run_id, sequence)
);
CREATE INDEX IF NOT EXISTS idx_agent_events_type
ON agent_events(run_id, event, sequence);
""",
] ]
+11 -7
View File
@@ -1,8 +1,12 @@
from app.extensions.runtime import ( from app.extensions.errors import ExtensionError
AgentConfiguration, from app.extensions.runtime import AgentConfiguration, PluginRuntime, SkillRuntime
ExtensionError, from app.extensions.mcp import McpBridge, McpBridgeError
PluginRuntime,
SkillRuntime,
)
__all__ = ["AgentConfiguration", "ExtensionError", "PluginRuntime", "SkillRuntime"] __all__ = [
"AgentConfiguration",
"ExtensionError",
"McpBridge",
"McpBridgeError",
"PluginRuntime",
"SkillRuntime",
]
+847
View File
@@ -0,0 +1,847 @@
"""Plugin Command Registry 与 Settings/Secret 命名空间存储。"""
from __future__ import annotations
import asyncio
import hashlib
import inspect
import json
import math
import re
import threading
from collections import deque
from dataclasses import dataclass
from datetime import UTC, datetime
from pathlib import Path
from time import perf_counter
from typing import Any, Awaitable, Callable, Literal
from jsonschema import Draft202012Validator
from jsonschema.exceptions import SchemaError, ValidationError as JsonSchemaValidationError
from pydantic import BaseModel, ConfigDict, Field, model_validator
from app.config import get_settings
from app.contracts import (
PLUGIN_COMMAND_EFFECT_TYPES,
PluginCommand,
PluginCommandContext,
PluginCommandEffect,
PluginCommandLocation,
PluginCommandResult,
PluginSecretState,
PluginSecretStatus,
PluginSettingField,
PluginSettingType,
PluginSettingsSchema,
)
from app.extensions.errors import ExtensionError
from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore
from app.schema_security import (
SchemaReferenceError,
reject_external_schema_references,
)
_CONTRIBUTION_ID = re.compile(r"^[a-z0-9][a-z0-9._-]*$")
_SETTING_KEY = re.compile(r"^[a-z][a-z0-9._-]{0,127}$")
_HOST_ICONS = {"bolt", "document", "edit", "link", "refresh", "search", "setting"}
_WHEN_TOKENS = {
"workspace.has_vault",
"editor.has_note",
"editor.has_selection",
}
_CONTEXT_KEYS = {"vault_id", "note_id", "file_path", "selection"}
_WHEN_CONTEXT = {
"workspace.has_vault": "vault_id",
"editor.has_note": "note_id",
"editor.has_selection": "selection",
}
class PluginCommandSpec(BaseModel):
"""包内 commands.yaml 的宿主侧声明,不直接暴露 handler。"""
model_config = ConfigDict(extra="forbid")
command_id: str
title: str
description: str = ""
icon: str | None = None
locations: list[PluginCommandLocation] = Field(default_factory=list)
when: list[str] = Field(default_factory=list)
context: list[Literal["vault_id", "note_id", "file_path", "selection"]] = Field(
default_factory=list
)
parameters: dict[str, Any] = Field(
default_factory=lambda: {
"type": "object",
"properties": {},
"additionalProperties": False,
}
)
permission: str | None = None
secrets: list[str] = Field(default_factory=list)
handler: Literal["echo", "uppercase_selection"] | None = None
mcp_tool: str | None = None
timeout_seconds: int = Field(default=30, ge=1, le=120)
@model_validator(mode="after")
def validate_execution_target(self) -> "PluginCommandSpec":
if (self.handler is None) == (self.mcp_tool is None):
raise ValueError("Command must declare exactly one handler or mcp_tool target.")
return self
CommandExecutor = Callable[
[dict[str, Any], dict[str, Any]],
PluginCommandEffect | Awaitable[PluginCommandEffect],
]
PluginSecretResolver = Callable[[str], str | None]
@dataclass(slots=True)
class _RegisteredCommand:
command: PluginCommand
spec: PluginCommandSpec
executor: CommandExecutor
@dataclass(frozen=True, slots=True)
class PluginCommandAuditEvent:
"""不记录参数与上下文的轻量审计事件,避免把正文或 Secret 写入日志。"""
command_id: str
plugin_id: str
status: Literal["completed", "failed"]
duration_ms: int
error_code: str | None
created_at: datetime
class CommandRegistry:
"""只发布已启用 Plugin 的受控 Command Contribution。"""
def __init__(self) -> None:
self._commands: dict[str, _RegisteredCommand] = {}
self._audit: deque[PluginCommandAuditEvent] = deque(maxlen=500)
self._lock = threading.RLock()
def register(
self,
plugin_id: str,
spec: PluginCommandSpec,
executor: CommandExecutor,
) -> None:
validate_command_spec(plugin_id, spec)
command = PluginCommand(
command_id=spec.command_id,
plugin_id=plugin_id,
title=spec.title,
description=spec.description,
icon=spec.icon,
locations=spec.locations,
when=spec.when,
parameters=spec.parameters,
enabled=True,
)
with self._lock:
if spec.command_id in self._commands:
raise ExtensionError(
"PLUGIN_COMMAND_CONFLICT",
f"Plugin command is already registered: {spec.command_id}",
status_code=409,
details={"command_id": spec.command_id},
)
self._commands[spec.command_id] = _RegisteredCommand(command, spec, executor)
def unregister(self, command_id: str) -> None:
with self._lock:
self._commands.pop(command_id, None)
def contains(self, command_id: str) -> bool:
with self._lock:
return command_id in self._commands
def list(self, location: PluginCommandLocation | None = None) -> list[PluginCommand]:
with self._lock:
items = [
item.command.model_copy(deep=True)
for item in self._commands.values()
if location is None or location in item.command.locations
]
return sorted(items, key=lambda item: item.command_id)
def audit_events(self) -> list[PluginCommandAuditEvent]:
"""返回有界审计快照;事件刻意不包含 arguments/context/effect。"""
with self._lock:
return list(self._audit)
async def execute(
self,
command_id: str,
arguments: dict[str, Any],
context: PluginCommandContext,
) -> PluginCommandResult:
with self._lock:
registered = self._commands.get(command_id)
if registered is None:
raise ExtensionError(
"PLUGIN_COMMAND_NOT_FOUND",
f"Plugin command is not registered or enabled: {command_id}",
status_code=404,
details={"command_id": command_id},
)
started_at = perf_counter()
try:
Draft202012Validator(registered.spec.parameters).validate(arguments)
except JsonSchemaValidationError as exc:
error = ExtensionError(
"PLUGIN_COMMAND_ARGUMENT_INVALID",
"Plugin command arguments do not match the declared schema.",
details={"command_id": command_id, "path": list(exc.path)},
)
self._record_audit(registered, started_at, error.code)
raise error from exc
raw_context = context.model_dump(exclude_none=True)
missing = [
token
for token in registered.spec.when
if not raw_context.get(_WHEN_CONTEXT[token])
]
if missing:
error = ExtensionError(
"PLUGIN_COMMAND_CONTEXT_INVALID",
"Plugin command context does not satisfy its when conditions.",
details={"command_id": command_id, "missing": missing},
)
self._record_audit(registered, started_at, error.code)
raise error
scoped_context = {
key: raw_context[key]
for key in registered.spec.context
if key in raw_context
}
try:
effect = registered.executor(dict(arguments), scoped_context)
if inspect.isawaitable(effect):
effect = await asyncio.wait_for(
effect, timeout=registered.spec.timeout_seconds
)
except TimeoutError as exc:
error = ExtensionError(
"PLUGIN_COMMAND_TIMEOUT",
"Plugin command execution timed out.",
status_code=504,
details={"command_id": command_id},
)
self._record_audit(registered, started_at, error.code)
raise error from exc
except ExtensionError as exc:
self._record_audit(registered, started_at, exc.code)
raise
except Exception as exc:
error = ExtensionError(
"PLUGIN_COMMAND_EXECUTION_FAILED",
"Plugin command execution failed.",
status_code=502,
details={"command_id": command_id},
)
self._record_audit(registered, started_at, error.code)
raise error from exc
if not isinstance(effect, PLUGIN_COMMAND_EFFECT_TYPES):
error = ExtensionError(
"PLUGIN_COMMAND_RESULT_INVALID",
"Plugin command returned an invalid effect.",
status_code=502,
details={"command_id": command_id},
)
self._record_audit(registered, started_at, error.code)
raise error
try:
encoded_effect = json.dumps(effect.model_dump(mode="json"), ensure_ascii=False)
except (TypeError, ValueError) as exc:
error = ExtensionError(
"PLUGIN_COMMAND_RESULT_INVALID",
"Plugin command returned a non-serializable effect.",
status_code=502,
details={"command_id": command_id},
)
self._record_audit(registered, started_at, error.code)
raise error from exc
if len(encoded_effect.encode("utf-8")) > 64 * 1024:
error = ExtensionError(
"PLUGIN_COMMAND_RESULT_TOO_LARGE",
"Plugin command effect exceeds the 64 KiB response limit.",
status_code=502,
details={"command_id": command_id},
)
self._record_audit(registered, started_at, error.code)
raise error
self._record_audit(registered, started_at, None)
return PluginCommandResult(command_id=command_id, effect=effect)
def _record_audit(
self,
registered: _RegisteredCommand,
started_at: float,
error_code: str | None,
) -> None:
event = PluginCommandAuditEvent(
command_id=registered.command.command_id,
plugin_id=registered.command.plugin_id,
status="failed" if error_code else "completed",
duration_ms=max(0, round((perf_counter() - started_at) * 1000)),
error_code=error_code,
created_at=datetime.now(UTC),
)
with self._lock:
self._audit.append(event)
class PluginSettingsDefinition(BaseModel):
model_config = ConfigDict(extra="forbid")
section_id: str
schema_version: int = Field(ge=1)
fields: list[PluginSettingField] = Field(default_factory=list)
class PluginSettingsStore:
"""非敏感值写入插件命名空间;Secret 只保存加密凭据引用。"""
def __init__(self, credentials: EncryptedCredentialStore) -> None:
self.credentials = credentials
self._lock = threading.RLock()
@staticmethod
def _path() -> Path:
return get_settings().data_dir / "plugins" / "settings.json"
def get(
self, plugin_id: str, definition: PluginSettingsDefinition
) -> PluginSettingsSchema:
with self._lock:
entry = self._entry(self._read(), plugin_id)
stored_values = entry.get("values", {})
secret_refs = entry.get("secret_refs", {})
if not isinstance(stored_values, dict) or not isinstance(secret_refs, dict):
raise self._storage_format_error(plugin_id)
validated_refs = self._validate_secret_refs(plugin_id, secret_refs)
values = {
field.key: field.default
for field in definition.fields
if field.type != PluginSettingType.secret and field.default is not None
}
allowed_values = {
field.key
for field in definition.fields
if field.type != PluginSettingType.secret
}
fields = {field.key: field for field in definition.fields}
for key, value in stored_values.items():
if key not in allowed_values:
continue
try:
_validate_setting_value(fields[key], value)
except ExtensionError as exc:
raise self._storage_format_error(plugin_id) from exc
values[key] = value
secrets: dict[str, PluginSecretState] = {}
for field in definition.fields:
if field.type != PluginSettingType.secret:
continue
reference = validated_refs.get(field.key)
secrets[field.key] = PluginSecretState(
configured=isinstance(reference, str) and self._has_secret(reference)
)
return PluginSettingsSchema(
plugin_id=plugin_id,
schema_version=definition.schema_version,
fields=definition.fields,
values=values,
secrets=secrets,
)
def runtime_values(
self, plugin_id: str, definition: PluginSettingsDefinition
) -> dict[str, Any]:
"""返回可供 Command 使用的完整普通设置,并拦截未配置的必填项。"""
schema = self.get(plugin_id, definition)
missing = [
field.key
for field in definition.fields
if field.required
and field.type != PluginSettingType.secret
and field.key not in schema.values
]
if missing:
raise ExtensionError(
"PLUGIN_SETTINGS_REQUIRED",
"Required Plugin settings have not been configured.",
status_code=409,
details={"plugin_id": plugin_id, "fields": missing},
)
return schema.values
def update(
self,
plugin_id: str,
definition: PluginSettingsDefinition,
schema_version: int,
values: dict[str, Any],
) -> PluginSettingsSchema:
if schema_version != definition.schema_version:
raise ExtensionError(
"PLUGIN_SETTINGS_VERSION_CONFLICT",
"Plugin settings schema version is out of date.",
status_code=409,
details={
"plugin_id": plugin_id,
"requested_version": schema_version,
"current_version": definition.schema_version,
},
)
fields = {field.key: field for field in definition.fields}
unknown = sorted(set(values) - set(fields))
if unknown:
raise ExtensionError(
"PLUGIN_SETTINGS_FIELD_INVALID",
"Plugin settings contain unknown fields.",
details={"plugin_id": plugin_id, "fields": unknown},
)
secret_keys = sorted(
key for key in values if fields[key].type == PluginSettingType.secret
)
if secret_keys:
raise ExtensionError(
"PLUGIN_SETTINGS_FIELD_INVALID",
"Secret fields must use the dedicated Secret endpoint.",
details={"plugin_id": plugin_id, "fields": secret_keys},
)
for key, value in values.items():
_validate_setting_value(fields[key], value)
with self._lock:
data = self._read()
entry = self._entry(data, plugin_id, create=True)
current = entry.get("values", {})
if not isinstance(current, dict):
raise self._storage_format_error(plugin_id)
entry["values"] = current
current.update(values)
effective = {
field.key: field.default
for field in definition.fields
if field.type != PluginSettingType.secret and field.default is not None
}
effective.update(current)
missing = [
field.key
for field in definition.fields
if field.required
and field.type != PluginSettingType.secret
and field.key not in effective
]
if missing:
raise ExtensionError(
"PLUGIN_SETTINGS_FIELD_INVALID",
"Required Plugin settings are missing.",
details={"plugin_id": plugin_id, "fields": missing},
)
entry["schema_version"] = definition.schema_version
self._write(data)
return self.get(plugin_id, definition)
def put_secret(
self,
plugin_id: str,
definition: PluginSettingsDefinition,
key: str,
secret: str,
) -> PluginSecretStatus:
_secret_field(definition, plugin_id, key)
if not secret:
raise ExtensionError(
"PLUGIN_SECRET_VALUE_INVALID",
"Plugin secret cannot be empty.",
details={"plugin_id": plugin_id, "key": key},
)
if len(secret.encode("utf-8")) > 64 * 1024:
raise ExtensionError(
"PLUGIN_SECRET_VALUE_INVALID",
"Plugin secret exceeds the 64 KiB limit.",
details={"plugin_id": plugin_id, "key": key},
)
reference = _secret_reference(plugin_id, key)
with self._lock:
data = self._read()
entry = self._entry(data, plugin_id, create=True)
refs = entry.get("secret_refs", {})
if not isinstance(refs, dict):
raise self._storage_format_error(plugin_id)
self._validate_secret_refs(plugin_id, refs)
entry["secret_refs"] = refs
try:
previous = self.credentials.resolve(reference)
self.credentials.put(reference, secret)
except CredentialStoreError as exc:
raise ExtensionError(
"PLUGIN_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc
refs[key] = reference
entry["schema_version"] = definition.schema_version
try:
self._write(data)
except ExtensionError:
# 普通设置落盘失败时恢复凭据旧值,避免产生不可达的新 Secret。
try:
if previous is None:
self.credentials.delete(reference)
else:
self.credentials.put(reference, previous)
except CredentialStoreError:
pass
raise
return PluginSecretStatus(plugin_id=plugin_id, key=key, configured=True)
def delete_secret(
self,
plugin_id: str,
definition: PluginSettingsDefinition,
key: str,
) -> PluginSecretStatus:
_secret_field(definition, plugin_id, key)
with self._lock:
data = self._read()
entry = self._entry(data, plugin_id)
refs = entry.get("secret_refs", {})
if not isinstance(refs, dict):
raise self._storage_format_error(plugin_id)
self._validate_secret_refs(plugin_id, refs)
reference = _secret_reference(plugin_id, key)
had_reference = refs.pop(key, None) is not None
if plugin_id in data and had_reference:
self._write(data)
try:
self.credentials.delete(reference)
except CredentialStoreError as exc:
if had_reference:
refs[key] = reference
try:
self._write(data)
except ExtensionError as rollback_exc:
raise ExtensionError(
"PLUGIN_STORAGE_ERROR",
"Plugin Secret deletion failed and its reference could not be restored.",
status_code=500,
details={"plugin_id": plugin_id, "key": key},
) from rollback_exc
raise ExtensionError(
"PLUGIN_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc
return PluginSecretStatus(plugin_id=plugin_id, key=key, configured=False)
def resolve_secret(
self, plugin_id: str, definition: PluginSettingsDefinition, key: str
) -> str | None:
_secret_field(definition, plugin_id, key)
with self._lock:
entry = self._entry(self._read(), plugin_id)
refs = entry.get("secret_refs", {})
if not isinstance(refs, dict):
raise self._storage_format_error(plugin_id)
reference = self._validate_secret_refs(plugin_id, refs).get(key)
try:
return self.credentials.resolve(reference) if isinstance(reference, str) else None
except CredentialStoreError as exc:
raise ExtensionError(
"PLUGIN_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc
def remove_plugin(self, plugin_id: str) -> None:
with self._lock:
data = self._read()
entry = data.pop(plugin_id, None)
references: list[str] = []
if entry is not None and not isinstance(entry, dict):
raise self._storage_format_error(plugin_id)
if entry is not None:
refs = entry.get("secret_refs", {})
if not isinstance(refs, dict):
raise self._storage_format_error(plugin_id)
references = list(self._validate_secret_refs(plugin_id, refs).values())
if entry is not None:
self._write(data)
try:
self.credentials.delete_many(references)
except CredentialStoreError as exc:
if entry is not None:
data[plugin_id] = entry
try:
self._write(data)
except ExtensionError as rollback_exc:
raise ExtensionError(
"PLUGIN_STORAGE_ERROR",
"Plugin uninstall failed and its Settings namespace could not be restored.",
status_code=500,
details={"plugin_id": plugin_id},
) from rollback_exc
raise ExtensionError(
"PLUGIN_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc
def _validate_secret_refs(
self, plugin_id: str, refs: dict[Any, Any]
) -> dict[str, str]:
validated: dict[str, str] = {}
for key, reference in refs.items():
if (
not isinstance(key, str)
or not _SETTING_KEY.fullmatch(key)
or not isinstance(reference, str)
or reference != _secret_reference(plugin_id, key)
):
raise self._storage_format_error(plugin_id)
validated[key] = reference
return validated
def _has_secret(self, reference: str) -> bool:
try:
return self.credentials.has(reference)
except CredentialStoreError as exc:
raise ExtensionError(
"PLUGIN_SECRET_STORE_ERROR", str(exc), status_code=500
) from exc
@staticmethod
def _storage_format_error(plugin_id: str) -> ExtensionError:
return ExtensionError(
"PLUGIN_STORAGE_ERROR",
"Plugin settings namespace has an invalid format.",
status_code=500,
details={"plugin_id": plugin_id},
)
def _entry(
self,
data: dict[str, dict[str, Any]],
plugin_id: str,
*,
create: bool = False,
) -> dict[str, Any]:
entry = data.get(plugin_id)
if entry is None:
if create:
data[plugin_id] = {}
return data[plugin_id]
return {}
if not isinstance(entry, dict):
raise self._storage_format_error(plugin_id)
return entry
def _read(self) -> dict[str, dict[str, Any]]:
path = self._path()
if not path.exists():
return {}
try:
value = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise ExtensionError(
"PLUGIN_STORAGE_ERROR",
"Plugin settings storage cannot be loaded.",
status_code=500,
) from exc
if not isinstance(value, dict):
raise ExtensionError(
"PLUGIN_STORAGE_ERROR",
"Plugin settings storage has an invalid format.",
status_code=500,
)
return value
def _write(self, value: dict[str, dict[str, Any]]) -> None:
path = self._path()
temporary = path.with_suffix(".tmp")
try:
path.parent.mkdir(parents=True, exist_ok=True)
temporary.write_text(
json.dumps(value, ensure_ascii=False, sort_keys=True),
encoding="utf-8",
)
temporary.replace(path)
except OSError as exc:
try:
temporary.unlink(missing_ok=True)
except OSError:
pass
raise ExtensionError(
"PLUGIN_STORAGE_ERROR",
"Plugin settings storage cannot be written.",
status_code=500,
) from exc
def validate_settings_definition(
plugin_id: str, definition: PluginSettingsDefinition
) -> None:
if not _CONTRIBUTION_ID.fullmatch(definition.section_id):
raise _settings_schema_error(plugin_id, "Settings section id is invalid.")
if not definition.section_id.startswith(f"{plugin_id}."):
raise _settings_schema_error(
plugin_id, "Settings section id must use the Plugin namespace."
)
keys: set[str] = set()
for field in definition.fields:
if not _SETTING_KEY.fullmatch(field.key) or field.key in keys:
raise _settings_schema_error(plugin_id, f"Invalid or duplicate setting key: {field.key}")
keys.add(field.key)
if field.type == PluginSettingType.select and not field.options:
raise _settings_schema_error(plugin_id, f"Select setting requires options: {field.key}")
if field.type != PluginSettingType.select and field.options:
raise _settings_schema_error(plugin_id, f"Only select settings accept options: {field.key}")
if field.type != PluginSettingType.number and (
field.minimum is not None or field.maximum is not None
):
raise _settings_schema_error(plugin_id, f"Only number settings accept bounds: {field.key}")
if any(
bound is not None and not math.isfinite(bound)
for bound in (field.minimum, field.maximum)
):
raise _settings_schema_error(
plugin_id, f"Number setting bounds must be finite: {field.key}"
)
if field.minimum is not None and field.maximum is not None and field.minimum > field.maximum:
raise _settings_schema_error(plugin_id, f"Setting bounds are reversed: {field.key}")
if field.type == PluginSettingType.secret and field.default is not None:
raise _settings_schema_error(plugin_id, f"Secret settings cannot declare defaults: {field.key}")
if field.default is not None:
try:
_validate_setting_value(field, field.default)
except ExtensionError as exc:
raise _settings_schema_error(plugin_id, exc.message) from exc
def validate_command_spec(plugin_id: str, spec: PluginCommandSpec) -> None:
if not _CONTRIBUTION_ID.fullmatch(spec.command_id) or not spec.command_id.startswith(
f"{plugin_id}."
):
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command id must be valid and use the Plugin namespace.",
details={"plugin_id": plugin_id, "command_id": spec.command_id},
)
if not spec.locations:
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command must declare at least one location.",
details={"command_id": spec.command_id},
)
if len(spec.locations) != len(set(spec.locations)):
raise ExtensionError("PLUGIN_COMMAND_INVALID", "Plugin command locations must be unique.")
if len(spec.when) != len(set(spec.when)) or len(spec.context) != len(set(spec.context)):
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command when/context entries must be unique.",
)
if len(spec.secrets) != len(set(spec.secrets)):
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command Secret entries must be unique.",
details={"command_id": spec.command_id},
)
unknown_when = sorted(set(spec.when) - _WHEN_TOKENS)
if unknown_when:
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command declares unsupported when tokens.",
details={"command_id": spec.command_id, "when": unknown_when},
)
required_context = {_WHEN_CONTEXT[token] for token in spec.when}
if not required_context.issubset(set(spec.context)):
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command context must include every field required by when.",
details={"command_id": spec.command_id},
)
if not set(spec.context).issubset(_CONTEXT_KEYS):
raise ExtensionError("PLUGIN_COMMAND_INVALID", "Plugin command context is invalid.")
if spec.icon and spec.icon not in _HOST_ICONS:
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command icon is not a supported Host icon.",
details={"command_id": spec.command_id, "icon": spec.icon},
)
if spec.parameters.get("type", "object") != "object":
raise ExtensionError("PLUGIN_COMMAND_INVALID", "Command parameters must be an object schema.")
try:
Draft202012Validator.check_schema(spec.parameters)
reject_external_schema_references(spec.parameters)
except (SchemaReferenceError, SchemaError) as exc:
message = exc.message if isinstance(exc, SchemaError) else str(exc)
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
f"Plugin command parameters contain invalid JSON Schema: {message}",
) from exc
def _validate_setting_value(field: PluginSettingField, value: Any) -> None:
valid = False
if field.type == PluginSettingType.string:
valid = isinstance(value, str) and len(value.encode("utf-8")) <= 64 * 1024
elif field.type == PluginSettingType.number:
valid = (
(isinstance(value, int) and not isinstance(value, bool))
or (isinstance(value, float) and math.isfinite(value))
)
elif field.type == PluginSettingType.boolean:
valid = isinstance(value, bool)
elif field.type == PluginSettingType.select:
valid = isinstance(value, str) and value in field.options
if not valid:
raise ExtensionError(
"PLUGIN_SETTINGS_FIELD_INVALID",
f"Plugin setting has an invalid value: {field.key}",
details={"key": field.key},
)
if field.type == PluginSettingType.number:
if field.minimum is not None and value < field.minimum:
raise ExtensionError(
"PLUGIN_SETTINGS_FIELD_INVALID",
f"Plugin setting is below its minimum: {field.key}",
details={"key": field.key, "minimum": field.minimum},
)
if field.maximum is not None and value > field.maximum:
raise ExtensionError(
"PLUGIN_SETTINGS_FIELD_INVALID",
f"Plugin setting is above its maximum: {field.key}",
details={"key": field.key, "maximum": field.maximum},
)
def _secret_field(
definition: PluginSettingsDefinition, plugin_id: str, key: str
) -> PluginSettingField:
field = next((item for item in definition.fields if item.key == key), None)
if field is None or field.type != PluginSettingType.secret:
raise ExtensionError(
"PLUGIN_SECRET_FIELD_NOT_FOUND",
f"Plugin secret field does not exist: {key}",
status_code=404,
details={"plugin_id": plugin_id, "key": key},
)
return field
def _secret_reference(plugin_id: str, key: str) -> str:
digest = hashlib.sha256(f"{plugin_id}\0{key}".encode("utf-8")).hexdigest()
return f"plugin.{digest}"
def _settings_schema_error(plugin_id: str, message: str) -> ExtensionError:
return ExtensionError(
"PLUGIN_SETTINGS_SCHEMA_INVALID",
message,
details={"plugin_id": plugin_id},
)
+21
View File
@@ -0,0 +1,21 @@
from __future__ import annotations
from typing import Any
class ExtensionError(RuntimeError):
"""Extension Core 对 API 暴露的稳定领域错误。"""
def __init__(
self,
code: str,
message: str,
*,
status_code: int = 422,
details: dict[str, Any] | None = None,
) -> None:
super().__init__(message)
self.code = code
self.message = message
self.status_code = status_code
self.details = details or {}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+696 -49
View File
@@ -1,21 +1,44 @@
from __future__ import annotations from __future__ import annotations
import re import re
import threading
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import Any, Literal from typing import Any, Literal
from uuid import uuid4
import yaml import yaml
from jsonschema import Draft202012Validator from jsonschema import Draft202012Validator
from jsonschema.exceptions import SchemaError from jsonschema.exceptions import (
from pydantic import BaseModel, ConfigDict, Field, ValidationError, create_model SchemaError,
ValidationError as JsonSchemaValidationError,
)
from pydantic import (
BaseModel,
ConfigDict,
Field,
TypeAdapter,
ValidationError,
create_model,
)
from app.agent.tools import ToolExecutionContext, ToolRegistry from app.agent.tools import ToolExecutionContext, ToolExecutionError, ToolRegistry
from app.agent.permissions import KNOWN_PERMISSIONS from app.agent.permissions import KNOWN_PERMISSIONS
from app.contracts import ( from app.contracts import (
ModelCapability, ModelCapability,
Plugin, Plugin,
PluginCommand,
PluginCommandContext,
PluginCommandEffect,
PluginNoEffect,
PluginNotificationEffect,
PluginCommandLocation,
PluginCommandResult,
PluginManifest, PluginManifest,
PluginHostStatus,
PluginSecretStatus,
PluginSettingType,
PluginSettingsSchema,
PluginStatus, PluginStatus,
RetrievalConfig, RetrievalConfig,
Skill, Skill,
@@ -23,26 +46,26 @@ from app.contracts import (
SkillStatus, SkillStatus,
ToolDefinition, ToolDefinition,
) )
from app.extensions.contributions import (
CommandRegistry,
PluginCommandSpec,
PluginSecretResolver,
PluginSettingsDefinition,
PluginSettingsStore,
validate_command_spec,
validate_settings_definition,
)
from app.extensions.errors import ExtensionError
from app.extensions.mcp import McpBridge, McpBridgeError, McpDiscoveredTool
from app.providers.credentials import EncryptedCredentialStore
from app.schema_security import (
SchemaReferenceError,
reject_external_schema_references,
)
_EXTENSION_ID = re.compile(r"^[a-z0-9][a-z0-9._-]*$") _EXTENSION_ID = re.compile(r"^[a-z0-9][a-z0-9._-]*$")
class ExtensionError(RuntimeError):
def __init__(
self,
code: str,
message: str,
*,
status_code: int = 422,
details: dict[str, Any] | None = None,
) -> None:
super().__init__(message)
self.code = code
self.message = message
self.status_code = status_code
self.details = details or {}
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class AgentConfiguration: class AgentConfiguration:
skill_id: str skill_id: str
@@ -67,6 +90,7 @@ class SkillRuntime:
self._records: dict[str, _SkillRecord] = {} self._records: dict[str, _SkillRecord] = {}
def install(self, package_path: str | Path) -> Skill: def install(self, package_path: str | Path) -> Skill:
# TODO(extension): 将安装记录持久化,应用重启后从可信包目录恢复状态。
root = _package_dir(package_path) root = _package_dir(package_path)
raw = _read_yaml(root / "skill.yaml") raw = _read_yaml(root / "skill.yaml")
if "id" in raw and "skill_id" not in raw: if "id" in raw and "skill_id" not in raw:
@@ -234,24 +258,70 @@ class DeclarativePluginHost:
return {"text": str(values.get("text", "")).upper()} return {"text": str(values.get("text", "")).upper()}
raise ExtensionError("PLUGIN_HANDLER_UNSUPPORTED", f"Unsupported handler: {handler}") raise ExtensionError("PLUGIN_HANDLER_UNSUPPORTED", f"Unsupported handler: {handler}")
async def execute_command(
self,
handler: str,
arguments: dict[str, Any],
context: dict[str, Any],
settings: dict[str, Any],
resolve_secret: PluginSecretResolver,
) -> PluginCommandEffect:
"""执行宿主内置的白名单 Command handler,不导入 Plugin Python 代码。"""
if handler == "echo":
message = str(arguments.get("message", context.get("selection", "")))
if not message:
return PluginNoEffect()
return PluginNotificationEffect(
payload={"level": "info", "message": message},
)
if handler == "uppercase_selection":
text = str(arguments.get("text", context.get("selection", "")))
limit = int(settings.get("result_limit", 100))
return PluginNotificationEffect(
payload={"level": "success", "message": text[:limit].upper()},
)
raise ExtensionError(
"PLUGIN_HANDLER_UNSUPPORTED", f"Unsupported command handler: {handler}"
)
@dataclass(slots=True) @dataclass(slots=True)
class _PluginRecord: class _PluginRecord:
plugin: Plugin plugin: Plugin
tools: list[DeclarativeToolSpec] tools: list[DeclarativeToolSpec]
commands: list[PluginCommandSpec]
settings_definition: PluginSettingsDefinition | None
package_path: Path package_path: Path
registered_tools: list[str] registered_tools: list[str]
registered_commands: list[str]
mcp_remote_names: dict[str, str]
mcp_command_schemas: dict[str, dict[str, Any]]
class PluginRuntime: class PluginRuntime:
"""Plugin Manifest、生命周期及 Tool Contribution 注册。""" """Plugin Manifest、生命周期及 Tool Contribution 注册。"""
def __init__(self, tools: ToolRegistry, host: DeclarativePluginHost | None = None) -> None: def __init__(
self,
tools: ToolRegistry,
host: DeclarativePluginHost | None = None,
mcp_bridge: McpBridge | None = None,
credentials: EncryptedCredentialStore | None = None,
*,
allow_unsandboxed_mcp: bool = False,
) -> None:
self.registry = tools self.registry = tools
self.host = host or DeclarativePluginHost() self.host = host or DeclarativePluginHost()
self.mcp = mcp_bridge or McpBridge()
self.commands = CommandRegistry()
self.settings = PluginSettingsStore(credentials or EncryptedCredentialStore())
self.allow_unsandboxed_mcp = allow_unsandboxed_mcp
self._records: dict[str, _PluginRecord] = {} self._records: dict[str, _PluginRecord] = {}
self._lock = threading.RLock()
def install(self, package_path: str | Path) -> Plugin: def install(self, package_path: str | Path) -> Plugin:
# 安装阶段只读取清单;MCP 子进程必须在权限授予后的 enable 阶段启动。
root = _package_dir(package_path) root = _package_dir(package_path)
raw = _read_yaml(root / "plugin.yaml") raw = _read_yaml(root / "plugin.yaml")
if "id" in raw and "plugin_id" not in raw: if "id" in raw and "plugin_id" not in raw:
@@ -269,7 +339,11 @@ class PluginRuntime:
status_code=409, status_code=409,
) )
specs = self._load_tools(root) _validate_backend(manifest)
specs = [] if manifest.backend.type == "mcp" else self._load_tools(root)
command_specs = self._load_commands(root)
settings_definition = self._load_settings(root)
if manifest.backend.type != "mcp":
declared = set(manifest.contributes.tools) declared = set(manifest.contributes.tools)
actual = {spec.name for spec in specs} actual = {spec.name for spec in specs}
if declared != actual: if declared != actual:
@@ -287,6 +361,89 @@ class PluginRuntime:
f"Tool permission is not declared by Plugin: {spec.permission}", f"Tool permission is not declared by Plugin: {spec.permission}",
details={"tool": spec.name, "permission": spec.permission}, details={"tool": spec.name, "permission": spec.permission},
) )
declared_commands = set(manifest.contributes.commands)
actual_commands = {spec.command_id for spec in command_specs}
if (
declared_commands != actual_commands
or len(manifest.contributes.commands) != len(declared_commands)
or len(command_specs) != len(actual_commands)
):
raise ExtensionError(
"PLUGIN_CONTRIBUTION_INVALID",
"plugin.yaml command contributions must exactly match commands.yaml",
details={
"declared": sorted(declared_commands),
"actual": sorted(actual_commands),
},
)
for spec in command_specs:
validate_command_spec(manifest.plugin_id, spec)
if spec.permission and spec.permission not in manifest.permissions:
raise ExtensionError(
"PLUGIN_PERMISSION_UNDECLARED",
f"Command permission is not declared by Plugin: {spec.permission}",
details={"command": spec.command_id, "permission": spec.permission},
)
declared_sections = set(manifest.contributes.settings_sections)
actual_sections = (
{settings_definition.section_id} if settings_definition is not None else set()
)
if (
declared_sections != actual_sections
or len(manifest.contributes.settings_sections) != len(declared_sections)
):
raise ExtensionError(
"PLUGIN_CONTRIBUTION_INVALID",
"plugin.yaml settings contributions must exactly match settings.yaml",
details={
"declared": sorted(declared_sections),
"actual": sorted(actual_sections),
},
)
if settings_definition is not None:
validate_settings_definition(manifest.plugin_id, settings_definition)
secret_fields = (
{
field.key
for field in settings_definition.fields
if field.type == PluginSettingType.secret
}
if settings_definition is not None
else set()
)
for spec in command_specs:
unknown_secrets = sorted(set(spec.secrets) - secret_fields)
if unknown_secrets:
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"Plugin command references undeclared Secret settings.",
details={
"command_id": spec.command_id,
"secrets": unknown_secrets,
},
)
if spec.secrets and "secrets.use" not in manifest.permissions:
raise ExtensionError(
"PLUGIN_PERMISSION_UNDECLARED",
"Commands using Secret settings require the secrets.use permission.",
details={"command_id": spec.command_id},
)
if spec.mcp_tool is not None:
_validate_id("MCP command target", spec.mcp_tool)
if manifest.backend.type != "mcp" or not spec.mcp_tool.startswith(
f"{manifest.plugin_id}."
):
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"MCP Command target must use the current Plugin namespace.",
details={"command_id": spec.command_id},
)
if spec.mcp_tool in manifest.contributes.tools:
raise ExtensionError(
"PLUGIN_COMMAND_INVALID",
"MCP Command target cannot also be exposed as an Agent Tool.",
details={"command_id": spec.command_id},
)
record = _PluginRecord( record = _PluginRecord(
plugin=Plugin( plugin=Plugin(
@@ -298,8 +455,13 @@ class PluginRuntime:
), ),
), ),
tools=specs, tools=specs,
commands=command_specs,
settings_definition=settings_definition,
package_path=root, package_path=root,
registered_tools=[], registered_tools=[],
registered_commands=[],
mcp_remote_names={},
mcp_command_schemas={},
) )
self._records[manifest.plugin_id] = record self._records[manifest.plugin_id] = record
return record.plugin.model_copy(deep=True) return record.plugin.model_copy(deep=True)
@@ -311,17 +473,14 @@ class PluginRuntime:
return self._record(plugin_id).plugin.model_copy(deep=True) return self._record(plugin_id).plugin.model_copy(deep=True)
def enable(self, plugin_id: str) -> Plugin: def enable(self, plugin_id: str) -> Plugin:
# Host 启动和 Tool 批量注册必须串行,避免并发 enable 产生重复进程或半注册状态。
with self._lock:
return self._enable(plugin_id)
def _enable(self, plugin_id: str) -> Plugin:
record = self._record(plugin_id) record = self._record(plugin_id)
if record.plugin.enabled: if record.plugin.enabled:
return record.plugin.model_copy(deep=True) return record.plugin.model_copy(deep=True)
if record.plugin.manifest.backend.type == "mcp":
record.plugin.status = PluginStatus.dependency_missing
raise ExtensionError(
"PLUGIN_HOST_UNAVAILABLE",
"MCP Plugin Host is reserved for the second development phase.",
status_code=501,
details={"plugin_id": plugin_id, "backend": "mcp"},
)
missing_grants = sorted( missing_grants = sorted(
set(record.plugin.manifest.permissions) - set(record.plugin.granted_permissions) set(record.plugin.manifest.permissions) - set(record.plugin.granted_permissions)
) )
@@ -333,7 +492,20 @@ class PluginRuntime:
status_code=409, status_code=409,
details={"plugin_id": plugin_id, "permissions": missing_grants}, details={"plugin_id": plugin_id, "permissions": missing_grants},
) )
conflicts = [spec.name for spec in record.tools if self.registry.contains(spec.name)] if (
record.plugin.manifest.backend.type == "mcp"
and not self.allow_unsandboxed_mcp
):
raise ExtensionError(
"MCP_TRUST_APPROVAL_REQUIRED",
"Unsandboxed MCP Hosts are disabled outside development mode.",
status_code=403,
details={"plugin_id": plugin_id},
)
if record.settings_definition is not None:
self.settings.runtime_values(plugin_id, record.settings_definition)
declared_tools = list(record.plugin.manifest.contributes.tools)
conflicts = [name for name in declared_tools if self.registry.contains(name)]
if conflicts: if conflicts:
raise ExtensionError( raise ExtensionError(
"PLUGIN_TOOL_CONFLICT", "PLUGIN_TOOL_CONFLICT",
@@ -341,8 +513,50 @@ class PluginRuntime:
status_code=409, status_code=409,
details={"plugin_id": plugin_id, "tools": conflicts}, details={"plugin_id": plugin_id, "tools": conflicts},
) )
command_conflicts = [
spec.command_id for spec in record.commands if self.commands.contains(spec.command_id)
]
if command_conflicts:
raise ExtensionError(
"PLUGIN_COMMAND_CONFLICT",
"Plugin commands are already registered.",
status_code=409,
details={"plugin_id": plugin_id, "commands": command_conflicts},
)
record.plugin.status = PluginStatus.starting record.plugin.status = PluginStatus.starting
try: try:
if record.plugin.manifest.backend.type == "mcp":
discovered = self._start_mcp(record)
actual = {item.definition.name for item in discovered}
declared = set(declared_tools)
command_targets = {
spec.mcp_tool for spec in record.commands if spec.mcp_tool is not None
}
expected = declared | command_targets
if actual != expected:
raise ExtensionError(
"PLUGIN_CONTRIBUTION_INVALID",
"Discovered MCP tools must exactly match Tool and Command targets.",
details={"declared": sorted(expected), "actual": sorted(actual)},
)
for item in discovered:
if item.definition.name in declared:
self._register_mcp_tool(record, item)
else:
record.mcp_remote_names[item.definition.name] = item.remote_name
record.mcp_command_schemas[item.definition.name] = (
item.definition.parameters
)
for spec in (
command
for command in record.commands
if command.mcp_tool == item.definition.name
):
_validate_mcp_command_target_schema(
item.definition.parameters,
spec.command_id,
)
else:
for spec in record.tools: for spec in record.tools:
arguments_model = _arguments_model(spec) arguments_model = _arguments_model(spec)
@@ -365,19 +579,163 @@ class PluginRuntime:
executor, executor,
) )
record.registered_tools.append(spec.name) record.registered_tools.append(spec.name)
for spec in record.commands:
async def command_executor(
arguments: dict[str, Any],
context: dict[str, Any],
_spec: PluginCommandSpec = spec,
_record: _PluginRecord = record,
) -> PluginCommandEffect:
if (
not _record.plugin.enabled
or _record.plugin.status != PluginStatus.ready
):
raise ExtensionError(
"PLUGIN_COMMAND_NOT_FOUND",
"Plugin command is not available while its Plugin is inactive.",
status_code=404,
details={"command_id": _spec.command_id},
)
settings = (
self.settings.runtime_values(
_record.plugin.manifest.plugin_id,
_record.settings_definition,
)
if _record.settings_definition is not None
else {}
)
def resolve_secret(key: str) -> str | None:
if key not in _spec.secrets:
raise ExtensionError(
"PLUGIN_SECRET_ACCESS_DENIED",
"Command cannot access an undeclared Plugin Secret.",
status_code=403,
details={
"command_id": _spec.command_id,
"key": key,
},
)
if "secrets.use" not in _record.plugin.granted_permissions:
raise ExtensionError(
"PLUGIN_SECRET_ACCESS_DENIED",
"Plugin no longer has permission to access Secret settings.",
status_code=403,
details={"command_id": _spec.command_id, "key": key},
)
if _record.settings_definition is None:
return None
value = self.settings.resolve_secret(
_record.plugin.manifest.plugin_id,
_record.settings_definition,
key,
)
field = next(
item
for item in _record.settings_definition.fields
if item.key == key
)
if field.required and value is None:
raise ExtensionError(
"PLUGIN_SECRET_REQUIRED",
"A required Plugin Secret has not been configured.",
status_code=409,
details={"command_id": _spec.command_id, "key": key},
)
return value
if _spec.mcp_tool is not None:
remote_name = _record.mcp_remote_names[_spec.mcp_tool]
secret_values = {
key: value
for key in _spec.secrets
if (value := resolve_secret(key)) is not None
}
envelope = _mcp_command_envelope(
_spec,
arguments=arguments,
context=context,
settings=settings,
secrets=secret_values,
)
_validate_mcp_command_envelope(
_record.mcp_command_schemas[_spec.mcp_tool],
envelope,
_spec.command_id,
)
try:
effect = await self.mcp.call_tool(
_record.plugin.manifest.plugin_id,
remote_name,
envelope,
request_id=f"command:{uuid4().hex}",
)
except ToolExecutionError as exc:
raise ExtensionError(
exc.code,
"MCP Command target execution failed.",
status_code=502,
details={"command_id": _spec.command_id},
) from exc
try:
return TypeAdapter(PluginCommandEffect).validate_python(effect)
except ValidationError as exc:
raise ExtensionError(
"PLUGIN_COMMAND_RESULT_INVALID",
"MCP Command target returned an invalid effect.",
status_code=502,
details={"command_id": _spec.command_id},
) from exc
return await self.host.execute_command(
_spec.handler,
arguments,
context,
settings,
resolve_secret,
)
self.commands.register(plugin_id, spec, command_executor)
record.registered_commands.append(spec.command_id)
except Exception as exc: except Exception as exc:
# 注册过程必须具备回滚语义,防止半启用插件污染全局工具表。
for name in record.registered_tools: for name in record.registered_tools:
self.registry.unregister(name) self.registry.unregister(name)
record.registered_tools.clear() record.registered_tools.clear()
for command_id in record.registered_commands:
self.commands.unregister(command_id)
record.registered_commands.clear()
record.mcp_remote_names.clear()
record.mcp_command_schemas.clear()
self.mcp.stop(plugin_id)
record.plugin.status = PluginStatus.error record.plugin.status = PluginStatus.error
record.plugin.error_message = str(exc) record.plugin.error_message = _safe_extension_message(exc)
if isinstance(exc, ExtensionError):
raise raise
if isinstance(exc, McpBridgeError):
raise ExtensionError(
exc.code,
exc.message,
status_code=exc.status_code,
details={"plugin_id": plugin_id},
) from exc
raise ExtensionError(
"PLUGIN_HOST_START_FAILED",
record.plugin.error_message,
status_code=503,
details={"plugin_id": plugin_id},
) from exc
record.plugin.enabled = True record.plugin.enabled = True
record.plugin.status = PluginStatus.ready record.plugin.status = PluginStatus.ready
record.plugin.error_message = None record.plugin.error_message = None
return record.plugin.model_copy(deep=True) return record.plugin.model_copy(deep=True)
def set_permissions(self, plugin_id: str, permissions: list[str]) -> Plugin: def set_permissions(self, plugin_id: str, permissions: list[str]) -> Plugin:
with self._lock:
return self._set_permissions(plugin_id, permissions)
def _set_permissions(self, plugin_id: str, permissions: list[str]) -> Plugin:
record = self._record(plugin_id) record = self._record(plugin_id)
requested = set(permissions) requested = set(permissions)
declared = set(record.plugin.manifest.permissions) declared = set(record.plugin.manifest.permissions)
@@ -399,15 +757,174 @@ class PluginRuntime:
return record.plugin.model_copy(deep=True) return record.plugin.model_copy(deep=True)
def disable(self, plugin_id: str) -> Plugin: def disable(self, plugin_id: str) -> Plugin:
with self._lock:
return self._disable(plugin_id)
def _disable(self, plugin_id: str) -> Plugin:
record = self._record(plugin_id) record = self._record(plugin_id)
for name in record.registered_tools: for name in record.registered_tools:
self.registry.unregister(name) self.registry.unregister(name)
record.registered_tools.clear() record.registered_tools.clear()
for command_id in record.registered_commands:
self.commands.unregister(command_id)
record.registered_commands.clear()
record.mcp_remote_names.clear()
record.mcp_command_schemas.clear()
if record.plugin.manifest.backend.type == "mcp":
self.mcp.stop(plugin_id)
record.plugin.enabled = False record.plugin.enabled = False
record.plugin.status = PluginStatus.disabled record.plugin.status = PluginStatus.disabled
return record.plugin.model_copy(deep=True) return record.plugin.model_copy(deep=True)
def get_host_status(self, plugin_id: str) -> PluginHostStatus:
record = self._record(plugin_id)
return self.mcp.status(plugin_id, record.plugin.manifest.backend)
def list_commands(
self, location: PluginCommandLocation | None = None
) -> list[PluginCommand]:
return self.commands.list(location)
async def execute_command(
self,
command_id: str,
arguments: dict[str, Any],
context: PluginCommandContext,
) -> PluginCommandResult:
return await self.commands.execute(command_id, arguments, context)
def get_settings(self, plugin_id: str) -> PluginSettingsSchema:
record = self._record(plugin_id)
definition = self._settings_definition(record)
return self.settings.get(plugin_id, definition)
def update_settings(
self, plugin_id: str, schema_version: int, values: dict[str, Any]
) -> PluginSettingsSchema:
record = self._record(plugin_id)
definition = self._settings_definition(record)
return self.settings.update(plugin_id, definition, schema_version, values)
def put_setting_secret(
self, plugin_id: str, key: str, secret: str
) -> PluginSecretStatus:
record = self._record(plugin_id)
definition = self._settings_definition(record)
return self.settings.put_secret(plugin_id, definition, key, secret)
def delete_setting_secret(self, plugin_id: str, key: str) -> PluginSecretStatus:
record = self._record(plugin_id)
definition = self._settings_definition(record)
return self.settings.delete_secret(plugin_id, definition, key)
def restart_host(self, plugin_id: str) -> PluginHostStatus:
with self._lock:
return self._restart_host(plugin_id)
def _restart_host(self, plugin_id: str) -> PluginHostStatus:
record = self._record(plugin_id)
if record.plugin.manifest.backend.type != "mcp":
raise ExtensionError(
"PLUGIN_HOST_UNAVAILABLE",
"Plugin does not use an MCP Host.",
status_code=409,
details={"plugin_id": plugin_id},
)
if record.plugin.status in {
PluginStatus.installed,
PluginStatus.disabled,
PluginStatus.permission_required,
}:
raise ExtensionError(
"PLUGIN_HOST_UNAVAILABLE",
"Disabled or inactive MCP Plugins must be started with Enable.",
status_code=409,
details={"plugin_id": plugin_id, "status": record.plugin.status.value},
)
for name in record.registered_tools:
self.registry.unregister(name)
record.registered_tools.clear()
for command_id in record.registered_commands:
self.commands.unregister(command_id)
record.registered_commands.clear()
record.mcp_remote_names.clear()
record.mcp_command_schemas.clear()
self.mcp.stop(plugin_id)
record.plugin.enabled = False
record.plugin.status = PluginStatus.installed
record.plugin.error_message = None
self.enable(plugin_id)
return self.get_host_status(plugin_id)
def shutdown(self) -> None:
"""关闭所有隔离 Host;用于 FastAPI lifespan 和测试清理。"""
with self._lock:
for plugin_id, record in list(self._records.items()):
if record.plugin.manifest.backend.type == "mcp":
self.mcp.stop(plugin_id)
def _start_mcp(self, record: _PluginRecord) -> list[McpDiscoveredTool]:
manifest = record.plugin.manifest
return self.mcp.start(
manifest.plugin_id,
manifest.backend,
record.package_path,
manifest.permissions,
self._handle_mcp_unavailable,
)
def _register_mcp_tool(
self, record: _PluginRecord, discovered: McpDiscoveredTool
) -> None:
definition = discovered.definition
arguments_model = _arguments_model_from_schema(
definition.name, definition.parameters
)
plugin_id = record.plugin.manifest.plugin_id
remote_name = discovered.remote_name
async def executor(
arguments: BaseModel,
context: ToolExecutionContext,
) -> Any:
return await self.mcp.call_tool(
plugin_id,
remote_name,
# 省略的可选字段不能被补成 null;显式传入的 null 仍由
# model_fields_set 保留并交给 MCP Server。
arguments.model_dump(exclude_unset=True),
request_id=context.tool_call_id or f"{context.run_id}:{definition.name}",
)
self.registry.register(definition, arguments_model, executor)
record.registered_tools.append(definition.name)
record.mcp_remote_names[definition.name] = remote_name
def _handle_mcp_unavailable(self, plugin_id: str, message: str) -> None:
with self._lock:
record = self._records.get(plugin_id)
if record is None:
return
for name in record.registered_tools:
self.registry.unregister(name)
record.registered_tools.clear()
for command_id in record.registered_commands:
self.commands.unregister(command_id)
record.registered_commands.clear()
record.mcp_remote_names.clear()
record.mcp_command_schemas.clear()
record.plugin.enabled = False
record.plugin.status = PluginStatus.error
record.plugin.error_message = message
def uninstall(self, plugin_id: str, dependent_skills: list[str] | None = None) -> None: def uninstall(self, plugin_id: str, dependent_skills: list[str] | None = None) -> None:
with self._lock:
self._uninstall(plugin_id, dependent_skills)
def _uninstall(
self, plugin_id: str, dependent_skills: list[str] | None = None
) -> None:
record = self._record(plugin_id) record = self._record(plugin_id)
if dependent_skills: if dependent_skills:
raise ExtensionError( raise ExtensionError(
@@ -416,8 +933,14 @@ class PluginRuntime:
status_code=409, status_code=409,
details={"plugin_id": plugin_id, "skills": dependent_skills}, details={"plugin_id": plugin_id, "skills": dependent_skills},
) )
is_mcp = record.plugin.manifest.backend.type == "mcp"
if record.plugin.enabled: if record.plugin.enabled:
self.disable(plugin_id) self.disable(plugin_id)
if is_mcp:
# stop 只结束本次进程并保留状态供故障诊断;真正卸载时必须连同
# 历史状态一起遗忘,避免同 ID 重装继承旧协商信息。
self.mcp.remove(plugin_id)
self.settings.remove_plugin(plugin_id)
del self._records[plugin_id] del self._records[plugin_id]
def _record(self, plugin_id: str) -> _PluginRecord: def _record(self, plugin_id: str) -> _PluginRecord:
@@ -439,6 +962,52 @@ class PluginRuntime:
except ValidationError as exc: except ValidationError as exc:
raise _manifest_error("plugin tool", exc) from exc raise _manifest_error("plugin tool", exc) from exc
@staticmethod
def _load_commands(root: Path) -> list[PluginCommandSpec]:
path = root / "commands.yaml"
if not path.exists():
return []
raw = _read_yaml(path)
items = raw.get("commands", [])
if not isinstance(items, list):
raise ExtensionError(
"EXTENSION_MANIFEST_INVALID",
"Invalid plugin command manifest: commands must be an array.",
)
try:
return [
PluginCommandSpec.model_validate(item)
for item in items
]
except ValidationError as exc:
raise _manifest_error("plugin command", exc) from exc
@staticmethod
def _load_settings(root: Path) -> PluginSettingsDefinition | None:
path = root / "settings.yaml"
if not path.exists():
return None
raw = _read_yaml(path)
try:
return PluginSettingsDefinition.model_validate(raw)
except ValidationError as exc:
raise ExtensionError(
"PLUGIN_SETTINGS_SCHEMA_INVALID",
"Invalid Plugin settings schema.",
details={"errors": exc.errors(include_url=False)},
) from exc
@staticmethod
def _settings_definition(record: _PluginRecord) -> PluginSettingsDefinition:
if record.settings_definition is None:
raise ExtensionError(
"PLUGIN_SETTINGS_NOT_FOUND",
"Plugin does not contribute a Settings section.",
status_code=404,
details={"plugin_id": record.plugin.manifest.plugin_id},
)
return record.settings_definition
def _package_dir(package_path: str | Path) -> Path: def _package_dir(package_path: str | Path) -> Path:
root = Path(package_path).expanduser().resolve() root = Path(package_path).expanduser().resolve()
@@ -494,34 +1063,85 @@ def _manifest_error(kind: str, exc: ValidationError) -> ExtensionError:
def _arguments_model(spec: DeclarativeToolSpec) -> type[BaseModel]: def _arguments_model(spec: DeclarativeToolSpec) -> type[BaseModel]:
schema = spec.parameters or {"type": "object", "properties": {}} schema = spec.parameters or {"type": "object", "properties": {}}
return _arguments_model_from_schema(spec.name, schema)
def _mcp_command_envelope(
spec: PluginCommandSpec,
*,
arguments: dict[str, Any],
context: dict[str, Any],
settings: dict[str, Any],
secrets: dict[str, str],
) -> dict[str, Any]:
return {
"_notesagent": {
"command_id": spec.command_id,
"arguments": arguments,
"context": context,
"settings": settings,
"secrets": secrets,
}
}
def _validate_mcp_command_envelope(
schema: dict[str, Any],
envelope: dict[str, Any],
command_id: str,
) -> None:
"""执行前用目标 Tool Schema 校验包含真实业务数据的宿主信封。"""
try:
Draft202012Validator(schema).validate(envelope)
except JsonSchemaValidationError as exc:
raise ExtensionError(
"PLUGIN_COMMAND_TARGET_SCHEMA_MISMATCH",
"MCP Command envelope does not match the target inputSchema.",
status_code=502,
details={"command_id": command_id, "path": list(exc.path)},
) from exc
def _validate_mcp_command_target_schema(
schema: dict[str, Any], command_id: str
) -> None:
"""启用时只检查稳定信封入口,避免用伪造业务值误判合法 Schema。"""
properties = schema.get("properties")
envelope_schema = (
properties.get("_notesagent") if isinstance(properties, dict) else None
)
if not isinstance(envelope_schema, dict) or envelope_schema.get("type") != "object":
raise ExtensionError(
"PLUGIN_CONTRIBUTION_INVALID",
"MCP Command target inputSchema must directly declare "
"_notesagent with type object.",
details={"command_id": command_id},
)
def _arguments_model_from_schema(
tool_name: str, schema: dict[str, Any]
) -> type[BaseModel]:
if schema.get("type", "object") != "object": if schema.get("type", "object") != "object":
raise ExtensionError("PLUGIN_TOOL_SCHEMA_INVALID", "Tool parameters must be an object schema.") raise ExtensionError("PLUGIN_TOOL_SCHEMA_INVALID", "Tool parameters must be an object schema.")
properties = schema.get("properties", {}) model_name = "PluginArgs_" + re.sub(r"\W+", "_", tool_name)
required = set(schema.get("required", [])) # 完整 JSON Schema 已在 ToolRegistry 中先行校验。参数载体不重复声明字段,
fields: dict[str, tuple[Any, Any]] = {} # 从而完整保留 model_dump、连字符键、联合类型和动态属性等合法 JSON 键值。
types = { return create_model(model_name, __config__=ConfigDict(extra="allow"))
"string": str,
"number": float,
"integer": int,
"boolean": bool,
"array": list[Any],
"object": dict[str, Any],
}
for name, field_schema in properties.items():
annotation = types.get(field_schema.get("type"), Any)
fields[name] = (annotation, ... if name in required else None)
model_name = "PluginArgs_" + re.sub(r"\W+", "_", spec.name)
return create_model(model_name, __config__=ConfigDict(extra="forbid"), **fields)
def _validate_tool_schema(spec: DeclarativeToolSpec) -> None: def _validate_tool_schema(spec: DeclarativeToolSpec) -> None:
schema = spec.parameters or {"type": "object", "properties": {}} schema = spec.parameters or {"type": "object", "properties": {}}
try: try:
Draft202012Validator.check_schema(schema) Draft202012Validator.check_schema(schema)
except SchemaError as exc: reject_external_schema_references(schema)
except (SchemaReferenceError, SchemaError) as exc:
message = exc.message if isinstance(exc, SchemaError) else str(exc)
raise ExtensionError( raise ExtensionError(
"PLUGIN_TOOL_SCHEMA_INVALID", "PLUGIN_TOOL_SCHEMA_INVALID",
f"Invalid JSON Schema for tool {spec.name}: {exc.message}", f"Invalid JSON Schema for tool {spec.name}: {message}",
details={"tool": spec.name}, details={"tool": spec.name},
) from exc ) from exc
if schema.get("type", "object") != "object" or not isinstance( if schema.get("type", "object") != "object" or not isinstance(
@@ -532,3 +1152,30 @@ def _validate_tool_schema(spec: DeclarativeToolSpec) -> None:
"Tool parameters must be an object schema with object properties.", "Tool parameters must be an object schema with object properties.",
details={"tool": spec.name}, details={"tool": spec.name},
) )
def _validate_backend(manifest: PluginManifest) -> None:
backend = manifest.backend
if backend.type == "mcp":
if backend.transport != "stdio":
raise ExtensionError(
"MCP_CAPABILITY_UNSUPPORTED",
"Phase C MCP Plugins must use stdio transport.",
status_code=501,
)
if not backend.command or not backend.command.strip():
raise ExtensionError(
"EXTENSION_MANIFEST_INVALID",
"MCP stdio backend requires a command.",
)
elif backend.command is not None or backend.args:
raise ExtensionError(
"EXTENSION_MANIFEST_INVALID",
"Only MCP stdio backends may declare command or args.",
)
def _safe_extension_message(exc: Exception) -> str:
if isinstance(exc, (ExtensionError, McpBridgeError)):
return exc.message
return f"Plugin Host operation failed: {type(exc).__name__}."
+13
View File
@@ -1,19 +1,32 @@
from contextlib import asynccontextmanager
from fastapi import FastAPI from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from starlette.exceptions import HTTPException as StarletteHttpException from starlette.exceptions import HTTPException as StarletteHttpException
from app.config import get_settings from app.config import get_settings
from app.container import container
from app.errors import ApiError, api_error_handler, http_error_handler, validation_error_handler from app.errors import ApiError, api_error_handler, http_error_handler, validation_error_handler
from app.routes import router as api_router from app.routes import router as api_router
from app.schemas import HealthResponse, ServiceStatusResponse from app.schemas import HealthResponse, ServiceStatusResponse
settings = get_settings() settings = get_settings()
@asynccontextmanager
async def lifespan(_: FastAPI):
yield
# 第三方 MCP Server 必须跟随 AI Core 退出,不能遗留孤儿进程。
container.plugins.shutdown()
container.mcp_servers.shutdown()
app = FastAPI( app = FastAPI(
title=settings.name, title=settings.name,
version=settings.version, version=settings.version,
description="AI 笔记软件的本地 AI Core 与 Agent Core 服务。", description="AI 笔记软件的本地 AI Core 与 Agent Core 服务。",
lifespan=lifespan,
) )
app.add_middleware( app.add_middleware(
+91 -9
View File
@@ -1,16 +1,19 @@
"""Provider 凭据解析及本地加密存储。"""
import json import json
import os import os
import re import re
import threading import threading
from pathlib import Path from pathlib import Path
from typing import Protocol from typing import ClassVar, Protocol
from cryptography.fernet import Fernet, InvalidToken from cryptography.fernet import Fernet, InvalidToken
from app.config import get_settings from app.config import get_settings
_CREDENTIAL_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$") _CREDENTIAL_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")
_PLUGIN_CREDENTIAL_PREFIX = "plugin."
_MCP_CREDENTIAL_PREFIX = "mcp."
class CredentialStoreError(RuntimeError): class CredentialStoreError(RuntimeError):
@@ -21,10 +24,21 @@ class CredentialResolver(Protocol):
def resolve(self, credential_id: str | None) -> str | None: ... def resolve(self, credential_id: str | None) -> str | None: ...
def validate_provider_credential_id(credential_id: str | None) -> None:
"""阻止 Provider 和通用凭据 API 跨入 Plugin 私有命名空间。"""
if credential_id and credential_id.casefold().startswith(_PLUGIN_CREDENTIAL_PREFIX):
raise CredentialStoreError(
"Credential namespace is reserved for Plugin settings."
)
if credential_id and credential_id.casefold().startswith(_MCP_CREDENTIAL_PREFIX):
raise CredentialStoreError("Credential namespace is reserved for MCP settings.")
class EnvironmentCredentialResolver: class EnvironmentCredentialResolver:
"""解析由桌面 Host 注入 Sidecar 进程的临时凭证上下文。""" """解析由桌面 Host 注入 Sidecar 进程的临时凭证上下文。"""
_development_aliases = { _development_aliases: ClassVar[dict[str, str]] = {
"openai": "OPENAI_API_KEY", "openai": "OPENAI_API_KEY",
"deepseek": "DEEPSEEK_API_KEY", "deepseek": "DEEPSEEK_API_KEY",
} }
@@ -43,6 +57,8 @@ class EnvironmentCredentialResolver:
class EncryptedCredentialStore: class EncryptedCredentialStore:
"""将本地开发凭据作为 Fernet 密文存储,Provider 使用时按 ID 解密。""" """将本地开发凭据作为 Fernet 密文存储,Provider 使用时按 ID 解密。"""
# TODO(security): 桌面 Host 接入后将主密钥迁移到系统钥匙串/凭据保险库。
def __init__(self) -> None: def __init__(self) -> None:
self._lock = threading.RLock() self._lock = threading.RLock()
@@ -70,11 +86,14 @@ class EncryptedCredentialStore:
try: try:
return Fernet(environment_key.encode("ascii")) return Fernet(environment_key.encode("ascii"))
except (ValueError, UnicodeEncodeError) as exc: except (ValueError, UnicodeEncodeError) as exc:
raise CredentialStoreError("APP_CREDENTIAL_MASTER_KEY is invalid.") from exc raise CredentialStoreError(
"APP_CREDENTIAL_MASTER_KEY is invalid."
) from exc
key_path.parent.mkdir(parents=True, exist_ok=True) key_path.parent.mkdir(parents=True, exist_ok=True)
self._restrict(key_path.parent, 0o700) self._restrict(key_path.parent, 0o700)
if not key_path.exists(): if not key_path.exists():
# 先写临时文件再原子替换,避免异常退出留下半截主密钥。
temporary = key_path.with_suffix(".tmp") temporary = key_path.with_suffix(".tmp")
temporary.write_bytes(Fernet.generate_key()) temporary.write_bytes(Fernet.generate_key())
self._restrict(temporary, 0o600) self._restrict(temporary, 0o600)
@@ -86,7 +105,9 @@ class EncryptedCredentialStore:
try: try:
return Fernet(key_path.read_bytes().strip()) return Fernet(key_path.read_bytes().strip())
except (OSError, ValueError) as exc: except (OSError, ValueError) as exc:
raise CredentialStoreError("Credential master key cannot be loaded.") from exc raise CredentialStoreError(
"Credential master key cannot be loaded."
) from exc
def _read_tokens(self) -> dict[str, str]: def _read_tokens(self) -> dict[str, str]:
_, store_path = self._paths() _, store_path = self._paths()
@@ -95,25 +116,40 @@ class EncryptedCredentialStore:
try: try:
data = json.loads(store_path.read_text(encoding="utf-8")) data = json.loads(store_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc: except (OSError, json.JSONDecodeError) as exc:
raise CredentialStoreError("Encrypted credential store cannot be loaded.") from exc raise CredentialStoreError(
"Encrypted credential store cannot be loaded."
) from exc
if not isinstance(data, dict) or not all( if not isinstance(data, dict) or not all(
isinstance(key, str) and isinstance(value, str) for key, value in data.items() isinstance(key, str) and isinstance(value, str)
for key, value in data.items()
): ):
raise CredentialStoreError("Encrypted credential store has an invalid format.") raise CredentialStoreError(
"Encrypted credential store has an invalid format."
)
return data return data
def _write_tokens(self, tokens: dict[str, str]) -> None: def _write_tokens(self, tokens: dict[str, str]) -> None:
_, store_path = self._paths() _, store_path = self._paths()
temporary = store_path.with_suffix(".tmp")
try:
store_path.parent.mkdir(parents=True, exist_ok=True) store_path.parent.mkdir(parents=True, exist_ok=True)
self._restrict(store_path.parent, 0o700) self._restrict(store_path.parent, 0o700)
temporary = store_path.with_suffix(".tmp")
temporary.write_text( temporary.write_text(
json.dumps(tokens, ensure_ascii=True, sort_keys=True), json.dumps(tokens, ensure_ascii=True, sort_keys=True),
encoding="utf-8", encoding="utf-8",
) )
self._restrict(temporary, 0o600) self._restrict(temporary, 0o600)
# 凭据表同样使用原子替换,确保并发读取只会看到完整 JSON。
temporary.replace(store_path) temporary.replace(store_path)
self._restrict(store_path, 0o600) self._restrict(store_path, 0o600)
except OSError as exc:
try:
temporary.unlink(missing_ok=True)
except OSError:
pass
raise CredentialStoreError(
"Encrypted credential store cannot be written."
) from exc
def put(self, credential_id: str, secret: str) -> None: def put(self, credential_id: str, secret: str) -> None:
self._validate_id(credential_id) self._validate_id(credential_id)
@@ -152,14 +188,60 @@ class EncryptedCredentialStore:
self._write_tokens(tokens) self._write_tokens(tokens)
return removed return removed
def delete_many(self, credential_ids: list[str]) -> set[str]:
"""用一次原子替换删除多个凭据,避免插件卸载只删除部分 Secret。"""
for credential_id in credential_ids:
self._validate_id(credential_id)
with self._lock:
tokens = self._read_tokens()
removed = {
credential_id
for credential_id in credential_ids
if credential_id in tokens
}
if removed:
for credential_id in removed:
del tokens[credential_id]
self._write_tokens(tokens)
return removed
def move_many(self, replacements: dict[str, str]) -> None:
"""原子迁移凭据 ID,直接移动密文且不覆盖已经写入的新凭据。"""
for old_id, new_id in replacements.items():
self._validate_id(old_id)
self._validate_id(new_id)
with self._lock:
tokens = self._read_tokens()
changed = False
for old_id, new_id in replacements.items():
if old_id != new_id and old_id in tokens:
tokens.setdefault(new_id, tokens.pop(old_id))
changed = True
if changed:
self._write_tokens(tokens)
class ChainedCredentialResolver: class ChainedCredentialResolver:
def __init__(self, *resolvers: CredentialResolver) -> None: def __init__(self, *resolvers: CredentialResolver) -> None:
self._resolvers = resolvers self._resolvers = resolvers
def resolve(self, credential_id: str | None) -> str | None: def resolve(self, credential_id: str | None) -> str | None:
# 顺序即优先级:调用方可让 Host 注入值覆盖本地开发凭据。
for resolver in self._resolvers: for resolver in self._resolvers:
value = resolver.resolve(credential_id) value = resolver.resolve(credential_id)
if value: if value:
return value return value
return None return None
class ProviderCredentialResolver:
"""Provider 专用防御层,避免配置绕过 HTTP 校验读取 Plugin Secret。"""
def __init__(self, delegate: CredentialResolver) -> None:
self._delegate = delegate
def resolve(self, credential_id: str | None) -> str | None:
validate_provider_credential_id(credential_id)
return self._delegate.resolve(credential_id)
+4 -2
View File
@@ -1,6 +1,6 @@
from app.contracts import ModelCapability, ProviderConfig, ProviderPreset, ProviderType from app.contracts import ModelCapability, ProviderConfig, ProviderPreset, ProviderType
from app.providers.base import ModelProvider from app.providers.base import ModelProvider
from app.providers.credentials import CredentialResolver from app.providers.credentials import CredentialResolver, ProviderCredentialResolver
from app.providers.ollama import OllamaProvider from app.providers.ollama import OllamaProvider
from app.providers.openai_compatible import OpenAICompatibleProvider from app.providers.openai_compatible import OpenAICompatibleProvider
@@ -11,7 +11,9 @@ class UnsupportedProviderError(ValueError):
class ProviderFactory: class ProviderFactory:
def __init__(self, credentials: CredentialResolver) -> None: def __init__(self, credentials: CredentialResolver) -> None:
self.credentials = credentials # ProviderFactory 是所有可配置 Provider 的创建边界,在此统一禁止
# Provider 借用 Plugin Secret 引用,避免调用方漏包安全 Resolver。
self.credentials = ProviderCredentialResolver(credentials)
def build(self, config: ProviderConfig) -> ModelProvider: def build(self, config: ProviderConfig) -> ModelProvider:
if config.provider_type in { if config.provider_type in {
+134 -12
View File
@@ -63,6 +63,14 @@ class FtsHit:
bm25: float bm25: float
@dataclass(frozen=True, slots=True)
class NoteLocation:
note_id: str
title: str
file_path: str
folder: str
def replace_note_metadata( def replace_note_metadata(
*, *,
conn: sqlite3.Connection, conn: sqlite3.Connection,
@@ -219,11 +227,61 @@ def fts_search(match: str, limit: int = 100) -> list[FtsHit]:
conn.close() conn.close()
def fts_search_page( def list_note_locations(*, conn: sqlite3.Connection | None = None) -> list[NoteLocation]:
"""返回 Workspace 构树和目录事务所需的最小笔记位置集合。"""
owns = conn is None
conn = conn or connect()
try:
rows = conn.execute(
"SELECT note_id, title, file_path, folder FROM notes ORDER BY file_path"
).fetchall()
return [
NoteLocation(
note_id=row["note_id"],
title=row["title"],
file_path=row["file_path"],
folder=row["folder"],
)
for row in rows
]
finally:
if owns:
conn.close()
def update_note_location(
*, *,
conn: sqlite3.Connection,
note_id: str,
title: str,
file_path: str,
folder: str,
updated_at: datetime,
) -> None:
"""更新文件位置和展示标题;Block/FTS/向量内容不变,无需重新生成。"""
cursor = conn.execute(
"""
UPDATE notes
SET title = ?, file_path = ?, folder = ?, updated_at = ?
WHERE note_id = ?
""",
(title, file_path, folder, _iso(updated_at), note_id),
)
if cursor.rowcount != 1:
raise LookupError(note_id)
_FTS_FROM = """
FROM blocks_fts
JOIN blocks AS b ON b.block_id = blocks_fts.block_id
JOIN notes AS n ON n.note_id = b.note_id
"""
def _fts_where(
match: str, match: str,
limit: int,
offset: int,
folders: list[str], folders: list[str],
note_ids: list[str], note_ids: list[str],
tags: list[str], tags: list[str],
@@ -231,8 +289,11 @@ def fts_search_page(
created_to: datetime | None, created_to: datetime | None,
updated_from: datetime | None, updated_from: datetime | None,
updated_to: datetime | None, updated_to: datetime | None,
) -> tuple[list[FtsHit], int]: ) -> tuple[str, list[object]]:
"""执行带元数据过滤的 FTS 精确分页,并返回过滤后的完整命中数。""" """构建 FTS 过滤 WHERE 子句(不含 WHERE 关键字),返回 (where_sql, params)。
fts_search_page fts_score_bounds 共用保证计数与取数口径一致
"""
where = ["blocks_fts MATCH ?"] where = ["blocks_fts MATCH ?"]
params: list[object] = [match] params: list[object] = [match]
@@ -263,22 +324,44 @@ def fts_search_page(
where.append(f"julianday({column}) <= julianday(?)") where.append(f"julianday({column}) <= julianday(?)")
params.append(_iso(upper)) params.append(_iso(upper))
from_sql = """ return " AND ".join(where), params
FROM blocks_fts
JOIN blocks AS b ON b.block_id = blocks_fts.block_id
JOIN notes AS n ON n.note_id = b.note_id def fts_search_page(
*,
match: str,
limit: int,
offset: int,
folders: list[str],
note_ids: list[str],
tags: list[str],
created_from: datetime | None,
created_to: datetime | None,
updated_from: datetime | None,
updated_to: datetime | None,
bm25_max: float | None = None,
) -> tuple[list[FtsHit], int]:
"""执行带元数据过滤的 FTS 精确分页,并返回过滤后的完整命中数。
bm25_max 非空时按 bm25 截止值过滤用于阈值过滤的精确分页计数与取数同口径
""" """
where_sql = " AND ".join(where) where_sql, params = _fts_where(
match, folders, note_ids, tags,
created_from, created_to, updated_from, updated_to,
)
if bm25_max is not None:
where_sql += " AND bm25(blocks_fts) <= ?"
params.append(bm25_max)
conn = connect() conn = connect()
try: try:
total = conn.execute( total = conn.execute(
f"SELECT COUNT(*) {from_sql} WHERE {where_sql}", params f"SELECT COUNT(*) {_FTS_FROM} WHERE {where_sql}", params
).fetchone()[0] ).fetchone()[0]
rows = conn.execute( rows = conn.execute(
f""" f"""
SELECT blocks_fts.block_id, blocks_fts.note_id, bm25(blocks_fts) AS rank SELECT blocks_fts.block_id, blocks_fts.note_id, bm25(blocks_fts) AS rank
{from_sql} {_FTS_FROM}
WHERE {where_sql} WHERE {where_sql}
ORDER BY rank ORDER BY rank
LIMIT ? OFFSET ? LIMIT ? OFFSET ?
@@ -294,6 +377,45 @@ def fts_search_page(
conn.close() conn.close()
def fts_score_bounds(
*,
match: str,
folders: list[str],
note_ids: list[str],
tags: list[str],
created_from: datetime | None,
created_to: datetime | None,
updated_from: datetime | None,
updated_to: datetime | None,
) -> tuple[float, float] | None:
"""返回 metadata 过滤后的 FTS 命中集里 bm25 的 (min, max),无命中时返回 None。
用于阈值过滤min-max 归一化是 bm25 的线性函数据此可把阈值换算为 bm25 截止值
"""
where_sql, params = _fts_where(
match, folders, note_ids, tags,
created_from, created_to, updated_from, updated_to,
)
conn = connect()
try:
# bm25() 不能作为聚合函数参数,也不能用在被聚合的子查询里;改用 ORDER BY 取首尾两行
lo_row = conn.execute(
f"SELECT bm25(blocks_fts) AS rank {_FTS_FROM} WHERE {where_sql}"
" ORDER BY rank ASC LIMIT 1",
params,
).fetchone()
if lo_row is None or lo_row["rank"] is None:
return None
hi_row = conn.execute(
f"SELECT bm25(blocks_fts) AS rank {_FTS_FROM} WHERE {where_sql}"
" ORDER BY rank DESC LIMIT 1",
params,
).fetchone()
return (float(lo_row["rank"]), float(hi_row["rank"]))
finally:
conn.close()
def get_block_hits(block_ids: list[str]) -> list[BlockHit]: def get_block_hits(block_ids: list[str]) -> list[BlockHit]:
if not block_ids: if not block_ids:
return [] return []
+2
View File
@@ -19,6 +19,7 @@ class EmbeddingProvider(Protocol):
"""统一 Embedding 接口(与文档一致)。""" """统一 Embedding 接口(与文档一致)。"""
model_id: str model_id: str
version: str
dim: int dim: int
async def embed_documents(self, texts: list[str]) -> list[list[float]]: ... async def embed_documents(self, texts: list[str]) -> list[list[float]]: ...
@@ -33,6 +34,7 @@ class HashEmbeddingProvider:
""" """
model_id = "hash-v1" model_id = "hash-v1"
version = "1"
dim = EMBEDDING_DIM dim = EMBEDDING_DIM
async def embed_documents(self, texts: list[str]) -> list[list[float]]: async def embed_documents(self, texts: list[str]) -> list[list[float]]:
+62 -11
View File
@@ -56,7 +56,7 @@ class RetrievalEngine:
# 候选池至少覆盖本次请求的 offset+limit,保证分页能取到目标页;设上限防内存失控 # 候选池至少覆盖本次请求的 offset+limit,保证分页能取到目标页;设上限防内存失控
window = min(request.offset + request.limit, MAX_CANDIDATE_POOL) window = min(request.offset + request.limit, MAX_CANDIDATE_POOL)
pool_size = max(CANDIDATE_POOL, window) pool_size = max(CANDIDATE_POOL, window)
# 带过滤时放大召回;FTS 则一次性取全量命中(≤FTS_FETCH_LIMIT)避免截断漏召回 # 带过滤时放大召回,缓解「先截断候选池再过滤」造成的漏召回
recall = min(pool_size * OVERSCAN_FACTOR, MAX_CANDIDATE_POOL) if has_filters else pool_size recall = min(pool_size * OVERSCAN_FACTOR, MAX_CANDIDATE_POOL) if has_filters else pool_size
# 1. 按模式收集候选(FTS 与 Vector 各产出「按相关性降序」的 block_id 列表) # 1. 按模式收集候选(FTS 与 Vector 各产出「按相关性降序」的 block_id 列表)
@@ -84,7 +84,7 @@ class RetrievalEngine:
elif request.mode == SearchMode.vector: elif request.mode == SearchMode.vector:
candidate_scores = vec_scores candidate_scores = vec_scores
else: # hybridRRF 融合 else: # hybridRRF 融合
candidate_scores = rrf_fuse([fts_ranked, vec_ranked]) candidate_scores = rrf_fuse([fts_ranked, vec_ranked], k=request.rrf_k)
if not candidate_scores: if not candidate_scores:
return self._empty(request) return self._empty(request)
@@ -97,14 +97,23 @@ class RetrievalEngine:
if not filtered: if not filtered:
return self._empty(request) return self._empty(request)
# 4. 排序 / 精排 # 4. 排序 / 精排:hybrid 先按融合分预排序,再对前 rerank_candidates 个候选做精排,
# 剩余候选按融合分排在精排结果之后;rerank=False 时跳过精排直接按融合分排序。
if request.mode == SearchMode.hybrid: if request.mode == SearchMode.hybrid:
pre_sorted = sorted(filtered, key=lambda h: -candidate_scores[h.block_id])
if request.rerank:
limit = request.rerank_candidates
pool = pre_sorted if limit is None else pre_sorted[:limit]
rest = [] if limit is None else pre_sorted[limit:]
candidates = [ candidates = [
RankedCandidate(block_id=h.block_id, score=candidate_scores[h.block_id], text=h.content) RankedCandidate(block_id=h.block_id, score=candidate_scores[h.block_id], text=h.content)
for h in filtered for h in pool
] ]
ranked = await self.reranker.rerank(request.query, candidates) ranked = await self.reranker.rerank(request.query, candidates)
ordered = [(c.block_id, c.score) for c in ranked] ordered = [(c.block_id, c.score) for c in ranked]
ordered += [(h.block_id, candidate_scores[h.block_id]) for h in rest]
else:
ordered = [(h.block_id, candidate_scores[h.block_id]) for h in pre_sorted]
else: else:
ordered = sorted( ordered = sorted(
((h.block_id, candidate_scores[h.block_id]) for h in filtered), ((h.block_id, candidate_scores[h.block_id]) for h in filtered),
@@ -112,8 +121,10 @@ class RetrievalEngine:
) )
ordered = normalize_scores(ordered) ordered = normalize_scores(ordered)
# score_threshold:归一化后过滤低分结果(默认 0 不过滤)
ordered = [(bid, score) for bid, score in ordered if score >= request.score_threshold]
# 5. 分页:total = 过滤后候选集大小。fts 已取全量(≤FTS_FETCH_LIMIT)故为真实命中数; # 5. 分页:total = 过滤后候选集大小。fts 走数据库精确分页,total 为真实命中数;
# vector/hybrid 为 KNN 候选集,无全局 total。 # vector/hybrid 为 KNN 候选集,无全局 total。
total = len(ordered) total = len(ordered)
page = ordered[request.offset : request.offset + request.limit] page = ordered[request.offset : request.offset + request.limit]
@@ -126,11 +137,41 @@ class RetrievalEngine:
) )
def _search_fts(self, request: SearchRequest) -> SearchResponse: def _search_fts(self, request: SearchRequest) -> SearchResponse:
"""FTS 专用路径:过滤、COUNT 与分页全部在 SQLite 中完成。""" """FTS 专用路径:在数据库侧完成过滤、计数与分页,不取全量后再截断。
阈值过滤时min-max 归一化是 bm25 的线性函数据此把 score_threshold 换算为
bm25 截止值bm25_max使过滤计数与分页口径一致无阈值时走数据库原生分页
total 始终为过滤后的真实命中数不再受固定截断影响
"""
match = match_query(request.query) match = match_query(request.query)
if not match: if not match:
return self._empty(request) return self._empty(request)
bounds = repository.fts_score_bounds(
match=match,
folders=request.folders,
note_ids=request.note_ids,
tags=request.tags,
created_from=request.created_from,
created_to=request.created_to,
updated_from=request.updated_from,
updated_to=request.updated_to,
)
if bounds is None:
return self._empty(request)
lo, hi = bounds
span = hi - lo
bm25_max: float | None = None
if request.score_threshold > 0:
if span == 0:
# 全部命中 bm25 相同,归一化后皆为 1.0;阈值超过 1.0 时无命中
if request.score_threshold > 1.0:
return self._empty(request)
else:
# norm = (hi - bm25) / spannorm >= threshold ⟺ bm25 <= hi - threshold * span
bm25_max = hi - request.score_threshold * span
fts_hits, total = repository.fts_search_page( fts_hits, total = repository.fts_search_page(
match=match, match=match,
limit=request.limit, limit=request.limit,
@@ -142,19 +183,29 @@ class RetrievalEngine:
created_to=request.created_to, created_to=request.created_to,
updated_from=request.updated_from, updated_from=request.updated_from,
updated_to=request.updated_to, updated_to=request.updated_to,
bm25_max=bm25_max,
) )
if not fts_hits: if not fts_hits:
# 本页无结果:offset 越过末页时 total 仍为真实命中数(>0),需保留而非归零
return SearchResponse( return SearchResponse(
query=request.query, query=request.query,
mode=request.mode, mode=request.mode,
items=[],
page=PageMeta(total=total, limit=request.limit, offset=request.offset), page=PageMeta(total=total, limit=request.limit, offset=request.offset),
) )
hits = {h.block_id: h for h in repository.get_block_hits([hit.block_id for hit in fts_hits])} # 分数按全局 bm25 上下界归一化(与取全量后 normalize_scores 等价),保证跨页一致
ordered = normalize_scores( span = hi - lo
[(hit.block_id, -hit.bm25) for hit in fts_hits if hit.block_id in hits] if span == 0:
) ordered = [(hit.block_id, 1.0) for hit in fts_hits]
items = [self._build_result(hits[block_id], request, score) for block_id, score in ordered] else:
ordered = [(hit.block_id, round((hi - hit.bm25) / span, 6)) for hit in fts_hits]
hits = {h.block_id: h for h in repository.get_block_hits([bid for bid, _ in ordered])}
items = [
self._build_result(hits[block_id], request, score)
for block_id, score in ordered
if block_id in hits
]
return SearchResponse( return SearchResponse(
query=request.query, query=request.query,
mode=request.mode, mode=request.mode,
+2
View File
@@ -24,6 +24,7 @@ class RerankerProvider(Protocol):
"""统一 Reranker 接口:输入候选块,输出按相关性重排后的候选块。""" """统一 Reranker 接口:输入候选块,输出按相关性重排后的候选块。"""
model_id: str model_id: str
version: str
async def rerank(self, query: str, candidates: list[RankedCandidate]) -> list[RankedCandidate]: ... async def rerank(self, query: str, candidates: list[RankedCandidate]) -> list[RankedCandidate]: ...
@@ -32,6 +33,7 @@ class LexicalReranker:
"""轻量精排:query 与块正文的词面重叠度,与归一化后的原始分数加权求和。""" """轻量精排:query 与块正文的词面重叠度,与归一化后的原始分数加权求和。"""
model_id = "lexical-v1" model_id = "lexical-v1"
version = "1"
def __init__(self, lexical_weight: float = 0.5) -> None: def __init__(self, lexical_weight: float = 0.5) -> None:
self.lexical_weight = lexical_weight self.lexical_weight = lexical_weight
+8
View File
@@ -35,6 +35,7 @@ class VectorStore(Protocol):
async def upsert(self, records: list[VectorRecord]) -> None: ... async def upsert(self, records: list[VectorRecord]) -> None: ...
async def delete(self, ids: list[str]) -> None: ... async def delete(self, ids: list[str]) -> None: ...
async def search(self, vector: list[float], *, top_k: int) -> list[VectorHit]: ... async def search(self, vector: list[float], *, top_k: int) -> list[VectorHit]: ...
async def count(self) -> int: ...
class SqliteVecStore: class SqliteVecStore:
@@ -92,3 +93,10 @@ class SqliteVecStore:
conn.execute("DELETE FROM vec_blocks") conn.execute("DELETE FROM vec_blocks")
finally: finally:
conn.close() conn.close()
async def count(self) -> int:
conn = connect()
try:
return conn.execute("SELECT COUNT(*) FROM vec_blocks").fetchone()[0]
finally:
conn.close()
+610 -31
View File
@@ -1,34 +1,67 @@
import asyncio
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from datetime import datetime, timezone from datetime import datetime, timezone
from uuid import uuid4 from uuid import uuid4
from fastapi import APIRouter, Query from fastapi import APIRouter, Header, Query
from fastapi.responses import StreamingResponse from fastapi.responses import StreamingResponse
from app.agent import AgentCapacityError, AgentRunNotFoundError
from app.container import container
from app.contracts import ( from app.contracts import (
AgentRun, AgentRun,
AgentRunCreateRequest, AgentRunCreateRequest,
AgentRunListResponse, AgentRunListResponse,
AgentTraceResponse,
ChatRequest, ChatRequest,
BenchmarkDatasetListResponse,
BenchmarkEventType,
BenchmarkKind,
BenchmarkReport,
BenchmarkRun,
BenchmarkRunListResponse,
BenchmarkStatus,
RAGRunRequest,
CredentialStatus, CredentialStatus,
CredentialWriteRequest, CredentialWriteRequest,
ExtensionInstallRequest, ExtensionInstallRequest,
FolderCreateRequest,
FolderDeleteRequest,
FolderRenameRequest,
IndexJob, IndexJob,
IndexRebuildRequest, IndexRebuildRequest,
IndexStatus, IndexStatus,
McpServer,
McpServerCreateRequest,
McpServerListResponse,
McpServerSecretStatus,
McpServerSecretWriteRequest,
McpServerTrustRequest,
McpServerUpdateRequest,
McpToolSummaryListResponse,
ModelEvent, ModelEvent,
ModelEventType, ModelEventType,
Note, Note,
NoteCreateRequest, NoteCreateRequest,
NoteListResponse, NoteListResponse,
NoteMoveRequest, NoteMoveRequest,
NoteRenameRequest,
NoteUpdateRequest, NoteUpdateRequest,
OperationResponse, OperationResponse,
PageMeta, PageMeta,
PermissionDecisionRequest, PermissionDecisionRequest,
Plugin, Plugin,
PluginCommandExecuteRequest,
PluginCommandListResponse,
PluginCommandLocation,
PluginCommandResult,
PluginHostStatus,
PluginListResponse, PluginListResponse,
PluginPermissionGrantRequest, PluginPermissionGrantRequest,
PluginSecretStatus,
PluginSecretWriteRequest,
PluginSettingsSchema,
PluginSettingsUpdateRequest,
ProviderConfig, ProviderConfig,
ProviderCreateRequest, ProviderCreateRequest,
ProviderListResponse, ProviderListResponse,
@@ -48,27 +81,59 @@ from app.contracts import (
ToolListResponse, ToolListResponse,
TranscriptionJob, TranscriptionJob,
TranscriptionRequest, TranscriptionRequest,
WorkspaceEntry,
WorkspaceInfo,
WorkspaceOpenRequest,
WorkspaceSnapshot,
) )
from app.agent import AgentCapacityError, AgentRunNotFoundError from app.agent import AgentCapacityError, AgentRunNotFoundError
from app.benchmarks import datasets as benchmark_datasets
from app.benchmarks import service as benchmark_service
from app.container import container from app.container import container
from app.errors import ApiError from app.errors import ApiError
from app.extensions import ExtensionError from app.extensions import ExtensionError
from app.providers.registry import ProviderNotFoundError from app.extensions.mcp_registry import McpRegistryError
from app.providers.factory import UnsupportedProviderError
from app.providers.base import ProviderError from app.providers.base import ProviderError
from app.providers.credentials import CredentialStoreError from app.providers.credentials import (
CredentialStoreError,
validate_provider_credential_id,
)
from app.providers.factory import UnsupportedProviderError
from app.providers.registry import ProviderNotFoundError
from app.retrieval.engine import engine from app.retrieval.engine import engine
from app.services import index_service, note_service, task_service, transcription_service from app.services import (
index_service,
note_service,
task_service,
transcription_service,
workspace_service,
)
router = APIRouter(prefix="/api") router = APIRouter(prefix="/api")
async def mcp_call_async(operation):
"""Even registry reads can wait on lifecycle locks; keep all MCP work off the event loop."""
try:
return await asyncio.to_thread(operation)
except McpRegistryError as exc:
raise ApiError(exc.status_code, exc.code, exc.message) from exc
def utc_now() -> datetime: def utc_now() -> datetime:
return datetime.now(timezone.utc) return datetime.now(timezone.utc)
def as_sse(event: str, payload: str) -> str: def validate_public_credential_id(credential_id: str | None) -> None:
return f"event: {event}\ndata: {payload}\n\n" try:
validate_provider_credential_id(credential_id)
except CredentialStoreError as exc:
raise ApiError(422, "CREDENTIAL_NAMESPACE_RESERVED", str(exc)) from exc
def as_sse(event: str, payload: str, *, event_id: int | None = None) -> str:
id_line = f"id: {event_id}\n" if event_id is not None else ""
return f"{id_line}event: {event}\ndata: {payload}\n\n"
def provider_or_404(provider_id: str): def provider_or_404(provider_id: str):
@@ -114,6 +179,50 @@ def extension_call(operation):
raise ApiError(exc.status_code, exc.code, exc.message, exc.details) from exc raise ApiError(exc.status_code, exc.code, exc.message, exc.details) from exc
async def extension_call_async(operation):
"""进程启动/关闭可能等待 stdio Host,移出 FastAPI 事件循环。"""
try:
return await asyncio.to_thread(operation)
except ExtensionError as exc:
raise ApiError(exc.status_code, exc.code, exc.message, exc.details) from exc
# Workspace (single configured Vault in Web development mode)
@router.get("/workspace", response_model=WorkspaceInfo, tags=["Workspace"])
async def get_workspace() -> WorkspaceInfo:
return workspace_service.get_workspace_info()
@router.post("/workspace/open", response_model=WorkspaceSnapshot, tags=["Workspace"])
async def open_workspace(request: WorkspaceOpenRequest) -> WorkspaceSnapshot:
return await workspace_service.open_workspace(request.path)
@router.get("/workspace/tree", response_model=list[WorkspaceEntry], tags=["Workspace"])
async def get_workspace_tree() -> list[WorkspaceEntry]:
return workspace_service.get_workspace_tree()
@router.post("/workspace/folders", response_model=WorkspaceEntry, tags=["Workspace"])
async def create_workspace_folder(request: FolderCreateRequest) -> WorkspaceEntry:
return await workspace_service.create_folder(request.parent, request.name)
@router.post(
"/workspace/folders/rename", response_model=WorkspaceEntry, tags=["Workspace"]
)
async def rename_workspace_folder(request: FolderRenameRequest) -> WorkspaceEntry:
return await workspace_service.rename_folder(request.path, request.new_name)
@router.post(
"/workspace/folders/delete", response_model=OperationResponse, tags=["Workspace"]
)
async def delete_workspace_folder(request: FolderDeleteRequest) -> OperationResponse:
return await workspace_service.delete_folder(request.path)
# Notes # Notes
@router.get("/notes", response_model=NoteListResponse, tags=["Notes"]) @router.get("/notes", response_model=NoteListResponse, tags=["Notes"])
async def list_notes( async def list_notes(
@@ -122,14 +231,21 @@ async def list_notes(
folder: str | None = None, folder: str | None = None,
tag: str | None = None, tag: str | None = None,
) -> NoteListResponse: ) -> NoteListResponse:
items, total = note_service.list_notes(limit=limit, offset=offset, folder=folder, tag=tag) items, total = note_service.list_notes(
return NoteListResponse(items=items, page=PageMeta(total=total, limit=limit, offset=offset)) limit=limit, offset=offset, folder=folder, tag=tag
)
return NoteListResponse(
items=items, page=PageMeta(total=total, limit=limit, offset=offset)
)
@router.post("/notes", response_model=Note, tags=["Notes"]) @router.post("/notes", response_model=Note, tags=["Notes"])
async def create_note(request: NoteCreateRequest) -> Note: async def create_note(request: NoteCreateRequest) -> Note:
return await note_service.create_note( return await note_service.create_note(
title=request.title, markdown=request.markdown, folder=request.folder, tags=request.tags title=request.title,
markdown=request.markdown,
folder=request.folder,
tags=request.tags,
) )
@@ -137,7 +253,9 @@ async def create_note(request: NoteCreateRequest) -> Note:
async def get_note(note_id: str) -> Note: async def get_note(note_id: str) -> Note:
note = await note_service.get_note(note_id) note = await note_service.get_note(note_id)
if note is None: if note is None:
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}) raise ApiError(
404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}
)
return note return note
@@ -151,7 +269,9 @@ async def update_note(note_id: str, request: NoteUpdateRequest) -> Note:
@router.delete("/notes/{note_id}", response_model=OperationResponse, tags=["Notes"]) @router.delete("/notes/{note_id}", response_model=OperationResponse, tags=["Notes"])
async def delete_note(note_id: str) -> OperationResponse: async def delete_note(note_id: str) -> OperationResponse:
if not await note_service.delete_note(note_id): if not await note_service.delete_note(note_id):
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}) raise ApiError(
404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}
)
return OperationResponse(status="completed", resource_id=note_id, message="deleted") return OperationResponse(status="completed", resource_id=note_id, message="deleted")
@@ -160,6 +280,11 @@ async def move_note(note_id: str, request: NoteMoveRequest) -> Note:
return await note_service.move_note(note_id, folder=request.folder) return await note_service.move_note(note_id, folder=request.folder)
@router.post("/notes/{note_id}/rename", response_model=Note, tags=["Notes"])
async def rename_note(note_id: str, request: NoteRenameRequest) -> Note:
return await note_service.rename_note(note_id, file_name=request.file_name)
# Retrieval and chat # Retrieval and chat
@router.post("/search", response_model=SearchResponse, tags=["Search"]) @router.post("/search", response_model=SearchResponse, tags=["Search"])
async def search_notes(request: SearchRequest) -> SearchResponse: async def search_notes(request: SearchRequest) -> SearchResponse:
@@ -190,7 +315,9 @@ async def chat(request: ChatRequest) -> StreamingResponse:
data={"code": "PROVIDER_ERROR", "message": str(exc)}, data={"code": "PROVIDER_ERROR", "message": str(exc)},
timestamp=utc_now(), timestamp=utc_now(),
) )
done = ModelEvent(event=ModelEventType.done, sequence=1, timestamp=utc_now()) done = ModelEvent(
event=ModelEventType.done, sequence=1, timestamp=utc_now()
)
yield as_sse(error.event.value, error.model_dump_json()) yield as_sse(error.event.value, error.model_dump_json())
yield as_sse(done.event.value, done.model_dump_json()) yield as_sse(done.event.value, done.model_dump_json())
@@ -260,16 +387,65 @@ async def cancel_agent_run(run_id: str) -> OperationResponse:
}, },
tags=["Agent"], tags=["Agent"],
) )
async def agent_events(run_id: str) -> StreamingResponse: async def agent_events(
run_id: str,
after_sequence: int | None = Query(default=None, ge=-1),
last_event_id: str | None = Header(default=None, alias="Last-Event-ID"),
) -> StreamingResponse:
agent_run_or_404(run_id) agent_run_or_404(run_id)
cursor = after_sequence
if cursor is None and last_event_id is not None:
try:
cursor = int(last_event_id)
except ValueError as exc:
raise ApiError(
400,
"TRACE_CURSOR_INVALID",
"Last-Event-ID must be an integer sequence.",
{"last_event_id": last_event_id},
) from exc
if cursor < -1:
raise ApiError(
400,
"TRACE_CURSOR_INVALID",
"Last-Event-ID must be greater than or equal to -1.",
)
cursor = cursor if cursor is not None else -1
async def stream() -> AsyncIterator[str]: async def stream() -> AsyncIterator[str]:
async for event in container.agent.events(run_id): async for event in container.agent.events(run_id, after_sequence=cursor):
yield as_sse(event.event.value, event.model_dump_json()) yield as_sse(
event.event.value,
event.model_dump_json(),
event_id=event.sequence,
)
return StreamingResponse(stream(), media_type="text/event-stream") return StreamingResponse(stream(), media_type="text/event-stream")
@router.get(
"/agent/runs/{run_id}/trace",
response_model=AgentTraceResponse,
tags=["Agent"],
)
async def get_agent_trace(
run_id: str,
after_sequence: int = Query(default=-1, ge=-1),
limit: int = Query(default=200, ge=1, le=500),
) -> AgentTraceResponse:
try:
return container.agent.get_trace(
run_id, after_sequence=after_sequence, limit=limit
)
except AgentRunNotFoundError as exc:
raise ApiError(
404,
"AGENT_RUN_NOT_FOUND",
f"Agent run does not exist: {run_id}",
{"run_id": run_id},
) from exc
@router.post( @router.post(
"/agent/runs/{run_id}/permissions/{request_id}", "/agent/runs/{run_id}/permissions/{request_id}",
response_model=OperationResponse, response_model=OperationResponse,
@@ -302,9 +478,7 @@ async def list_skills() -> SkillListResponse:
return SkillListResponse(items=container.skills.list()) return SkillListResponse(items=container.skills.list())
@router.get( @router.get("/skills/{skill_id}", response_model=Skill, tags=["Skills"])
"/skills/{skill_id}", response_model=Skill, tags=["Skills"]
)
async def get_skill(skill_id: str) -> Skill: async def get_skill(skill_id: str) -> Skill:
return extension_call(lambda: container.skills.get(skill_id)) return extension_call(lambda: container.skills.get(skill_id))
@@ -344,7 +518,120 @@ async def disable_skill(skill_id: str) -> Skill:
) )
async def uninstall_skill(skill_id: str) -> OperationResponse: async def uninstall_skill(skill_id: str) -> OperationResponse:
extension_call(lambda: container.skills.uninstall(skill_id)) extension_call(lambda: container.skills.uninstall(skill_id))
return OperationResponse(status="completed", resource_id=skill_id, message="uninstalled") return OperationResponse(
status="completed", resource_id=skill_id, message="uninstalled"
)
# Independent MCP Server Registry
@router.get("/mcp/servers", response_model=McpServerListResponse, tags=["MCP Servers"])
async def list_mcp_servers() -> McpServerListResponse:
return McpServerListResponse(items=await mcp_call_async(container.mcp_servers.list))
@router.post(
"/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"]
)
async def create_mcp_server(request: McpServerCreateRequest) -> McpServer:
return await mcp_call_async(lambda: container.mcp_servers.create(request))
@router.get("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"])
async def get_mcp_server(server_id: str) -> McpServer:
return await mcp_call_async(lambda: container.mcp_servers.get(server_id))
@router.get(
"/mcp/servers/{server_id}/tools",
response_model=McpToolSummaryListResponse,
tags=["MCP Servers"],
)
async def list_mcp_server_tools(server_id: str) -> McpToolSummaryListResponse:
return McpToolSummaryListResponse(
items=await mcp_call_async(lambda: container.mcp_servers.list_tools(server_id))
)
@router.put("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"])
async def update_mcp_server(
server_id: str, request: McpServerUpdateRequest
) -> McpServer:
return await mcp_call_async(
lambda: container.mcp_servers.update(server_id, request)
)
@router.delete(
"/mcp/servers/{server_id}", response_model=OperationResponse, tags=["MCP Servers"]
)
async def delete_mcp_server(server_id: str) -> OperationResponse:
await mcp_call_async(lambda: container.mcp_servers.delete(server_id))
return OperationResponse(
status="completed", resource_id=server_id, message="deleted"
)
@router.post(
"/mcp/servers/{server_id}/trust", response_model=McpServer, tags=["MCP Servers"]
)
async def trust_mcp_server(server_id: str, request: McpServerTrustRequest) -> McpServer:
return await mcp_call_async(
lambda: container.mcp_servers.trust(server_id, request.command_digest)
)
@router.post(
"/mcp/servers/{server_id}/test", response_model=McpServer, tags=["MCP Servers"]
)
async def test_mcp_server(server_id: str) -> McpServer:
return await mcp_call_async(lambda: container.mcp_servers.test(server_id))
@router.post(
"/mcp/servers/{server_id}/enable", response_model=McpServer, tags=["MCP Servers"]
)
async def enable_mcp_server(server_id: str) -> McpServer:
return await mcp_call_async(lambda: container.mcp_servers.enable(server_id))
@router.post(
"/mcp/servers/{server_id}/disable", response_model=McpServer, tags=["MCP Servers"]
)
async def disable_mcp_server(server_id: str) -> McpServer:
return await mcp_call_async(lambda: container.mcp_servers.disable(server_id))
@router.put(
"/mcp/servers/{server_id}/secrets/{key}",
response_model=McpServerSecretStatus,
tags=["MCP Servers"],
)
async def put_mcp_server_secret(
server_id: str,
key: str,
request: McpServerSecretWriteRequest,
kind: str = Query(default="environment", pattern="^(environment|header)$"),
) -> McpServerSecretStatus:
return await mcp_call_async(
lambda: container.mcp_servers.put_secret(
server_id, key, request.secret.get_secret_value(), kind=kind
)
)
@router.delete(
"/mcp/servers/{server_id}/secrets/{key}",
response_model=McpServerSecretStatus,
tags=["MCP Servers"],
)
async def delete_mcp_server_secret(
server_id: str,
key: str,
kind: str = Query(default="environment", pattern="^(environment|header)$"),
) -> McpServerSecretStatus:
return await mcp_call_async(
lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind)
)
# Plugins # Plugins
@@ -378,7 +665,7 @@ async def install_plugin(request: ExtensionInstallRequest) -> Plugin:
tags=["Plugins"], tags=["Plugins"],
) )
async def enable_plugin(plugin_id: str) -> Plugin: async def enable_plugin(plugin_id: str) -> Plugin:
return extension_call(lambda: container.plugins.enable(plugin_id)) return await extension_call_async(lambda: container.plugins.enable(plugin_id))
@router.post( @router.post(
@@ -387,7 +674,7 @@ async def enable_plugin(plugin_id: str) -> Plugin:
tags=["Plugins"], tags=["Plugins"],
) )
async def disable_plugin(plugin_id: str) -> Plugin: async def disable_plugin(plugin_id: str) -> Plugin:
return extension_call(lambda: container.plugins.disable(plugin_id)) return await extension_call_async(lambda: container.plugins.disable(plugin_id))
@router.put( @router.put(
@@ -398,11 +685,37 @@ async def disable_plugin(plugin_id: str) -> Plugin:
async def set_plugin_permissions( async def set_plugin_permissions(
plugin_id: str, request: PluginPermissionGrantRequest plugin_id: str, request: PluginPermissionGrantRequest
) -> Plugin: ) -> Plugin:
return extension_call( return await extension_call_async(
lambda: container.plugins.set_permissions(plugin_id, request.permissions) lambda: container.plugins.set_permissions(plugin_id, request.permissions)
) )
@router.get(
"/plugins/{plugin_id}/host",
response_model=PluginHostStatus,
tags=["Plugins"],
)
async def get_plugin_host_status(plugin_id: str) -> PluginHostStatus:
return extension_call(lambda: container.plugins.get_host_status(plugin_id))
@router.post(
"/plugins/{plugin_id}/host/restart",
response_model=OperationResponse,
status_code=202,
tags=["Plugins"],
)
async def restart_plugin_host(plugin_id: str) -> OperationResponse:
status = await extension_call_async(
lambda: container.plugins.restart_host(plugin_id)
)
return OperationResponse(
status="accepted",
resource_id=plugin_id,
message=f"Plugin Host status: {status.status.value}",
)
@router.delete( @router.delete(
"/plugins/{plugin_id}", "/plugins/{plugin_id}",
response_model=OperationResponse, response_model=OperationResponse,
@@ -410,9 +723,93 @@ async def set_plugin_permissions(
) )
async def uninstall_plugin(plugin_id: str) -> OperationResponse: async def uninstall_plugin(plugin_id: str) -> OperationResponse:
plugin = extension_call(lambda: container.plugins.get(plugin_id)) plugin = extension_call(lambda: container.plugins.get(plugin_id))
dependent_skills = container.skills.depending_on_tools(plugin.manifest.contributes.tools) dependent_skills = container.skills.depending_on_tools(
extension_call(lambda: container.plugins.uninstall(plugin_id, dependent_skills)) plugin.manifest.contributes.tools
return OperationResponse(status="completed", resource_id=plugin_id, message="uninstalled") )
await extension_call_async(
lambda: container.plugins.uninstall(plugin_id, dependent_skills)
)
return OperationResponse(
status="completed", resource_id=plugin_id, message="uninstalled"
)
# Plugin Command / Settings Contributions
@router.get(
"/plugin-contributions/commands",
response_model=PluginCommandListResponse,
tags=["Plugins"],
)
async def list_plugin_commands(
location: PluginCommandLocation | None = Query(default=None),
) -> PluginCommandListResponse:
return PluginCommandListResponse(items=container.plugins.list_commands(location))
@router.post(
"/plugin-contributions/commands/{command_id}/execute",
response_model=PluginCommandResult,
tags=["Plugins"],
)
async def execute_plugin_command(
command_id: str, request: PluginCommandExecuteRequest
) -> PluginCommandResult:
try:
return await container.plugins.execute_command(
command_id, request.arguments, request.context
)
except ExtensionError as exc:
raise ApiError(exc.status_code, exc.code, exc.message, exc.details) from exc
@router.get(
"/plugins/{plugin_id}/settings",
response_model=PluginSettingsSchema,
tags=["Plugins"],
)
async def get_plugin_settings(plugin_id: str) -> PluginSettingsSchema:
return extension_call(lambda: container.plugins.get_settings(plugin_id))
@router.put(
"/plugins/{plugin_id}/settings",
response_model=PluginSettingsSchema,
tags=["Plugins"],
)
async def update_plugin_settings(
plugin_id: str, request: PluginSettingsUpdateRequest
) -> PluginSettingsSchema:
return extension_call(
lambda: container.plugins.update_settings(
plugin_id, request.schema_version, request.values
)
)
@router.put(
"/plugins/{plugin_id}/settings/{key}/secret",
response_model=PluginSecretStatus,
tags=["Plugins"],
)
async def put_plugin_setting_secret(
plugin_id: str, key: str, request: PluginSecretWriteRequest
) -> PluginSecretStatus:
return extension_call(
lambda: container.plugins.put_setting_secret(
plugin_id, key, request.secret.get_secret_value()
)
)
@router.delete(
"/plugins/{plugin_id}/settings/{key}/secret",
response_model=PluginSecretStatus,
tags=["Plugins"],
)
async def delete_plugin_setting_secret(plugin_id: str, key: str) -> PluginSecretStatus:
return extension_call(
lambda: container.plugins.delete_setting_secret(plugin_id, key)
)
# Providers # Providers
@@ -422,6 +819,7 @@ async def uninstall_plugin(plugin_id: str) -> OperationResponse:
tags=["Providers"], tags=["Providers"],
) )
async def get_credential_status(credential_id: str) -> CredentialStatus: async def get_credential_status(credential_id: str) -> CredentialStatus:
validate_public_credential_id(credential_id)
try: try:
configured = container.credentials.has(credential_id) configured = container.credentials.has(credential_id)
except CredentialStoreError as exc: except CredentialStoreError as exc:
@@ -437,6 +835,7 @@ async def get_credential_status(credential_id: str) -> CredentialStatus:
async def put_credential( async def put_credential(
credential_id: str, request: CredentialWriteRequest credential_id: str, request: CredentialWriteRequest
) -> CredentialStatus: ) -> CredentialStatus:
validate_public_credential_id(credential_id)
try: try:
container.credentials.put(credential_id, request.api_key.get_secret_value()) container.credentials.put(credential_id, request.api_key.get_secret_value())
except CredentialStoreError as exc: except CredentialStoreError as exc:
@@ -450,6 +849,7 @@ async def put_credential(
tags=["Providers"], tags=["Providers"],
) )
async def delete_credential(credential_id: str) -> CredentialStatus: async def delete_credential(credential_id: str) -> CredentialStatus:
validate_public_credential_id(credential_id)
try: try:
container.credentials.delete(credential_id) container.credentials.delete(credential_id)
except CredentialStoreError as exc: except CredentialStoreError as exc:
@@ -486,6 +886,7 @@ async def get_provider(provider_id: str) -> ProviderConfig:
tags=["Providers"], tags=["Providers"],
) )
async def create_provider(request: ProviderCreateRequest) -> ProviderConfig: async def create_provider(request: ProviderCreateRequest) -> ProviderConfig:
validate_public_credential_id(request.credential_id)
config = ProviderConfig( config = ProviderConfig(
provider_id=f"provider_{uuid4().hex}", provider_id=f"provider_{uuid4().hex}",
provider_type=request.provider_type, provider_type=request.provider_type,
@@ -518,7 +919,9 @@ async def update_provider(
) -> ProviderConfig: ) -> ProviderConfig:
current = configurable_provider_or_404(provider_id).config current = configurable_provider_or_404(provider_id).config
if provider_id == "mock": if provider_id == "mock":
raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified.") raise ApiError(
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified."
)
fields = request.model_fields_set fields = request.model_fields_set
if ("name" in fields and request.name is None) or ( if ("name" in fields and request.name is None) or (
"enabled" in fields and request.enabled is None "enabled" in fields and request.enabled is None
@@ -529,6 +932,8 @@ async def update_provider(
"name and enabled cannot be null when explicitly provided.", "name and enabled cannot be null when explicitly provided.",
) )
updates = {name: getattr(request, name) for name in fields} updates = {name: getattr(request, name) for name in fields}
if "credential_id" in fields:
validate_public_credential_id(request.credential_id)
config = ProviderConfig.model_validate( config = ProviderConfig.model_validate(
{**current.model_dump(mode="python"), **updates} {**current.model_dump(mode="python"), **updates}
) )
@@ -545,7 +950,9 @@ async def update_provider(
async def delete_provider(provider_id: str) -> OperationResponse: async def delete_provider(provider_id: str) -> OperationResponse:
configurable_provider_or_404(provider_id) configurable_provider_or_404(provider_id)
if provider_id == "mock": if provider_id == "mock":
raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted.") raise ApiError(
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted."
)
container.providers.unregister(provider_id) container.providers.unregister(provider_id)
return OperationResponse(status="completed", resource_id=provider_id) return OperationResponse(status="completed", resource_id=provider_id)
@@ -588,6 +995,7 @@ async def list_provider_models(provider_id: str) -> ProviderModelsResponse:
async def test_provider(request: ProviderTestRequest) -> ProviderTestResponse: async def test_provider(request: ProviderTestRequest) -> ProviderTestResponse:
registered = configurable_provider_or_404(request.provider_id) registered = configurable_provider_or_404(request.provider_id)
if request.credential_context_id: if request.credential_context_id:
validate_public_credential_id(request.credential_context_id)
temporary_config = registered.config.model_copy( temporary_config = registered.config.model_copy(
update={"credential_id": request.credential_context_id, "enabled": True} update={"credential_id": request.credential_context_id, "enabled": True}
) )
@@ -622,7 +1030,9 @@ async def create_task(request: TaskCreateRequest) -> Task:
async def get_task(task_id: str) -> Task: async def get_task(task_id: str) -> Task:
task = task_service.get_task(task_id) task = task_service.get_task(task_id)
if task is None: if task is None:
raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}) raise ApiError(
404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}
)
return task return task
@@ -638,7 +1048,9 @@ async def update_task(task_id: str, request: TaskUpdateRequest) -> Task:
) )
async def delete_task(task_id: str) -> OperationResponse: async def delete_task(task_id: str) -> OperationResponse:
if not task_service.delete_task(task_id): if not task_service.delete_task(task_id):
raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}) raise ApiError(
404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}
)
return OperationResponse(status="completed", resource_id=task_id, message="deleted") return OperationResponse(status="completed", resource_id=task_id, message="deleted")
@@ -688,5 +1100,172 @@ async def rebuild_index(request: IndexRebuildRequest) -> IndexJob:
async def get_index_job(job_id: str) -> IndexJob: async def get_index_job(job_id: str) -> IndexJob:
job = index_service.get_job(job_id) job = index_service.get_job(job_id)
if job is None: if job is None:
raise ApiError(404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id}) raise ApiError(
404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id}
)
return job return job
# Benchmark
@router.get(
"/benchmarks/datasets",
response_model=BenchmarkDatasetListResponse,
tags=["Benchmark"],
)
async def list_benchmark_datasets(
kind: BenchmarkKind = Query(default=BenchmarkKind.rag),
) -> BenchmarkDatasetListResponse:
return BenchmarkDatasetListResponse(items=benchmark_datasets.list_datasets(kind))
@router.post(
"/benchmarks/rag/runs",
response_model=BenchmarkRun,
status_code=202,
tags=["Benchmark"],
)
async def create_rag_benchmark(request: RAGRunRequest) -> BenchmarkRun:
return await benchmark_service.create_rag_run(request)
@router.get(
"/benchmarks/runs",
response_model=BenchmarkRunListResponse,
tags=["Benchmark"],
)
async def list_benchmark_runs(
kind: BenchmarkKind | None = Query(default=None),
status: BenchmarkStatus | None = Query(default=None),
limit: int = Query(default=50, ge=1, le=100),
offset: int = Query(default=0, ge=0),
) -> BenchmarkRunListResponse:
items, total = benchmark_service.list_runs(
kind=kind, status=status, limit=limit, offset=offset
)
return BenchmarkRunListResponse(
items=items, page=PageMeta(total=total, limit=limit, offset=offset)
)
@router.get(
"/benchmarks/runs/{run_id}",
response_model=BenchmarkRun,
tags=["Benchmark"],
)
async def get_benchmark_run(run_id: str) -> BenchmarkRun:
run = benchmark_service.get_run(run_id)
if run is None:
raise ApiError(
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id}
)
return run
@router.post(
"/benchmarks/runs/{run_id}/cancel",
response_model=OperationResponse,
tags=["Benchmark"],
)
async def cancel_benchmark_run(run_id: str) -> OperationResponse:
run = benchmark_service.cancel_run(run_id)
if run is None:
raise ApiError(
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id}
)
return OperationResponse(
status="accepted",
resource_id=run_id,
message=f"Benchmark run status: {run.status.value}",
)
@router.get(
"/benchmarks/runs/{run_id}/events",
response_class=StreamingResponse,
responses={
200: {
"description": "BenchmarkEvent Server-Sent Events stream",
"content": {"text/event-stream": {}},
}
},
tags=["Benchmark"],
)
async def benchmark_events(
run_id: str,
after_sequence: int = Query(default=-1, ge=-1),
last_event_id: str | None = Header(default=None, alias="Last-Event-ID"),
) -> StreamingResponse:
if benchmark_service.get_run(run_id) is None:
raise ApiError(
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id}
)
# SSE 断线重连:Last-Event-ID 优先于 after_sequence,用于从上次收到的事件继续
cursor = after_sequence
if last_event_id is not None:
try:
cursor = int(last_event_id)
except ValueError as exc:
raise ApiError(
400,
"BENCHMARK_EVENT_CURSOR_INVALID",
"Last-Event-ID must be an integer sequence.",
{"last_event_id": last_event_id},
) from exc
if cursor < -1:
raise ApiError(
400,
"BENCHMARK_EVENT_CURSOR_INVALID",
"Last-Event-ID must be greater than or equal to -1.",
)
async def stream() -> AsyncIterator[str]:
# 先订阅(保证订阅之后产生的事件也能收到),再回放历史事件,最后实时输出新事件
terminal = (
BenchmarkEventType.run_completed,
BenchmarkEventType.run_failed,
BenchmarkEventType.run_cancelled,
)
queue = benchmark_service.subscribe(run_id)
try:
last_sequence = cursor
# 回放按订阅时刻的快照长度遍历,避免列表在回放期间被追加;终止事件同样要结束流,
# 防止回放完成后进入实时队列却因序号去重跳过同一终止事件而永久等待。
history = benchmark_service.get_events(run_id)
for index in range(len(history)):
event = history[index]
if event.sequence <= cursor:
continue
yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence)
last_sequence = event.sequence
if event.event in terminal:
return
if queue is None:
return
while True:
event = await queue.get()
if event.sequence <= last_sequence:
continue
yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence)
last_sequence = event.sequence
if event.event in terminal:
return
finally:
if queue is not None:
benchmark_service.unsubscribe(run_id, queue)
return StreamingResponse(stream(), media_type="text/event-stream")
@router.get(
"/benchmarks/runs/{run_id}/report",
response_model=BenchmarkReport,
tags=["Benchmark"],
)
async def get_benchmark_report(run_id: str) -> BenchmarkReport:
report = benchmark_service.get_report(run_id)
if report is None:
raise ApiError(
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark report not found", {"run_id": run_id}
)
return report
+60
View File
@@ -0,0 +1,60 @@
"""共享 JSON Schema 安全约束。"""
from __future__ import annotations
from typing import Any
from urllib.parse import urljoin
from referencing import Registry
from referencing.exceptions import Unresolvable
from referencing.jsonschema import DRAFT202012
_SCHEMA_BASE_URI = "https://notesagent.invalid/local-schema"
class SchemaReferenceError(ValueError):
"""Schema 引用不符合宿主的离线、文档内解析约束。"""
class ExternalSchemaReferenceError(SchemaReferenceError):
def __init__(self, keyword: str, reference: Any) -> None:
super().__init__(f"External JSON Schema reference is not allowed: {reference!r}")
self.keyword = keyword
self.reference = reference
class UnresolvableLocalSchemaReferenceError(SchemaReferenceError):
def __init__(self, reference: str) -> None:
super().__init__(f"Local JSON Schema reference cannot be resolved: {reference!r}")
self.reference = reference
def reject_external_schema_references(schema: Any) -> None:
"""只允许可解析的文档内 Fragment,并按 JSON Schema Resource 作用域解析。"""
root = DRAFT202012.create_resource(schema)
root_uri = urljoin(_SCHEMA_BASE_URI, root.id() or "")
registry = Registry().with_resource(_SCHEMA_BASE_URI, root).crawl()
resolver = registry.resolver(root_uri)
_validate_resource_references(root, resolver)
def _validate_resource_references(resource, resolver: Any) -> None:
contents = resource.contents
if isinstance(contents, dict):
for keyword in ("$ref", "$dynamicRef"):
if keyword not in contents:
continue
reference = contents[keyword]
if not isinstance(reference, str) or not reference.startswith("#"):
raise ExternalSchemaReferenceError(keyword, reference)
try:
resolver.lookup(reference)
except Unresolvable as exc:
raise UnresolvableLocalSchemaReferenceError(reference) from exc
for subresource in resource.subresources():
_validate_resource_references(
subresource,
resolver.in_subresource(subresource),
)
+72 -52
View File
@@ -6,13 +6,11 @@ Markdown 文件是笔记正文的持久化载体(Vault),SQLite/FTS5/向量
from __future__ import annotations from __future__ import annotations
import re
from datetime import datetime, timezone from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from uuid import uuid4 from uuid import uuid4
from app import repository from app import repository
from app.config import get_settings
from app.contracts import Note, NoteBlock, NoteSummary from app.contracts import Note, NoteBlock, NoteSummary
from app.database.db import connect, transaction from app.database.db import connect, transaction
from app.errors import ApiError from app.errors import ApiError
@@ -20,74 +18,40 @@ from app.knowledge.parser import ParsedNote, parse_note
from app.retrieval.embedding import HashEmbeddingProvider from app.retrieval.embedding import HashEmbeddingProvider
from app.retrieval.vectorstore import SqliteVecStore, VectorRecord from app.retrieval.vectorstore import SqliteVecStore, VectorRecord
from app.services.coordination import serialized_vault_mutation from app.services.coordination import serialized_vault_mutation
from app.services.vault_paths import (
normalize_entry_name,
normalize_folder,
resolve_in_vault,
safe_note_filename,
)
# 轻量实现实例(无状态,可直接复用);接入真实模型后替换为对应 Provider # 轻量实现实例(无状态,可直接复用);接入真实模型后替换为对应 Provider
embedding = HashEmbeddingProvider() embedding = HashEmbeddingProvider()
vector_store = SqliteVecStore() vector_store = SqliteVecStore()
def _vault() -> Path:
return get_settings().vault_path
def _safe_name(title: str) -> str:
name = re.sub(r'[\\/:*?"<>|]', "_", title).strip()
return name or "untitled"
def _normalize_folder(folder: str | None) -> str:
"""清洗 folder 为安全的相对目录,拒绝 `..`/`.`/绝对路径/盘符/空字节,防路径逃逸。"""
if not folder:
return ""
if "\x00" in folder:
raise ApiError(400, "INVALID_PATH", "folder must not contain NUL bytes", {"folder": folder})
segments: list[str] = []
for part in re.split(r"[\\/]+", folder):
if part == "":
continue
if part in (".", ".."):
raise ApiError(400, "INVALID_PATH", "folder must not contain '.' or '..'", {"folder": folder})
if ":" in part:
raise ApiError(400, "INVALID_PATH", "folder must be a relative path", {"folder": folder})
segments.append(part)
return "/".join(segments)
def _rel_path(folder: str | None, title: str) -> tuple[str, str]: def _rel_path(folder: str | None, title: str) -> tuple[str, str]:
"""由 folder + title 生成安全的相对路径,返回 (rel_path, 清洗后的 folder)。""" """由 folder + title 生成安全的相对路径,返回 (rel_path, 清洗后的 folder)。"""
clean_folder = _normalize_folder(folder) clean_folder = normalize_folder(folder)
name = _safe_name(title) name = safe_note_filename(title)
if not name.endswith(".md"):
name += ".md"
rel = f"{clean_folder}/{name}" if clean_folder else name rel = f"{clean_folder}/{name}" if clean_folder else name
return rel, clean_folder return rel, clean_folder
def _abs_path(rel_path: str) -> Path:
"""把相对路径解析为 Vault 内的绝对路径;越界即报 400,杜绝路径逃逸。"""
if not rel_path or "\x00" in rel_path:
raise ApiError(400, "INVALID_PATH", "invalid file path", {"file_path": rel_path})
root = _vault().resolve()
candidate = (_vault() / rel_path).resolve()
if not candidate.is_relative_to(root):
raise ApiError(400, "INVALID_PATH", "path escapes vault", {"file_path": rel_path})
return candidate
def _read_markdown(rel_path: str) -> str: def _read_markdown(rel_path: str) -> str:
path = _abs_path(rel_path) path = resolve_in_vault(rel_path)
return path.read_text(encoding="utf-8") if path.exists() else "" return path.read_text(encoding="utf-8") if path.exists() else ""
def _write_markdown(rel_path: str, markdown: str) -> None: def _write_markdown(rel_path: str, markdown: str) -> None:
path = _abs_path(rel_path) path = resolve_in_vault(rel_path)
path.parent.mkdir(parents=True, exist_ok=True) path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(markdown, encoding="utf-8") path.write_text(markdown, encoding="utf-8")
def _create_markdown(rel_path: str, markdown: str) -> None: def _create_markdown(rel_path: str, markdown: str) -> None:
"""排他创建 Markdown;目标已存在时返回资源冲突,不覆盖用户文件。""" """排他创建 Markdown;目标已存在时返回资源冲突,不覆盖用户文件。"""
path = _abs_path(rel_path) path = resolve_in_vault(rel_path)
path.parent.mkdir(parents=True, exist_ok=True) path.parent.mkdir(parents=True, exist_ok=True)
try: try:
with path.open("x", encoding="utf-8") as handle: with path.open("x", encoding="utf-8") as handle:
@@ -102,7 +66,7 @@ def _create_markdown(rel_path: str, markdown: str) -> None:
def _delete_markdown(rel_path: str) -> None: def _delete_markdown(rel_path: str) -> None:
path = _abs_path(rel_path) path = resolve_in_vault(rel_path)
if path.exists(): if path.exists():
path.unlink() path.unlink()
@@ -222,7 +186,7 @@ async def move_note(note_id: str, *, folder: str) -> Note:
if record is None: if record is None:
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}) raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
clean_folder = _normalize_folder(folder) clean_folder = normalize_folder(folder)
filename = Path(record.file_path).name filename = Path(record.file_path).name
new_rel_path = f"{clean_folder}/{filename}" if clean_folder else filename new_rel_path = f"{clean_folder}/{filename}" if clean_folder else filename
if new_rel_path == record.file_path: if new_rel_path == record.file_path:
@@ -230,8 +194,8 @@ async def move_note(note_id: str, *, folder: str) -> Note:
assert note is not None assert note is not None
return note return note
source = _abs_path(record.file_path) source = resolve_in_vault(record.file_path)
target = _abs_path(new_rel_path) target = resolve_in_vault(new_rel_path)
if not source.is_file(): if not source.is_file():
raise ApiError( raise ApiError(
409, "NOTE_FILE_MISSING", "note file is missing from the Vault", 409, "NOTE_FILE_MISSING", "note file is missing from the Vault",
@@ -267,13 +231,69 @@ async def move_note(note_id: str, *, folder: str) -> Note:
) )
@serialized_vault_mutation
async def rename_note(note_id: str, *, file_name: str) -> Note:
"""重命名 Markdown 文件并保留 note_id、Block 与向量身份。"""
record = repository.get_note_record(note_id)
if record is None:
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
normalized = normalize_entry_name(file_name, markdown=True)
source = resolve_in_vault(record.file_path)
folder = normalize_folder(record.folder)
new_file_path = f"{folder}/{normalized}" if folder else normalized
target = resolve_in_vault(new_file_path)
if new_file_path == record.file_path:
note = await get_note(note_id)
assert note is not None
return note
if not source.is_file():
raise ApiError(
409,
"NOTE_FILE_MISSING",
"note file is missing from the Vault",
{"note_id": note_id, "file_path": record.file_path},
)
if target.exists():
raise ApiError(
409,
"RESOURCE_CONFLICT",
"a note already exists with the requested file name",
{"note_id": note_id, "file_path": new_file_path},
)
source.replace(target)
now = datetime.now(timezone.utc)
conn = connect()
try:
with transaction(conn):
repository.update_note_location(
conn=conn,
note_id=note_id,
title=Path(normalized).stem,
file_path=new_file_path,
folder=folder,
updated_at=now,
)
except BaseException:
target.replace(source)
raise
finally:
conn.close()
note = await get_note(note_id)
assert note is not None
return note
@serialized_vault_mutation @serialized_vault_mutation
async def delete_note(note_id: str) -> bool: async def delete_note(note_id: str) -> bool:
record = repository.get_note_record(note_id) record = repository.get_note_record(note_id)
if record is None: if record is None:
return False return False
path = _abs_path(record.file_path) path = resolve_in_vault(record.file_path)
tombstone = path.with_name(f".{path.name}.{uuid4().hex}.deleting") if path.exists() else None tombstone = path.with_name(f".{path.name}.{uuid4().hex}.deleting") if path.exists() else None
if tombstone is not None: if tombstone is not None:
path.replace(tombstone) path.replace(tombstone)
@@ -1,3 +1,5 @@
"""转写适配层;第一阶段消费文本附件或桌面 Host 预生成的旁路文本。"""
from __future__ import annotations from __future__ import annotations
from collections import OrderedDict from collections import OrderedDict
@@ -13,6 +15,7 @@ MAX_JOBS = 100
def create_transcription(attachment_id: str, language: str | None = None) -> TranscriptionJob: def create_transcription(attachment_id: str, language: str | None = None) -> TranscriptionJob:
# TODO(ai-core): 第二阶段接入本地 ASR 队列后,保留相同 Job 契约替换此同步降级实现。
del language # 预生成 transcript 暂不需要语言识别。 del language # 预生成 transcript 暂不需要语言识别。
source = attachment_path(attachment_id) source = attachment_path(attachment_id)
transcript = source if source.suffix.lower() in {".txt", ".md"} else Path(f"{source}.txt") transcript = source if source.suffix.lower() in {".txt", ".md"} else Path(f"{source}.txt")
+78
View File
@@ -0,0 +1,78 @@
"""Vault 相对路径校验;所有文件操作必须先经过本模块。"""
from __future__ import annotations
import re
from pathlib import Path
from app.config import get_settings
from app.errors import ApiError
_INVALID_FILE_CHARS = re.compile(r'[\\/:*?"<>|]')
def normalize_folder(folder: str | None) -> str:
"""返回使用 `/` 的安全相对目录;根目录表示为空字符串。"""
if not folder or folder in {"/", "\\"}:
return ""
if "\x00" in folder:
raise ApiError(400, "INVALID_PATH", "folder must not contain NUL bytes")
segments: list[str] = []
for part in re.split(r"[\\/]+", folder):
if not part:
continue
if part in {".", ".."} or ":" in part:
raise ApiError(
400,
"INVALID_PATH",
"folder must be a relative path without '.' or '..' segments",
{"folder": folder},
)
segments.append(part)
return "/".join(segments)
def normalize_entry_name(name: str, *, markdown: bool = False) -> str:
"""校验单个目录项名称;不静默接受路径分隔符或保留段。"""
value = name.strip()
if not value or value in {".", ".."} or "\x00" in value:
raise ApiError(400, "INVALID_PATH", "entry name is invalid", {"name": name})
if _INVALID_FILE_CHARS.search(value):
raise ApiError(
400,
"INVALID_PATH",
"entry name contains unsupported characters",
{"name": name},
)
if markdown and not value.lower().endswith(".md"):
value += ".md"
return value
def safe_note_filename(title: str) -> str:
"""为创建笔记保留原有的宽松清洗行为。"""
value = _INVALID_FILE_CHARS.sub("_", title).strip() or "untitled"
return value if value.lower().endswith(".md") else f"{value}.md"
def resolve_in_vault(relative_path: str) -> Path:
"""把相对路径解析到当前 Vault,并拒绝符号链接/`..` 导致的越界。"""
if not relative_path or "\x00" in relative_path:
raise ApiError(
400, "INVALID_PATH", "invalid Vault-relative path", {"path": relative_path}
)
root = get_settings().vault_path.resolve()
candidate = (root / relative_path.replace("\\", "/").lstrip("/")).resolve()
if not candidate.is_relative_to(root):
raise ApiError(
400, "INVALID_PATH", "path escapes Vault", {"path": relative_path}
)
return candidate
def relative_to_vault(path: Path) -> str:
return path.resolve().relative_to(get_settings().vault_path.resolve()).as_posix()
+249
View File
@@ -0,0 +1,249 @@
"""Web 联调 Workspace:把单一配置 Vault 映射为前端可用的真实文件树。"""
from __future__ import annotations
import hashlib
import shutil
from datetime import datetime, timezone
from pathlib import Path
from uuid import uuid4
from app import repository
from app.config import get_settings
from app.contracts import (
IndexRebuildRequest,
OperationResponse,
WorkspaceEntry,
WorkspaceInfo,
WorkspaceSnapshot,
)
from app.database.db import connect, transaction
from app.errors import ApiError
from app.retrieval.vectorstore import SqliteVecStore
from app.services import index_service
from app.services.coordination import serialized_vault_mutation
from app.services.vault_paths import normalize_entry_name, normalize_folder, resolve_in_vault
vector_store = SqliteVecStore()
def _entry_id(kind: str, path: str) -> str:
digest = hashlib.sha256(f"{kind}:{path}".encode("utf-8")).hexdigest()[:16]
return f"{kind}_{digest}"
def _disk_markdown_paths() -> set[str]:
root = get_settings().vault_path
if not root.exists():
return set()
resolved_root = root.resolve()
paths: set[str] = set()
for path in root.rglob("*.md"):
if path.is_symlink():
continue
resolved = path.resolve()
if resolved.is_file() and resolved.is_relative_to(resolved_root):
paths.add(resolved.relative_to(resolved_root).as_posix())
return paths
def get_workspace_info() -> WorkspaceInfo:
root = get_settings().vault_path.resolve()
disk_paths = _disk_markdown_paths()
indexed_paths = {item.file_path for item in repository.list_note_locations()}
return WorkspaceInfo(
name=root.name or "Vault",
path=str(root),
file_count=len(disk_paths),
indexed_note_count=len(indexed_paths),
requires_refresh=disk_paths != indexed_paths,
)
def _tree(directory: Path, locations: dict[str, repository.NoteLocation]) -> list[WorkspaceEntry]:
if not directory.exists():
return []
root = get_settings().vault_path.resolve()
entries: list[WorkspaceEntry] = []
children = sorted(
directory.iterdir(), key=lambda item: (not item.is_dir(), item.name.casefold())
)
for child in children:
if child.name.startswith(".") or child.is_symlink():
continue
resolved = child.resolve()
if not resolved.is_relative_to(root):
continue
relative = resolved.relative_to(root).as_posix()
public_path = f"/{relative}"
if resolved.is_dir():
entries.append(
WorkspaceEntry(
entry_id=_entry_id("folder", relative),
name=child.name,
path=public_path,
type="folder",
children=_tree(resolved, locations),
)
)
elif resolved.is_file() and child.suffix.lower() == ".md":
location = locations.get(relative)
entries.append(
WorkspaceEntry(
entry_id=location.note_id if location else _entry_id("file", relative),
note_id=location.note_id if location else None,
name=child.name,
path=public_path,
type="file",
)
)
return entries
def get_workspace_tree() -> list[WorkspaceEntry]:
locations = {item.file_path: item for item in repository.list_note_locations()}
return _tree(get_settings().vault_path.resolve(), locations)
async def open_workspace(requested_path: str | None) -> WorkspaceSnapshot:
"""打开当前配置 Vault;发现未索引文件时先执行一次安全全量刷新。"""
root = get_settings().vault_path.resolve()
if requested_path and Path(requested_path).resolve() != root:
raise ApiError(
409,
"WORKSPACE_PATH_MISMATCH",
"Web development mode can only open the backend configured Vault.",
{"configured_path": str(root)},
)
root.mkdir(parents=True, exist_ok=True)
info = get_workspace_info()
if info.requires_refresh:
await index_service.rebuild(IndexRebuildRequest())
info = get_workspace_info()
return WorkspaceSnapshot(workspace=info, items=get_workspace_tree())
@serialized_vault_mutation
async def create_folder(parent: str, name: str) -> WorkspaceEntry:
clean_parent = normalize_folder(parent)
clean_name = normalize_entry_name(name)
relative = f"{clean_parent}/{clean_name}" if clean_parent else clean_name
target = resolve_in_vault(relative)
if not clean_parent:
get_settings().vault_path.mkdir(parents=True, exist_ok=True)
if target.exists():
raise ApiError(
409, "RESOURCE_CONFLICT", "folder already exists", {"path": relative}
)
if not target.parent.is_dir():
raise ApiError(
404,
"RESOURCE_NOT_FOUND",
"parent folder not found",
{"parent": clean_parent},
)
target.mkdir(parents=False)
return WorkspaceEntry(
entry_id=_entry_id("folder", relative),
name=clean_name,
path=f"/{relative}",
type="folder",
)
@serialized_vault_mutation
async def rename_folder(path: str, new_name: str) -> WorkspaceEntry:
old_folder = normalize_folder(path)
if not old_folder:
raise ApiError(400, "INVALID_PATH", "the Vault root cannot be renamed")
clean_name = normalize_entry_name(new_name)
parent = Path(old_folder).parent.as_posix()
parent = "" if parent == "." else parent
new_folder = f"{parent}/{clean_name}" if parent else clean_name
source = resolve_in_vault(old_folder)
target = resolve_in_vault(new_folder)
if not source.is_dir() or source.is_symlink():
raise ApiError(404, "RESOURCE_NOT_FOUND", "folder not found", {"path": path})
if target.exists():
raise ApiError(
409, "RESOURCE_CONFLICT", "target folder already exists", {"path": new_folder}
)
affected = [
item
for item in repository.list_note_locations()
if item.folder == old_folder or item.folder.startswith(f"{old_folder}/")
]
source.replace(target)
conn = connect()
now = datetime.now(timezone.utc)
try:
with transaction(conn):
for item in affected:
file_suffix = item.file_path[len(old_folder) :].lstrip("/")
folder_suffix = item.folder[len(old_folder) :].lstrip("/")
repository.update_note_location(
conn=conn,
note_id=item.note_id,
title=item.title,
file_path=f"{new_folder}/{file_suffix}",
folder=(
f"{new_folder}/{folder_suffix}" if folder_suffix else new_folder
),
updated_at=now,
)
except BaseException:
target.replace(source)
raise
finally:
conn.close()
return WorkspaceEntry(
entry_id=_entry_id("folder", new_folder),
name=clean_name,
path=f"/{new_folder}",
type="folder",
children=_tree(target, {item.file_path: item for item in repository.list_note_locations()}),
)
@serialized_vault_mutation
async def delete_folder(path: str) -> OperationResponse:
folder = normalize_folder(path)
if not folder:
raise ApiError(400, "INVALID_PATH", "the Vault root cannot be deleted")
source = resolve_in_vault(folder)
if not source.is_dir() or source.is_symlink():
raise ApiError(404, "RESOURCE_NOT_FOUND", "folder not found", {"path": path})
affected = [
item
for item in repository.list_note_locations()
if item.folder == folder or item.folder.startswith(f"{folder}/")
]
tombstone = source.with_name(f".{source.name}.{uuid4().hex}.deleting")
source.replace(tombstone)
conn = connect()
try:
with transaction(conn):
block_ids: list[str] = []
for item in affected:
block_ids.extend(repository.delete_note(item.note_id, conn=conn))
await vector_store.delete(block_ids, conn=conn)
except BaseException:
tombstone.replace(source)
raise
finally:
conn.close()
try:
shutil.rmtree(tombstone)
except OSError:
# 已提交的删除不回滚;隐藏 tombstone 可由后续维护任务清理。
pass
return OperationResponse(
status="completed",
resource_id=_entry_id("folder", folder),
message=f"deleted folder and {len(affected)} indexed notes",
)
+48
View File
@@ -0,0 +1,48 @@
{
"dataset_id": "rag-core-v1",
"kind": "rag",
"version": "1.0.0",
"description": "基础中文笔记检索集(对应 backend/data/vault 内置语料,重建索引后即可复现)",
"cases": [
{
"case_id": "rag-vector-sim",
"query": "向量数据库如何进行相似度检索",
"expected_note_ids": ["note_c1454740a0e55ef5"],
"expected_block_ids": ["blk_07c4c6bce0ec4d12", "blk_605fb3593809f224"],
"citation_required": true,
"tags": ["向量数据库", "检索"]
},
{
"case_id": "rag-python-func",
"query": "Python 如何定义函数",
"expected_note_ids": ["note_424c3742c6f0e555"],
"expected_block_ids": ["blk_45d48cae2fed40fe", "blk_0768d9c25c2ecf07"],
"citation_required": true,
"tags": ["python"]
},
{
"case_id": "rag-citation",
"query": "搜索结果如何定位到原文位置",
"expected_note_ids": ["note_0c619caa30b1614c"],
"expected_block_ids": ["blk_3f6fcead71c25fc6", "blk_9af7b12e9ce909fc"],
"citation_required": true,
"tags": ["RAG"]
},
{
"case_id": "rag-hybrid",
"query": "混合检索怎么融合全文和向量",
"expected_note_ids": ["note_c1454740a0e55ef5"],
"expected_block_ids": ["blk_82b45418dba9f720"],
"citation_required": true,
"tags": ["检索"]
},
{
"case_id": "rag-tech-stack",
"query": "这个项目用什么后端和检索技术",
"expected_note_ids": ["note_3327e6cf18f3701f"],
"expected_block_ids": ["blk_feb2a9c42e7d31ad"],
"citation_required": false,
"tags": ["项目"]
}
]
}
@@ -0,0 +1,20 @@
commands:
- command_id: mcp-fixture.notify
title: MCP 通知
description: 通过隔离 MCP Host 返回宿主白名单通知 effect。
icon: bolt
locations:
- command_palette
when:
- editor.has_selection
context:
- selection
secrets:
- api_key
mcp_tool: mcp-fixture.command
parameters:
type: object
properties:
message:
type: string
additionalProperties: false
@@ -0,0 +1,26 @@
id: mcp-fixture
name: MCP Fixture
version: 1.0.0
description: 阶段 C/D 离线联调 Fixture,覆盖 MCP Tool、Command 与错误边界。
permissions:
- notes.read
- secrets.use
contributes:
tools:
- mcp-fixture.echo
- mcp-fixture.fail
- mcp-fixture.sleep
- mcp-fixture.large
- mcp-fixture.environment
- mcp-fixture.exit
commands:
- mcp-fixture.notify
settings_sections:
- mcp-fixture.general
backend:
type: mcp
transport: stdio
command: python
args: [server.py]
startup_timeout_seconds: 5
tool_timeout_seconds: 1
@@ -0,0 +1,248 @@
"""确定性的 MCP stdio 测试 Server;仅使用标准库,不依赖产品代码。"""
from __future__ import annotations
import json
import os
import sys
import threading
import time
from typing import Any
WRITE_LOCK = threading.Lock()
CANCELLED: dict[int, threading.Event] = {}
MODE = sys.argv[1] if len(sys.argv) > 1 else "normal"
def send(message: dict[str, Any]) -> None:
with WRITE_LOCK:
sys.stdout.write(json.dumps(message, ensure_ascii=False, separators=(",", ":")) + "\n")
sys.stdout.flush()
def respond(request_id: int, result: dict[str, Any]) -> None:
send({"jsonrpc": "2.0", "id": request_id, "result": result})
def tool(name: str, description: str, properties: dict[str, Any] | None = None) -> dict[str, Any]:
return {
"name": name,
"description": description,
"inputSchema": {
"type": "object",
"properties": properties or {},
"required": list(properties or {}),
"additionalProperties": False,
},
}
TOOLS = {
"echo": {
**tool(
"echo",
"Return the provided text.",
{
"text": {"type": "string"},
"suffix": {"type": ["string", "null"]},
},
),
"_meta": {"notesagent/permission": "notes.read"},
},
"fail": tool("fail", "Return an MCP business error."),
"sleep": tool("sleep", "Wait until completed or cancelled.", {"seconds": {"type": "number"}}),
"large": tool("large", "Return a result larger than the host limit."),
"environment": tool("environment", "Report whether host secrets leaked into the process."),
"exit": tool("exit", "Terminate the fixture process."),
"command": tool(
"command",
"Execute a NotesAgent Plugin Command envelope.",
{"_notesagent": {"type": "object"}},
),
}
# suffix 是可选字段,用于验证 Host 不会把缺省值擅自补成 null。
TOOLS["echo"]["inputSchema"]["required"] = ["text"]
def call_tool(request_id: int, params: dict[str, Any]) -> None:
name = params.get("name")
arguments = params.get("arguments") or {}
if name == "command":
envelope = arguments.get("_notesagent") or {}
command_arguments = envelope.get("arguments") or {}
context = envelope.get("context") or {}
settings = envelope.get("settings") or {}
secrets = envelope.get("secrets") or {}
if not isinstance(secrets.get("api_key"), str):
respond(
request_id,
{
"content": [{"type": "text", "text": "declared secret missing"}],
"isError": True,
},
)
return
message = command_arguments.get("message") or context.get("selection") or ""
message = f"{settings.get('message_prefix', '')}{message}"
respond(
request_id,
{
"content": [{"type": "text", "text": "command completed"}],
"structuredContent": {
"type": "notification",
"payload": {
"level": "success",
"message": str(message),
},
},
"isError": False,
},
)
return
if name == "echo":
text = str(arguments.get("text", ""))
structured_content = {"echo": text}
if "suffix" in arguments:
structured_content["suffix"] = arguments["suffix"]
respond(
request_id,
{
"content": [{"type": "text", "text": text}],
"structuredContent": structured_content,
"isError": False,
},
)
return
if name == "fail":
respond(
request_id,
{
"content": [{"type": "text", "text": "fixture failure"}],
"isError": True,
},
)
return
if name == "large":
respond(
request_id,
{
"content": [{"type": "text", "text": "x" * 300_000}],
"isError": False,
},
)
return
if name == "environment":
respond(
request_id,
{
"content": [{"type": "text", "text": "environment checked"}],
"structuredContent": {
"has_openai_key": "OPENAI_API_KEY" in os.environ,
"has_app_db_path": "APP_DB_PATH" in os.environ,
},
"isError": False,
},
)
return
if name == "exit":
os._exit(17)
if name == "sleep":
cancelled = CANCELLED.setdefault(request_id, threading.Event())
seconds = max(0.0, min(float(arguments.get("seconds", 0)), 30.0))
if cancelled.wait(seconds):
respond(
request_id,
{
"content": [{"type": "text", "text": "cancelled"}],
"isError": True,
},
)
else:
respond(
request_id,
{
"content": [{"type": "text", "text": "completed"}],
"structuredContent": {"slept": seconds},
"isError": False,
},
)
CANCELLED.pop(request_id, None)
return
send(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {"code": -32602, "message": f"Unknown tool: {name}"},
}
)
def main() -> None:
for line in sys.stdin:
message = json.loads(line)
method = message.get("method")
request_id = message.get("id")
params = message.get("params") or {}
if method == "initialize" and isinstance(request_id, int):
if MODE == "invalid-result":
send({"jsonrpc": "2.0", "id": request_id, "result": None})
continue
if MODE == "oversized-stdout":
# 不带换行,验证 Host 在读取完整内容前执行硬上限。
sys.stdout.write("x" * (2 * 1024 * 1024 + 1))
sys.stdout.flush()
time.sleep(10)
return
respond(
request_id,
{
"protocolVersion": params.get("protocolVersion"),
"capabilities": (
{} if MODE == "no-tools" else {"tools": {"listChanged": False}}
),
"serverInfo": {"name": "notesagent-mcp-fixture", "version": "1.0.0"},
},
)
elif method == "tools/list" and isinstance(request_id, int):
if MODE == "invalid-schema":
respond(
request_id,
{
"tools": [
{
"name": "broken",
"description": "invalid schema",
"inputSchema": {"type": "string"},
}
]
},
)
elif params.get("cursor") == "page-2":
respond(
request_id,
{
"tools": [
TOOLS["large"],
TOOLS["environment"],
TOOLS["exit"],
TOOLS["command"],
]
},
)
else:
respond(
request_id,
{"tools": [TOOLS["echo"], TOOLS["fail"], TOOLS["sleep"]], "nextCursor": "page-2"},
)
elif method == "tools/call" and isinstance(request_id, int):
threading.Thread(target=call_tool, args=(request_id, params), daemon=True).start()
elif method == "notifications/cancelled":
cancelled_id = params.get("requestId")
if isinstance(cancelled_id, int):
CANCELLED.setdefault(cancelled_id, threading.Event()).set()
elif method == "ping" and isinstance(request_id, int):
respond(request_id, {})
if __name__ == "__main__":
main()
@@ -0,0 +1,11 @@
section_id: mcp-fixture.general
schema_version: 1
fields:
- key: message_prefix
label: Message Prefix
type: string
default: ""
- key: api_key
label: Fixture API Key
type: secret
required: true
@@ -0,0 +1,19 @@
commands:
- command_id: text-tools.uppercase-selection
title: 转为大写
description: 将当前选区或传入文本转换为大写并显示通知。
icon: edit
locations:
- command_palette
- context_menu
when:
- editor.has_selection
context:
- selection
handler: uppercase_selection
parameters:
type: object
properties:
text:
type: string
additionalProperties: false
@@ -6,6 +6,10 @@ permissions: []
contributes: contributes:
tools: tools:
- text.uppercase - text.uppercase
commands:
- text-tools.uppercase-selection
settings_sections:
- text-tools.general
backend: backend:
type: internal_rpc type: internal_rpc
transport: none transport: none
@@ -0,0 +1,31 @@
section_id: text-tools.general
schema_version: 1
fields:
- key: result_limit
label: 结果字符数
description: Command 通知中最多保留的字符数。
type: number
required: true
default: 100
minimum: 1
maximum: 1000
- key: label_prefix
label: 标签前缀
type: string
default: ""
- key: output_style
label: 输出样式
type: select
default: notification
options:
- notification
- compact
- key: enabled_hint
label: 显示提示
type: boolean
default: true
- key: api_key
label: API Key
description: Secret 示例字段;普通 Settings API 永不返回明文。
type: secret
required: false
+1
View File
@@ -10,6 +10,7 @@ dependencies = [
"httpx>=0.28,<1.0", "httpx>=0.28,<1.0",
"jsonschema>=4.25,<5.0", "jsonschema>=4.25,<5.0",
"pyyaml>=6.0,<7.0", "pyyaml>=6.0,<7.0",
"referencing>=0.36,<1.0",
"sqlite-vec>=0.1.9", "sqlite-vec>=0.1.9",
"uvicorn[standard]>=0.35,<1.0", "uvicorn[standard]>=0.35,<1.0",
] ]
+200
View File
@@ -1,10 +1,18 @@
import asyncio import asyncio
from datetime import datetime, timezone
import pytest
from app.agent.trace_repository import AgentTraceRepository
from app.agent.permissions import PermissionMode from app.agent.permissions import PermissionMode
from app.agent.tools import ToolExecutionContext from app.agent.tools import ToolExecutionContext
from app.container import build_container from app.container import build_container
from app.database.db import connect
from app.errors import ApiError
from app.routes import agent_events
from app.contracts import ( from app.contracts import (
AgentEventType, AgentEventType,
AgentRun,
AgentRunCreateRequest, AgentRunCreateRequest,
AgentRunStatus, AgentRunStatus,
ToolCall, ToolCall,
@@ -108,8 +116,200 @@ def test_permission_confirmation_resumes_agent() -> None:
created.run_id, request_id, "allow_once" created.run_id, request_id, "allow_once"
) )
completed = await container.agent.wait(created.run_id) completed = await container.agent.wait(created.run_id)
events = [event async for event in container.agent.events(created.run_id)]
assert completed.status == AgentRunStatus.completed assert completed.status == AgentRunStatus.completed
assert completed.tool_results[0].success is True assert completed.tool_results[0].success is True
assert AgentEventType.permission_resolved in {event.event for event in events}
run(scenario())
def test_agent_trace_persists_and_replays_from_sequence() -> None:
async def scenario() -> None:
first = build_container()
created = await first.agent.create_run(
AgentRunCreateRequest(
input="persistent trace",
provider_id="mock",
model="mock-1",
metadata={"suite": "agent-benchmark-v1"},
)
)
completed = await first.agent.wait(created.run_id)
restarted = build_container()
restored = restarted.agent.get_run(created.run_id)
first_page = restarted.agent.get_trace(
created.run_id, after_sequence=-1, limit=2
)
second_page = restarted.agent.get_trace(
created.run_id,
after_sequence=first_page.next_sequence,
limit=100,
)
replay = [
event
async for event in restarted.agent.events(
created.run_id, after_sequence=first_page.next_sequence
)
]
assert completed.status == restored.status == AgentRunStatus.completed
assert first_page.has_more is True
assert [item.sequence for item in first_page.items] == [0, 1]
assert second_page.items[0].sequence == 2
assert replay == second_page.items
assert first_page.summary.model_calls == 1
assert first_page.summary.token_usage == completed.token_usage
assert first_page.config_snapshot["metadata"] == {
"suite": "agent-benchmark-v1"
}
assert second_page.items[-1].event == AgentEventType.run_completed
run(scenario())
def test_interrupted_persisted_run_is_closed_after_restart() -> None:
now = datetime.now(timezone.utc)
request = AgentRunCreateRequest(
input="interrupted",
provider_id="mock",
model="mock-1",
)
persisted = AgentRun(
run_id="run_interrupted",
status=AgentRunStatus.running,
input=request.input,
provider_id=request.provider_id,
model=request.model,
max_steps=request.max_steps,
created_at=now,
updated_at=now,
)
AgentTraceRepository().create_run(persisted, request, {"model": "mock-1"})
restarted = build_container()
recovered = restarted.agent.get_run(persisted.run_id)
events = run(
_collect_events(restarted.agent.events(persisted.run_id, after_sequence=-1))
)
assert recovered.status == AgentRunStatus.failed
assert recovered.error_code == "AGENT_PROCESS_RESTARTED"
assert events[-1].event == AgentEventType.run_failed
assert events[-1].sequence == 0
def test_trace_redacts_secrets_and_truncates_large_values() -> None:
async def scenario() -> None:
container = build_container()
secret = "sk-should-not-be-stored"
created = await container.agent.create_run(
AgentRunCreateRequest(
input=f'/tool system.echo {{"text":"{"x" * 4200}","api_key":"{secret}"}}',
provider_id="mock",
model="mock-1",
allowed_tools=["system.echo"],
metadata={"authorization": secret},
)
)
await container.agent.wait(created.run_id)
trace = container.agent.get_trace(
created.run_id, after_sequence=-1, limit=100
)
tool_call = next(
item for item in trace.items if item.event == AgentEventType.tool_call
)
assert tool_call.data["arguments"]["api_key"] == "[REDACTED]"
assert str(tool_call.data["arguments"]["text"]).endswith("...[TRUNCATED]")
assert trace.config_snapshot["metadata"]["authorization"] == "[REDACTED]"
assert secret not in trace.model_dump_json()
conn = connect()
try:
stored_row = conn.execute(
"""
SELECT run_json, request_json, config_snapshot_json
FROM agent_runs WHERE run_id = ?
""",
(created.run_id,),
).fetchone()
stored = "\n".join(str(value) for value in stored_row)
finally:
conn.close()
assert secret not in stored
run(scenario())
def test_persisted_agent_run_preserves_long_input_and_output() -> None:
"""审计事件可以限长,但重启后读取的 AgentRun 不能丢失正文。"""
now = datetime.now(timezone.utc)
long_input = "输入" * 2_500
long_output = "输出" * 2_500
request = AgentRunCreateRequest(
input=long_input,
provider_id="mock",
model="mock-1",
)
persisted = AgentRun(
run_id="run_long_content",
status=AgentRunStatus.completed,
input=long_input,
output=long_output,
provider_id=request.provider_id,
model=request.model,
max_steps=request.max_steps,
created_at=now,
updated_at=now,
)
repository = AgentTraceRepository()
repository.create_run(persisted, request, {"model": request.model})
restored = repository.get_run(persisted.run_id)
assert restored is not None
assert restored.input == long_input
assert restored.output == long_output
async def _collect_events(iterator):
return [event async for event in iterator]
def test_agent_sse_uses_last_event_id_and_emits_event_ids(monkeypatch) -> None:
async def scenario() -> None:
test_container = build_container()
monkeypatch.setattr("app.routes.container", test_container)
created = await test_container.agent.create_run(
AgentRunCreateRequest(
input="resume sse",
provider_id="mock",
model="mock-1",
)
)
await test_container.agent.wait(created.run_id)
response = await agent_events(
created.run_id, after_sequence=None, last_event_id="1"
)
chunks = [chunk async for chunk in response.body_iterator]
body = "".join(
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in chunks
)
assert "id: 0\n" not in body
assert "id: 1\n" not in body
assert "id: 2\n" in body
assert "event: RunCompleted" in body
with pytest.raises(ApiError) as error:
await agent_events(
created.run_id, after_sequence=None, last_event_id="invalid"
)
assert error.value.code == "TRACE_CURSOR_INVALID"
run(scenario()) run(scenario())
+253 -23
View File
@@ -1,26 +1,12 @@
import asyncio import asyncio
import threading
from types import SimpleNamespace
import pytest
from app.main import health, service_status
from app.routes import (
get_index_status,
list_notes,
list_plugins,
list_provider_presets,
list_providers,
list_skills,
)
from app.routes import (
create_provider,
create_task,
delete_provider,
delete_task,
get_provider,
get_task,
list_tasks,
update_provider,
update_task,
)
from app.contracts import ( from app.contracts import (
McpServerSecretStatus,
McpServerSecretWriteRequest,
ProviderCreateRequest, ProviderCreateRequest,
ProviderType, ProviderType,
ProviderUpdateRequest, ProviderUpdateRequest,
@@ -28,6 +14,61 @@ from app.contracts import (
TaskStatus, TaskStatus,
TaskUpdateRequest, TaskUpdateRequest,
) )
from app.main import health, service_status
from app.routes import (
create_provider,
create_task,
delete_provider,
delete_task,
get_index_status,
get_provider,
get_task,
list_notes,
list_plugins,
list_provider_presets,
list_providers,
list_skills,
list_tasks,
update_provider,
update_task,
)
def test_mcp_secret_routes_offload_blocking_lifecycle_work(monkeypatch) -> None:
from app import routes
caller_thread = threading.get_ident()
worker_threads: list[int] = []
class FakeMcpRegistry:
def put_secret(self, server_id, key, secret, *, kind):
worker_threads.append(threading.get_ident())
return McpServerSecretStatus(key=key, configured=True)
def delete_secret(self, server_id, key, *, kind):
worker_threads.append(threading.get_ident())
return McpServerSecretStatus(key=key, configured=False)
monkeypatch.setattr(
routes,
"container",
SimpleNamespace(mcp_servers=FakeMcpRegistry()),
)
written = asyncio.run(
routes.put_mcp_server_secret(
"server-1",
"TOKEN",
McpServerSecretWriteRequest(secret="hidden"),
kind="environment",
)
)
deleted = asyncio.run(
routes.delete_mcp_server_secret("server-1", "TOKEN", kind="environment")
)
assert written.configured is True
assert deleted.configured is False
assert worker_threads and all(item != caller_thread for item in worker_threads)
def test_health() -> None: def test_health() -> None:
@@ -36,6 +77,172 @@ def test_health() -> None:
assert response.model_dump() == {"status": "ok"} assert response.model_dump() == {"status": "ok"}
def test_mcp_create_and_trust_are_not_executed_on_event_loop(monkeypatch) -> None:
from app import routes
from app.contracts import McpServerCreateRequest, McpServerTrustRequest
caller = threading.get_ident()
workers = []
class Registry:
def create(self, request):
workers.append(threading.get_ident())
return "created"
def trust(self, server_id, digest):
workers.append(threading.get_ident())
return "trusted"
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=Registry()))
assert (
asyncio.run(
routes.create_mcp_server(McpServerCreateRequest(name="test", command="uvx"))
)
== "created"
)
assert (
asyncio.run(
routes.trust_mcp_server(
"test", McpServerTrustRequest(command_digest="a" * 64)
)
)
== "trusted"
)
assert len(workers) == 2
assert all(worker != caller for worker in workers)
def test_mcp_split_config_and_secret_requests_persist_without_plaintext(
monkeypatch,
) -> None:
from fastapi.testclient import TestClient
from app import routes
from app.agent.tools import ToolRegistry
from app.config import get_settings
from app.extensions.mcp_registry import McpServerRegistry
from app.main import app
from app.providers.credentials import EncryptedCredentialStore
service = McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
get_settings().data_dir,
allow_process_launch=True,
)
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service))
client = TestClient(app)
config = {
"name": "MiniMax configuration test",
"command": "uvx",
"environment": {"MINIMAX_API_HOST": "https://api.minimaxi.com"},
"secret_environment_keys": ["MINIMAX_API_KEY"],
"startup_timeout_seconds": 120,
"tool_timeout_seconds": 300,
}
# Reproduce the old frontend payload. The backend still enforces separation.
invalid = client.post(
"/api/mcp/servers",
json={
**config,
"environment": {
**config["environment"],
"MINIMAX_API_KEY": "synthetic-only",
},
},
)
assert invalid.status_code == 422
assert invalid.json()["error"]["code"] == "MCP_ENVIRONMENT_INVALID"
created = client.post("/api/mcp/servers", json=config)
assert created.status_code == 201
server_id = created.json()["server_id"]
saved = client.put(
f"/api/mcp/servers/{server_id}/secrets/MINIMAX_API_KEY",
json={"secret": "synthetic-only"},
)
assert saved.status_code == 200
current = client.get(f"/api/mcp/servers/{server_id}")
assert current.json()["secret_environment"] == {"MINIMAX_API_KEY": True}
assert "synthetic-only" not in current.text
assert "synthetic-only" not in service._path.read_text(encoding="utf-8")
_, credentials_path = service.credentials._paths()
assert "synthetic-only" not in credentials_path.read_text(encoding="utf-8")
assert not current.json()["enabled"] # Saving never starts a third-party process.
client.close()
@pytest.mark.parametrize("operation", ["create", "trust"])
def test_mcp_lifecycle_lock_contention_keeps_event_loop_responsive(
monkeypatch,
operation,
) -> None:
from app import routes
from app.agent.tools import ToolRegistry
from app.config import get_settings
from app.contracts import McpServerCreateRequest, McpServerTrustRequest
from app.extensions.mcp_registry import McpServerRegistry
from app.providers.credentials import EncryptedCredentialStore
service = McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
get_settings().data_dir,
allow_process_launch=True,
)
request = McpServerCreateRequest(
name="Lock contention fixture", command="not-executed"
)
server = service.create(request)
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service))
entered = threading.Event()
locked = threading.Event()
release = threading.Event()
original = getattr(service, operation)
def observed(*args):
entered.set()
return original(*args)
def hold_lifecycle_lock():
with service._lifecycle_lock:
locked.set()
release.wait(timeout=5)
monkeypatch.setattr(service, operation, observed)
holder = threading.Thread(target=hold_lifecycle_lock, daemon=True)
holder.start()
# An independent watchdog lets the test fail rather than hang if a regression
# blocks the event loop itself (an asyncio timeout alone cannot catch that).
watchdog = threading.Timer(5, release.set)
watchdog.start()
async def exercise():
pending = asyncio.create_task(
routes.create_mcp_server(request)
if operation == "create"
else routes.trust_mcp_server(
server.server_id,
McpServerTrustRequest(command_digest=server.command_digest),
)
)
try:
assert await asyncio.to_thread(entered.wait, 2)
assert not pending.done()
assert not release.is_set()
assert (await health()).status == "ok"
finally:
release.set()
await pending
try:
assert locked.wait(timeout=2)
asyncio.run(exercise())
finally:
release.set()
watchdog.cancel()
holder.join(timeout=2)
def test_service_status() -> None: def test_service_status() -> None:
response = asyncio.run(service_status()) response = asyncio.run(service_status())
@@ -52,7 +259,9 @@ def test_core_collections_are_typed() -> None:
assert notes.items == [] assert notes.items == []
assert notes.page.limit == 20 assert notes.page.limit == 20
assert [skill.manifest.skill_id for skill in skills.items] == ["knowledge-assistant"] assert [skill.manifest.skill_id for skill in skills.items] == [
"knowledge-assistant"
]
assert skills.items[0].status == "ready" assert skills.items[0].status == "ready"
assert [plugin.manifest.plugin_id for plugin in plugins.items] == ["text-tools"] assert [plugin.manifest.plugin_id for plugin in plugins.items] == ["text-tools"]
assert plugins.items[0].status == "ready" assert plugins.items[0].status == "ready"
@@ -73,9 +282,15 @@ def test_provider_presets_include_openai_and_deepseek() -> None:
def test_provider_presets_static_route_precedes_provider_id_route() -> None: def test_provider_presets_static_route_precedes_provider_id_route() -> None:
from app.routes import router from app.routes import router
get_paths = [route.path for route in router.routes if "GET" in getattr(route, "methods", set())] get_paths = [
route.path
for route in router.routes
if "GET" in getattr(route, "methods", set())
]
assert get_paths.index("/api/providers/presets") < get_paths.index("/api/providers/{provider_id}") assert get_paths.index("/api/providers/presets") < get_paths.index(
"/api/providers/{provider_id}"
)
def test_openapi_contains_documented_frontend_interfaces() -> None: def test_openapi_contains_documented_frontend_interfaces() -> None:
@@ -89,11 +304,26 @@ def test_openapi_contains_documented_frontend_interfaces() -> None:
"/api/agent/runs", "/api/agent/runs",
"/api/agent/runs/{run_id}/cancel", "/api/agent/runs/{run_id}/cancel",
"/api/agent/runs/{run_id}/events", "/api/agent/runs/{run_id}/events",
"/api/agent/runs/{run_id}/trace",
"/api/skills", "/api/skills",
"/api/plugins", "/api/plugins",
"/api/plugins/install", "/api/plugins/install",
"/api/plugins/{plugin_id}/host",
"/api/plugins/{plugin_id}/host/restart",
"/api/plugin-contributions/commands",
"/api/plugin-contributions/commands/{command_id}/execute",
"/api/plugins/{plugin_id}/settings",
"/api/plugins/{plugin_id}/settings/{key}/secret",
"/api/plugins/{plugin_id}/enable", "/api/plugins/{plugin_id}/enable",
"/api/plugins/{plugin_id}/disable", "/api/plugins/{plugin_id}/disable",
"/api/mcp/servers",
"/api/mcp/servers/{server_id}",
"/api/mcp/servers/{server_id}/tools",
"/api/mcp/servers/{server_id}/trust",
"/api/mcp/servers/{server_id}/test",
"/api/mcp/servers/{server_id}/enable",
"/api/mcp/servers/{server_id}/disable",
"/api/mcp/servers/{server_id}/secrets/{key}",
"/api/providers/test", "/api/providers/test",
"/api/providers/presets", "/api/providers/presets",
"/api/credentials/{credential_id}", "/api/credentials/{credential_id}",
+566
View File
@@ -0,0 +1,566 @@
"""Benchmark 服务的单元与端到端测试。
沿用 conftest 的隔离机制APP_DATA_DIR / DB / Vault 都指向临时目录benchmark
数据集也落在临时目录settings.benchmark_datasets_path不读写真实数据
运行采用创建即 queued + 后台 Task 执行的异步模型测试通过 _run 在同一事件循环内
创建并等待后台任务结束得到终态 BenchmarkRun 后再断言
"""
from __future__ import annotations
import asyncio
import json
import pytest
from pydantic import ValidationError
from app.benchmarks import datasets, metrics as m, service
from app.config import get_settings
from app.contracts import (
BenchmarkKind,
BenchmarkRun,
BenchmarkStatus,
RAGRunRequest,
SearchMode,
)
from app.errors import ApiError
def _write_dataset(dataset_id: str, cases: list[dict], *, kind: str = "rag") -> None:
directory = get_settings().benchmark_datasets_path
directory.mkdir(parents=True, exist_ok=True)
payload = {
"dataset_id": dataset_id,
"kind": kind,
"version": "1.0.0",
"description": "test dataset",
"cases": cases,
}
(directory / f"{dataset_id}.json").write_text(
json.dumps(payload, ensure_ascii=False), encoding="utf-8"
)
def _write_raw(dataset_id: str, raw: dict) -> None:
directory = get_settings().benchmark_datasets_path
directory.mkdir(parents=True, exist_ok=True)
(directory / f"{dataset_id}.json").write_text(
json.dumps(raw, ensure_ascii=False), encoding="utf-8"
)
def _run(request: RAGRunRequest):
"""创建运行并在同一事件循环内等待后台任务结束,返回终态 BenchmarkRun。"""
from app.contracts import BenchmarkRun
async def _execute() -> BenchmarkRun:
run = await service.create_rag_run(request)
return await service.wait_for_run(run.run_id)
return asyncio.run(_execute())
# --------------------------------------------------------------------------- #
# 指标纯函数
# --------------------------------------------------------------------------- #
def test_hit_at_k_and_recall() -> None:
retrieved = ["a", "b", "c"]
expected = {"b", "z"}
assert m.hit_at_k(retrieved, expected, 1) is False
assert m.hit_at_k(retrieved, expected, 2) is True
assert m.recall_at_k(retrieved, expected, 5) == 0.5 # 只召回 b
def test_recall_at_k_dedups_duplicate_notes() -> None:
# 同一 Note 经多个 Block 重复出现,去重后 Recall 不应超过 1
assert m.recall_at_k(["note-a", "note-a"], {"note-a"}, 2) == 1.0
assert m.recall_at_k(["note-a", "note-a", "note-b"], {"note-a"}, 3) == 1.0
def test_reciprocal_rank_and_citation_hit() -> None:
assert m.reciprocal_rank(["x", "a", "b"], {"b"}) == 1 / 3
assert m.reciprocal_rank(["x"], {"b"}) == 0.0
assert m.citation_hit(["blk_1"], {"blk_1"}) is True
assert m.citation_hit(["blk_2"], {"blk_1"}) is False
assert m.citation_hit([], {"blk_1"}) is False
def test_percentile() -> None:
assert m.percentile([1.0, 2.0, 3.0, 4.0], 50.0) == 2.5
assert m.percentile([], 50.0) == 0.0
assert m.percentile([7.0], 95.0) == 7.0
# --------------------------------------------------------------------------- #
# Dataset 注册与校验
# --------------------------------------------------------------------------- #
def test_list_datasets_empty_by_default() -> None:
assert datasets.list_datasets(BenchmarkKind.rag) == []
def test_load_missing_dataset_raises() -> None:
with pytest.raises(ApiError) as exc:
datasets.load_dataset("does-not-exist", BenchmarkKind.rag)
assert exc.value.status_code == 404
assert exc.value.code == "BENCHMARK_DATASET_NOT_FOUND"
def test_dataset_without_expected_ids_is_invalid() -> None:
_write_dataset("bad-v1", [{"case_id": "x", "query": "q", "citation_required": False}])
with pytest.raises(ApiError) as exc:
datasets.load_dataset("bad-v1", BenchmarkKind.rag)
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
def test_dataset_kind_mismatch_is_invalid() -> None:
_write_dataset("agent-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}], kind="agent")
with pytest.raises(ApiError) as exc:
datasets.load_dataset("agent-v1", BenchmarkKind.rag)
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
def test_citation_required_requires_expected_block_ids() -> None:
# citation_required=true 却没有 expected_block_ids,无法计算 Citation Hit Rate,应拒绝
_write_dataset(
"cit-req-v1",
[{"case_id": "x", "query": "q", "expected_note_ids": ["n"], "citation_required": True}],
)
with pytest.raises(ApiError) as exc:
datasets.load_dataset("cit-req-v1", BenchmarkKind.rag)
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
def test_list_datasets_skips_corrupted_structure() -> None:
# 合法 JSON 但字段结构错误(cases: 42),列表接口应隔离该文件而非整体 500
_write_raw("bad-structure", {"dataset_id": "bad-structure", "kind": "rag", "cases": 42})
_write_dataset("good-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}])
infos = datasets.list_datasets(BenchmarkKind.rag)
ids = {info.dataset_id for info in infos}
assert "good-v1" in ids
assert "bad-structure" not in ids
# --------------------------------------------------------------------------- #
# 请求校验(空 / 重复 modes)
# --------------------------------------------------------------------------- #
def test_empty_modes_rejected() -> None:
with pytest.raises(ValidationError):
RAGRunRequest(dataset_id="x", modes=[])
def test_duplicate_modes_rejected() -> None:
with pytest.raises(ValidationError):
RAGRunRequest(dataset_id="x", modes=[SearchMode.fts, SearchMode.fts])
# --------------------------------------------------------------------------- #
# RAG Benchmark 端到端
# --------------------------------------------------------------------------- #
def _single_note_case() -> tuple[str, str, dict]:
from app.services import note_service
note = asyncio.run(
note_service.create_note(
title="向量库",
markdown="向量数据库用于存储高维向量并支持近似最近邻检索。",
folder="",
tags=["向量"],
)
)
case = {
"case_id": "c1",
"query": "向量数据库相似度检索",
"expected_note_ids": [note.note_id],
"expected_block_ids": [note.blocks[0].block_id],
"citation_required": True,
"tags": ["向量"],
}
return note.note_id, note.blocks[0].block_id, case
def test_rag_benchmark_end_to_end() -> None:
_, _, case = _single_note_case()
_write_dataset("e2e-v1", [case])
run = _run(RAGRunRequest(dataset_id="e2e-v1", modes=[SearchMode.fts]))
assert run.status.value == "completed"
assert run.dataset_hash.startswith("sha256:")
assert run.metrics is not None
fts = run.metrics["fts"]
assert fts["hit_at_1"] == 1.0
assert fts["recall_at_k"] == 1.0
assert fts["mrr"] == 1.0
assert fts["citation_hit_rate"] == 1.0
assert fts["p50_latency_ms"] >= 0.0
assert fts["p95_latency_ms"] >= fts["p50_latency_ms"]
def test_rag_benchmark_all_modes_produce_metrics() -> None:
_, _, case = _single_note_case()
_write_dataset("e2e-modes-v1", [case])
run = _run(RAGRunRequest(dataset_id="e2e-modes-v1"))
assert run.status.value == "completed"
for mode in ("fts", "vector", "hybrid"):
assert mode in run.metrics
for key in ("hit_at_1", "hit_at_5", "recall_at_k", "mrr", "citation_hit_rate"):
assert 0.0 <= run.metrics[mode][key] <= 1.0
def test_config_snapshot_records_index_and_models() -> None:
_, _, case = _single_note_case()
_write_dataset("snapshot-v1", [case])
run = _run(RAGRunRequest(dataset_id="snapshot-v1", modes=[SearchMode.fts]))
snapshot = run.config_snapshot
assert snapshot["index_meta"] is not None
assert snapshot["embedding"]["version"]
assert snapshot["embedding"]["dim"]
assert snapshot["reranker"]["version"]
assert snapshot["retrieval"]["rrf_k"] == 60
def test_benchmark_report_and_events() -> None:
_, _, case = _single_note_case()
_write_dataset("report-v1", [case])
run = _run(RAGRunRequest(dataset_id="report-v1", modes=[SearchMode.fts]))
report = service.get_report(run.run_id)
events = service.get_events(run.run_id)
assert report is not None
assert report.run_id == run.run_id
assert len(report.cases) == 1
assert report.cases[0].case_id == "c1"
assert report.cases[0].hit_at_1 is True
assert events, "运行应产生事件"
assert events[0].event.value == "RunStarted"
assert events[-1].event.value == "RunCompleted"
def test_cancel_completed_run_keeps_status() -> None:
_, _, case = _single_note_case()
_write_dataset("cancel-v1", [case])
run = _run(RAGRunRequest(dataset_id="cancel-v1", modes=[SearchMode.fts]))
assert run.status.value == "completed"
cancelled = service.cancel_run(run.run_id)
assert cancelled.status.value == "completed" # 已结束,不再变 cancelled
def test_cancel_queued_run_marks_cancelled() -> None:
_, _, case = _single_note_case()
_write_dataset("cancel-queued-v1", [case])
async def _scenario():
run = await service.create_rag_run(
RAGRunRequest(dataset_id="cancel-queued-v1", modes=[SearchMode.fts])
)
service.cancel_run(run.run_id)
return await service.wait_for_run(run.run_id)
run = asyncio.run(_scenario())
assert run.status.value == "cancelled"
# --------------------------------------------------------------------------- #
# 指标聚合:Citation Hit Rate 只统计 citation_required 样本
# --------------------------------------------------------------------------- #
def test_citation_hit_rate_only_counts_citation_required() -> None:
from app.benchmarks import rag as rag_module
from app.contracts import RAGCaseResult
cases = [
RAGCaseResult(
case_id="a", mode=SearchMode.fts, repeat=0, latency_ms=1.0,
citation_hit=True, citation_applicable=True,
),
RAGCaseResult(
case_id="b", mode=SearchMode.fts, repeat=0, latency_ms=1.0,
citation_hit=False, citation_applicable=False,
),
]
metrics = rag_module._aggregate(cases, SearchMode.fts)
# 只有 citation_applicablecitation_required=true)的样本计入分母
assert metrics.citation_hit_rate == 1.0
# --------------------------------------------------------------------------- #
# 路由接入
# --------------------------------------------------------------------------- #
def test_benchmark_routes_wired() -> None:
from app import routes
_, _, case = _single_note_case()
_write_dataset("route-v1", [case])
async def _scenario():
listed = await routes.list_benchmark_datasets(BenchmarkKind.rag)
assert any(item.dataset_id == "route-v1" for item in listed.items)
run = await routes.create_rag_benchmark(
RAGRunRequest(dataset_id="route-v1", modes=[SearchMode.fts])
)
assert run.status.value == "queued"
return await service.wait_for_run(run.run_id)
run = asyncio.run(_scenario())
assert run.status.value == "completed"
got = asyncio.run(routes.get_benchmark_run(run.run_id))
assert got.run_id == run.run_id
report = asyncio.run(routes.get_benchmark_report(run.run_id))
assert report.cases[0].case_id == "c1"
def test_benchmark_run_not_found_raises() -> None:
from app import routes
with pytest.raises(ApiError) as exc:
asyncio.run(routes.get_benchmark_run("benchmark_missing"))
assert exc.value.code == "BENCHMARK_RUN_NOT_FOUND"
# --------------------------------------------------------------------------- #
# 审阅回归:索引兼容 / 容量 / 失败样本 / 取消事件 / 数据集隔离
# --------------------------------------------------------------------------- #
def test_create_rag_run_requires_built_index() -> None:
# 空索引(无已索引 block)会让所有模式得到全 0 指标,应在创建时拒绝而非跑出误导结果
_write_dataset("empty-index-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}])
with pytest.raises(ApiError) as exc:
asyncio.run(
service.create_rag_run(
RAGRunRequest(dataset_id="empty-index-v1", modes=[SearchMode.fts])
)
)
assert exc.value.status_code == 409
assert exc.value.code == "BENCHMARK_INDEX_INCOMPATIBLE"
def test_capacity_exceeded_when_all_runs_active(monkeypatch) -> None:
# 满容量且全为活动(非终态)run 时,无法淘汰,应拒绝创建而非删掉正在运行的 run
_, _, case = _single_note_case()
_write_dataset("capacity-v1", [case])
monkeypatch.setattr(service, "MAX_RUNS", 1)
fake_id = "benchmark_fake_active"
service._runs[fake_id] = BenchmarkRun(
run_id=fake_id,
kind=BenchmarkKind.rag,
dataset_id="capacity-v1",
dataset_hash="sha256:fake",
status=BenchmarkStatus.queued,
created_at=service._now(),
)
try:
with pytest.raises(ApiError) as exc:
asyncio.run(
service.create_rag_run(
RAGRunRequest(dataset_id="capacity-v1", modes=[SearchMode.fts])
)
)
assert exc.value.status_code == 429
assert exc.value.code == "BENCHMARK_CAPACITY_EXCEEDED"
finally:
service._runs.pop(fake_id, None)
def test_failed_samples_counted_as_zero_in_aggregate() -> None:
from app.benchmarks import rag as rag_module
from app.contracts import RAGCaseResult
cases = [
RAGCaseResult(
case_id="ok", mode=SearchMode.fts, repeat=0, latency_ms=10.0,
hit_at_1=True, recall=1.0, reciprocal_rank=1.0,
citation_hit=True, citation_applicable=True,
),
RAGCaseResult(
case_id="boom", mode=SearchMode.fts, repeat=0, latency_ms=0.0,
error="RAG case evaluation failed.",
error_code="BENCHMARK_CASE_EVALUATION_FAILED",
),
]
metrics = rag_module._aggregate(cases, SearchMode.fts)
assert metrics.total_cases == 2
assert metrics.successful_cases == 1
assert metrics.failed_cases == 1
assert metrics.failure_rate == 0.5
# 失败样本按零分计入质量指标分母,汇总不虚高
assert metrics.hit_at_1 == 0.5
assert metrics.recall_at_k == 0.5
# 延迟只统计成功样本
assert metrics.p50_latency_ms == 10.0
def test_cancel_emits_run_cancelled_event() -> None:
_, _, case = _single_note_case()
_write_dataset("cancel-event-v1", [case])
async def _scenario():
run = await service.create_rag_run(
RAGRunRequest(dataset_id="cancel-event-v1", modes=[SearchMode.fts])
)
service.cancel_run(run.run_id)
return await service.wait_for_run(run.run_id)
run = asyncio.run(_scenario())
assert run.status.value == "cancelled"
events = service.get_events(run.run_id)
assert events[-1].event.value == "RunCancelled"
def test_load_dataset_ignores_corrupted_unrelated_files() -> None:
# 无关文件损坏(非法 JSON / 顶层非对象)不应阻断目标数据集加载
directory = get_settings().benchmark_datasets_path
directory.mkdir(parents=True, exist_ok=True)
(directory / "broken.json").write_text("{ not valid json", encoding="utf-8")
(directory / "array.json").write_text('["a", "b"]', encoding="utf-8")
_write_dataset("ok-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}])
dataset = datasets.load_dataset("ok-v1", BenchmarkKind.rag)
assert dataset.dataset_id == "ok-v1"
assert len(dataset.cases) == 1
def test_load_dataset_top_level_must_be_object() -> None:
_write_raw("array-top", ["a", "b"])
with pytest.raises(ApiError) as exc:
datasets.load_dataset("array-top", BenchmarkKind.rag)
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
# --------------------------------------------------------------------------- #
# 审阅回归:运行中取消 / 仅块标注 / SSE 终止事件
# --------------------------------------------------------------------------- #
def test_cancel_running_benchmark_stops_early() -> None:
"""运行中取消应在样本边界及时生效,而非跑完全部样本(审阅 P1)。"""
from app.benchmarks import service
from app.services import note_service
note = asyncio.run(
note_service.create_note(
title="取消回归", markdown="向量数据库用于存储高维向量。", folder="", tags=["向量"]
)
)
cases = [
{
"case_id": f"c{i}",
"query": "向量数据库",
"expected_note_ids": [note.note_id],
"expected_block_ids": [note.blocks[0].block_id],
"citation_required": True,
}
for i in range(50)
]
_write_dataset("cancel-running-v1", cases)
async def _scenario():
run = await service.create_rag_run(
RAGRunRequest(dataset_id="cancel-running-v1", modes=[SearchMode.fts])
)
async def _cancel_after_start():
# 取消通过事件循环调度(独立 Task),而非同步直调,才能复现事件循环饥饿
while service.get_run(run.run_id).status == BenchmarkStatus.queued:
await asyncio.sleep(0)
service.cancel_run(run.run_id)
cancel_task = asyncio.create_task(_cancel_after_start())
finished = await service.wait_for_run(run.run_id)
await cancel_task
return finished
run = asyncio.run(_scenario())
assert run.status.value == "cancelled"
completed = sum(
1 for e in service.get_events(run.run_id) if e.event.value == "CaseCompleted"
)
assert completed < 50 # 未跑完全部样本,证明取消在样本边界生效
def test_block_only_annotation_resolves_note_and_scores() -> None:
"""仅标注 expected_block_ids 的样本应按块反查笔记评分,而非零分(审阅 P2)。"""
from app.services import note_service
note = asyncio.run(
note_service.create_note(
title="仅块标注", markdown="向量数据库存储高维向量。", folder="", tags=["向量"]
)
)
_write_dataset("block-only-v1", [{
"case_id": "c1",
"query": "向量数据库",
"expected_block_ids": [note.blocks[0].block_id],
"citation_required": False,
}])
run = _run(RAGRunRequest(dataset_id="block-only-v1", modes=[SearchMode.fts]))
assert run.status.value == "completed"
fts = run.metrics["fts"]
assert fts["hit_at_1"] == 1.0
assert fts["recall_at_k"] == 1.0
assert fts["mrr"] == 1.0
def test_sse_stream_ends_on_terminal_event_in_replay() -> None:
"""历史回放期间遇到终止事件时流应立即结束,而非进入实时队列永久等待(审阅 P2)。"""
from app import routes
from app.benchmarks import service
from app.contracts import BenchmarkEvent, BenchmarkEventType
run_id = "benchmark_sse_replay"
now = service._now()
# 模拟「回放期间运行完成」:run 仍为 runningsubscribe 返回非空队列),
# 但历史事件里已含 RunCompleted 终止事件。
service._runs[run_id] = BenchmarkRun(
run_id=run_id,
kind=BenchmarkKind.rag,
dataset_id="d",
dataset_hash="sha256:x",
status=BenchmarkStatus.running,
created_at=now,
)
service._events[run_id] = [
BenchmarkEvent(
event=BenchmarkEventType.run_started, run_id=run_id, sequence=0,
data={}, timestamp=now,
),
BenchmarkEvent(
event=BenchmarkEventType.run_completed, run_id=run_id, sequence=1,
data={}, timestamp=now,
),
]
try:
# 直调路由函数时 FastAPI 不解析 Query/Header 默认值,需显式传 None 覆盖 Header 哨兵
response = asyncio.run(
routes.benchmark_events(run_id, after_sequence=-1, last_event_id=None)
)
async def _collect() -> list[str]:
out: list[str] = []
async for chunk in response.body_iterator:
out.append(chunk)
return out
# 加超时防止回归(旧实现会永久挂起)
chunks = asyncio.run(asyncio.wait_for(_collect(), timeout=5))
finally:
service._forget(run_id)
events = [
line for chunk in chunks for line in chunk.splitlines() if line.startswith("event: ")
]
assert events == ["event: RunStarted", "event: RunCompleted"]
+59 -1
View File
@@ -1,16 +1,21 @@
import asyncio import asyncio
from pathlib import Path
import httpx import httpx
import pytest
from app.config import get_settings from app.config import get_settings
from app.contracts import CredentialWriteRequest from app.contracts import CredentialWriteRequest
from app.errors import ApiError
from app.providers.credentials import ( from app.providers.credentials import (
ChainedCredentialResolver, ChainedCredentialResolver,
CredentialStoreError,
EncryptedCredentialStore, EncryptedCredentialStore,
EnvironmentCredentialResolver, EnvironmentCredentialResolver,
) )
from app.providers.factory import ProviderFactory
from app.providers.openai_compatible import OpenAICompatibleProvider from app.providers.openai_compatible import OpenAICompatibleProvider
from app.routes import get_credential_status, put_credential from app.routes import delete_credential, get_credential_status, put_credential
def test_encrypted_credential_store_round_trip_without_plaintext_on_disk() -> None: def test_encrypted_credential_store_round_trip_without_plaintext_on_disk() -> None:
@@ -31,6 +36,33 @@ def test_encrypted_credential_store_round_trip_without_plaintext_on_disk() -> No
assert store.resolve("deepseek") is None assert store.resolve("deepseek") is None
def test_encrypted_credential_store_deletes_multiple_credentials_atomically() -> None:
store = EncryptedCredentialStore()
store.put("plugin.first", "first")
store.put("plugin.second", "second")
store.put("openai", "keep")
removed = store.delete_many(["plugin.first", "plugin.second"])
assert removed == {"plugin.first", "plugin.second"}
assert store.resolve("plugin.first") is None
assert store.resolve("plugin.second") is None
assert store.resolve("openai") == "keep"
def test_credential_write_os_error_uses_stable_store_error(monkeypatch) -> None:
store = EncryptedCredentialStore()
store.put("existing", "value")
def fail_replace(_path: Path, _target: Path) -> Path:
raise OSError("injected replace failure")
monkeypatch.setattr(Path, "replace", fail_replace)
with pytest.raises(CredentialStoreError, match="cannot be written"):
store.put("new", "value")
def test_credential_api_never_returns_secret() -> None: def test_credential_api_never_returns_secret() -> None:
written = asyncio.run( written = asyncio.run(
put_credential( put_credential(
@@ -72,3 +104,29 @@ def test_saved_credential_takes_precedence_over_environment_fallback(monkeypatch
resolver = ChainedCredentialResolver(store, EnvironmentCredentialResolver()) resolver = ChainedCredentialResolver(store, EnvironmentCredentialResolver())
assert resolver.resolve("deepseek") == "saved-key" assert resolver.resolve("deepseek") == "saved-key"
def test_public_credential_api_rejects_plugin_namespace() -> None:
operations = [
get_credential_status("plugin.text-tools.api_key"),
put_credential(
"plugin.text-tools.api_key",
CredentialWriteRequest(api_key="must-not-write"),
),
delete_credential("plugin.text-tools.api_key"),
]
for operation in operations:
with pytest.raises(ApiError) as exc:
asyncio.run(operation)
assert exc.value.code == "CREDENTIAL_NAMESPACE_RESERVED"
assert EncryptedCredentialStore().resolve("plugin.text-tools.api_key") is None
def test_provider_resolver_cannot_read_plugin_secret() -> None:
store = EncryptedCredentialStore()
store.put("plugin.text-tools.api_key", "private-plugin-secret")
resolver = ProviderFactory(store).credentials
with pytest.raises(CredentialStoreError, match="reserved for Plugin settings"):
resolver.resolve("plugin.text-tools.api_key")
+478 -1
View File
@@ -1,4 +1,7 @@
import asyncio import asyncio
import shutil
import threading
import time
import pytest import pytest
@@ -8,18 +11,39 @@ from app.container import build_container
from app.contracts import ( from app.contracts import (
AgentRunCreateRequest, AgentRunCreateRequest,
AgentRunStatus, AgentRunStatus,
PluginCommandContext,
SkillStatus, SkillStatus,
ToolCall, ToolCall,
) )
from app.extensions import ExtensionError from app.extensions import ExtensionError
from app.extensions.mcp import McpStdioClient
from app.extensions.runtime import (
_arguments_model_from_schema,
_validate_mcp_command_target_schema,
)
from app.services import note_service from app.services import note_service
from app.config import get_settings from app.config import BACKEND_DIR, get_settings
MCP_FIXTURE = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo"
def run(coroutine): def run(coroutine):
return asyncio.run(coroutine) return asyncio.run(coroutine)
@pytest.fixture
def mcp_container():
container = build_container()
installed = container.plugins.install(MCP_FIXTURE)
assert installed.status == "permission_required"
container.plugins.set_permissions("mcp-fixture", ["notes.read", "secrets.use"])
try:
yield container
finally:
container.plugins.shutdown()
def test_bundled_plugin_registers_tool_and_skill_is_ready() -> None: def test_bundled_plugin_registers_tool_and_skill_is_ready() -> None:
async def scenario() -> None: async def scenario() -> None:
container = build_container() container = build_container()
@@ -297,3 +321,456 @@ def test_attachment_and_transcription_tools_use_host_storage() -> None:
assert transcription.output["text"] == "会议转写内容" assert transcription.output["text"] == "会议转写内容"
run(scenario()) run(scenario())
def test_mcp_stdio_host_discovers_namespaced_tools_and_maps_results(
mcp_container, monkeypatch
) -> None:
async def scenario() -> None:
monkeypatch.setenv("OPENAI_API_KEY", "must-not-enter-plugin-host")
enabled = mcp_container.plugins.enable("mcp-fixture")
status = mcp_container.plugins.get_host_status("mcp-fixture")
definition = mcp_container.tools.get("mcp-fixture.echo").definition
result = await mcp_container.tools.execute(
ToolCall(
tool_call_id="call_mcp_echo",
name="mcp-fixture.echo",
arguments={"text": "hello mcp"},
),
ToolExecutionContext(
run_id="run_mcp_fixture", tool_call_id="call_mcp_echo"
),
)
assert enabled.status == "ready" and enabled.enabled is True
assert status.status == "ready"
environment = await mcp_container.tools.execute(
ToolCall(
tool_call_id="call_mcp_environment",
name="mcp-fixture.environment",
arguments={},
),
ToolExecutionContext(run_id="run_mcp_fixture"),
)
assert status.tools_count == 7
assert status.protocol_version == "2025-11-25"
assert status.server_name == "notesagent-mcp-fixture"
assert definition.permission == "notes.read"
assert result.success is True
assert result.output == {"echo": "hello mcp"}
explicit_null = await mcp_container.tools.execute(
ToolCall(
tool_call_id="call_mcp_explicit_null",
name="mcp-fixture.echo",
arguments={"text": "null stays explicit", "suffix": None},
),
ToolExecutionContext(run_id="run_mcp_fixture"),
)
assert explicit_null.success is True
assert explicit_null.output == {
"echo": "null stays explicit",
"suffix": None,
}
assert environment.success is True
assert environment.output == {
"has_openai_key": False,
"has_app_db_path": False,
}
disabled = mcp_container.plugins.disable("mcp-fixture")
assert disabled.status == "disabled"
assert mcp_container.plugins.get_host_status("mcp-fixture").status == "stopped"
assert not mcp_container.tools.contains("mcp-fixture.echo")
with pytest.raises(ExtensionError) as exc:
mcp_container.plugins.restart_host("mcp-fixture")
assert exc.value.code == "PLUGIN_HOST_UNAVAILABLE"
assert mcp_container.plugins.get("mcp-fixture").status == "disabled"
assert not mcp_container.tools.contains("mcp-fixture.echo")
mcp_container.plugins.uninstall("mcp-fixture")
reinstalled = mcp_container.plugins.install(MCP_FIXTURE)
fresh_status = mcp_container.plugins.get_host_status("mcp-fixture")
assert reinstalled.status == "permission_required"
assert fresh_status.status == "stopped"
assert fresh_status.started_at is None
assert fresh_status.protocol_version is None
assert fresh_status.server_name is None
run(scenario())
def test_mcp_command_target_receives_scoped_context_and_declared_secret(
mcp_container,
) -> None:
async def scenario() -> None:
mcp_container.plugins.enable("mcp-fixture")
assert not mcp_container.tools.contains("mcp-fixture.command")
with pytest.raises(ExtensionError) as missing:
await mcp_container.plugins.execute_command(
"mcp-fixture.notify",
{},
PluginCommandContext(selection="来自选区"),
)
assert missing.value.code == "PLUGIN_SECRET_REQUIRED"
mcp_container.plugins.put_setting_secret(
"mcp-fixture", "api_key", "mcp-command-secret"
)
mcp_container.plugins.update_settings(
"mcp-fixture", 1, {"message_prefix": "Fixture: "}
)
result = await mcp_container.plugins.execute_command(
"mcp-fixture.notify",
{},
PluginCommandContext(
note_id="must-not-enter-envelope",
selection="来自选区",
),
)
assert result.effect.type == "notification"
assert result.effect.payload.model_dump() == {
"level": "success",
"message": "Fixture: 来自选区",
}
assert "mcp-command-secret" not in repr(
mcp_container.plugins.commands.audit_events()
)
run(scenario())
def test_mcp_command_target_rejects_incompatible_envelope_schema(tmp_path) -> None:
package = tmp_path / "mcp-bad-command"
shutil.copytree(MCP_FIXTURE, package)
for filename in ("plugin.yaml", "commands.yaml", "settings.yaml"):
path = package / filename
path.write_text(
path.read_text(encoding="utf-8").replace(
"mcp-fixture", "mcp-bad-command"
),
encoding="utf-8",
)
server_path = package / "server.py"
server_path.write_text(
server_path.read_text(encoding="utf-8").replace(
'{"_notesagent": {"type": "object"}}',
'{"unexpected": {"type": "string"}}',
),
encoding="utf-8",
)
container = build_container()
container.plugins.install(package)
container.plugins.set_permissions(
"mcp-bad-command", ["notes.read", "secrets.use"]
)
try:
with pytest.raises(ExtensionError) as exc:
container.plugins.enable("mcp-bad-command")
assert exc.value.code == "PLUGIN_CONTRIBUTION_INVALID"
assert container.plugins.get("mcp-bad-command").status == "error"
finally:
container.plugins.shutdown()
def test_mcp_command_target_enable_check_only_requires_protocol_marker() -> None:
# `not`/`oneOf` 等完整语义由实际调用前的官方 Validator 处理;启用检查
# 只确认不可被引用或组合隐藏的稳定宿主入口,避免维护不完整的求解器。
_validate_mcp_command_target_schema(
{
"type": "object",
"properties": {
"_notesagent": {
"type": "object",
"not": {"type": "object"},
}
},
},
"marker.run",
)
invalid_markers = [
{
"$defs": {"envelope": {"type": "object"}},
"properties": {"_notesagent": {"$ref": "#/$defs/envelope"}},
},
{
"allOf": [
{"properties": {"_notesagent": {"type": "object"}}},
]
},
]
for schema in invalid_markers:
with pytest.raises(ExtensionError) as exc:
_validate_mcp_command_target_schema(schema, "marker.run")
assert exc.value.code == "PLUGIN_CONTRIBUTION_INVALID"
def test_mcp_command_validates_actual_envelope_before_call(tmp_path) -> None:
package = tmp_path / "mcp-runtime-schema"
shutil.copytree(MCP_FIXTURE, package)
for filename in ("plugin.yaml", "commands.yaml", "settings.yaml"):
path = package / filename
path.write_text(
path.read_text(encoding="utf-8").replace(
"mcp-fixture", "mcp-runtime-schema"
),
encoding="utf-8",
)
server_path = package / "server.py"
server_path.write_text(
server_path.read_text(encoding="utf-8").replace(
'{"_notesagent": {"type": "object"}}',
'{"_notesagent": {"type": "object", "properties": '
'{"arguments": {"type": "object", "maxProperties": 0}, '
'"context": {"type": "object", "properties": '
'{"selection": {"type": "string"}}, "required": ["selection"]}}, '
'"required": ["arguments", "context"]}}',
),
encoding="utf-8",
)
container = build_container()
container.plugins.install(package)
container.plugins.set_permissions(
"mcp-runtime-schema", ["notes.read", "secrets.use"]
)
try:
# context.selection 是 Command 的 when/context 契约保证的真实字段;
# 启用期结构检查不得因没有伪造该业务值而拒绝目标 Schema。
container.plugins.enable("mcp-runtime-schema")
container.plugins.put_setting_secret(
"mcp-runtime-schema", "api_key", "configured"
)
with pytest.raises(ExtensionError) as exc:
run(
container.plugins.execute_command(
"mcp-runtime-schema.notify",
{"message": "must be rejected locally"},
PluginCommandContext(selection="visible"),
)
)
assert exc.value.code == "PLUGIN_COMMAND_TARGET_SCHEMA_MISMATCH"
finally:
container.plugins.shutdown()
def test_agent_calls_mcp_tool_through_registry_and_writes_trace(mcp_container) -> None:
async def scenario() -> None:
mcp_container.plugins.enable("mcp-fixture")
created = await mcp_container.agent.create_run(
AgentRunCreateRequest(
input='/tool mcp-fixture.echo {"text":"agent mcp"}',
provider_id="mock",
model="mock-1",
allowed_tools=["mcp-fixture.echo"],
)
)
completed = await mcp_container.agent.wait(created.run_id)
trace = mcp_container.agent.get_trace(
created.run_id, after_sequence=-1, limit=100
)
assert completed.status == AgentRunStatus.completed
assert completed.tool_results[0].success is True
assert completed.tool_results[0].output == {"echo": "agent mcp"}
assert any(
item.event == "ToolCall" and item.data.get("name") == "mcp-fixture.echo"
for item in trace.items
)
run(scenario())
def test_mcp_business_error_size_limit_and_timeout_are_structured(mcp_container) -> None:
async def scenario() -> None:
mcp_container.plugins.enable("mcp-fixture")
context = ToolExecutionContext(run_id="run_mcp_errors")
failed = await mcp_container.tools.execute(
ToolCall(tool_call_id="call_fail", name="mcp-fixture.fail", arguments={}),
context,
)
oversized = await mcp_container.tools.execute(
ToolCall(tool_call_id="call_large", name="mcp-fixture.large", arguments={}),
context,
)
timed_out = await mcp_container.tools.execute(
ToolCall(
tool_call_id="call_sleep",
name="mcp-fixture.sleep",
arguments={"seconds": 5},
),
ToolExecutionContext(
run_id="run_mcp_errors", tool_call_id="call_sleep"
),
)
recovered = await mcp_container.tools.execute(
ToolCall(
tool_call_id="call_after_timeout",
name="mcp-fixture.echo",
arguments={"text": "still ready"},
),
context,
)
assert failed.success is False
assert failed.error_code == "MCP_TOOL_CALL_FAILED"
assert failed.error_message == "fixture failure"
assert oversized.success is False
assert oversized.error_code == "MCP_TOOL_RESULT_TOO_LARGE"
assert timed_out.success is False
assert timed_out.error_code == "MCP_TOOL_CALL_FAILED"
assert recovered.success is True
assert mcp_container.plugins.get_host_status("mcp-fixture").status == "ready"
run(scenario())
def test_mcp_cancel_releases_blocking_response_thread(
mcp_container, monkeypatch
) -> None:
async def scenario() -> None:
mcp_container.plugins.enable("mcp-fixture")
released = threading.Event()
original_wait = McpStdioClient.wait_response
def tracked_wait(self, *args, **kwargs):
try:
return original_wait(self, *args, **kwargs)
finally:
released.set()
monkeypatch.setattr(McpStdioClient, "wait_response", tracked_wait)
task = asyncio.create_task(
mcp_container.plugins.mcp.call_tool(
"mcp-fixture",
"sleep",
{"seconds": 5},
request_id="call_cancel_release",
)
)
await asyncio.sleep(0.05)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
deadline = time.monotonic() + 0.5
while not released.is_set() and time.monotonic() < deadline:
await asyncio.sleep(0.01)
assert released.is_set(), "cancelled MCP wait must not occupy a worker until timeout"
run(scenario())
def test_mcp_argument_model_preserves_json_schema_additional_properties() -> None:
arguments_model = _arguments_model_from_schema(
"mcp-fixture.dynamic",
{
"type": "object",
"properties": {"model_dump": {"type": "string"}},
"required": ["model_dump"],
"additionalProperties": {"type": "string"},
},
)
arguments = arguments_model.model_validate(
{"model_dump": "method name remains data", "dynamic-key": "value"}
)
assert arguments.model_dump() == {
"model_dump": "method name remains data",
"dynamic-key": "value",
}
def test_production_rejects_unsandboxed_mcp_host(monkeypatch) -> None:
monkeypatch.setenv("APP_ENVIRONMENT", "production")
get_settings.cache_clear()
container = build_container()
installed = container.plugins.install(MCP_FIXTURE)
assert installed.status == "permission_required"
container.plugins.set_permissions("mcp-fixture", ["notes.read", "secrets.use"])
try:
with pytest.raises(ExtensionError) as exc:
container.plugins.enable("mcp-fixture")
assert exc.value.code == "MCP_TRUST_APPROVAL_REQUIRED"
assert container.plugins.get_host_status("mcp-fixture").status == "stopped"
assert not container.tools.contains("mcp-fixture.echo")
finally:
container.plugins.shutdown()
get_settings.cache_clear()
def test_mcp_abnormal_exit_unregisters_tools_and_restart_recovers(mcp_container) -> None:
async def scenario() -> None:
mcp_container.plugins.enable("mcp-fixture")
crashed = await mcp_container.tools.execute(
ToolCall(tool_call_id="call_exit", name="mcp-fixture.exit", arguments={}),
ToolExecutionContext(run_id="run_mcp_exit", tool_call_id="call_exit"),
)
deadline = time.monotonic() + 2
while mcp_container.tools.contains("mcp-fixture.echo") and time.monotonic() < deadline:
await asyncio.sleep(0.02)
plugin = mcp_container.plugins.get("mcp-fixture")
status = mcp_container.plugins.get_host_status("mcp-fixture")
assert crashed.success is False
assert crashed.error_code == "PLUGIN_HOST_UNAVAILABLE"
assert plugin.status == "error" and plugin.enabled is False
assert status.status == "unhealthy"
assert not mcp_container.tools.contains("mcp-fixture.echo")
restarted = mcp_container.plugins.restart_host("mcp-fixture")
assert restarted.status == "ready"
assert restarted.tools_count == 7
assert mcp_container.tools.contains("mcp-fixture.echo")
run(scenario())
@pytest.mark.parametrize(
("mode", "contributions", "expected_code"),
[
("no-tools", "[]", "MCP_CAPABILITY_UNSUPPORTED"),
("invalid-schema", "[mcp-invalid.broken]", "MCP_TOOL_SCHEMA_INVALID"),
("invalid-result", "[]", "MCP_INITIALIZE_FAILED"),
("oversized-stdout", "[]", "PLUGIN_HOST_UNAVAILABLE"),
],
)
def test_mcp_rejects_invalid_initialization_and_discovery(
tmp_path, mode, contributions, expected_code
) -> None:
package = tmp_path / f"mcp-{mode}"
package.mkdir()
shutil.copyfile(MCP_FIXTURE / "server.py", package / "server.py")
(package / "plugin.yaml").write_text(
f"""
id: mcp-invalid
name: Invalid MCP Fixture
version: 1.0.0
contributes:
tools: {contributions}
backend:
type: mcp
transport: stdio
command: python
args: [server.py, {mode}]
startup_timeout_seconds: 5
tool_timeout_seconds: 1
""".strip(),
encoding="utf-8",
)
container = build_container()
container.plugins.install(package)
try:
with pytest.raises(ExtensionError) as exc:
container.plugins.enable("mcp-invalid")
assert exc.value.code == expected_code
assert container.plugins.get("mcp-invalid").status == "error"
assert container.plugins.get_host_status("mcp-invalid").status == "error"
assert not container.tools.contains("mcp-invalid.broken")
finally:
container.plugins.shutdown()
+769
View File
@@ -0,0 +1,769 @@
import asyncio
import hashlib
import json
import sys
import threading
import time
from concurrent.futures import ThreadPoolExecutor
import httpx
import pytest
from app.agent.tools import ToolExecutionContext, ToolRegistry
from app.config import BACKEND_DIR, get_settings
from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall
from app.extensions.mcp import McpLegacySseClient
from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry
from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore
SERVER = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" / "server.py"
def request(**overrides) -> McpServerCreateRequest:
values = {
"name": "Echo MCP",
"command": sys.executable,
"args": [str(SERVER)],
"permissions": ["notes.read", "secrets.use"],
"secret_environment_keys": ["TEST_MCP_SECRET"],
}
values.update(overrides)
return McpServerCreateRequest(**values)
def registry(*, launch: bool = True) -> McpServerRegistry:
return McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
get_settings().data_dir,
allow_process_launch=launch,
)
def test_registry_requires_current_trust_and_never_returns_secret() -> None:
service = registry()
created = service.create(request())
assert created.trusted is False
assert created.secret_environment == {"TEST_MCP_SECRET": False}
service.put_secret(created.server_id, "TEST_MCP_SECRET", "do-not-return")
configured = service.get(created.server_id)
assert configured.secret_environment == {"TEST_MCP_SECRET": True}
assert "do-not-return" not in configured.model_dump_json()
with pytest.raises(McpRegistryError, match="approve"):
service.test(created.server_id)
service.trust(created.server_id, created.command_digest)
tested = service.test(created.server_id)
assert tested.status == "stopped"
assert tested.last_test_succeeded is True
assert tested.tools_count > 0
service.shutdown()
def test_secret_change_disables_server_and_requires_a_new_connection_test() -> None:
service = registry()
created = service.create(request())
service.put_secret(created.server_id, "TEST_MCP_SECRET", "first")
service.trust(created.server_id, created.command_digest)
service.test(created.server_id)
service.enable(created.server_id)
service.put_secret(created.server_id, "TEST_MCP_SECRET", "second")
current = service.get(created.server_id)
assert current.enabled is False
assert current.last_test_succeeded is None
assert not any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in service.tools.definitions()
)
with pytest.raises(McpRegistryError) as error:
service.enable(created.server_id)
assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED"
service.shutdown()
def test_update_disables_server_and_revokes_command_trust() -> None:
service = registry()
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
service.test(created.server_id)
enabled = service.enable(created.server_id)
assert enabled.enabled is True
assert any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in service.tools.definitions()
)
updated = service.update(
created.server_id,
McpServerUpdateRequest(
**request(name="Changed", secret_environment_keys=[]).model_dump(),
version=enabled.version,
),
)
assert updated.enabled is False
assert updated.trusted is False
assert not any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in service.tools.definitions()
)
service.shutdown()
def test_update_remains_retryable_when_removed_secret_cleanup_fails(
monkeypatch,
) -> None:
service = registry()
created = service.create(request())
service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry")
def fail_delete_many(_secret_ids: list[str]) -> set[str]:
raise CredentialStoreError("credential store unavailable")
monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many)
with pytest.raises(McpRegistryError) as error:
service.update(
created.server_id,
McpServerUpdateRequest(
**request(secret_environment_keys=[]).model_dump(),
version=created.version,
),
)
current = service.get(created.server_id)
assert error.value.code == "MCP_SECRET_STORE_ERROR"
assert current.version == created.version
assert current.secret_environment == {"TEST_MCP_SECRET": True}
service.shutdown()
def test_delete_keeps_server_retryable_when_secret_cleanup_fails(monkeypatch) -> None:
service = registry()
created = service.create(request())
service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry")
def fail_delete_many(_secret_ids: list[str]) -> set[str]:
raise CredentialStoreError("credential store unavailable")
monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many)
with pytest.raises(McpRegistryError) as error:
service.delete(created.server_id)
current = service.get(created.server_id)
assert error.value.code == "MCP_SECRET_STORE_ERROR"
assert current.server_id == created.server_id
assert current.secret_environment == {"TEST_MCP_SECRET": True}
service.shutdown()
def test_unavailable_server_removes_bridge_host(monkeypatch) -> None:
service = registry()
created = service.create(request(secret_environment_keys=[]))
with service._lock:
service._records[created.server_id] = {
**service._records[created.server_id],
"enabled": True,
}
removed: list[str] = []
monkeypatch.setattr(service.bridge, "remove", removed.append)
generation = object()
service._generations[created.server_id] = generation
service._unavailable(created.server_id, generation, "connection lost")
current = service.get(created.server_id)
assert removed == [f"mcp.{created.server_id}"]
assert current.enabled is False
assert current.status == "unhealthy"
service.shutdown()
def test_old_failure_callback_cannot_stop_replacement_host(monkeypatch) -> None:
service = registry()
callbacks = []
original_start = service.bridge.start
def capture_callback(*args, **kwargs):
callbacks.append(args[4])
return original_start(*args, **kwargs)
monkeypatch.setattr(service.bridge, "start", capture_callback)
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
callback_thread = None
try:
service.test(created.server_id)
service.enable(created.server_id)
old_callback = callbacks[-1]
callback_started = threading.Event()
callback_finished = threading.Event()
def delayed_failure():
callback_started.set()
old_callback(f"mcp.{created.server_id}", "delayed old failure")
callback_finished.set()
# Queue the old callback while a replacement owns the lifecycle lock.
with service._lifecycle_lock:
callback_thread = threading.Thread(target=delayed_failure, daemon=True)
callback_thread.start()
assert callback_started.wait(timeout=2)
service.disable(created.server_id)
service.enable(created.server_id)
assert callback_finished.wait(timeout=2)
assert service.get(created.server_id).enabled is True
assert service.get(created.server_id).status == "ready"
assert service.tools.definitions()
callbacks[-1](f"mcp.{created.server_id}", "current failure")
assert service.get(created.server_id).enabled is False
assert service.get(created.server_id).status == "unhealthy"
finally:
service.shutdown()
if callback_thread is not None:
callback_thread.join(timeout=2)
def test_header_case_only_rename_preserves_secret() -> None:
service = registry()
config = {
"name": "HTTP",
"transport": "streamable_http",
"url": "https://example.test/mcp",
"secret_header_keys": ["Authorization"],
}
created = service.create(McpServerCreateRequest(**config))
service.put_secret(created.server_id, "Authorization", "synthetic", kind="header")
config["secret_header_keys"] = ["authorization"]
updated = service.update(
created.server_id, McpServerUpdateRequest(**config, version=created.version)
)
assert updated.secret_headers == {"authorization": True}
assert (
service.credentials.resolve(
service._secret_id(created.server_id, "authorization", "header")
)
== "synthetic"
)
def test_environment_secrets_are_case_sensitive_and_delete_independently() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
service.put_secret(created.server_id, "TOKEN", "upper")
service.put_secret(created.server_id, "token", "lower")
assert (
service.credentials.resolve(service._secret_id(created.server_id, "TOKEN"))
== "upper"
)
assert (
service.credentials.resolve(service._secret_id(created.server_id, "token"))
== "lower"
)
service.delete_secret(created.server_id, "TOKEN")
assert service.get(created.server_id).secret_environment == {
"TOKEN": False,
"token": True,
}
def test_legacy_environment_credential_migration_is_idempotent() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN"]))
suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20]
legacy_id = f"mcp.{created.server_id}.{suffix}"
service.credentials.put(legacy_id, "legacy-value")
service._records[created.server_id]["secret_environment_version"] = 1
service._write()
migrated = registry()
assert migrated.get(created.server_id).secret_environment == {"TOKEN": True}
assert (
migrated.credentials.resolve(migrated._secret_id(created.server_id, "TOKEN"))
== "legacy-value"
)
assert not migrated.credentials.has(legacy_id)
migrated.put_secret(created.server_id, "TOKEN", "new-value")
assert (
registry().credentials.resolve(migrated._secret_id(created.server_id, "TOKEN"))
== "new-value"
)
def test_ambiguous_legacy_credentials_are_not_assigned_to_two_variables() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20]
legacy_id = f"mcp.{created.server_id}.{suffix}"
service.credentials.put(legacy_id, "cannot-reconstruct-originals")
service._records[created.server_id]["secret_environment_version"] = 1
service._write()
migrated = registry()
current = migrated.get(created.server_id)
assert current.secret_environment == {"TOKEN": False, "token": False}
assert current.enabled is False
assert current.last_test_succeeded is None
assert migrated.credentials.has(
legacy_id
) # Keep the original ciphertext recoverable.
migrated.put_secret(created.server_id, "TOKEN", "upper")
migrated.put_secret(created.server_id, "token", "lower")
assert registry().get(created.server_id).secret_environment == {
"TOKEN": True,
"token": True,
}
migrated.delete(created.server_id)
assert not migrated.credentials.has(legacy_id)
def test_credential_id_migration_keeps_new_values_and_is_atomic(monkeypatch) -> None:
credentials = EncryptedCredentialStore()
credentials.put("mcp.old", "old-value")
credentials.put("mcp.new", "new-value")
original_write = credentials._write_tokens
def fail_write(_tokens):
raise CredentialStoreError("synthetic failure")
monkeypatch.setattr(credentials, "_write_tokens", fail_write)
with pytest.raises(CredentialStoreError):
credentials.move_many({"mcp.old": "mcp.new"})
assert credentials.resolve("mcp.old") == "old-value"
assert credentials.resolve("mcp.new") == "new-value"
monkeypatch.setattr(credentials, "_write_tokens", original_write)
credentials.move_many({"mcp.old": "mcp.new"})
assert credentials.resolve("mcp.old") is None
assert credentials.resolve("mcp.new") == "new-value"
def test_ambiguous_legacy_secret_is_not_resurrected_after_removing_a_key() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
legacy_id = service._legacy_environment_secret_id(created.server_id, "TOKEN")
service.credentials.put(legacy_id, "ambiguous-old-value")
service._records[created.server_id]["secret_environment_version"] = 1
service._write()
migrated = registry()
migrated.update(
created.server_id,
McpServerUpdateRequest(
**request(secret_environment_keys=["token"]).model_dump(),
version=created.version,
),
)
assert registry().get(created.server_id).secret_environment == {"token": False}
def test_production_rejects_process_launch_even_after_approval() -> None:
service = registry(launch=False)
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.enable(created.server_id)
assert error.value.code == "MCP_SANDBOX_REQUIRED"
@pytest.mark.parametrize("startup,tool", [(120, 300), (1.5, 2.5)])
def test_server_timeouts_survive_bridge_adaptation_and_reload(startup, tool) -> None:
service = registry()
created = service.create(
request(
secret_environment_keys=[],
startup_timeout_seconds=startup,
tool_timeout_seconds=tool,
)
)
service.trust(created.server_id, created.command_digest)
try:
tested = service.test(created.server_id)
assert tested.last_test_succeeded is True
assert tested.startup_timeout_seconds == startup
assert tested.tool_timeout_seconds == tool
restored = registry().get(created.server_id)
assert restored.startup_timeout_seconds == startup
assert restored.tool_timeout_seconds == tool
finally:
service.shutdown()
def test_enable_requires_successful_test_and_update_checks_version() -> None:
service = registry()
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.enable(created.server_id)
assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED"
with pytest.raises(McpRegistryError) as error:
service.update(
created.server_id,
McpServerUpdateRequest(
**request(secret_environment_keys=[]).model_dump(), version=99
),
)
assert error.value.code == "MCP_SERVER_VERSION_CONFLICT"
def test_http_transport_rejects_invalid_cross_transport_fields() -> None:
service = registry()
with pytest.raises(McpRegistryError) as error:
service.create(
request(
transport="streamable_http",
url="https://example.invalid/mcp",
secret_environment_keys=[],
)
)
assert error.value.code == "MCP_CONFIG_INVALID"
def test_registry_rejects_corrupt_persisted_json(tmp_path) -> None:
path = tmp_path / "mcp"
path.mkdir()
(path / "servers.json").write_text("{broken", encoding="utf-8")
with pytest.raises(McpRegistryError) as error:
McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
tmp_path,
allow_process_launch=True,
)
assert error.value.code == "MCP_REGISTRY_INVALID"
def test_registry_rejects_structurally_invalid_record(tmp_path) -> None:
path = tmp_path / "mcp"
path.mkdir()
(path / "servers.json").write_text(
json.dumps({"server-1": {"name": "Broken", "transport": "stdio"}}),
encoding="utf-8",
)
with pytest.raises(McpRegistryError) as error:
McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
tmp_path,
allow_process_launch=True,
)
assert error.value.code == "MCP_REGISTRY_INVALID"
def test_registry_rejects_create_before_exceeding_persisted_limit(
monkeypatch,
) -> None:
service = registry()
service.create(request(name="Only server"))
monkeypatch.setattr("app.extensions.mcp_registry._MAX_MCP_SERVERS", 1)
with pytest.raises(McpRegistryError) as error:
service.create(request(name="One too many"))
assert error.value.code == "MCP_SERVER_LIMIT_REACHED"
assert len(service.list()) == 1
service.shutdown()
def test_stdio_command_is_not_parsed_as_a_shell_string() -> None:
service = registry()
created = service.create(
request(
command=f'"{sys.executable}" "{SERVER}"',
args=[],
secret_environment_keys=[],
)
)
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.test(created.server_id)
assert error.value.code == "PLUGIN_HOST_START_FAILED"
assert service.get(created.server_id).last_test_succeeded is False
service.shutdown()
def test_enabled_server_is_restored_from_persisted_registry() -> None:
first = registry()
created = first.create(request(secret_environment_keys=[]))
first.trust(created.server_id, created.command_digest)
first.test(created.server_id)
first.enable(created.server_id)
first.shutdown()
restored = registry()
restored.restore_enabled()
current = restored.get(created.server_id)
assert current.enabled is True
assert current.status == "ready"
assert any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in restored.tools.definitions()
)
restored.shutdown()
def test_lifecycle_operations_are_serialized_and_tool_names_are_isolated() -> None:
service = registry()
servers = [
service.create(request(name=f"Echo {index}", secret_environment_keys=[]))
for index in range(2)
]
for server in servers:
service.trust(server.server_id, server.command_digest)
service.test(server.server_id)
with ThreadPoolExecutor(max_workers=4) as pool:
enabled = list(
pool.map(lambda item: service.enable(item.server_id), servers * 2)
)
assert all(item.enabled for item in enabled)
names = [
item.name for item in service.tools.definitions() if item.source == "mcp_server"
]
assert len(names) == len(set(names))
assert all(
any(name.startswith(f"mcp.{item.server_id}.") for name in names)
for item in servers
)
with ThreadPoolExecutor(max_workers=4) as pool:
list(pool.map(lambda item: service.disable(item.server_id), servers * 2))
assert not any(item.source == "mcp_server" for item in service.tools.definitions())
service.shutdown()
def _http_result(request_id: int, result: dict) -> httpx.Response:
return httpx.Response(
200,
headers={"content-type": "application/json"},
json={"jsonrpc": "2.0", "id": request_id, "result": result},
)
def test_streamable_http_supports_session_headers_secrets_and_tool_summary(
monkeypatch,
) -> None:
requests: list[httpx.Request] = []
request_timeouts: dict[str, float] = {}
def handler(request_value: httpx.Request) -> httpx.Response:
requests.append(request_value)
if request_value.method == "GET":
return httpx.Response(405)
if request_value.method == "DELETE":
return httpx.Response(405)
payload = json.loads(request_value.content)
timeout = request_value.extensions.get("timeout", {}).get("read")
if isinstance(timeout, (int, float)):
request_timeouts[payload.get("method", "notification")] = float(timeout)
if payload.get("method") == "initialize":
response = _http_result(
payload["id"],
{
"protocolVersion": "2025-11-25",
"capabilities": {"tools": {}},
"serverInfo": {"name": "HTTP Fixture", "version": "1"},
},
)
response.headers["MCP-Session-Id"] = "session-test"
return response
if payload.get("method") == "tools/list":
return _http_result(
payload["id"],
{
"tools": [
{
"name": "echo",
"description": "Echo over HTTP",
"inputSchema": {"type": "object", "properties": {}},
}
]
},
)
if payload.get("method") == "tools/call":
return _http_result(
payload["id"], {"structuredContent": {"transport": "http"}}
)
return httpx.Response(202)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
service = registry()
created = service.create(
McpServerCreateRequest(
name="Remote MCP",
transport="streamable_http",
url="https://mcp.example.test/mcp",
headers={"X-Client": "NotesAgent"},
secret_header_keys=["Authorization"],
)
)
service.put_secret(
created.server_id, "Authorization", "Bearer hidden", kind="header"
)
service.trust(created.server_id, created.command_digest)
tested = service.test(created.server_id)
assert tested.last_test_succeeded is True
assert tested.secret_headers == {"Authorization": True}
assert "Bearer hidden" not in tested.model_dump_json()
assert service.list_tools(created.server_id)[0].remote_name == "echo"
assert any(
request.headers.get("mcp-session-id") == "session-test" for request in requests
)
assert any(
request.headers.get("mcp-protocol-version") == "2025-11-25"
for request in requests
)
assert all(
request.headers.get("authorization") == "Bearer hidden" for request in requests
)
assert request_timeouts["initialize"] == 15
assert request_timeouts["notifications/initialized"] == 15
assert request_timeouts["tools/list"] == 15
enabled = service.enable(created.server_id)
tool_name = service.list_tools(created.server_id)[0].name
result = asyncio.run(
service.tools.execute(
ToolCall(tool_call_id="call-1", name=tool_name, arguments={}),
ToolExecutionContext(run_id="run-1"),
)
)
assert enabled.enabled is True
assert result.success is True
assert result.output == {"transport": "http"}
assert request_timeouts["tools/call"] == 30
service.disable(created.server_id)
service.shutdown()
class _LegacyEventStream(httpx.SyncByteStream):
def __init__(self) -> None:
self.closed = threading.Event()
def __iter__(self):
yield b"event: endpoint\ndata: /messages\n\n"
time.sleep(0.1)
initialize = {
"jsonrpc": "2.0",
"id": 1,
"result": {
"protocolVersion": "2024-11-05",
"capabilities": {"tools": {}},
"serverInfo": {"name": "Legacy Fixture"},
},
}
yield f"data: {json.dumps(initialize)}\n\n".encode()
time.sleep(0.1)
tools = {
"jsonrpc": "2.0",
"id": 2,
"result": {"tools": []},
}
yield f"data: {json.dumps(tools)}\n\n".encode()
self.closed.wait()
def close(self) -> None:
self.closed.set()
def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None:
posted_urls: list[str] = []
event_stream = _LegacyEventStream()
def handler(request_value: httpx.Request) -> httpx.Response:
if request_value.method == "GET":
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=event_stream,
)
posted_urls.append(str(request_value.url))
return httpx.Response(202)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
service = registry()
created = service.create(
McpServerCreateRequest(
name="Legacy MCP",
transport="sse",
url="https://legacy.example.test/sse",
)
)
service.trust(created.server_id, created.command_digest)
tested = service.test(created.server_id)
assert tested.last_test_succeeded is True
assert posted_urls and all(
url == "https://legacy.example.test/messages" for url in posted_urls
)
service.shutdown()
event_stream.close()
class _EndingLegacyEventStream(httpx.SyncByteStream):
def __iter__(self):
yield b"event: endpoint\ndata: /messages\n\n"
def test_legacy_sse_eof_marks_client_unavailable(monkeypatch) -> None:
def handler(_request_value: httpx.Request) -> httpx.Response:
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=_EndingLegacyEventStream(),
)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
broken = threading.Event()
client = McpLegacySseClient(
"https://legacy.example.test/sse",
headers={},
startup_timeout_seconds=1,
on_seen=lambda: None,
on_broken=lambda _message: broken.set(),
on_tools_changed=lambda: None,
)
client.start()
assert broken.wait(timeout=1)
client.stop()
class _CrossOriginLegacyEventStream(httpx.SyncByteStream):
def __iter__(self):
yield b"event: endpoint\ndata: https://attacker.example/messages\n\n"
def test_legacy_sse_rejects_cross_origin_message_endpoint(monkeypatch) -> None:
def handler(request_value: httpx.Request) -> httpx.Response:
assert request_value.method == "GET"
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=_CrossOriginLegacyEventStream(),
)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
service = registry()
created = service.create(
McpServerCreateRequest(
name="Unsafe legacy MCP",
transport="sse",
url="https://legacy.example.test/sse",
)
)
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.test(created.server_id)
assert error.value.code == "MCP_HTTP_RESPONSE_INVALID"
service.shutdown()
+83
View File
@@ -0,0 +1,83 @@
import json
from contextlib import closing
import httpx
import pytest
from app.extensions import mcp
class ChunkStream(httpx.SyncByteStream):
def __init__(self, chunks):
self.chunks = chunks
self.bytes_read = 0
def __iter__(self):
for chunk in self.chunks:
self.bytes_read += len(chunk)
yield chunk
def test_sse_rejects_unterminated_line_before_reading_entire_stream(monkeypatch):
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 1024)
stream = ChunkStream([b"x" * 256] * 256)
with (
closing(httpx.Response(200, stream=stream)) as response,
pytest.raises(mcp.McpBridgeError, match="too large"),
):
list(mcp._iter_sse(response))
assert stream.bytes_read == 1280
def test_sse_limits_combined_event_before_partial_line_is_complete(monkeypatch):
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 32)
stream = ChunkStream(
[b"data: 123456789\n", b"data: 123456789\n", b"x", b"not-read"]
)
with (
closing(httpx.Response(200, stream=stream)) as response,
pytest.raises(mcp.McpBridgeError, match="too large"),
):
list(mcp._iter_sse(response))
assert stream.bytes_read == 33
@pytest.mark.parametrize("separator", [b"\n", b"\r", b"\r\n"])
@pytest.mark.parametrize("chunk_size", [1, 2, 7, 1024])
def test_sse_preserves_utf8_and_line_endings_across_chunks(separator, chunk_size):
payload = json.dumps(
{"jsonrpc": "2.0", "id": 1, "result": {"text": "中文"}}, ensure_ascii=False
)
wire = b"\xef\xbb\xbf" + separator.join(
[
b": heartbeat",
b"event: message",
b"id: replay-1",
("data: " + payload).encode(),
b"",
b"",
]
)
stream = ChunkStream(
[wire[index : index + chunk_size] for index in range(0, len(wire), chunk_size)]
)
with closing(httpx.Response(200, stream=stream)) as response:
assert list(mcp._iter_sse(response)) == [("message", "replay-1", payload)]
def test_sse_event_limit_resets_between_events(monkeypatch):
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 16)
with closing(
httpx.Response(200, stream=ChunkStream([b"data: one\n\ndata: two\r\r"]))
) as response:
assert list(mcp._iter_sse(response)) == [
("message", None, "one"),
("message", None, "two"),
]
def test_sse_preserves_multiline_data_and_final_unterminated_line():
with closing(
httpx.Response(200, stream=ChunkStream([b"data: first\ndata: last"]))
) as response:
assert list(mcp._iter_sse(response)) == [("message", None, "first\nlast")]
+653
View File
@@ -0,0 +1,653 @@
import asyncio
import json
from pathlib import Path
import pytest
from pydantic import TypeAdapter, ValidationError
from app.agent import ToolRegistry
from app.config import BACKEND_DIR, get_settings
from app.container import build_container
from app.contracts import (
PluginCommandContext,
PluginCommandEffect,
PluginNoEffect,
PluginSettingType,
)
from app.extensions import ExtensionError, PluginRuntime
from app.extensions.contributions import _secret_reference
from app.extensions.runtime import DeclarativePluginHost
from app.providers.credentials import CredentialStoreError
TEXT_TOOLS = BACKEND_DIR / "extensions" / "plugins" / "text-tools"
def run(coroutine):
return asyncio.run(coroutine)
def test_command_list_filter_and_lifecycle() -> None:
container = build_container()
commands = container.plugins.list_commands()
palette = container.plugins.list_commands(location="command_palette")
assert [item.command_id for item in commands] == ["text-tools.uppercase-selection"]
assert palette[0].plugin_id == "text-tools"
assert palette[0].icon == "edit"
assert palette[0].when == ["editor.has_selection"]
container.plugins.disable("text-tools")
assert container.plugins.list_commands() == []
with pytest.raises(ExtensionError) as exc:
run(
container.plugins.execute_command(
"text-tools.uppercase-selection",
{},
PluginCommandContext(selection="hello"),
)
)
assert exc.value.code == "PLUGIN_COMMAND_NOT_FOUND"
container.plugins.enable("text-tools")
assert len(container.plugins.list_commands()) == 1
def test_command_executes_with_scoped_context_and_settings() -> None:
container = build_container()
container.plugins.update_settings("text-tools", 1, {"result_limit": 4})
result = run(
container.plugins.execute_command(
"text-tools.uppercase-selection",
{},
PluginCommandContext(
vault_id="default",
note_id="note_private",
file_path="private.md",
selection="abcdef",
),
)
)
assert result.status == "completed"
assert result.effect.type == "notification"
assert result.effect.payload.model_dump() == {
"level": "success",
"message": "ABCD",
}
def test_echo_command_returns_none_for_empty_message() -> None:
host = DeclarativePluginHost()
empty = run(host.execute_command("echo", {}, {}, {}, lambda _: None))
populated = run(
host.execute_command("echo", {"message": "hello"}, {}, {}, lambda _: None)
)
assert isinstance(empty, PluginNoEffect)
assert populated.type == "notification"
assert populated.payload.message == "hello"
@pytest.mark.parametrize(
("effect_type", "payload"),
[
("none", {"unexpected": True}),
("notification", {"level": "debug", "message": "invalid"}),
("navigate", {"route": "https://example.com"}),
("refresh", {"scope": "everything"}),
("job", {"job_id": "invalid job id"}),
],
)
def test_command_effect_rejects_untrusted_payloads(effect_type, payload) -> None:
with pytest.raises(ValidationError):
TypeAdapter(PluginCommandEffect).validate_python(
{"type": effect_type, "payload": payload}
)
def test_command_rejects_missing_context_and_invalid_arguments() -> None:
container = build_container()
with pytest.raises(ExtensionError) as context_error:
run(
container.plugins.execute_command(
"text-tools.uppercase-selection", {}, PluginCommandContext()
)
)
assert context_error.value.code == "PLUGIN_COMMAND_CONTEXT_INVALID"
with pytest.raises(ExtensionError) as argument_error:
run(
container.plugins.execute_command(
"text-tools.uppercase-selection",
{"unknown": True},
PluginCommandContext(selection="hello"),
)
)
assert argument_error.value.code == "PLUGIN_COMMAND_ARGUMENT_INVALID"
audit = container.plugins.commands.audit_events()
assert [event.error_code for event in audit[-2:]] == [
"PLUGIN_COMMAND_CONTEXT_INVALID",
"PLUGIN_COMMAND_ARGUMENT_INVALID",
]
# 审计事件不得携带参数、正文选区或返回 effect。
assert "hello" not in repr(audit)
def test_command_only_receives_declared_context() -> None:
class CapturingHost(DeclarativePluginHost):
def __init__(self) -> None:
self.context = None
async def execute_command(
self, handler, arguments, context, settings, resolve_secret
):
self.context = context
return PluginNoEffect()
host = CapturingHost()
runtime = PluginRuntime(ToolRegistry(), host=host)
runtime.install(TEXT_TOOLS)
runtime.enable("text-tools")
run(
runtime.execute_command(
"text-tools.uppercase-selection",
{},
PluginCommandContext(
vault_id="default", note_id="note_private", selection="visible"
),
)
)
assert host.context == {"selection": "visible"}
def test_command_resolves_only_declared_plugin_secrets(tmp_path: Path) -> None:
class SecretHost(DeclarativePluginHost):
def __init__(self) -> None:
self.secret = None
self.denied_code = None
async def execute_command(
self, handler, arguments, context, settings, resolve_secret
):
self.secret = resolve_secret("api_key")
try:
resolve_secret("undeclared")
except ExtensionError as exc:
self.denied_code = exc.code
return PluginNoEffect()
host = SecretHost()
runtime = PluginRuntime(ToolRegistry(), host=host)
package = tmp_path / "secret-command"
package.mkdir()
(package / "plugin.yaml").write_text(
"""
id: secret-command
name: Secret Command
version: 1.0.0
permissions: [secrets.use]
contributes:
commands: [secret-command.run]
settings_sections: [secret-command.general]
backend:
type: internal_rpc
transport: none
""".strip(),
encoding="utf-8",
)
(package / "commands.yaml").write_text(
"""
commands:
- command_id: secret-command.run
title: Secret Command
locations: [command_palette]
secrets: [api_key]
handler: echo
""".strip(),
encoding="utf-8",
)
(package / "settings.yaml").write_text(
"""
section_id: secret-command.general
schema_version: 1
fields:
- key: api_key
label: API Key
type: secret
""".strip(),
encoding="utf-8",
)
runtime.install(package)
runtime.set_permissions("secret-command", ["secrets.use"])
runtime.enable("secret-command")
runtime.put_setting_secret("secret-command", "api_key", "runtime-only-secret")
run(
runtime.execute_command(
"secret-command.run",
{},
PluginCommandContext(selection="visible"),
)
)
assert host.secret == "runtime-only-secret"
assert host.denied_code == "PLUGIN_SECRET_ACCESS_DENIED"
assert "runtime-only-secret" not in repr(runtime.commands.audit_events())
def test_settings_schema_contains_defaults_and_hides_secret() -> None:
container = build_container()
schema = container.plugins.get_settings("text-tools")
by_key = {field.key: field for field in schema.fields}
assert schema.schema_version == 1
assert schema.values == {
"result_limit": 100,
"label_prefix": "",
"output_style": "notification",
"enabled_hint": True,
}
assert "api_key" not in schema.values
assert schema.secrets["api_key"].configured is False
assert by_key["api_key"].type == PluginSettingType.secret
def test_settings_update_validates_version_type_bounds_and_secret_boundary() -> None:
container = build_container()
updated = container.plugins.update_settings(
"text-tools", 1, {"result_limit": 20, "output_style": "compact"}
)
assert updated.values["result_limit"] == 20
assert updated.values["output_style"] == "compact"
cases = [
(2, {}, "PLUGIN_SETTINGS_VERSION_CONFLICT"),
(1, {"result_limit": 0}, "PLUGIN_SETTINGS_FIELD_INVALID"),
(1, {"enabled_hint": "yes"}, "PLUGIN_SETTINGS_FIELD_INVALID"),
(1, {"output_style": "unknown"}, "PLUGIN_SETTINGS_FIELD_INVALID"),
(1, {"api_key": "plaintext"}, "PLUGIN_SETTINGS_FIELD_INVALID"),
(1, {"unknown": True}, "PLUGIN_SETTINGS_FIELD_INVALID"),
]
for version, values, code in cases:
with pytest.raises(ExtensionError) as exc:
container.plugins.update_settings("text-tools", version, values)
assert exc.value.code == code
def test_required_plain_setting_blocks_enable_until_configured(tmp_path: Path) -> None:
package = tmp_path / "required-setting"
package.mkdir()
(package / "plugin.yaml").write_text(
"""
id: required-setting
name: Required Setting
version: 1.0.0
contributes:
commands: [required-setting.run]
settings_sections: [required-setting.general]
backend:
type: internal_rpc
transport: none
""".strip(),
encoding="utf-8",
)
(package / "commands.yaml").write_text(
"""
commands:
- command_id: required-setting.run
title: Required Setting
locations: [command_palette]
handler: echo
""".strip(),
encoding="utf-8",
)
(package / "settings.yaml").write_text(
"""
section_id: required-setting.general
schema_version: 1
fields:
- key: endpoint
label: Endpoint
type: string
required: true
""".strip(),
encoding="utf-8",
)
runtime = PluginRuntime(ToolRegistry())
runtime.install(package)
with pytest.raises(ExtensionError) as exc:
runtime.enable("required-setting")
assert exc.value.code == "PLUGIN_SETTINGS_REQUIRED"
assert runtime.get("required-setting").status == "installed"
runtime.update_settings("required-setting", 1, {"endpoint": "local"})
assert runtime.enable("required-setting").status == "ready"
def test_secret_roundtrip_never_enters_plain_settings_storage() -> None:
container = build_container()
plaintext = "stage-d-secret-value"
status = container.plugins.put_setting_secret("text-tools", "api_key", plaintext)
schema = container.plugins.get_settings("text-tools")
settings_path = get_settings().data_dir / "plugins" / "settings.json"
credentials_path = get_settings().data_dir / "credentials" / "credentials.json"
assert status.configured is True
assert schema.secrets["api_key"].configured is True
assert "api_key" not in schema.values
assert plaintext not in settings_path.read_text(encoding="utf-8")
assert plaintext not in credentials_path.read_text(encoding="utf-8")
stored_settings = json.loads(settings_path.read_text(encoding="utf-8"))
reference = stored_settings["text-tools"]["secret_refs"]["api_key"]
assert reference.startswith("plugin.")
assert len(reference) == 71
assert "text-tools" not in reference and "api_key" not in reference
assert container.credentials.resolve(reference) == plaintext
deleted = container.plugins.delete_setting_secret("text-tools", "api_key")
assert deleted.configured is False
assert container.credentials.resolve(reference) is None
def test_uninstall_removes_plugin_settings_and_secret_namespace() -> None:
container = build_container()
container.plugins.update_settings("text-tools", 1, {"result_limit": 12})
container.plugins.put_setting_secret("text-tools", "api_key", "temporary")
settings_path = get_settings().data_dir / "plugins" / "settings.json"
reference = json.loads(settings_path.read_text(encoding="utf-8"))[
"text-tools"
]["secret_refs"]["api_key"]
container.plugins.uninstall("text-tools")
stored = json.loads(settings_path.read_text(encoding="utf-8"))
assert "text-tools" not in stored
assert container.credentials.resolve(reference) is None
def test_plugin_secret_reference_has_fixed_credential_safe_length() -> None:
reference = _secret_reference("p" * 512, "k" * 128)
assert reference.startswith("plugin.")
assert len(reference) <= 128
def test_tampered_secret_reference_cannot_cross_credential_namespace() -> None:
container = build_container()
container.credentials.put("openai", "provider-private-secret")
settings_path = get_settings().data_dir / "plugins" / "settings.json"
settings_path.parent.mkdir(parents=True, exist_ok=True)
settings_path.write_text(
json.dumps(
{
"text-tools": {
"schema_version": 1,
"values": {},
"secret_refs": {"api_key": "openai"},
}
}
),
encoding="utf-8",
)
with pytest.raises(ExtensionError) as read_error:
container.plugins.get_settings("text-tools")
with pytest.raises(ExtensionError) as uninstall_error:
container.plugins.uninstall("text-tools")
assert read_error.value.code == "PLUGIN_STORAGE_ERROR"
assert uninstall_error.value.code == "PLUGIN_STORAGE_ERROR"
assert container.credentials.resolve("openai") == "provider-private-secret"
def test_secret_delete_restores_reference_when_credential_delete_fails(
monkeypatch,
) -> None:
container = build_container()
container.plugins.put_setting_secret("text-tools", "api_key", "keep-me")
settings_path = get_settings().data_dir / "plugins" / "settings.json"
original = settings_path.read_text(encoding="utf-8")
reference = _secret_reference("text-tools", "api_key")
def fail_delete(_credential_id: str) -> bool:
raise CredentialStoreError("injected delete failure")
monkeypatch.setattr(container.credentials, "delete", fail_delete)
with pytest.raises(ExtensionError) as exc:
container.plugins.delete_setting_secret("text-tools", "api_key")
assert exc.value.code == "PLUGIN_SECRET_STORE_ERROR"
assert settings_path.read_text(encoding="utf-8") == original
assert container.credentials.resolve(reference) == "keep-me"
def test_uninstall_restores_settings_when_atomic_secret_delete_fails(
monkeypatch,
) -> None:
container = build_container()
container.plugins.update_settings("text-tools", 1, {"result_limit": 12})
container.plugins.put_setting_secret("text-tools", "api_key", "keep-me")
settings_path = get_settings().data_dir / "plugins" / "settings.json"
original = settings_path.read_text(encoding="utf-8")
reference = _secret_reference("text-tools", "api_key")
def fail_delete_many(_credential_ids: list[str]) -> set[str]:
raise CredentialStoreError("injected batch delete failure")
monkeypatch.setattr(container.credentials, "delete_many", fail_delete_many)
with pytest.raises(ExtensionError) as exc:
container.plugins.uninstall("text-tools")
assert exc.value.code == "PLUGIN_SECRET_STORE_ERROR"
assert settings_path.read_text(encoding="utf-8") == original
assert container.credentials.resolve(reference) == "keep-me"
assert container.plugins.get("text-tools").manifest.plugin_id == "text-tools"
def test_invalid_command_and_settings_manifest_are_rejected(tmp_path: Path) -> None:
invalid_command = tmp_path / "invalid-command"
invalid_command.mkdir()
(invalid_command / "plugin.yaml").write_text(
"""
id: invalid-command
name: Invalid Command
version: 1.0.0
contributes:
commands: [other.run]
""".strip(),
encoding="utf-8",
)
(invalid_command / "commands.yaml").write_text(
"""
commands:
- command_id: other.run
title: Invalid
locations: [command_palette]
handler: echo
""".strip(),
encoding="utf-8",
)
invalid_settings = tmp_path / "invalid-settings"
invalid_settings.mkdir()
(invalid_settings / "plugin.yaml").write_text(
"""
id: invalid-settings
name: Invalid Settings
version: 1.0.0
contributes:
settings_sections: [invalid-settings.general]
""".strip(),
encoding="utf-8",
)
(invalid_settings / "settings.yaml").write_text(
"""
section_id: invalid-settings.general
schema_version: 1
fields:
- key: token
label: Token
type: secret
default: leaked-default
""".strip(),
encoding="utf-8",
)
runtime = PluginRuntime(ToolRegistry())
with pytest.raises(ExtensionError) as command_error:
runtime.install(invalid_command)
assert command_error.value.code == "PLUGIN_COMMAND_INVALID"
with pytest.raises(ExtensionError) as settings_error:
runtime.install(invalid_settings)
assert settings_error.value.code == "PLUGIN_SETTINGS_SCHEMA_INVALID"
@pytest.mark.parametrize("bound", [".nan", ".inf", "-.inf"])
def test_non_finite_setting_bounds_are_rejected(tmp_path: Path, bound: str) -> None:
package = tmp_path / f"invalid-bound-{bound.replace('.', 'dot').replace('-', 'neg')}"
package.mkdir()
(package / "plugin.yaml").write_text(
"""
id: invalid-bound
name: Invalid Bound
version: 1.0.0
contributes:
settings_sections: [invalid-bound.general]
backend:
type: none
transport: none
""".strip(),
encoding="utf-8",
)
(package / "settings.yaml").write_text(
f"""
section_id: invalid-bound.general
schema_version: 1
fields:
- key: limit
label: Limit
type: number
minimum: {bound}
""".strip(),
encoding="utf-8",
)
with pytest.raises(ExtensionError) as exc:
PluginRuntime(ToolRegistry()).install(package)
assert exc.value.code == "PLUGIN_SETTINGS_SCHEMA_INVALID"
assert "must be finite" in exc.value.message
def test_null_command_list_returns_stable_manifest_error(tmp_path: Path) -> None:
package = tmp_path / "null-commands"
package.mkdir()
(package / "plugin.yaml").write_text(
"""
id: null-commands
name: Null Commands
version: 1.0.0
contributes:
commands: []
backend:
type: internal_rpc
transport: none
""".strip(),
encoding="utf-8",
)
(package / "commands.yaml").write_text("commands:\n", encoding="utf-8")
with pytest.raises(ExtensionError) as exc:
PluginRuntime(ToolRegistry()).install(package)
assert exc.value.code == "EXTENSION_MANIFEST_INVALID"
def test_external_command_schema_reference_is_rejected(tmp_path: Path) -> None:
package = tmp_path / "external-ref"
package.mkdir()
(package / "plugin.yaml").write_text(
"""
id: external-ref
name: External Ref
version: 1.0.0
contributes:
commands: [external-ref.run]
backend:
type: internal_rpc
transport: none
""".strip(),
encoding="utf-8",
)
(package / "commands.yaml").write_text(
"""
commands:
- command_id: external-ref.run
title: External Ref
locations: [command_palette]
handler: echo
parameters:
$ref: file:///host/private-schema.json
""".strip(),
encoding="utf-8",
)
with pytest.raises(ExtensionError) as exc:
PluginRuntime(ToolRegistry()).install(package)
assert exc.value.code == "PLUGIN_COMMAND_INVALID"
assert "External JSON Schema reference" in exc.value.message
def test_settings_missing_and_secret_field_errors_are_stable() -> None:
container = build_container()
with pytest.raises(ExtensionError) as missing:
container.plugins.get_settings("does-not-exist")
assert missing.value.code == "PLUGIN_NOT_FOUND"
with pytest.raises(ExtensionError) as field:
container.plugins.put_setting_secret("text-tools", "result_limit", "secret")
assert field.value.code == "PLUGIN_SECRET_FIELD_NOT_FOUND"
with pytest.raises(ExtensionError) as empty:
container.plugins.put_setting_secret("text-tools", "api_key", "")
assert empty.value.code == "PLUGIN_SECRET_VALUE_INVALID"
def test_corrupted_plugin_settings_namespace_returns_stable_error() -> None:
container = build_container()
settings_path = get_settings().data_dir / "plugins" / "settings.json"
settings_path.parent.mkdir(parents=True, exist_ok=True)
settings_path.write_text('{"text-tools": []}', encoding="utf-8")
with pytest.raises(ExtensionError) as exc:
container.plugins.get_settings("text-tools")
assert exc.value.code == "PLUGIN_STORAGE_ERROR"
with pytest.raises(ExtensionError) as secret_exc:
container.plugins.put_setting_secret("text-tools", "api_key", "must-not-orphan")
assert secret_exc.value.code == "PLUGIN_STORAGE_ERROR"
credentials_path = get_settings().data_dir / "credentials" / "credentials.json"
credential_ids = (
json.loads(credentials_path.read_text(encoding="utf-8")).keys()
if credentials_path.exists()
else []
)
assert not any(item.startswith("plugin.") for item in credential_ids)
+77
View File
@@ -436,6 +436,83 @@ def test_fts_pagination_is_not_truncated_at_one_thousand(vault) -> None:
assert len(response.items) == 10 assert len(response.items) == 10
def test_fts_score_threshold_filters_before_total(vault) -> None:
"""score_threshold 先于计数与分页生效:total 反映过滤后数量,与 items 一致。
高阈值过滤掉全部结果时 total==0 items 为空杜绝空页但 total>0
不一致审阅 P2-7
"""
from app.retrieval.engine import engine
from app.services import note_service
# 10 个 block,含「目标」次数递增,bm25 分数各异,min-max 归一化后分数落在 [0,1]
markdown = "\n\n".join(f"{'目标' * i} 分隔内容" for i in range(1, 11))
asyncio.run(
note_service.create_note(title="阈值过滤", markdown=markdown, folder="", tags=[])
)
all_hits = asyncio.run(
engine.search(
SearchRequest(query="目标", mode=SearchMode.fts, limit=20, score_threshold=0.0)
)
)
filtered = asyncio.run(
engine.search(
SearchRequest(query="目标", mode=SearchMode.fts, limit=20, score_threshold=0.5)
)
)
none = asyncio.run(
engine.search(
SearchRequest(query="目标", mode=SearchMode.fts, limit=20, score_threshold=2.0)
)
)
assert all_hits.page.total >= 10
assert 0 < filtered.page.total < all_hits.page.total # 阈值过滤掉部分而非全部
assert filtered.page.total == len(filtered.items)
assert none.page.total == 0
assert none.items == []
def test_fts_offset_beyond_end_reports_real_total(vault) -> None:
"""offset 越过末页时 items 为空,但 total 仍为真实命中数而非归零。"""
from app.retrieval.engine import engine
from app.services import note_service
asyncio.run(
note_service.create_note(title="越界分页", markdown="检索 检索 检索 检索", folder="", tags=[])
)
resp = asyncio.run(
engine.search(SearchRequest(query="检索", mode=SearchMode.fts, limit=10, offset=100))
)
assert resp.page.total >= 1
assert resp.items == []
def test_fts_not_truncated_at_five_thousand(vault) -> None:
"""FTS 结果不再被 5000 条上限截断:>5000 命中时 total 为真实计数,末页仍可访问。"""
from app.retrieval.engine import engine
from app.services import note_service
markdown = "\n\n".join(f"共同词 q{i}" for i in range(5010))
asyncio.run(
note_service.create_note(title="五千条分页", markdown=markdown, folder="", tags=[])
)
first = asyncio.run(
engine.search(SearchRequest(query="共同词", mode=SearchMode.fts, limit=10, offset=0))
)
assert first.page.total == 5010
assert len(first.items) == 10
last = asyncio.run(
engine.search(SearchRequest(query="共同词", mode=SearchMode.fts, limit=10, offset=5005))
)
assert last.page.total == 5010
assert len(last.items) == 5
# --------------------------------------------------------------------------- # # --------------------------------------------------------------------------- #
# 审阅回归:PATCH tags 语义 / 向量-块一致性 / 过滤漏召回 / rebuild 语义与回滚 # 审阅回归:PATCH tags 语义 / 向量-块一致性 / 过滤漏召回 / rebuild 语义与回滚
# --------------------------------------------------------------------------- # # --------------------------------------------------------------------------- #
+65
View File
@@ -0,0 +1,65 @@
import pytest
from app.schema_security import (
ExternalSchemaReferenceError,
UnresolvableLocalSchemaReferenceError,
reject_external_schema_references,
)
@pytest.mark.parametrize(
"schema",
[
{"$ref": "file:///host/private-schema.json"},
{"properties": {"value": {"$ref": "https://schema.invalid/value.json"}}},
{"allOf": [{"$dynamicRef": "https://schema.invalid/dynamic"}]},
],
)
def test_external_json_schema_references_are_rejected(schema) -> None:
with pytest.raises(ExternalSchemaReferenceError):
reject_external_schema_references(schema)
def test_local_json_schema_fragment_reference_is_allowed() -> None:
reject_external_schema_references(
{
"$defs": {"value": {"type": "string"}},
"properties": {"value": {"$ref": "#/$defs/value"}},
}
)
@pytest.mark.parametrize("reference", ["#/$defs/missing", "#missing-anchor"])
def test_unresolvable_local_schema_reference_is_rejected(reference: str) -> None:
with pytest.raises(UnresolvableLocalSchemaReferenceError):
reject_external_schema_references({"type": "object", "$ref": reference})
def test_root_reference_cannot_use_anchor_from_nested_schema_resource() -> None:
schema = {
"$defs": {
"nested": {
"$id": "nested",
"$anchor": "inside",
"type": "string",
}
},
"properties": {"value": {"$ref": "#inside"}},
}
with pytest.raises(UnresolvableLocalSchemaReferenceError):
reject_external_schema_references(schema)
def test_nested_schema_resource_can_resolve_its_own_anchor() -> None:
schema = {
"$defs": {
"nested": {
"$id": "nested",
"$anchor": "inside",
"allOf": [{"$ref": "#inside"}],
}
}
}
reject_external_schema_references(schema)
+116
View File
@@ -0,0 +1,116 @@
import asyncio
import pytest
from app.config import get_settings
from app.contracts import (
FolderCreateRequest,
FolderDeleteRequest,
FolderRenameRequest,
NoteCreateRequest,
NoteRenameRequest,
WorkspaceOpenRequest,
)
from app.errors import ApiError
from app.routes import (
create_note,
create_workspace_folder,
delete_workspace_folder,
get_note,
get_workspace_tree,
open_workspace,
rename_note,
rename_workspace_folder,
)
def test_open_workspace_indexes_real_markdown_and_returns_tree() -> None:
vault = get_settings().vault_path
note_path = vault / "课程" / "操作系统.md"
note_path.parent.mkdir(parents=True)
note_path.write_text("# 操作系统\n\n进程调度。\n", encoding="utf-8")
snapshot = asyncio.run(open_workspace(WorkspaceOpenRequest()))
assert snapshot.workspace.path == str(vault.resolve())
assert snapshot.workspace.requires_refresh is False
assert snapshot.workspace.file_count == snapshot.workspace.indexed_note_count == 1
folder = snapshot.items[0]
assert folder.path == "/课程"
assert folder.children[0].path == "/课程/操作系统.md"
assert folder.children[0].note_id is not None
def test_open_workspace_rejects_unconfigured_path() -> None:
with pytest.raises(ApiError) as error:
asyncio.run(open_workspace(WorkspaceOpenRequest(path="C:/another-vault")))
assert error.value.code == "WORKSPACE_PATH_MISMATCH"
def test_note_rename_preserves_identity_and_content() -> None:
created = asyncio.run(
create_note(
NoteCreateRequest(
title="旧名称", markdown="# 标题不变\n\n真实正文。\n", folder="课程"
)
)
)
renamed = asyncio.run(
rename_note(created.note_id, NoteRenameRequest(file_name="新名称.md"))
)
assert renamed.note_id == created.note_id
assert renamed.file_path == "课程/新名称.md"
assert renamed.title == "新名称"
assert renamed.markdown == "# 标题不变\n\n真实正文。\n"
assert not (get_settings().vault_path / "课程" / "旧名称.md").exists()
def test_folder_lifecycle_updates_database_and_vectors() -> None:
folder = asyncio.run(
create_workspace_folder(FolderCreateRequest(parent="/", name="课程"))
)
created = asyncio.run(
create_note(
NoteCreateRequest(title="网络", markdown="# 网络\n\nTCP。\n", folder="课程")
)
)
renamed_folder = asyncio.run(
rename_workspace_folder(
FolderRenameRequest(path=folder.path, new_name="计算机课程")
)
)
moved_note = asyncio.run(get_note(created.note_id))
assert renamed_folder.path == "/计算机课程"
assert moved_note.note_id == created.note_id
assert moved_note.file_path == "计算机课程/网络.md"
assert asyncio.run(get_workspace_tree())[0].children[0].note_id == created.note_id
response = asyncio.run(
delete_workspace_folder(FolderDeleteRequest(path=renamed_folder.path))
)
assert response.status == "completed"
with pytest.raises(ApiError) as error:
asyncio.run(get_note(created.note_id))
assert error.value.code == "RESOURCE_NOT_FOUND"
assert asyncio.run(get_workspace_tree()) == []
def test_workspace_openapi_paths_are_published() -> None:
from app.main import app
paths = app.openapi()["paths"]
assert {
"/api/workspace",
"/api/workspace/open",
"/api/workspace/tree",
"/api/workspace/folders",
"/api/workspace/folders/rename",
"/api/workspace/folders/delete",
"/api/notes/{note_id}/rename",
} <= paths.keys()
+2
View File
@@ -374,6 +374,7 @@ dependencies = [
{ name = "httpx" }, { name = "httpx" },
{ name = "jsonschema" }, { name = "jsonschema" },
{ name = "pyyaml" }, { name = "pyyaml" },
{ name = "referencing" },
{ name = "sqlite-vec" }, { name = "sqlite-vec" },
{ name = "uvicorn", extra = ["standard"] }, { name = "uvicorn", extra = ["standard"] },
] ]
@@ -390,6 +391,7 @@ requires-dist = [
{ name = "httpx", specifier = ">=0.28,<1.0" }, { name = "httpx", specifier = ">=0.28,<1.0" },
{ name = "jsonschema", specifier = ">=4.25,<5.0" }, { name = "jsonschema", specifier = ">=4.25,<5.0" },
{ name = "pyyaml", specifier = ">=6.0,<7.0" }, { name = "pyyaml", specifier = ">=6.0,<7.0" },
{ name = "referencing", specifier = ">=0.36,<1.0" },
{ name = "sqlite-vec", specifier = ">=0.1.9" }, { name = "sqlite-vec", specifier = ">=0.1.9" },
{ name = "uvicorn", extras = ["standard"], specifier = ">=0.35,<1.0" }, { name = "uvicorn", extras = ["standard"], specifier = ">=0.35,<1.0" },
] ]
+1 -1
View File
@@ -4,7 +4,7 @@
<meta charset="UTF-8" /> <meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" /> <meta name="viewport" content="width=device-width, initial-scale=1.0" />
<meta name="theme-color" content="#171717" /> <meta name="theme-color" content="#171717" />
<title>Notes Agent</title> <title>NotesAgent</title>
</head> </head>
<body> <body>
<div id="app"></div> <div id="app"></div>
+2
View File
@@ -0,0 +1,2 @@
allowBuilds:
esbuild: true
+2 -1
View File
@@ -73,7 +73,7 @@ defineExpose({ openCitation })
flex-direction: column; flex-direction: column;
height: 100vh; height: 100vh;
width: 100vw; width: 100vw;
background: var(--color-background-primary); background: var(--color-background-secondary);
color: var(--color-text-primary); color: var(--color-text-primary);
} }
@@ -89,5 +89,6 @@ defineExpose({ openCitation })
min-width: 0; min-width: 0;
overflow: hidden; overflow: hidden;
background: var(--color-background-primary); background: var(--color-background-primary);
isolation: isolate;
} }
</style> </style>
@@ -0,0 +1,77 @@
// @vitest-environment happy-dom
import { flushPromises, mount } from '@vue/test-utils'
import { createPinia, setActivePinia } from 'pinia'
import { createMemoryHistory, createRouter } from 'vue-router'
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import * as pluginService from '@/services/pluginService'
import { useEditorStore } from '@/stores/editor'
import { useWorkspaceStore } from '@/stores/workspace'
import CommandPalette from './CommandPalette.vue'
vi.mock('@/services/pluginService', async (loadOriginal) => {
const original = await loadOriginal<typeof import('@/services/pluginService')>()
return { ...original, listPluginCommands: vi.fn(), executePluginCommand: vi.fn() }
})
beforeEach(() => {
setActivePinia(createPinia())
vi.mocked(pluginService.listPluginCommands).mockResolvedValue([{
command_id: 'demo.selection',
plugin_id: 'demo',
title: '处理选区',
description: '',
icon: null,
locations: ['command_palette'],
when: ['workspace.has_vault', 'editor.has_note', 'editor.has_selection'],
parameters: { type: 'object', properties: {}, additionalProperties: false },
enabled: true,
}])
vi.mocked(pluginService.executePluginCommand).mockResolvedValue({
command_id: 'demo.selection',
status: 'completed',
effect: { type: 'notification', payload: { level: 'success', message: '完成' } },
})
})
afterEach(() => {
document.body.innerHTML = ''
vi.restoreAllMocks()
})
describe('CommandPalette Plugin Command', () => {
it('filters by when context and sends stable backend identities plus the captured selection', async () => {
const workspace = useWorkspaceStore()
workspace.hasVault = true
workspace.vaultId = 'vault-default'
const editor = useEditorStore()
editor.currentNoteId = 'note-1'
editor.currentFilePath = '/note.md'
vi.spyOn(window, 'getSelection').mockReturnValue({
toString: () => 'selected text',
} as Selection)
const router = createRouter({
history: createMemoryHistory(),
routes: [{ path: '/', component: { template: '<div />' } }],
})
await router.push('/')
const wrapper = mount(CommandPalette, { attachTo: document.body, global: { plugins: [router] } })
window.dispatchEvent(new KeyboardEvent('keydown', { key: 'p', ctrlKey: true }))
await flushPromises()
const command = Array.from(document.querySelectorAll('button')).find((button) => button.textContent?.includes('处理选区'))
expect(command).toBeTruthy()
command!.click()
await flushPromises()
expect(pluginService.executePluginCommand).toHaveBeenCalledWith('demo.selection', {}, {
vault_id: 'vault-default',
note_id: 'note-1',
file_path: '/note.md',
selection: 'selected text',
})
expect(document.body.textContent).toContain('完成')
wrapper.unmount()
})
})
@@ -5,18 +5,26 @@ import { useEditorStore } from '@/stores/editor'
import { useThemeStore } from '@/stores/theme' import { useThemeStore } from '@/stores/theme'
import { useWorkspaceStore } from '@/stores/workspace' import { useWorkspaceStore } from '@/stores/workspace'
import * as workspaceService from '@/services/workspaceService' import * as workspaceService from '@/services/workspaceService'
import * as pluginService from '@/services/pluginService'
import type { PluginCommand, PluginCommandEffect } from '@/contracts'
import { usePluginStore } from '@/stores/plugin'
const router = useRouter() const router = useRouter()
const editorStore = useEditorStore() const editorStore = useEditorStore()
const themeStore = useThemeStore() const themeStore = useThemeStore()
const workspaceStore = useWorkspaceStore() const workspaceStore = useWorkspaceStore()
const pluginStore = usePluginStore()
const open = ref(false) const open = ref(false)
const query = ref('') const query = ref('')
const input = ref<HTMLInputElement | null>(null) const input = ref<HTMLInputElement | null>(null)
const pluginCommands = ref<PluginCommand[]>([])
const commandError = ref('')
const commandNotice = ref('')
const selectionSnapshot = ref<string | null>(null)
interface Command { id: string; label: string; hint: string; run: () => void | Promise<void> } interface Command { id: string; label: string; hint: string; run: () => void | Promise<void> }
const commands = computed<Command[]>(() => [ const builtinCommands = computed<Command[]>(() => [
{ id: 'workspace', label: '打开工作区', hint: '导航', run: () => router.push('/workspace') }, { id: 'workspace', label: '打开工作区', hint: '导航', run: () => router.push('/workspace') },
{ id: 'search', label: '全局搜索', hint: '导航', run: () => router.push('/search') }, { id: 'search', label: '全局搜索', hint: '导航', run: () => router.push('/search') },
{ id: 'chat', label: '打开 AI 对话', hint: '导航', run: () => router.push('/chat') }, { id: 'chat', label: '打开 AI 对话', hint: '导航', run: () => router.push('/chat') },
@@ -28,14 +36,37 @@ const commands = computed<Command[]>(() => [
{ id: 'new-note', label: '创建笔记', hint: '工作区', run: createNote }, { id: 'new-note', label: '创建笔记', hint: '工作区', run: createNote },
]) ])
const commands = computed<Command[]>(() => [
...builtinCommands.value,
...pluginCommands.value.filter(isPluginCommandAvailable).map((command) => ({
id: 'plugin:' + command.command_id,
label: command.title,
hint: 'Plugin · ' + command.plugin_id,
run: () => executePluginCommand(command),
})),
])
function isPluginCommandAvailable(command: PluginCommand) {
if (!command.enabled) return false
return command.when.every((condition) => {
if (condition === 'workspace.has_vault') return Boolean(workspaceStore.vaultId)
if (condition === 'editor.has_note') return Boolean(editorStore.currentNoteId)
if (condition === 'editor.has_selection') return Boolean(selectionSnapshot.value)
return false
})
}
const filteredCommands = computed(() => { const filteredCommands = computed(() => {
const value = query.value.trim().toLocaleLowerCase() const value = query.value.trim().toLocaleLowerCase()
return value ? commands.value.filter((command) => `${command.label} ${command.hint}`.toLocaleLowerCase().includes(value)) : commands.value return value ? commands.value.filter((command) => `${command.label} ${command.hint}`.toLocaleLowerCase().includes(value)) : commands.value
}) })
function show() { function show() {
selectionSnapshot.value = window.getSelection()?.toString() || null
open.value = true open.value = true
query.value = '' query.value = ''
commandError.value = ''
void loadPluginCommands()
void nextTick(() => input.value?.focus()) void nextTick(() => input.value?.focus())
} }
@@ -44,7 +75,11 @@ function hide() { open.value = false }
async function execute(command: Command | undefined) { async function execute(command: Command | undefined) {
if (!command) return if (!command) return
hide() hide()
try {
await command.run() await command.run()
} catch (error) {
commandNotice.value = error instanceof Error ? error.message : '命令执行失败'
}
} }
async function createNote() { async function createNote() {
@@ -58,6 +93,55 @@ async function createNote() {
await router.push('/workspace') await router.push('/workspace')
} }
async function loadPluginCommands() {
try {
pluginCommands.value = await pluginService.listPluginCommands('command_palette')
} catch (error) {
commandError.value = error instanceof Error ? error.message : 'Plugin 命令加载失败'
}
}
function hasRequiredArguments(command: PluginCommand) {
return Array.isArray(command.parameters.required) && command.parameters.required.length > 0
}
async function executePluginCommand(command: PluginCommand) {
if (hasRequiredArguments(command)) {
pluginStore.selectPlugin(command.plugin_id)
await router.push('/extensions/plugins')
commandNotice.value = '请在 Plugin 详情页填写参数后执行“' + command.title + '”。'
return
}
const result = await pluginService.executePluginCommand(command.command_id, {}, {
vault_id: workspaceStore.hasVault ? workspaceStore.vaultId : null,
note_id: editorStore.currentNoteId,
file_path: editorStore.currentFilePath,
selection: selectionSnapshot.value,
})
await applyPluginEffect(result.effect)
}
async function applyPluginEffect(effect: PluginCommandEffect) {
if (effect.type === 'notification') { commandNotice.value = effect.payload.message; return }
if (effect.type === 'navigate') {
const routes: Record<string, string> = {
'vault-entry': '/', workspace: '/workspace', search: '/search', chat: '/chat',
agent: '/agent/runs', tasks: '/tasks', skills: '/extensions/skills',
plugins: '/extensions/plugins', themes: '/themes', settings: '/settings',
}
await router.push(routes[effect.payload.route])
return
}
if (effect.type === 'refresh') {
if (effect.payload.scope === 'plugins') await pluginStore.loadPlugins()
if (effect.payload.scope === 'commands') await loadPluginCommands()
commandNotice.value = '相关数据已刷新。'
return
}
if (effect.type === 'job') { commandNotice.value = '后台任务已创建:' + effect.payload.job_id; return }
commandNotice.value = 'Plugin 命令执行完成。'
}
function handleKeydown(event: KeyboardEvent) { function handleKeydown(event: KeyboardEvent) {
if ((event.ctrlKey || event.metaKey) && event.key.toLocaleLowerCase() === 'p') { if ((event.ctrlKey || event.metaKey) && event.key.toLocaleLowerCase() === 'p') {
event.preventDefault() event.preventDefault()
@@ -72,10 +156,14 @@ onBeforeUnmount(() => window.removeEventListener('keydown', handleKeydown))
</script> </script>
<template> <template>
<div v-if="commandNotice" class="command-toast" role="status">
<span>{{ commandNotice }}</span><button aria-label="关闭通知" @click="commandNotice = ''">×</button>
</div>
<Teleport to="body"> <Teleport to="body">
<div v-if="open" class="command-backdrop" @click.self="hide"> <div v-if="open" class="command-backdrop" @click.self="hide">
<section class="command-palette" role="dialog" aria-modal="true" aria-label="命令面板"> <section class="command-palette" role="dialog" aria-modal="true" aria-label="命令面板">
<input ref="input" v-model="query" class="command-input" placeholder="输入命令…" @keydown.enter.prevent="execute(filteredCommands[0])" /> <input ref="input" v-model="query" class="command-input" placeholder="输入命令…" @keydown.enter.prevent="execute(filteredCommands[0])" />
<p v-if="commandError" class="command-error">{{ commandError }}</p>
<div class="command-list"> <div class="command-list">
<button v-for="command in filteredCommands" :key="command.id" type="button" @click="execute(command)"> <button v-for="command in filteredCommands" :key="command.id" type="button" @click="execute(command)">
<span>{{ command.label }}</span><small>{{ command.hint }}</small> <span>{{ command.label }}</span><small>{{ command.hint }}</small>
@@ -89,13 +177,20 @@ onBeforeUnmount(() => window.removeEventListener('keydown', handleKeydown))
</template> </template>
<style scoped> <style scoped>
.command-backdrop { position: fixed; inset: 0; z-index: var(--z-modal); display: flex; justify-content: center; align-items: flex-start; padding-top: 12vh; background: var(--color-background-overlay); } .command-backdrop { position: fixed; inset: 0; z-index: var(--z-modal); display: flex; justify-content: center; align-items: flex-start; padding-top: 12vh; background: var(--color-background-overlay); animation: command-backdrop-in var(--motion-fast) both; }
.command-palette { width: min(600px, calc(100vw - 32px)); overflow: hidden; border: 1px solid var(--color-border-default); border-radius: var(--radius-lg); background: var(--color-surface-elevated); box-shadow: var(--shadow-xl); } .command-palette { width: min(620px, calc(100vw - 32px)); overflow: hidden; border: 1px solid var(--color-border-default); border-radius: var(--radius-xl); background: var(--color-surface-elevated); box-shadow: var(--shadow-xl); animation: command-palette-in var(--motion-normal) both; }
.command-input { width: 100%; padding: var(--space-lg); border: 0; border-bottom: 1px solid var(--color-border-default); outline: 0; background: transparent; font-size: var(--font-size-xl); } .command-input { width: 100%; padding: var(--space-xl); border: 0; border-bottom: 1px solid var(--color-border-default); outline: 0; background: transparent; color: var(--color-text-primary); font-size: var(--font-size-xl); }
.command-list { max-height: 360px; overflow: auto; padding: var(--space-sm); } .command-list { max-height: 360px; overflow: auto; padding: var(--space-sm); }
.command-list button { display: flex; justify-content: space-between; width: 100%; padding: var(--space-md); border-radius: var(--radius-md); text-align: left; } .command-list button { display: flex; justify-content: space-between; width: 100%; padding: var(--space-md) var(--space-lg); border-radius: var(--radius-md); text-align: left; transition: color var(--motion-fast), background-color var(--motion-fast), transform var(--motion-fast); }
.command-list button:hover, .command-list button:focus { outline: 0; background: var(--color-accent-soft); color: var(--color-accent-primary); } .command-list button:hover, .command-list button:focus { outline: 0; background: var(--color-accent-soft); color: var(--color-accent-primary); }
.command-list button:hover { transform: translateX(2px); }
.command-list small, .command-list p, footer { color: var(--color-text-tertiary); } .command-list small, .command-list p, footer { color: var(--color-text-tertiary); }
.command-list p { padding: var(--space-xl); text-align: center; } .command-list p { padding: var(--space-xl); text-align: center; }
footer { display: flex; gap: var(--space-lg); padding: var(--space-sm) var(--space-lg); border-top: 1px solid var(--color-border-subtle); font-size: var(--font-size-xs); } footer { display: flex; gap: var(--space-lg); padding: var(--space-sm) var(--space-lg); border-top: 1px solid var(--color-border-subtle); font-size: var(--font-size-xs); }
.command-error { margin: var(--space-sm); padding: var(--space-sm) var(--space-md); border-radius: var(--radius-md); background: var(--color-error-soft); color: var(--color-error); font-size: var(--font-size-sm); }
.command-toast { position: fixed; top: 48px; right: var(--space-xl); z-index: calc(var(--z-modal) + 1); display: flex; align-items: center; gap: var(--space-lg); max-width: min(420px, calc(100vw - 32px)); padding: var(--space-md) var(--space-lg); border: 1px solid var(--color-border-default); border-radius: var(--radius-lg); background: var(--color-surface-elevated); box-shadow: var(--shadow-lg); animation: notice-in var(--motion-normal) both; }
.command-toast button { color: var(--color-text-tertiary); font-size: var(--font-size-xl); }
@keyframes command-backdrop-in { from { opacity: 0; } to { opacity: 1; } }
@keyframes command-palette-in { from { opacity: 0; transform: translateY(-8px) scale(.99); } to { opacity: 1; transform: translateY(0) scale(1); } }
</style> </style>
@@ -17,24 +17,33 @@ watch(() => props.source, async (source) => {
<div class="markdown-content" v-html="html" /> <div class="markdown-content" v-html="html" />
</template> </template>
<style scoped> <style>
.markdown-content { white-space: normal; user-select: text; } .markdown-content { white-space: normal; user-select: text; }
.markdown-content :deep(p), .markdown-content :deep(ul), .markdown-content :deep(ol), .markdown-content :deep(pre), .markdown-content :deep(blockquote) { margin: .65em 0; } .markdown-content p, .markdown-content ul, .markdown-content ol, .markdown-content pre, .markdown-content blockquote { margin: .65em 0; }
.markdown-content :deep(h1), .markdown-content :deep(h2), .markdown-content :deep(h3) { margin: 1em 0 .5em; line-height: var(--line-height-tight); } .markdown-content h1, .markdown-content h2, .markdown-content h3 { margin: 1em 0 .5em; line-height: var(--line-height-tight); }
.markdown-content :deep(ul) { padding-left: 1.5em; list-style: disc; } .markdown-content ul { padding-left: 1.5em; list-style: disc; }
.markdown-content :deep(ol) { padding-left: 1.5em; list-style: decimal; } .markdown-content ol { padding-left: 1.5em; list-style: decimal; }
.markdown-content :deep(.shiki) { overflow: auto; padding: var(--space-md); border: 1px solid var(--color-border-subtle); border-radius: var(--radius-md); } .markdown-content li::marker { color: var(--color-markdown-marker); font-weight: 700; }
.markdown-content :deep(code) { padding: .1em .3em; border-radius: var(--radius-sm); background: var(--color-background-tertiary); font-family: var(--font-ui-mono); } .markdown-content .shiki { overflow: auto; margin: .85em 0; padding: 16px; border: 1px solid var(--color-code-border); border-radius: 6px; background: var(--color-code-background) !important; color: var(--color-code-text); font-family: var(--font-ui-mono); font-size: .875em; line-height: 1.45; tab-size: 4; }
.markdown-content :deep(pre code) { padding: 0; background: transparent; } .markdown-content code { padding: .1em .3em; border-radius: var(--radius-sm); background: var(--color-background-tertiary); font-family: var(--font-ui-mono); }
.markdown-content :deep(blockquote) { padding-left: 1em; border-left: 3px solid var(--color-accent-primary); color: var(--color-text-secondary); } .markdown-content .shiki code { display: block; min-width: max-content; padding: 0; background: transparent; font: inherit; }
.markdown-content :deep(table) { width: 100%; margin: .65em 0; border-collapse: collapse; } .markdown-content .shiki .line { display: block; min-height: 1.45em; }
.markdown-content :deep(th), .markdown-content :deep(td) { padding: .45em .65em; border: 1px solid var(--color-border-default); text-align: left; } .markdown-content blockquote { padding-left: 1em; border-left: 3px solid var(--color-accent-primary); color: var(--color-text-secondary); }
.markdown-content :deep(img) { max-width: 100%; } .markdown-content table { width: 100%; margin: .65em 0; border-collapse: collapse; }
.markdown-content :deep(hr) { margin: 1em 0; border: 0; border-top: 1px solid var(--color-border-default); } .markdown-content th, .markdown-content td { padding: .45em .65em; border: 1px solid var(--color-markdown-grid); text-align: left; }
:global([data-theme='dark']) .markdown-content :deep(.shiki), .markdown-content th { background: var(--color-markdown-table-header); font-weight: 700; }
:global([data-theme='dark']) .markdown-content :deep(.shiki span) { .markdown-content img { max-width: 100%; }
.markdown-content hr { margin: 1em 0; border: 0; border-top: 1px solid var(--color-border-default); }
[data-code-theme='github-light'] .markdown-content .shiki,
[data-code-theme='github-light'] .markdown-content .shiki span {
color: var(--shiki-light) !important;
font-style: var(--shiki-light-font-style) !important;
font-weight: var(--shiki-light-font-weight) !important;
text-decoration: var(--shiki-light-text-decoration) !important;
}
[data-code-theme='github-dark'] .markdown-content .shiki,
[data-code-theme='github-dark'] .markdown-content .shiki span {
color: var(--shiki-dark) !important; color: var(--shiki-dark) !important;
background-color: var(--shiki-dark-bg) !important;
font-style: var(--shiki-dark-font-style) !important; font-style: var(--shiki-dark-font-style) !important;
font-weight: var(--shiki-dark-font-weight) !important; font-weight: var(--shiki-dark-font-weight) !important;
text-decoration: var(--shiki-dark-text-decoration) !important; text-decoration: var(--shiki-dark-text-decoration) !important;
@@ -1,7 +1,7 @@
<script setup lang="ts"> <script setup lang="ts">
import { useRoute, useRouter } from 'vue-router' import { useRoute, useRouter } from 'vue-router'
import { computed, ref } from 'vue' import { computed, ref } from 'vue'
import { ArrowLeftBold, ArrowRightBold, Brush, ChatDotRound, CircleCheck, Connection, Cpu, FolderOpened, Lightning, Search, Setting } from '@element-plus/icons-vue' import { ArrowLeftBold, ArrowRightBold, Brush, ChatDotRound, CircleCheck, Connection, Cpu, FolderOpened, Lightning, Monitor, Search, Setting } from '@element-plus/icons-vue'
import AppIcon from './AppIcon.vue' import AppIcon from './AppIcon.vue'
const route = useRoute() const route = useRoute()
@@ -16,6 +16,7 @@ const navItems = [
{ name: 'tasks', icon: CircleCheck, label: '任务' }, { name: 'tasks', icon: CircleCheck, label: '任务' },
{ name: 'skills', icon: Lightning, label: 'Skill' }, { name: 'skills', icon: Lightning, label: 'Skill' },
{ name: 'plugins', icon: Connection, label: 'Plugin' }, { name: 'plugins', icon: Connection, label: 'Plugin' },
{ name: 'mcp-servers', icon: Monitor, label: 'MCP' },
{ name: 'themes', icon: Brush, label: '主题' }, { name: 'themes', icon: Brush, label: '主题' },
{ name: 'settings', icon: Setting, label: '设置' }, { name: 'settings', icon: Setting, label: '设置' },
] ]
@@ -60,13 +61,14 @@ function toggleExpanded() {
<style scoped> <style scoped>
.primary-sidebar { .primary-sidebar {
width: 56px; width: var(--sidebar-primary-width);
background: var(--color-background-secondary); background: var(--color-background-secondary);
border-right: 1px solid var(--color-border-subtle); border-right: 1px solid var(--color-border-subtle);
display: flex; display: flex;
flex-direction: column; flex-direction: column;
flex-shrink: 0; flex-shrink: 0;
z-index: var(--z-sidebar); z-index: var(--z-sidebar);
transition: width var(--motion-normal), background-color var(--motion-normal);
} }
.primary-sidebar.expanded { width: var(--sidebar-primary-width-expanded); } .primary-sidebar.expanded { width: var(--sidebar-primary-width-expanded); }
@@ -76,7 +78,7 @@ function toggleExpanded() {
.nav-list { .nav-list {
flex: 1; flex: 1;
padding: var(--space-sm) 0; padding: var(--space-md) 0;
display: flex; display: flex;
flex-direction: column; flex-direction: column;
gap: 2px; gap: 2px;
@@ -87,31 +89,34 @@ function toggleExpanded() {
flex-direction: column; flex-direction: column;
align-items: center; align-items: center;
justify-content: center; justify-content: center;
height: 50px; height: 48px;
margin: 0 4px; margin: 0 6px;
border-radius: var(--radius-md); border-radius: var(--radius-md);
cursor: pointer; cursor: pointer;
color: var(--color-text-secondary); color: var(--color-text-secondary);
transition: all var(--motion-fast); border: 1px solid transparent;
transition: color var(--motion-fast), background-color var(--motion-fast), border-color var(--motion-fast), transform var(--motion-fast);
position: relative; position: relative;
&:hover { &:hover {
background: var(--color-background-hover); background: var(--color-background-hover);
color: var(--color-text-primary); color: var(--color-text-primary);
transform: translateX(2px);
} }
&.active { &.active {
background: var(--color-accent-soft); background: var(--color-accent-soft);
color: var(--color-accent-primary); color: var(--color-accent-primary);
border-color: color-mix(in srgb, var(--color-accent-primary) 16%, transparent);
&::before { &::before {
content: ''; content: '';
position: absolute; position: absolute;
left: -4px; left: -7px;
top: 50%; top: 50%;
transform: translateY(-50%); transform: translateY(-50%);
width: 3px; width: 4px;
height: 24px; height: 22px;
border-radius: 0 var(--radius-sm) var(--radius-sm) 0; border-radius: 0 var(--radius-sm) var(--radius-sm) 0;
background: var(--color-accent-primary); background: var(--color-accent-primary);
} }
@@ -127,6 +132,7 @@ function toggleExpanded() {
.nav-label { .nav-label {
font-size: 10px; font-size: 10px;
line-height: 1.2; line-height: 1.2;
font-weight: 550;
} }
.sidebar-footer { .sidebar-footer {
@@ -134,5 +140,5 @@ function toggleExpanded() {
border-top: 1px solid var(--color-border-subtle); border-top: 1px solid var(--color-border-subtle);
} }
.collapse-button { width: calc(100% - 8px); } .collapse-button { width: calc(100% - 12px); }
</style> </style>
@@ -53,7 +53,7 @@ const showSkillToggle = computed(() => routeName.value === 'skills' || routeName
<style scoped> <style scoped>
.secondary-sidebar { .secondary-sidebar {
width: var(--sidebar-secondary-width); width: var(--sidebar-secondary-width);
background: var(--color-background-primary); background: var(--color-surface-secondary);
border-right: 1px solid var(--color-border-default); border-right: 1px solid var(--color-border-default);
display: flex; display: flex;
flex-direction: column; flex-direction: column;
@@ -62,36 +62,36 @@ const showSkillToggle = computed(() => routeName.value === 'skills' || routeName
} }
.sidebar-header { .sidebar-header {
padding: var(--space-md) var(--space-lg); padding: var(--space-lg);
border-bottom: 1px solid var(--color-border-subtle); border-bottom: 1px solid var(--color-border-subtle);
flex-shrink: 0; flex-shrink: 0;
} }
.sidebar-title { .sidebar-title {
font-size: var(--font-size-sm); font-size: var(--font-size-lg);
font-weight: 600; font-weight: 700;
color: var(--color-text-primary); color: var(--color-text-primary);
margin: 0 0 var(--space-sm) 0; margin: 0 0 var(--space-md) 0;
} }
.sidebar-tabs { .sidebar-tabs {
display: flex; display: flex;
gap: 2px; gap: 2px;
background: var(--color-background-secondary); background: var(--color-background-secondary);
padding: 2px; padding: 3px;
border-radius: var(--radius-md); border-radius: var(--radius-md);
} }
.tab { .tab {
flex: 1; flex: 1;
text-align: center; text-align: center;
padding: 4px 8px; padding: 6px 8px;
font-size: var(--font-size-xs); font-size: var(--font-size-xs);
color: var(--color-text-secondary); color: var(--color-text-secondary);
border-radius: var(--radius-sm); border-radius: var(--radius-sm);
cursor: pointer; cursor: pointer;
text-decoration: none; text-decoration: none;
transition: all var(--motion-fast); transition: color var(--motion-fast), background-color var(--motion-fast), box-shadow var(--motion-fast);
&.active { &.active {
background: var(--color-surface-primary); background: var(--color-surface-primary);
@@ -108,6 +108,7 @@ const showSkillToggle = computed(() => routeName.value === 'skills' || routeName
flex: 1; flex: 1;
overflow-y: auto; overflow-y: auto;
overflow-x: hidden; overflow-x: hidden;
scrollbar-gutter: stable;
} }
</style> </style>
+9 -2
View File
@@ -107,8 +107,8 @@ const showEditorInfo = computed(() => route.name === 'workspace')
display: flex; display: flex;
align-items: center; align-items: center;
justify-content: space-between; justify-content: space-between;
padding: 0 var(--space-md); padding: 0 var(--space-lg);
background: var(--color-background-secondary); background: var(--color-surface-secondary);
border-top: 1px solid var(--color-border-subtle); border-top: 1px solid var(--color-border-subtle);
font-size: var(--font-size-xs); font-size: var(--font-size-xs);
color: var(--color-text-secondary); color: var(--color-text-secondary);
@@ -129,6 +129,7 @@ const showEditorInfo = computed(() => route.name === 'workspace')
gap: 6px; gap: 6px;
white-space: nowrap; white-space: nowrap;
cursor: default; cursor: default;
transition: color var(--motion-fast);
&:hover { &:hover {
color: var(--color-text-primary); color: var(--color-text-primary);
@@ -140,6 +141,7 @@ const showEditorInfo = computed(() => route.name === 'workspace')
height: 6px; height: 6px;
border-radius: 50%; border-radius: 50%;
flex-shrink: 0; flex-shrink: 0;
box-shadow: 0 0 0 2px var(--color-surface-secondary);
} }
.agent-status { .agent-status {
@@ -162,4 +164,9 @@ const showEditorInfo = computed(() => route.name === 'workspace')
.provider-info { .provider-info {
color: var(--color-text-tertiary); color: var(--color-text-tertiary);
} }
@media (max-width: 760px) {
.statusbar-left, .statusbar-right { gap: var(--space-sm); }
.provider-info { display: none; }
}
</style> </style>
+23 -10
View File
@@ -21,11 +21,11 @@ const pageTitle = computed(() => {
agent: '智能体执行轨迹', agent: '智能体执行轨迹',
tasks: '任务', tasks: '任务',
skills: 'Skill 管理', skills: 'Skill 管理',
plugins: 'Plugin 管理', plugins: 'Plugin 与 MCP',
themes: '主题管理', themes: '主题管理',
settings: '设置', settings: '设置',
} }
return titles[name] || '知笔知己' return titles[name] || 'NotesAgent'
}) })
const currentFileName = computed(() => { const currentFileName = computed(() => {
@@ -50,7 +50,7 @@ const isDirty = computed(() => editorStore.saveStatus === 'dirty' || editorStore
</span> </span>
</div> </div>
<div class="titlebar-center"> <div class="titlebar-center">
<span class="app-name">知笔知己</span> <span class="app-name">NotesAgent</span>
</div> </div>
<div class="titlebar-right"> <div class="titlebar-right">
<button class="icon-btn" @click="themeStore.toggleTheme()" :title="themeStore.isDark ? '切换浅色主题' : '切换深色主题'"> <button class="icon-btn" @click="themeStore.toggleTheme()" :title="themeStore.isDark ? '切换浅色主题' : '切换深色主题'">
@@ -71,8 +71,8 @@ const isDirty = computed(() => editorStore.saveStatus === 'dirty' || editorStore
display: flex; display: flex;
align-items: center; align-items: center;
justify-content: space-between; justify-content: space-between;
padding: 0 var(--space-md); padding: 0 var(--space-lg);
background: var(--color-background-secondary); background: var(--color-surface-secondary);
border-bottom: 1px solid var(--color-border-subtle); border-bottom: 1px solid var(--color-border-subtle);
font-size: var(--font-size-sm); font-size: var(--font-size-sm);
flex-shrink: 0; flex-shrink: 0;
@@ -129,7 +129,13 @@ const isDirty = computed(() => editorStore.saveStatus === 'dirty' || editorStore
} }
.app-name { .app-name {
font-weight: 500; padding: 3px 10px;
border: 1px solid var(--color-border-subtle);
border-radius: var(--radius-full);
background: var(--color-surface-primary);
color: var(--color-text-secondary);
font-weight: 650;
letter-spacing: .04em;
} }
.titlebar-right { .titlebar-right {
@@ -141,8 +147,8 @@ const isDirty = computed(() => editorStore.saveStatus === 'dirty' || editorStore
} }
.icon-btn { .icon-btn {
width: 28px; width: 30px;
height: 28px; height: 30px;
display: flex; display: flex;
align-items: center; align-items: center;
justify-content: center; justify-content: center;
@@ -150,11 +156,12 @@ const isDirty = computed(() => editorStore.saveStatus === 'dirty' || editorStore
color: var(--color-text-secondary); color: var(--color-text-secondary);
font-size: 14px; font-size: 14px;
-webkit-app-region: no-drag; -webkit-app-region: no-drag;
transition: background var(--motion-fast); transition: background-color var(--motion-fast), color var(--motion-fast), transform var(--motion-fast);
&:hover { &:hover {
background: var(--color-background-hover); background: var(--color-background-hover);
color: var(--color-text-primary); color: var(--color-text-primary);
transform: rotate(8deg);
} }
} }
@@ -179,7 +186,7 @@ const isDirty = computed(() => editorStore.saveStatus === 'dirty' || editorStore
color: var(--color-text-secondary); color: var(--color-text-secondary);
border-radius: var(--radius-sm); border-radius: var(--radius-sm);
cursor: pointer; cursor: pointer;
transition: background var(--motion-fast); transition: background-color var(--motion-fast), color var(--motion-fast);
&:hover { &:hover {
background: var(--color-background-hover); background: var(--color-background-hover);
@@ -190,4 +197,10 @@ const isDirty = computed(() => editorStore.saveStatus === 'dirty' || editorStore
color: white; color: white;
} }
} }
@media (max-width: 760px) {
.titlebar-left, .titlebar-right { min-width: 0; }
.titlebar-center, .window-controls { display: none; }
.file-name { max-width: 42vw; }
}
</style> </style>
+197 -2
View File
@@ -24,6 +24,7 @@ export interface NoteBlock {
export interface FileNode { export interface FileNode {
id: string id: string
note_id?: string
name: string name: string
path: string path: string
type: 'file' | 'folder' type: 'file' | 'folder'
@@ -142,6 +143,10 @@ export type AgentEventType =
| 'PermissionRequired' | 'PermissionRequired'
| 'Usage' | 'Usage'
| 'Citation' | 'Citation'
| 'ModelCallStarted'
| 'ModelCallCompleted'
| 'ModelCallFailed'
| 'PermissionResolved'
| 'RunCompleted' | 'RunCompleted'
| 'RunFailed' | 'RunFailed'
| 'RunCancelled' | 'RunCancelled'
@@ -154,6 +159,24 @@ export interface AgentEvent {
timestamp: string timestamp: string
} }
export interface AgentTraceSummary {
model_calls: number
tool_calls: number
duration_ms: number
token_usage: number
errors: number
}
export interface AgentTraceResponse {
run_id: string
status: AgentRunStatus
items: AgentEvent[]
next_sequence: number
has_more: boolean
summary: AgentTraceSummary
config_snapshot: Record<string, unknown>
}
export interface ToolCall { export interface ToolCall {
tool_call_id: string tool_call_id: string
name: string name: string
@@ -170,7 +193,7 @@ export interface ToolDefinition {
name: string name: string
description: string description: string
parameters: Record<string, unknown> parameters: Record<string, unknown>
source?: 'builtin' | 'plugin' source?: 'builtin' | 'plugin' | 'mcp_server'
plugin_id?: string plugin_id?: string
} }
@@ -232,6 +255,102 @@ export type PluginStatus =
| 'dependency_missing' | 'dependency_missing'
| 'permission_required' | 'permission_required'
export type PluginHostState = 'stopped' | 'starting' | 'ready' | 'unhealthy' | 'error'
export interface PluginHostStatus {
plugin_id: string
backend_type: 'mcp' | 'internal_rpc' | 'none'
transport: 'stdio' | 'http' | 'none'
status: PluginHostState
tools_count: number
started_at?: string | null
last_seen_at?: string | null
protocol_version?: string | null
server_name?: string | null
server_version?: string | null
error?: string | null
}
export type PluginCommandLocation = 'command_palette' | 'context_menu' | 'toolbar'
export interface PluginCommand {
command_id: string
plugin_id: string
title: string
description: string
icon?: string | null
locations: PluginCommandLocation[]
when: string[]
parameters: Record<string, unknown>
enabled: boolean
}
export interface PluginCommandContext {
vault_id?: string | null
note_id?: string | null
file_path?: string | null
selection?: string | null
}
export type PluginCommandEffect =
| { type: 'none'; payload: Record<string, never> }
| {
type: 'notification'
payload: { level: 'info' | 'success' | 'warning' | 'error'; message: string }
}
| {
type: 'navigate'
payload: {
route:
| 'vault-entry'
| 'workspace'
| 'search'
| 'chat'
| 'agent'
| 'tasks'
| 'skills'
| 'plugins'
| 'themes'
| 'settings'
}
}
| { type: 'refresh'; payload: { scope: 'workspace' | 'commands' | 'settings' | 'plugins' } }
| { type: 'job'; payload: { job_id: string } }
export interface PluginCommandResult {
command_id: string
status: 'completed'
effect: PluginCommandEffect
}
export type PluginSettingType = 'string' | 'number' | 'boolean' | 'select' | 'secret'
export interface PluginSettingField {
key: string
label: string
description: string
type: PluginSettingType
required: boolean
default?: unknown
minimum?: number | null
maximum?: number | null
options: string[]
}
export interface PluginSettingsSchema {
plugin_id: string
schema_version: number
fields: PluginSettingField[]
values: Record<string, unknown>
secrets: Record<string, { configured: boolean }>
}
export interface PluginSecretStatus {
plugin_id: string
key: string
configured: boolean
}
export interface PluginContribution { export interface PluginContribution {
type: 'tool' | 'command' | 'importer' | 'exporter' | 'sidebar_panel' | 'settings_section' type: 'tool' | 'command' | 'importer' | 'exporter' | 'sidebar_panel' | 'settings_section'
id: string id: string
@@ -329,6 +448,7 @@ export interface ThemeConfig {
is_dark: boolean is_dark: boolean
author?: string author?: string
builtin: boolean builtin: boolean
code_theme?: 'github-light' | 'github-dark'
} }
// ============ Index ============ // ============ Index ============
@@ -386,12 +506,80 @@ export interface PageMeta {
offset: number offset: number
} }
export interface ApiWorkspaceInfo {
vault_id: string
name: string
path: string
file_count: number
indexed_note_count: number
requires_refresh: boolean
}
export interface ApiWorkspaceEntry {
entry_id: string
name: string
path: string
type: 'file' | 'folder'
note_id?: string | null
children: ApiWorkspaceEntry[]
}
export interface ApiWorkspaceSnapshot {
workspace: ApiWorkspaceInfo
items: ApiWorkspaceEntry[]
}
export interface OperationResponse { export interface OperationResponse {
status: 'accepted' | 'completed' status: 'accepted' | 'completed'
resource_id?: string | null resource_id?: string | null
message?: string | null message?: string | null
} }
export type McpServerTransport = 'stdio' | 'streamable_http' | 'sse'
export type McpServerState = 'stopped' | 'starting' | 'ready' | 'unhealthy' | 'error'
export interface McpServerInput {
version?: number
name: string
transport: McpServerTransport
command?: string | null
args: string[]
url?: string | null
headers: Record<string, string>
environment: Record<string, string>
secret_environment_keys: string[]
secret_header_keys: string[]
permissions: string[]
startup_timeout_seconds: number
tool_timeout_seconds: number
}
export interface McpServer extends Omit<McpServerInput, 'secret_environment_keys' | 'secret_header_keys'> {
server_id: string
version: number
secret_environment: Record<string, boolean>
secret_headers: Record<string, boolean>
enabled: boolean
trusted: boolean
command_digest: string
command_summary: string
status: McpServerState
tools_count: number
protocol_version?: string | null
remote_server_name?: string | null
remote_server_version?: string | null
error?: string | null
last_tested_at?: string | null
last_test_succeeded?: boolean | null
}
export interface McpToolSummary {
name: string
remote_name: string
description: string
permission?: string | null
}
export interface ApiNoteBlock { export interface ApiNoteBlock {
block_id: string block_id: string
note_id: string note_id: string
@@ -492,7 +680,14 @@ export interface ApiPlugin {
panels: string[] panels: string[]
settings_sections: string[] settings_sections: string[]
} }
backend: { type: 'mcp' | 'internal_rpc' | 'none'; transport: 'stdio' | 'http' | 'none' } backend: {
type: 'mcp' | 'internal_rpc' | 'none'
transport: 'stdio' | 'http' | 'none'
command?: string | null
args?: string[]
startup_timeout_seconds?: number
tool_timeout_seconds?: number
}
} }
status: PluginStatus status: PluginStatus
enabled: boolean enabled: boolean
+9 -3
View File
@@ -108,16 +108,22 @@ function eventText(event: AgentEvent) {
</template> </template>
<style scoped> <style scoped>
.run-form { display: grid; gap: var(--space-xl); max-width: 980px; } .agent-page > * { width: min(100%, 1080px); margin-inline: auto; }
.run-form { display: grid; gap: var(--space-xl); }
.tool-grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(230px, 1fr)); gap: var(--space-sm); } .tool-grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(230px, 1fr)); gap: var(--space-sm); }
.tool-option { display: flex; gap: var(--space-sm); padding: var(--space-sm); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); } .tool-option { display: flex; gap: var(--space-sm); padding: var(--space-md); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); background: var(--color-surface-primary); cursor: pointer; transition: border-color var(--motion-fast), background-color var(--motion-fast), transform var(--motion-fast), box-shadow var(--motion-fast); }
.tool-option:hover { border-color: var(--color-accent-secondary); transform: translateY(-1px); box-shadow: var(--shadow-sm); }
.tool-option:has(input:checked) { border-color: var(--color-accent-primary); background: var(--color-accent-soft); box-shadow: 0 0 0 2px color-mix(in srgb, var(--color-accent-primary) 10%, transparent); }
.tool-option small { display: block; color: var(--color-text-secondary); } .tool-option small { display: block; color: var(--color-text-secondary); }
.tool-option code { display: block; margin: 2px 0; color: var(--color-text-tertiary); font-size: var(--font-size-xs); } .tool-option code { display: block; margin: 2px 0; color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
.network { display: flex; gap: var(--space-sm); } .network { display: flex; gap: var(--space-sm); }
.trace-layout { display: grid; gap: var(--space-lg); } .trace-layout { display: grid; gap: var(--space-lg); }
.run-summary, .event-head { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); } .run-summary, .event-head { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); }
.run-summary h2 { margin-top: var(--space-sm); font-family: var(--font-ui-mono); font-size: var(--font-size-lg); } .run-summary h2 { margin-top: var(--space-sm); font-family: var(--font-ui-mono); font-size: var(--font-size-lg); }
.timeline { display: grid; gap: var(--space-md); } .timeline { position: relative; display: grid; gap: var(--space-md); padding-left: var(--space-md); }
.timeline::before { content: ''; position: absolute; top: 10px; bottom: 10px; left: 1px; width: 2px; border-radius: var(--radius-full); background: var(--color-border-default); }
.event-card { position: relative; }
.event-card::before { content: ''; position: absolute; top: 20px; left: calc(-1 * var(--space-md) - 5px); width: 8px; height: 8px; border: 2px solid var(--color-surface-primary); border-radius: var(--radius-full); background: var(--color-accent-primary); box-shadow: 0 0 0 1px var(--color-accent-secondary); }
.event-head { color: var(--color-text-tertiary); font-size: var(--font-size-xs); } .event-head { color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
.event-text { margin-top: var(--space-md); white-space: pre-wrap; line-height: var(--line-height-relaxed); } .event-text { margin-top: var(--space-md); white-space: pre-wrap; line-height: var(--line-height-relaxed); }
pre { margin-top: var(--space-md); max-height: 260px; overflow: auto; padding: var(--space-md); border-radius: var(--radius-md); background: var(--color-background-secondary); font-family: var(--font-ui-mono); font-size: var(--font-size-xs); white-space: pre-wrap; user-select: text; } pre { margin-top: var(--space-md); max-height: 260px; overflow: auto; padding: var(--space-md); border-radius: var(--radius-md); background: var(--color-background-secondary); font-family: var(--font-ui-mono); font-size: var(--font-size-xs); white-space: pre-wrap; user-select: text; }
+8
View File
@@ -18,6 +18,10 @@ const eventLabels: Record<AgentEventType, string> = {
PermissionRequired: '请求权限', PermissionRequired: '请求权限',
Usage: '用量统计', Usage: '用量统计',
Citation: '引用来源', Citation: '引用来源',
ModelCallStarted: '模型调用开始',
ModelCallCompleted: '模型调用完成',
ModelCallFailed: '模型调用失败',
PermissionResolved: '权限已处理',
RunCompleted: '运行完成', RunCompleted: '运行完成',
RunFailed: '运行失败', RunFailed: '运行失败',
RunCancelled: '运行取消', RunCancelled: '运行取消',
@@ -86,6 +90,10 @@ const detailLabels: Record<string, string> = {
total_tokens: '令牌总数', total_tokens: '令牌总数',
status: '状态', status: '状态',
duration_ms: '耗时(毫秒)', duration_ms: '耗时(毫秒)',
model_call_id: '模型调用 ID',
parent_model_call_id: '上级模型调用 ID',
finish_reason: '结束原因',
decision: '授权决定',
} }
export function runStatusLabel(status?: AgentRunStatus): string { export function runStatusLabel(status?: AgentRunStatus): string {
+11 -6
View File
@@ -94,24 +94,29 @@ async function openCitation(citation: Citation) {
</template> </template>
<style scoped> <style scoped>
.chat-page { display: grid; grid-template-rows: auto auto 1fr auto; height: 100%; min-height: 0; background: var(--color-background-primary); } .chat-page { display: grid; grid-template-rows: auto auto 1fr auto; height: 100%; min-height: 0; background: radial-gradient(circle at 85% -10%, var(--color-accent-soft), transparent 30%), var(--color-background-primary); }
.chat-toolbar { display: flex; align-items: end; flex-wrap: wrap; gap: var(--space-md); padding: var(--space-md) var(--space-xl); border-bottom: 1px solid var(--color-border-default); } .chat-toolbar { display: flex; align-items: end; flex-wrap: wrap; gap: var(--space-md); padding: var(--space-md) var(--space-xl); border-bottom: 1px solid var(--color-border-default); background: var(--color-surface-secondary); box-shadow: var(--shadow-sm); z-index: 1; }
.compact { min-width: 160px; } .compact { min-width: 160px; }
.rag-toggle { display: flex; align-items: center; gap: var(--space-xs); min-height: 36px; color: var(--color-text-secondary); } .rag-toggle { display: flex; align-items: center; gap: var(--space-xs); min-height: 36px; color: var(--color-text-secondary); }
.chat-error { margin: var(--space-md) var(--space-xl) 0; } .chat-error { margin: var(--space-md) var(--space-xl) 0; }
.message-timeline { min-height: 0; overflow: auto; padding: var(--space-xl) max(var(--space-xl), calc((100% - 820px) / 2)); user-select: text; } .message-timeline { min-height: 0; overflow: auto; padding: var(--space-xl) max(var(--space-xl), calc((100% - 820px) / 2)); user-select: text; }
.message { display: grid; grid-template-columns: 36px 1fr; gap: var(--space-md); margin-bottom: var(--space-xl); } .message { display: grid; grid-template-columns: 36px 1fr; gap: var(--space-md); margin-bottom: var(--space-xl); animation: message-in var(--motion-normal) both; }
.avatar { display: grid; place-items: center; width: 34px; height: 34px; border-radius: var(--radius-full); background: var(--color-background-tertiary); font-weight: 700; } .avatar { display: grid; place-items: center; width: 34px; height: 34px; border: 1px solid var(--color-border-default); border-radius: var(--radius-full); background: var(--color-background-tertiary); box-shadow: var(--shadow-sm); font-weight: 700; }
.assistant .avatar { background: var(--color-accent-soft); color: var(--color-accent-primary); } .assistant .avatar { background: var(--color-accent-soft); color: var(--color-accent-primary); }
.message-body { min-width: 0; padding: var(--space-md) var(--space-lg); border: 1px solid var(--color-border-subtle); border-radius: 4px var(--radius-lg) var(--radius-lg) var(--radius-lg); background: color-mix(in srgb, var(--color-surface-primary) 88%, transparent); box-shadow: var(--shadow-sm); }
.user .message-body { background: var(--color-accent-soft); border-color: color-mix(in srgb, var(--color-accent-primary) 14%, transparent); }
.message-content { white-space: pre-wrap; line-height: var(--line-height-relaxed); } .message-content { white-space: pre-wrap; line-height: var(--line-height-relaxed); }
.thinking { margin-bottom: var(--space-sm); color: var(--color-text-secondary); }.thinking p { margin-top: var(--space-sm); white-space: pre-wrap; } .thinking { margin-bottom: var(--space-sm); color: var(--color-text-secondary); }.thinking p { margin-top: var(--space-sm); white-space: pre-wrap; }
.tool-calls { display: grid; gap: var(--space-sm); margin-top: var(--space-md); }.tool-calls .item-card { display: grid; gap: var(--space-xs); }.tool-calls pre { overflow: auto; font-size: var(--font-size-xs); } .tool-calls { display: grid; gap: var(--space-sm); margin-top: var(--space-md); }.tool-calls .item-card { display: grid; gap: var(--space-xs); }.tool-calls pre { overflow: auto; font-size: var(--font-size-xs); }
.usage { display: block; margin-top: var(--space-xs); color: var(--color-text-tertiary); } .usage { display: block; margin-top: var(--space-xs); color: var(--color-text-tertiary); }
.message time { display: block; margin-top: var(--space-sm); color: var(--color-text-tertiary); font-size: var(--font-size-xs); } .message time { display: block; margin-top: var(--space-sm); color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
.citations { display: grid; gap: var(--space-sm); margin-top: var(--space-md); } .citations { display: grid; gap: var(--space-sm); margin-top: var(--space-md); }
.citation-card { display: flex; align-items: flex-start; gap: var(--space-sm); padding: var(--space-sm); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); text-align: left; } .citation-card { display: flex; align-items: flex-start; gap: var(--space-sm); padding: var(--space-md); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); background: var(--color-surface-primary); text-align: left; transition: border-color var(--motion-fast), transform var(--motion-fast), box-shadow var(--motion-fast); }
.citation-card:hover { border-color: var(--color-accent-secondary); transform: translateY(-1px); box-shadow: var(--shadow-sm); }
.citation-card small { display: block; margin-top: 2px; color: var(--color-text-secondary); } .citation-card small { display: block; margin-top: 2px; color: var(--color-text-secondary); }
.composer { padding: var(--space-md) max(var(--space-xl), calc((100% - 820px) / 2)); border-top: 1px solid var(--color-border-default); background: var(--color-surface-primary); } .composer { padding: var(--space-md) max(var(--space-xl), calc((100% - 820px) / 2)); border-top: 1px solid var(--color-border-default); background: var(--color-surface-secondary); box-shadow: 0 -8px 24px color-mix(in srgb, var(--color-text-primary) 5%, transparent); }
.composer .textarea { min-height: 72px; } .composer .textarea { min-height: 72px; }
.composer-actions { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); margin-top: var(--space-sm); } .composer-actions { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); margin-top: var(--space-sm); }
@keyframes message-in { from { opacity: 0; transform: translateY(5px); } to { opacity: 1; transform: translateY(0); } }
</style> </style>
@@ -1,10 +1,11 @@
// @vitest-environment happy-dom // @vitest-environment happy-dom
import { afterEach, beforeEach, describe, expect, it } from 'vitest' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import { mount, type VueWrapper } from '@vue/test-utils' import { mount, type VueWrapper } from '@vue/test-utils'
import { createPinia, setActivePinia } from 'pinia' import { createPinia, setActivePinia } from 'pinia'
import { nextTick } from 'vue' import { nextTick } from 'vue'
import EditorPane from './EditorPane.vue' import EditorPane from './EditorPane.vue'
import { useEditorStore } from '@/stores/editor' import { useEditorStore } from '@/stores/editor'
import * as workspaceService from '@/services/workspaceService'
let wrapper: VueWrapper | null = null let wrapper: VueWrapper | null = null
@@ -19,20 +20,31 @@ async function waitForText(text: string) {
beforeEach(() => { beforeEach(() => {
localStorage.clear() localStorage.clear()
setActivePinia(createPinia()) setActivePinia(createPinia())
vi.spyOn(workspaceService, 'readFileContent').mockImplementation(async (filePath) => {
if (filePath === '/欢迎使用 NotesAgent.md') {
return '# 欢迎使用 NotesAgent\n\n祝你写作愉快'
}
if (filePath === '/数据结构/红黑树.md') return '# 红黑树\n\n新的文件内容'
throw new Error(`Unexpected file path: ${filePath}`)
})
vi.spyOn(workspaceService, 'getNoteId').mockImplementation(async (filePath) =>
filePath.includes('红黑树') ? 'note-rbt' : 'note-welcome'
)
}) })
afterEach(() => { afterEach(() => {
wrapper?.unmount() wrapper?.unmount()
wrapper = null wrapper = null
document.body.innerHTML = '' document.body.innerHTML = ''
vi.restoreAllMocks()
}) })
describe('EditorPane file switching', () => { describe('EditorPane file switching', () => {
it('recreates the visual editor with the newly loaded file content', async () => { it('recreates the visual editor with the newly loaded file content', async () => {
const store = useEditorStore() const store = useEditorStore()
await store.loadFile('/欢迎使用知笔知己.md') await store.loadFile('/欢迎使用 NotesAgent.md')
wrapper = mount(EditorPane, { attachTo: document.body }) wrapper = mount(EditorPane, { attachTo: document.body })
await waitForText('欢迎使用知笔知己') await waitForText('欢迎使用 NotesAgent')
await store.loadFile('/数据结构/红黑树.md') await store.loadFile('/数据结构/红黑树.md')
await nextTick() await nextTick()
+3 -1
View File
@@ -1,10 +1,12 @@
<script setup lang="ts"> <script setup lang="ts">
import { useEditorStore } from '@/stores/editor' import { useEditorStore } from '@/stores/editor'
import { useSettingsStore } from '@/stores/settings' import { useSettingsStore } from '@/stores/settings'
import { useThemeStore } from '@/stores/theme'
import VisualMarkdownEditor from './VisualMarkdownEditor.vue' import VisualMarkdownEditor from './VisualMarkdownEditor.vue'
const editorStore = useEditorStore() const editorStore = useEditorStore()
const settingsStore = useSettingsStore() const settingsStore = useSettingsStore()
const themeStore = useThemeStore()
function updateContent(event: Event) { function updateContent(event: Event) {
editorStore.updateContent((event.target as HTMLTextAreaElement).value) editorStore.updateContent((event.target as HTMLTextAreaElement).value)
editorStore.scheduleAutoSave(settingsStore.autoSaveInterval) editorStore.scheduleAutoSave(settingsStore.autoSaveInterval)
@@ -12,7 +14,7 @@ function updateContent(event: Event) {
</script> </script>
<template> <template>
<VisualMarkdownEditor v-if="editorStore.mode === 'wysiwyg'" :key="editorStore.currentFilePath ?? 'empty'" <VisualMarkdownEditor v-if="editorStore.mode === 'wysiwyg'" :key="`${editorStore.currentFilePath ?? 'empty'}:${themeStore.resolvedCodeBlockTheme}`"
:initial-content="editorStore.content" /> :initial-content="editorStore.content" />
<textarea v-else class="editor-pane source" :value="editorStore.content" :spellcheck="false" <textarea v-else class="editor-pane source" :value="editorStore.content" :spellcheck="false"
aria-label="Markdown 源码编辑器" @input="updateContent" /> aria-label="Markdown 源码编辑器" @input="updateContent" />
@@ -2,6 +2,7 @@
import { onBeforeUnmount, onMounted, ref } from 'vue' import { onBeforeUnmount, onMounted, ref } from 'vue'
import { Link } from '@element-plus/icons-vue' import { Link } from '@element-plus/icons-vue'
import { Crepe } from '@milkdown/crepe' import { Crepe } from '@milkdown/crepe'
import { oneDark } from '@codemirror/theme-one-dark'
import { import {
createCodeBlockCommand, createCodeBlockCommand,
toggleEmphasisCommand, toggleEmphasisCommand,
@@ -19,6 +20,7 @@ import { callCommand } from '@milkdown/kit/utils'
import AppIcon from '@/components/common/AppIcon.vue' import AppIcon from '@/components/common/AppIcon.vue'
import { useEditorStore } from '@/stores/editor' import { useEditorStore } from '@/stores/editor'
import { useSettingsStore } from '@/stores/settings' import { useSettingsStore } from '@/stores/settings'
import { useThemeStore } from '@/stores/theme'
import { applyMarkdownFontSize, fontSizeMarkdownPlugin } from './fontSizeMarkdown' import { applyMarkdownFontSize, fontSizeMarkdownPlugin } from './fontSizeMarkdown'
import '@milkdown/crepe/theme/common/style.css' import '@milkdown/crepe/theme/common/style.css'
import '@milkdown/crepe/theme/frame.css' import '@milkdown/crepe/theme/frame.css'
@@ -26,6 +28,7 @@ import '@milkdown/crepe/theme/frame.css'
const props = defineProps<{ initialContent: string }>() const props = defineProps<{ initialContent: string }>()
const editorStore = useEditorStore() const editorStore = useEditorStore()
const settingsStore = useSettingsStore() const settingsStore = useSettingsStore()
const themeStore = useThemeStore()
const editorRoot = ref<HTMLElement | null>(null) const editorRoot = ref<HTMLElement | null>(null)
const loading = ref(true) const loading = ref(true)
const fontSizeInput = ref(16) const fontSizeInput = ref(16)
@@ -36,6 +39,7 @@ type ToolbarCommand = 'bold' | 'italic' | 'ordered-list' | 'bullet-list' | 'inli
function runCommand(command: ToolbarCommand) { function runCommand(command: ToolbarCommand) {
const editor = crepe?.editor const editor = crepe?.editor
if (!editor) return if (!editor) return
// Milkdown
const actions = { const actions = {
bold: callCommand(toggleStrongCommand.key), bold: callCommand(toggleStrongCommand.key),
italic: callCommand(toggleEmphasisCommand.key), italic: callCommand(toggleEmphasisCommand.key),
@@ -52,6 +56,7 @@ function runCommand(command: ToolbarCommand) {
function applyLink() { function applyLink() {
if (!crepe) return if (!crepe) return
// TODO(editor): Element Plus prompt URL
const href = window.prompt('请输入链接地址', 'https://')?.trim() const href = window.prompt('请输入链接地址', 'https://')?.trim()
if (!href) return if (!href) return
@@ -104,6 +109,7 @@ onMounted(async () => {
featureConfigs: { featureConfigs: {
[Crepe.Feature.Placeholder]: { text: '开始记录你的想法…' }, [Crepe.Feature.Placeholder]: { text: '开始记录你的想法…' },
[Crepe.Feature.CodeMirror]: { [Crepe.Feature.CodeMirror]: {
theme: themeStore.resolvedCodeBlockTheme === 'github-dark' ? oneDark : [],
previewOnlyByDefault: false, previewOnlyByDefault: false,
searchPlaceholder: '搜索语言', searchPlaceholder: '搜索语言',
noResultText: '没有匹配的语言', noResultText: '没有匹配的语言',
@@ -158,6 +164,7 @@ onMounted(async () => {
crepe.editor.use(fontSizeMarkdownPlugin) crepe.editor.use(fontSizeMarkdownPlugin)
crepe.on((listener) => { crepe.on((listener) => {
listener.markdownUpdated((_ctx, markdown, previousMarkdown) => { listener.markdownUpdated((_ctx, markdown, previousMarkdown) => {
// /
if (markdown === previousMarkdown || markdown === editorStore.content) return if (markdown === previousMarkdown || markdown === editorStore.content) return
editorStore.updateContent(markdown) editorStore.updateContent(markdown)
editorStore.scheduleAutoSave(settingsStore.autoSaveInterval) editorStore.scheduleAutoSave(settingsStore.autoSaveInterval)
@@ -250,7 +257,7 @@ defineExpose({ getEditor: () => crepe?.editor })
--crepe-color-surface-low: var(--color-background-secondary); --crepe-color-surface-low: var(--color-background-secondary);
--crepe-color-on-surface: var(--color-text-primary); --crepe-color-on-surface: var(--color-text-primary);
--crepe-color-on-surface-variant: var(--color-text-secondary); --crepe-color-on-surface-variant: var(--color-text-secondary);
--crepe-color-outline: var(--color-border-default); --crepe-color-outline: var(--color-markdown-grid);
--crepe-color-primary: var(--color-accent-primary); --crepe-color-primary: var(--color-accent-primary);
--crepe-color-secondary: var(--color-accent-soft); --crepe-color-secondary: var(--color-accent-soft);
--crepe-color-on-secondary: var(--color-text-primary); --crepe-color-on-secondary: var(--color-text-primary);
@@ -269,14 +276,21 @@ defineExpose({ getEditor: () => crepe?.editor })
.milkdown-host :deep(.ProseMirror p) { font-weight: 400; } .milkdown-host :deep(.ProseMirror p) { font-weight: 400; }
.milkdown-host :deep(.ProseMirror h1), .milkdown-host :deep(.ProseMirror h2), .milkdown-host :deep(.ProseMirror h3), .milkdown-host :deep(.ProseMirror h4), .milkdown-host :deep(.ProseMirror h5), .milkdown-host :deep(.ProseMirror h6) { font-weight: 700; } .milkdown-host :deep(.ProseMirror h1), .milkdown-host :deep(.ProseMirror h2), .milkdown-host :deep(.ProseMirror h3), .milkdown-host :deep(.ProseMirror h4), .milkdown-host :deep(.ProseMirror h5), .milkdown-host :deep(.ProseMirror h6) { font-weight: 700; }
.milkdown-host :deep(.font-size-marker) { display: none; } .milkdown-host :deep(.font-size-marker) { display: none; }
.milkdown-host :deep(.milkdown-code-block) { overflow: hidden; border: 1px solid var(--color-code-border); border-radius: 6px; background: var(--color-code-background); color: var(--color-code-text); }
.milkdown-host :deep(.milkdown-code-block .cm-editor),
.milkdown-host :deep(.milkdown-code-block .cm-gutters),
.milkdown-host :deep(.milkdown-code-block .cm-panel) { background: var(--color-code-background); }
.milkdown-host :deep(.milkdown-code-block .cm-content) { caret-color: var(--color-code-text); font-family: var(--font-editor-mono); }
.milkdown-host :deep(.milkdown-code-block .language-button) { color: var(--color-code-muted); }
:global(.milkdown-toolbar) { border: 1px solid var(--color-border-default) !important; background: var(--color-surface-elevated) !important; box-shadow: var(--shadow-md) !important; } :global(.milkdown-toolbar) { border: 1px solid var(--color-border-default) !important; background: var(--color-surface-elevated) !important; box-shadow: var(--shadow-md) !important; }
:global(.milkdown-toolbar .toolbar-item svg), :global(.milkdown-toolbar .toolbar-item.active svg) { color: var(--color-text-primary) !important; fill: var(--color-text-primary) !important; opacity: 1 !important; } :global(.milkdown-toolbar .toolbar-item svg), :global(.milkdown-toolbar .toolbar-item.active svg) { color: var(--color-text-primary) !important; fill: var(--color-text-primary) !important; opacity: 1 !important; }
:global(.milkdown-toolbar .toolbar-item:hover svg), :global(.milkdown-toolbar .toolbar-item.active svg) { color: var(--color-accent-primary) !important; fill: var(--color-accent-primary) !important; } :global(.milkdown-toolbar .toolbar-item:hover svg), :global(.milkdown-toolbar .toolbar-item.active svg) { color: var(--color-accent-primary) !important; fill: var(--color-accent-primary) !important; }
:global([data-theme='light']) .milkdown-host :deep(.milkdown-table-block th), .milkdown-host :deep(.milkdown-table-block th),
:global([data-theme='light']) .milkdown-host :deep(.milkdown-table-block td) { border-color: var(--color-text-tertiary); } .milkdown-host :deep(.milkdown-table-block td) { border-color: var(--color-markdown-grid); }
:global([data-theme='light']) .milkdown-host :deep(.milkdown-list-item-block li .label-wrapper) { color: var(--color-text-secondary); font-weight: 600; } .milkdown-host :deep(.milkdown-table-block th) { background: var(--color-markdown-table-header); font-weight: 700; }
:global([data-theme='light']) .milkdown-host :deep(.milkdown-list-item-block li .label-wrapper svg) { fill: var(--color-text-secondary); } .milkdown-host :deep(.milkdown-list-item-block li .label-wrapper) { color: var(--color-markdown-marker); font-weight: 700; }
.milkdown-host :deep(.milkdown-list-item-block li .label-wrapper svg) { fill: var(--color-markdown-marker); }
.milkdown-host :deep(code) { font-family: var(--font-editor-mono); } .milkdown-host :deep(code) { font-family: var(--font-editor-mono); }
:global([data-theme='dark']) .milkdown-host :deep(.milkdown) { color-scheme: dark; } :global([data-theme='dark'] .milkdown-host .milkdown) { color-scheme: dark; }
@media (max-width: 680px) { .toolbar-select select { min-width: 46px; width: 46px; } } @media (max-width: 680px) { .toolbar-select select { min-width: 46px; width: 46px; } }
</style> </style>
@@ -0,0 +1,169 @@
// @vitest-environment happy-dom
import { flushPromises, mount } from '@vue/test-utils'
import { beforeEach, describe, expect, it, vi } from 'vitest'
import type { McpServer } from '@/contracts'
import * as service from '@/services/mcpServerService'
import McpServersView from './McpServersView.vue'
vi.mock('@/services/mcpServerService', () => ({
listMcpServers: vi.fn(), createMcpServer: vi.fn(), updateMcpServer: vi.fn(),
deleteMcpServer: vi.fn(), trustMcpServer: vi.fn(), testMcpServer: vi.fn(),
enableMcpServer: vi.fn(), disableMcpServer: vi.fn(), putMcpServerSecret: vi.fn(),
}))
const server: McpServer = {
server_id: 'server-1', version: 2, name: 'Remote', transport: 'streamable_http',
command: null, args: [], url: 'https://mcp.example.test/mcp', headers: {}, environment: {},
secret_environment: {}, secret_headers: { Authorization: false }, permissions: [],
startup_timeout_seconds: 15, tool_timeout_seconds: 30, enabled: false, trusted: true,
command_digest: 'a'.repeat(64), command_summary: 'https://mcp.example.test/mcp',
status: 'stopped', tools_count: 1, last_test_succeeded: false,
}
async function render(items: McpServer[] = []) {
vi.mocked(service.listMcpServers).mockResolvedValue(items)
const wrapper = mount(McpServersView, { global: { stubs: { AppIcon: true } } })
await flushPromises()
return wrapper
}
beforeEach(() => {
vi.clearAllMocks()
vi.stubGlobal('confirm', vi.fn(() => true))
})
describe('McpServersView', () => {
it('switches transport templates and round-trips the JSON configuration mode', async () => {
const wrapper = await render()
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === 'Streamable HTTP')!.trigger('click')
expect(wrapper.find('input[placeholder="https://example.com/mcp"]').exists()).toBe(true)
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
const raw = (wrapper.get('.json-editor').element as HTMLTextAreaElement).value
expect(JSON.parse(raw)).toMatchObject({ transport: 'streamable_http', command: null })
expect(raw).not.toContain('secret_value')
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
expect(wrapper.text()).toContain('MCP URL')
})
it('rejects invalid JSON without sending a create request', async () => {
const wrapper = await render()
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
await wrapper.get('.json-editor').setValue('{invalid')
await flushPromises()
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(wrapper.text()).toContain('服务器配置不是有效 JSON')
expect(service.createMcpServer).not.toHaveBeenCalled()
})
it('keeps secrets request-only, exposes test failures, and confirms deletion', async () => {
const wrapper = await render([server])
const password = wrapper.get('input[type="password"]')
await password.setValue('request-only-secret')
vi.mocked(service.putMcpServerSecret).mockResolvedValue({} as never)
await wrapper.findAll('button').find(button => button.text() === '保存')!.trigger('click')
await flushPromises()
expect(service.putMcpServerSecret).toHaveBeenCalledWith('server-1', 'Authorization', 'request-only-secret', 'header')
expect((password.element as HTMLInputElement).value).toBe('')
vi.mocked(service.testMcpServer).mockRejectedValue(new Error('连接失败'))
await wrapper.findAll('button').find(button => button.text().includes('测试连接'))!.trigger('click')
await flushPromises()
expect(wrapper.text()).toContain('连接失败')
vi.mocked(service.deleteMcpServer).mockResolvedValue({ status: 'completed' })
await wrapper.findAll('button').find(button => button.text().includes('删除'))!.trigger('click')
await flushPromises()
expect(confirm).toHaveBeenCalled()
expect(service.deleteMcpServer).toHaveBeenCalledWith('server-1')
})
it('confirms permission changes before updating an existing server', async () => {
const wrapper = await render([server])
vi.mocked(service.updateMcpServer).mockResolvedValue(server)
await wrapper.findAll('button').find(button => button.text().includes('编辑'))!.trigger('click')
await wrapper.get('input[placeholder="network.request, notes.read"]').setValue('notes.read')
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(confirm).toHaveBeenCalledWith(expect.stringContaining('旧测试与授权会失效'))
expect(service.updateMcpServer).toHaveBeenCalled()
})
it('saves an environment API key via the encrypted endpoint, not the config body', async () => {
const wrapper = await render()
vi.mocked(service.createMcpServer).mockResolvedValue({ ...server, server_id: 'new-server' })
vi.mocked(service.putMcpServerSecret).mockResolvedValue({})
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
await wrapper.get('.json-editor').setValue(JSON.stringify({ command: 'uvx', environment: { MINIMAX_API_KEY: 'synthetic-only' }, secret_environment_keys: ['MINIMAX_API_KEY'] }))
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(service.createMcpServer).toHaveBeenCalledWith(expect.objectContaining({ environment: {}, secret_environment_keys: ['MINIMAX_API_KEY'] }))
expect(JSON.stringify(vi.mocked(service.createMcpServer).mock.calls)).not.toContain('synthetic-only')
expect(service.putMcpServerSecret).toHaveBeenCalledWith('new-server', 'MINIMAX_API_KEY', 'synthetic-only', 'environment')
expect(wrapper.find('.modal-backdrop').exists()).toBe(false)
})
it('retains imported keys over mode switches and retries partial saves without duplicates', async () => {
const wrapper = await render()
vi.mocked(service.createMcpServer).mockResolvedValue({ ...server, server_id: 'new-server', version: 1 })
vi.mocked(service.updateMcpServer).mockResolvedValue({ ...server, server_id: 'new-server', version: 2 })
vi.mocked(service.putMcpServerSecret).mockRejectedValueOnce(new Error('credential store unavailable')).mockResolvedValue({})
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
await wrapper.get('.json-editor').setValue(JSON.stringify({ command: 'uvx', env: { API_KEY: 'retry-value' } }))
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
expect(wrapper.text()).toContain('已识别 1 项密钥')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
expect((wrapper.get('.json-editor').element as HTMLTextAreaElement).value).not.toContain('retry-value')
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(wrapper.get('.modal-card [role="alert"]').text()).toContain('服务器配置已保存,但密钥保存失败')
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(service.createMcpServer).toHaveBeenCalledTimes(1)
expect(service.updateMcpServer).toHaveBeenCalledWith('new-server', expect.objectContaining({ version: 1 }))
expect(service.putMcpServerSecret).toHaveBeenCalledTimes(2)
expect(wrapper.find('.modal-backdrop').exists()).toBe(false)
})
it('clears staged keys on cancel and accepts minimal JSON while editing', async () => {
const wrapper = await render([server])
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
await wrapper.get('.json-editor').setValue('{"command":"uvx","env":{"API_KEY":"cancelled-value"}}')
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === '取消')!.trigger('click')
await wrapper.findAll('button').find(button => button.text().includes('编辑'))!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
await wrapper.get('.json-editor').setValue('{"name":"Minimal","url":"https://example.test/mcp"}')
vi.mocked(service.updateMcpServer).mockResolvedValue(server)
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(service.updateMcpServer).toHaveBeenCalledWith('server-1', expect.objectContaining({ version: 2, headers: {}, args: [] }))
expect(service.putMcpServerSecret).not.toHaveBeenCalled()
})
it('saves an imported Header secret after a case-only declaration rename', async () => {
const wrapper = await render()
vi.mocked(service.createMcpServer).mockResolvedValue({ ...server, secret_headers: { authorization: false } })
vi.mocked(service.putMcpServerSecret).mockResolvedValue({})
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
await wrapper.get('.json-editor').setValue(JSON.stringify({ url: 'https://example.test/mcp', headers: { Authorization: 'synthetic-draft' } }))
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
await wrapper.get('textarea[placeholder="Authorization"]').setValue('authorization')
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
expect(wrapper.text()).toContain('已识别 1 项密钥')
expect((wrapper.get('.json-editor').element as HTMLTextAreaElement).value).not.toContain('synthetic-draft')
await wrapper.get('form').trigger('submit')
await flushPromises()
expect(service.createMcpServer).toHaveBeenCalledWith(expect.objectContaining({ headers: {}, secret_header_keys: ['authorization'] }))
expect(service.putMcpServerSecret).toHaveBeenCalledWith('server-1', 'authorization', 'synthetic-draft', 'header')
expect(wrapper.find('.modal-backdrop').exists()).toBe(false)
})
})
@@ -0,0 +1,275 @@
<script setup lang="ts">
import { computed, onMounted, reactive, ref } from 'vue'
import { Connection, Delete, EditPen, Plus, Refresh, VideoPlay } from '@element-plus/icons-vue'
import AppIcon from '@/components/common/AppIcon.vue'
import type { McpServer, McpServerInput, McpServerTransport } from '@/contracts'
import * as service from '@/services/mcpServerService'
import { emptyMcpConfig, mergeImportedSecrets, normalizeMcpConfig, parseMcpJson, type ImportedSecret, type SecretKind } from './configuration'
const servers = ref<McpServer[]>([])
const busy = ref('')
const error = ref('')
const dialogOpen = ref(false)
const editingId = ref<string | null>(null)
const editingOriginal = ref<McpServer | null>(null)
const editorMode = ref<'form' | 'json'>('form')
const argsText = ref('')
const environmentText = ref('{}')
const headersText = ref('{}')
const secretKeysText = ref('')
const secretHeaderKeysText = ref('')
const permissionsText = ref('')
const rawConfig = ref('')
const secretDrafts = reactive<Record<string, string>>({})
const form = reactive<McpServerInput>(emptyMcpConfig())
const importedSecrets = ref<ImportedSecret[]>([])
const dialogTitle = computed(() => editingId.value ? '编辑 MCP 服务器' : '新增 MCP 服务器')
async function load() {
error.value = ''
try { servers.value = await service.listMcpServers() }
catch (cause) { error.value = message(cause, '读取 MCP 服务器失败') }
}
function resetEditor(input: McpServerInput) {
Object.assign(form, emptyMcpConfig(), { version: undefined }, input)
argsText.value = input.args.join('\n')
environmentText.value = JSON.stringify(input.environment, null, 2)
headersText.value = JSON.stringify(input.headers, null, 2)
secretKeysText.value = input.secret_environment_keys.join('\n')
secretHeaderKeysText.value = input.secret_header_keys.join('\n')
permissionsText.value = input.permissions.join(', ')
editorMode.value = 'form'
rawConfig.value = ''
}
function openCreate() {
if (busy.value) return
error.value = ''
importedSecrets.value = []
editingId.value = null
editingOriginal.value = null
resetEditor(emptyMcpConfig())
dialogOpen.value = true
}
function openEdit(server: McpServer) {
if (busy.value) return
error.value = ''
importedSecrets.value = []
editingId.value = server.server_id
editingOriginal.value = server
resetEditor({
version: server.version, name: server.name, transport: server.transport,
command: server.command, args: [...server.args], url: server.url,
headers: { ...server.headers }, environment: { ...server.environment },
secret_environment_keys: Object.keys(server.secret_environment),
secret_header_keys: Object.keys(server.secret_headers), permissions: [...server.permissions],
startup_timeout_seconds: server.startup_timeout_seconds,
tool_timeout_seconds: server.tool_timeout_seconds,
})
dialogOpen.value = true
}
function applyTemplate(transport: McpServerTransport) {
form.transport = transport
if (transport === 'stdio') {
form.command = 'uvx'; form.url = null
argsText.value = '--isolated\n--from\npackage-name==1.0.0\nserver-command'
} else {
form.command = null; argsText.value = ''; form.url = transport === 'sse' ? 'http://127.0.0.1:3000/sse' : 'http://127.0.0.1:3000/mcp'
}
}
function parseObject(value: string, label: string): Record<string, string> {
let parsed: unknown
try { parsed = JSON.parse(value || '{}') } catch { throw new Error(`${label}必须是 JSON 对象`) }
if (!parsed || Array.isArray(parsed) || typeof parsed !== 'object' || Object.values(parsed).some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串键值 JSON 对象`)
return parsed as Record<string, string>
}
function formPayload(): McpServerInput {
const stdio = form.transport === 'stdio'
return {
version: form.version,
name: form.name.trim(), transport: form.transport,
command: stdio ? form.command?.trim() : null,
args: stdio ? argsText.value.split('\n').map(value => value.trim()).filter(Boolean) : [],
url: stdio ? null : form.url?.trim(),
headers: stdio ? {} : parseObject(headersText.value, '普通 Header'),
environment: stdio ? parseObject(environmentText.value, '普通环境变量') : {},
secret_environment_keys: stdio ? splitKeys(secretKeysText.value) : [],
secret_header_keys: stdio ? [] : splitKeys(secretHeaderKeysText.value),
permissions: permissionsText.value.split(',').map(value => value.trim()).filter(Boolean),
startup_timeout_seconds: form.startup_timeout_seconds,
tool_timeout_seconds: form.tool_timeout_seconds,
}
}
function payload(requireConnection = true): McpServerInput {
const { config, secrets } = editorMode.value === 'form'
? normalizeMcpConfig(formPayload(), '', requireConnection) : parseMcpJson(rawConfig.value, form.name, requireConnection)
// Keep only still-declared drafts. A mode switch must not discard imported keys,
// and editing the declaration must not later send a removed key to the Secret API.
importedSecrets.value = mergeImportedSecrets(config, importedSecrets.value, secrets)
if (editingId.value) config.version = form.version
if (editorMode.value === 'json') rawConfig.value = JSON.stringify(config, null, 2)
else {
environmentText.value = JSON.stringify(config.environment, null, 2)
headersText.value = JSON.stringify(config.headers, null, 2)
secretKeysText.value = config.secret_environment_keys.join('\n')
secretHeaderKeysText.value = config.secret_header_keys.join('\n')
}
return config
}
function switchMode(mode: 'form' | 'json') {
try {
if (mode === editorMode.value) return
error.value = ''
if (mode === 'json') rawConfig.value = JSON.stringify(payload(false), null, 2)
else resetEditor(payload(false))
editorMode.value = mode
} catch (cause) { error.value = message(cause, '配置转换失败') }
}
async function save() {
if (busy.value) return
let saved: McpServer | undefined
try {
error.value = ''
const input = payload()
if (!input.name || (input.transport === 'stdio' ? !input.command : !input.url)) throw new Error('请填写服务器名称和连接地址')
if (editingOriginal.value && executionChanged(editingOriginal.value, input) && !confirm('连接命令、地址或认证配置已变化,保存后旧测试与授权会失效。是否保存?')) return
busy.value = 'save'
saved = editingId.value ? await service.updateMcpServer(editingId.value, input) : await service.createMcpServer(input)
// Commit the returned ID/version before saving secrets so a partial failure can
// retry this server instead of creating a duplicate or sending a stale version.
editingId.value = saved.server_id
editingOriginal.value = saved
resetEditor({ ...input, version: saved.version })
for (const item of [...importedSecrets.value]) {
await service.putMcpServerSecret(saved.server_id, item.key, item.value, item.kind)
importedSecrets.value = importedSecrets.value.filter(candidate => candidate !== item)
}
closeEditor()
await load()
} catch (cause) {
if (saved) await load()
error.value = `${saved ? '服务器配置已保存,但密钥保存失败;可点击保存重试。' : ''}${message(cause, '保存失败')}`
}
finally { busy.value = '' }
}
function closeEditor() {
importedSecrets.value = []
rawConfig.value = ''
environmentText.value = '{}'
headersText.value = '{}'
dialogOpen.value = false
}
function executionChanged(server: McpServer, input: McpServerInput) {
const sortedEntries = (value: Record<string, string>) => Object.entries(value).sort(([left], [right]) => left.localeCompare(right))
const current = [
server.transport, server.command, server.args, server.url,
sortedEntries(server.headers), sortedEntries(server.environment),
Object.keys(server.secret_headers).sort(), Object.keys(server.secret_environment).sort(),
[...server.permissions].sort(), server.startup_timeout_seconds, server.tool_timeout_seconds,
]
const next = [
input.transport, input.command, input.args, input.url,
sortedEntries(input.headers), sortedEntries(input.environment),
[...input.secret_header_keys].sort(), [...input.secret_environment_keys].sort(),
[...input.permissions].sort(), input.startup_timeout_seconds, input.tool_timeout_seconds,
]
return JSON.stringify(current) !== JSON.stringify(next)
}
async function approve(server: McpServer): Promise<McpServer | null> {
if (server.trusted) return server
const localWarning = server.transport === 'stdio' ? '\n\n本机进程尚无系统级沙箱,仅应运行可信服务器。' : '\n\n连接可能向该地址发送配置的 Header。'
if (!confirm(`请确认 MCP 连接:\n\n${server.command_summary}${localWarning}\n\n是否继续?`)) return null
return service.trustMcpServer(server)
}
async function test(server: McpServer) { await act(server, 'test', current => service.testMcpServer(current.server_id)) }
async function toggle(server: McpServer) { await act(server, 'toggle', current => current.enabled ? service.disableMcpServer(current.server_id) : service.enableMcpServer(current.server_id)) }
async function act(server: McpServer, action: string, operation: (server: McpServer) => Promise<McpServer>) {
busy.value = `${action}:${server.server_id}`; error.value = ''
try { const current = action === 'toggle' && server.enabled ? server : await approve(server); if (!current) return; await operation(current); await load() }
catch (cause) { error.value = message(cause, '操作失败') }
finally { busy.value = '' }
}
async function remove(server: McpServer) {
if (!confirm(`删除“${server.name}”及其加密凭据?`)) return
try { busy.value = `delete:${server.server_id}`; await service.deleteMcpServer(server.server_id); await load() }
catch (cause) { error.value = message(cause, '删除失败') } finally { busy.value = '' }
}
async function saveSecret(server: McpServer, key: string, kind: SecretKind) {
const draftKey = `${server.server_id}:${kind}:${key}`
const value = secretDrafts[draftKey]?.trim()
if (!value) return
try { busy.value = `secret:${draftKey}`; await service.putMcpServerSecret(server.server_id, key, value, kind); secretDrafts[draftKey] = ''; await load() }
catch (cause) { error.value = message(cause, '保存密钥失败') } finally { busy.value = '' }
}
function splitKeys(value: string) { return value.split(/[\n,]/).map(item => item.trim()).filter(Boolean) }
function message(cause: unknown, fallback: string) { return cause instanceof Error ? cause.message : fallback }
onMounted(load)
</script>
<template>
<section class="feature-page mcp-page">
<header class="feature-header"><div><h1>MCP 服务器</h1><p>管理独立 MCP Server 的连接凭据与工具生命周期</p></div><div class="inline-actions"><button class="button-secondary" :disabled="!!busy" @click="load"><AppIcon :icon="Refresh" /> 刷新</button><button class="button-primary" @click="openCreate"><AppIcon :icon="Plus" /> 新增服务器</button></div></header>
<div class="notice-banner">stdio 本机进程仅在开发环境开放Streamable HTTP 为首选远程传输SSE 仅用于兼容旧服务器uvx 隔离依赖但不是安全沙箱</div>
<div v-if="error" class="error-banner">{{ error }}</div>
<div v-if="!servers.length" class="panel empty"><AppIcon :icon="Connection" :size="34" /><h2>尚未配置 MCP 服务器</h2><p>添加 Server,测试连接成功后才能启用工具。</p><button class="button-primary" @click="openCreate">新增服务器</button></div>
<div v-else class="server-list">
<article v-for="server in servers" :key="server.server_id" class="panel server-card">
<div class="server-main"><div class="server-title"><AppIcon :icon="Connection" :size="24" /><div><h2>{{ server.name }}</h2><code>{{ server.command_summary }}</code></div></div><span class="badge" :class="{ success: server.status === 'ready', error: ['error','unhealthy'].includes(server.status) }">{{ server.status }}</span></div>
<div class="metadata"><span>{{ server.transport }}</span><span>v{{ server.version }}</span><span>{{ server.tools_count }} 个工具</span><span>{{ server.trusted ? '连接已确认' : '等待确认连接' }}</span><span v-if="server.last_test_succeeded">当前配置测试成功</span><span v-if="server.remote_server_name">{{ server.remote_server_name }} {{ server.remote_server_version }}</span></div>
<div v-if="server.error" class="error-banner compact">{{ server.error }}</div>
<div v-if="Object.keys(server.secret_environment).length || Object.keys(server.secret_headers).length" class="secrets">
<label v-for="(configured, key) in server.secret_environment" :key="`env:${key}`"><span>环境变量 · {{ key }} <small>{{ configured ? '已加密保存' : '未配置' }}</small></span><span class="secret-input"><input v-model="secretDrafts[`${server.server_id}:environment:${key}`]" type="password" autocomplete="new-password" placeholder="输入后保存(不会回显)"><button class="button-secondary" @click="saveSecret(server, key, 'environment')">保存</button></span></label>
<label v-for="(configured, key) in server.secret_headers" :key="`header:${key}`"><span>HTTP Header · {{ key }} <small>{{ configured ? '已加密保存' : '未配置' }}</small></span><span class="secret-input"><input v-model="secretDrafts[`${server.server_id}:header:${key}`]" type="password" autocomplete="new-password" placeholder="输入后保存(不会回显)"><button class="button-secondary" @click="saveSecret(server, key, 'header')">保存</button></span></label>
</div>
<footer class="card-actions"><button class="button-secondary" :disabled="!!busy || server.enabled" @click="test(server)"><AppIcon :icon="VideoPlay" /> 测试连接</button><button class="button-secondary" :disabled="!!busy" @click="openEdit(server)"><AppIcon :icon="EditPen" /> 编辑</button><button class="button-danger" :disabled="!!busy" @click="remove(server)"><AppIcon :icon="Delete" /> 删除</button><button class="button-primary" :disabled="!!busy || (!server.enabled && !server.last_test_succeeded)" :title="!server.enabled && !server.last_test_succeeded ? '请先测试当前配置' : ''" @click="toggle(server)">{{ server.enabled ? '停用' : '启用' }}</button></footer>
</article>
</div>
<div v-if="dialogOpen" class="modal-backdrop" @click.self="!busy && closeEditor()">
<form class="modal-card" @submit.prevent="save">
<fieldset :disabled="!!busy" class="editor-fields">
<header><h2><AppIcon :icon="Plus" /> {{ dialogTitle }}</h2><button type="button" class="close" @click="closeEditor">×</button></header>
<div v-if="error" class="error-banner" role="alert">{{ error }}</div>
<div v-if="importedSecrets.length" class="notice-banner">已识别 {{ importedSecrets.length }} 项密钥保存时将单独加密不会写入普通服务器配置取消将清除未保存密钥</div>
<div class="mode-tabs"><button type="button" :class="{ active: editorMode === 'form' }" @click="switchMode('form')">表单配置</button><button type="button" :class="{ active: editorMode === 'json' }" @click="switchMode('json')">JSON 配置</button></div>
<template v-if="editorMode === 'form'">
<label>服务器名称<input v-model="form.name" maxlength="80" placeholder="例如:文件系统工具"></label>
<div class="template-row"><span>服务器配置</span><button type="button" class="template" :class="{ active: form.transport === 'stdio' }" @click="applyTemplate('stdio')">stdio 模板</button><button type="button" class="template" :class="{ active: form.transport === 'streamable_http' }" @click="applyTemplate('streamable_http')">Streamable HTTP</button><button type="button" class="template" :class="{ active: form.transport === 'sse' }" @click="applyTemplate('sse')">SSE兼容</button></div>
<template v-if="form.transport === 'stdio'"><label>可执行命令<input v-model="form.command" placeholder="uvx、npx 或可信可执行文件路径"></label><label>参数(每行一项)<textarea v-model="argsText" rows="5"></textarea></label><div class="two-columns"><label>普通环境变量(JSON<textarea v-model="environmentText" rows="5"></textarea></label><label>敏感环境变量名(每行一项)<textarea v-model="secretKeysText" rows="5" placeholder="API_KEY"></textarea></label></div></template>
<template v-else><label>MCP URL<input v-model="form.url" placeholder="https://example.com/mcp"></label><div class="two-columns"><label>普通 HeaderJSON<textarea v-model="headersText" rows="5" placeholder='{"X-Client":"NotesAgent"}'></textarea></label><label>敏感 Header 名(每行一项)<textarea v-model="secretHeaderKeysText" rows="5" placeholder="Authorization"></textarea></label></div></template>
<label>声明权限逗号分隔可选<input v-model="permissionsText" placeholder="network.request, notes.read"></label>
<div class="two-columns"><label>启动超时<input v-model.number="form.startup_timeout_seconds" type="number" min="1" max="120"></label><label>工具超时<input v-model.number="form.tool_timeout_seconds" type="number" min="1" max="300"></label></div>
</template>
<label v-else>服务器 JSON 配置<textarea v-model="rawConfig" class="json-editor" rows="22" spellcheck="false"></textarea><small>支持 NotesAgent 配置command/args/env 和单服务器 mcpServers 配置已声明的 Secret 及常见 API KeyTokenAuthorization 会拆分后加密保存其他敏感值请显式声明不要把密钥放入命令或参数</small><small>兼容导入 timeout 为启动超时sse_read_timeout 为工具等待上限不保留 SSE 读取超时语义</small></label>
<footer><button type="button" class="button-secondary" @click="closeEditor">取消</button><button class="button-primary" :disabled="busy === 'save'">保存</button></footer>
</fieldset>
</form>
</div>
</section>
</template>
<style scoped>
.editor-fields { display: grid; gap: var(--space-lg); border: 0; padding: 0; margin: 0; min-width: 0; }
.mcp-page { overflow: auto; }.notice-banner,.error-banner { margin-bottom: var(--space-lg); }.server-list { display: grid; gap: var(--space-lg); }.server-card { display: grid; gap: var(--space-md); }
.server-main,.server-title,.metadata,.card-actions,.inline-actions,.template-row,.modal-card header,.modal-card footer { display: flex; align-items: center; gap: var(--space-sm); }.server-main { justify-content: space-between; }.server-title { align-items: flex-start; }.server-title h2 { margin-bottom: 4px; }.server-title code { color: var(--color-text-secondary); overflow-wrap: anywhere; }.metadata { flex-wrap: wrap; color: var(--color-text-tertiary); font-size: var(--font-size-sm); }.metadata span + span::before { content: '·'; margin-right: var(--space-sm); }.compact { margin: 0; }
.card-actions { justify-content: flex-end; border-top: 1px solid var(--color-border-subtle); padding-top: var(--space-md); }.empty { text-align: center; place-items: center; display: grid; gap: var(--space-md); padding: 64px; }.secrets { border: 1px solid var(--color-border-subtle); border-radius: var(--radius-md); padding: var(--space-md); display: grid; gap: var(--space-sm); }.secrets label { display: grid; grid-template-columns: minmax(220px,.7fr) 1fr; align-items: center; gap: var(--space-md); }.secrets small,.modal-card small { color: var(--color-text-tertiary); }.secret-input { display: flex; gap: var(--space-sm); }.secret-input input { flex: 1; }
.modal-backdrop { position: fixed; inset: 0; z-index: 1000; background: rgb(0 0 0 / .48); display: grid; place-items: center; padding: var(--space-xl); }.modal-card { width: min(800px,100%); max-height: calc(100vh - 48px); overflow: auto; background: var(--color-background-primary); border: 1px solid var(--color-border-default); border-radius: var(--radius-xl); box-shadow: var(--shadow-xl); padding: var(--space-xl); display: grid; gap: var(--space-lg); animation: modal-in var(--motion-normal) ease-out; }.modal-card header,.modal-card footer { justify-content: space-between; }.modal-card footer { justify-content: flex-end; }.modal-card label { display: grid; gap: var(--space-xs); font-weight: 600; }.modal-card input,.modal-card textarea { width: 100%; border: 1px solid var(--color-border-default); border-radius: var(--radius-md); padding: 10px 12px; color: var(--color-text-primary); background: var(--color-background-secondary); font: inherit; }.modal-card textarea { resize: vertical; font-family: var(--font-family-mono); font-size: var(--font-size-sm); }.json-editor { line-height: 1.55; }.close { border: 0; background: transparent; color: var(--color-text-secondary); font-size: 28px; cursor: pointer; }
.template-row { flex-wrap: wrap; }.template-row > span { margin-right: auto; font-weight: 600; }.template,.mode-tabs button { border: 1px solid var(--color-border-default); background: var(--color-background-secondary); color: var(--color-text-secondary); padding: 7px 10px; border-radius: var(--radius-md); cursor: pointer; }.template.active,.mode-tabs button.active { color: var(--color-accent-primary); border-color: var(--color-accent-primary); background: var(--color-accent-soft); }.mode-tabs { display: inline-flex; justify-self: start; gap: 2px; padding: 3px; border-radius: var(--radius-md); background: var(--color-background-secondary); }.two-columns { display: grid; grid-template-columns: 1fr 1fr; gap: var(--space-md); }
@keyframes modal-in { from { opacity: 0; transform: translateY(8px) scale(.99); } } @media (max-width:720px) { .two-columns,.secrets label { grid-template-columns:1fr; }.card-actions { flex-wrap:wrap; } }
</style>
@@ -0,0 +1,63 @@
import { describe, expect, it } from 'vitest'
import { emptyMcpConfig, mergeImportedSecrets, normalizeMcpConfig, parseMcpJson } from './configuration'
describe('MCP configuration normalization', () => {
it('retains renamed HTTP drafts with the latest spelling and value', () => {
const config = { ...emptyMcpConfig(), secret_header_keys: ['authorization'] }
const previous = [{ kind: 'header' as const, key: 'Authorization', value: 'old-value' }]
expect(mergeImportedSecrets(config, previous, [])).toEqual([{ kind: 'header', key: 'authorization', value: 'old-value' }])
expect(mergeImportedSecrets(config, previous, [{ kind: 'header', key: 'AUTHORIZATION', value: 'new-value' }])).toEqual([{ kind: 'header', key: 'authorization', value: 'new-value' }])
expect(mergeImportedSecrets(emptyMcpConfig(), previous, [])).toEqual([])
})
it('does not transfer an environment draft across a case-only rename', () => {
const config = { ...emptyMcpConfig(), secret_environment_keys: ['TOKEN', 'token'] }
const previous = [{ kind: 'environment' as const, key: 'TOKEN', value: 'upper' }, { kind: 'environment' as const, key: 'token', value: 'lower' }]
expect(mergeImportedSecrets(config, previous, [])).toEqual(previous)
expect(mergeImportedSecrets({ ...config, secret_environment_keys: ['token'] }, [previous[0]!], [])).toEqual([])
})
it('fills backend defaults for minimal JSON', () => {
const { config } = parseMcpJson('{"name":"demo","command":"uvx"}')
expect(config).toMatchObject({ transport: 'stdio', args: [], headers: {}, environment: {}, permissions: [], secret_header_keys: [] })
})
it('extracts a key pasted into environment despite its existing secret declaration', () => {
const { config, secrets } = normalizeMcpConfig({
name: 'MiniMax', command: 'uvx', secret_environment_keys: ['MINIMAX_API_KEY'],
environment: { MINIMAX_API_KEY: 'synthetic-key', MINIMAX_API_HOST: 'https://api.minimaxi.com' },
})
expect(config.environment).toEqual({ MINIMAX_API_HOST: 'https://api.minimaxi.com' })
expect(config.secret_environment_keys).toEqual(['MINIMAX_API_KEY'])
expect(JSON.stringify(config)).not.toContain('synthetic-key')
expect(secrets).toEqual([{ kind: 'environment', key: 'MINIMAX_API_KEY', value: 'synthetic-key' }])
})
it('imports a standard single-server wrapper and legacy timeouts', () => {
const { config, secrets } = normalizeMcpConfig({ mcpServers: { MiniMax: {
command: 'uvx', args: ['--with', 'mcp<2', 'minimax-coding-plan-mcp', '-y'],
env: { MINIMAX_API_KEY: 'synthetic-key' }, timeout: 120, sse_read_timeout: 300,
} } })
expect(config).toMatchObject({ name: 'MiniMax', transport: 'stdio', environment: {}, startup_timeout_seconds: 120, tool_timeout_seconds: 300 })
expect(secrets).toHaveLength(1)
})
it('extracts case-insensitive HTTP credentials without duplicate declarations', () => {
const { config, secrets } = normalizeMcpConfig({ url: 'https://example.test/mcp', headers: { authorization: 'synthetic' }, secret_header_keys: ['Authorization'] })
expect(config.headers).toEqual({})
expect(config.secret_header_keys).toEqual(['Authorization'])
expect(secrets[0]?.key).toBe('Authorization')
})
it.each([
[{ command: 'uvx', args: 'not-array' }, 'args'],
[{ command: 'uvx', environment: [] }, 'environment'],
[{ command: 'uvx', timeout: 121 }, '启动超时'],
[{ command: 'uvx', args: ['[https://example.test](https://example.test)'] }, '纯 URL'],
[{ command: 'uvx', api_key: 'do-not-echo' }, '顶层'],
[{ command: 'uvx', env: {}, environment: {} }, '只保留一个'],
[{ mcpServers: { one: {}, two: {} } }, '一次导入一个'],
])('rejects invalid fields without leaking their values', (input, hint) => {
expect(() => normalizeMcpConfig(input)).toThrow(hint)
try { normalizeMcpConfig(input) } catch (error) { expect(String(error)).not.toContain('do-not-echo') }
})
})
+139
View File
@@ -0,0 +1,139 @@
import type { McpServerInput } from '@/contracts'
export type SecretKind = 'environment' | 'header'
export interface ImportedSecret { kind: SecretKind; key: string; value: string }
export function mergeImportedSecrets(config: McpServerInput, previous: ImportedSecret[], incoming: ImportedSecret[]): ImportedSecret[] {
const merged = new Map<string, ImportedSecret>()
for (const item of [...previous, ...incoming]) {
const normalize = (key: string) => item.kind === 'header' ? key.toLowerCase() : key
const keys = item.kind === 'header' ? config.secret_header_keys : config.secret_environment_keys
const declared = keys.find(key => normalize(key) === normalize(item.key))
if (declared === undefined) continue
// HTTP identity is case-insensitive, but the Secret API requires the current
// declared spelling. New inline values replace older drafts of that identity.
merged.set(`${item.kind}:${normalize(declared)}`, { ...item, key: declared })
}
return [...merged.values()]
}
export function emptyMcpConfig(): McpServerInput {
return {
name: '', transport: 'stdio', command: '', args: [], url: null, headers: {},
environment: {}, secret_environment_keys: [], secret_header_keys: [], permissions: [],
startup_timeout_seconds: 15, tool_timeout_seconds: 30,
}
}
function object(value: unknown, label: string): Record<string, unknown> {
if (!value || Array.isArray(value) || typeof value !== 'object') throw new Error(`${label}必须是 JSON 对象`)
return value as Record<string, unknown>
}
function strings(value: unknown, label: string): string[] {
if (value === undefined) return []
if (!Array.isArray(value) || value.some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串数组`)
return [...value]
}
function entries(value: unknown, label: string): Record<string, string> {
if (value === undefined) return {}
const result = object(value, label)
if (Object.values(result).some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串键值 JSON 对象`)
return { ...result } as Record<string, string>
}
function timeout(value: unknown, fallback: number, max: number, label: string): number {
if (value === undefined) return fallback
if (typeof value !== 'number' || !Number.isFinite(value) || value < 1 || value > max) throw new Error(`${label}必须是 1${max} 秒之间的数字`)
return value
}
// Do not silently rewrite executable arguments or secret values copied from chat.
function checkUrl(value: string, label: string) {
if (/^\[https?:\/\//i.test(value)) throw new Error(`${label}请填写纯 URL,不要粘贴 Markdown 链接`)
}
export function parseMcpJson(raw: string, fallbackName = '', requireConnection = true) {
let parsed: unknown
try { parsed = JSON.parse(raw) }
catch { throw new Error('服务器配置不是有效 JSON;请检查逗号、引号和无效的 \\_ 转义') }
return normalizeMcpConfig(parsed, fallbackName, requireConnection)
}
/** Normalize external client JSON before it reaches either the form or the API.
* Inline secrets leave the public config here and are sent only to the Secret API.
*/
export function normalizeMcpConfig(parsed: unknown, fallbackName = '', requireConnection = true) {
let raw = object(parsed, '服务器配置')
if ('mcpServers' in raw) {
const servers = Object.entries(object(raw.mcpServers, 'mcpServers'))
if (servers.length !== 1) throw new Error('请一次导入一个 MCP 服务器')
fallbackName = servers[0]![0]
raw = object(servers[0]![1], '服务器配置')
}
const allowed = new Set([...Object.keys(emptyMcpConfig()), 'version', 'env', 'type', 'timeout', 'sse_read_timeout'])
if (Object.keys(raw).some(key => !allowed.has(key))) {
// Never echo arbitrary unknown keys: pasted secrets sometimes become JSON keys.
throw new Error('服务器配置含不支持的字段;API Key 请放在 env/environment 的对应变量中,不要放在顶层')
}
if (raw.env !== undefined && raw.environment !== undefined) throw new Error('env 与 environment 请只保留一个,避免覆盖配置')
const transport = raw.transport ?? raw.type ?? (raw.url ? 'streamable_http' : 'stdio')
if (!['stdio', 'streamable_http', 'sse'].includes(transport as string)) throw new Error('transport 必须是 stdio、streamable_http 或 sse')
const config = emptyMcpConfig()
config.transport = transport as McpServerInput['transport']
const name = raw.name ?? (fallbackName || (typeof raw.command === 'string' ? raw.command : 'MCP 服务器'))
if (typeof name !== 'string' || (requireConnection && !name.trim()) || name.trim().length > 80) throw new Error('服务器名称必须为 180 个字符')
config.name = name.trim()
for (const key of ['command', 'url'] as const) {
const value = raw[key]
if (value !== undefined && value !== null && typeof value !== 'string') throw new Error(`${key}必须是字符串`)
config[key] = typeof value === 'string' ? value.trim() : null
}
config.args = strings(raw.args, 'args')
if (config.args.length > 64) throw new Error('args 最多允许 64 项')
for (const value of config.args) checkUrl(value, 'args 中的地址')
config.environment = entries(raw.environment ?? raw.env, 'environment/env')
config.headers = entries(raw.headers, 'headers')
config.secret_environment_keys = [...new Set(strings(raw.secret_environment_keys, 'secret_environment_keys'))]
config.secret_header_keys = [...new Set(strings(raw.secret_header_keys, 'secret_header_keys'))]
config.permissions = strings(raw.permissions, 'permissions')
config.startup_timeout_seconds = timeout(raw.startup_timeout_seconds ?? raw.timeout, 15, 120, '启动超时')
// Compatibility policy: legacy read timeout becomes the tool wait budget, not an SSE transport setting.
config.tool_timeout_seconds = timeout(raw.tool_timeout_seconds ?? raw.sse_read_timeout, 30, 300, '工具超时')
if (config.transport === 'stdio') {
if (requireConnection && !config.command) throw new Error('stdio 配置必须填写 command')
if (config.url || Object.keys(config.headers).length || config.secret_header_keys.length) throw new Error('stdio 配置不能包含 URL 或 HTTP Header')
} else {
if (requireConnection && !config.url) throw new Error('HTTP/SSE 配置必须填写 url')
if (config.url) {
checkUrl(config.url, 'url')
let url: URL
try { url = new URL(config.url) } catch { throw new Error('url 必须是有效的 HTTP(S) 地址') }
if (!['http:', 'https:'].includes(url.protocol) || url.username || url.password || url.hash) throw new Error('url 必须为不含账号密码或片段的 HTTP(S) 地址')
}
if (config.command || config.args.length || Object.keys(config.environment).length || config.secret_environment_keys.length) throw new Error('HTTP/SSE 配置不能包含 command、args 或环境变量')
}
const secrets: ImportedSecret[] = []
for (const kind of ['environment', 'header'] as const) {
const values = kind === 'environment' ? config.environment : config.headers
const keys = kind === 'environment' ? config.secret_environment_keys : config.secret_header_keys
const identity = (key: string) => kind === 'header' ? key.toLowerCase() : key
const allKeys = [...Object.keys(values), ...keys]
if (kind === 'header' && (new Set(keys.map(identity)).size !== keys.length || new Set(Object.keys(values).map(identity)).size !== Object.keys(values).length)) throw new Error('HTTP Header 名称不能仅大小写不同而重复声明')
const validKey = kind === 'environment' ? /^[A-Za-z_][A-Za-z0-9_]{0,127}$/ : /^[!#$%&'*+.^_`|~0-9A-Za-z-]{1,128}$/
if (allKeys.some(key => !validKey.test(key))) throw new Error(`${kind === 'environment' ? '环境变量' : 'Header'}名称无效;敏感变量名只能填名称,不能填密钥值`)
for (const [key, value] of Object.entries(values)) {
const declared = keys.find(item => identity(item) === identity(key))
const sensitive = /api[_-]?key|token|secret|password|authorization|cookie|credential/i.test(key)
if (declared || sensitive) {
if (!value || value.length > 32768) throw new Error('密钥值必须为 132768 个字符')
const secretKey = declared ?? key
if (!declared) keys.push(key)
secrets.push({ kind, key: secretKey, value })
delete values[key]
} else if (/host|url|endpoint/i.test(key)) checkUrl(value, '环境变量或 Header 地址')
}
}
return { config, secrets }
}
@@ -0,0 +1,109 @@
// @vitest-environment happy-dom
import { flushPromises, mount } from '@vue/test-utils'
import { createPinia } from 'pinia'
import { createMemoryHistory, createRouter } from 'vue-router'
import { beforeEach, describe, expect, it, vi } from 'vitest'
import type { Plugin } from '@/contracts'
import * as pluginService from '@/services/pluginService'
import PluginMcpPanel from './PluginMcpPanel.vue'
import { useWorkspaceStore } from '@/stores/workspace'
vi.mock('@/services/pluginService', async (loadOriginal) => {
const original = await loadOriginal<typeof import('@/services/pluginService')>()
return {
...original,
getPluginHostStatus: vi.fn(),
restartPluginHost: vi.fn(),
getPluginSettings: vi.fn(),
updatePluginSettings: vi.fn(),
putPluginSecret: vi.fn(),
deletePluginSecret: vi.fn(),
listPluginCommands: vi.fn(),
executePluginCommand: vi.fn(),
}
})
const plugin: Plugin = {
plugin_id: 'mcp-demo',
name: 'MCP Demo',
version: '1.0.0',
description: 'demo',
status: 'ready',
enabled: true,
permissions: [],
contributions: [
{ type: 'settings_section', id: 'mcp-demo.general', name: 'settings' },
{ type: 'command', id: 'mcp-demo.run', name: 'run' },
],
backend_type: 'mcp',
transport: 'stdio',
}
async function render() {
const pinia = createPinia()
const router = createRouter({
history: createMemoryHistory(),
routes: [{ path: '/', component: { template: '<div />' } }],
})
await router.push('/')
const wrapper = mount(PluginMcpPanel, {
props: { plugin },
global: { plugins: [pinia, router], stubs: { AppIcon: true } },
})
const workspaceStore = useWorkspaceStore(pinia)
workspaceStore.vaultId = 'default'
workspaceStore.hasVault = true
return wrapper
}
beforeEach(() => {
vi.clearAllMocks()
vi.mocked(pluginService.getPluginHostStatus).mockResolvedValue({
plugin_id: 'mcp-demo', backend_type: 'mcp', transport: 'stdio',
status: 'ready', tools_count: 2, server_name: 'demo',
})
vi.mocked(pluginService.getPluginSettings).mockResolvedValue({
plugin_id: 'mcp-demo',
schema_version: 1,
fields: [
{ key: 'limit', label: '数量', description: '', type: 'number', required: true, options: [] },
{ key: 'api_key', label: 'API Key', description: '', type: 'secret', required: true, options: [] },
],
values: { limit: 5 },
secrets: { api_key: { configured: false } },
})
vi.mocked(pluginService.putPluginSecret).mockResolvedValue({
plugin_id: 'mcp-demo', key: 'api_key', configured: true,
})
vi.mocked(pluginService.listPluginCommands).mockResolvedValue([])
})
describe('PluginMcpPanel', () => {
it('loads MCP Host status and exposes restart controls', async () => {
const wrapper = await render()
await flushPromises()
expect(pluginService.getPluginHostStatus).toHaveBeenCalledWith('mcp-demo')
expect(wrapper.text()).toContain('demo')
expect(wrapper.text()).toContain('工具数量')
})
it('builds settings fields from schema and writes secrets separately', async () => {
const wrapper = await render()
const settingsTab = wrapper.findAll('button').find((button) => button.text() === '设置与密钥')
expect(settingsTab).toBeTruthy()
await settingsTab!.trigger('click')
await flushPromises()
expect(wrapper.text()).toContain('数量')
expect(wrapper.text()).toContain('API Key')
await wrapper.get('input[type="password"]').setValue('secret-only-in-request')
const secretButton = wrapper.findAll('button').find((button) => button.text() === '安全保存')
expect(secretButton).toBeTruthy()
await secretButton!.trigger('click')
await flushPromises()
expect(pluginService.putPluginSecret).toHaveBeenCalledWith('mcp-demo', 'api_key', 'secret-only-in-request')
expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
expect(wrapper.text()).toContain('已配置')
})
})
@@ -0,0 +1,275 @@
<script setup lang="ts">
import { Key, Refresh, VideoPlay } from '@element-plus/icons-vue'
import { computed, ref, watch } from 'vue'
import { useRouter } from 'vue-router'
import AppIcon from '@/components/common/AppIcon.vue'
import type { Plugin, PluginCommand, PluginHostStatus, PluginSettingField, PluginSettingsSchema } from '@/contracts'
import * as pluginService from '@/services/pluginService'
import { useEditorStore } from '@/stores/editor'
import { usePluginStore } from '@/stores/plugin'
import { useWorkspaceStore } from '@/stores/workspace'
const props = defineProps<{ plugin: Plugin }>()
const pluginStore = usePluginStore()
const editorStore = useEditorStore()
const workspaceStore = useWorkspaceStore()
const router = useRouter()
const activeTab = ref<'host' | 'settings' | 'commands'>('host')
const host = ref<PluginHostStatus | null>(null)
const schema = ref<PluginSettingsSchema | null>(null)
const values = ref<Record<string, unknown>>({})
//
const secrets = ref<Record<string, string>>({})
const commands = ref<PluginCommand[]>([])
const argumentsByCommand = ref<Record<string, Record<string, unknown>>>({})
const loading = ref(false)
const busy = ref('')
const error = ref('')
const notice = ref('')
let loadVersion = 0
const hasSettings = computed(() => props.plugin.contributions.some((item) => item.type === 'settings_section'))
const tabs = computed(() => [
...(props.plugin.backend_type === 'mcp' ? [{ id: 'host' as const, label: 'MCP Host' }] : []),
...(hasSettings.value ? [{ id: 'settings' as const, label: '设置与密钥' }] : []),
{ id: 'commands' as const, label: '插件命令' },
])
watch(() => props.plugin.plugin_id, () => {
loadVersion++
activeTab.value = props.plugin.backend_type === 'mcp' ? 'host' : hasSettings.value ? 'settings' : 'commands'
host.value = null
schema.value = null
values.value = {}
secrets.value = {}
commands.value = []
void loadActive()
}, { immediate: true })
function feedback(message = '') { error.value = message; notice.value = '' }
function message(reason: unknown, fallback: string) { return reason instanceof Error ? reason.message : fallback }
function formatTime(value?: string | null) { return value ? new Date(value).toLocaleString() : '—' }
async function selectTab(tab: typeof activeTab.value) {
activeTab.value = tab
await loadActive()
}
async function loadActive() {
const version = ++loadVersion
const pluginId = props.plugin.plugin_id
const tab = activeTab.value
feedback()
loading.value = true
try {
if (tab === 'host') {
const loadedHost = await pluginService.getPluginHostStatus(pluginId)
if (version === loadVersion) host.value = loadedHost
}
if (tab === 'settings') {
const loadedSchema = await pluginService.getPluginSettings(pluginId)
if (version === loadVersion) {
schema.value = loadedSchema
values.value = { ...loadedSchema.values }
}
}
if (tab === 'commands') {
const loadedCommands = (await pluginService.listPluginCommands()).filter((command) => command.plugin_id === pluginId)
if (version === loadVersion) {
commands.value = loadedCommands
for (const command of loadedCommands) argumentsByCommand.value[command.command_id] = {}
}
}
} catch (reason) {
if (version === loadVersion) feedback(message(reason, 'MCP 数据加载失败'))
} finally {
if (version === loadVersion) loading.value = false
}
}
async function restartHost() {
busy.value = 'host'
feedback()
try {
await pluginService.restartPluginHost(props.plugin.plugin_id)
host.value = await pluginService.getPluginHostStatus(props.plugin.plugin_id)
await pluginStore.loadPlugins()
notice.value = 'MCP Host 已重启。'
} catch (reason) { feedback(message(reason, 'MCP Host 重启失败')) } finally { busy.value = '' }
}
function updateValue(field: PluginSettingField, raw: string | boolean) {
values.value[field.key] = field.type === 'number' && typeof raw === 'string' ? (raw === '' ? null : Number(raw)) : raw
}
async function saveSettings() {
if (!schema.value) return
busy.value = 'settings'
feedback()
try {
schema.value = await pluginService.updatePluginSettings(props.plugin.plugin_id, schema.value.schema_version, values.value)
values.value = { ...schema.value.values }
notice.value = '普通设置已保存。'
} catch (reason) { feedback(message(reason, '设置保存失败')) } finally { busy.value = '' }
}
async function saveSecret(field: PluginSettingField) {
const secret = secrets.value[field.key]?.trim()
if (!secret) { feedback('请输入' + field.label); return }
busy.value = 'secret:' + field.key
feedback()
try {
const state = await pluginService.putPluginSecret(props.plugin.plugin_id, field.key, secret)
if (schema.value) schema.value.secrets[field.key] = { configured: state.configured }
secrets.value[field.key] = ''
notice.value = field.label + '已加密保存。'
} catch (reason) { feedback(message(reason, '密钥保存失败')) } finally { busy.value = '' }
}
async function deleteSecret(field: PluginSettingField) {
if (!confirm('删除已保存的' + field.label + '')) return
busy.value = 'secret:' + field.key
feedback()
try {
const state = await pluginService.deletePluginSecret(props.plugin.plugin_id, field.key)
if (schema.value) schema.value.secrets[field.key] = { configured: state.configured }
secrets.value[field.key] = ''
notice.value = field.label + '已删除。'
} catch (reason) { feedback(message(reason, '密钥删除失败')) } finally { busy.value = '' }
}
function properties(command: PluginCommand): Record<string, Record<string, unknown>> {
const result = command.parameters.properties
return result && typeof result === 'object' && !Array.isArray(result) ? result as Record<string, Record<string, unknown>> : {}
}
function required(command: PluginCommand, key: string) {
return Array.isArray(command.parameters.required) && command.parameters.required.includes(key)
}
function commandAvailable(command: PluginCommand) {
if (!command.enabled) return false
return command.when.every((condition) => {
if (condition === 'workspace.has_vault') return Boolean(workspaceStore.vaultId)
if (condition === 'editor.has_note') return Boolean(editorStore.currentNoteId)
// Plugin
if (condition === 'editor.has_selection') return false
return false
})
}
function updateArgument(commandId: string, key: string, raw: string, definition: Record<string, unknown>) {
const target = argumentsByCommand.value[commandId] ??= {}
if (definition.type === 'number' || definition.type === 'integer') target[key] = raw === '' ? undefined : Number(raw)
else if (definition.type === 'boolean') target[key] = raw === 'true'
else target[key] = raw
}
async function execute(command: PluginCommand) {
busy.value = command.command_id
feedback()
try {
const result = await pluginService.executePluginCommand(command.command_id, argumentsByCommand.value[command.command_id] ?? {}, {
vault_id: workspaceStore.hasVault ? workspaceStore.vaultId : null,
note_id: editorStore.currentNoteId,
file_path: editorStore.currentFilePath,
selection: null,
})
if (result.effect.type === 'notification') notice.value = result.effect.payload.message
else if (result.effect.type === 'job') notice.value = '后台任务已创建:' + result.effect.payload.job_id
else if (result.effect.type === 'navigate') {
const routes: Record<string, string> = {
'vault-entry': '/', workspace: '/workspace', search: '/search', chat: '/chat',
agent: '/agent/runs', tasks: '/tasks', skills: '/extensions/skills',
plugins: '/extensions/plugins', themes: '/themes', settings: '/settings',
}
await router.push(routes[result.effect.payload.route])
} else if (result.effect.type === 'refresh') {
await loadActive()
notice.value = '相关数据已刷新。'
} else notice.value = '命令执行完成。'
} catch (reason) { feedback(message(reason, '命令执行失败')) } finally { busy.value = '' }
}
</script>
<template>
<section class="mcp-panel">
<nav class="mcp-tabs" aria-label="MCP Plugin 配置">
<button v-for="tab in tabs" :key="tab.id" :class="{ active: activeTab === tab.id }" @click="selectTab(tab.id)">{{ tab.label }}</button>
</nav>
<div v-if="error" class="error-banner">{{ error }}</div>
<div v-if="notice" class="notice-banner">{{ notice }}</div>
<div v-if="activeTab === 'host'" class="mcp-section">
<div class="section-head"><div><h3>MCP Host 状态</h3><p>查看协议协商运行状态与 Host 错误</p></div><div class="inline-actions"><button class="button-secondary" :disabled="loading" @click="loadActive"><AppIcon :icon="Refresh" :size="15" />刷新</button><button class="button-primary" :disabled="busy === 'host' || !plugin.enabled" @click="restartHost">{{ busy === 'host' ? '重启中' : '重启 Host' }}</button></div></div>
<div v-if="host" class="status-grid">
<div><span>状态</span><strong><i class="status-dot" :class="host.status"></i>{{ host.status }}</strong></div>
<div><span>服务</span><strong>{{ host.server_name || '—' }} {{ host.server_version || '' }}</strong></div>
<div><span>协议版本</span><strong>{{ host.protocol_version || '—' }}</strong></div>
<div><span>工具数量</span><strong>{{ host.tools_count }}</strong></div>
<div><span>启动时间</span><strong>{{ formatTime(host.started_at) }}</strong></div>
<div><span>最后心跳</span><strong>{{ formatTime(host.last_seen_at) }}</strong></div>
</div>
<div v-else-if="loading" class="empty-state">正在读取 Host 状态</div>
<div v-if="host?.error" class="error-banner host-error">{{ host.error }}</div>
<p class="security-hint">当前仅运行插件清单声明的 stdio MCP Server不开放任意 Shell 命令和环境变量编辑</p>
</div>
<div v-else-if="activeTab === 'settings'" class="mcp-section">
<div class="section-head"><div><h3>设置与密钥</h3><p>表单由后端 Schema 生成密钥不会被读取或回显</p></div><button class="button-primary" :disabled="!schema || busy === 'settings'" @click="saveSettings">{{ busy === 'settings' ? '保存中' : '保存普通设置' }}</button></div>
<div v-if="schema" class="settings-list">
<div v-for="field in schema.fields" :key="field.key" class="setting-row">
<div class="field-copy"><label :for="'plugin-setting-' + field.key"><AppIcon v-if="field.type === 'secret'" :icon="Key" :size="15" />{{ field.label }}<em v-if="field.required">必填</em></label><p>{{ field.description || (field.type === 'secret' ? '加密保存,不在页面回显。' : '') }}</p></div>
<template v-if="field.type === 'secret'">
<div class="secret-control"><input :id="'plugin-setting-' + field.key" :value="secrets[field.key] || ''" class="input" type="password" autocomplete="new-password" :placeholder="schema.secrets[field.key]?.configured ? '已配置;输入新值可替换' : '输入密钥'" @input="secrets[field.key] = ($event.target as HTMLInputElement).value"><button class="button-secondary" :disabled="!secrets[field.key]?.trim() || busy === 'secret:' + field.key" @click="saveSecret(field)">安全保存</button><button v-if="schema.secrets[field.key]?.configured" class="button-danger" @click="deleteSecret(field)">删除</button></div>
<span class="secret-state" :class="{ configured: schema.secrets[field.key]?.configured }">{{ schema.secrets[field.key]?.configured ? '已配置' : '未配置' }}</span>
</template>
<template v-else-if="field.type === 'boolean'"><label class="check-control"><input :id="'plugin-setting-' + field.key" type="checkbox" :checked="Boolean(values[field.key])" @change="updateValue(field, ($event.target as HTMLInputElement).checked)">{{ values[field.key] ? '开启' : '关闭' }}</label></template>
<template v-else-if="field.type === 'select'"><select :id="'plugin-setting-' + field.key" class="select" :value="values[field.key]" @change="updateValue(field, ($event.target as HTMLSelectElement).value)"><option v-for="option in field.options" :key="option" :value="option">{{ option }}</option></select></template>
<template v-else><input :id="'plugin-setting-' + field.key" class="input" :type="field.type === 'number' ? 'number' : 'text'" :min="field.minimum ?? undefined" :max="field.maximum ?? undefined" :required="field.required" :value="values[field.key] ?? ''" @input="updateValue(field, ($event.target as HTMLInputElement).value)"></template>
</div>
</div>
<div v-else-if="loading" class="empty-state">正在读取 Plugin 设置</div>
</div>
<div v-else class="mcp-section">
<div class="section-head"><div><h3>Plugin 命令</h3><p>执行该 Plugin 注册的受控 Command Contribution</p></div><button class="button-secondary" :disabled="loading" @click="loadActive"><AppIcon :icon="Refresh" :size="15" />刷新</button></div>
<div v-if="commands.length" class="command-list">
<article v-for="command in commands" :key="command.command_id" class="item-card command-card">
<div class="command-head"><div><strong>{{ command.title }}</strong><p>{{ command.description || command.command_id }}</p></div><span class="badge" :class="{ success: commandAvailable(command), warning: command.enabled && !commandAvailable(command) }">{{ commandAvailable(command) ? '可执行' : command.enabled ? '缺少上下文' : '不可用' }}</span></div>
<div v-if="Object.keys(properties(command)).length" class="command-fields">
<label v-for="(definition, key) in properties(command)" :key="key" class="field"><span>{{ String(definition.title || key) }}<em v-if="required(command, key)">必填</em></span><select v-if="Array.isArray(definition.enum)" class="select" @change="updateArgument(command.command_id, key, ($event.target as HTMLSelectElement).value, definition)"><option value="">请选择</option><option v-for="option in definition.enum" :key="String(option)" :value="String(option)">{{ option }}</option></select><select v-else-if="definition.type === 'boolean'" class="select" @change="updateArgument(command.command_id, key, ($event.target as HTMLSelectElement).value, definition)"><option value="false">否</option><option value="true">是</option></select><input v-else class="input" :type="definition.type === 'number' || definition.type === 'integer' ? 'number' : 'text'" @input="updateArgument(command.command_id, key, ($event.target as HTMLInputElement).value, definition)"></label>
</div>
<button class="button-primary command-run" :disabled="!commandAvailable(command) || busy === command.command_id" @click="execute(command)"><AppIcon :icon="VideoPlay" :size="15" />{{ busy === command.command_id ? '执行中…' : '执行命令' }}</button>
</article>
</div>
<div v-else-if="!loading" class="empty-state"><div><strong>没有可用命令</strong><p>启用 Plugin 已注册的命令会出现在这里</p></div></div>
</div>
</section>
</template>
<style scoped>
.mcp-panel { margin-top: var(--space-xl); padding-top: var(--space-xl); border-top: 1px solid var(--color-border-default); }
.mcp-tabs { display: flex; gap: var(--space-xs); margin-bottom: var(--space-xl); padding: var(--space-xs); border: 1px solid var(--color-border-default); border-radius: var(--radius-lg); background: var(--color-background-secondary); }
.mcp-tabs button { padding: 9px var(--space-md); border-radius: var(--radius-md); color: var(--color-text-secondary); }
.mcp-tabs button:hover { background: var(--color-background-hover); }
.mcp-tabs button.active { background: var(--color-surface-primary); color: var(--color-accent-primary); box-shadow: var(--shadow-sm); }
.mcp-section { min-height: 220px; }
.section-head, .command-head { display: flex; align-items: flex-start; justify-content: space-between; gap: var(--space-md); margin-bottom: var(--space-lg); }
.section-head p, .command-head p { margin-top: var(--space-xs); color: var(--color-text-tertiary); font-size: var(--font-size-sm); }
.section-head button, .command-run { display: inline-flex; align-items: center; gap: var(--space-xs); }
.status-grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(165px, 1fr)); gap: var(--space-sm); }
.status-grid > div { display: grid; gap: var(--space-xs); padding: var(--space-md); border: 1px solid var(--color-border-subtle); border-radius: var(--radius-md); background: var(--color-background-secondary); }
.status-grid span { color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
.status-grid strong { display: flex; align-items: center; gap: var(--space-xs); font-size: var(--font-size-sm); }
.status-dot { width: 8px; height: 8px; border-radius: 50%; background: var(--color-text-tertiary); }
.status-dot.ready { background: var(--color-success); box-shadow: 0 0 0 4px var(--color-success-soft); }
.status-dot.error, .status-dot.unhealthy { background: var(--color-error); box-shadow: 0 0 0 4px var(--color-error-soft); }
.status-dot.starting { background: var(--color-warning); box-shadow: 0 0 0 4px var(--color-warning-soft); }
.security-hint { margin-top: var(--space-lg); padding: var(--space-md); border-left: 3px solid var(--color-info); background: var(--color-info-soft); color: var(--color-text-secondary); font-size: var(--font-size-sm); }
.host-error { margin-top: var(--space-lg); }
.settings-list { display: grid; }
.setting-row { display: grid; grid-template-columns: minmax(180px, .9fr) minmax(260px, 1.1fr) auto; align-items: center; gap: var(--space-lg); padding: var(--space-lg) 0; border-bottom: 1px solid var(--color-border-subtle); }
.field-copy label { display: flex; align-items: center; gap: var(--space-xs); font-weight: 650; }
.field-copy p { margin-top: var(--space-xs); color: var(--color-text-tertiary); font-size: var(--font-size-sm); }
em { margin-left: var(--space-xs); color: var(--color-error); font-size: var(--font-size-xs); font-style: normal; }
.secret-control { display: flex; gap: var(--space-xs); }
.secret-state { color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
.secret-state.configured { color: var(--color-success); }
.check-control { display: flex; align-items: center; gap: var(--space-sm); color: var(--color-text-secondary); }
.check-control input { width: 18px; height: 18px; accent-color: var(--color-accent-primary); }
.command-list, .command-card { display: grid; gap: var(--space-sm); }
.command-card:hover { transform: none; }
.command-fields { display: grid; grid-template-columns: repeat(auto-fit, minmax(180px, 1fr)); gap: var(--space-md); }
.command-run { justify-self: end; }
@media (max-width: 800px) { .mcp-tabs { overflow-x: auto; } .mcp-tabs button { flex: 0 0 auto; } .setting-row { grid-template-columns: 1fr; gap: var(--space-sm); } .secret-control { flex-wrap: wrap; } }
</style>
@@ -1,6 +1,7 @@
<script setup lang="ts"> <script setup lang="ts">
import { Connection } from '@element-plus/icons-vue' import { Connection } from '@element-plus/icons-vue'
import AppIcon from '@/components/common/AppIcon.vue' import AppIcon from '@/components/common/AppIcon.vue'
import PluginMcpPanel from './PluginMcpPanel.vue'
import { onMounted, ref } from 'vue' import { onMounted, ref } from 'vue'
import { usePluginStore } from '@/stores/plugin' import { usePluginStore } from '@/stores/plugin'
@@ -16,7 +17,7 @@ async function uninstall(id: string, name: string) { if (!confirm(`卸载“${na
<template> <template>
<section class="feature-page"> <section class="feature-page">
<header class="feature-header"><div><h1>Plugin 管理</h1><p>管理插件生命周期权限和受控 Contribution</p></div><button class="button-primary" @click="install">安装 Plugin</button></header> <header class="feature-header"><div><h1>Plugin MCP</h1><p>管理插件生命周期MCP Host权限和受控 Contribution</p></div><button class="button-primary" @click="install">安装 Plugin</button></header>
<div v-if="pluginStore.error || actionError" class="error-banner">{{ pluginStore.error || actionError }}</div> <div v-if="pluginStore.error || actionError" class="error-banner">{{ pluginStore.error || actionError }}</div>
<div v-if="pluginStore.selectedPlugin" class="panel"> <div v-if="pluginStore.selectedPlugin" class="panel">
<div class="detail-head"><div><span class="badge" :class="{ success: pluginStore.selectedPlugin.status === 'ready', error: pluginStore.selectedPlugin.status === 'error', warning: pluginStore.selectedPlugin.status === 'permission_required' }">{{ pluginStore.selectedPlugin.status }}</span><h2>{{ pluginStore.selectedPlugin.icon }} {{ pluginStore.selectedPlugin.name }}</h2><p class="muted">v{{ pluginStore.selectedPlugin.version }} · {{ pluginStore.selectedPlugin.backend_type || 'none' }}/{{ pluginStore.selectedPlugin.transport || 'none' }}</p></div><div class="inline-actions"><button v-if="pluginStore.selectedPlugin.status === 'permission_required'" class="button-primary" @click="grant(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.permissions)">授权权限</button><button class="button-secondary" @click="toggle(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.enabled)">{{ pluginStore.selectedPlugin.enabled ? '停用' : '启用' }}</button><button class="button-danger" @click="uninstall(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.name)">卸载</button></div></div> <div class="detail-head"><div><span class="badge" :class="{ success: pluginStore.selectedPlugin.status === 'ready', error: pluginStore.selectedPlugin.status === 'error', warning: pluginStore.selectedPlugin.status === 'permission_required' }">{{ pluginStore.selectedPlugin.status }}</span><h2>{{ pluginStore.selectedPlugin.icon }} {{ pluginStore.selectedPlugin.name }}</h2><p class="muted">v{{ pluginStore.selectedPlugin.version }} · {{ pluginStore.selectedPlugin.backend_type || 'none' }}/{{ pluginStore.selectedPlugin.transport || 'none' }}</p></div><div class="inline-actions"><button v-if="pluginStore.selectedPlugin.status === 'permission_required'" class="button-primary" @click="grant(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.permissions)">授权权限</button><button class="button-secondary" @click="toggle(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.enabled)">{{ pluginStore.selectedPlugin.enabled ? '停用' : '启用' }}</button><button class="button-danger" @click="uninstall(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.name)">卸载</button></div></div>
@@ -24,6 +25,7 @@ async function uninstall(id: string, name: string) { if (!confirm(`卸载“${na
<div class="detail-grid"><div><h3>权限</h3><div class="tag-list"><span v-for="permission in pluginStore.selectedPlugin.permissions" :key="permission" class="badge warning">{{ permission }}</span></div></div><div><h3>Contribution</h3><div class="contribution-list"><div v-for="item in pluginStore.selectedPlugin.contributions" :key="item.id" class="item-card"><span class="badge info">{{ item.type }}</span><strong>{{ item.name }}</strong><p class="subtle">{{ item.description || item.id }}</p></div></div></div></div> <div class="detail-grid"><div><h3>权限</h3><div class="tag-list"><span v-for="permission in pluginStore.selectedPlugin.permissions" :key="permission" class="badge warning">{{ permission }}</span></div></div><div><h3>Contribution</h3><div class="contribution-list"><div v-for="item in pluginStore.selectedPlugin.contributions" :key="item.id" class="item-card"><span class="badge info">{{ item.type }}</span><strong>{{ item.name }}</strong><p class="subtle">{{ item.description || item.id }}</p></div></div></div></div>
<div v-if="pluginStore.selectedPlugin.last_error" class="error-banner last-error">{{ pluginStore.selectedPlugin.last_error }}</div> <div v-if="pluginStore.selectedPlugin.last_error" class="error-banner last-error">{{ pluginStore.selectedPlugin.last_error }}</div>
<div v-if="pluginStore.selectedPlugin.dependent_skills?.length" class="notice-banner last-error">依赖此插件的 Skill{{ pluginStore.selectedPlugin.dependent_skills.join('') }}</div> <div v-if="pluginStore.selectedPlugin.dependent_skills?.length" class="notice-banner last-error">依赖此插件的 Skill{{ pluginStore.selectedPlugin.dependent_skills.join('') }}</div>
<PluginMcpPanel :plugin="pluginStore.selectedPlugin" />
</div> </div>
<div v-else class="feature-grid"><article v-for="plugin in pluginStore.plugins" :key="plugin.plugin_id" class="item-card extension-card" @click="pluginStore.selectPlugin(plugin.plugin_id)"><div class="extension-title"><AppIcon :icon="Connection" :size="22" /><div><strong>{{ plugin.name }}</strong><p>v{{ plugin.version }}</p></div><span class="badge" :class="{ success: plugin.status === 'ready', error: plugin.status === 'error', warning: plugin.status === 'permission_required' }">{{ plugin.status }}</span></div><p class="muted">{{ plugin.description }}</p><p class="subtle">{{ plugin.permissions.length }} 项权限 · {{ plugin.contributions.length }} Contribution</p></article></div> <div v-else class="feature-grid"><article v-for="plugin in pluginStore.plugins" :key="plugin.plugin_id" class="item-card extension-card" @click="pluginStore.selectPlugin(plugin.plugin_id)"><div class="extension-title"><AppIcon :icon="Connection" :size="22" /><div><strong>{{ plugin.name }}</strong><p>v{{ plugin.version }}</p></div><span class="badge" :class="{ success: plugin.status === 'ready', error: plugin.status === 'error', warning: plugin.status === 'permission_required' }">{{ plugin.status }}</span></div><p class="muted">{{ plugin.description }}</p><p class="subtle">{{ plugin.permissions.length }} 项权限 · {{ plugin.contributions.length }} Contribution</p></article></div>
</section> </section>
+4 -1
View File
@@ -66,13 +66,16 @@ async function openResult(result: SearchResult) {
</template> </template>
<style scoped> <style scoped>
.search-page > * { width: min(100%, 1040px); margin-inline: auto; }
.search-form { display: grid; grid-template-columns: 1fr auto; gap: var(--space-md); margin-bottom: var(--space-lg); } .search-form { display: grid; grid-template-columns: 1fr auto; gap: var(--space-md); margin-bottom: var(--space-lg); }
.search-input { height: 44px; font-size: var(--font-size-lg); } .search-input { height: 44px; font-size: var(--font-size-lg); }
.advanced { grid-column: 1 / -1; } .advanced { grid-column: 1 / -1; }
.results-header, .result-title, .result-meta { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); } .results-header, .result-title, .result-meta { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); }
.results-header { margin: var(--space-xl) 0 var(--space-md); color: var(--color-text-secondary); } .results-header { margin: var(--space-xl) 0 var(--space-md); color: var(--color-text-secondary); }
.result-list { display: grid; gap: var(--space-md); } .result-list { display: grid; gap: var(--space-md); }
.result-card { cursor: pointer; } .result-card { position: relative; cursor: pointer; overflow: hidden; }
.result-card::before { content: ''; position: absolute; inset: 0 auto 0 0; width: 3px; background: var(--color-accent-primary); opacity: 0; transform: scaleY(.45); transition: opacity var(--motion-fast), transform var(--motion-fast); }
.result-card:hover::before { opacity: 1; transform: scaleY(1); }
.snippet { margin: var(--space-md) 0; line-height: var(--line-height-relaxed); } .snippet { margin: var(--space-md) 0; line-height: var(--line-height-relaxed); }
.result-meta { color: var(--color-text-tertiary); font-size: var(--font-size-xs); } .result-meta { color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
@media (max-width: 700px) { .search-form { grid-template-columns: 1fr; } .advanced { grid-column: auto; } } @media (max-width: 700px) { .search-form { grid-template-columns: 1fr; } .advanced { grid-column: auto; } }
@@ -181,11 +181,12 @@ async function chooseDefaultModel(provider: ProviderConfig, event: Event) {
.settings-page { max-width: 1120px; margin: 0 auto; } .settings-page { max-width: 1120px; margin: 0 auto; }
.settings-section { display: grid; gap: var(--space-md); } .settings-section { display: grid; gap: var(--space-md); }
.settings-section h2 { margin-bottom: var(--space-sm); } .settings-section h2 { margin-bottom: var(--space-sm); }
.setting-row { display: flex; align-items: center; justify-content: space-between; gap: var(--space-xl); min-height: 54px; padding: var(--space-sm) 0; border-bottom: 1px solid var(--color-border-subtle); } .setting-row { display: flex; align-items: center; justify-content: space-between; gap: var(--space-xl); min-height: 58px; padding: var(--space-sm) var(--space-md); border-bottom: 1px solid var(--color-border-subtle); border-radius: var(--radius-md); transition: background-color var(--motion-fast); }
.setting-row:hover { background: var(--color-background-secondary); }
.setting-row small { display: block; color: var(--color-text-tertiary); }.short { width: min(220px, 45%); } .setting-row small { display: block; color: var(--color-text-tertiary); }.short { width: min(220px, 45%); }
.section-head { display: flex; align-items: center; justify-content: space-between; margin-bottom: var(--space-lg); } .section-head { display: flex; align-items: center; justify-content: space-between; margin-bottom: var(--space-lg); }
.provider-list { display: grid; gap: var(--space-md); }.provider-card { display: flex; align-items: center; justify-content: space-between; gap: var(--space-xl); }.provider-main { min-width: 0; flex: 1; }.provider-card p, .provider-card .tag-list { margin-top: var(--space-sm); } .provider-list { display: grid; gap: var(--space-md); }.provider-card { display: flex; align-items: center; justify-content: space-between; gap: var(--space-xl); }.provider-main { min-width: 0; flex: 1; }.provider-card p, .provider-card .tag-list { margin-top: var(--space-sm); }
.model-picker { display: flex; align-items: center; gap: var(--space-sm); margin-top: var(--space-md); }.model-picker label { white-space: nowrap; font-weight: 600; }.model-picker .select { width: min(360px, 100%); }.provider-actions { flex-wrap: wrap; justify-content: flex-end; }.error-text { color: var(--color-danger, #d33); } .model-picker { display: flex; align-items: center; gap: var(--space-sm); margin-top: var(--space-md); }.model-picker label { white-space: nowrap; font-weight: 600; }.model-picker .select { width: min(360px, 100%); }.provider-actions { flex-wrap: wrap; justify-content: flex-end; }.error-text { color: var(--color-error); }
.test-result { color: var(--color-info); }.index-summary, .diagnostic-grid { display: grid; grid-template-columns: repeat(3, 1fr); gap: var(--space-md); }.index-summary > div { padding: var(--space-lg); border-radius: var(--radius-md); background: var(--color-background-secondary); }.index-summary strong, .index-summary small { display: block; }.index-summary strong { font-size: var(--font-size-3xl); } .test-result { color: var(--color-info); }.index-summary, .diagnostic-grid { display: grid; grid-template-columns: repeat(3, 1fr); gap: var(--space-md); }.index-summary > div { padding: var(--space-lg); border-radius: var(--radius-md); background: var(--color-background-secondary); }.index-summary strong, .index-summary small { display: block; }.index-summary strong { font-size: var(--font-size-3xl); }
.section-description { margin-top: calc(-1 * var(--space-md)); }.diagnostic-grid { grid-template-columns: repeat(2, 1fr); }.diagnostic-grid h3 { margin: var(--space-md) 0 var(--space-xs); }.diagnostic-actions { margin-top: var(--space-md); } .section-description { margin-top: calc(-1 * var(--space-md)); }.diagnostic-grid { grid-template-columns: repeat(2, 1fr); }.diagnostic-grid h3 { margin: var(--space-md) 0 var(--space-xs); }.diagnostic-actions { margin-top: var(--space-md); }
@media (max-width: 700px) { .provider-card, .setting-row, .model-picker { align-items: flex-start; flex-direction: column; }.short, .model-picker .select { width: 100%; }.index-summary, .diagnostic-grid { grid-template-columns: 1fr; }.provider-actions { justify-content: flex-start; } } @media (max-width: 700px) { .provider-card, .setting-row, .model-picker { align-items: flex-start; flex-direction: column; }.short, .model-picker .select { width: 100%; }.index-summary, .diagnostic-grid { grid-template-columns: 1fr; }.provider-actions { justify-content: flex-start; } }
+4 -3
View File
@@ -50,10 +50,11 @@ async function remove(task: TaskItem) {
</template> </template>
<style scoped> <style scoped>
.task-list { display: grid; gap: var(--space-md); } .task-list { display: grid; gap: var(--space-md); width: min(100%, 980px); margin-inline: auto; }
.task-card { display: grid; grid-template-columns: auto 1fr auto; align-items: center; gap: var(--space-md); } .task-card { display: grid; grid-template-columns: auto 1fr auto; align-items: center; gap: var(--space-md); }
.status-check { width: 26px; height: 26px; border: 2px solid var(--color-border-default); border-radius: var(--radius-full); } .status-check { width: 28px; height: 28px; border: 2px solid var(--color-border-default); border-radius: var(--radius-full); transition: border-color var(--motion-fast), background-color var(--motion-fast), color var(--motion-fast), transform var(--motion-fast); }
.status-check.done { border-color: var(--color-success); background: var(--color-success); color: white; } .status-check:hover { border-color: var(--color-success); transform: scale(1.06); }
.status-check.done { border-color: var(--color-success); background: var(--color-success); color: white; box-shadow: 0 3px 10px color-mix(in srgb, var(--color-success) 24%, transparent); }
.task-title { display: flex; align-items: center; flex-wrap: wrap; gap: var(--space-sm); } .task-title { display: flex; align-items: center; flex-wrap: wrap; gap: var(--space-sm); }
.task-content p { margin: var(--space-xs) 0; } .task-content p { margin: var(--space-xs) 0; }
.task-content .subtle { display: flex; flex-wrap: wrap; gap: var(--space-md); } .task-content .subtle { display: flex; flex-wrap: wrap; gap: var(--space-md); }
+27 -2
View File
@@ -1,6 +1,15 @@
<script setup lang="ts"> <script setup lang="ts">
import { computed } from 'vue'
import MarkdownContent from '@/components/common/MarkdownContent.vue'
import { useThemeStore } from '@/stores/theme' import { useThemeStore } from '@/stores/theme'
const themeStore = useThemeStore() const themeStore = useThemeStore()
const shikiPreview = `\`\`\`typescript
const notes = await search('本地优先')
\`\`\``
const codeThemeLabel = computed(() => themeStore.resolvedCodeBlockTheme === 'github-dark'
? 'Shiki · GitHub Dark'
: 'Shiki · GitHub Light')
</script> </script>
<template> <template>
@@ -13,7 +22,20 @@ const themeStore = useThemeStore()
<p class="subtle">v{{ theme.version }} · {{ theme.builtin ? '内置主题' : theme.author }}</p> <p class="subtle">v{{ theme.version }} · {{ theme.builtin ? '内置主题' : theme.author }}</p>
</button> </button>
</div> </div>
<div class="panel preference-panel"><h2 class="panel-title">编辑器外观</h2><div class="form-grid"><div class="field"><label>字号{{ themeStore.fontEditorSize }}px</label><input v-model.number="themeStore.fontEditorSize" type="range" min="12" max="24" /></div><div class="field"><label>行高{{ themeStore.lineHeight }}</label><input v-model.number="themeStore.lineHeight" type="range" min="1.2" max="2.2" step="0.1" /></div><div class="field"><label>字体</label><select v-model="themeStore.fontEditorFamily" class="select"><option value="system-ui">系统字体</option><option value="serif">衬线字体</option><option value="var(--font-ui-mono)">等宽字体</option></select></div></div><div class="editor-preview" :style="{ fontSize: `${themeStore.fontEditorSize}px`, lineHeight: themeStore.lineHeight, fontFamily: themeStore.fontEditorFamily }"><h3>主题预览</h3><p>知识的价值不只在于保存更在于被重新发现和使用</p><code>const notes = await search('本地优先')</code></div></div> <div class="panel preference-panel">
<h2 class="panel-title">编辑器外观</h2>
<div class="form-grid">
<div class="field"><label>字号{{ themeStore.fontEditorSize }}px</label><input v-model.number="themeStore.fontEditorSize" type="range" min="12" max="24" /></div>
<div class="field"><label>行高{{ themeStore.lineHeight }}</label><input v-model.number="themeStore.lineHeight" type="range" min="1.2" max="2.2" step="0.1" /></div>
<div class="field"><label>字体</label><select v-model="themeStore.fontEditorFamily" class="select"><option value="system-ui">系统字体</option><option value="serif">衬线字体</option><option value="var(--font-ui-mono)">等宽字体</option></select></div>
<div class="field"><label>代码块样式</label><select v-model="themeStore.codeBlockTheme" class="select"><option value="auto">跟随主题</option><option value="github-light">GitHub Light</option><option value="github-dark">GitHub Dark</option></select><small>Markdown 渲染使用对应的 Shiki GitHub 主题</small></div>
</div>
<div class="editor-preview" :style="{ fontSize: `${themeStore.fontEditorSize}px`, lineHeight: themeStore.lineHeight, fontFamily: themeStore.fontEditorFamily }">
<div class="preview-heading"><h3>主题预览</h3><span class="badge info">{{ codeThemeLabel }}</span></div>
<p>知识的价值不只在于保存更在于被重新发现和使用</p>
<MarkdownContent class="code-theme-preview" :source="shikiPreview" />
</div>
</div>
</section> </section>
</template> </template>
@@ -28,5 +50,8 @@ const themeStore = useThemeStore()
.theme-info { display: flex; justify-content: space-between; gap: var(--space-md); } .theme-info { display: flex; justify-content: space-between; gap: var(--space-md); }
.preference-panel { display: grid; gap: var(--space-xl); } .preference-panel { display: grid; gap: var(--space-xl); }
.editor-preview { padding: var(--space-xl); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); background: var(--color-background-secondary); } .editor-preview { padding: var(--space-xl); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); background: var(--color-background-secondary); }
.editor-preview p { margin: var(--space-sm) 0; }.editor-preview code { color: var(--color-accent-primary); } .editor-preview p { margin: var(--space-sm) 0; }
.preview-heading { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); }
.field small { color: var(--color-text-tertiary); }
.code-theme-preview { margin-top: var(--space-md); }
</style> </style>
+36 -111
View File
@@ -4,7 +4,7 @@ import { useRouter } from 'vue-router'
import { useWorkspaceStore } from '@/stores/workspace' import { useWorkspaceStore } from '@/stores/workspace'
import { useThemeStore } from '@/stores/theme' import { useThemeStore } from '@/stores/theme'
import { useSettingsStore } from '@/stores/settings' import { useSettingsStore } from '@/stores/settings'
import { ArrowRight, Document, Folder, FolderOpened, Moon, Plus, Sunny } from '@element-plus/icons-vue' import { ArrowRight, Document, Folder, FolderOpened, Moon, Sunny } from '@element-plus/icons-vue'
import AppIcon from '@/components/common/AppIcon.vue' import AppIcon from '@/components/common/AppIcon.vue'
const router = useRouter() const router = useRouter()
@@ -13,17 +13,19 @@ const themeStore = useThemeStore()
const settingsStore = useSettingsStore() const settingsStore = useSettingsStore()
const isLoading = ref(false) const isLoading = ref(false)
const showCreateDialog = ref(false)
const newVaultName = ref('')
const newVaultPath = ref('')
const aiCoreStatus = ref<'checking' | 'running' | 'stopped'>('checking') const aiCoreStatus = ref<'checking' | 'running' | 'stopped'>('checking')
onMounted(async () => { onMounted(async () => {
await Promise.all([workspaceStore.loadRecentVaults(), settingsStore.loadDiagnostics()]) await Promise.allSettled([workspaceStore.loadRecentVaults(), settingsStore.loadDiagnostics()])
const lastVaultPath = localStorage.getItem('last-vault-path') const lastVaultPath = localStorage.getItem('last-vault-path')
if (settingsStore.restoreLastVault && lastVaultPath) { if (settingsStore.restoreLastVault && lastVaultPath) {
try {
await openVault(lastVaultPath) await openVault(lastVaultPath)
return return
} catch {
// Mock Vault
localStorage.removeItem('last-vault-path')
}
} }
setTimeout(() => { setTimeout(() => {
aiCoreStatus.value = settingsStore.aiCoreStatus === 'running' ? 'running' : 'stopped' aiCoreStatus.value = settingsStore.aiCoreStatus === 'running' ? 'running' : 'stopped'
@@ -41,24 +43,8 @@ async function openVault(path: string) {
} }
async function openFolderPicker() { async function openFolderPicker() {
// In Tauri this would use the native dialog const configured = workspaceStore.recentVaults[0]
// For web dev, simulate if (configured) await openVault(configured.path)
const path = prompt('请输入 Vault 路径(开发模式)', '/Users/demo/Documents/MyVault')
if (path) {
await openVault(path)
}
}
async function createVault() {
if (!newVaultName.value || !newVaultPath.value) return
isLoading.value = true
try {
await workspaceStore.createVault(newVaultPath.value, newVaultName.value)
router.push('/workspace')
} finally {
isLoading.value = false
showCreateDialog.value = false
}
} }
</script> </script>
@@ -68,13 +54,13 @@ async function createVault() {
<div class="entry-container"> <div class="entry-container">
<div class="brand-section"> <div class="brand-section">
<div class="logo"><AppIcon :icon="Document" :size="56" /></div> <div class="logo"><AppIcon :icon="Document" :size="56" /></div>
<h1 class="app-title">知笔知己</h1> <h1 class="app-title">NotesAgent</h1>
<p class="app-subtitle">本地优先的 AI 笔记软件</p> <p class="app-subtitle">本地优先的 AI 笔记软件</p>
</div> </div>
<div class="vault-card"> <div class="vault-card">
<h2 class="card-title">选择知识库</h2> <h2 class="card-title">选择知识库</h2>
<p class="card-desc">选择一个本地 Vault 开始你的知识之旅</p> <p class="card-desc">Web 联调模式连接 AI Core 当前配置的 Vault</p>
<div v-if="workspaceStore.recentVaults.length" class="recent-vaults"> <div v-if="workspaceStore.recentVaults.length" class="recent-vaults">
<div class="section-label">最近打开</div> <div class="section-label">最近打开</div>
@@ -97,11 +83,8 @@ async function createVault() {
</div> </div>
<div class="actions"> <div class="actions">
<button class="btn btn-primary" @click="openFolderPicker" :disabled="isLoading"> <button class="btn btn-primary" @click="openFolderPicker" :disabled="isLoading || !workspaceStore.recentVaults.length">
<AppIcon :icon="FolderOpened" /> 打开本地 Vault <AppIcon :icon="FolderOpened" /> 打开后端 Vault
</button>
<button class="btn btn-secondary" @click="showCreateDialog = true" :disabled="isLoading">
<AppIcon :icon="Plus" /> 创建新 Vault
</button> </button>
</div> </div>
@@ -122,24 +105,6 @@ async function createVault() {
</div> </div>
</div> </div>
<!-- Create Vault Dialog -->
<div v-if="showCreateDialog" class="dialog-overlay" @click.self="showCreateDialog = false">
<div class="dialog">
<h3>创建新 Vault</h3>
<div class="form-group">
<label>Vault 名称</label>
<input v-model="newVaultName" type="text" placeholder="我的知识库" />
</div>
<div class="form-group">
<label>存储路径</label>
<input v-model="newVaultPath" type="text" placeholder="/path/to/vault" />
</div>
<div class="dialog-actions">
<button class="btn btn-secondary" @click="showCreateDialog = false">取消</button>
<button class="btn btn-primary" @click="createVault" :disabled="!newVaultName || !newVaultPath">创建</button>
</div>
</div>
</div>
</div> </div>
</template> </template>
@@ -161,7 +126,7 @@ async function createVault() {
background: background:
radial-gradient(circle at 20% 30%, var(--color-accent-soft) 0%, transparent 50%), radial-gradient(circle at 20% 30%, var(--color-accent-soft) 0%, transparent 50%),
radial-gradient(circle at 80% 70%, var(--color-info-soft) 0%, transparent 50%); radial-gradient(circle at 80% 70%, var(--color-info-soft) 0%, transparent 50%);
opacity: 0.5; opacity: 0.62;
} }
.entry-container { .entry-container {
@@ -173,6 +138,7 @@ async function createVault() {
gap: 32px; gap: 32px;
max-width: 480px; max-width: 480px;
width: 90%; width: 90%;
animation: entry-in var(--motion-slow) both;
} }
.brand-section { .brand-section {
@@ -180,8 +146,16 @@ async function createVault() {
} }
.logo { .logo {
font-size: 64px; display: inline-grid;
margin-bottom: 12px; place-items: center;
width: 84px;
height: 84px;
margin-bottom: 14px;
border: 1px solid color-mix(in srgb, var(--color-accent-primary) 18%, transparent);
border-radius: 24px;
background: var(--color-surface-primary);
color: var(--color-accent-primary);
box-shadow: var(--shadow-lg);
} }
.app-title { .app-title {
@@ -206,7 +180,7 @@ async function createVault() {
border: 1px solid var(--color-border-default); border: 1px solid var(--color-border-default);
border-radius: var(--radius-xl); border-radius: var(--radius-xl);
padding: var(--space-2xl); padding: var(--space-2xl);
box-shadow: var(--shadow-lg); box-shadow: var(--shadow-xl);
} }
.card-title { .card-title {
@@ -247,11 +221,13 @@ async function createVault() {
border-radius: var(--radius-md); border-radius: var(--radius-md);
cursor: pointer; cursor: pointer;
text-align: left; text-align: left;
transition: all var(--motion-fast); transition: background-color var(--motion-fast), border-color var(--motion-fast), box-shadow var(--motion-fast), transform var(--motion-fast);
&:hover { &:hover {
background: var(--color-accent-soft); background: var(--color-accent-soft);
border-color: var(--color-accent-secondary); border-color: var(--color-accent-secondary);
box-shadow: var(--shadow-sm);
transform: translateY(-1px);
} }
&:disabled { &:disabled {
@@ -288,8 +264,11 @@ async function createVault() {
.vault-arrow { .vault-arrow {
color: var(--color-text-tertiary); color: var(--color-text-tertiary);
font-size: 20px; font-size: 20px;
transition: color var(--motion-fast), transform var(--motion-fast);
} }
.vault-item:hover .vault-arrow { color: var(--color-accent-primary); transform: translateX(3px); }
.actions { .actions {
display: flex; display: flex;
flex-direction: column; flex-direction: column;
@@ -307,7 +286,7 @@ async function createVault() {
font-size: 14px; font-size: 14px;
font-weight: 500; font-weight: 500;
cursor: pointer; cursor: pointer;
transition: all var(--motion-fast); transition: background-color var(--motion-fast), border-color var(--motion-fast), box-shadow var(--motion-fast), transform var(--motion-fast);
border: 1px solid transparent; border: 1px solid transparent;
&:disabled { &:disabled {
@@ -321,6 +300,8 @@ async function createVault() {
&:hover:not(:disabled) { &:hover:not(:disabled) {
background: var(--color-accent-primary-hover); background: var(--color-accent-primary-hover);
transform: translateY(-1px);
box-shadow: 0 7px 18px color-mix(in srgb, var(--color-accent-primary) 25%, transparent);
} }
} }
@@ -392,61 +373,5 @@ async function createVault() {
} }
} }
.dialog-overlay { @keyframes entry-in { from { opacity: 0; transform: translateY(8px); } to { opacity: 1; transform: translateY(0); } }
position: fixed;
inset: 0;
background: var(--color-background-overlay);
display: flex;
align-items: center;
justify-content: center;
z-index: var(--z-modal);
}
.dialog {
background: var(--color-surface-primary);
border-radius: var(--radius-lg);
padding: var(--space-xl);
width: 90%;
max-width: 400px;
box-shadow: var(--shadow-xl);
}
.dialog h3 {
margin: 0 0 var(--space-lg) 0;
font-size: 18px;
}
.form-group {
margin-bottom: var(--space-md);
label {
display: block;
font-size: 13px;
color: var(--color-text-secondary);
margin-bottom: var(--space-xs);
}
input {
width: 100%;
padding: 8px 12px;
background: var(--color-background-secondary);
border: 1px solid var(--color-border-default);
border-radius: var(--radius-md);
font-size: 14px;
color: var(--color-text-primary);
outline: none;
transition: border-color var(--motion-fast);
&:focus {
border-color: var(--color-border-focus);
}
}
}
.dialog-actions {
display: flex;
justify-content: flex-end;
gap: var(--space-sm);
margin-top: var(--space-lg);
}
</style> </style>
@@ -1,11 +1,12 @@
// @vitest-environment happy-dom // @vitest-environment happy-dom
import { afterEach, beforeEach, describe, expect, it } from 'vitest' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import { mount, type VueWrapper } from '@vue/test-utils' import { mount, type VueWrapper } from '@vue/test-utils'
import { createPinia, setActivePinia } from 'pinia' import { createPinia, setActivePinia } from 'pinia'
import { createMemoryHistory, createRouter } from 'vue-router' import { createMemoryHistory, createRouter } from 'vue-router'
import FileTreePanel from './FileTreePanel.vue' import FileTreePanel from './FileTreePanel.vue'
import { useEditorStore } from '@/stores/editor' import { useEditorStore } from '@/stores/editor'
import { useWorkspaceStore } from '@/stores/workspace' import { useWorkspaceStore } from '@/stores/workspace'
import * as workspaceService from '@/services/workspaceService'
let wrapper: VueWrapper | null = null let wrapper: VueWrapper | null = null
@@ -21,12 +22,29 @@ async function waitForPath(path: string) {
beforeEach(() => { beforeEach(() => {
localStorage.clear() localStorage.clear()
setActivePinia(createPinia()) setActivePinia(createPinia())
vi.spyOn(workspaceService, 'openVault').mockResolvedValue({ vault_id: 'default', path: 'C:/vault', name: 'vault' })
vi.spyOn(workspaceService, 'getFileTree').mockResolvedValue([
{
id: 'folder-data', name: '数据结构', path: '/数据结构', type: 'folder', is_open: true,
children: [
{ id: 'note-rbt', note_id: 'note-rbt', name: '红黑树.md', path: '/数据结构/红黑树.md', type: 'file' },
{ id: 'note-bst', note_id: 'note-bst', name: '二叉搜索树.md', path: '/数据结构/二叉搜索树.md', type: 'file' },
],
},
])
vi.spyOn(workspaceService, 'readFileContent').mockImplementation(async (path) =>
path.includes('红黑树') ? '# 红黑树\n' : '# 二叉搜索树\n'
)
vi.spyOn(workspaceService, 'getNoteId').mockImplementation(async (path) =>
path.includes('红黑树') ? 'note-rbt' : 'note-bst'
)
}) })
afterEach(() => { afterEach(() => {
wrapper?.unmount() wrapper?.unmount()
wrapper = null wrapper = null
document.body.innerHTML = '' document.body.innerHTML = ''
vi.restoreAllMocks()
}) })
describe('FileTreePanel file switching', () => { describe('FileTreePanel file switching', () => {
@@ -40,7 +58,7 @@ describe('FileTreePanel file switching', () => {
const workspaceStore = useWorkspaceStore() const workspaceStore = useWorkspaceStore()
const editorStore = useEditorStore() const editorStore = useEditorStore()
await workspaceStore.openVault('/mock-vault') await workspaceStore.openVault('C:/vault')
wrapper = mount(FileTreePanel, { attachTo: document.body, global: { plugins: [router] } }) wrapper = mount(FileTreePanel, { attachTo: document.body, global: { plugins: [router] } })
const findNode = (name: string) => wrapper!.findAll('.tree-node').find((node) => node.text().includes(name))! const findNode = (name: string) => wrapper!.findAll('.tree-node').find((node) => node.text().includes(name))!
@@ -48,10 +66,41 @@ describe('FileTreePanel file switching', () => {
await waitForPath('/数据结构/红黑树.md') await waitForPath('/数据结构/红黑树.md')
expect(workspaceStore.activeFilePath).toBe('/数据结构/红黑树.md') expect(workspaceStore.activeFilePath).toBe('/数据结构/红黑树.md')
expect(editorStore.content).toContain('# 红黑树') expect(editorStore.content).toContain('# 红黑树')
expect(editorStore.currentNoteId).toBe('note-rbt')
await findNode('二叉搜索树.md').trigger('click') await findNode('二叉搜索树.md').trigger('click')
await waitForPath('/数据结构/二叉搜索树.md') await waitForPath('/数据结构/二叉搜索树.md')
expect(workspaceStore.activeFilePath).toBe('/数据结构/二叉搜索树.md') expect(workspaceStore.activeFilePath).toBe('/数据结构/二叉搜索树.md')
expect(editorStore.content).toContain('# 二叉搜索树') expect(editorStore.content).toContain('# 二叉搜索树')
expect(editorStore.currentNoteId).toBe('note-bst')
})
it('creates a Markdown note inside the selected folder', async () => {
const router = createRouter({
history: createMemoryHistory(),
routes: [{ path: '/workspace', component: { template: '<div />' } }],
})
await router.push('/workspace')
await router.isReady()
const workspaceStore = useWorkspaceStore()
await workspaceStore.openVault('C:/vault')
const createFile = vi.spyOn(workspaceService, 'createFile').mockResolvedValue({
id: 'note-new', note_id: 'note-new', name: '新笔记.md',
path: '/数据结构/新笔记.md', type: 'file',
})
wrapper = mount(FileTreePanel, { attachTo: document.body, global: { plugins: [router] } })
await wrapper.findAll('.tree-node').find((node) => node.text().includes('数据结构'))!.trigger('click')
await wrapper.get('button[aria-label="新建笔记"]').trigger('click')
await wrapper.get('.new-item input').setValue('新笔记')
await wrapper.get('.new-item').trigger('submit')
await waitForPath('/数据结构/新笔记.md')
await vi.waitFor(() => {
expect(workspaceStore.activeFilePath).toBe('/数据结构/新笔记.md')
})
expect(createFile).toHaveBeenCalledWith('/数据结构', '新笔记.md', '# 新笔记\n\n')
expect(wrapper.findAll('.tree-node').some((node) => node.classes().includes('active') && node.text().includes('新笔记.md'))).toBe(true)
}) })
}) })
@@ -1,5 +1,5 @@
<script setup lang="ts"> <script setup lang="ts">
import { ref } from 'vue' import { ref, watch } from 'vue'
import { useRouter } from 'vue-router' import { useRouter } from 'vue-router'
import type { FileNode } from '@/contracts' import type { FileNode } from '@/contracts'
import * as workspaceService from '@/services/workspaceService' import * as workspaceService from '@/services/workspaceService'
@@ -15,9 +15,19 @@ const router = useRouter()
const newItemType = ref<'file' | 'folder' | null>(null) const newItemType = ref<'file' | 'folder' | null>(null)
const newItemName = ref('') const newItemName = ref('')
const parentPath = ref('/') const parentPath = ref('/')
const selectedTreePath = ref(workspaceStore.activeFilePath ?? '/')
const selectedFolderPath = ref(
workspaceStore.activeFilePath ? containingFolder(workspaceStore.activeFilePath) : '/',
)
const contextTarget = ref<FileNode | null>(null) const contextTarget = ref<FileNode | null>(null)
const contextMenuPosition = ref({ x: 0, y: 0 }) const contextMenuPosition = ref({ x: 0, y: 0 })
watch(() => workspaceStore.activeFilePath, (path) => {
if (!path) return
selectedTreePath.value = path
selectedFolderPath.value = containingFolder(path)
})
function beginCreate(type: 'file' | 'folder', parent = '/') { function beginCreate(type: 'file' | 'folder', parent = '/') {
newItemType.value = type newItemType.value = type
newItemName.value = '' newItemName.value = ''
@@ -31,19 +41,28 @@ async function createItem() {
const name = rawName.endsWith('.md') ? rawName : `${rawName}.md` const name = rawName.endsWith('.md') ? rawName : `${rawName}.md`
const file = await workspaceService.createFile(parentPath.value, name, `# ${rawName}\n\n`) const file = await workspaceService.createFile(parentPath.value, name, `# ${rawName}\n\n`)
workspaceStore.addFileToTree(parentPath.value, file) workspaceStore.addFileToTree(parentPath.value, file)
selectedTreePath.value = file.path
selectedFolderPath.value = parentPath.value
await editorStore.loadFile(file.path) await editorStore.loadFile(file.path)
workspaceStore.openFile(file.path) workspaceStore.openFile(file.path)
await router.push('/workspace') await router.push('/workspace')
} else { } else {
const folder = await workspaceService.createFolder(parentPath.value, rawName) const folder = await workspaceService.createFolder(parentPath.value, rawName)
workspaceStore.addFileToTree(parentPath.value, folder) workspaceStore.addFileToTree(parentPath.value, folder)
selectedTreePath.value = folder.path
selectedFolderPath.value = folder.path
} }
newItemType.value = null newItemType.value = null
newItemName.value = '' newItemName.value = ''
} }
async function openNode(node: FileNode) { async function openNode(node: FileNode) {
if (node.type === 'folder') return workspaceStore.toggleFolder(node.path) selectedTreePath.value = node.path
if (node.type === 'folder') {
selectedFolderPath.value = node.path
return workspaceStore.toggleFolder(node.path)
}
selectedFolderPath.value = containingFolder(node.path)
// //
const previousPath = workspaceStore.activeFilePath const previousPath = workspaceStore.activeFilePath
const wasOpen = workspaceStore.openFiles.includes(node.path) const wasOpen = workspaceStore.openFiles.includes(node.path)
@@ -61,6 +80,8 @@ async function openNode(node: FileNode) {
function openContextMenu(event: MouseEvent, node: FileNode) { function openContextMenu(event: MouseEvent, node: FileNode) {
event.preventDefault() event.preventDefault()
event.stopPropagation() event.stopPropagation()
selectedTreePath.value = node.path
selectedFolderPath.value = node.type === 'folder' ? node.path : containingFolder(node.path)
contextTarget.value = node contextTarget.value = node
contextMenuPosition.value = { x: event.clientX, y: event.clientY } contextMenuPosition.value = { x: event.clientX, y: event.clientY }
} }
@@ -79,6 +100,12 @@ async function renameTarget() {
await workspaceService.renameFile(oldPath, normalizedName) await workspaceService.renameFile(oldPath, normalizedName)
workspaceStore.renamePath(oldPath, newPath, normalizedName) workspaceStore.renamePath(oldPath, newPath, normalizedName)
editorStore.renameFilePath(oldPath, newPath) editorStore.renameFilePath(oldPath, newPath)
if (selectedTreePath.value === oldPath || selectedTreePath.value.startsWith(`${oldPath}/`)) {
selectedTreePath.value = `${newPath}${selectedTreePath.value.slice(oldPath.length)}`
}
if (selectedFolderPath.value === oldPath || selectedFolderPath.value.startsWith(`${oldPath}/`)) {
selectedFolderPath.value = `${newPath}${selectedFolderPath.value.slice(oldPath.length)}`
}
} }
closeContextMenu() closeContextMenu()
} }
@@ -90,19 +117,28 @@ async function deleteTarget() {
await workspaceService.deleteFile(node.path) await workspaceService.deleteFile(node.path)
const activeWasRemoved = workspaceStore.closePath(node.path) const activeWasRemoved = workspaceStore.closePath(node.path)
workspaceStore.removeFromTree(node.path) workspaceStore.removeFromTree(node.path)
if (selectedTreePath.value === node.path || selectedTreePath.value.startsWith(`${node.path}/`)) {
selectedTreePath.value = containingFolder(node.path)
selectedFolderPath.value = selectedTreePath.value
}
if (activeWasRemoved) { if (activeWasRemoved) {
editorStore.closeFile() editorStore.closeFile()
if (workspaceStore.activeFilePath) await editorStore.loadFile(workspaceStore.activeFilePath) if (workspaceStore.activeFilePath) await editorStore.loadFile(workspaceStore.activeFilePath)
} }
closeContextMenu() closeContextMenu()
} }
function containingFolder(path: string): string {
const separator = path.lastIndexOf('/')
return separator > 0 ? path.slice(0, separator) : '/'
}
</script> </script>
<template> <template>
<section class="file-tree-panel" @click="closeContextMenu"> <section class="file-tree-panel" @click="closeContextMenu">
<div class="toolbar"> <div class="toolbar">
<button type="button" title="新建笔记" aria-label="新建笔记" @click.stop="beginCreate('file')"><AppIcon :icon="DocumentAdd" /></button> <button type="button" title="新建笔记" aria-label="新建笔记" @click.stop="beginCreate('file', selectedFolderPath)"><AppIcon :icon="DocumentAdd" /></button>
<button type="button" title="新建文件夹" aria-label="新建文件夹" @click.stop="beginCreate('folder')"><AppIcon :icon="FolderAdd" /></button> <button type="button" title="新建文件夹" aria-label="新建文件夹" @click.stop="beginCreate('folder', selectedFolderPath)"><AppIcon :icon="FolderAdd" /></button>
</div> </div>
<form v-if="newItemType" class="new-item" @submit.prevent="createItem"> <form v-if="newItemType" class="new-item" @submit.prevent="createItem">
<input v-model="newItemName" :placeholder="newItemType === 'file' ? '笔记名称' : '文件夹名称'" autofocus /> <input v-model="newItemName" :placeholder="newItemType === 'file' ? '笔记名称' : '文件夹名称'" autofocus />
@@ -111,7 +147,7 @@ async function deleteTarget() {
</form> </form>
<div class="tree"> <div class="tree">
<FileTreeNode v-for="node in workspaceStore.fileTree" :key="node.id" :node="node" <FileTreeNode v-for="node in workspaceStore.fileTree" :key="node.id" :node="node"
:active-path="workspaceStore.activeFilePath" @open="openNode" @context-menu="openContextMenu" /> :active-path="selectedTreePath" @open="openNode" @context-menu="openContextMenu" />
</div> </div>
<Teleport to="body"> <Teleport to="body">
<div v-if="contextTarget" class="context-menu" <div v-if="contextTarget" class="context-menu"
@@ -133,5 +169,5 @@ button:hover { background: var(--color-background-secondary); }
.tree { padding: var(--space-xs); } .tree { padding: var(--space-xs); }
.context-menu { position: fixed; z-index: 1000; display: grid; min-width: 130px; padding: var(--space-xs); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); background: var(--color-background-primary); box-shadow: var(--shadow-md); } .context-menu { position: fixed; z-index: 1000; display: grid; min-width: 130px; padding: var(--space-xs); border: 1px solid var(--color-border-default); border-radius: var(--radius-md); background: var(--color-background-primary); box-shadow: var(--shadow-md); }
.context-menu button { text-align: left; } .context-menu button { text-align: left; }
.context-menu .danger { color: var(--color-danger, #d33); } .context-menu .danger { color: var(--color-error); }
</style> </style>
@@ -17,16 +17,16 @@ afterEach(() => {
wrapper = null wrapper = null
}) })
describe('WorkspaceView initial file', () => { describe('WorkspaceView empty state', () => {
it('does not overwrite a file selected while the welcome note is loading', async () => { it('does not fabricate a Mock welcome note when no backend file is selected', async () => {
const workspaceStore = useWorkspaceStore() const workspaceStore = useWorkspaceStore()
wrapper = mount(WorkspaceView, { wrapper = mount(WorkspaceView, {
global: { stubs: { EditorHeader: true, EditorPane: true } }, global: { stubs: { EditorHeader: true, EditorPane: true } },
}) })
workspaceStore.openFile('/数据结构/红黑树.md')
await new Promise((resolve) => setTimeout(resolve, 0)) await new Promise((resolve) => setTimeout(resolve, 0))
expect(workspaceStore.activeFilePath).toBe('/数据结构/红黑树.md') expect(workspaceStore.activeFilePath).toBeNull()
expect(wrapper.find('.empty-workspace').exists()).toBe(true)
}) })
}) })
@@ -1,28 +1,11 @@
<script setup lang="ts"> <script setup lang="ts">
import { onMounted } from 'vue'
import { useWorkspaceStore } from '@/stores/workspace' import { useWorkspaceStore } from '@/stores/workspace'
import { useEditorStore } from '@/stores/editor'
import EditorHeader from '@/features/editor/EditorHeader.vue' import EditorHeader from '@/features/editor/EditorHeader.vue'
import EditorPane from '@/features/editor/EditorPane.vue' import EditorPane from '@/features/editor/EditorPane.vue'
import { EditPen } from '@element-plus/icons-vue' import { EditPen } from '@element-plus/icons-vue'
import AppIcon from '@/components/common/AppIcon.vue' import AppIcon from '@/components/common/AppIcon.vue'
const workspaceStore = useWorkspaceStore() const workspaceStore = useWorkspaceStore()
const editorStore = useEditorStore()
onMounted(() => {
if (!workspaceStore.fileTree.length && workspaceStore.hasVault) {
// Already loaded
}
if (!workspaceStore.activeFilePath && workspaceStore.fileTree.length === 0) {
void editorStore.loadFile('/欢迎使用知笔知己.md').then(() => {
//
if (!workspaceStore.activeFilePath && editorStore.currentFilePath === '/欢迎使用知笔知己.md') {
workspaceStore.openFile('/欢迎使用知笔知己.md')
}
})
}
})
</script> </script>
<template> <template>
@@ -59,6 +42,11 @@ onMounted(() => {
.empty-content { .empty-content {
text-align: center; text-align: center;
padding: var(--space-3xl);
border: 1px dashed var(--color-border-default);
border-radius: var(--radius-xl);
background: var(--color-background-secondary);
animation: workspace-empty-in var(--motion-normal) both;
h2 { h2 {
font-size: 18px; font-size: 18px;
@@ -75,4 +63,6 @@ onMounted(() => {
font-size: 48px; font-size: 48px;
opacity: 0.5; opacity: 0.5;
} }
@keyframes workspace-empty-in { from { opacity: 0; transform: translateY(5px); } to { opacity: 1; transform: translateY(0); } }
</style> </style>
+12 -8
View File
@@ -44,11 +44,17 @@ const routes = [
component: () => import('@/features/skills/SkillsView.vue'), component: () => import('@/features/skills/SkillsView.vue'),
meta: { title: 'Skill 管理', requiresVault: true }, meta: { title: 'Skill 管理', requiresVault: true },
}, },
{
path: '/extensions/mcp',
name: 'mcp-servers',
component: () => import('@/features/mcp/McpServersView.vue'),
meta: { title: 'MCP 服务器', requiresVault: true },
},
{ {
path: '/extensions/plugins', path: '/extensions/plugins',
name: 'plugins', name: 'plugins',
component: () => import('@/features/plugins/PluginsView.vue'), component: () => import('@/features/plugins/PluginsView.vue'),
meta: { title: 'Plugin 管理', requiresVault: true }, meta: { title: 'Plugin 与 MCP', requiresVault: true },
}, },
{ {
path: '/themes', path: '/themes',
@@ -69,21 +75,19 @@ const router = createRouter({
routes, routes,
}) })
router.beforeEach((to, _from, next) => { router.beforeEach((to) => {
const workspaceStore = useWorkspaceStore() const workspaceStore = useWorkspaceStore()
if (to.meta.requiresVault && !workspaceStore.hasVault) { if (to.meta.requiresVault && !workspaceStore.hasVault) {
next({ path: '/' }) return { path: '/' }
return
} }
if (to.path === '/' && workspaceStore.hasVault) { if (to.path === '/' && workspaceStore.hasVault) {
next({ path: '/workspace' }) return { path: '/workspace' }
return
} }
next() return true
}) })
router.afterEach((to) => { router.afterEach((to) => {
const baseTitle = '知笔知己' const baseTitle = 'NotesAgent'
const title = to.meta.title as string | undefined const title = to.meta.title as string | undefined
document.title = title ? `${title} · ${baseTitle}` : baseTitle document.title = title ? `${title} · ${baseTitle}` : baseTitle
}) })
+14 -3
View File
@@ -1,8 +1,9 @@
import apiClient from './apiClient' import apiClient from './apiClient'
import { SseClient } from './sseClient' import { SseClient } from './sseClient'
import type { AgentRun, AgentEvent, ApiAgentRun, OperationResponse, PageMeta, ToolDefinition, PermissionRequest } from '@/contracts' import type { AgentRun, AgentEvent, AgentTraceResponse, ApiAgentRun, OperationResponse, PageMeta, ToolDefinition, PermissionRequest } from '@/contracts'
function toAgentRun(run: ApiAgentRun): AgentRun { function toAgentRun(run: ApiAgentRun): AgentRun {
// API 的 token_usage 是累计值,UI 模型预留了输入/输出拆分字段。
return { return {
run_id: run.run_id, run_id: run.run_id,
status: run.status, status: run.status,
@@ -53,6 +54,13 @@ export async function cancelAgentRun(runId: string): Promise<OperationResponse>
return apiClient.post(`/api/agent/runs/${runId}/cancel`) return apiClient.post(`/api/agent/runs/${runId}/cancel`)
} }
export async function getAgentTrace(
runId: string,
params?: { after_sequence?: number; limit?: number },
): Promise<AgentTraceResponse> {
return apiClient.get(`/api/agent/runs/${runId}/trace`, { params })
}
export async function listTools(): Promise<ToolDefinition[]> { export async function listTools(): Promise<ToolDefinition[]> {
const response = await apiClient.get<{ items: ToolDefinition[] }>('/api/tools') const response = await apiClient.get<{ items: ToolDefinition[] }>('/api/tools')
return response.items return response.items
@@ -65,11 +73,14 @@ export function streamAgentEvents(
onError?: (error: Error) => void onError?: (error: Error) => void
onDone?: () => void onDone?: () => void
onOpen?: () => void onOpen?: () => void
} },
afterSequence = -1,
): SseClient { ): SseClient {
// 将通用 SSE 包装成领域事件,Store 无需了解传输层 envelope。
const client = new SseClient({ const client = new SseClient({
url: `/api/agent/runs/${runId}/events`, url: `/api/agent/runs/${runId}/events?after_sequence=${afterSequence}`,
method: 'GET', method: 'GET',
lastEventId: afterSequence >= 0 ? String(afterSequence) : undefined,
onEvent: (eventName, data) => { onEvent: (eventName, data) => {
handlers.onEvent?.({ handlers.onEvent?.({
event: eventName as AgentEvent['event'], event: eventName as AgentEvent['event'],
+2
View File
@@ -1,5 +1,6 @@
import type { ApiError, ErrorResponse } from '@/contracts' import type { ApiError, ErrorResponse } from '@/contracts'
// 所有 HTTP 请求都经过此边界,以统一地址、请求追踪和错误契约。
const BASE_URL = import.meta.env.VITE_API_BASE_URL ?? import.meta.env.VITE_API_BASE ?? '' const BASE_URL = import.meta.env.VITE_API_BASE_URL ?? import.meta.env.VITE_API_BASE ?? ''
export function resolveApiUrl(path: string): string { export function resolveApiUrl(path: string): string {
@@ -63,6 +64,7 @@ async function request<T>(path: string, options: RequestOptions = {}): Promise<T
return resp as unknown as T return resp as unknown as T
} }
// 后端约定返回 ErrorResponse;代理或网关的非 JSON 错误仍降级为 HTTP 状态码。
let errBody: ErrorResponse | null = null let errBody: ErrorResponse | null = null
try { try {
errBody = (await resp.json()) as ErrorResponse errBody = (await resp.json()) as ErrorResponse
+1
View File
@@ -8,6 +8,7 @@ export * as chatService from './chatService'
export * as agentService from './agentService' export * as agentService from './agentService'
export * as skillService from './skillService' export * as skillService from './skillService'
export * as pluginService from './pluginService' export * as pluginService from './pluginService'
export * as mcpServerService from './mcpServerService'
export * as providerService from './providerService' export * as providerService from './providerService'
export * as taskService from './taskService' export * as taskService from './taskService'
export * as indexService from './indexService' export * as indexService from './indexService'
+18
View File
@@ -0,0 +1,18 @@
import apiClient from './apiClient'
import type { McpServer, McpServerInput, McpToolSummary, OperationResponse } from '@/contracts'
const base = '/api/mcp/servers'
export async function listMcpServers(): Promise<McpServer[]> {
return (await apiClient.get<{ items: McpServer[] }>(base)).items
}
export const createMcpServer = (input: McpServerInput) => apiClient.post<McpServer>(base, input)
export const updateMcpServer = (id: string, input: McpServerInput) => apiClient.put<McpServer>(`${base}/${id}`, input)
export const listMcpServerTools = async (id: string) => (await apiClient.get<{ items: McpToolSummary[] }>(`${base}/${id}/tools`)).items
export const deleteMcpServer = (id: string) => apiClient.delete<OperationResponse>(`${base}/${id}`)
export const trustMcpServer = (server: McpServer) => apiClient.post<McpServer>(`${base}/${server.server_id}/trust`, { command_digest: server.command_digest })
export const testMcpServer = (id: string) => apiClient.post<McpServer>(`${base}/${id}/test`)
export const enableMcpServer = (id: string) => apiClient.post<McpServer>(`${base}/${id}/enable`)
export const disableMcpServer = (id: string) => apiClient.post<McpServer>(`${base}/${id}/disable`)
export const putMcpServerSecret = (id: string, key: string, secret: string, kind: 'environment' | 'header' = 'environment') => apiClient.put(`${base}/${id}/secrets/${encodeURIComponent(key)}?kind=${kind}`, { secret })
export const deleteMcpServerSecret = (id: string, key: string, kind: 'environment' | 'header' = 'environment') => apiClient.delete(`${base}/${id}/secrets/${encodeURIComponent(key)}?kind=${kind}`)
+4
View File
@@ -37,3 +37,7 @@ export async function deleteNote(noteId: string): Promise<OperationResponse> {
export async function moveNote(noteId: string, folder: string): Promise<ApiNote> { export async function moveNote(noteId: string, folder: string): Promise<ApiNote> {
return apiClient.post(`/api/notes/${noteId}/move`, { folder }) return apiClient.post(`/api/notes/${noteId}/move`, { folder })
} }
export async function renameNote(noteId: string, fileName: string): Promise<ApiNote> {
return apiClient.post(`/api/notes/${noteId}/rename`, { file_name: fileName })
}
@@ -0,0 +1,85 @@
// @vitest-environment happy-dom
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import * as pluginService from './pluginService'
function jsonResponse(body: unknown) {
return new Response(JSON.stringify(body), {
status: 200,
headers: { 'Content-Type': 'application/json' },
})
}
beforeEach(() => {
vi.stubGlobal('fetch', vi.fn())
})
afterEach(() => {
vi.unstubAllGlobals()
vi.restoreAllMocks()
})
describe('pluginService contribution adapter', () => {
it('lists and executes Plugin Commands with scoped wire fields', async () => {
const fetchMock = vi.mocked(fetch)
fetchMock
.mockResolvedValueOnce(jsonResponse({ items: [{ command_id: 'text-tools.uppercase-selection' }] }))
.mockResolvedValueOnce(jsonResponse({
command_id: 'text-tools.uppercase-selection',
status: 'completed',
effect: { type: 'notification', payload: { level: 'success', message: 'HELLO' } },
}))
const commands = await pluginService.listPluginCommands('command_palette')
const result = await pluginService.executePluginCommand(
'text-tools.uppercase-selection',
{},
{ note_id: 'note-1', selection: 'hello' },
)
expect(commands[0].command_id).toBe('text-tools.uppercase-selection')
if (result.effect.type !== 'notification') throw new Error('expected notification effect')
expect(result.effect.payload.message).toBe('HELLO')
expect(fetchMock.mock.calls[0][0]).toBe(
'/api/plugin-contributions/commands?location=command_palette',
)
expect(JSON.parse(String(fetchMock.mock.calls[1][1]?.body))).toEqual({
arguments: {},
context: { note_id: 'note-1', selection: 'hello' },
})
})
it('uses separate Settings and Secret endpoints', async () => {
const fetchMock = vi.mocked(fetch)
fetchMock
.mockResolvedValueOnce(jsonResponse({
plugin_id: 'text-tools', schema_version: 1, fields: [],
values: { result_limit: 10 }, secrets: { api_key: { configured: false } },
}))
.mockResolvedValueOnce(jsonResponse({
plugin_id: 'text-tools', schema_version: 1, fields: [],
values: { result_limit: 20 }, secrets: { api_key: { configured: false } },
}))
.mockResolvedValueOnce(jsonResponse({ plugin_id: 'text-tools', key: 'api_key', configured: true }))
.mockResolvedValueOnce(jsonResponse({ plugin_id: 'text-tools', key: 'api_key', configured: false }))
await pluginService.getPluginSettings('text-tools')
await pluginService.updatePluginSettings('text-tools', 1, { result_limit: 20 })
await pluginService.putPluginSecret('text-tools', 'api_key', 'request-only-secret')
await pluginService.deletePluginSecret('text-tools', 'api_key')
expect(fetchMock.mock.calls.map(([url]) => url)).toEqual([
'/api/plugins/text-tools/settings',
'/api/plugins/text-tools/settings',
'/api/plugins/text-tools/settings/api_key/secret',
'/api/plugins/text-tools/settings/api_key/secret',
])
expect(JSON.parse(String(fetchMock.mock.calls[1][1]?.body))).toEqual({
schema_version: 1,
values: { result_limit: 20 },
})
expect(JSON.parse(String(fetchMock.mock.calls[2][1]?.body))).toEqual({
secret: 'request-only-secret',
})
expect(fetchMock.mock.calls[3][1]?.method).toBe('DELETE')
})
})
+72 -3
View File
@@ -1,5 +1,17 @@
import apiClient from './apiClient' import apiClient from './apiClient'
import type { ApiPlugin, OperationResponse, Plugin, PluginContribution } from '@/contracts' import type {
ApiPlugin,
OperationResponse,
Plugin,
PluginCommand,
PluginCommandContext,
PluginCommandLocation,
PluginCommandResult,
PluginContribution,
PluginHostStatus,
PluginSecretStatus,
PluginSettingsSchema,
} from '@/contracts'
function toPlugin(plugin: ApiPlugin): Plugin { function toPlugin(plugin: ApiPlugin): Plugin {
const { manifest } = plugin const { manifest } = plugin
@@ -54,6 +66,63 @@ export async function grantPluginPermissions(pluginId: string, permissions: stri
return toPlugin(await apiClient.put<ApiPlugin>(`/api/plugins/${pluginId}/permissions`, { permissions })) return toPlugin(await apiClient.put<ApiPlugin>(`/api/plugins/${pluginId}/permissions`, { permissions }))
} }
export async function getPluginHostStatus(pluginId: string): Promise<PluginHostStatus> {
return apiClient.get(`/api/plugins/${pluginId}/host`)
}
export async function restartPluginHost(pluginId: string): Promise<OperationResponse> {
return apiClient.post(`/api/plugins/${pluginId}/host/restart`)
}
export async function listPluginCommands(location?: PluginCommandLocation): Promise<PluginCommand[]> {
const query = location ? `?location=${encodeURIComponent(location)}` : ''
const response = await apiClient.get<{ items: PluginCommand[] }>(`/api/plugin-contributions/commands${query}`)
return response.items
}
export async function executePluginCommand(
commandId: string,
argumentsValue: Record<string, unknown> = {},
context: PluginCommandContext = {},
): Promise<PluginCommandResult> {
return apiClient.post(`/api/plugin-contributions/commands/${encodeURIComponent(commandId)}/execute`, {
arguments: argumentsValue,
context,
})
}
export async function getPluginSettings(pluginId: string): Promise<PluginSettingsSchema> {
return apiClient.get(`/api/plugins/${encodeURIComponent(pluginId)}/settings`)
}
export async function updatePluginSettings(
pluginId: string,
schemaVersion: number,
values: Record<string, unknown>,
): Promise<PluginSettingsSchema> {
return apiClient.put(`/api/plugins/${encodeURIComponent(pluginId)}/settings`, {
schema_version: schemaVersion,
values,
})
}
export async function putPluginSecret(
pluginId: string,
key: string,
secret: string,
): Promise<PluginSecretStatus> {
return apiClient.put(
`/api/plugins/${encodeURIComponent(pluginId)}/settings/${encodeURIComponent(key)}/secret`,
{ secret },
)
}
export async function deletePluginSecret(pluginId: string, key: string): Promise<PluginSecretStatus> {
return apiClient.delete(
`/api/plugins/${encodeURIComponent(pluginId)}/settings/${encodeURIComponent(key)}/secret`,
)
}
export async function uninstallPlugin(pluginId: string): Promise<OperationResponse> { export async function uninstallPlugin(pluginId: string): Promise<OperationResponse> {
return apiClient.delete(`/api/plugins/${pluginId}`) return apiClient.delete(`/api/plugins/${pluginId}`)
} }
@@ -65,7 +134,7 @@ export const mockPlugins: Plugin[] = [
version: '1.3.2', version: '1.3.2',
description: '接入 GitHub API,支持搜索 Issue、查看 PR 和管理仓库', description: '接入 GitHub API,支持搜索 Issue、查看 PR 和管理仓库',
icon: '', icon: '',
author: '知笔知己团队', author: 'NotesAgent 团队',
status: 'ready', status: 'ready',
enabled: true, enabled: true,
permissions: ['notes.read', 'network.request'], permissions: ['notes.read', 'network.request'],
@@ -117,7 +186,7 @@ export const mockPlugins: Plugin[] = [
version: '2.1.0', version: '2.1.0',
description: '导入 PDF 文档,提取文本和目录结构生成笔记', description: '导入 PDF 文档,提取文本和目录结构生成笔记',
icon: '', icon: '',
author: '知笔知己团队', author: 'NotesAgent 团队',
status: 'error', status: 'error',
enabled: false, enabled: false,
permissions: ['notes.write', 'attachments.read'], permissions: ['notes.write', 'attachments.read'],
+2 -2
View File
@@ -50,7 +50,7 @@ export const mockSkills: Skill[] = [
version: '1.0.0', version: '1.0.0',
description: '根据课程笔记生成复习要点和练习题,帮助高效备考', description: '根据课程笔记生成复习要点和练习题,帮助高效备考',
icon: '', icon: '',
author: '知笔知己团队', author: 'NotesAgent 团队',
permissions: ['notes.search', 'notes.read', 'tasks.create'], permissions: ['notes.search', 'notes.read', 'tasks.create'],
tools: ['notes.search', 'notes.read', 'tasks.create'], tools: ['notes.search', 'notes.read', 'tasks.create'],
retrieval_config: { top_k: 10, rerank: true, citation: true }, retrieval_config: { top_k: 10, rerank: true, citation: true },
@@ -64,7 +64,7 @@ export const mockSkills: Skill[] = [
version: '1.1.0', version: '1.1.0',
description: '从音频或文本中提取会议要点、行动项和待办任务', description: '从音频或文本中提取会议要点、行动项和待办任务',
icon: '', icon: '',
author: '知笔知己团队', author: 'NotesAgent 团队',
permissions: ['notes.search', 'notes.write', 'tasks.write', 'attachments.read'], permissions: ['notes.search', 'notes.write', 'tasks.write', 'attachments.read'],
tools: ['notes.search', 'notes.create', 'tasks.create', 'attachments.read'], tools: ['notes.search', 'notes.create', 'tasks.create', 'attachments.read'],
retrieval_config: { top_k: 5, rerank: false, citation: true }, retrieval_config: { top_k: 5, rerank: false, citation: true },
+42
View File
@@ -0,0 +1,42 @@
import { afterEach, describe, expect, it, vi } from 'vitest'
import { SseClient } from './sseClient'
afterEach(() => {
vi.unstubAllGlobals()
vi.restoreAllMocks()
})
describe('SseClient resumable event transport', () => {
it('sends Last-Event-ID and exposes the returned SSE id', async () => {
const fetchMock = vi.fn().mockResolvedValue(
new Response(
'id: 3\nevent: ModelCallCompleted\ndata: {"sequence":3,"data":{"duration_ms":12}}\n\n',
{ status: 200, headers: { 'Content-Type': 'text/event-stream' } },
),
)
vi.stubGlobal('fetch', fetchMock)
const received = vi.fn()
const client = new SseClient({
url: '/api/agent/runs/run-1/events?after_sequence=2',
method: 'GET',
lastEventId: '2',
onEvent: received,
})
await client.connect()
expect(fetchMock).toHaveBeenCalledWith(
'/api/agent/runs/run-1/events?after_sequence=2',
expect.objectContaining({
method: 'GET',
headers: expect.objectContaining({ 'Last-Event-ID': '2' }),
}),
)
expect(received).toHaveBeenCalledWith(
'ModelCallCompleted',
{ sequence: 3, data: { duration_ms: 12 } },
'3',
)
})
})

Some files were not shown because too many files have changed in this diff Show More