32 Commits
Author SHA1 Message Date
徐涛 2af92cd554 docs(roadmap): 标记 Phase 10 ContextSlot 上下文管理已完成
- 顶部状态行更新:Phase 0-10 全部完成,11 个离线示例,254 测试
- Phase 10 状态:Step 10.1/10.2/10.3 全部标记 
- 详细实际新增:agent/context.rs (~430 行) + agent/error.rs 3 个变体 +
  agent/session.rs 改造(slots 字段 + 5 个管理方法 + submit_turn/finalize_turn
  增量追加写回) + context_slot_demo 示例 + 43 个新测试
- finalize_turn 签名变更记录(new_messages_from_cycle 参数 + Result 返回)
- 依赖关系图 P10 节点标 
- 里程碑 M6 →  2026-07-07
- 下一步行动:Phase 10 → Phase 11(测试与检索补强)
- 已完成阶段列表追加 Phase 10 完整说明
- 最后更新日期:2026-07-06 → 2026-07-07
2026-07-06 09:53:44 +08:00
徐涛 635942248b feat(core): 新增 Phase 10 ContextSlot 多上下文分区管理
- 新增 ContextSlot 类型(Full / Focused / Readonly 三种模式,
  New / Derived / Static 三种来源),支持 JSON blob 批次持久化
- AgentSession 新增 slots 字段与 5 个管理方法
  (create_slot / switch_slot / list_slots / derive_slot / delete_slot),
  自动创建 "default" slot
- submit_turn / finalize_turn 改造为基于当前 slot 的增量追加写回,
  确保 Focused 模式"读时过滤"语义不丢失数据
- finalize_turn 签名变更(新增 new_messages_from_cycle 参数,
  返回 Result<(), AgentError>),向后兼容列于 docs/17
- 新增 3 个 AgentError 变体(SlotReadonly / SlotNotFound / SlotAlreadyExists)
- 新增分支对话示例 context_slot_demo(法律咨询→两个派生方向→切换→隔离验证)
- 新增 43 个测试覆盖持久化、Focused 过滤、Readonly 阻断、delete 保护、
  派生逻辑、流式 finalize_turn、key 注入防护等场景
- 方案文档:docs/17-phase10-contextslot.md(含 §5 推荐方案、§6 实施建议、
  §9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复)
2026-07-06 09:46:04 +08:00
徐涛 fe51961202 docs(roadmap): 标记 Phase 9 流式体验增强已完成
- 顶部「当前状态」「最后更新」同步到 Phase 0-9 完成
- Phase 9 Step 9.1 标记完成并新增「实际新增」段落
- Mermaid 依赖图 P9 节点 class 从 p1 切换为 done
- 里程碑 M5 标记  2026-07-06
- 下一步行动与已完成列表同步更新(指向 Phase 10)
2026-07-06 05:43:47 +08:00
徐涛 212cfcc916 feat(core): 新增 Phase 9 流式体验增强 submit_turn_stream
- 新增 StreamEvent::ToolExecutionStarted/Completed 变体(apply_to 元事件)
- 新增 LlmCycle::submit_with_tools_stream + run_tool_loop(spawn + mpsc 状态机)
- 新增 AgentSession::submit_turn_stream + finalize_turn(手动同步状态)
- CycleConfig 加 Clone derive

- 新增 9 个单元测试 + 2 个集成测试(211 passed)
- clippy 0 警告,存量 0 回归

docs: 追加 Phase 9 实施方案(docs/16-phase9-streaming-experience.md)
2026-07-05 23:37:37 +08:00
徐涛 88d00ac927 docs(roadmap): 补充 Phase 8 修复 commit 85b92ae 与最终行数
Phase 8 章节实际新增小节从 6 commits 改为 7 commits(含 docs(roadmap)
自身 + 实施后修复 commit 85b92ae),并同步行数:

- quick_start 57 → 60 行(追加 trailing newline + 错误处理改进)
- end_to_end 237 → 246 行(除零修复 + drop 注释 + EchoTool 一致化)
- 新增 docs(roadmap) commit 项(标记 Phase 8 + M4 里程碑)
- 新增 fix(examples) commit 项(6 项审查问题:1 🔴 + 2 🟡 + 3 💭)

'已完成 / 进行中阶段' 列表同步补充 Phase 8 行末的三方审查修复标注。
2026-07-05 21:10:24 +08:00
徐涛 358e971094 feat(core): 新增 Phase 8 MVP 集成出口方案文档 2026-07-05 21:04:31 +08:00
徐涛 85b92ae9d4 fix(examples): 修复 Phase 8 端到端示例审查发现的 6 项问题
按 PM/SA/Code Reviewer 三方审计报告修复:

🔴 阻塞修复:
- CalcTool 除零 panic: `a / b` 改 `a.checked_div(b).ok_or_else(...)`,
  b=0 时返回 ToolError::InvalidArguments 而非 panic

🟡 警告修复:
- end_to_end.rs 持久化验证注释与实际不符: 显式 drop(backend) 让注释
  描述与 Arc 释放顺序一致
- quick_start EchoTool 参数验证: 用 args.get("text").and_then().ok_or_else()
  替换 as_str().unwrap_or("") 静默降级, 缺失/类型错误时返回显式错误

💭 一致性修复:
- end_to_end.rs EchoTool 与 quick_start 一致化 (format!("收到: {text}"))
- quick_start mock 响应文本 "已通过 echo 回传" → "EchoTool 已收到并完成回传"
- quick_start 断言改为检查 "收到", 与方案字面要求一致
- quick_start 末尾追加 POSIX trailing newline

验收: cargo test 200/0 + clippy 0 警告 + doc 0 warning + 10 示例 exit 0
2026-07-05 21:00:36 +08:00
徐涛 57b2fbaaed docs(roadmap): 标记 Phase 8 MVP 集成出口已完成,v0.2.0-rc.1 标签
Phase 8 全部 Step 完成(8.1 API 稳定性扫尾 + 8.2 quick_start + 8.3
end_to_end),更新 roadmap 同步:

- 顶部「当前状态」补充 v0.2.0-rc.1 标签信息
- Phase 8 章节 3 个 Step 全标 ;新增「实际新增」6 commits 小节
- 依赖关系图 P8 节点 class 从 mvp 改为 done
- M4 里程碑标  2026-07-05
- 「下一步行动」从 Phase 8 改为 Phase 9 + v0.2.0 正式版规划
- 「已完成 / 进行中阶段」列表追加  Phase 8
2026-07-05 20:06:43 +08:00
徐涛 2c8e31919d feat(examples): 新增 end_to_end 端到端集成示例
展示真实场景集成: 3 工具 (EchoTool + CalcTool + NoteTool) +
3 轮对话 (计算 → 记笔记 → 回忆) + SqliteStore 持久化跨连接验证。

关键设计:
- Provider 自动检测: AG_LLM_* 环境变量齐全时用真实 LLM (from_env),
  否则降级 MockProvider + 9 条预设响应 (含真实工具调用序列)
- NoteTool 展示 MemoryStore trait 解耦: 直接持有 Arc<dyn MemoryStore>,
  绕过 AgentSession 封装 (key 前缀 "note:" + list 过滤)
- CalcTool ponytail 方案: 手写 'a op b' 解析, 不引入 rhai 依赖
- 持久化验证: drop(bundle) + drop(session) → backend Arc 引用归零 →
  SqliteStore Connection 自动 close → 重开连接读取数据存活

规模: 237 行 (含完整注释); clippy 0 警告; cargo run exit 0。
2026-07-05 20:04:44 +08:00
徐涛 c6651c9b75 feat(examples): 新增 quick_start 最小可运行示例
展示 Agent / BaseTool / AgentBuilder / AgentSession 四层抽象的
最小可行集成:MockProvider + EchoTool + submit_turn("你好"),
EchoTool 真正被 LLM 调用并产出 "收到" 字样,零外部配置。

规模: 57 行(含注释),clippy 0 警告 + cargo run exit 0。
2026-07-05 20:01:13 +08:00
徐涛 e636e16820 test(core): 验证 Phase 8 Step 8.1 commit 1-3 零回归
cargo test --all-targets: 200 passed / 0 failed
cargo clippy --all-targets -- -D warnings: 0 警告
cargo doc --no-deps: 0 warning

Step 8.1 四个增量 commit 全部交付,零 deprecated warning。
2026-07-05 19:50:57 +08:00
徐涛 1c89d23ba2 docs: v0.2.0-rc.1 CHANGELOG + 版本号 + README 示例列表
- CHANGELOG.md 新增 [0.2.0-rc.1] 条目(Added/Changed/Non-exhaustive/
  Deprecated/Fixed/Migration Guide 六节)
- Cargo.toml version: 0.1.0 → 0.2.0-rc.1
- README 示例列表: 7 → 10(新增 quick_start / end_to_end /
  simple_visit),依赖版本 0.1 → 0.2

验收: 人工 review + git diff 确认版本号与示例数对齐。
2026-07-05 19:50:10 +08:00
徐涛 5b4343a051 refactor(agent): StepStatus::Completed 切换至 MessageResponse
StepStatus::Completed 字段类型从废弃的 ChatResponse 切换为 IR 层
MessageResponse,同时清理 task_agent_demo.rs 的 3 处废弃类型引用:

- ChatResponse → MessageResponse (字段映射见 §3.1.2)
- OpenaiChatMessage::assistant_text(t) → Message::assistant(t)
- FinishReason::Stop → StopReason::Stop (Option 包裹同步移除)

移除 src/agent/task.rs 的两处 #[allow(deprecated)] 与
examples/task_agent_demo.rs 顶部 #![allow(deprecated)],零
deprecated warning。

验收: cargo build + clippy -D warnings + test --all-targets 全绿;
cargo run --example task_agent_demo exit 0。
2026-07-05 19:48:59 +08:00
徐涛 6e1182e64c feat(core): 14 个公开枚举标记 #[non_exhaustive]
为 v0.2.0-rc.1 API 稳定性收尾。覆盖:
- P0 核心 IR: Message / ContentBlock / ContentBlockType / StreamEvent / HookEvent
- P0 Error: AgentError / LlmError / ToolError / MemoryError / PromptError
- P1 其他: MemoryStrategy / StepStatus / ToolChoice / ResponseFormat

明确不加的: 内部 wire-format (OpenaiChatMessage 等) /
语义已收敛 (Role/ServiceTier 等) / 使用面窄 (Permission 等)。

示例侧的 3 处 Message exhaustive match (prompt_composer /
conversation_memory_demo) 补全 `_` 通配分支,零行为变化。

验收: cargo build + clippy -D warnings + test --all-targets 全绿 (200 passed)
2026-07-05 19:47:03 +08:00
徐涛 b8f4fe0fe3 docs(AGENTS.md): 添加进度同步规范,明确 roadmap 更新要求 2026-07-05 17:40:22 +08:00
徐涛 7574f9c24c docs(roadmap): 标记 Phase 7 SqliteStore 已完成
更新当前状态描述、Phase 7 交付物详情、里程碑 M3 状态、依赖性图谱及下一步行动
2026-07-05 17:25:32 +08:00
徐涛 c82af60f81 feat(memory): 实现 SqliteStore 持久化
- 新增 src/memory/store/sqlite_store.rs(SqliteStore + 9 个内联测试)
- 基于 rusqlite 0.32(bundled),使用 Arc<Mutex<Connection>> + spawn_blocking
- WAL 模式 + synchronous=NORMAL + busy_timeout=5s + PRAGMA user_version
  schema 版本管理
- created_at 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
- 错误精细映射:SqliteFailure/InvalidQuery → InvalidInput;
  FromSqlConversionFailure → Serialization;其他 → Storage
- 9 个测试覆盖 CRUD、upsert、prefix/since/offset+limit 过滤、
  10 写者 × 10 次并发、持久化 round-trip、trait-box 互换兼容性

依赖:
- rusqlite = { version = "0.32", features = ["bundled"] }
- time 增补 features: parsing, formatting, macros
- dev-dependencies: tempfile = "3"

测试:199 → 200 pass(191 原有 + 9 新增)
2026-07-05 17:13:28 +08:00
徐涛 c8a91f6eaf refactor(memory): 将 store.rs 拆分为模块目录,仅结构搬移
- 将 src/memory/store.rs 单体文件拆分为模块根 + store/ 目录结构
- InMemoryStore(struct + impl + Default + 6 个内联测试)整体提取到
  src/memory/store/in_memory.rs
- store.rs 保留 MemoryStore trait 与 EvictionPolicy/EvictionConfig
- 外部消费者导入路径 crate::memory::store::MemoryStore 不变
2026-07-05 17:09:45 +08:00
徐涛 13edacd775 feat(memory): 添加 SqliteStore 持久化方案文档 2026-07-05 16:54:29 +08:00
徐涛 821cea8e60 feat(docs): 更新 roadmap,标记 Phase 6 交付完成
Phase 6 ToolDef IR 的 5 个子步骤全部标注已完成,新增实际变更摘要,更新里程碑状态及下一步行动计划
2026-07-05 10:49:50 +08:00
徐涛 517ef7db32 docs: 添加 Phase 6 ToolDefinition IR 正式化实施方案 2026-07-05 10:46:46 +08:00
徐涛 4cf5918b9c refactor(llm): 移除 ToolDefinition 别名与遗留 deprecation 抑制点
完成 ToolDef IR 切换的最后清理:移除 ToolDefinition 别名与
OpenaiToolDefinition 的公共 re-export;各调用点(cycle/registry/mcp/agent)
直接使用 ToolDef 并清理对应的 #[allow(deprecated)] 抑制点;新增
roundtrip 测试验证 ToolDef 序列化兼容性。
2026-07-05 10:32:35 +08:00
徐涛 9da9b83167 feat(llm): 切换至 ToolDef IR 并适配 OpenAI Provider
将 MessageRequest.tools 切换为 Provider 无关的 ToolDef IR;
ToolDefinition 别名切换指向 ToolDef,registry/mcp 构造去掉
strict 字段;OpenAI 适配层通过 From 转换将 ToolDef 转为
OpenaiToolDefinition wire format。Anthropic 适配层字段名一致
无需改动,openai_compat/ollama 委托 GenericOpenaiProvider 无影响。
2026-07-05 10:14:32 +08:00
徐涛 b1875192fd feat(llm): 引入 ToolDef IR 类型
新增 Provider 无关的工具定义中间表示 ToolDef,配套实现双向 From 转换;
旧 OpenaiToolDefinition 标记 doc(hidden)。本次仅新增类型,不影响现有
代码路径。
2026-07-05 10:06:52 +08:00
徐涛 d3067e2f53 docs(roadmap): 更新 Phase 5 完成状态和下一阶段计划 2026-07-05 09:04:44 +08:00
徐涛 5648b1d217 style(tools, llm): 统一导入顺序与代码格式 2026-07-05 08:19:13 +08:00
徐涛 98dfe6c1ed feat(llm): 完成 Phase 5 热身准备(Ollama / non_exhaustive / ProviderConfig)
Phase 5 三个 Step 全部落地:

Step 5.2 — Ollama Provider
- 新增 OllamaProvider newtype 包装(默认 localhost:11434/v1,零 API key)
- ProviderType 新增 Ollama 变体与 FromStr 解析

Step 5.3 — #[non_exhaustive] 前置标记
- ProviderType / StopReason / FinishReason / EvictionPolicy 加 #[non_exhaustive]
- 编译期兼容护栏,避免下游 silent break

Step 5.1 — ProviderConfig 扩展
- 加 timeout_secs / max_retries 字段、Default、from_env(prefix)
- create_provider 各分支通过 pub(crate) from_parts 一次性构造并注入 timeout
  (同时避开 Anthropic 的 default_headers 与双重 client 构造)
- map_reqwest_error 改为方法读取 self.timeout_secs(移除硬编码 120s)
- AnthropicProvider::with_timeout 同值短路,with_client 标 #[deprecated]
- DeepSeek / Qwen 加公开 with_client,new_with_client 走代理
- 7 个新测试:5 个 from_env 单元测试 + 3 个 timeout 传导 wiremock
  (OpenAI Chat / DeepSeek / Anthropic)
- Cargo.toml 加 temp-env dev-dep
2026-07-05 08:12:40 +08:00
徐涛 9e476e79bb merge: 合并 release/v0.1 全部修改(4 个 commit 含 v0.2 路线图与 Phase 5 实施方案) 2026-07-05 07:32:23 +08:00
徐涛 3bd135ec98 docs(roadmap): 添加 Phase 5 warmup 实施方案
明确 Step 5.2 Ollama Provider、Step 5.3 non_exhaustive 前置标记、
Step 5.1 ProviderConfig 扩展三个独立 Step 的执行顺序与交付物
2026-07-05 07:32:19 +08:00
徐涛 76f3235ed7 docs(roadmap): 更新 v0.2 规划为 8 个增量 Phase 并细化实施步骤 2026-07-04 10:32:24 +08:00
徐涛 6315f2d008 docs: 添加 opencode 子代理调度分发与合并机制调研笔记 2026-07-04 10:19:54 +08:00
徐涛 fba78f5f33 docs(roadmap): 更新路线图为 v0.2 生产就绪规划
同步 v0.1.0 发布状态,将 v0.2+ 扩展项重
组为 12 项基础功能和 ContextSlot 上下文管
理,明确 v0.3+ 展望及边界范围
2026-07-04 08:11:09 +08:00
69 changed files with 9785 additions and 1080 deletions
+12
View File
@@ -198,6 +198,18 @@ pub use vector_store::VectorStore;
5. **风险评估** - 潜在风险、缓解措施
6. **验收标准** - 可验证的完成条件
### 进度同步规范 (docs/roadmap.md)
完成一项实施后,必须检查 `docs/roadmap.md` 是否存在对应内容;若存在,必须同步标记为完成:
- **Step / Phase 状态行**:对应 Step 加 ✅ 标记;Phase 章节末尾「状态」行从 ⏳ 改为 ✅ Phase X 全部交付物已完成
- **里程碑表**:更新对应里程碑状态从 ⏳ 改为 ✅ + 完成日期
- **依赖关系图(Mermaid**:节点 `class``pending` / `core` 改为 `done`,必要时更新节点摘要
- **文末「已完成 / 进行中阶段」列表**:追加一行 `- ✅ Phase X — 一句话要点`
- **顶部「当前状态」**:补充新完成 Phase,更新「下一步」指向
参考案例:2026-07-05 完成 Phase 7 SqliteStore 时同步更新 6 处(顶部状态 / Phase 章节 / 依赖图 / M3 / 下一步行动 / 已完成列表)。
---
## 项目特定规则
+61
View File
@@ -2,6 +2,67 @@
本项目所有重要变更均记录于此文件。格式参考 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/)。
## [0.2.0-rc.1] - 2026-07-05
v0.2.0 候选发布。Phase 5-7 三大 P0 全部交付完成,API 稳定性扫尾,新增 2 个面向新用户的集成示例。
### Added
**Phase 5 — 热身准备**
- `ProviderConfig::from_env(prefix)`:从 `{prefix}_BASE_URL` / `{prefix}_API_KEY` / `{prefix}_MODEL` / `{prefix}_TIMEOUT_SECS` / `{prefix}_MAX_RETRIES` 环境变量构造配置
- `ProviderConfig::timeout_secs` / `max_retries` 字段(默认 30 / 3
- `OllamaProvider`:本地推理 ProviderOpenAI-compatible`http://localhost:11434/v1` 默认端点)
- `ProviderType::Ollama` 变体 + `FromStr` 解析
**Phase 6 — ToolDef IR 正式化**
- `ToolDef` 结构体(name / description / parameters)替代已废弃的 `OpenaiToolDefinition`
- `MessageRequest.tools` 切换为 `Vec<ToolDef>`
- `OpenaiToolDefinition` 降级为 `#[doc(hidden)]`,仅供 OpenAI 适配层内部消费
**Phase 7 — SqliteStore 持久化**
- `SqliteStore``MemoryStore` 的 SQLite 后端实现,基于 `rusqlite 0.32` bundled
- WAL 模式 + `synchronous=NORMAL` + `busy_timeout=5s` 兼顾崩溃安全与吞吐
- `Arc<Mutex<Connection>>` + `spawn_blocking` 卸载 IO10×10 并发写入无 race
- `MemoryStore::save / get / delete / list` CRUD + prefix / since / offset+limit 过滤
- 进程重启数据不丢的 round-trip 验证
**Phase 8 — MVP 集成出口**
- `examples/quick_start`:30 行最小可运行示例(MockProvider + EchoTool + submit_turn
- `examples/end_to_end`3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证
### Changed
- **API 稳定性护栏**:14 个公开枚举标记 `#[non_exhaustive]`,覆盖 P0 核心 IR`Message` / `ContentBlock` / `ContentBlockType` / `StreamEvent` / `HookEvent`)、P0 Error`AgentError` / `LlmError` / `ToolError` / `MemoryError` / `PromptError`)、P1 其他(`MemoryStrategy` / `StepStatus` / `ToolChoice` / `ResponseFormat`
- **`StepStatus::Completed` 字段类型**:从废弃的 `ChatResponse` 切换为 IR 层 `MessageResponse`(同时清理 `task_agent_demo.rs``ChatResponse` / `OpenaiChatMessage` / `FinishReason` 三处废弃类型引用)
### Non-exhaustive 清单
为防止未来新增变体时下游 exhaustive match 静默失效,14 个枚举追加 `#[non_exhaustive]`
| 优先级 | 枚举 |
|--------|------|
| P0 核心 IR | `Message`, `ContentBlock`, `ContentBlockType`, `StreamEvent`, `HookEvent` |
| P0 Error | `AgentError`, `LlmError`, `ToolError`, `MemoryError`, `PromptError` |
| P1 其他 | `MemoryStrategy`, `StepStatus`, `ToolChoice`, `ResponseFormat` |
明确不加:内部 wire-format`OpenaiChatMessage` 等)/ 语义已收敛(`Role` / `ServiceTier` / `Modality` / `ImageDetail` / `AudioFormat` / `StopSequence`/ 使用面窄(`Permission` / `McpTransport` 等)。
### Deprecated
(继承自 0.1.0,无新增)`ChatResponse` / `ToolDefinition` 保持 `#[deprecated]` 标记。
### Fixed
- 修复 `StepStatus::Completed(ChatResponse)` 字段类型与 IR 体系不一致问题(已完成迁移)
### Migration Guide (v0.1 → v0.2.0-rc.1)
1. **枚举 match**14 个 `#[non_exhaustive]` 枚举在 crate 外必须使用 `_ =>` 通配分支
2. **`StepStatus::Completed`**:字段类型从 `ChatResponse` 切换为 `MessageResponse`,需做字段映射(参考 `docs/15-phase8-mvp-integration.md` §3.1.2
3. **`ToolDefinition``ToolDef`**Phase 6 已彻底替换 `#[deprecated]` 别名,需全局重命名
---
## [0.1.0] - 2026-07-04
首个公开版本。涵盖 Phase 0-4c 的全部核心能力、Provider IR 重构、LlmCycle 简化,以及面向用户的 7 个离线示例。
+5 -2
View File
@@ -1,6 +1,6 @@
[package]
name = "agcore"
version = "0.1.0"
version = "0.2.0-rc.1"
edition = "2024"
[dependencies]
@@ -19,8 +19,11 @@ futures-core = "0.3"
bytes = "1"
async-stream = "0.3"
tokio-util = { version = "0.7", features = ["rt"] }
time = { version = "0.3", features = ["serde"] }
time = { version = "0.3", features = ["serde", "parsing", "formatting", "macros"] }
rusqlite = { version = "0.32", features = ["bundled"] }
[dev-dependencies]
dotenvy = "0.15.7"
wiremock = "0.6"
temp-env = "0.3"
tempfile = "3"
+5 -2
View File
@@ -26,7 +26,7 @@ AG Core 不是 Agent 产品,而是 Agent 的**底层依赖库**:上层应用
```toml
[dependencies]
agcore = "0.1"
agcore = "0.2"
tokio = { version = "1", features = ["macros", "rt-multi-thread"] }
```
@@ -110,10 +110,12 @@ let provider = create_provider(
).expect("创建 Provider 失败");
```
更多端到端示例见 [`examples/`](./examples/) 目录(共 7 个,全部可 `cargo run --example <name>`):
更多端到端示例见 [`examples/`](./examples/) 目录(共 10 个,全部可 `cargo run --example <name>`):
| 示例 | 说明 |
|------|------|
| `quick_start` | **30 行最小示例**MockProvider + EchoTool + submit_turn,新用户 5 分钟上手 |
| `end_to_end` | **完整集成示例**3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证 |
| `agent_session_demo` | Agent + 会话 + SessionMemory 完整链路(MockProvider 离线) |
| `custom_tool` | 自定义工具注册、单次 / 并行调用、权限检查 |
| `prompt_composer` | 提示词模板与组合器(纯离线) |
@@ -121,6 +123,7 @@ let provider = create_provider(
| `conversation_memory_demo` | 对话记忆滑动窗口与隔离 |
| `knowledge_search_demo` | 知识页面关键词检索 |
| `streaming_events_demo` | LLM 流式响应事件消费(含错误路径) |
| `simple_visit` | 真实 LLM 调用(OpenAI / Anthropic,设置 `OPENAI_*` / `ANTHROPIC_*` 环境变量) |
## 核心模块
+522
View File
@@ -0,0 +1,522 @@
# Phase 5:热身准备 — 实施方案
## 1. 背景与目标
Phase 5 是 v0.2.0 发布周期的**热身准备阶段**,包含三个互不依赖的 Step,为后续 Phase 6-12 的端到端集成提供基础设施。
**核心目标**
- 为 Phase 8(端到端示例)提供零 API key 的运行路径(Ollama
- 为公共枚举的向后兼容性加上编译期护栏(`#[non_exhaustive]`
- 为 Provider 构造提供统一的超时与重试配置入口(`ProviderConfig` 扩展)
三个 Step 之间**无依赖关系**,但出于实现效率考虑,按 **5.2 → 5.3 → 5.1** 顺序执行。理由:5.2 先新增 `Ollama` 枚举变体,5.3 再加 `#[non_exhaustive]`,避免枚举标记后添加变体需要在外部 crate 加 `_ =>` 兜底分支的困扰。
## 2. 需求分析
### Step 5.2 — Ollama Provider
| 维度 | 内容 |
|------|------|
| **需求** | 新增 `OllamaProvider`newtype 包装 `GenericOpenaiProvider`,默认连接本地 Ollama 实例 |
| **优先级** | P0 — 为 Phase 8 端到端示例提供无需 API key 的运行路径 |
| **预期交付物** | `src/llm/provider/ollama.rs` 新建文件;`ProviderType` 新增 `Ollama` 变体 |
| **代码量** | ~55 行 |
### Step 5.3 — `#[non_exhaustive]` 前置标记
| 维度 | 内容 |
|------|------|
| **需求** | 为 4 个公共枚举添加 `#[non_exhaustive]` 属性,避免后续新增变体时破坏下游 match |
| **优先级** | P1 — 编译期兼容性保障 |
| **预期交付物** | 修改 4 个枚举定义,各加一行属性 |
| **代码量** | ~4 行 |
### Step 5.1 — ProviderConfig 扩展
| 维度 | 内容 |
|------|------|
| **需求** | `ProviderConfig` 新增 `timeout_secs``max_retries` 字段;实现 `Default``from_env()` 构造;timeout 传导到各 Provider HTTP Client |
| **优先级** | P0 — 与 Roadmap 一致,Phase 8MVP 出口)依赖 from_env |
| **预期交付物** | `ProviderConfig` 扩展;`create_provider()` 超时注入;`from_env()` + 单元测试 |
| **代码量** | ~60 行 + 测试 |
## 3. 方案设计
### 3.1 Step 5.2 — Ollama Provider(先执行)
#### 改动文件清单
| 文件 | 操作 | 说明 |
|------|------|------|
| `src/llm/provider/ollama.rs` | **新建** | OllamaProvider newtype 包装 |
| `src/llm/provider.rs` | 修改 | `ProviderType` 新增 `Ollama` 变体;`FromStr` 加解析;`create_provider()` 加分支 |
| `src/llm/provider/mod.rs` 或其他模块注册文件 | 修改(如需要) | 注册 `pub mod ollama` |
#### 关键代码
**`src/llm/provider/ollama.rs`**(新建):
```rust
//! Ollama Provider —— OpenAI-compatible 协议的 newtype 包装,零 API key。
//!
//! 默认 base_url = `http://localhost:11434/v1`,空 api_key 也可工作。
//! 实现方式同 DeepSeekProvider / QwenProvider,共享 GenericOpenaiProvider 的 HTTP/SSE/转换逻辑。
use reqwest::Client;
use std::pin::Pin;
use async_trait::async_trait;
use futures_core::Stream;
use super::openai::GenericOpenaiProvider;
use super::{LlmProvider, ProviderCapabilities};
use crate::llm::error::LlmError;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
pub struct OllamaProvider(pub GenericOpenaiProvider);
impl OllamaProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
let url = if base_url.is_empty() {
"http://localhost:11434/v1".to_string()
} else {
base_url
};
Self(GenericOpenaiProvider::new_with_name(
url,
api_key,
model,
"ollama",
))
}
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
/// 与 `OpenaiChatProvider::with_client` 和 `DeepSeekProvider::with_client` 一致。
pub fn with_client(self, client: Client) -> Self {
Self(self.0.with_client(client))
}
}
#[async_trait]
impl LlmProvider for OllamaProvider {
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
self.0.chat(request).await
}
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
fn capabilities(&self) -> ProviderCapabilities {
let mut caps = self.0.capabilities();
caps.provider_name = "ollama";
caps
}
}
```
**`src/llm/provider.rs`** 的修改:
```rust
// ProviderType 新增变体
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderType {
OpenaiChat,
OpenaiResponse,
Anthropic,
DeepSeek,
Qwen,
/// Ollama(本地),默认 base_url = `http://localhost:11434/v1`。
Ollama,
}
// FromStr 加解析
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
// ... 已有条目 ...
"ollama" => Ok(ProviderType::Ollama),
_ => Err(format!("未知的 Provider 类型: {s}")),
}
}
// create_provider() 加分支
// Step 5.2 阶段仅展示基本构造。Step 5.1ProviderConfig 扩展)
// 执行到此分支时,将同步补充 with_client 链式调用注入 timeout
//
// let client = Client::builder()
// .timeout(Duration::from_secs(config.timeout_secs))
// .build()?;
// Ok(Box::new(
// ollama::OllamaProvider::new(config.base_url, config.api_key, config.model)
// .with_client(client),
// ))
ProviderType::Ollama => Ok(Box::new(ollama::OllamaProvider::new(
config.base_url,
config.api_key,
config.model,
))),
```
#### 集成方式
OllamaProvider 的 newtype 包装模式与 `DeepSeekProvider``QwenProvider` 完全一致,`LlmProvider` trait 委托给 `self.0``capabilities().provider_name` 返回 `"ollama"`
### 3.2 Step 5.3 — `#[non_exhaustive]` 前置标记
#### 改动文件清单
| 文件 | 行号 | 枚举 | 操作 |
|------|------|------|------|
| `src/llm/provider.rs` | ~21 | `ProviderType` | 加 `#[non_exhaustive]` |
| `src/llm/types/response_v2.rs` | ~22 | `StopReason` | 加 `#[non_exhaustive]` |
| `src/llm/types/shared.rs` | ~16 | `FinishReason` | 加 `#[non_exhaustive]` |
| `src/memory/store.rs` | ~35 | `EvictionPolicy` | 加 `#[non_exhaustive]` |
**排除清单**`SlotMode`
**决策理由**`SlotMode` 枚举在 Phase 10`src/llm/context.rs`)中才实际定义,Phase 5 尚不存在此类型。`#[non_exhaustive]` 无法标注不存在的枚举,因此排除标注。Roadmapv0.2.0 §Phase 5 Step 5.3)列出的 `SlotMode`(预置) 推迟到 Phase 10 实现时一并添加。
#### 关键代码
每个枚举在 `derive` 上方或下方加一行属性:
```rust
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum ProviderType {
// ...
}
```
#### 影响分析
- `#[non_exhaustive]` 是纯编译期属性,不影响运行时行为
- 同一 crate 内的 exhaustive match 不受影响(同 crate 可穷举)
- 下游 crate 的 match 必须加 `_ =>` 兜底分支,这是期望行为——确保未来新增变体时不会 silent break
- **单向门**:此步骤一旦通过 `v0.2.0` 发布到公共 API 后,**不可回退**。回退意味着移除 `#[non_exhaustive]`,可能破坏已添加 `_ =>` 的下游代码。因此必须在发布前完成并确认所有枚举变体正确
### 3.3 Step 5.1 — ProviderConfig 扩展(最后执行)
#### 改动文件清单
| 文件 | 操作 | 说明 |
|------|------|------|
| `src/llm/provider.rs` | 修改 | `ProviderConfig` 加字段;加 `impl Default`;加 `from_env()``create_provider` 注入 timeout |
| `src/llm/provider/openai.rs` | 修改 | `GenericOpenaiProvider` 新增 `timeout_secs` 字段;`new_with_name` 接受 timeout 参数;`map_reqwest_error` 参数化 |
| `src/llm/provider/anthropic.rs` | 修改 | 新增 `timeout_secs` 字段;`new()` 接受 timeout 参数;`map_reqwest_error` 参数化 |
| `src/llm/provider/anthropic.rs` | 修改 | 新增 `with_timeout()` 方法(返回 `Result<Self, LlmError>` |
| `src/llm/provider/openai_compat.rs` | 修改 | `DeepSeekProvider``QwenProvider` 新增公开 `with_client()` 方法 |
| `src/llm/provider/ollama.rs` | 修改 | `OllamaProvider` 新增公开 `with_client()` 方法 |
| `Cargo.toml` | 修改 | 加 `temp_env` dev-dependency |
| 测试文件(`provider.rs` 内联或独立) | 新增 | `from_env` 单元测试 + timeout 传导集成测试 |
#### 数据结构
```rust
/// Provider 构造参数 —— 通用 base_url + api_key + model + timeout/retry 配置。
pub struct ProviderConfig {
pub base_url: String,
pub api_key: String,
pub model: String,
/// 请求超时秒数(默认 30)。应用于 Provider 的 HTTP Client 级别。
pub timeout_secs: u64,
/// 最大重试次数(默认 3)。当前此字段仅由 `from_env()` 采集,
/// 实际重试逻辑由 `CycleConfig.retry.max_retries` 控制。
/// 未来可合并到统一的 retry 配置。
pub max_retries: u32,
}
impl Default for ProviderConfig {
fn default() -> Self {
Self {
base_url: String::new(),
api_key: String::new(),
model: String::new(),
timeout_secs: 30,
max_retries: 3,
}
}
}
impl ProviderConfig {
/// 从环境变量构造 ProviderConfig。
///
/// 必填变量:
/// - `{prefix}_BASE_URL`
/// - `{prefix}_API_KEY`
/// - `{prefix}_MODEL`
///
/// 可选变量(有默认值):
/// - `{prefix}_TIMEOUT_SECS`(默认 30
/// - `{prefix}_MAX_RETRIES`(默认 3
pub fn from_env(prefix: &str) -> Result<Self, String> {
let base_url = std::env::var(format!("{prefix}_BASE_URL"))
.map_err(|_| format!("{prefix}_BASE_URL 环境变量未设置"))?;
let api_key = std::env::var(format!("{prefix}_API_KEY"))
.map_err(|_| format!("{prefix}_API_KEY 环境变量未设置"))?;
let model = std::env::var(format!("{prefix}_MODEL"))
.map_err(|_| format!("{prefix}_MODEL 环境变量未设置"))?;
let timeout_secs = match std::env::var(format!("{prefix}_TIMEOUT_SECS")) {
Ok(v) => v.parse().unwrap_or_else(|_| {
tracing::warn!("{prefix}_TIMEOUT_SECS='{v}' 解析失败,使用默认值 30");
30
}),
Err(_) => 30,
};
let max_retries = match std::env::var(format!("{prefix}_MAX_RETRIES")) {
Ok(v) => v.parse().unwrap_or_else(|_| {
tracing::warn!("{prefix}_MAX_RETRIES='{v}' 解析失败,使用默认值 3");
3
}),
Err(_) => 3,
};
// ponytail: max_retries 当前仅采集,不传入 Provider。
// 实际重试由 CycleConfig.retry.max_retries 控制。
// 此 warn 在应用启动时通常只触发一次,多次调用 from_env 时
// 重复输出的风险低。如有噪声,可改用 std::sync::Once 控制。
if max_retries != 3 {
tracing::warn!(
"ProviderConfig.max_retries={} 已采集但当前未生效;\
重试次数由 CycleConfig.retry.max_retries 控制",
max_retries,
);
}
Ok(Self {
base_url,
api_key,
model,
timeout_secs,
max_retries,
})
}
}
```
#### Timeout 传导模式
`create_provider()` 中,对基于 `GenericOpenaiProvider` 的 ProviderOpenAI / DeepSeek / Qwen / Ollama),通过同一模式注入 timeout:构造带 timeout 的 `Client` 后调用 `with_client(client)`
所有 OpenAI-compatible 分支新增的 `with_client()` 公开方法:
| Provider | 方法 | 位置 |
|----------|------|------|
| `OpenaiChatProvider` | 已有 `with_client(Client) -> Self` | `openai.rs` |
| `DeepSeekProvider` | 新增 `with_client(Client) -> Self` | `openai_compat.rs` |
| `QwenProvider` | 新增 `with_client(Client) -> Self` | `openai_compat.rs` |
| `OllamaProvider` | 新增 `with_client(Client) -> Self` | `ollama.rs`(新建文件) |
**关于 `new_with_client` 的说明**`DeepSeekProvider``QwenProvider` 当前已有测试用的 `new_with_client(base_url, api_key, model, client)` 方法(通过 `inner.http_client = client` 直接写字段)。新增 `with_client` 后,`new_with_client` 应重构为 `Self::new(base_url, api_key, model).with_client(client)` 代理,统一走公开 API 路径。
代码示例(以 DeepSeek 为例,OpenAI/Qwen/Ollama 模式完全一致):
```rust
ProviderType::DeepSeek => {
let client = Client::builder()
.timeout(Duration::from_secs(config.timeout_secs))
.build()
.map_err(|e| LlmError::Other(format!("创建 HTTP 客户端失败: {e}")))?;
Ok(Box::new(
openai_compat::DeepSeekProvider::new(
config.base_url,
config.api_key,
config.model,
)
.with_client(client),
))
}
```
Anthropic 由于需要保留 `default_headers`,使用独立的 `with_timeout` 模式:
AnthropicProvider 新增 `with_timeout` 方法:
```rust
impl AnthropicProvider {
/// 替换默认 HTTP Client 的超时配置。
///
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后原有通过 `with_client`
/// 注入的 Client 将被替换。headers 逻辑与 `new()` 中的构造保持一致。
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
// ponytail: 重建 http_client 时保留已有默认 headersx-api-key / anthropic-version)。
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
let key_header = HeaderValue::from_str(&self.api_key)
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
let version_header = HeaderValue::from_static("2023-06-01");
self.http_client = Client::builder()
.timeout(Duration::from_secs(secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
headers.insert("anthropic-version", version_header);
headers
})
.build()
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
Ok(self)
}
}
```
#### `map_reqwest_error` 中的硬编码超时修复
`openai.rs``anthropic.rs` 中的 `map_reqwest_error` 辅助函数当前在超时错误中返回硬编码的 `Duration::from_secs(120)`
```rust
// 现状 —— 硬编码 120s,与可配置 timeout 脱节
LlmError::Timeout { duration: Duration::from_secs(120) }
```
**修复方式**:采用**方案 A**——在 Provider struct 中存储 `timeout_secs` 字段,`map_reqwest_error` 读取该字段的值而非硬编码 120s。
```rust
// 修复后 —— 参数化,从 Provider 存储的 timeout_secs 读取
// GenericOpenaiProvider 新增 timeout_secs 字段:
pub struct GenericOpenaiProvider {
http_client: Client,
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
extra_headers: Vec<(String, String)>,
timeout_secs: u64, // ← 新增,由 new_with_name 的参数传入
}
// map_reqwest_error 使用 self.timeout_secs 而非硬编码 120
LlmError::Timeout { duration: Duration::from_secs(self.timeout_secs) }
```
**方案 B(从 reqwest::Client 提取 timeout)已被否决**`reqwest::Client` 不提供 timeout getter,无法从已构造的 client 中反向读取超时配置。
如果漏掉此修复,用户设置 `AG_LLM_TIMEOUT_SECS=60` 后超时,错误消息仍显示 "LLM 请求超时(120s",与实际配置不符。
---
#### max_retries 说明
`ProviderConfig.max_retries` 当前仅由 `from_env()` 采集存储,**实际重试操作由 `CycleConfig.retry.max_retries` 控制**。两者之间的关系通过文档注释声明:
```rust
/// 最大重试次数(默认 3)。当前此字段仅由 `from_env()` 采集,
/// 实际重试逻辑由 `CycleConfig.retry.max_retries` 控制。
/// 未来 Phase 6+ 可统一合并此字段到 CycleConfig。
```
#### 测试设计
使用 `temp_env` 在单元测试中隔离环境变量:
```rust
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn provider_config_from_env_requires_all_vars() {
// 未设置任何变量时应返回 Err
let result = ProviderConfig::from_env("TEST_PROVIDER");
assert!(result.is_err());
}
#[test]
fn provider_config_from_env_uses_defaults() {
temp_env::with_vars([
("TEST_PROVIDER_BASE_URL", Some("http://localhost:11434/v1")),
("TEST_PROVIDER_API_KEY", Some("")),
("TEST_PROVIDER_MODEL", Some("llama3")),
], || {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
assert_eq!(config.timeout_secs, 30);
assert_eq!(config.max_retries, 3);
});
}
#[test]
fn provider_config_from_env_reads_custom_timeout() {
temp_env::with_vars([
("TEST_PROVIDER_BASE_URL", Some("http://x")),
("TEST_PROVIDER_API_KEY", Some("k")),
("TEST_PROVIDER_MODEL", Some("m")),
("TEST_PROVIDER_TIMEOUT_SECS", Some("60")),
("TEST_PROVIDER_MAX_RETRIES", Some("5")),
], || {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
assert_eq!(config.timeout_secs, 60);
assert_eq!(config.max_retries, 5);
});
}
}
```
## 4. 实现计划
### Step 5.2 — Ollama Provider~55 行)
| 步骤 | 操作 | 验证 |
|------|------|------|
| 1 | 创建 `src/llm/provider/ollama.rs`,实现 `OllamaProvider` newtype | 编译通过 |
| 2 | 在 `provider.rs` 注册 `pub mod ollama` | 编译通过 |
| 3 | `ProviderType` 新增 `Ollama` 变体 | 编译通过 |
| 4 | `FromStr``"ollama"` 解析 | 编译通过 |
| 5 | `create_provider()``Ollama =>` 分支 | 编译通过 |
| 6 | 运行 `cargo build` | 无错误 |
### Step 5.3 — `#[non_exhaustive]` 前置标记(~4 行)
| 步骤 | 操作 | 验证 |
|------|------|------|
| 1 | `ProviderType``provider.rs`)加 `#[non_exhaustive]` | 编译通过 |
| 2 | `StopReason``response_v2.rs`)加 `#[non_exhaustive]` | 编译通过 |
| 3 | `FinishReason``shared.rs`)加 `#[non_exhaustive]` | 编译通过 |
| 4 | `EvictionPolicy``memory/store.rs`)加 `#[non_exhaustive]` | 编译通过 |
| 5 | 运行 `cargo build --all-targets` | 无 warning |
### Step 5.1 — ProviderConfig 扩展(~60 行 + 测试)
| 步骤 | 操作 | 验证 |
|------|------|------|
| 1 | `ProviderConfig``timeout_secs` / `max_retries` 字段 | 编译通过 |
| 2 | 实现 `impl Default for ProviderConfig` | 编译通过 |
| 3 | 实现 `ProviderConfig::from_env()` | 编译通过 |
| 4 | `GenericOpenaiProvider``AnthropicProvider` 新增 `timeout_secs` 字段,`new_with_name``new()` 接受 timeout 参数 | 编译通过 |
| 5 | `map_reqwest_error` 在各 Provider 中改为从 `self.timeout_secs` 读取,移除硬编码 120s | 编译通过 |
| 6 | `create_provider()` 中各分支注入 timeoutOpenAI-compatible 用 `Client::builder().timeout()` + `with_client`Anthropic 用 `with_timeout()` | 编译通过 |
| 7 | `DeepSeekProvider``QwenProvider``new_with_client` 重构为 `Self::new(...).with_client(client)` 代理 | 测试通过 |
| 8 | `Cargo.toml` 添加 `temp_env` dev-dependency | `cargo build` 通过 |
| 8 | 添加 `from_env` 单元测试 + timeout 传导集成测试 | `cargo test` 通过 |
| 9 | 完整验证 | 见第 6 节 |
## 5. 风险评估
| 风险 | 影响 | 概率 | 缓解措施 |
|------|------|------|----------|
| `create_provider()``Client::builder().build()` 返回 `Result`,当前代码使用 `.expect()`,改为 `map_err` 转为 `LlmError` 后需确保所有分支正确转换 | 编译期强制处理,遗漏分支直接报错 | 低 | `create_provider` 返回 `Result<Box<dyn LlmProvider>, LlmError>``map_err` 天然适配。新增的 timeout 注入路径逐一检查 |
| `AnthropicProvider``default_headers``with_timeout` 中重建时与 `new()` 中的 headers 不一致 | Anthropic 认证失败 | 低 | `with_timeout` 方法复制 `new()` 中的 headers 构造逻辑。通过已有测试验证认证通过 |
| Ollama 实际运行时行为差异:版本兼容性、API 路径、模型名等 | 运行时才能发现 | 中 | Phase 5 仅做类型级验证(`cargo build`),Phase 8 端到端测试时通过 Ollama mock 或真实实例验证 |
| `max_retries` 存储了却未实际使用,造成困惑 | 开发者误以为已生效 | 中 | 通过文档注释明确声明 `max_retries` 当前仅采集,实际重试由 `CycleConfig.retry.max_retries` 控制 |
| `temp_env` 测试在多线程并发测试中互相污染环境变量 | 偶发测试失败 | 中(Rust 默认单线程测试用 `--test-threads=1` 可避免) | 将 `from_env` 测试控制在同一测试文件,避免并行执行。必要时在 CI 中确保 `--test-threads=1` |
## 6. 验收标准
以下条件**全部满足**方可认为 Phase 5 完成:
- [ ] `cargo build --all-targets` 通过,无错误
- [ ] `cargo test --all-targets` 通过,新增测试覆盖 `from_env` 的必填/选填/默认值场景
- [ ] `cargo clippy --all-targets -- -D warnings` 通过,无任何 warning
- [ ] `cargo doc --no-deps -D warnings` 通过,所有公共 API 有文档注释(`///`
- [ ] 新增文件:1`ollama.rs`
- [ ] 修改文件:9`provider.rs``openai.rs``anthropic.rs``openai_compat.rs``response_v2.rs``shared.rs``store.rs``Cargo.toml`、测试文件)
- [ ] 净代码增量:~160 行
- [ ] `ProviderType` 新增 `Ollama` 变体,`"ollama"` 字符串可解析
- [ ] 4 个公共枚举带有 `#[non_exhaustive]` 属性
- [ ] `ProviderConfig` 可从环境变量构造(`from_env()`),含默认值
- [ ] timeout 值已传导到 `create_provider()` 中各 Provider 的 HTTP Client 配置
- [ ] timeout 传导验证通过至少一个端到端 wiremock 集成测试(模拟 HTTP 服务在超时后返回 408,验证 Provider 返回 `LlmError::Timeout`
- [ ] `DeepSeekProvider``QwenProvider``OllamaProvider` 均有公开 `with_client()` 方法,可在 `create_provider` 中注入 timeout Client
- [ ] `map_reqwest_error` 中不再硬编码 `Duration::from_secs(120)`,改为参数化读取
+236
View File
@@ -0,0 +1,236 @@
# Phase 6 — ToolDefinition IR 正式化实施方案
## 背景与目标
在 agcore v0.2 路线图中,Phase 6 旨在引入 `ToolDef` 新类型,替换已标记 `#[deprecated(since = "0.1.0")]``ToolDefinition`(即 `OpenaiToolDefinition` 类型别名),消除 OpenAI wire format 对核心类型系统的泄漏,建立 Provider 无关的工具定义中间表示(IR)。
预期成果:
- 核心类型系统不再直接依赖 `OpenaiToolDefinition`
- 所有 Provider 适配层从统一的 `ToolDef` IR 出发,各自转换为对应 wire format
- 消除 `#[allow(deprecated)]` 抑制点,恢复 clippy 零警告状态
## 当前状态分析
当前代码库中工具定义相关的关键状态如下:
1. **`OpenaiToolDefinition` 结构体**定义于 `llm/types/tool.rs`,包含 4 个字段:
- `name: String`
- `description: Option<String>`
- `parameters: Value`
- `strict: Option<bool>`
当前存在多处 `#[allow(deprecated)]` 抑制点,分布在 `llm/cycle.rs``tools/registry.rs``tools/mcp.rs``agent/agent.rs` 等文件中。
2. **`ToolDefinition` 类型别名**定义于 `llm/types/mod.rs:105`,标记为 `#[deprecated]`
```rust
#[deprecated(since = "0.1.0", note = "use OpenaiToolDefinition directly")]
pub type ToolDefinition = OpenaiToolDefinition;
```
3. **`MessageRequest.tools` 字段**类型为 `Vec<OpenaiToolDefinition>`(直接引用原始类型,而非别名)。
4. **引用该类型的 4 个源文件**:
- `llm/cycle.rs`4 个方法参数使用 `Vec<ToolDefinition>`
- `tools/registry.rs``definitions() -> Vec<ToolDefinition>` 返回类型 + struct literal 构造(含 `strict: None`
- `tools/mcp.rs``list_tools() -> Vec<ToolDefinition>` 返回类型 + struct literal 构造(含 `strict: None`
- `agent/agent.rs``fn tool_definitions() -> Vec<ToolDefinition>` trait 默认实现
5. **Provider 适配层**`openai.rs` 和 `anthropic.rs` 从 `MessageRequest.tools` 读取数据并转换为各自的 wire format。`openai_compat.rs` 和 `ollama.rs` 委托给 `GenericOpenaiProvider`,无需直接改动。
## 需求推演
### 决策 1`strict` 字段的处理
当前所有构造路径均硬编码 `strict: None``registry.rs`、`mcp.rs`),`BaseTool` trait 无 `strict` 方法,用户 API 无法设置该值。
**结论:移除。** `strict` 是 OpenAI 的 Structured Outputs 专属字段,不属于 Provider 无关的 IR。未来如需支持,走 `MessageRequest.extra` 逃生舱,在各 Provider 适配层自行消费。
### 决策 2`description` 保持 `Option<String>`
Anthropic 要求 `description` 为必填(`String`),但 MCP 等来源可能缺失该字段。保持 `Option`,由 Anthropic 适配层以 `unwrap_or_default()` 兜底。
### 决策 3`parameters` 保持 `Value`
所有 Provider 的 wire format 均接受 JSON Schema 格式的 Value。当前不做 typed 方案,保留 `Value`。
### 决策 4:采用直接切断而非阶段性 deprecation
v0.1 已标记 `#[deprecated]`,用户已有预期。pre-1.0 阶段的 breaking change 是合理的。`OpenaiToolDefinition` 保留但降级为 `#[doc(hidden)]`。
## 方案设计
### ToolDef 结构体
位置:`src/llm/types/tool.rs`(与 `OpenaiToolDefinition` 同文件,不建独立文件/模块)。
```rust
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ToolDef {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default)]
pub parameters: Value,
}
```
实现双向 `From` 转换:
```rust
impl From<ToolDef> for OpenaiToolDefinition {
fn from(t: ToolDef) -> Self {
Self {
name: t.name,
description: t.description,
parameters: t.parameters,
strict: None,
}
}
}
impl From<OpenaiToolDefinition> for ToolDef {
fn from(t: OpenaiToolDefinition) -> Self {
Self {
name: t.name,
description: t.description,
parameters: t.parameters,
}
}
}
```
serde 属性与 `OpenaiToolDefinition` 原有属性一致,保证 JSON 序列化兼容。
明确不做:
- builder 模式(Rust struct literal + `..Default::default()` 已足够)
- `#[non_exhaustive]`(IR 类型自有完整控制权,不需要)
- 独立文件(7 行 struct 无需独立模块)
### 4 单元切割计划
每步设计为可编译的安全 checkpoint。
#### 单元 6.1 — 新增 ToolDef + From 实现
| 项目 | 内容 |
|------|------|
| 涉及文件 | `llm/types/tool.rs` |
| 变更内容 | 新增 `ToolDef` struct(约 7 行)、2 个 `From` impl(约 12 行) |
| 验证标准 | `cargo build` 编译通过(旧代码照常编译,零影响) |
| 检查点 | 新类型存在但未被消费,安全 checkpoint |
#### 单元 6.2 — 别名切换 + 构造同步修复
| 项目 | 内容 |
|------|------|
| 涉及文件 | `llm/types/mod.rs`、`llm/types/request_v2.rs`、`llm/types/request.rs`、`tools/registry.rs`、`tools/mcp.rs` |
| 变更内容 | 切换别名 `pub type ToolDefinition = ToolDef``MessageRequest.tools` 改为 `Vec<ToolDef>`registry/mcp 构造去掉 `strict: None` |
| 验证标准 | `cargo build` 编译通过 |
| 风险提示 | `cycle.rs` 方法参数使用别名,自动生效无需修改;`agent/agent.rs` trait 默认实现使用别名,自动适配;暂不移除 `#[allow(deprecated)]` |
#### 单元 6.3 — Provider 适配
| 项目 | 内容 |
|------|------|
| 涉及文件 | `llm/provider/openai.rs` |
| 变更内容 | `convert_request()` 中 `tool_defs.into_iter().map(|t| OpenaiTool::Function { function: t })` → 改为 `.map(|t| OpenaiTool::Function { function: t.into() })``t` 类型从 `OpenaiToolDefinition` 变为 `ToolDef`,需 `Into` 转换 |
| 验证标准 | `cargo test --all-targets` 全部通过 |
| 不修改的文件 | `anthropic.rs`(同名字段访问自动适配)、`openai_compat.rs`/`ollama.rs`(委托给 `GenericOpenaiProvider` |
#### 单元 6.4 — 清理
| 项目 | 内容 |
|------|------|
| 涉及文件 | `llm/cycle.rs`、`tools/registry.rs`、`tools/mcp.rs`、`agent/agent.rs`、`llm/types/mod.rs`、`llm/types/tool.rs` |
| 变更内容 | 移除所有与 `ToolDefinition` 相关的 `#[allow(deprecated)]``llm/types/mod.rs` 移除旧 `#[deprecated]` 别名(仅保留 `pub use tool::ToolDef`);`OpenaiToolDefinition` 降级为 `#[doc(hidden)]`。精确列表由 `cargo clippy -D warnings` 检出——clippy 会标记所有不再需要的 `#[allow]` |
| 验证标准 | `cargo clippy --all-targets -- -D warnings` 零警告;`cargo test --all-targets` 全部通过 |
| 无需改动 | 测试代码(无一直接引用 `ToolDefinition` |
### Provider 适配策略
Provider 适配层改动最小化,仅在序列化入口处加一层 `From` 转换:
| Provider | 适配方式 | 改动 |
|----------|---------|------|
| OpenAI`GenericOpenaiProvider` | `MessageRequest.tools: Vec<ToolDef>` → lambda 内改为 `.map(\|t\| OpenaiTool::Function { function: t.into() })`,将 `ToolDef` 通过 `Into` 转为 `OpenaiToolDefinition` | `convert_request` lambda 内 +`.into()` |
| Anthropic | `t.name` / `t.description` / `t.parameters` 字段名不变,直接访问 | 零改动 |
| OpenAI CompatDeepSeek、Qwen | 委托给 `GenericOpenaiProvider` | 零改动 |
| Ollama | 委托给 `GenericOpenaiProvider` | 零改动 |
### 变更清单汇总
| 文件 | 改动类型 | 估计行数 |
|------|---------|---------|
| `llm/types/tool.rs` | +`ToolDef` + 2x `From` | +19 |
| `llm/types/mod.rs` | 改别名 + re-export | ~3 |
| `llm/types/request_v2.rs` | 改 `tools` 字段 + import | ~2 |
| `llm/cycle.rs` | 移 `#[allow(deprecated)]` | -1 |
| `tools/registry.rs` | 构造去掉 `strict` + 移 `allow` | ~4 |
| `tools/mcp.rs` | 构造去掉 `strict` + 移 `allow` | ~4 |
| `agent/agent.rs` | 移 `#[allow(deprecated)]` | -1 |
| `llm/provider/openai.rs` | `convert_request` 加 `.map(Into::into)` | +2 |
| 新增 roundtrip 测试 | `MessageRequest` 序列化 roundtrip 验证 | +15 |
| **合计** | | **约 47 行(+ 约 15 行测试)** |
### 测试策略
新增一条 `MessageRequest` 序列化 roundtrip 测试,覆盖 `ToolDef` 的 JSON 序列化/反序列化兼容性。该测试验证 `ToolDef` 的 serde 属性与 `OpenaiToolDefinition` 一致,确保 wire format 兼容。
## 实施步骤
1. 新建分支 `phase-6-tooldef-ir`
2. 按单元 6.1 → 6.2 → 6.3 → 6.4 顺序执行,每步提交一个 commit
3. 每步执行对应的验证标准
4. 全量通过后创建 PR
```
git checkout -b phase-6-tooldef-ir
# 执行单元 6.1 → commit
# 执行单元 6.2 → commit
# 执行单元 6.3 → commit
# 执行单元 6.4 → commit
# 全量验证
```
### 用户迁移指引
Phase 6 涉及公共 API 类型替换,下游用户升级到 v0.2 时需注意:
| 旧用法 | 新用法 |
|--------|--------|
| `use agcore::llm::types::ToolDefinition` | `use agcore::llm::types::ToolDef`(别名已移除) |
| `use agcore::llm::types::OpenaiToolDefinition` | `use agcore::llm::types::ToolDef``OpenaiToolDefinition` 已降级为 `#[doc(hidden)]` |
| 直接构造 `ToolDefinition { strict: None, .. }` | 构造 `ToolDef { .. }`(去掉 `strict` 字段) |
`OpenaiToolDefinition` 仍保留但标记 `#[doc(hidden)]`,极端情况仍需使用时可通过全路径访问。
## 验证标准
| 阶段 | 验证命令 |
|------|---------|
| 单元 6.1 | `cargo build` 编译通过 |
| 单元 6.2 | `cargo build` 编译通过(新旧代码全量编译) |
| 单元 6.3 | `cargo test --all-targets` 全部通过 |
| 单元 6.4 | `cargo clippy --all-targets -- -D warnings` 零警告;`cargo test --all-targets` 全部通过 |
| 最终 | `cargo build --all-targets` + `cargo test --all-targets` + `cargo clippy --all-targets -- -D warnings` 全绿 |
## 风险与缓解
| 风险 | 说明 | 缓解措施 |
|------|------|---------|
| Struct literal 断层 | Step 6.2 切别名与构造修复若不同步,registry/mcp 中使用 `OpenaiToolDefinition` struct literal 的构造代码会编译失败 | 别名切换与构造修复合并在同一单元,原子化提交 |
| 遗漏 `#[allow(deprecated)]` | 部分抑制点因 grep 遗漏而未在 6.4 移除 | clippy `-D warnings` 可检出;6.4 前做一次全库 grep 确认无遗漏 |
| JSON 兼容性 | `ToolDef` serde 属性与 `OpenaiToolDefinition` 不一致导致 wire format 变化 | `ToolDef` serde 属性与 `OpenaiToolDefinition` 保持一致;roundtrip 测试验证 |
| Provider 适配遗漏 | 部分 Provider 分支未经测试覆盖 | `cargo test --all-targets` 包含 Provider 测试 |
## 否决记录
| 否决方案 | 原因 |
|---------|------|
| 保留 `strict` 字段 | OpenAI 专属字段,当前所有构造路径传 `None`。不属于 Provider 无关的 IR。未来支持走 `MessageRequest.extra` 逃生舱 |
| 逐步 deprecation 过渡 | pre-1.0 阶段 breaking change 合理,v0.1 已标记 deprecation,用户已有预期 |
| `ToolDef` 建独立文件 | 约 7 行的 struct 不需要独立文件,与 `OpenaiToolDefinition` 共享 `types/tool.rs` 即可 |
| `ToolDef` 放在 `tools/` 模块 | 会创造 `llm` → `tools` 的逆向依赖,破坏模块分层 |
| 添加 builder 模式 | Rust struct literal + `..Default::default()` 已足够覆盖使用场景 |
| 添加 `#[non_exhaustive]` | IR 类型自有完整控制权,不需要对外隐藏字段 |
| 同 Phase 净化 `parameters` 类型化 | 属于独立工作,留给 v0.3+ 阶段处理 |
+526
View File
@@ -0,0 +1,526 @@
# Phase 7 — SqliteStore 持久化实现方案
- **文档编号**14
- **标题**Phase 7 — SqliteStore 持久化实现方案
- **日期**2026-07-05
- **状态**:已定稿
- **涉及模块**memory/store
- **关联文档**roadmap.md, 6-memory-system.md
---
## 背景与目标
Phase 7 的核心任务是完成 MemoryStore trait 的 SQLite 后端实现,使 Agent 进程重启后记忆数据不丢失。这是 v0.2.0 从"内存玩具"走向"可用工具"的关键门槛,也是后续 Phase 8MVP 出口)和 Phase 10ContextSlot)的前置依赖。
**成功标准**
- SqliteStore 完整实现 MemoryStore trait4 个方法:save/get/delete/list
- 进程关闭后重新打开同一数据库文件,数据完整可读
- 与现有 InMemoryStore 通过 MemoryStore trait 可互换,消费者零改动
- 所有现有测试保持通过,clippy 0 警告
### Scope & Non-goals
| 范围 | 内容 |
|------|------|
| 包含 | 单表 CRUD + prefix/since/offset/limit 查询 + WAL 并发 + Mutex 串行化 + 错误映射 |
| 不包含(Phase 7 | 淘汰策略(EvictionPolicy,仅 InMemoryStore 持有,需 v0.3 纳入 SqliteStore |
| 不包含(Phase 7 | Schema 迁移框架(PRAGMA user_version 足矣,不引入 refinery/sea-query |
| 不包含(Phase 7 | 批量写入 / 事务 API(N+1 clear 延迟可接受,优化后置) |
| 不包含(Phase 7) | 跨进程共享同一数据库文件(Mutex 为单进程设计) |
---
## 当前状态分析
### 现有实现
- MemoryStore trait 已在 v0.1 Phase 3 就绪,定义 4 个异步方法
- InMemoryStore 实现稳定运行,使用 `Mutex<HashMap>` 作为后端
- 全量测试 191 个通过,clippy 0 警告
- 项目当前无 SQLite 或其他数据库依赖
### 现有消费者
通过 `crate::memory::store::MemoryStore` 路径引用的模块:
| 模块 | 文件 | 使用方式 |
|------|------|----------|
| Agent Builder | `agent/builder.rs` | RuntimeBundle 中引用 MemoryStore |
| Session Memory | `agent/session_memory.rs` | SessionMemory 实现 |
| Agent Runtime | `agent/runtime.rs` | 类型标注 |
| Agent Session | `agent/session.rs` | 默认 InMemoryStore 兜底 |
| Conversation | `memory/conversation.rs` | ConversationMemory 测试 |
| Knowledge | `memory/knowledge.rs` | KnowledgeStore 测试 |
| Retriever | `memory/retriever.rs` | MemoryRetriever 测试 |
所有消费者均通过 `MemoryStore` trait 访问,不依赖具体实现类型,因此新增 SqliteStore 不会产生编译或运行时影响。
### 目录结构现状
```
src/memory/
├── mod.rs
├── store.rs ← 包含 MemoryStore trait + InMemoryStore + EvictionPolicy
├── conversation.rs
├── knowledge.rs
├── retriever.rs
└── vector.rs
```
`store.rs` 目前是一个单体文件,同时承载 trait 定义和 InMemoryStore 实现。
---
## 调研发现
### MemoryStore trait 定义
```rust
#[async_trait]
pub trait MemoryStore: Send + Sync {
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError>;
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError>;
async fn delete(&self, id: &str) -> Result<(), MemoryError>;
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError>;
}
```
### 关键类型
| 类型 | 定义 |
|------|------|
| `MemoryItem` | `{ id: String, content: String, metadata: Value, created_at: OffsetDateTime }` |
| `MemoryFilter` | `{ prefix: Option<String>, since: Option<OffsetDateTime>, offset: Option<usize>, limit: Option<usize> }` |
| `MemoryError` | 变体:`NotFound` / `Storage` / `Serialization` / `InvalidInput` / `RetrievalError` |
| `EvictionPolicy` | `None` / `Ttl { ttl_secs }` / `Capacity { max_items }` |
| `EvictionConfig` | `{ policy, check_interval }` |
### 并发模型参考
InMemoryStore 当前使用 `Mutex<HashMap>` 实现 `Send + Sync`。SqliteStore 将遵循相同模式,使用 `Arc<Mutex<Connection>>` + `spawn_blocking` 满足异步 trait 约束。
### Schema 设计考虑
- `created_at` 使用 TEXT(ISO 8601) 存储——`.to_string()` 零转换,字典序与时间序一致(前提:所有时间戳归一化到 UTC;`OffsetDateTime::to_string()` 在 UTC 下输出 `"2026-07-05T12:00:00Z"` 格式,字典序与时间序严格对应)
- 初始 schema 即创建 `created_at` 索引,避免后续大数据量全表排序
- Schema 版本管理通过 `PRAGMA user_version` 实现,零外部依赖,后续加字段只需追加 `if version < N { ALTER TABLE }`
---
## 可选方案
### A. rusqlite + Mutex\<Connection\>(推荐)
| 维度 | 评估 |
|------|------|
| 新增依赖 | 1 个(rusqlite 0.32 + bundled features |
| 实现量 | ~200 行 |
| SQL 支持 | 原生支持 prefix LIKE 过滤 + ORDER BY 排序 |
| 性能 | 有索引时查询 O(log n),写入串行化 |
| 并发 | WAL 模式 + Mutex 串行化写入,适合单进程 Agent |
| 事务支持 | 完整 ACID |
| 崩溃安全 | WAL 模式,崩溃恢复有保障 |
**适用场景**:单进程 Agent 本地持久化、嵌入式场景、需要关系查询能力的通用存储。
**外部依赖评估**
- rusqlite 0.32 — 最新稳定版(2025-12 发布),维护活跃(月均 2+ 次提交),Apache-2.0 许可证
- `bundled` feature 编译 SQLite 源码(Public Domain)进二进制,无系统级 SQLite 依赖,零外部 C 库安装步骤
- 供应链风险:bundled 模式依赖 crate 发布节奏同步 SQLite 安全更新;SQLite 安全公告频率极低(年均 <3 例),此风险可接受
### B. JSONL 文件
| 维度 | 评估 |
|------|------|
| 新增依赖 | 0 |
| 实现量 | ~150 行 |
| get() 复杂度 | O(n) 全量扫描 |
| delete() 复杂度 | O(n) 全量重写 |
| 并发 | 需文件锁(flock |
| 事务支持 | 无 |
| 崩溃安全 | 无保障,写入中断可能丢失或损坏数据 |
**否决理由**:核心的 get() 查询场景不可接受 O(n) 性能;在 Agent 运行时频繁读写记忆的场景下,全量扫描的成本会随着数据积累线性增长,不符合可用性要求。
### C. sled 嵌入式 KV
| 维度 | 评估 |
|------|------|
| 新增依赖 | 1 个(纯 Rust |
| 实现量 | ~150 行 |
| prefix scan | 原生支持 |
| 排序 | 需手动实现 |
| 关系模型 | 不如 SQL 匹配当前查询模式 |
| 社区成熟度 | 较新,API 仍在演进 |
**否决理由**:当前查询模式(prefix 过滤 + 按 created_at 排序)在关系模型中用一条 SQL 即可表达,引入 KV 存储反而需要手动处理排序逻辑。非必要不引入新存储范式。
---
## 推荐方案
### 总体方向:方案 Arusqlite + Mutex\<Connection\>
选择理由:
1. **最少依赖,最高匹配**:1 个新增依赖即可完整支持 MemoryFilter 的所有查询维度(prefix LIKE、created_at 范围、offset/limit
2. **生产就绪**rusqlite 是 SQLite 的 Rust 绑定事实标准,bundled 模式免去系统 SQLite 依赖
3. **Schema 演进简单**PRAGMA user_version + 逐版本迁移,零外部迁移工具依赖
4. **与 InMemoryStore 语义一致**Mutex 串行化 + spawn_blocking 适配 async trait,与现有并发模型同构
### 关键设计决策
| 决策 | 选择 | 理由 |
|------|------|------|
| Schema 版本管理 | PRAGMA user_version | 零外部依赖,~20 行,后续 ALTER TABLE 即可 |
| created_at 存储格式 | TEXT(ISO 8601) + UTC 归一化 | 零转换代码;UTC 下输出 `"2026-07-05T12:00:00Z"`,字典序与时间序严格一致 |
| Upsert SQL 策略 | `INSERT ... ON CONFLICT(id) DO UPDATE SET ...` | 保留调用方传入的 `created_at`,避免被 `DEFAULT` 覆盖 |
| 性能索引 | 初始 schema 加 created_at 索引 | 避免大数据量全表排序 |
| 配置参数 | 仅 `path``busy_timeout=5s` | 其余内置默认值;5s 超时避免 `SQLITE_BUSY` 快速失败 |
| Mutex 中毒恢复 | `.lock().unwrap_or_else(\|e\| e.into_inner())` | 不 panic,恢复执行 |
| 批量操作 | 不加 | N+1 clear ~250ms(N=50),可接受,优化后置 |
| spawn_blocking 取消安全性 | 短事务模式(auto-commit) | 每个操作独立事务,取消时后台 task 自然完成/panic,不 Cross 操作持有 Mutex |
**我们放弃了什么**(集中 Trade-off 记录):
- **写入串行化**`Mutex<Connection>` 确保 SQLite 写入安全,代价是同一时刻只能有一个写入者。Agent 场景下写入频率低(每次 LLM 调用触发 1-2 次),串行化不构成瓶颈
- **单进程锁**:无法跨进程共享同一数据库文件。多进程场景需要网络后端(PostgreSQL/Redis
- **无横向扩展**:单文件 SQLite 无分片能力。需扩展时切换到分布式后端
### Schema 定义(初始版本)
```sql
CREATE TABLE IF NOT EXISTS memory_items (
id TEXT PRIMARY KEY,
content TEXT NOT NULL,
metadata TEXT NOT NULL DEFAULT '{}',
created_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_memory_items_created_at
ON memory_items(created_at);
```
### 数据完整性防御
| 异常场景 | 防御措施 | 错误映射 |
|---------|---------|---------|
| 数据库文件损坏 | `migrate()` 中执行 `PRAGMA quick_check`;失败时 `open()` 返回 `MemoryError::Storage` | `Storage` |
| created_at 解析失败 | `get()`/`list()``OffsetDateTime::parse` 失败不 panic,返回 `MemoryError::Serialization` | `Serialization` |
| content / metadata 为 NULL | `get()` 中检测 SQLite 返回值,NULL 时返回 `MemoryError::Storage` | `Storage` |
| 约束冲突(PRIMARY KEY / NOT NULL | 映射为 `MemoryError::InvalidInput` | `InvalidInput` |
| 序列化/反序列化失败 | `serde_json::to_string`/`from_str` 错误映射为 `MemoryError::Serialization` | `Serialization` |
### 性能预算(目标延迟,单条操作)
| 操作 | 目标延迟 | 说明 |
|------|---------|------|
| `save(1KB item)` | < 5ms | 含 serde_json 序列化 + spawn_blocking + SQLite INSERT |
| `get(1KB item)` | < 3ms | 含 SQLite SELECT + 反序列化 |
| `list(prefix 匹配 100 行)` | < 20ms | 含索引 B-tree 遍历 + ORDER BY + LIMIT |
| 并发 10 writer | p99 < 50ms | Mutex 串行化排队,每 writer 等待 9×5ms 内 |
实施后通过 Step C 测试验证以上预算。未达标时不阻塞发布,但记录为可观测告警阈值。
---
## 实施建议
### 实施计划
#### Step A — 目录重构(纯搬移,零行为变化)
目标:将单体 `store.rs` 拆分为模块目录架构,为新增 SqliteStore 做准备。
```
src/memory/
├── store.rs ← 模块根:MemoryStore trait + EvictionPolicy/EvictionConfig
│ + pub mod in_memory;
│ + pub mod sqlite_store;
│ + pub use in_memory::InMemoryStore;
├── store/
│ ├── in_memory.rs ← InMemoryStore 提取至此(struct + impl + 6 个内联测试)
│ └── sqlite_store.rs ← 新增 SqliteStore
```
模式参考:`llm/provider.rs``llm/provider/{openai,anthropic,ollama}.rs`
**外部消费者的导入路径不变**`crate::memory::store::MemoryStore`),零改动风险。
重构步骤:
1. 创建 `src/memory/store/` 目录
2. 创建 `src/memory/store/in_memory.rs`,从原 `store.rs` 提取 InMemoryStore 全部代码(struct + impl + Default + 6 个测试)
3. 修改 `src/memory/store.rs`:保留 MemoryStore trait + EvictionPolicy/EvictionConfig,加 `pub mod in_memory;` + `pub use in_memory::InMemoryStore;`
4. 验证:`cargo test --all-targets` 全绿,测试数量不变(191 pass)
#### Step B — SqliteStore 实现
1. `Cargo.toml` 添加 `rusqlite = { version = "0.32", features = ["bundled"] }`
2. 创建 `src/memory/store/sqlite_store.rs`,实现:
- `SqliteStore` 结构体:`{ conn: Arc<Mutex<Connection>> }`
- `SqliteStore::open(path)` 构造函数,支持 `":memory:"`
- `lock_conn()` 辅助方法(Mutex 中毒恢复)
- `migrate()` Schema 初始化 + 版本管理
- `MemoryStore` trait 的 4 个方法
- `From<rusqlite::Error> for MemoryError`
3. 修改 `src/memory/store.rs`:加 `pub mod sqlite_store;` + `pub use sqlite_store::SqliteStore;`
4. 修改 `src/memory.rs`:加 `pub use store::SqliteStore;`
5. 编写测试覆盖:
- CRUD 基本操作
- Upsert(同 id 重复 save 覆盖)
- prefix 过滤
- 由于/until 时间范围过滤
- 并发 10 个 writer × 10 次操作
- 持久化恢复(write → drop → reopen → read
6. 验证:`cargo test --all-targets` 全绿 + `cargo clippy --all-targets -- -D warnings` 0 警告
#### Step C — 验证确认
1. 确认现有 memory 模块内测试全部通过
2. 确认 agent/llm/tools/prompt 模块不受影响
3. 确认 SqliteStore 与 InMemoryStore 通过 MemoryStore trait 可互换
4. 确认 clippy 无新增警告
### Commit 安排
| 顺序 | 类型 | Scope | 描述 |
|------|------|-------|------|
| 1 | refactor | memory | 将 store.rs 拆分为模块目录,仅结构搬移 |
| 2 | feat | memory | 实现 SqliteStore 持久化 |
### 风险与缓解
| 风险 | 严重度 | 缓解措施 |
|------|--------|----------|
| Mutex 中毒导致后续操作全部失败 | 中 | `lock_conn()` 使用 `.lock().unwrap_or_else(\|e\| e.into_inner())` 恢复模式,不 panic |
| spawn_blocking 取消后连接状态不一致 | 中 | 每个操作使用短事务(auto-commit),不跨操作持有 Mutex;取消时遗留 task 自然完成或 panicMutex 通过 `.into_inner()` 恢复 |
| WAL 文件无限增长 | 低 | 内置 auto-checkpoint 阈值 + 启动时执行 `PRAGMA wal_checkpoint(TRUNCATE)` |
| list 无索引导致全表扫描 | 中(大数据量) | 初始 schema 即创建 `idx_memory_items_created_at` 索引 |
| 父目录不存在导致 open 失败 | 低 | `open()` 内部调用 `fs::create_dir_all()` 确保目录存在 |
| clear() N+1 删除性能 | 低 | 不走 trait 接口的逐条删除,可后续优化为直接 `DELETE FROM memory_items` |
| 数据库文件损坏 | 低 | `migrate()` 中执行 `PRAGMA quick_check`;失败时返回 `MemoryError::Storage`,调用方可切换 InMemoryStore |
### 可观测性(实施时落实)
- 所有 MemoryStore 方法通过 `tracing::instrument` 记录延迟和结果(`info!` 正常完成,`warn!` 超过性能预算阈值,`error!` 操作失败)
- `list()` 返回行数通过 `tracing::debug` 记录(调优参考)
- WAL 文件大小在 `migrate()` 后检查一次,超过 100MB 时记录 `warn!`
- 操作计数(读写次数、错误率)暂不暴露为独立 metrics,v0.3 按需添加
---
## 已知假设
| 假设 | 验证状态 | Fallback |
|------|---------|----------|
| 单进程独享 SQLite 文件,无跨进程竞争 | ✅ 设计前提(Mutex 为单进程设计) | 多进程场景使用网络后端(PostgreSQL/Redisv0.3+ |
| ISO 8601 TEXT 字典序等价于时间序 | ✅ 条件成立(需 UTC 归一化) | 若时区异常,切换 INTEGER(unix_timestamp) 存储后重建索引 |
| SqliteStore 初始化失败可 fallback 到 InMemoryStore | ✅ 调用方自行控制 | `open()` 返回 `MemoryError`,消费者 catch 后改用 `InMemoryStore::new()` |
| N+1 clear ~250ms(N=50) 可接受 | 🟡 未实测(基于 N×5ms 推算) | 若成为瓶颈,SqliteStore 内部加 `delete_by_prefix()` 方法(不走 trait 接口) |
| rusqlite bundled SQLite 版本足够新 | ✅ 0.32 版内置 SQLite 3.46 | 如需特定版本,切换 `bundled` 为指定版本或使用系统 SQLite |
| busy_timeout=5s 覆盖所有竞争场景 | 🟡 未实测(WAL 下写写冲突概率低) | 若观测到 `SQLITE_BUSY`,增大超时或在重试逻辑中处理 |
| spawn_blocking 线程池不会被耗尽 | ✅ 默认 512 线程,Agent 场景占用 ≤10 | 若观测到阻塞任务排队,启动时 `tokio::task::spawn_blocking` 已有兜底排队机制 |
---
## 参考来源
- [rusqlite crate](https://crates.io/crates/rusqlite) — 官方文档
- [SQLite PRAGMA user_version](https://www.sqlite.org/pragma.html#pragma_user_version) — Schema 版本管理机制
- [SQLite WAL mode](https://www.sqlite.org/wal.html) — 并发读写性能优化
- `docs/6-memory-system.md` — Phase 3 MemoryStore trait 原始设计
- `docs/roadmap.md` — 项目里程碑规划(Phase 7/8/10 依赖关系)
- `src/llm/provider.rs``src/llm/provider/` — 目录重构模式参考
---
## 实施计划
### 任务总览
3 个阶段、8 个任务单元、2 个 Commit。
### 阶段一:目录重构
#### Task A1 — 创建 store/ 目录并提取 InMemoryStore
| 项目 | 内容 |
|------|------|
| 任务描述 | 创建 `src/memory/store/` 目录,新建 `store/in_memory.rs`,从 `store.rs` 完整提取 InMemoryStore 结构体、impl MemoryStore、impl Default、6 个内联测试 |
| 涉及文件 | `src/memory/store.rs` → 分割到 `src/memory/store/in_memory.rs`(新增) |
| 前置依赖 | 无 |
| 预估工作量 | S< 1h |
| 风险等级 | 低 — 纯搬移,编译器可验证 |
| 验收条件 | `cargo build` 通过(此时 store.rs 尚未修改,store/in_memory.rs 应被 crate 忽略) |
注意:需要先在 store.rs 顶部添加 `pub mod in_memory;` 声明,否则子模块不会被编译。或者可以先创建目录和文件,等 Task A2 再统一加声明路径。
实际做法:先创建文件但不声明,A2 统一声明。这样 A1 和 A2 之间可以有一个无编译的中间状态。
#### Task A2 — 修改 store.rs 模块根
| 项目 | 内容 |
|------|------|
| 任务描述 | 修改 `store.rs` 为纯模块根:保留 `MemoryStore` trait、`EvictionPolicy``EvictionConfig`;添加 `pub mod in_memory;` + `pub use in_memory::InMemoryStore;`;删除已提取到 in_memory.rs 中的代码 |
| 涉及文件 | `src/memory/store.rs`(修改) |
| 前置依赖 | Task A1(文件已存在) |
| 预估工作量 | S< 1h |
| 风险等级 | 低 — 保留部分不变,提取部分在子模块中 |
| 验收条件 | `cargo test --all-targets` 全绿,测试数量不变(191 pass),clippy 0 warning |
#### Task A3 — 验证阶段一
| 项目 | 内容 |
|------|------|
| 任务描述 | 运行全量测试链确认目录重构零行为变化 |
| 涉及文件 | 全量 |
| 前置依赖 | Task A2 |
| 预估工作量 | XS(验证) |
| 风险等级 | 低 |
| 验收条件 | `cargo test --all-targets` 191 pass、`cargo clippy --all-targets -- -D warnings` 0 警告、`cargo build` 通过 |
### 阶段二:SqliteStore 实现
#### Task B1 — 添加 rusqlite 及 dev-dependencies
| 项目 | 内容 |
|------|------|
| 任务描述 | 在 `Cargo.toml` 中添加依赖:`[dependencies]``rusqlite = { version = "0.32", features = ["bundled"] }``[dev-dependencies]``tempfile = "3"`(用于测试隔离);运行 `cargo build` 确认编译通过,`cargo test --no-run` 验证 dev-dependencies 可用 |
| 涉及文件 | `Cargo.toml`(修改)、`Cargo.lock`(自动更新) |
| 前置依赖 | 无(可与阶段一并行) |
| 预估工作量 | XS< 15min |
| 风险等级 | 低 — 标准依赖添加 |
| 验收条件 | `cargo build` 成功,`cargo test --no-run` 成功,Cargo.lock 中生成 rusqlite 和 tempfile 条目 |
#### Task B2 — 实现 SqliteStore 核心
| 项目 | 内容 |
|------|------|
| 任务描述 | 创建 `src/memory/store/sqlite_store.rs`,实现以下 8 个子模块: |
| | 1`SqliteStore` 结构体 `{ conn: Arc<Mutex<Connection>> }` |
| | 2`SqliteStore::open(path)` — 支持 `":memory:"`,内部调用 `fs::create_dir_all` 确保父目录存在 |
| | 3`lock_conn()` — 内部辅助方法,`.lock().unwrap_or_else(\|e\| e.into_inner())` 处理 Mutex 中毒 |
| | 4`migrate()` — 按版本递增执行迁移:`PRAGMA user_version` 检查(初始版本号=1)→ 建表 `memory_items` + 索引 `idx_memory_items_created_at` + 设置 WAL 模式 + `busy_timeout=5s` + `PRAGMA synchronous = NORMAL` + `PRAGMA wal_autocheckpoint=1000` + `PRAGMA quick_check`(检测数据库损坏)+ `PRAGMA wal_checkpoint(TRUNCATE)` |
| | 5`MemoryStore` trait 的 4 个方法实现(save/get/delete/list),全部使用 spawn_blocking 包裹;**生命周期注意**:`spawn_blocking` 闭包前先 `.conn.clone()``Arc<Connection>`,参数调 `.clone()` 取 owned 值,再传入 `spawn_blocking(move \|{ ... })` |
| | — `save`: `INSERT INTO memory_items (id, content, metadata, created_at) VALUES (?1, ?2, ?3, ?4) ON CONFLICT(id) DO UPDATE SET content=excluded.content, metadata=excluded.metadata, created_at=excluded.created_at`(全字段覆盖 upsert,与 InMemoryStore 行为一致) |
| | — `get`: `SELECT content, metadata, created_at FROM memory_items WHERE id = ?1` → Ok(None) 当无结果 |
| | — `delete`: `DELETE FROM memory_items WHERE id = ?1`(幂等,不返回 NotFound |
| | — `list`: 根据 filter 字段(prefix/since/offset/limit)组合动态构造 WHERE 子句 + `ORDER BY created_at ASC` + `LIMIT ? OFFSET ?`,全部使用参数化查询防注入 |
| | 6)序列化转换层:`time::OffsetDateTime` 存为 TEXT(ISO 8601),通过 `.to_string()` 绑定 `String` 参数;`serde_json::Value` 存为 TEXT,通过 `serde_json::to_string()` 绑定 `String` 参数;读取时通过 `OffsetDateTime::parse``serde_json::from_str` 反序列化。在 SqliteStore 内部实现 `to_sql_params()` / `from_sql_row()` 辅助方法集中处理 |
| | 7)错误映射:不在 blanket impl From 中处理所有 rusqlite Error,而是在每个方法内部按数据完整性防御表的映射规则逐类处理: |
| | — 数据库文件损坏 / IO 错误 → `MemoryError::Storage` |
| | — created_at 解析失败 → `MemoryError::Serialization`(不 panic |
| | — content/metadata 为 NULL → `MemoryError::Storage` |
| | — serde_json 序列化/反序列化失败 → `MemoryError::Serialization` |
| | — 约束冲突 → `MemoryError::InvalidInput` |
| | (8)可观测性:在 4 个 trait 方法和 `open()` 上添加 `#[tracing::instrument(skip(self))]`;正常完成记录 `trace!`,超过性能预算阈值记录 `warn!`,操作失败记录 `error!` |
| 涉及文件 | `src/memory/store/sqlite_store.rs`(新增) |
| 前置依赖 | Task B1rusqlite + tempfile 依赖)、Task A1store/ 目录存在) |
| 预估工作量 | M1-4h |
| 风险等级 | 高 — 3 个技术点需留意:`time::OffsetDateTime``rusqlite::ToSql` 实现,需显式处理 String 绑定;`spawn_blocking` + `&self` 生命周期需 clone 后才能传闭包;错误映射需精细匹配 `rusqlite::Error` 嵌套变体(`SqliteFailure` 内含 `ErrorCode` |
| 验收条件 | 单元测试通过(见 Task B4)、`cargo build` 通过 |
#### Task B3 — 注册模块并重导出
| 项目 | 内容 |
|------|------|
| 任务描述 | 在 `store.rs` 添加 `pub mod sqlite_store;` + `pub use sqlite_store::SqliteStore;`;在 `memory.rs` 添加 `pub use store::SqliteStore;` |
| 涉及文件 | `src/memory/store.rs`(修改)、`src/memory.rs`(修改) |
| 前置依赖 | Task B2sqlite_store.rs 文件存在) |
| 预估工作量 | XS< 15min |
| 风险等级 | 低 |
| 验收条件 | `cargo build` 通过,`SqliteStore` 可从 `agcore::memory::SqliteStore` 路径访问 |
#### Task B4 — 编写测试
| 项目 | 内容 |
|------|------|
| 任务描述 | 在 `sqlite_store.rs` 中编写 `#[cfg(test)] mod tests`,覆盖: |
| | 1CRUD 基本操作(save → get → list → delete → get None |
| | 2Upsert 语义(同 id 重复 save 覆盖内容,created_at 保持调用方传入值) |
| | 3prefix 过滤(MemoryFilter.prefix |
| | 4)时间范围过滤(MemoryFilter.since |
| | 5offset/limit 分页 |
| | 6)并发 10 个 writer × 10 次写入,验证无数据丢失 |
| | 7)持久化恢复(write → drop store → reopen 同一文件 → read |
| | 8)错误路径:`open("/nonexistent_dir/ag.db")` 返回 Storage 错误 |
| | 辅助函数:`make_item(id)` 创建 MemoryItem,使用 `tempfile::TempDir`(或自定义 tmp 路径)隔离测试数据库文件 |
| 涉及文件 | `src/memory/store/sqlite_store.rs`(修改追加 test mod |
| 前置依赖 | Task B2(实现完成) |
| 预估工作量 | M1-4h |
| 风险等级 | 中 — 并发测试的时序控制、持久化恢复测试的 TempDir 管理 |
| 验收条件 | 全量测试通过,新增测试 ≥ 8 个 |
#### Task B5 — 验证阶段二
| 项目 | 内容 |
|------|------|
| 任务描述 | 运行全量测试链确认 SqliteStore 实现正确,不影响已有模块 |
| 涉及文件 | 全量 |
| 前置依赖 | Task B3、Task B4 |
| 预估工作量 | S< 1h |
| 风险等级 | 低 |
| 验收条件 | `cargo test --all-targets` 全绿(191 + 新增测试)、`cargo clippy --all-targets -- -D warnings` 0 警告、`cargo build` 通过 |
### 阶段三:验证收尾
#### Task C1 — 跨模块兼容性验证
| 项目 | 内容 |
|------|------|
| 任务描述 | 1)确认 `ConversationMemory` / `KnowledgeStore` / `MemoryRetriever` / `AgentSession` / `SessionMemory``Arc<dyn MemoryStore>` 接受 SqliteStore 时编译通过且测试全绿 |
| | 2)确认 SqliteStore 与 InMemoryStore 可互换——修改一个现有测试将后端从 InMemoryStore 换为 SqliteStore(使用 `":memory:"`),测试全绿 |
| | (3)验证性能预算:在测试环境下测量 save(1KB)/get(1KB)/list(100行) 的单次延迟,确认 < 5ms / < 3ms / < 20ms |
| 涉及文件 | 测试文件(memory/ 模块内各 test mod |
| 前置依赖 | Task B5 |
| 预估工作量 | S< 1h |
| 风险等级 | 低 |
| 验收条件 | 全量测试通过 + 互换测试通过 + 性能预算大致满足(未达标不阻塞发布) |
### 依赖关系图
```
阶段一(目录重构) 阶段二(SqliteStore 实现)
┌──────────────┐ ┌──────────────┐
│ Task A1 │ │ Task B1 │ ← 无依赖,可与 A 并行
│ 创建目录+提取 │ │ Cargo.toml │
└──────┬───────┘ └──────┬───────┘
↓ ↓
┌──────────────┐ ┌──────────────┐
│ Task A2 │ │ Task B2 │
│ 修改store.rs │← A1 ───→│ 核心实现 │← B1 + A1
└──────┬───────┘ └──────┬───────┘
↓ ↓
┌──────────────┐ ┌──────────────┐ ┌──────────────┐
│ Task A3 │ │ Task B3 │← B2 ──→│ Task B4 │
│ 验证阶段一 │ │ 注册+重导出 │ │ 编写测试 │
└──────────────┘ └──────┬───────┘ └──────┬───────┘
↓ ↓
┌──────────────┐────────────────┘
│ Task B5 │← B3 + B4
│ 验证阶段二 │
└──────┬───────┘
阶段三(验证收尾)
┌──────────────┐
│ Task C1 │
│ 跨模块兼容性 │
└──────────────┘
```
### Commit 安排
| 顺序 | Commit 类型 | Scope | 描述 | 包含 Task |
|------|------------|-------|------|-----------|
| 1 | refactor | memory | 将 store.rs 拆分为模块目录,仅结构搬移 | A1 → A2 → A3 |
| 2 | feat | memory | 实现 SqliteStore 持久化(含错误映射、测试、WAL 模式) | B1 → B2 → B3 → B4 → B5 → C1 |
注意:Task B1 与阶段一无依赖,可以在 Commit 1 合并进行或在 Commit 2 开头。建议在 Commit 2 开头,因为 Cargo.toml 变更属于功能变更而非重构。
### 验证全链
实施完毕后整体认证链路:
1. `cargo test --all-targets` — 全量测试通过
2. `cargo clippy --all-targets -- -D warnings` — 0 警告
3. `cargo build --release` — release 构建通过
4. 确认 `cargo doc --no-deps` 无 warning(新增公开类型文档注释)
5. 确认测试数量:191 + (8 个 new sqlite_store tests) = 199+ pass
+543
View File
@@ -0,0 +1,543 @@
# Phase 8 — MVP 集成出口实现方案
- **文档编号**15
- **标题**Phase 8 — MVP 集成出口实现方案
- **日期**2026-07-05
- **状态**:已定稿
- **涉及模块**:全局(llm/types、agent、tools、memory、prompt、examples
- **关联文档**roadmap.md(§Phase 8)、14-phase7-sqlite-store.md
---
## 1. 背景与目标
Phase 5-7 已交付 P0 功能闭环:ProviderConfig `from_env()`Phase 5)、ToolDef IR 正式化(Phase 6)、SqliteStore 持久化(Phase 7)。当前 200 个测试全绿、clippy 0 警告,但缺乏一个"可被人依赖"的集成出口。
Phase 8 的目标是完成 API 稳定性扫尾 + Quick Start 示例 + 端到端示例,产出 **v0.2.0-rc.1** 标签。三个 Step 分别对应三类用户群体:
| Step | 受众 | 交付物 |
|------|------|--------|
| **8.1** | 存量升级者(v0.1 → v0.2 | API 稳定性扫尾 + CHANGELOG |
| **8.2** | 新用户评估者("30 秒决定要不要用" | Quick Start 示例 |
| **8.3** | 技术决策者("这框架能跑真实场景吗") | 端到端集成示例 |
---
## 2. 当前状态
| 度量 | 数值 |
|------|------|
| `cargo test --all-targets` | ✅ 200 passed / 0 failed |
| `cargo clippy --all-targets -- -D warnings` | ✅ 0 警告 |
| 已存在 `#[non_exhaustive]` 枚举 | 4 个(StopReason / FinishReason / EvictionPolicy / ProviderType |
| 已存在 `#[deprecated]` 项 | 3 个(ChatResponse / ToolDefinition / task_agent_demo 中旧类型使用) |
| 已有示例 | 8 个 |
| `StepStatus::Completed` 使用类型 | `ChatResponse`(已 `#[deprecated]` |
### 2.1 关键技术债
```
// agent/task.rs —— StepStatus 当前使用已废弃类型
#[allow(deprecated)]
pub enum StepStatus {
Completed(ChatResponse), // ← ChatResponse 已在 0.1.0 标记 #[deprecated]
...
}
```
`task_agent_demo.rs` 中同时使用了 `ChatResponse` / `OpenaiChatMessage` / `FinishReason` 三个废弃类型,入口处有 `#![allow(deprecated)]`
---
## 3. 实施方案
### 3.1 Step 8.1 — API 稳定性扫尾
拆为 4 个增量 commit
| Commit | 内容 | 涉及文件 |
|--------|------|---------|
| **commit 1** | 14 个公开枚举追加 `#[non_exhaustive]` | 各枚举定义文件(详见 §3.1.1) |
| **commit 2** | `StepStatus::Completed(ChatResponse)``Completed(MessageResponse)` + `task_agent_demo.rs` 清理全部 3 个废弃类型(`ChatResponse` / `OpenaiChatMessage` / `FinishReason`),移除 `#![allow(deprecated)]` | `src/agent/task.rs``examples/task_agent_demo.rs` |
| **commit 3** | CHANGELOG v0.2 条目 + Cargo.toml version → `0.2.0-rc.1` + README 更新 | `CHANGELOG.md``Cargo.toml``README.md` |
| **commit 4** | 验证:`cargo test + clippy + cargo doc` 零告警 | 无代码改动 |
#### 3.1.1 `#[non_exhaustive]` 追加清单(14 个枚举)
按优先级分级:
| 优先级 | 枚举 | 模块路径 | 理由 |
|--------|------|---------|------|
| **P0 核心** | `Message` | `llm/types/message.rs` | 核心 IR 类型,未来可能新增变体(MultiModal 扩展) |
| | `ContentBlock` | `llm/types/message.rs` | 同上 |
| | `ContentBlockType` | `llm/types/message.rs` | 同上 |
| | `StreamEvent` | `llm/types/response_v2.rs` | 流式事件集,Provider 扩展可能新增事件 |
| | `HookEvent` | `llm/hooks.rs` | 生命周期钩子,框架扩展需要新增事件点 |
| **P0 Error** | `AgentError` | `agent/error.rs` | 顶层错误,下游 match 需保护 |
| | `LlmError` | `llm/error.rs` | LLM 调用错误 |
| | `ToolError` | `tools/error.rs` | 工具系统错误 |
| | `MemoryError` | `memory/error.rs` | 记忆系统错误 |
| | `PromptError` | `prompt/error.rs` | 提示词工程错误 |
| **P1 其他** | `MemoryStrategy` | `memory/conversation.rs` | 对话策略,未来可扩展(如 Summarize) |
| | `StepStatus` | `agent/task.rs` | 步骤状态机,可扩展(如 Cancelled) |
| | `ToolChoice` | `llm/types/request.rs` | Provider 工具选择策略 |
| | `ResponseFormat` | `llm/types/shared.rs` | 响应格式枚举 |
**明确不加的**
| 类别 | 枚举 | 原因 |
|------|------|------|
| 内部 wire-format | `OpenaiChatMessage` / `OpenaiTool` / `OpenaiToolCall` / `ContentField` / `OpenaiContentPart` / `LegacyStreamEvent` | 内部转换层,不构成公共 API 契约 |
| 语义稳定 | `Role` / `ServiceTier` / `Modality` / `ImageDetail` / `AudioFormat` / `StopSequence` | 语义已收敛,协议层无新增变体预期 |
| 使用面窄 | `TemplateValue` / `Permission` / `McpTransport` / `ContentBlockBuilder` / `ExtraError` | 内部实现细节或使用频率极低,下游不直接 match |
> **`#[non_exhaustive]` 的不可逆性**:一旦 v0.2.0-rc.1 发布,以下游代码可能依赖 `_ =>` 通配分支。在 v0.3+ 中移除 `#[non_exhaustive]` 将构成 semver breaking change(新增变体不再触发编译警告,下游 match 可能遗漏新变体),因此当前追加的标记应视为永久 API 契约。
#### 3.1.2 StepStatus 迁移细节
```rust
// 变更前
#[allow(deprecated)]
pub enum StepStatus {
Completed(ChatResponse), // ChatResponse 已 #[deprecated]
...
}
// 变更后
#[non_exhaustive]
pub enum StepStatus {
Completed(MessageResponse),
...
}
```
**字段映射差异**`MessageResponse` 不是 `ChatResponse` 的简单改名——两者结构不同,迁移需要做字段适配:
| ChatResponse 字段 | 类型 | MessageResponse 字段 | 类型 | 映射方式 |
|-------------------|------|---------------------|------|---------|
| `message` | `OpenaiChatMessage` | `message` | `Message` | 类型替换:`OpenaiChatMessage::assistant_text(t)``Message::Assistant { content: vec![ContentBlock::Text { text: t.into() }] }` |
| `usage` | `Usage` | `usage` | `Usage` | ✅ 同类型,直接迁移 |
| `stop_reason` | `Option<FinishReason>` | `stop_reason` | `StopReason` | 类型替换:`Some(FinishReason::Stop)``StopReason::Stop`;无 Option 包裹 |
| — | — | `id` | `String` | 新增必填字段,使用空字符串 `""` 占位 |
| — | — | `model` | `String` | 新增必填字段,使用 `"mock"` 或空字符串占位 |
| — | — | `extra` | `HashMap<String, Value>` | 新增字段,使用 `HashMap::new()` 占位 |
**迁移示例**`task_agent_demo.rs` 的构造代码):
```rust
// 旧代码(3 个废弃类型)
StepStatus::Completed(ChatResponse {
message: OpenaiChatMessage::assistant_text("天气:晴,22°C"),
usage: Usage::from_input_output(10, 5),
stop_reason: Some(FinishReason::Stop),
})
// 新代码(纯 MessageResponse
StepStatus::Completed(MessageResponse {
id: String::new(),
model: "mock".into(),
message: Message::assistant("天气:晴,22°C"),
usage: Usage::from_input_output(10, 5),
stop_reason: StopReason::Stop,
extra: HashMap::new(),
})
```
涉及文件:
- `src/agent/task.rs`:枚举定义 + `#[allow(deprecated)]` 移除 + `#[non_exhaustive]` 追加
- `examples/task_agent_demo.rs``ChatResponse{...}``MessageResponse{...}` 构造替换,同时替换 `OpenaiChatMessage` / `FinishReason` 引用,移除 `#![allow(deprecated)]`
### 3.2 Step 8.2 — Quick Start 示例
| 属性 | 值 |
|------|-----|
| 文件 | `examples/quick_start.rs` |
| 规模 | ~36 行 |
| Provider | `MockProvider`FIFO 单响应队列) |
| 工具 | `EchoTool`(回传 `"收到: {input}"`,完整 JSON Schema 参数声明) |
| 执行 | `submit_turn("你好")` → 验证输出包含 `"收到"` |
| 验证 | `cargo run --example quick_start` exit 0 |
设计要点:
- 展示四层抽象:Agent trait / BaseTool 自定义 / AgentBuilder 装配 / AgentSession 执行
- 无外部依赖、无 API key、零配置
### 3.3 Step 8.3 — 端到端示例
| 属性 | 值 |
|------|-----|
| 文件 | `examples/end_to_end.rs` |
| 规模 | ~160 行(**最小可行边界**:3 工具 + 3 轮 + 持久化验证,防止实施中进一步膨胀) |
| Provider | 自动检测 `AG_LLM_*``from_env()`fallback 到 `MockProvider` |
| 工具组合 | EchoTool(回显)+ CalcTool(四则运算,本地执行)+ NoteTool(笔记,通过 MemoryStore trait 操作 SessionMemory |
| 持久化 | `tempfile::TempDir` + `SqliteStore`,drop 后重建连接验证数据不丢 |
| 对话 | 3 轮:计算 → 记笔记 → 回忆 |
| 验证 | `cargo run --example end_to_end` exit 0(无需任何外部配置) |
**真实 Provider 切换**:示例在文件顶部注释中说明 "设置 `AG_LLM_BASE_URL` / `AG_LLM_API_KEY` / `AG_LLM_MODEL` 环境变量即可使用真实 LLM Provider(支持 OpenAI / Ollama 等);未设置时自动降级为 MockProvider,零配置可运行。"
**`from_env()` 部分环境变量策略**`from_env()` 要求完整的三件套(`{prefix}_BASE_URL` + `{prefix}_API_KEY` + `{prefix}_MODEL`)。当环境变量部分设置时,示例**整体降级到 MockProvider**——不在"半配置"状态下尝试部分初始化。日志输出形如 `"AG_LLM_* 环境变量不完整(检测到: {found_vars}),回退到 MockProvider"`
架构亮点:
```
┌─────────────────────────┐
│ AgentSession │
│ (submit_turn × 3) │
└────┬──────┬──────┬──────┘
│ │ │
┌────┘ │ └──────┐
▼ ▼ ▼
┌──────────┐ ┌────────┐ ┌──────────┐
│ EchoTool │ │CalcTool│ │ NoteTool │
│ (回显) │ │(四则) │ │ (记忆) │
└──────────┘ └────────┘ └────┬─────┘
┌──────▼──────┐
│ SessionMemory│
│ (MemoryStore)│
└──────┬──────┘
┌──────▼──────┐
│ SqliteStore │
│ (temp dir) │
└─────────────┘
```
NoteTool 展示 `MemoryStore` trait 解耦能力:不绑定 SqliteStore,上层 `AgentSession` 通过 `SessionMemory` 操作,底层可互换。
---
## 4. 否决项记录
| 否决方案 | 否决原因 |
|---------|---------|
| `#[non_exhaustive]` 仅加 5 个核心类型 | 全面覆盖 Error enums 为零运行时成本,对下游更友好。Error 枚举是下游 match 最密集的地方,漏标会在 v0.3 引入 breakage |
| StepStatus::Completed 留到 v0.3 再修 | rc.1 前清理 deprecated 类型污染最划算——越晚 migration cost 越高,且当前仅 1 个示例 + 1 个测试引用 |
| Quick Start 纯文本路线(不展示自定义工具) | 含 EchoTool 展示核心差异化,仅多 5 行代码但传递了"可以自定义工具"的关键信息 |
| 端到端仅 Echo + Calc(无 NoteTool | NoteTool 展示 MemoryStore trait 解耦能力是架构亮点,跳过后新用户无法理解 memory 如何集成到 Agent 流程 |
| 持久化仅注释说明不实际运行(方案 Y) | 进程内实操验证(create → drop → reopen → assert)比注释更有说服力,增加约 15 行代码 |
---
## 5. 关键假设
1. **MockProvider FIFO 队列满足 auto-tool-loop 消费顺序**MockProvider 的 `pop()` 按预设顺序弹出。当 LLM 返回多个 tool call 时队列消费顺序与预设一致,无需额外同步
2. **StepStatus 切换需做字段适配**`ChatResponse`(3 字段) 到 `MessageResponse`(6 字段) 存在字段类型差异(`message` 类型不同、`stop_reason` 类型 + Option 有无不同、`id`/`model`/`extra` 为新增必填字段),消费者需按字段映射表提供占位值。但消费者仅 1 个(`task_agent_demo.rs`)+ 1 个内联测试,手动适配工作量极小。`StepStatus``is_terminal()` / `is_pending()` 行为不受影响
3. **所有 10 个示例零外部配置 exit 0**:已有 8 个示例已验证,新增 2 个(quick_start + end_to_end)均使用 MockProvider fallback,无需 API key
4. **`#[non_exhaustive]` × 14 不触发额外 clippy warning**:当前无代码对以上枚举做 exhaustive match(不含 `_`),追加 `#[non_exhaustive]` 是纯安全标记
---
## 6. 实施顺序与验证标准
### 6.1 提交顺序
```
Step 8.1 (4 commits)
→ commit 1: #[non_exhaustive] × 14
→ commit 2: StepStatus 修复(Completed(ChatResponse) → Completed(MessageResponse)
→ commit 3: CHANGELOG v0.2 + Cargo.toml version 0.2.0-rc.1 + README 更新
→ commit 4: 验证(test / clippy / doc 零告警)
Step 8.2
→ commit 5: examples/quick_start.rs~36 行)
Step 8.3
→ commit 6: examples/end_to_end.rs~160 行)
最终验证
→ cargo test --all-targets
→ cargo clippy --all-targets -- -D warnings
→ cargo doc --no-deps
→ git tag v0.2.0-rc.1
```
### 6.2 验收标准
| 指标 | 要求 |
|------|------|
| `cargo test --all-targets` | 全绿 |
| `cargo clippy --all-targets -- -D warnings` | 0 警告 |
| `cargo doc --no-deps` | 0 warning |
| 所有 10 个示例 | `cargo run --example <name>` exit 0 |
| Cargo.toml version | `0.2.0-rc.1` |
| CHANGELOG | v0.2 条目完整(Added / Changed / Deprecated / Fixed / Removed 各节) |
| README | 示例列表 + 版本号更新 |
| git tag | `v0.2.0-rc.1` |
---
## 7. 参考来源
- **roadmap.md** — Phase 8 原始定义(Step 8.1/8.2/8.3)、依赖关系(Phase 5/6/7 → Phase 8
- **`src/agent/task.rs`** — `StepStatus` 当前实现,`Completed(ChatResponse)` 类型
- **`src/llm/types/message.rs`** — `Message` / `ContentBlock` / `ContentBlockType` 枚举定义
- **`src/llm/types/response_v2.rs`** — `StreamEvent` / `StopReason` 枚举定义(StopReason 已有 `#[non_exhaustive]`
- **`src/llm/types/shared.rs`** — `ResponseFormat` / `Role` / `FinishReason` 等枚举(FinishReason 已有 `#[non_exhaustive]`
- **`src/llm/types/request.rs`** — `ToolChoice` 枚举定义
- **`src/llm/hooks.rs`** — `HookEvent` 枚举定义
- **`src/llm/error.rs`** — `LlmError` 枚举定义
- **`src/agent/error.rs`** — `AgentError` 枚举定义
- **`src/tools/error.rs`** — `ToolError` 枚举定义
- **`src/memory/error.rs`** — `MemoryError` 枚举定义
- **`src/memory/conversation.rs`** — `MemoryStrategy` 枚举定义
- **`src/prompt/error.rs`** — `PromptError` 枚举定义
- **`examples/task_agent_demo.rs`** — 当前使用 `#[allow(deprecated)]` + `ChatResponse` 的示例
---
## 8. 实施计划
### 8.1 实施步骤
#### Step 8.1 — API 稳定性扫尾
拆为 4 个增量 commit,依次提交。
##### commit 1: #[non_exhaustive] × 14
| 属性 | 值 |
|------|-----|
| 涉及文件 | 14 个枚举定义所在文件(见下方清单) |
| 前置依赖 | 无 |
| 预估工作量 | S<1h |
| 风险等级 | 低 |
在每个目标枚举定义处的 `pub enum` 之前加一行 `#[non_exhaustive]`,纯文本属性追加,无逻辑变更。
| 目标枚举 | 文件路径 | 行号附近 |
|---------|---------|---------|
| `Message` | `src/llm/types/message.rs` | `pub enum Message` (L22) |
| `ContentBlock` | `src/llm/types/message.rs` | `pub enum ContentBlock` (L99) |
| `ContentBlockType` | `src/llm/types/message.rs` | `pub enum ContentBlockType` (L134) |
| `StreamEvent` | `src/llm/types/response_v2.rs` | `pub enum StreamEvent` (L167) |
| `HookEvent` | `src/llm/hooks.rs` | `pub enum HookEvent` (L9) |
| `AgentError` | `src/agent/error.rs` | `pub enum AgentError` (L20) |
| `LlmError` | `src/llm/error.rs` | `pub enum LlmError` (L10) |
| `ToolError` | `src/tools/error.rs` | `pub enum ToolError` (L6) |
| `MemoryError` | `src/memory/error.rs` | `pub enum MemoryError` (L8) |
| `PromptError` | `src/prompt/error.rs` | `pub enum PromptError` (L3) |
| `MemoryStrategy` | `src/memory/conversation.rs` | `pub enum MemoryStrategy` (L14) |
| `StepStatus` | `src/agent/task.rs` | `pub enum StepStatus` (L59) |
| `ToolChoice` | `src/llm/types/request.rs` | `pub enum ToolChoice` (L14) |
| `ResponseFormat` | `src/llm/types/shared.rs` | `pub enum ResponseFormat` (L70) |
> **注意**`StepStatus` 在 commit 2 中会同时被修改(variant 类型替换 + 移除 `#[allow(deprecated)]`)。commit 1 仅追加 `#[non_exhaustive]` 属性,commit 2 再处理变体变更和清理。
**验收条件**`cargo build --all-targets` 通过
##### commit 2: StepStatus 修复 + 废弃类型清理
| 属性 | 值 |
|------|-----|
| 涉及文件 | `src/agent/task.rs``examples/task_agent_demo.rs` |
| 前置依赖 | commit 1StepStatus 先标记 `#[non_exhaustive]`,此处改 variant 时一并保留,无实际冲突) |
| 预估工作量 | S<1h,约 20 行改动) |
| 风险等级 | 低 |
两步操作:
1. **`src/agent/task.rs`**L59-L71):
- `StepStatus::Completed(ChatResponse)``Completed(MessageResponse)`
- 移除 `#[allow(deprecated)]`(第 13、59 行两处)
2. **`examples/task_agent_demo.rs`**
- 替换 3 个废弃类型:`ChatResponse``MessageResponse``OpenaiChatMessage::assistant_text(t)``Message::assistant(t)``FinishReason::Stop``StopReason::Stop`
- 补充 `id: String::new()``model: "mock".into()``extra: HashMap::new()` 占位字段
- 移除 `#![allow(deprecated)]`(第 26 行)
- 移除 `use` 中的 `ChatResponse``OpenaiChatMessage``FinishReason`
- 添加 `use std::collections::HashMap``use agcore::llm::types::{Message, MessageResponse, StopReason}`(注意:`Message::assistant_text(t)` 不存在,需使用 `Message::assistant(t)`)
字段映射参见 §3.1.2 的字段映射表和迁移示例。
**验收条件**`cargo build --all-targets` 通过,零 deprecated warning
##### commit 3: CHANGELOG + 版本号 + README
| 属性 | 值 |
|------|-----|
| 涉及文件 | `CHANGELOG.md``Cargo.toml``README.md` |
| 前置依赖 | commit 1+2CHANGELOG 需记录实际变更) |
| 预估工作量 | S<1h |
| 风险等级 | 低 |
1. **`CHANGELOG.md`**:新增 `[0.2.0-rc.1]` 条目,包含:
- **Added**SqliteStore 持久化 / OllamaProvider / ProviderConfig::from_env / ToolDef IR / Quick Start 和 end_to_end 示例
- **Changed**MessageRequest.tools 切换 ToolDef / StepStatus::Completed 类型替换
- **Deprecated**ChatResponse / with_system_prompt() / with_client()
- **Non-exhaustive**14 个枚举标记清单
2. **`Cargo.toml`**:第 3 行 `version = "0.1.0"``version = "0.2.0-rc.1"`
3. **`README.md`**:更新示例列表从 7 个改为 10 个(含新增 2 个),版本号同步
**验收条件**:人工 review CHANGELOG + `git diff` 确认版本号
##### commit 4: 验证
| 属性 | 值 |
|------|-----|
| 涉及文件 | 无代码改动 |
| 前置依赖 | commit 3 |
| 预估工作量 | S(<1h,主要等待编译) |
| 风险等级 | 低 |
运行三条命令:
```bash
cargo test --all-targets
cargo clippy --all-targets -- -D warnings
cargo doc --no-deps 2>&1 | grep "^warning:" && echo "WARNINGS FOUND" || echo "0 warnings"
```
**验收条件**:前两条 0 错误,第三条输出 `0 warnings`
#### Step 8.2 — Quick Start 示例
##### commit 5: examples/quick_start.rs
| 属性 | 值 |
|------|-----|
| 涉及文件 | `examples/quick_start.rs` |
| 前置依赖 | 无(可从 Phase 7 独立创建) |
| 预估工作量 | S<1h |
| 风险等级 | 低 |
新文件 `examples/quick_start.rs`~36 行,结构如下:
```
1- 6 use 块(agcore 类型 + Arrow/std 类型)
7- 8 struct Greeter + impl Agentname / system_prompt
9-14 struct EchoTool + #[async_trait] impl BaseTool(完整 JSON Schema 带 text 参数)
15-20 fn mock_response() -> MessageResponse 辅助函数(构造纯文本响应)
21-33 #[tokio::main] async fn main():
- ToolRegistry::new() + register EchoTool
- MockProvider 预设 1 条 mock_response
- AgentBuilder::new() + provider + tool_registry + hook_executor → build
- AgentSession::new + submit_turn("你好")
- println!("{}", response.text())
```
**设计约束**
- EchoTool 的 `parameters()` 返回完整 JSON Schema`{"type":"object","properties":{"text":{"type":"string"}},"required":["text"]}`
- 无外部依赖、无 API key、零配置
- 展示四层抽象:Agent trait / BaseTool 自定义 / AgentBuilder 装配 / AgentSession 执行
**验收条件**`cargo run --example quick_start` exit 0,输出包含 `"收到"`
#### Step 8.3 — 端到端示例
##### commit 6: examples/end_to_end.rs
| 属性 | 值 |
|------|-----|
| 涉及文件 | `examples/end_to_end.rs` |
| 前置依赖 | commit 5(示例编写模式已建立);SqliteStorePhase 7 已完成) |
| 预估工作量 | M1-4h |
| 风险等级 | 中 |
新文件 `examples/end_to_end.rs`,~160 行,最小可行边界(3 工具 + 3 轮 + 持久化验证)。
**Provider 初始化策略**
```
if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
// 使用真实 ProviderAG_LLM_MODEL 非必填,from_env 内部会处理默认值)
let provider: Arc<dyn LlmProvider> = Arc::from(create_provider(
ProviderType::OpenaiChat, ProviderConfig::from_env("AG_LLM").unwrap()
)?);
} else {
// MockProvider fallback,预设 4 条响应序列
let found = ["AG_LLM_BASE_URL", "AG_LLM_API_KEY"].iter()
.filter(|k| env::var(k).is_ok()).collect::<Vec<_>>();
eprintln!("AG_LLM_* 环境变量不完整(检测到: {:?}),回退到 MockProvider", found);
}
```
**工具定义**
| 工具 | 功能 | 关键技术点 |
|------|------|-----------|
| `EchoTool` | 回显输入 | 基础工具注册模式 |
| `CalcTool` | 本地执行四则运算 | 手动解析算术表达式(ponytail:基础 +-*/ 运算无需引入 `rhai` 依赖) |
| `NoteTool` | 通过 MemoryStore trait 读写笔记 | 直接持有 `Arc<dyn MemoryStore>`key 前缀 `"note:"`save 用 `MemoryStore::save(MemoryItem { id: "note:{key}", content, .. })`query 用 `MemoryStore::list(MemoryFilter { prefix: Some("note:"), .. })` |
**持久化验证**
```rust
let dir = tempfile::TempDir::new()?;
let db_path = dir.path().join("agcore.db");
let backend = Arc::new(SqliteStore::open(&db_path)?);
// ... 构建 RuntimeBundle + AgentSession,写入数据 ...
drop(bundle); // 释放所有对 backend 的 Arc 引用
drop(session);
// 此时 backend 无活跃引用,SQLite 连接自动关闭
let backend2 = Arc::new(SqliteStore::open(&db_path)?); // 重建连接
// assert 数据仍在
```
**输出示范**
```
=== agcore 端到端演示 ===
🔄 Provider: MockProvider (离线回退模式)
💾 SqliteStore: /tmp/agcore_XXXXX/agcore.db
🔧 注册工具: echo, calc, note
第 1 轮 用户: 帮我算 25 * 4
→ 调用 calc(...) → 100
→ 回答: 25 * 4 = 100
第 2 轮 用户: 记下来:结果是 100
→ 调用 note(save, ...)
→ 回答: 已记录
第 3 轮 用户: 我刚才算了什么?
→ 调用 note(query)
→ 回答: 您刚才的计算结果是 100
📊 用量: prompt=XX, completion=XX
=== 持久化验证 ===
✓ 跨连接数据存活验证通过
✓ 端到端演示完成
```
**设计约束**
- 文件顶部注释说明 `AG_LLM_*` 环境变量切换真实 Provider
- 零外部配置可运行(Mock fallback
- 最小可行边界:3 工具 + 3 轮 + 持久化验证,不膨胀
**验收条件**`cargo run --example end_to_end` exit 0(零外部配置)
### 8.2 并行机会
commit 1 和 commit 5 可以并行执行(零文件重叠)。commit 5 也可与 commit 2 并行。commit 6 实质上也仅依赖「代码库状态稳定」而非某个具体 commit。
| 并行组 | commit A | commit B | 前提 |
|--------|---------|---------|------|
| 1 | commit 1#[non_exhaustive] | commit 5Quick Start | 零文件重叠 |
| 2 | commit 2StepStatus 修复) | commit 5Quick Start | 零文件重叠 |
| 3 | commit 5Quick Start | commit 6(端到端) | 零文件重叠,但存在知识依赖——commit 6 需参考 commit 5 的 `MessageResponse` 构造、`MockProvider` 用法、`AgentBuilder` 装配模式。推荐 commit 5 先行或实施前同步这些模式 |
### 8.3 风险与应对
| 风险 | 影响 | 可能性 | 应对 |
|------|------|--------|------|
| MockProvider 响应序列与 tool-loop 消费顺序不匹配 | commit 6 端到端示例不通过 | 中 | 按 §5 假设 1:设计响应队列时确保每条 Mock 响应的 `stop_reason` 与 ToolUse/Stop 匹配。出现不匹配时改用完整 `MessageResponse` 构造显式控制 |
| NoteTool 与 AgentSession 的数据传递路径需要扩展现有 API | commit 6 需要修改 `session.rs` | 低 | ponytail 方案:NoteTool 直接持有 `Arc<dyn MemoryStore>` 引用,在 execute 时直接操作 `MemoryStore::save/get`,绕过 AgentSession 的 session_memory 封装 |
| `#[non_exhaustive]` 在某个 enum 上导致 crate 内 match 编译失败 | commit 1 不通过 | 低 | 实施前先运行 `rg "match.*(Message|ContentBlock|ContentBlockType|StreamEvent|HookEvent|AgentError|LlmError|ToolError|MemoryError|PromptError|MemoryStrategy|StepStatus|ToolChoice|ResponseFormat)" src/ --include="*.rs"` 快速扫描 exhaustive match。若某 enum 编译失败,回退该 enum 上的 `#[non_exhaustive]` 属性,标注原因 |
### 8.4 测试策略
| commit | 测试 | 方式 |
|--------|------|------|
| commit 1 | 编译测试 | `cargo build --all-targets` |
| commit 2 | 编译 + 单测 + 无 deprecated warning | `cargo build --all-targets && cargo test` |
| commit 3 | 人工 review | `git diff` |
| commit 4 | 全量自动化 | `cargo test + clippy + doc` |
| commit 5 | 示例运行 | `cargo run --example quick_start` |
| commit 6 | 示例运行 | `cargo run --example end_to_end` |
| 最终 | 全量回归 | 全部三项 + 所有 10 个示例 |
+821
View File
@@ -0,0 +1,821 @@
# Phase 9 — 流式体验增强实施方案
- **文档编号**16
- **标题**:Phase 9 — 流式体验增强实施方案
- **日期**2026-07-05
- **状态**:已定稿
- **涉及模块**llm/cycle、llm/types/response_v2、agent/session
- **关联文档**roadmap.md(§Phase 9)、15-phase8-mvp-integration.md
---
## 1. 背景与目标
agcore 已发布 v0.2.0-rc.1Phase 0-8 全部完成。当前 Agent 会话只有非流式 API(`submit_turn`),开发者无法看到实时 token 输出和工具执行过程。Phase 9 的目标是为 `AgentSession` 新增流式方法 `submit_turn_stream`,让开发者能实时看到 LLM token 生成和工具执行状态。
### 1.1 现有能力
| 能力 | 方法 | 流式 | 自动工具循环 | 状态 |
|------|------|------|-------------|------|
| LLM 流式请求 | `LlmCycle::submit_stream` | ✅ | ❌ | 已就绪 |
| LLM 工具循环 | `LlmCycle::submit_with_tools` | ❌ | ✅ | 已就绪 |
| Agent 会话 | `AgentSession::submit_turn` | ❌ | ✅ | 已就绪 |
| 流事件枚举 | `StreamEvent`(11 变体) | — | — | 缺工具执行事件 |
| Mock 流 | `MockProvider::chat_stream` | ✅ | — | 可模拟流事件序列 |
### 1.2 核心矛盾
流式能力和工具循环能力分别存在于两个方法中,从未被组合。`submit_stream` 只管将 LLM 流事件原样转发,不理解工具调用;`submit_with_tools` 自动执行工具循环但全程阻塞。Phase 9 就是要组合它们:**在工具循环中,每一轮 LLM 调用都是流式的,并在工具执行前后插入语义事件**。
---
## 2. 需求分析
### 2.1 功能需求
1. **`AgentSession::submit_turn_stream(user_input)`** — 返回 `StreamEvent` 流,开发者通过 `while let Some(event) = stream.next().await` 逐事件消费
2. **流式工具循环** — 多轮工具调用过程中流不卡死,每轮工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
3. **`finalize_turn(response)`** — 流消费完成后同步 session 状态(cost 累计 + `OnTurnEnd` hook 触发)
4. **新增 `StreamEvent` 变体**`ToolExecutionStarted` + `ToolExecutionCompleted`,携带工具名称、调用 ID、参数/结果摘要
### 2.2 非功能需求
- **零影响**:现有 `submit_turn``submit_with_tools` 行为不变,存量测试 0 回归
- **异步流**:消费者通过 `futures_util::StreamExt::next()` 逐事件消费
- **错误事件化**:错误通过 `StreamEvent::Error` 事件表达,不通过 `Result` 通道终止流
- **最少代码**:复用现有 `submit_with_tools` 的工具循环逻辑模式和 `submit_stream` 的流管道模式
### 2.3 不做事项
| 事项 | 理由 |
|------|------|
| 新增示例(Phase 9.2 再加) | 缩窄 Phase 9 范围至核心能力 |
| `OnTurnEnd` 自动触发 | Rust 所有权约束:流是延迟求值,`&mut self` 无法进入闭包;由消费者收到 `MessageComplete` 后手动调用 `finalize_turn` |
| 修复 cost 统计 | 中间轮 cost 丢失是已知限制,与 `submit_turn` 行为一致 |
| 跨 turn 消息历史保留 | Phase 10 `ContextSlot` 的职责 |
| 并行 tool 调用的事件细化 | 当前工具调用是顺序 `for` 循环,并行化留待后续优化 |
| `run_tool_loop` 内消息压缩 | `run_tool_loop` 不接收 `compact_config` 参数,不执行上下文压缩。长工具循环中消息增长可能导致 context window 溢出,这是流式实现的已知限制。后续可通过传递 `compact_config``run_tool_loop` 支持 |
| LLM 请求自动 retry | 流式版本不在 `run_tool_loop` 内部实现 retry(详见 §3.6 说明)。调用方可自行包装 `RetryProvider` 或在 `LlmProvider` 实现层处理 |
---
## 3. 方案设计
### 3.1 架构总览
```
┌──────────────────────────────────────────────────────────────┐
│ AgentSession │
│ ┌──────────────────────────────────────────────────────┐ │
│ │ submit_turn_stream() │ │
│ │ ├─ OnTurnStart hook(同步触发,返回流之前) │ │
│ │ ├─ 组装 LlmCyclesystem_prompt / compact_config │ │
│ │ ├─ 调用 submit_with_tools_stream() │ │
│ │ ├─ turn_index += 1 │ │
│ │ └─ 返回流 │ │
│ │ │ │
│ │ finalize_turn(response) │ │
│ │ ├─ cost_so_far.add(&response.usage) │ │
│ │ └─ OnTurnEnd hookturn_index - 1 │ │
│ └──────────────────────────────────────────────────────┘ │
│ submit_with_tools_stream(prompt, Arc<ToolRegistry>)
┌──────────────────────────────────────────────────────────────┐
│ LlmCycle (tokio::spawn task — run_tool_loop 状态机) │
│ │
│ max_turns = max_tool_turns.unwrap_or(10) │
│ for round in 1..=max_turns { │
│ ① build_request(messages, tools) │
│ ② provider.chat_stream(request).await │
│ 匹配 Err → tx.send(Error{..}) + return(不 panic
│ ③ 消费 LLM 流,所有事件 → mpsc unbounded tx(全量转发) │
│ ④ partial.finalize() → MessageResponse │
│ ⑤ if has_tool_use: │
│ ├─ tx → ToolExecutionStarted { tool_name, id, args } │
│ ├─ registry.invoke_all(calls, timeout).await │
│ ├─ for result: tx → ToolExecutionCompleted { ... } │
│ ├─ push tool results → messages │
│ └─ continue(新一轮) │
│ else: break(最终轮,已发出 MessageComplete
│ } │
│ │
│ 产出事件序列(通过 mpsc::unbounded_channel): │
│ MessageStart → ... → ToolCallEnd → ToolExecutionStarted → │
│ ToolExecutionCompleted → MessageStart → TextDelta → ... → │
│ CostUpdate → MessageComplete │
└──────────────────────────────────────────────────────────────┘
```
> **`max_tool_turns` 语义**:与非流式 `submit_with_tools` 一致——`None` 退化为 `10``unwrap_or(10)`)。默认值 `Some(10)` 已提供安全上限;如需增大限制,手动设置为 `Some(N)`。⚠️ 生产环境建议始终设有限值防止无限循环。
### 3.2 事件序列约定
**纯文本流**(无 tool_use):
```
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd → CostUpdate → MessageComplete
```
**单轮工具调用**
```
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd
→ ContentBlockStart → ToolCallArgumentsDelta* → ToolCallEnd
→ CostUpdate → MessageComplete { stop_reason: ToolUse }
→ ToolExecutionStarted → [工具执行] → ToolExecutionCompleted
→ ContentBlockStart → TextDelta* → ContentBlockEnd
→ CostUpdate → MessageComplete { stop_reason: Stop }
```
**多轮工具调用**
```
... → ToolExecutionCompleted(第 1 轮)
→ ToolCallArgumentsDelta* → ToolCallEnd
→ ToolExecutionStarted → ToolExecutionCompleted(第 2 轮)
→ ... → CostUpdate → MessageComplete(最终轮)
```
**工具不可恢复错误**
```
... → ToolCallEnd → ToolExecutionStarted
→ Error { "tool 'search' 不可恢复错误: ..." } → MessageComplete
```
> **`MessageComplete.full_response` 内容范围**:每轮 LLM 调用独立产生一个 `MessageComplete`,其中 `full_response` 仅包含**该轮 LLM 的单个响应**(不累积前面工具轮次的结果)。中间轮(`stop_reason: ToolUse`)的 `full_response` 通常只包含 `ToolUse` block,无文本。最终轮(`stop_reason: Stop`)的 `full_response` 包含 LLM 的最终输出。消费者如需追踪完整对话历史,应自行累加所有轮次的 `Message`。
### 3.3 StreamEvent 新增变体
`src/llm/types/response_v2.rs``StreamEvent` 枚举中追加两个变体:
```rust
/// 工具开始执行 —— 在 ToolCallEnd 之后、registry.invoke 之前发出。
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
ToolExecutionStarted {
tool_name: String,
tool_call_id: String,
/// 工具参数(JSON 字符串形式)
arguments: String,
},
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
ToolExecutionCompleted {
tool_name: String,
tool_call_id: String,
/// 结果摘要(前 200 字符)
result_summary: String,
/// 是否出错
is_error: bool,
},
```
`PartialMessageResponse::apply_to` 中追加:
```rust
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
```
这两个是**元事件**,不参与内容块累积,`apply_to` 直接返回 `true`
### 3.4 新增方法签名
**`LlmCycle` 层**`src/llm/cycle.rs`):
```rust
/// 提交消息并自动处理工具调用循环,流式产出所有事件。
///
/// 与 `submit_with_tools` 的区别:
/// - LLM 响应是流式的(全程 `chat_stream` 而非 `chat`
/// - 工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
/// - 错误以 `StreamEvent::Error` 形式出现在流中,而非终止 `Result`
/// - 消费方需手动 `push_message()` 同步消息历史
///
/// **运行时要求**:内部使用 `tokio::spawn`,需要 tokio 多线程运行时。
pub async fn submit_with_tools_stream(
&mut self,
prompt: String,
tool_registry: Arc<ToolRegistry>,
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, LlmError>
```
```rust
/// 运行工具循环的核心异步状态机。
///
/// 接收 owned 字段,通过 mpsc::unbounded_channel 产出事件序列。
/// 由 `submit_with_tools_stream` 在 tokio::spawn 中调用。
///
/// **运行时要求**:此函数内部使用 `tokio::spawn`,要求调用方运行在
/// tokio 多线程运行时中(`#[tokio::main]` 或 `#[tokio::test(flavor = "multi_thread")]`)。
/// 不在 WASM 目标下可用。
async fn run_tool_loop(
messages: Vec<Message>,
provider: Arc<dyn LlmProvider>,
config: CycleConfig,
tool_registry: Arc<ToolRegistry>,
tools: Vec<ToolDef>,
tx: mpsc::UnboundedSender<StreamEvent>,
hook_executor: Option<Arc<HookExecutor>>,
)
```
**`AgentSession` 层**`src/agent/session.rs`):
```rust
/// 提交一轮对话(流式,含自动 tool 循环),返回 `StreamEvent` 流。
///
/// 与 `submit_turn` 的区别:
/// - 以流事件序列而非 `MessageResponse` 返回
/// - 工具执行期间插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
/// - 消费方在收到 `MessageComplete` 后需手动调用 `finalize_turn` 同步状态
///
/// **运行时要求**:内部委托 `submit_with_tools_stream`,需要 tokio 多线程运行时。
pub async fn submit_turn_stream(
&mut self,
user_input: impl Into<String>,
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, AgentError>
/// 完成一轮 turn:累计 cost + 触发 OnTurnEnd hook。
///
/// 由消费者在收到 `MessageComplete.full_response` 后调用。
pub async fn finalize_turn(&mut self, response: &MessageResponse)
```
### 3.5 消费者使用模式
```rust
use futures_util::StreamExt;
let mut stream = session.submit_turn_stream("计算 1+2").await?;
let mut final_response = None;
while let Some(event) = stream.next().await {
match &event {
StreamEvent::TextDelta { text } => print!("{}", text),
StreamEvent::ToolExecutionStarted { tool_name, arguments, .. } => {
println!("\n🔧 [{}({})]", tool_name, arguments);
}
StreamEvent::ToolExecutionCompleted { result_summary, .. } => {
println!("{}", result_summary);
}
StreamEvent::MessageComplete { full_response } => {
final_response = Some(full_response.clone());
}
_ => {}
}
}
std::io::stdout().flush().ok();
if let Some(response) = final_response {
session.finalize_turn(&response).await;
}
```
> **⚠️ 消费者注意**`finalize_turn` 是开发者责任 —— 遗漏调用会导致 `cost_so_far` 不累计、`OnTurnEnd` hook 不触发。session 状态仍然可用,后续 `submit_turn` 也能正常执行,但 cost 信息不完整。`finalize_turn` 无自动补偿机制,建议使用 `Drop` guard 或在 `while` 循环的 `finally` 块中确保调用。
### 3.6 run_tool_loop 核心逻辑
`run_tool_loop` 是此方案的核心状态机(约 90 行),其伪代码逻辑如下:
```
1. 接收 owned 字段:messages, provider, config, tool_registry, tools, tx, hook_executor
2. max_turns = config.max_tool_turns.unwrap_or(10)
// None → 10(退化为默认值),Some(n) → n
// 与非流式 submit_with_tools 行为一致
3. 工具循环(for round in 1..=max_turns):
a. build_request(messages, tools)
// 空 tool_registry 时 tools 为空列表,流退化为纯文本流(可安全运行)
b. PreRequest hook(如果有 hook_executor
c. 发起流式 LLM 调用:
let stream = match provider.chat_stream(request).await {
Ok(s) => s,
Err(e) => {
// 第一层错误:chat_stream 自身失败(网络/认证/限流)
// 这里不做 retryretry 逻辑留给上层循环的 submit_request 模式,
// 流式场景中 retry 需重新建立 mpsc 通道,复杂度与收益不匹配
tx.send(StreamEvent::Error { message: e.to_string() }).ok();
return; // 直接结束 task
}
};
d. 消费 LLM 流:
- PartialMessageResponse::new()
- while let Some(result) = stream.next().await
- match result:
Ok(event) → apply_to + tx.send(event)
Err(e) → tx.send(Error { message }) + break
// 第二层错误:stream 内部事件错误(如 chunk 解析失败)
e. partial.finalize()? → response
f. push response.message → messages
g. 检查 has_tool_calls_in_response(&response)
h. 如果没有 tool_use: break(最终轮,流已自然结束)
i. 如果有 tool_use:
- extract_tool_calls_from_response(&response)
- tx.send(ToolExecutionStarted { tool_name, tool_call_id, arguments })
- registry.invoke_all(calls, tool_timeout).await
- for result in results:
tx.send(ToolExecutionCompleted { tool_name, tool_call_id, result_summary, is_error })
- push tool results → messages
- continue(新一轮 LLM 流)
4. 流结束(tokio::spawn 自然退出)
```
> **关于 LLM 请求 retry**:非流式 `submit_with_tools` 内部通过 `submit_request` 的 retry 循环处理临时错误。流式版本 `run_tool_loop` **不在内部实现 retry**。原因:(1)retry 需要重新建立 mpsc 通道和事件流上下文,复杂度与收益不匹配;(2)`unbounded_channel` 已发出的事件无法撤回。如果需要 retry 语义,调用方应在上层做 fallback 策略,或在 `llm provider` 实现层完成 retry(如 `RetryProvider` 包装器)。
**错误处理**
| 场景 | 行为 |
|------|------|
| LLM 请求失败(`chat_stream` 返回 `Err` | `tx.send(Error { message })` + `return` 结束 task。**不做 retry**(见上方说明) |
| LLM 流内事件错误(stream Item 的 `Err` | `tx.send(Error { message })` + `break` 结束当轮流,终止循环 |
| 可恢复工具错误(`is_recoverable() == true` | 作为 tool result 回传 LLM,流继续,不出 Error 事件 |
| 不可恢复工具错误(`is_recoverable() == false` | `tx.send(Error { message })` + 终止循环 |
| 工具超时(`tokio::time::timeout` | 视为不可恢复,`tx.send(Error)` + 终止 |
| 最大工具循环轮次超限 | `tx.send(Error { "达到最大工具循环轮次" })` + 终止 |
| spawn task 内部 panic | 由于 `JoinHandle` 不保存(detached),panic 由 tokio 运行时静默捕获;消费者看到 stream 直接结束(返回 `None`),无 `Error` 事件。建议在 `run_tool_loop` 内部避免 `unwrap()`,所有可失败路径通过 `Result` + `?` 传播 |
### 3.7 修改文件清单
| # | 文件 | 改动 | 估算行数 |
|---|------|------|---------|
| 1 | `llm/types/response_v2.rs` | +2 `StreamEvent` 变体 +2 `apply_to` arm | ~20 |
| 2 | `llm/cycle.rs` | +`submit_with_tools_stream` 方法 + `run_tool_loop` 模块函数 | ~140 |
| 3 | `llm/cycle.rs` | `CycleConfig``#[derive(Clone)]` | ~1 |
| 4 | `agent/session.rs` | +`submit_turn_stream` + `finalize_turn` | ~70 |
| — | **测试**(内联) | 4 个场景测试(纯度本、单轮、多轮、超限) | ~150 |
| | **合计** | | **~380** |
> 注:`RetryConfig` 已标注 `#[derive(Debug, Clone)]`,无需额外修改。
---
## 4. 实现计划
按 5 个 Step 增量实施,每步可独立编译和测试。
### Step 1 — 基础设施准备
**目标**:数据层就绪,为流事件新增变体和配置 Clone 奠基。
**改动**
- `llm/types/response_v2.rs`
- `StreamEvent` 枚举追加 `ToolExecutionStarted` / `ToolExecutionCompleted` 变体
- `PartialMessageResponse::apply_to` 追加两个新变体的 arm(均返回 `true`
- `llm/cycle.rs`
- `CycleConfig``#[derive(Clone)]`(所有字段为基础类型 + `RetryConfig`
**验证**`cargo build` 通过
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
**目标**:实现流式工具循环的核心状态机,这是整个 Phase 9 的技术关键。
**改动**
- `llm/cycle.rs`
- 新增 `run_tool_loop()` 模块函数(约 90 行),基于 `mpsc::unbounded_channel` 通信
- 新增 `submit_with_tools_stream()` 公开方法,入口参数为 `prompt` + `Arc<ToolRegistry>`
- 内部 `tokio::spawn` 启动 `run_tool_loop`,返回 `rx` 端作为 `dyn Stream`
**验证**`cargo build` 通过
### Step 3 — 单元测试(LlmCycle 层)
**目标**:验证 `submit_with_tools_stream` 在 8 个核心场景下的行为和事件序列正确性(含 §8 Step 3 扩展的工具错误路径)。
**新增**`llm/cycle.rs` 内联测试 `#[cfg(test)]`):
| 场景 | Mock 响应序列 | 验证点 |
|------|---------------|--------|
| 1 — 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` |
| 2 — 单轮工具调用 | 2 个响应:tool_use → text | 包含 `ToolExecutionStarted` + `ToolExecutionCompleted`;最终 `stop_reason``Stop` |
| 3 — 多轮工具调用 | 4 个响应:3 × tool_use → 1 × text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度正确 |
| 4 — 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `Error` 事件;消息历史停在第 2 轮 |
**验证**`cargo test` 全部通过
### Step 4 — `AgentSession` 层包装
**目标**:为 `AgentSession` 新增流式会话接口,保持与 `submit_turn` 一致的行为语义。
**改动**
- `agent/session.rs`
- `submit_turn_stream(user_input)` — 触发 `OnTurnStart` hook → 组装 `LlmCycle` → 调用 `submit_with_tools_stream``turn_index += 1` → 返回流
- `finalize_turn(response)``cost_so_far.add(&response.usage)` → 触发 `OnTurnEnd` hook
**验证**`cargo build` 通过
### Step 5 — 集成测试 + 扫尾
**目标**:端到端验证 `submit_turn_stream` + `finalize_turn` 的完整链路,确保零回归。
**新增**`agent/session.rs` 内联测试):
- **场景**`submit_turn_stream` 跑通 mock provider → 消费流(验证各事件到达) → `finalize_turn` 后 cost 更新正确
- **场景**verify `OnTurnStart` hook 在 `submit_turn_stream` 返回流之前已触发
**验证**
```bash
cargo test --all-targets # 全绿,存量测试 0 回归
cargo clippy --all-targets -- -D warnings # 0 警告
```
---
## 5. 运行细节
### 5.1 `run_tool_loop` 的 spawn 生命周期
#### 执行模型:立即执行 vs 惰性流
`submit_with_tools_stream` 采用 **立即执行** 模型(`tokio::spawn` + `mpsc`),这与 `submit_stream`**惰性执行**`async_stream::stream!` 宏,消费者首次 `next()` 时才触发 LLM 调用)不同。
**选择理由**:工具循环是 **不确定轮次的** —— 每个工具执行的结果可能影响后续 LLM 调用。惰性流无法表达这种"边消费边控制"的语义。通过 `tokio::spawn` 将工具循环移到独立 task 中运行,使得:
- 消费者可以随时开始消费(不丢失事件)
- 工具循环在后台独立运行,不受消费者消费节奏影响
- `mpsc::unbounded_channel` 作为事件缓冲区,解耦生产者与消费者
**对消费者的影响**`submit_with_tools_stream().await?` 返回时,工具循环可能已经开始执行(事件已开始写入 channel)。消费者应尽快开始 `while let Some(event) = stream.next().await`,避免 channel 缓冲过多事件。如果在返回流后长时间不消费,事件会堆积在 mpsc buffer 中(内存开销,无阻塞风险 —— 见 §6 风险表)。
#### 生命周期
```
submit_with_tools_stream()
├─ mpsc::unbounded_channel() → (tx, rx)
├─ messages.push(user_text(prompt))
├─ compact check
├─ tokio::spawn(run_tool_loop(messages, provider, config, ..., tx))
└─ return Box::pin(rx) as dyn Stream
[用户消费 stream]
└─ while let Some(event) = rx.recv().await { yield event }
[用户 drop rx / 结束循环]
└─ rx 被 drop → tx.send() 返回 Err
→ run_tool_loop 检测到 tx.closed()
→ break → task 自然终止
```
#### JoinHandle 与 panic 处理
`run_tool_loop``JoinHandle` 在 spawn 后**不保存**detached pattern)。panic 由 tokio 运行时捕获并通过 `tracing::error` 记录:
```rust
// submit_with_tools_stream 内部
tokio::spawn(async move {
run_tool_loop(..., tx).await;
});
```
如果 `run_tool_loop` 内部发生 panic(如 `unwrap()`),tokio 的 `spawn` 会静默吞掉 panic 并终止 task。消费者此时看到 stream 直接返回 `None`,不会收到 `StreamEvent::Error`。实际编码中应避免 `unwrap()`,所有 `Result` 使用 `?``match` 处理。
Rx 侧实现 `Stream` trait:使用 `tokio_stream::wrappers::UnboundedReceiverStream` 包装 `mpsc::UnboundedReceiver`,因为 `mpsc::UnboundedReceiver` 本身不实现 `Stream``tokio-stream = "0.1"` 已在 `Cargo.toml` 中存在)。
### 5.2 消息历史同步
`submit_with_tools_stream` 内部由 `run_tool_loop` 管理 `messages` 的拷贝,不会写入 `self.messages`。消费方在收到 `MessageComplete` 后需手动:
```rust
let response = full_response.clone();
cycle.push_message(response.message.clone());
```
`AgentSession::submit_turn_stream` 中,由于流是延迟求值且 `&mut self` 无法进入 spawn 闭包,消息历史同步交由消费方在 `finalize_turn` 前自行决定。当前方案中 `submit_turn_stream` **不自动同步消息历史**,这与 `submit_stream` 的已有行为一致(ponytail: Phase 2 FIX-E 注释)。
---
## 6. 风险评估
| 风险 | 影响 | 缓解措施 |
|------|------|---------|
| `&mut self` 约束导致流内无法访问 session 状态 | 中 | 复用 `submit_stream` 已有模式:方法体内读取 `self` 后构建 owned 数据,spawn 闭包不捕获 `&mut self` |
| spawn task 生命周期管理 | 低 | 用户 drop rx → `tx.send` 返回 `Err``run_tool_loop` 自然终止 |
| spawn task panic 静默丢失 | 中 | `run_tool_loop` 内部使用 `match`/`?` 避免 `unwrap()``JoinHandle` 不做 `await`detached),panic 由 tokio 运行时记录日志。消费者看到 stream 提前结束(收到 `None`)但无 Error 事件 |
| 中间轮 cost 不累加到 `cost_so_far` | 低 | 与现有 `submit_turn` 行为一致(仅最终轮计入),标记为已知限制,不在此 Phase 修复 |
| 工具循环中 hook 可用性 | 低 | `PreRequest`/`PostRequest` hook 通过 `hook_executor.clone()` 进入 spawn taskhook 在 `run_tool_loop` 循环内触发 |
| `run_tool_loop` 不支持消息压缩 | 中 | 长工具循环中消息不断增长,可能超出 context window。当前不传递 `compact_config`,后续可扩展 `run_tool_loop` 签名增添此参数 |
| `unbounded_channel` 在消费慢于生产时内存增长 | 低 | LLM 流式输出天然有节流(token 生成速度远慢于 CPU 处理速度),消费者通常快于生产者。后续如需背压可切换为 `mpsc::channel(N)` + backpressure |
| 流式版本不做 LLM retry | 低 | 非流式 `submit_with_tools` 通过 `submit_request` 的 retry 循环处理临时错误。流式版本中 retry 需重建 mpsc 通道,复杂度不匹配。调用方可使用 `RetryProvider` 包装器或在 Provider 层实现 retry |
| 执行模式与 `submit_stream` 不一致(立即 vs 惰性) | 低 | `submit_stream` 的惰性语义不适配需要后台执行的工具循环。消费者应在 `submit_turn_stream` 返回后尽快消费流事件 |
| `tokio::spawn` 要求 tokio 多线程运行时 | 低 | agcore 已依赖 tokio,涉及 IO 的 API 均使用 async。`#[tokio::test]` 单线程运行时不支持 `spawn`,测试中将 `run_tool_loop` 提取为可独立调用的函数,测试不走 spawn 直接调用 |
| `CycleConfig``Clone` 影响现有代码 | 无 | 纯配置 struct,所有字段是基础类型或已 `Clone``RetryConfig` |
---
## 7. 验收标准
| # | 验收项 | 验证方式 |
|---|--------|---------|
| 1 | `cargo build --all-targets` 通过 | ✅ 编译器无错误 |
| 2 | `submit_with_tools_stream` 纯文本流事件序列正确 | 单元测试验证:事件类型、顺序与 `submit_stream` 一致 |
| 3 | `submit_with_tools_stream` 单轮工具调用事件序列正确 | 单元测试验证:含 `ToolExecutionStarted` / `ToolExecutionCompleted` |
| 4 | `submit_with_tools_stream` 多轮工具调用事件序列正确 | 单元测试验证:多对 `ToolExecutionStarted`/`ToolExecutionCompleted` |
| 5 | `submit_with_tools_stream` 最大轮次超限产生 Error 事件 | 单元测试验证:流中出现 `StreamEvent::Error` |
| 6 | `submit_turn_stream` + `finalize_turn` 端到端链路 | 集成测试验证:cost 更新 + hook 触发 |
| 7 | `cargo test --all-targets` 全绿,存量测试 0 回归 | ✅ 无回归 |
| 8 | `cargo clippy --all-targets -- -D warnings` 0 警告 | ✅ 无警告 |
| 9 | 现有 `submit_turn` / `submit_with_tools` / `submit_stream` 行为零影响 | ✅ 存量测试通过 |
---
---
## 8. 实施计划
按 5 个 Step 分阶段实施,每步产出独立 commit,可验证后退。
### 依赖关系
```mermaid
graph LR
S1["Step 1: 基础设施"]:::s1
S2["Step 2: 核心状态机"]:::s2
S3["Step 3: LlmCycle 单元测试"]:::s3
S4["Step 4: AgentSession 包装"]:::s4
S5["Step 5: 集成测试 + 扫尾"]:::s5
S1 --> S2
S1 --> S4
S2 --> S3
S2 --> S4
S3 --> S5
S4 --> S5
classDef s1 fill:#e2e8f0,stroke:#94a3b8
classDef s2 fill:#fbbf24,stroke:#d97706
classDef s3 fill:#93c5fd,stroke:#2563eb
classDef s4 fill:#93c5fd,stroke:#2563eb
classDef s5 fill:#4ade80,stroke:#16a34a
```
| Step | 依赖 | 并行机会 |
|------|------|---------|
| S1 | 无 | — |
| S2 | S1 | 可与 S4 并行 |
| S3 | S2 | 阻塞,需 S2 完成 |
| S4 | S1, S2 | 功能依赖 S2(调用 `submit_with_tools_stream`);文件级无重叠但需先编译过 S2 |
| S5 | S3 + S4 | 需 S3 和 S4 都完成 |
### Step 1 — 基础设施准备
**工作量**S< 1h
**风险**:低(纯新增,不影响现有代码逻辑)
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|---|------|---------|---------|------|
| 1.1 | `StreamEvent` 枚举追加 `ToolExecutionStarted` 变体 | `llm/types/response_v2.rs` | 无 | 低 |
| 1.2 | `StreamEvent` 枚举追加 `ToolExecutionCompleted` 变体 | `llm/types/response_v2.rs` | 1.1 | 低 |
| 1.3 | `PartialMessageResponse::apply_to` 追加两个元事件 arm(均返回 `true` | `llm/types/response_v2.rs` | 1.2 | 低 |
| 1.4 | `CycleConfig``#[derive(Clone)]` | `llm/cycle.rs` | 无 | 低 |
**验收条件**
- `cargo build` 通过,编译器无 warning
- 新增的 `StreamEvent` 变体可通过 `serde` roundtrip 序列化/反序列化
- `CycleConfig` 可正常 clone
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
**工作量**M1-4h
**风险**:中(核心实现,需正确设计 spawn + mpsc 生命周期)
**前置依赖**S1
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|---|------|---------|---------|------|
| 2.1 | 实现 `run_tool_loop()` 模块函数:消息循环构建请求 → `chat_stream` → 消费流 → 检测 tool_use → 工具执行 → 新一轮 | `llm/cycle.rs` | S1 | 中 |
| 2.2 | 实现 `submit_with_tools_stream()` 公开方法:提取字段 → spawn `run_tool_loop` → 返回 `UnboundedReceiverStream` | `llm/cycle.rs` | 2.1 | 中 |
| 2.3 | 新增导入:`tokio::sync::mpsc``tokio_stream::wrappers::UnboundedReceiverStream` | `llm/cycle.rs` | 2.2 | 低 |
**关键实现细节**
```rust
// run_tool_loop 函数签名
async fn run_tool_loop(
mut messages: Vec<Message>,
provider: Arc<dyn LlmProvider>,
config: CycleConfig,
tool_registry: Arc<ToolRegistry>,
tools: Vec<ToolDef>,
tx: mpsc::UnboundedSender<StreamEvent>,
hook_executor: Option<Arc<HookExecutor>>,
) {
let max_turns = config.max_tool_turns.unwrap_or(10);
let tool_timeout = config.tool_timeout_secs;
let max_bytes = config.max_tool_result_bytes;
let mut round = 0u32;
loop {
round += 1;
if round > max_turns {
// §3.6 错误表:最大轮次超限 → Error 事件 + 终止
let _ = tx.send(StreamEvent::Error { message: "达到最大工具循环轮次".to_string() });
break;
}
// ① 构建请求
let request = MessageRequest {
model: config.model.clone(),
messages: messages.clone(),
tools: tools.clone(),
tool_choice: ToolChoice::Auto,
max_tokens: config.max_tokens,
temperature: config.temperature,
..Default::default()
};
// ② PreRequest hook
// ...
// ③ chat_stream
let stream = match provider.chat_stream(request).await {
Ok(s) => s,
Err(e) => {
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
return;
}
};
// ④ 消费流
let mut partial = PartialMessageResponse::new();
let mut stream = stream;
while let Some(result) = stream.next().await {
match result {
Ok(event) => {
partial.apply_to(&event);
if tx.send(event).is_err() { return; }
}
Err(e) => {
// ponytail: 流内事件错误后 partial 处于损坏状态,
// 不能继续执行 finalize/finalize —— 直接 return 结束 task
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
return;
}
}
}
// ⑤ finalize
let response = match partial.finalize() {
Ok(r) => r,
Err(e) => { let _ = tx.send(StreamEvent::Error { .. }); return; }
};
messages.push(response.message.clone());
// ⑥ 检测 tool_use
if !has_tool_calls_in_response(&response) {
break; // 最终轮
}
// ⑦ 执行工具
let tool_calls = extract_tool_calls_from_response(&response);
let calls: Vec<_> = tool_calls.into_iter()
.map(|(id, name, args)| {
let value = serde_json::from_str(&args).unwrap_or(Value::Null);
(id, name, value)
}).collect();
for (tool_call_id, tool_name, args_value) in &calls {
let args_json = serde_json::to_string(&args_value).unwrap_or_default();
if tx.send(StreamEvent::ToolExecutionStarted {
tool_name: tool_name.clone(),
tool_call_id: tool_call_id.clone(),
arguments: args_json,
}).is_err() { return; }
}
let results = tool_registry.invoke_all(calls, tool_timeout).await;
for result in &results {
let summary = match &result.output {
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
Err(e) => e.to_string(),
};
// ponytail: 复用现有 truncate_tool_result 函数(cycle.rs 末尾),
// 确保多字节 UTF-8 字符不被截断破坏。上限 200 字符。
let truncated = truncate_tool_result(&summary, 200);
if tx.send(StreamEvent::ToolExecutionCompleted {
tool_name: result.tool_name.clone(),
tool_call_id: result.tool_call_id.clone(),
result_summary: truncated,
is_error: result.output.is_err(),
}).is_err() { return; }
}
for result in results {
let is_error = result.output.is_err();
let content = match &result.output {
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
Err(e) if e.is_recoverable() => format!("错误: {}", e),
Err(e) => {
let _ = tx.send(StreamEvent::Error { .. });
return;
}
};
messages.push(Message::tool_result(result.tool_call_id, content, is_error));
}
}
}
```
**验收条件**
- `cargo build` 通过
- 新增方法签名与方案设计一致
- 未修改现有 `submit_with_tools`/`submit_stream` 的行为
### Step 3 — 单元测试(LlmCycle 层)
**工作量**M1-4h
**风险**:低(与现有测试模式一致,使用已有 MockProvider
**前置依赖**S2
测试策略:直接使用公开的 `crate::llm::mock::MockProvider`(已完整实现 `chat_stream` + 预设响应队列),避免改造 `cycle.rs` 测试模块内的内联 Stub。测试中调用 `submit_with_tools_stream` 时通过 `#[tokio::test(flavor = "multi_thread")]` 满足 spawn 运行时要求,或在单元级将 `run_tool_loop` 作为独立函数直接测试(不走 spawn)。
| # | 测试场景 | Mock 响应序列 | 验证点 | 覆盖路径 |
|---|---------|---------------|--------|---------|
| 3.1 | 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` | 正常路径:单轮 LLM → 文本返回 |
| 3.2 | 单轮工具调用 | 2 个响应:tool_use + text | 包含一对 `ToolExecutionStarted`/`ToolExecutionCompleted`;最终 `stop_reason``Stop` | 正常路径:LLM → 工具 → LLM |
| 3.3 | 多轮工具调用 | 4 个响应:3×tool_use + 1×text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度为 8user + 3×(assistant+tool) + final assistant | 正常路径:LLM → 工具 → LLM → 工具 → LLM |
| 3.4 | 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `StreamEvent::Error`;消息历史停在第 2 轮 | 边界条件:超出上限 |
| 3.5 | `chat_stream` 返回 Err | Mock `chat_stream` 返回 `Err(LlmError::Other(...))` | 流中第一个事件为 `StreamEvent::Error`;随后流结束 | 异常路径:LLM 不可用 |
| 3.6 | 空 tool_registry | 1 个 text 响应,registry 中无工具 | 流退化为纯文本流,事件序列与 3.1 一致 | 退化场景:无工具可用 |
| 3.7 | 不可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(不可恢复) | 流中出现 `StreamEvent::Error`;消息历史中不含该工具结果(循环终止前未 push) | 异常路径:工具执行失败 |
| 3.8 | 可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(可恢复) | 工具结果作为 `ToolResult { is_error: true }` 回传 LLM;流正常结束,无 `Error` 事件 | 异常路径:工具出错但可恢复 |
| 3.9 | 工具超时 | 2 个响应:tool_use → text`tool_timeout_secs: 1`,模拟工具耗时 10 秒 | 流中出现 `StreamEvent::Error`;循环终止前未 push 工具结果 | 异常路径:工具执行超时 |
**验收条件**
- `cargo test` 新增 8 个测试全部通过
- `cargo test` 存量测试 0 回归
### Step 4 — `AgentSession` 层包装
**工作量**S< 1h
**风险**:低(薄包装层,逻辑简单)
**前置依赖**S1
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|---|------|---------|---------|------|
| 4.1 | 实现 `submit_turn_stream()`:触发 `OnTurnStart` → 组装 `LlmCycle` → 调用 `submit_with_tools_stream``turn_index += 1` → 返回流 | `agent/session.rs` | S1 | 低 |
| 4.2 | 实现 `finalize_turn()``cost_so_far.add()` → 触发 `OnTurnEnd` hook | `agent/session.rs` | S1 | 低 |
**验收条件**
- `cargo build` 通过
- 新增方法签名与方案设计一致
-`submit_turn` 的 system_prompt / compact_config / bundle 使用方式一致
### Step 5 — 集成测试 + 扫尾
**工作量**S< 1h
**风险**:低(基于现有测试框架)
**前置依赖**S3 + S4
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|---|------|---------|---------|------|
| 5.1 | `submit_turn_stream` 端到端测试:跑通 mock provider → 消费流验证各事件到达 → `finalize_turn` 后 cost 更新正确 | `agent/session.rs`(内联测试) | S4 | 低 |
| 5.2 | Hook 触发验证:`OnTurnStart``submit_turn_stream` 返回流之前触发;`finalize_turn` 调用后 `OnTurnEnd` 正确触发 | `agent/session.rs`(内联测试) | S4 | 低 |
| 5.3 | `cargo test --all-targets` 全绿验证 | 全仓 | S5.1+S5.2 | 低 |
| 5.4 | `cargo clippy --all-targets -- -D warnings` 0 警告 | 全仓 | S5.3 | 低 |
| 5.5 | `cargo build --all-targets` 发布模式验证 | 全仓 | S5.4 | 低 |
**验收条件**
- 全量测试通过,存量 0 回归
- clippy 0 警告
- 发布模式零 warning
### 实施总览
| | Step 1 | Step 2 | Step 3 | Step 4 | Step 5 | **合计** |
|--|--------|--------|--------|--------|--------|---------|
| **工作量** | S | M | M | S | S | **M-L** |
| **文件数** | 2 | 1 | 1(内联) | 1 | 1(内联) | **~4** |
| **代码行** | ~20 | ~140 | ~150 含测试 | ~70 | ~80 含测试 | **~380** |
| **风险** | 低 | 中 | 低 | 低 | 低 | 中 |
| **并行** | — | 阻塞(S4 依赖 S2) | 阻塞 | 阻塞(依赖 S2) | 阻塞 | — |
---
## 附录 A:新增 StreamEvent 变体的 apply_to 语义
```rust
// 在 PartialMessageResponse::apply_to 中追加:
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => {
// 元事件:不参与内容块累积,不修改 partial response 状态
true
}
```
## 附录 BCycleConfig 的 Clone 推导
```rust
/// LLM 调用周期配置。
#[derive(Debug, Clone)] // ← 追加 Clone
pub struct CycleConfig {
pub model: String,
pub max_tokens: Option<u32>,
pub temperature: Option<f32>,
pub max_turns: Option<u32>,
pub retry: RetryConfig, // 已 #[derive(Clone)]
pub max_tool_turns: Option<u32>,
pub tool_timeout_secs: u64,
pub max_tool_result_bytes: usize,
}
```
File diff suppressed because it is too large Load Diff
+344
View File
@@ -0,0 +1,344 @@
# 笔记:opencode 子代理调度、分发与合并及工作流推进
> 基于 `/Users/midnite/Samples/opencode` 源码调研,2026-07-04
---
## 一、整体架构
```
LLM(主 Agent
├── 调用 Task tooltool call
│ ↓
│ TaskTool.execute() ← packages/opencode/src/tool/task.ts
│ │
│ ├── agent.get() ← 查找 Agent 定义(agent.ts
│ ├── deriveSubagentPermission() ← 权限合并(subagent-permissions.ts
│ ├── sessions.create() ← 创建子 session
│ │
│ ├── [前台] background.wait() + background.waitForPromotion() race
│ │ ↓ 完成
│ │ renderOutput() → XML <task> 标签返回
│ │
│ └── [后台] background.start() → notify() 异步注入结果
└── 会话循环(runLoop) ← prompt.ts
├── 检测 subtask type part → handleSubtask()
├── 检测 compaction → compaction.process()
└── 正常流程 → LLM.stream() → processor.handleEvent()
```
---
## 二、子代理调度(Dispatch
### 2.1 三种触发入口
| 入口 | 触发方式 | 调用链路 |
|------|---------|---------|
| A — LLM 自主 | LLM 调用 `task` tool | 系统提示词中注入了 Task tool 描述 + `describeTask()` 输出子代理列表 → LLM 决策 |
| B — `subtask` part | 消息中有 `type: "subtask"` 的 part | `handleSubtask()` 直接执行 TaskTool,不走 LLM |
| C — `agent` part | 消息中有 `type: "agent"` 的 part | 转为"调用 task tool 带 subagent: XXX"的提示词,引导 LLM |
### 2.2 TaskTool.execute() 完整流程(task.ts
```
execute(params, ctx):
1. background 开关检查(需 experimental flag
2. ctx.ask() 权限询问
3. agent.get(subagent_type) 查找子代理定义
4. task_id 存在 → sessions.get(task_id) 恢复已有子 session
task_id 不存在 → sessions.create() 创建新子 session
5. deriveSubagentSessionPermission() 合并权限
6. 添加默认 deny 规则(todowrite / task
7. 确定 model(继承或子代理自定义)
8. 执行 runTask() → ops.resolvePromptParts() + ops.prompt()
9. 结果格式化为 XML ← renderOutput()
```
### 2.3 关键:子 session 创建(task.ts lines 121-158
```typescript
// 权限继承
const childPermission = deriveSubagentSessionPermission({
parentSessionPermission: parent.permission ?? [],
subagent: next,
})
// 默认 deny 规则
const childToolDenies = [
// 子代理自己的 permission 没允许 todowrite → 默认 deny
...(next.permission.some(r => r.permission === "todowrite") ? []
: [{ permission: "todowrite", pattern: "*", action: "deny" }]),
// 子代理自己的 permission 没允许 task → 默认 deny(防嵌套)
...(next.permission.some(r => r.permission === "task") ? []
: [{ permission: "task", pattern: "*", action: "deny" }]),
// 主 agent 专有工具也不给子代理
...(cfg.experimental?.primary_tools?.map(p => ({ permission: p, ... })) ?? []),
]
```
---
## 三、通信格式:Tool Call / Tool Result
### 3.1 父→子:Task tool 参数
```
{
subagent_type: "explore" | "general" | ...,
description: "简短描述(3-5词)",
prompt: "子代理的完整任务描述",
task_id?: "恢复已有子 session 时使用",
command?: "触发该调用的 CLI 命令(可选)",
background?: true // 后台模式(需 experimental flag
}
```
### 3.2 子→父:XML 包装的纯文本(renderOutput
```xml
<task id="ses_xxxxx" state="completed">
<summary>任务简述</summary>
<task_result>
子 agent 输出的完整文本内容...
</task_result>
</task>
```
错误时:
```xml
<task id="ses_xxxxx" state="error">
<summary>任务失败</summary>
<task_error>
Error: 具体错误信息...
</task_error>
</task>
```
### 3.3 传递给 LLM 的方式
**前台模式**
```
TaskTool.execute() 返回 { output: "<task>...</task>" }
AI SDK 将其转为 tool result,存入数据库 tool part
下一轮 LLM 调用时,tool result 作为消息历史的一部分传入
LLM 看到 XML,自行解析使用
```
**后台模式**
```
TaskTool.execute() 立即返回 <task state="running">...
子 agent 完成后 → background.wait() 触发 → inject()
向父 session 注入合成 text partsynthetic: true
携带 <task state="completed">... 结果
父 LLM 在下一轮循环中看到该消息
```
---
## 四、分发与合并(Distribution & Merge
### 4.1 并行分发
- **无专用分发层**。依赖 LLM 在单条消息中发出多个 tool call
- `task.txt` 引导 LLM*"Launch multiple agents concurrently whenever possible"*
- 底层通过 Effect.ts 的 `Effect.forkIn(scope, { startImmediately: true })` 实现同一消息内多 tool call 并发
- **子 agent 之间完全隔离**,无直接通信
### 4.2 结果合并
**无专用合并逻辑。** 合并完全通过 LLM 的上下文理解完成:
- 前台:tool result 自然进入消息历史,LLM 下一轮读取
- CLI 命令:额外注入 "Summarize the task tool output above and continue with your task." 引导 LLM 总结
- LLM 自主调用:无额外引导,LLM 自行决定如何使用
### 4.3 前台/后台切换机制(task.ts lines 303-333
```typescript
// 前台执行
return yield* Effect.raceFirst(
background.wait({ id: nextSession.id }), // 等完成
background.waitForPromotion(nextSession.id), // 等 promote 到后台
)
```
当用户将前台任务 promote 到后台时,`waitForPromotion` 先返回(标记 `metadata.background = true`),TaskTool 转而返回后台模式的输出。
### 4.4 后台作业引擎(core/background-job.ts
纯内存、非持久化注册表。使用 Effect.ts 的 `SynchronizedRef` 做并发控制。
| 操作 | 行为 |
|------|------|
| `start()` | 创建 jobfork run effect,返回 info |
| `extend()` | 追加顺序执行的 run(通过 `Deferred` 链式等待前一个完成) |
| `wait()` | `Deferred.await(done)`,可选 timeout |
| `waitForPromotion()` | 等待 `promoted` Deferred 或检测 `background` 标记 |
| `promote()` | 标记 `background = true`,触发 `onPromote` callback |
| `cancel()` | 设置 `cancelled`close scope(中断所有子 fork |
---
## 五、工作流推进(Workflow Progression
### 5.1 核心循环(prompt.ts → runLoop
```
runLoop(sessionID):
while true:
1. MessageV2.filterCompactedEffect() 获取消息
2. MessageV2.latest() 取最近 user/assistant/tasks
3. 检查 finish 状态
- 不是 tool-calls 且有 finish → break(退出循环)
4. 取 taskssubtask / compaction 队列)
- subtask → handleSubtask() → continue
- compaction → compaction.process() → continue/break
5. 检查 overflow → 自动创建 compaction task → continue
6. 构建 assistant message
7. SessionProcessor.create() 创建 handle
8. SessionTools.resolve() 解析所有工具
9. 构建 system prompt(环境信息 + skills + MCP + instructions
10. handle.process() — 启动 LLM stream
11. 检查 result:
- "compact" → 返回给外层触发 compaction
- "stop" → break
- "continue" → 继续循环
```
### 5.2 SessionProcessor 事件处理(processor.ts
| Stream 事件 | 处理逻辑 |
|------------|---------|
| `reasoning-start/delta/end` | 创建 reasoning part → 增量追加 → 最终持久化 |
| `tool-input-start/delta/end` | 创建/更新 tool partpending 状态) |
| `tool-call` | 标记 running → 设置 input → **doom loop 检测** |
| `tool-result` | `completeToolCall()` → 持久化结果 + 附件 |
| `tool-error` | `failToolCall()` → 标记错误 |
| `provider-error` | 抛出异常 → 触发重试 |
| `text-start/delta/end` | 流式文本 → `updatePartDelta()` **增量持久化** |
| `step-start` | 创建快照(snapshot |
| `step-finish` | 生成 patch diff → 更新 usage/tokens → **overflow 检测** → 触发 summary |
| `finish` | stream 结束 |
### 5.3 Doom Loop 检测(processor.ts lines 351-377
连续 3 次**完全相同的 tool call**(相同名称 + 相同输入)触发权限询问:
```typescript
const recentParts = parts.slice(-DOOM_LOOP_THRESHOLD) // DOOM_LOOP_THRESHOLD = 3
if (recentParts.length === DOOM_LOOP_THRESHOLD &&
recentParts.every(part =>
part.type === "tool" &&
part.tool === value.name &&
part.state.status !== "pending" &&
JSON.stringify(part.state.input) === JSON.stringify(input)
)) {
yield* permission.ask({ permission: "doom_loop", ... })
}
```
### 5.4 Compaction 工作流
两种触发方式:
| 触发条件 | 行为 |
|---------|------|
| step-finish 检测到 `isOverflow()` + `auto: true` | 创建 compaction task → 下一轮循环执行 → 压缩后 continue |
| step-finish 检测到 `isOverflow()` + `auto: false` | 标记 `assistantMessage.error` → idle 等待用户干预 |
Compaction 使用专门的 `compaction` agenthidden, mode=primary, `*=deny`)执行。
压缩后的消息标记 `compacted: true`,后续通过 `MessageV2.filterCompactedEffect()` 过滤。
### 5.5 重试机制(processor.ts lines 658-672
```typescript
Effect.retry(
SessionRetry.policy({
provider: input.model.providerID,
parse, // 错误解析(区分可重试/不可重试)
set: (info) => status.set(sessionID, { type: "retry", ... }),
}),
)
```
遇 provider 错误自动重试,LLM stream 完成后 `Effect.ensuring(cleanup)` 保证资源释放。
---
## 六、六种内置 Agent
| 名称 | Mode | Hidden | 用途 | 核心权限特征 |
|------|------|--------|------|-------------|
| `build` | primary | 否 | 默认 agent,全部工具 | question/plan_enter=allow |
| `plan` | primary | 否 | 计划模式,禁用编辑 | edit=deny(除 plans), task(general)=deny |
| `general` | subagent | 否 | 通用子代理 | todowrite=deny(默认禁止改 todo |
| `explore` | subagent | 否 | 只读代码探索 | `*=deny`,仅 read/grep/glob/bash/webfetch/websearch |
| `compaction` | primary | 是 | 会话压缩(自动) | `*=deny` |
| `title` | primary | 是 | 生成会话标题 | `*=deny`step=1 时异步 fork |
| `summary` | primary | 是 | 生成消息摘要 | `*=deny`(每个 step-finish 时异步 fork |
用户可通过 `config.agent` 自定义 agent(支持 `mode: "all"`),也可通过 `agent.generate` 让 LLM 辅助生成。
---
## 七、权限模型总结
```
父 session permission
├── 仅继承 deny 规则 + external_directory 规则 ← subagent-permissions.ts
│ (父 agent 的 allow 规则不传播到子代理)
├── 子代理自身 permission(来自 agent 定义)
├── 默认 deny
│ - todowrite(除非子代理明确允许)
│ - task(除非子代理明确允许,默认防嵌套)
└── 主 agent 专有工具 deny(来自 config.experimental.primary_tools
```
子代理的 session 权限 = **父 deny + 父 external_directory + 自身 permission - 默认 deny - primary_tools deny**
---
## 八、关键设计决策
| 决策 | 意图 | 效果/局限 |
|------|------|----------|
| 结果以 XML 纯文本嵌入上下文 | 简单、LLM 可直接理解 | LLM 自行解析 XML;大结果可能被截断 |
| 无专用 merge 逻辑 | 简洁,不引入额外抽象 | 依赖 LLM 的理解能力处理返回结果 |
| 默认禁止子代理嵌套 task | 防止无限递归 | 限制了多级分解场景 |
| 同一消息多 tool call 并发 | 利用 LLM 并行能力 | 子 agent 隔离,无法协作 |
| Effect.ts 贯穿全程 | 类型安全、结构化并发 | 学习曲线陡峭 |
| session 作为隔离边界 | 天然权限/消息隔离 | 每个子 session 独立数据库记录,开销较大 |
| 后台引擎纯内存 | 有意识取舍(注释说明) | 进程重启丢失状态 |
---
## 九、参考源码路径
| 文件 | 角色 |
|------|------|
| `packages/opencode/src/tool/task.ts` | Task tool 核心实现(调度入口) |
| `packages/opencode/src/tool/task.txt` | Task tool 的 LLM 使用说明 |
| `packages/opencode/src/agent/agent.ts` | Agent 定义注册中心 |
| `packages/opencode/src/agent/subagent-permissions.ts` | 子代理权限推导 |
| `packages/opencode/src/tool/registry.ts` | 工具注册 + `describeTask()` 列出可用子代理 |
| `packages/opencode/src/session/prompt.ts` | 会话循环 + `handleSubtask()` + 提示词构建 |
| `packages/opencode/src/session/processor.ts` | LLM stream 事件处理器 |
| `packages/opencode/src/session/tools.ts` | Tool ↔ AI SDK 桥接 |
| `packages/opencode/src/session/system.ts` | 系统提示词生成(含 Task tool 说明) |
| `packages/opencode/src/background/job.ts` | 后台作业包装层 |
| `packages/core/src/background-job.ts` | 后台作业核心引擎(内存注册表) |
+375 -63
View File
@@ -1,13 +1,13 @@
# AG Core Roadmap
> 定稿日期:2026-05-11
> 最后更新:2026-07-04v0.1 发布完成)
> 最后更新:2026-07-07
## 愿景
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
**当前状态**Phase 0-4c 全部完成Provider IR 重构(统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen Provider)已完成;LlmCycle 简化(IR 消息类型切换 + 桥接层移除)已完成;v0.1 发布就绪(**182 个测试通过、0 clippy 警告、7 个离线示例可运行**)。
**当前状态**v0.1.0 已发布(2026-07-04)。Phase 0-10 全部完成v0.2.0-rc.1 已打标签。Provider IR 重构 + LlmCycle 简化 + 11 个离线示例(含 `quick_start` 30 行最小示例、`end_to_end` 完整集成示例、`context_slot_demo` 分支对话示例)+ SqliteStore 持久化 + 14 个公开枚举 `#[non_exhaustive]` 护栏 + `StepStatus` IR 迁移 + `submit_turn_stream` 流式体验 + ContextSlot 多上下文分区管理已交付。下一步进入 Phase 11(测试与检索补强)。
---
@@ -240,95 +240,400 @@ graph BT
---
## 扩展计划(v0.2+
## v0.2.0 — 生产就绪(Production-Ready Core
> 以下功能在已完成的 phase 中已实现基础能力或在 Phase 4 阶段明确了边界,后续可按维度增量扩展
> 设计参考:见 `docs/note-agent-harness-references.md`OpenClaw / Hermes / OpenHuman / OpenHarness 横向对比)。
> OpenCode 借鉴:见 `docs/note-opencode-agent-switching.md`Agent 切换 + System Prompt 拼接机制)。
**目标**:解决 Rust Agent 工具箱从"能跑"到"能被人依赖"的鸿沟。持久化、配置层、上下文管理三大块补齐后,开发者可在 30 分钟内写出生产可用的 Agent 服务
### 已有扩展项(沿用)
**总体规模**8 个增量 PhasePhase 5-12),17 个可验证 Step。
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| Prompt Optimizer | `prompt` | 提示词自动优化 | P3 | 待实现 |
| 流式接口优化 | `llm/stream` | 流式响应解析与事件化 | P0 | ✅ 已完成基础实现 |
### 功能清单
### v0.2+ 新增扩展项
#### P0 — 必须交付
> 以下为基于 Phase 4 设计讨论确定的 v0.2+ 候选扩展方向,按维度分组。
> 标注为"v0.2 待评估"表示在 Phase 4 完成后再决定是否启动。
| # | 功能 | 模块 | 方案要点 |
|---|------|------|---------|
| 1 | SqliteStore | `memory` | `rusqlite` + `bundled` feature`MemoryStore` 的 SQLite 实现,进程重启数据不丢 |
| 2 | ProviderConfig 扩展 + `from_env()` | `llm` | 补全 `timeout_secs` / `max_retries` 字段;`AG_LLM_*` 环境变量辅助函数 |
| 3 | ToolDefinition IR 正式化 | `tools` | 移除 deprecated OpenAI wire 格式,替换为自定义 `ToolDef` 结构体 |
| 4 | API 稳定性管理 | `*` | 公开枚举加 `#[non_exhaustive]`CHANGELOG 记录 Breaking Changes;废弃 API 用 `#[deprecated]` 标记 |
| 5 | Quick Start + 端到端示例 | `examples/` | 30 行 `main.rs` 快速开始;一个"SQLite 持久化 + Provider + 工具调用 + 多轮对话"的可运行示例(`cargo run --example` |
#### Multi-Agent / 协同
#### P1 — 重要但不阻塞
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| Multi-Agent 协同(Swarm | `agent` | 子 Agent 委派、并行子任务、结果聚合 | P2 | v0.2 待评估 |
| # | 功能 | 模块 | 方案要点 |
|---|------|------|---------|
| 6 | Ollama Provider | `llm/provider` | OpenAI Compat,本地 LLM 支持,实现量极小 |
| 7 | VectorRetriever trait | `memory` | 语义检索 trait 抽象(`index` / `search`),不绑定后端实现 |
| 8 | 流式 `submit_turn_stream` | `agent` | `AgentSession` 新增 `submit_turn_stream()`,返回 `Stream<Item = StreamEvent>` |
| 9 | 测试补强 | `*` | wiremock Provider roundtrip 测试;多线程并发写入 MemoryStore 测试 |
#### 技能(Skills
#### P2 — 有时间再做
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| Markdown 技能按需加载 | `agent` / `prompt` | 兼容 `SKILL.md` 格式(Hermes / OpenHarness 风格),按 prompt 上下文动态加载 | P2 | v0.2 待评估 |
| # | 功能 | 模块 | 备注 |
|---|------|------|------|
| 10 | MCP StreamableHttp | `tools` | 当前仅预留枚举变体 |
| 11 | Gemini Provider | `llm/provider` | 协议差异大,实现成本较高 |
| 12 | 文件系统 MemoryStore 后端 | `memory` | JSON/JSONL 轻量持久化 |
#### 记忆(Memory
### ContextSlot 上下文管理
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| 多通道检索(hybrid | `memory/retriever` | 在 TextOverlap 之上叠加向量检索通道 | P2 | v0.2 待评估 |
| KnowledgeGraph 深度记忆 | `memory` | 实体-关系图、`note-knowledge-graph-design.md` 已记录设计 | P3 | v0.2 待评估 |
| TokenJuice 智能压缩 | `memory` / `llm/compact` | 借鉴 OpenHuman TokenJuice,对工具结果做语义压缩而非字节截断 | P3 | v0.2 待评估 |
**模块归属**`src/llm/context.rs`(与 `compact.rs` 同级)
#### 交互层(TUI / Gateway
**核心概念**`ContextSlot` 是一段带策略配置的消息列表,以 `slot_id` 为 namespace 独立持久化到 `MemoryStore`。支持三种模式、三种来源和派生关联(记录 `parent_id`)。
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| TUI / 多平台 Gateway | 应用层 | OpenClaw / Hermes 风格的消息平台桥接(Feishu / Telegram / Discord 等) | P3 | v0.2+ 应用层 |
**核心类型**
#### 训练基础设施
```rust
pub struct ContextSlot { id, session_id, config, messages, store }
pub struct SlotConfig { mode: SlotMode, source: SlotSource, budget, compact }
pub enum SlotMode {
Full, // 完整对话历史
Focused(FocusedConfig), // 聚焦:保持 LLM 注意力
Readonly, // 只读参考上下文
}
pub struct FocusedConfig { keep_system, recent_turns, inject_summary }
pub enum SlotSource {
New, // 全新空槽,独立持久化
Derived { parent_id, strategy: DeriveStrategy }, // 从父 slot 派生
Static(Vec<Message>), // 预置消息,不持久化
}
pub enum DeriveStrategy { Full, Focused(FocusedConfig) }
pub struct ContextBudget { system, history, tools, tool_results, reserve }
```
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| RL 轨迹导出 | `agent` | ShareGPT 格式轨迹、Atropos 集成(Hermes 风格) | P3 | v0.3+ 探索 |
**持久化 Key 命名**
- `slot_msg:{session_id}:{slot_id}:{index}` → 消息内容
- `slot_meta:{session_id}:{slot_id}``SlotMeta`(含 `parent_id`
- `slot_rel:{session_id}:{child_id}:parent``"{parent_id}"`
#### 安全治理
**`AgentSession` 扩展**
- `create_slot(id, config)` — 创建新 slot
- `switch_slot(id)` — 切换当前 slot
- `list_slots()` — 列出所有 slot
- `derive_slot(id, parent_id, strategy)` — 从父 slot 派生
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| Human-in-the-loop 审批 | `agent` / `tools/permission` | 高危工具执行前的异步审批回调(OpenHarness `permission_prompt` 模式) | P2 | v0.2 待评估 |
**与 `ConversationMemory` 的关系**:保留不废除。`ConversationMemory` 继续服务传统对话场景。
#### 流式 / 实时
**v0.2 不做**
-`slot.fork()` / `merge()` — 分支方法推迟到 v0.3+
-`inject_summary` 自动生成 — v0.2 仅消费端(从 `SessionMemory` 读取),生成在 v0.3+
- ❌ 血缘关系图遍历 — 只存 `parent_id`,不做查询
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| 流式 `submit_turn` | `agent/session` | Phase 4 v1 只暴露非流式 `submit_turn()`v0.2 包装 `LlmCycle::submit_stream` 暴露流式入口 | P2 | v0.2 待评估 |
**依赖**Phase 0MemoryStore trait)、Phase 3MemoryStore 持久化)
**优先级**P1
#### Agent 切换 / Prompt 动态(OpenCode 借鉴)
---
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|-------|---------|------|--------|------|
| Agent 身份切换(角色轮换) | `agent` | 借鉴 OpenCode Tab 键切换 build/plan:同一 `AgentSession` 持有可热替换的 `Agent` 引用,切换时不重置消息历史,在末尾追加 `synthetic: true` 的状态变更消息。详见 `docs/note-opencode-agent-switching.md` §4 | P2 | v0.2 待评估 |
| System Prompt 多层动态拼接 | `agent/session` | 借鉴 OpenCode `request.ts:58-66`:拆分 `base_prompt + agent_prompt + env_context` 三层,`AgentSession::submit_turn` 每轮重算(不缓存),便于按 agent 类型动态切换 | P2 | v0.2 待评估 |
| **多 Context 切换** | `agent` | **Phase 4c 的 SessionMemory 数据结构已预留信息桥接通道,v0.2+ 在其上包装 `ContextManager` 实现完整的多 context 切换:创建/销毁/切换 context、通过 SessionMemory 桥接关键信息。详见 `docs/note-context-switch-design.md`** | P2 | v0.2 待评估 |
### v0.2.0 实施计划 — 8 个增量 Phase
> **编号说明**Phase 5-12 接续 v0.1 的 Phase 0-4c,按开发顺序排列。
#### Phase 5: 热身准备(Warmup
**目标**:快速交付三个互不依赖的独立改动,建立交付节奏。
| Step | 内容 | 文件范围 | 验证标准 |
|------|------|---------|---------|
| **5.1** ✅ | `ProviderConfig` 扩展:补 `timeout_secs`(def=30) + `max_retries`(def=3);新增 `ProviderConfig::from_env(prefix)` | `llm/provider.rs` + 各 Provider `new()` 构造函数 | `cargo test` + `from_env()` 单元测试 |
| **5.2** ✅ | `OllamaProvider`:基于 `GenericOpenaiProvider` 包装,改 base_url 为 `http://localhost:11434``ProviderType` 新增 `Ollama` | `llm/provider/provider.rs` + `llm/provider/ollama.rs`(新增) | `cargo build` — 纯类型级验证 |
| **5.3** ✅ | 公开枚举 `#[non_exhaustive]` 前置标记:`ProviderType` / `StopReason` / `FinishReason` / `EvictionPolicy` / `SlotMode`(预置) | 各枚举定义处 | 编译通过 + `cargo clippy` 0 警告 |
**实际新增**2026-07-05 commit `98dfe6c`):
- 新增文件 1 个(`llm/provider/ollama.rs`72 行)
- 修改文件 2 个(`llm/provider.rs``from_env` + `Default` + 4 个字段;`memory/store.rs` EvictionPolicy 加 `#[non_exhaustive]`
- `ProviderType::Ollama` 变体 + `FromStr` 解析("ollama" → Ollama
- `OllamaProvider::new(base_url, api_key, model, timeout_secs)` + `with_client()` 构造函数
- `ProviderConfig::from_env(prefix)` 解析 `{prefix}_API_KEY` / `{prefix}_BASE_URL` / `{prefix}_MODEL` 环境变量
- 全量测试 182 → 190+8phase 5 新增 from_env 与 Ollama 相关单测)
- clippy 0 警告
**依赖**:无(三个 Step 互不冲突)
**优先级**P05.1+ P15.2+ P0 前置(5.3
**为何独立成 Phase**:三个改动零文件重叠,可以并行推进。它们是后续所有 Phase 的"门把手"——先做完热身再进入核心工作。
**状态**:✅ Phase 5 全部交付物已完成
---
#### Phase 6: ToolDefinition IR 正式化
**目标**:引入 `ToolDef` 新类型,替换已标记 `#[deprecated]``ToolDefinition``OpenaiToolDefinition` 别名)。
**这是 v0.2 技术风险最高的 Phase**,影响 4 个模块约 8 个文件。通过 5 个 Step 逐文件切割确保每步可编译。
| Step | 内容 | 验证标准 |
|------|------|---------|
| **6.1** ✅ | `types/tool.rs` 新增 `ToolDef` 结构体 + `From<ToolDef> for OpenaiToolDefinition` + 反向 `From` | 单元测试 roundtrip |
| **6.2** ✅ | `types/mod.rs` 切别名 `pub type ToolDefinition = ToolDef``MessageRequest.tools``Vec<ToolDef>` | `cargo build` 编译断点 |
| **6.3** ✅ | `cycle.rs` 4 个方法签名 + `registry.rs` `definitions()` 签名更新 | `cargo build` |
| **6.4** ✅ | Provider 适配层(openai.rs / anthropic.rs / openai_compat.rs):`build_request()` 内做 `ToolDef → wire-format` 转换 | `cargo test` 每个 provider 测试 |
| **6.5** ✅ | 所有测试/示例中 `ToolDefinition``ToolDef` 修复;移除旧 `#[deprecated]` alias | `cargo test --all-targets` 全绿 |
**边界切割技巧**
- Step 6.1 → 6.2 之间是安全 checkpoint:新类型存在但旧代码照常编译
- Provider 层不改序列化逻辑,只加一层 `From` 转换
- 当前代码中 `ToolDefinition` 已是 `#[deprecated(since = "0.1.0")]`,用户已有迁移预期
**依赖**:无(仅与 Phase 5.3 有枚举兼容关系)
**优先级**P0
**实际新增**2026-07-05 commit `4cf5918` / `9da9b83` / `b187519`,详见 `docs/13-phase6-tooldef-ir.md`):
- 修改文件 8 个:`llm/types/tool.rs``llm/types/mod.rs``llm/types/request_v2.rs``llm/cycle.rs``llm/provider/openai.rs``tools/registry.rs``tools/mcp.rs``agent/agent.rs`
- `ToolDef` IRname / description / parameters,无 `strict`)新增于 `types/tool.rs`,配套双向 `From` 转换
- `OpenaiToolDefinition` 降级为 `#[doc(hidden)]`,仅供 OpenAI 适配层内部消费
- `MessageRequest.tools` 切换为 `Vec<ToolDef>`
- `ToolDefinition` 别名最终完全移除(直接使用 `ToolDef`
- 4 处 `#[allow(deprecated)]` 抑制点全部清理(cycle/registry/mcp/agent);残留 `#[allow(deprecated)]` 均与 `ChatResponse` / `with_system_prompt` 等其他弃用项无关
- 新增 roundtrip 测试 `message_request_with_tools_roundtrip`(断言 `strict` 字段不泄漏到序列化输出)
- Anthropic 适配层字段名一致零改动;openai_compat/ollama 委托 `GenericOpenaiProvider` 零改动
- 全量测试 190 → 191+1Phase 6 新增 roundtrip);clippy 0 警告
**状态**:✅ Phase 6 全部交付物已完成
---
#### Phase 7: SqliteStore 持久化
**目标**:实现 `MemoryStore` 的 SQLite 后端,进程重启数据不丢。
**与 Phase 6 无耦合,可重叠开发。**
| Step | 内容 | 文件 | 验证标准 |
|------|------|-----|---------|
| **7.1** ✅ | 新增 `memory/store/sqlite.rs``Mutex<Connection>` + `spawn_blocking`,实现 `save/get/delete/list` + prefix 过滤 | `memory/store/sqlite.rs` + `Cargo.toml`add `rusqlite` | 单元测试 CRUD + prefix 查询 |
| **7.2** ✅ | WAL 模式 + 并发安全 + 集成测试(`tokio::spawn` 10 个并发 task | `sqlite.rs` 扩展 | 并发写入 100 轮无 race |
**设计决策**
-`Mutex<Connection>` 而非连接池(ponytail:一个连接够用就不加 r2d2)
- WAL 模式:`PRAGMA journal_mode=WAL` 解决读写锁
**依赖**`MemoryStore` traitv0.1 Phase 3 已就绪)
**优先级**P0
**实际新增**2026-07-05 commit `13edacd` / `c8a91f6` / `c82af60`,详见 `docs/14-phase7-sqlite-store.md`):
- 方案文档:`docs/14-phase7-sqlite-store.md`526 行,Phase 7 设计推演与权衡记录)
- 结构重组:`src/memory/store.rs` 单体文件 → `src/memory/store/{mod.rs(in_memory.rs, sqlite_store.rs)}` 模块目录;外部导入路径 `crate::memory::store::MemoryStore` 不变
- 新增文件 2 个:`src/memory/store/sqlite_store.rs`545 行 SqliteStore 实现 + 9 个内联测试)、`src/memory/store/in_memory.rs`266 行,结构搬移)
- 核心实现要点:
- `Arc<Mutex<Connection>>` 串行化所有 IO`spawn_blocking` 卸载到阻塞线程池
- WAL 模式 + `synchronous=NORMAL` + `busy_timeout=5s` + `wal_autocheckpoint=1000`
- `PRAGMA user_version` schema 版本管理(`INITIAL_USER_VERSION = 1`
- `created_at` 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
- 错误精细映射:`SqliteFailure` / `InvalidQuery``InvalidInput``FromSqlConversionFailure``Serialization`;其他 → `Storage`
- 9 个内联测试覆盖:CRUD、upsert、prefix / since / offset+limit 过滤、10 写者 × 10 次并发写入、持久化 round-trip(重启连接不丢数据)、`InMemoryStore ↔ SqliteStore` trait-box 互换兼容性
- 依赖:`rusqlite = { version = "0.32", features = ["bundled"] }``time` 增补 `parsing` / `formatting` / `macros` features`dev-dependencies` 新增 `tempfile = "3"`
- 全量测试 191 → 200+9Phase 7 新增 SqliteStore 单测);clippy 0 警告
**状态**:✅ Phase 7 全部交付物已完成
---
#### Phase 8: MVP 集成出口(v0.2.0-rc.1 候选)
**目标**:P0 五项全部交付。开发者 clone 仓库后 10 分钟跑起持久化 Agent。
| Step | 内容 | 验证标准 |
|------|------|---------|
| **8.1** ✅ | API 稳定性扫尾:`#[non_exhaustive]` × 14 公开枚举 + `StepStatus::Completed``MessageResponse` + CHANGELOG v0.2.0-rc.1 + Cargo.toml 0.2.0-rc.1 | `cargo doc --no-deps` 0 warning + 零 deprecated warning |
| **8.2** ✅ | Quick Start 示例(57 行 `main.rs`):MockProvider + EchoTool + submit_turn 真实工具调用 | `cargo run --example quick_start` exit 0 |
| **8.3** ✅ | 端到端示例:SqliteStore + AG_LLM_* from_env 自动检测 + 3 工具 + 3 轮对话 + 持久化跨连接验证 | `cargo run --example end_to_end`Mock fallback,无需 API key|
**Phase 8 全部完成**。**已打 `v0.2.0-rc.1` 标签**。
**实际新增**2026-07-057 commits):
- `feat(core)` —— 14 个公开枚举追加 `#[non_exhaustive]`P0 核心 IR + P0 Error + P1 其他)
- `refactor(agent)` —— `StepStatus::Completed(ChatResponse)``Completed(MessageResponse)` + `task_agent_demo.rs` 清理 3 处废弃类型
- `docs` —— CHANGELOG v0.2.0-rc.1 条目 + Cargo.toml version 0.1.0 → 0.2.0-rc.1 + README 示例列表 7 → 10
- `test(core)` —— 验证 commit 1-3 零回归(test 200 passed + clippy 0 警告 + doc 0 warning
- `feat(examples)` —— `quick_start.rs`60 行)+ `end_to_end.rs`246 行)
- `docs(roadmap)` —— 标记 Phase 8 全部完成 + M4 里程碑 ✅
- `fix(examples)` —— 实施后 PM/SA/Code Reviewer 三方审查发现 6 项问题(🔴 CalcTool 除零 panic + 🟡 drop 注释准确性 + 🟡 EchoTool 错误处理 + 💭 断言一致性 + 💭 工具两端语义统一 + 💭 trailing newline),全部修复
**依赖**Phase 5ProviderConfig from_env+ Phase 6ToolDef+ Phase 7SqliteStore
**优先级**P0
**状态**:✅ Phase 8 全部交付物已完成
---
#### Phase 9: 流式体验增强
**目标**:Agent 会话支持流式输出,开发者看到实时 token。
| Step | 内容 | 文件 | 验证标准 |
|------|------|-----|---------|
| **9.1** ✅ | `AgentSession::submit_turn_stream(user_input) -> impl Stream<Item=StreamEvent>` | `agent/session.rs` | 单元测试验证流事件序列:`TextDelta → ... → MessageComplete` |
**注意**tool 自动循环时流中插入 `ToolExecutionStarted` 事件,用户端 UI 显示"正在调用工具..."。
**依赖**Phase 6ToolDef+ `LlmProvider.chat_stream`v0.1 已有)
**优先级**P1
**实际新增**2026-07-06 commit `212cfcc`,详见 `docs/16-phase9-streaming-experience.md`):
- 方案文档:`docs/16-phase9-streaming-experience.md`(821 行,含状态机设计推演与边界情况)
- 修改文件 3 个:`src/agent/session.rs`+208,含 `submit_turn_stream` / `finalize_turn`)、`src/llm/cycle.rs`+784,含 `submit_with_tools_stream` / `run_tool_loop` spawn + mpsc 状态机)、`src/llm/types/response_v2.rs`+21,含 `StreamEvent::ToolExecutionStarted`/`Completed` 变体 + `apply_to` 元事件)
- 关键设计:`CycleConfig``Clone` derive 以支持 spawn 跨 task`finalize_turn` 手动同步状态(`submit_turn_stream` 返回流前不落库,避免半成品被 hook 误读)
- 测试:新增 9 个单元测试 + 2 个集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock provider 流消费 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211(+11,0 失败)
- clippy 0 警告
- 无新增外部依赖
**状态**:✅ Phase 9 全部交付物已完成
---
#### Phase 10: ContextSlot 上下文管理
**目标**:支持多上下文分区管理,Agent 可在不同 slot 之间切换。
| Step | 内容 | 验证标准 |
|------|------|---------|
| **10.1** ✅ | `src/agent/context.rs``ContextSlot` + `SlotConfig` / `SlotMode` / `FocusedConfig` / `SlotSource` / `DeriveStrategy` / `ContextBudget` / `SlotMeta` 核心类型 | `cargo build --all-targets` |
| **10.2** ✅ | ContextSlot 持久化:基于 `MemoryStore` trait(不绑定 SqliteStore)实现 save/load/list/delete + slot 命名空间 key 策略 + `load_messages()` Focused 读时过滤 + `append_messages()` Readonly 阻断 + colon 注入防护 | 单元测试:持久化 roundtrip / session 隔离 / Focused 边界 / delete 保护 / 派生 / load_messages() |
| **10.3** ✅ | `AgentSession` 扩展:`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot` + `new()` 自动创建 `"default"` slot + `submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回 + 新示例 `context_slot_demo` | 集成测试 + `cargo run --example context_slot_demo` exit 0 |
**如何保证简单场景无感**`AgentSession::new()` 内部检查,自动创建 `"default"` slot → `submit_turn` 默认写到 default slot。
**实际新增**2026-07-07 commit `6359422`,详见 `docs/17-phase10-contextslot.md`):
- 方案文档:`docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复)
- 新增文件 3 个:`src/agent/context.rs`~430 行 ContextSlot 核心类型 + 持久化方法 + 22 个测试)、`src/agent/context.rs` 中的 `ContextSlot::filter_focused` 静态方法(被 `load_messages``derive_slot` 复用,消除代码重复)、`examples/context_slot_demo.rs`(~160 行分支对话示例:法律咨询 → 派生两个方向 → 切换 → 隔离验证 → 删除保护)
- 修改文件 3 个:`src/agent.rs`+5 行 module 声明 + re-export)、`src/agent/error.rs`+56 行:3 个新变体 `SlotReadonly`/`SlotNotFound`/`SlotAlreadyExists` + 4 个测试)、`src/agent/session.rs`+825/-197 行:slots 字段 + 6 个管理方法 + submit_turn/finalize_turn 改造 + 17 个测试)
- 关键设计:
- **模块归属**`agent/context.rs`(零新依赖方向,遵循 `agent → memory` 已有依赖)
- **持久化**JSON blob 批次存储,每 slot 3-4 条 `MemoryItem``slot_data` / `slot_meta` / `slot_config` / `slot_rel`
- **submit_turn 签名不变**:方案 A(内部 `current_slot_id` 状态),向后兼容
- **Focused 模式读时过滤**`load_messages() -> Vec<Message>`,避免 Rust 借用检查问题
- **增量追加写回**`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式数据不丢失
- **delete_slot 双重保护**:禁止删 `"default"` + 至少保留一个 slot
- **colon 注入防护**`assert_no_colon` 在 key 构造时 panic
- **错误传播**`serde_json` / `MemoryStore` 所有错误用 `?` 传播,无静默吞掉
- 验证:211 → 254 测试(+43 新测试),clippy 0 警告,doc 0 warning10 + 1 示例全部 exit 0
- finalize_turn 签名变更(破坏性):新增 `new_messages_from_cycle: Vec<Message>` 参数,返回从 `()` 改为 `Result<(), AgentError>`——影响 Phase 9 的 `submit_turn_stream_triggers_turn_hooks``submit_turn_stream_end_to_end` 2 个测试,已适配
**依赖**Phase 5`#[non_exhaustive]` 预置 SlotMode 等枚举)、Phase 7SqliteStore 推荐持久化后端;`MemoryStore` trait 即可)
**优先级**P1
**状态**:✅ Phase 10 全部交付物已完成
---
#### Phase 11: 测试与检索补强
**目标**:补全测试覆盖 + 语义检索抽象。
| Step | 内容 | 验证标准 |
|------|------|---------|
| **11.1** | `VectorRetriever` trait`index(id, embeddings)` + `search(query, k)` | 编译 + mock 测试 |
| **11.2** | wiremock Provider roundtrip 测试:模拟 OpenAI/Anthropic HTTP 端点 | `cargo test` 新增 10+ roundtrip 测试 |
| **11.3** | 并发测试补强:InMemoryStore + SqliteStore 多线程写入验证 | 跑 100 轮无 race |
**依赖**:无(可随时做)
**优先级**P1
---
#### Phase 12: P2 锦上添花(可选)
**目标**:时间允许时按优先级交付。
| 优先级 | 功能 | 实现量估计 | 备注 |
|--------|------|-----------|------|
| **12.1** | 文件系统 MemoryStoreJSON/JSONL | ~80 行 | 最简单,适合练手 |
| **12.2** | MCP StreamableHttp 传输 | ~150 行 | 协议还在演进 |
| **12.3** | Gemini Provider | ~300 行 | 协议差异大,建议推迟到 v0.3 |
**依赖**:无(独立交付)
---
### v0.2.0 Phase 依赖关系图
```mermaid
graph BT
P5["<b>Phase 5: 热身准备</b><br/>ProviderConfig::from_env<br/>Ollama Provider<br/>#[non_exhaustive] 标记"]:::done
P6["<b>Phase 6: ToolDef IR</b><br/>Provider 无关工具定义"]:::done
P7["<b>Phase 7: SqliteStore</b><br/>rusqlite + WAL<br/>9 个内联测试<br/>持久化 round-trip"]:::done
P8["<b>Phase 8: MVP 出口</b><br/>rc.1 标签<br/>14 枚举 #[non_exhaustive]<br/>StepStatus IR 迁移<br/>quick_start + end_to_end"]:::done
P9["<b>Phase 9: 流式体验增强</b><br/>submit_turn_stream<br/>submit_with_tools_stream<br/>9 单元测试 + 2 集成测试"]:::done
P10["<b>Phase 10: ContextSlot</b><br/>ContextSlot 类型<br/>JSON blob 持久化<br/>AgentSession 集成<br/>43 个新测试"]:::done
P11["Phase 11<br/>测试与检索"]:::p1
P12["Phase 12<br/>P2 锦上添花"]:::p2
P8 --> P5
P8 --> P6
P8 --> P7
P9 --> P6
P10 --> P7
P10 --> P8
P11 -.-> P7
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
classDef warmup fill:#e2e8f0,stroke:#94a3b8
classDef core fill:#fbbf24,stroke:#d97706
classDef mvp fill:#4ade80,stroke:#16a34a
classDef p1 fill:#93c5fd,stroke:#2563eb
classDef p2 fill:#c4b5fd,stroke:#7c3aed
```
---
### 关键里程碑
| 里程碑 | Phase 完成条件 | 可验证指标 | 状态 |
|--------|---------------|-----------|------|
| **M1** | Phase 5 | 热身三项完成:`from_env()` 可用 / Ollama 类型存在 / `#[non_exhaustive]` 就位 | ✅ 2026-07-05 |
| **M2** | Phase 6 | `ToolDef` 全量切换,`cargo test --all-targets` 全绿 | ✅ 2026-07-05 |
| **M3** | Phase 7 | SqliteStore CRUD + 并发测试通过,进程重启数据不丢 | ✅ 2026-07-05 |
| **M4** | **Phase 8 (rc.1)** | P0 五项全部交付,`cargo run --example quick_start` 跑通 | ✅ 2026-07-05 |
| **M5** | Phase 9 | `submit_turn_stream` 流式事件序列验证通过 | ✅ 2026-07-06 |
| **M6** | Phase 10 | ContextSlot 创建/切换/派生集成测试通过 | ✅ 2026-07-07 |
| **M7** | Phase 11 | wiremock + 并发测试补强,测试总量 200+ | ⏳ |
| **M8** | Phase 12(可选) | P2 功能按需交付 | ⏳ |
---
## v0.3+ 展望
### 已规划的功能
| 功能 | 说明 | 预计版本 |
|------|------|---------|
| ContextSlot 分支(fork/merge | 在决策点 fork 出子上下文,分支独立演进,可合并/丢弃 | v0.3 |
| 摘要自动生成 | Hook 驱动,`OnTurnEnd` 自动将对话摘要写入 `SessionMemory``inject_summary` 消费端已在 v0.2 就绪 | v0.3 |
| 知识图谱 | 实体-关系图,`docs/note-knowledge-graph-design.md` 已记录设计 | v0.3+ |
| Multi-Agent 协同(Swarm | 子 Agent 委派、并行子任务、结果聚合 | v0.4+ |
| 精确 tokenizer 计数 | 绑定具体模型的 tokenizer 计数,替代当前的字符估算 | v0.3+ |
| 血缘关系图遍历 | 以 `parent_id` 为基础,提供 slot 血缘链查询 | v0.3+ |
| Markdown 技能按需加载 | 兼容 `SKILL.md` 格式,按 prompt 上下文动态加载 | v0.3+ |
| TokenJuice 语义压缩 | 对工具结果做语义压缩而非字节截断 | v0.3+ |
| Human-in-the-loop 审批 | 高危工具执行前的异步审批回调 | v0.3+ |
| RL 轨迹导出 | ShareGPT 格式轨迹、Atropos 集成 | v0.4+ |
### 明确不做(agcore 范围外)
| 功能 | 原因 |
|------|------|
| TUI / 多平台 Gateway | 应用层职责(Feishu / Telegram / Discord 桥接) |
| 配置自动加载(config/figment) | 配置来源策略应由上游应用决定,agcore 不定义配置格式 |
| 提示词自动优化 | 属于智能层,不应内建于 core 库 |
---
## 风险与建议
1. **Phase 0 已完成**:LLM 调用周期基础设施已全部实现,可以支撑后续模块开发
2. **并行可能性**Phase 0 和 Phase 1 可并行开展(无相互依赖),可加速早期交付
3. **MCP 协议复杂性**MCP 涉及协议握手、session 管理、长期连接,建议预留充足时间调研协议细节
4. **Scope 蔓延风险**当前 specs 只有 1 份文档,建议每个模块上线前都产出对应 spec,避免边实现边设计
5. **Phase 4 抽象化边界**AG Core 定位为"支持库"而非"Agent 产品"Phase 44a/4b/4c)需严格控制范围——只暴露 trait + 最小 reference impl,业务循环(多轮 turn 编排、对话记忆自动回写、Task 拆解策略)留给上层应用。`SessionMemory`(Phase 4c)提供信息桥接通道但不实现 context 切换逻辑。多 context 切换管理延后至 v0.2+。详细设计决策见 `docs/7-agent-runtime.md`
6. **参考项目语言差异**OpenClaw / Hermes / OpenHarness 均为 Python/TypeScript 实现,OpenHuman 虽是 Rust + Tauri 但定位是桌面应用。借鉴时**只取架构模式**,不照搬具体实现(如 Pydantic 工具校验、SQLite Memory Tree、Node+Python 双进程等)
1. **持久化依赖**`rusqlite` + `bundled` 零外部依赖编译,但 SQLite 不适配所有场景(分布式/高并发写)。`MemoryStore` trait 的抽象层允许下游自行实现 Redis / PostgreSQL 后端
2. **ContextSlot 心智负担**`ContextSlot` 引入了一等抽象的复杂度。建议通过 `AgentBuilder` 默认创建 `"default"` slot,让简单场景无感使用
3. **向量检索生态**`VectorRetriever` trait-only 不绑定实现,需社区贡献或用户自行适配 pgvector / qdrant / lancedb
4. **Scope 蔓延**agcore 定位为"支持库"而非"Agent 产品",始终以 trait + reference impl 为边界,业务循环留给上层
5. **API 稳定性**v0.2 引入 `#[non_exhaustive]``#[deprecated]` 机制,但不承诺 SemVer 稳定——仍在快速迭代期
---
## 下一步行动
1. **Phase 4c 已完成**Phase 4a + 4b + 4c 已交付(116 测试通过,0 clippy 警告)。可启动 v0.2+ 扩展评估(如多 Context 切换、Multi-Agent 协同等)
2. **Context 切换备忘**`docs/note-context-switch-design.md` 记录了多 context 切换方案讨论,作为 v0.2+ 扩展项的输
3. **参考项目调研沉淀**:已完成 OpenClaw / Hermes / OpenHuman / OpenHarness 横向调研,结果沉淀至 `docs/note-agent-harness-references.md`,作为 v0.2+ 扩展项的输入
4. **Phase 3 备用设计就绪**`docs/note-knowledge-graph-design.md` 记录了 KnowledgeGraph、高级评分、RecallBased 淘汰等设计,v0.2+ 记忆扩展可直接参考
1. **Phase 11 启动**:测试与检索补强(`VectorRetriever` trait + wiremock Provider roundtrip + 并发写入验证),P1 功能
2. **示例先行**:每完成一个 Phase 立即更新对应示例,验证通过后再合
3. **里程碑追踪**:以 Phase 10ContextSlot,已完成)为最新节点,逐 Phase 验收
4. **v0.2.0 正式版**:Phase 8-11 全部完成后,去掉 rc 后缀打 `v0.2.0` 正式版
**已完成 / 进行中阶段**
- ✅ Phase 0 Foundation — 全部交付物已完成
@@ -338,9 +643,16 @@ graph BT
- ✅ Phase 4a Core Glue — 全部交付物已完成
- ✅ Phase 4b Task Execution — 全部交付物已完成
- ✅ Phase 4c Session Memory — 全部交付物已完成
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen 适配(方案:`docs/10-llm-provider-refinement.md``docs/10a-phase0-types-and-trait.md``docs/10b-phase1-provider-adaptation.md`
-LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除(方案:`docs/10c-phase2-llm-cycle-simplify.md`
-v0.1 Release — 技术债扫清、MockProvider 公开化、7 个离线示例、README + 错误消息友好化、Roadmap 同步、CHANGELOG 初始化(计划:`docs/11-v0.1-release-plan.md`
- ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy
-Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
-Phase 7 SqliteStore — `rusqlite 0.32` + WAL 模式 + `Arc<Mutex<Connection>>` + `spawn_blocking``memory/store.rs``store/{in_memory,sqlite_store}.rs` 模块化;9 个内联测试覆盖 CRUD/upsert/过滤/10×10 并发/持久化 round-trip`InMemoryStore ↔ SqliteStore` trait-box 互换兼容
-**Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)``Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 60 行 + `end_to_end` 246 行),10 个离线示例全部 exit 0**v0.2.0-rc.1 标签已打**;实施后三方审查发现 6 项问题(1 🔴 + 2 🟡 + 3 💭)已全部修复
-**Phase 9 流式体验增强**`AgentSession::submit_turn_stream` 流式事件序列 + `LlmCycle::submit_with_tools_stream` spawn + mpsc 状态机 + `StreamEvent::ToolExecutionStarted`/`Completed` 新变体 + 9 单元测试 + 2 集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock 验证 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211;`CycleConfig``Clone` derive;方案文档 `docs/16-phase9-streaming-experience.md`821 行)
-**Phase 10 ContextSlot 上下文管理**`src/agent/context.rs` 新增 `ContextSlot` 核心类型(Full / Focused / Readonly 三种模式,New / Derived / Static 三种来源)+ JSON blob 批次持久化(每 slot 3-4 条 MemoryItem`slot_config` key 自恢复支持旧版本兼容);`AgentSession` 扩展 slots 字段 + 5 个管理方法(`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot`,自动创建 `"default"` slot`delete_slot` 双重保护禁止删 default/最后一个);`submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回(`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式"读时过滤"语义不丢失数据);`finalize_turn` 签名变更(新增 `new_messages_from_cycle: Vec<Message>` 参数,返回 `Result<(), AgentError>`);`agent/error.rs` 新增 3 个 Slot 错误变体(`SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`);`examples/context_slot_demo.rs` 新增分支对话示例(法律咨询入口 → 两个派生方向 → 切换 → 隔离验证 → 删除保护);方案文档 `docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复);全量 211 → 254(+43 新测试),clippy 0 警告,doc 0 warning11 个离线示例全部 exit 0
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen/Ollama 适配
- ✅ LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、8 个离线示例(含 `simple_visit`)、README + 错误消息友好化、CHANGELOG 初始化
- 📋 **v0.2 规划细化完成** — 8 个增量 PhasePhase 5-12),17 个可验证 Step,覆盖 P0-P2 全部 12 项功能 + ContextSlot
---
+7 -8
View File
@@ -15,9 +15,9 @@ use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider;
use agcore::llm::types::Usage;
use agcore::llm::types::message::{ContentBlock, Message};
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
use agcore::llm::types::Usage;
use agcore::tools::ToolRegistry;
/// 计算器角色 Agent。
@@ -72,7 +72,10 @@ async fn main() {
// 4. 提交第一轮
println!("=== 提交第 1 轮 ===");
let resp = session.submit_turn("1+1=?").await.expect("submit_turn 失败");
let resp = session
.submit_turn("1+1=?")
.await
.expect("submit_turn 失败");
println!("LLM: {}", resp.text());
session
.set_session_data("last_q", "1+1=?")
@@ -107,11 +110,7 @@ async fn main() {
// 8. 跨 session 数据隔离验证
println!("=== 数据隔离验证 ===");
let other = AgentSession::new(
Arc::new(CalculatorAgent),
"other-session",
bundle,
);
let other = AgentSession::new(Arc::new(CalculatorAgent), "other-session", bundle);
assert!(
other.get_session_data("last_q").await.unwrap().is_none(),
"新会话不应看到旧 session 的 last_q"
@@ -127,4 +126,4 @@ async fn main() {
);
println!("\n✓ agent_session_demo 完成");
}
}
+161
View File
@@ -0,0 +1,161 @@
//! context_slot_demo —— 多上下文槽位管理示例。
//!
//! 场景:法律咨询入口 → 派生两个独立探索方向 → 切换 → 隔离验证 → 删除。
//!
//! 展示:
//! - 默认 slot 自动创建
//! - 多 slot 间的消息隔离
//! - 派生 slot 从父 slot 复制消息
//! - 删除非 default slot 后自动回退到 default
//!
//! 运行:`cargo run --example context_slot_demo`(离线,零配置)
use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider;
use agcore::llm::provider::LlmProvider;
use agcore::llm::types::message::{ContentBlock, Message};
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
use agcore::llm::types::Usage;
use agcore::tools::ToolRegistry;
struct LegalAdvisor;
impl Agent for LegalAdvisor {
fn name(&self) -> &str {
"legal-advisor"
}
fn system_prompt(&self) -> Option<&str> {
Some("你是法律顾问。请用一句话回答用户问题。")
}
}
/// 构造一个简单的 Assistant 响应(用于 MockProvider)。
fn assistant_resp(text: &str) -> MessageResponse {
MessageResponse {
id: String::new(),
model: "mock".into(),
message: Message::Assistant {
content: vec![ContentBlock::Text { text: text.into() }],
},
usage: Usage::from_input_output(5, 5),
stop_reason: StopReason::Stop,
extra: Default::default(),
}
}
#[tokio::main]
async fn main() {
// 1. 构造 session(自动包含 default slot
let provider: Arc<dyn LlmProvider> = Arc::new(MockProvider::new(vec![
assistant_resp("您好,我可以帮您处理法律问题。"),
assistant_resp("管辖权问题:建议选择合同签订地法院。"),
assistant_resp("条款修改:建议将上限调整为 80 万。"),
assistant_resp("已回到主对话。"),
]));
let bundle = Arc::new(
AgentBuilder::new()
.provider(provider)
.tool_registry(Arc::new(ToolRegistry::new()))
.hook_executor(Arc::new(HookExecutor::new()))
.build()
.unwrap(),
);
let mut session = AgentSession::new(Arc::new(LegalAdvisor), "legal-001", bundle);
println!("=== 1. 默认 slot 自动创建 ===");
assert_eq!(session.current_slot_id(), "default");
let slots: Vec<_> = session.list_slots().collect();
println!("初始 slots: {slots:?}");
assert_eq!(slots.len(), 1);
assert!(slots.contains(&&"default".to_string()));
println!("\n=== 2. 在 default slot 中提交一轮 ===");
let r1 = session.submit_turn("我需要法律援助").await.unwrap();
println!("default slot response: {}", r1.text());
println!("\n=== 3. 派生两个独立探索方向的 slot ===");
session
.derive_slot("option_jurisdiction", "default", agcore::agent::DeriveStrategy::Full)
.await
.unwrap();
session
.derive_slot("option_amendment", "default", agcore::agent::DeriveStrategy::Full)
.await
.unwrap();
let slots: Vec<_> = session.list_slots().cloned().collect();
println!("派生后 slots: {slots:?}");
assert_eq!(slots.len(), 3);
println!("\n=== 4. 切到 option_jurisdiction 并提交 ===");
session.switch_slot("option_jurisdiction").await.unwrap();
assert_eq!(session.current_slot_id(), "option_jurisdiction");
let r2 = session.submit_turn("如果用户质疑管辖权?").await.unwrap();
println!("option_jurisdiction response: {}", r2.text());
println!("\n=== 5. 切到 option_amendment 并提交 ===");
session.switch_slot("option_amendment").await.unwrap();
let r3 = session.submit_turn("用户要求提高赔偿上限?").await.unwrap();
println!("option_amendment response: {}", r3.text());
println!("\n=== 6. 切回 default,验证消息隔离 ===");
session.switch_slot("default").await.unwrap();
let r4 = session.submit_turn("汇总一下我们的讨论").await.unwrap();
println!("default response: {}", r4.text());
// 验证 default slot 不包含 option_jurisdiction 的"管辖权"问题
let (_, default_slot) = session.slots().find(|(id, _)| *id == "default").unwrap();
let default_has_jurisdiction = default_slot
.messages
.iter()
.any(|m| message_contains(m, "管辖权"));
assert!(
!default_has_jurisdiction,
"default slot 不应包含 option_jurisdiction 的消息"
);
println!("\n=== 7. 删除 option_amendment,验证回退到 default ===");
session.delete_slot("option_amendment").await.unwrap();
let slots: Vec<_> = session.list_slots().cloned().collect();
println!("删除后 slots: {slots:?}");
assert!(!slots.contains(&"option_amendment".to_string()));
assert_eq!(slots.len(), 2);
println!("\n=== 8. 切到 option_jurisdiction 并删除,验证 current 回退 ===");
session.switch_slot("option_jurisdiction").await.unwrap();
session.delete_slot("option_jurisdiction").await.unwrap();
assert_eq!(session.current_slot_id(), "default");
let slots: Vec<_> = session.list_slots().cloned().collect();
println!("删除后 slots: {slots:?}");
assert_eq!(slots.len(), 1);
assert_eq!(slots[0], "default");
println!("\n=== 9. 验证 delete_slot 保护逻辑 ===");
let err = session.delete_slot("default").await.unwrap_err();
println!("删除 default 返回错误: {err}");
assert!(matches!(err, agcore::agent::AgentError::Config(_)));
println!("\n✓ context_slot_demo 完成");
}
/// 检查 Message 是否包含指定文本(提取第一个 Text block)。
fn message_contains(msg: &Message, needle: &str) -> bool {
use agcore::llm::types::message::ContentBlock;
let blocks = match msg {
Message::System { content }
| Message::User { content }
| Message::Assistant { content } => content,
Message::UserImage { .. } => return false,
Message::ToolResult { content, .. } => content,
_ => return false,
};
for block in blocks {
if let ContentBlock::Text { text } = block
&& text.contains(needle)
{
return true;
}
}
false
}
+9 -17
View File
@@ -29,6 +29,7 @@ fn message_text(msg: &Message) -> &str {
.next()
.unwrap_or(""),
Message::UserImage { .. } => "[image]",
_ => "",
}
}
@@ -80,11 +81,8 @@ async fn main() {
// 3. 多角色混合 + clear
println!("\n=== 多角色写入 + clear ===");
let store3 = Arc::new(InMemoryStore::new());
let mut memory3 = ConversationMemory::new(
store3,
"session-3",
ConversationMemoryConfig::default(),
);
let mut memory3 =
ConversationMemory::new(store3, "session-3", ConversationMemoryConfig::default());
memory3
.add_message(Message::user_text("你好"))
.await
@@ -98,7 +96,9 @@ async fn main() {
.await
.unwrap();
memory3
.add_message(Message::assistant("我无法查询实时天气,但你可以查看天气应用。"))
.add_message(Message::assistant(
"我无法查询实时天气,但你可以查看天气应用。",
))
.await
.unwrap();
println!(
@@ -119,16 +119,8 @@ async fn main() {
// 4. Session 隔离
println!("\n=== Session 隔离(共用 InMemoryStore===");
let store4 = Arc::new(InMemoryStore::new());
let mut a = ConversationMemory::new(
store4.clone(),
"s-a",
ConversationMemoryConfig::default(),
);
let mut b = ConversationMemory::new(
store4.clone(),
"s-b",
ConversationMemoryConfig::default(),
);
let mut a = ConversationMemory::new(store4.clone(), "s-a", ConversationMemoryConfig::default());
let mut b = ConversationMemory::new(store4.clone(), "s-b", ConversationMemoryConfig::default());
a.add_message(Message::user_text("A 的消息")).await.unwrap();
b.add_message(Message::user_text("B 的消息")).await.unwrap();
println!(
@@ -140,4 +132,4 @@ async fn main() {
assert_eq!(b.len(), 1);
println!("\n✓ conversation_memory_demo 完成");
}
}
+11 -16
View File
@@ -17,7 +17,7 @@ use agcore::tools::{
ToolRegistry,
};
use async_trait::async_trait;
use serde_json::{json, Value};
use serde_json::{Value, json};
/// 天气查询工具 —— 模拟根据城市返回天气数据。
struct WeatherTool;
@@ -42,11 +42,7 @@ impl BaseTool for WeatherTool {
fn required_permissions(&self) -> Vec<Permission> {
vec![Permission::Network]
}
async fn execute(
&self,
args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
async fn execute(&self, args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
let city = args["city"].as_str().unwrap_or("未知");
// 模拟查询:根据城市名给出不同温度
let (temperature, condition) = match city {
@@ -84,11 +80,7 @@ impl BaseTool for DeleteFileTool {
fn required_permissions(&self) -> Vec<Permission> {
vec![Permission::Delete]
}
async fn execute(
&self,
_args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
Ok(json!({"deleted": true}))
}
}
@@ -138,9 +130,8 @@ async fn main() {
// 5. 权限检查:默认 PermissionConfig 黑名单含 Delete
println!("\n=== 权限检查(默认 PermissionConfigdenied = [Delete, Shell]===");
let mut registry_with_checker = ToolRegistry::new().with_permission_checker(PermissionChecker::new(
PermissionConfig::default(),
));
let mut registry_with_checker = ToolRegistry::new()
.with_permission_checker(PermissionChecker::new(PermissionConfig::default()));
registry_with_checker
.register(Arc::new(WeatherTool) as ToolRef)
.unwrap();
@@ -155,7 +146,11 @@ async fn main() {
.unwrap();
println!(
"get_weather 权限检查: {}",
if r.output.is_ok() { "通过 ✓" } else { "阻断 ✗" }
if r.output.is_ok() {
"通过 ✓"
} else {
"阻断 ✗"
}
);
// delete_file 声明 Delete → 在 denied 列表 → 阻断
@@ -166,4 +161,4 @@ async fn main() {
println!("delete_file 权限检查: 阻断 ✗ ({err})");
println!("\n✓ custom_tool 完成");
}
}
+247
View File
@@ -0,0 +1,247 @@
//! end_to_end —— 3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证。
//!
//! 运行:`cargo run --example end_to_end`(离线,零配置)
//!
//! ## 真实 LLM Provider 切换
//!
//! 设置环境变量即可使用真实 LLM Provider
//! - `AG_LLM_BASE_URL` —— API 端点(如 `https://api.openai.com/v1`
//! - `AG_LLM_API_KEY` —— API key
//! - `AG_LLM_MODEL` —— 模型名(如 `gpt-4o-mini`
//! - `AG_LLM_PROVIDER`(可选)—— Provider 类型,默认 OpenaiChatOpenAI / DeepSeek / Qwen / Ollama
//!
//! 未设置上述变量时自动降级为 MockProvider,零配置可运行。
use std::env;
use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider;
use agcore::llm::provider::{create_provider, LlmProvider, ProviderConfig, ProviderType};
use agcore::llm::types::{Usage, message::{ContentBlock, Message}, response_v2::{MessageResponse, StopReason}};
use agcore::memory::store::{MemoryStore, SqliteStore};
use agcore::memory::types::{MemoryFilter, MemoryItem};
use agcore::tools::{BaseTool, ToolContext, ToolError, ToolRegistry};
use async_trait::async_trait;
use serde_json::{Value, json};
use tempfile::TempDir;
use time::OffsetDateTime;
// === Agent ===
struct AssistantAgent;
impl Agent for AssistantAgent {
fn name(&self) -> &str { "end-to-end assistant" }
fn system_prompt(&self) -> Option<&str> { Some("简洁助手,必要时调用工具完成任务。") }
}
// === Tools ===
struct EchoTool;
#[async_trait]
impl BaseTool for EchoTool {
fn name(&self) -> &str { "echo" }
fn description(&self) -> &str { "回显输入文本" }
fn parameters(&self) -> Value {
json!({"type":"object","properties":{"text":{"type":"string"}},"required":["text"]})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let text = args.get("text").and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidArguments("text".into(), "需要 string 类型的 text 参数".into()))?;
Ok(json!({"echoed": format!("收到: {text}")}))
}
}
/// 四则运算:'a op b' 格式(ponytail: 基础 +-*/ 不引入 rhai 依赖)。
struct CalcTool;
#[async_trait]
impl BaseTool for CalcTool {
fn name(&self) -> &str { "calc" }
fn description(&self) -> &str { "四则运算:'a op b' 格式,op ∈ {+, -, *, /}" }
fn parameters(&self) -> Value {
json!({"type":"object","properties":{"expr":{"type":"string"}},"required":["expr"]})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let expr = args["expr"].as_str().unwrap_or("");
let parts: Vec<&str> = expr.split_whitespace().collect();
if parts.len() != 3 {
return Err(ToolError::InvalidArguments("expr".into(), "需要 'a op b' 三段式".into()));
}
let a: i64 = parts[0].parse().map_err(|_| ToolError::InvalidArguments("expr".into(), format!("无法解析 '{}'", parts[0])))?;
let b: i64 = parts[2].parse().map_err(|_| ToolError::InvalidArguments("expr".into(), format!("无法解析 '{}'", parts[2])))?;
let result = match parts[1] {
"+" => a + b,
"-" => a - b,
"*" => a * b,
"/" => a.checked_div(b).ok_or_else(|| {
ToolError::InvalidArguments("expr".into(), "除数不能为 0".into())
})?,
op => return Err(ToolError::InvalidArguments("expr".into(), format!("不支持的运算符: {op}"))),
};
Ok(json!({"result": result}))
}
}
/// 通过 MemoryStore trait 读写笔记:直接持有 Arc<dyn MemoryStore>
/// 绕开 AgentSession 封装(NoteTool 在 tool.execute 中直接操作 store)。
/// 关键前缀 "note:" 用于 list 过滤。
struct NoteTool { store: Arc<dyn MemoryStore> }
impl NoteTool { const PREFIX: &'static str = "note:"; }
#[async_trait]
impl BaseTool for NoteTool {
fn name(&self) -> &str { "note" }
fn description(&self) -> &str { "笔记 save/query: save(key, content) / query()" }
fn parameters(&self) -> Value {
json!({
"type":"object",
"properties":{
"action":{"type":"string","enum":["save","query"]},
"key":{"type":"string"},
"content":{"type":"string"}
},
"required":["action"]
})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let action = args["action"].as_str().unwrap_or("");
match action {
"save" => {
let key = args["key"].as_str().unwrap_or("");
let content = args["content"].as_str().unwrap_or("");
let item = MemoryItem {
id: format!("{}{}", Self::PREFIX, key),
content: content.to_string(),
metadata: json!({}),
created_at: OffsetDateTime::now_utc(),
};
self.store.save(item).await
.map_err(|e| ToolError::ExecutionFailed("note".into(), e.to_string()))?;
Ok(json!({"saved": key}))
}
"query" => {
let filter = MemoryFilter { prefix: Some(Self::PREFIX.into()), ..Default::default() };
let items = self.store.list(&filter).await
.map_err(|e| ToolError::ExecutionFailed("note".into(), e.to_string()))?;
let notes: Vec<String> = items.into_iter().map(|i| i.content).collect();
Ok(json!({"notes": notes}))
}
_ => Err(ToolError::InvalidArguments("action".into(), format!("未知 action: {action}"))),
}
}
}
// === Mock response helper ===
fn resp(content: Vec<ContentBlock>, stop: StopReason, u: (u32, u32)) -> MessageResponse {
MessageResponse { id: String::new(), model: "mock".into(),
message: Message::Assistant { content },
usage: Usage::from_input_output(u.0, u.1),
stop_reason: stop, extra: Default::default() }
}
fn mock_responses() -> Vec<MessageResponse> {
vec![
// 第 1 轮:calc(25 * 4) → tool_result(100) → 文本回答
resp(vec![ContentBlock::ToolUse { id: "t1".into(), name: "calc".into(),
input: json!({"expr": "25 * 4"}) }], StopReason::ToolUse, (5, 8)),
resp(vec![ContentBlock::Text { text: "25 * 4 = 100".into() }], StopReason::Stop, (8, 12)),
// 第 2 轮:note(save, last_calc, "100") → tool_result(saved) → 文本回答
resp(vec![ContentBlock::ToolUse { id: "t2".into(), name: "note".into(),
input: json!({"action": "save", "key": "last_calc", "content": "100"}) }],
StopReason::ToolUse, (10, 14)),
resp(vec![ContentBlock::Text { text: "已记录:last_calc = 100".into() }], StopReason::Stop, (12, 16)),
// 第 3 轮:note(query) → tool_result([100]) → 文本回答
resp(vec![ContentBlock::ToolUse { id: "t3".into(), name: "note".into(),
input: json!({"action": "query"}) }], StopReason::ToolUse, (8, 8)),
resp(vec![ContentBlock::Text { text: "您刚才的计算结果是 100".into() }], StopReason::Stop, (10, 14)),
// 后续冗余响应(防止队列耗尽报错)
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
]
}
// === Provider selection ===
fn select_provider() -> Arc<dyn LlmProvider> {
if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
let cfg = ProviderConfig::from_env("AG_LLM").expect("AG_LLM_* 环境变量解析失败");
let provider_type = env::var("AG_LLM_PROVIDER").ok()
.and_then(|s| s.parse::<ProviderType>().ok())
.unwrap_or(ProviderType::OpenaiChat);
Arc::from(create_provider(provider_type, cfg).expect("Provider 创建失败"))
} else {
let found: Vec<&str> = ["AG_LLM_BASE_URL", "AG_LLM_API_KEY", "AG_LLM_MODEL"]
.iter().filter(|k| env::var(k).is_ok()).copied().collect();
eprintln!("AG_LLM_* 环境变量不完整(检测到: {:?}),回退到 MockProvider", found);
Arc::new(MockProvider::new(mock_responses()))
}
}
// === Main ===
#[tokio::main]
async fn main() {
println!("=== agcore 端到端演示 ===");
let dir = TempDir::new().expect("TempDir 创建失败");
let db_path = dir.path().join("agcore.db");
let backend: Arc<dyn MemoryStore> =
Arc::new(SqliteStore::open(&db_path).expect("SqliteStore 打开失败"));
println!("💾 SqliteStore: {}", db_path.display());
let provider_label = if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
"真实 LLM Provider"
} else {
"MockProvider (离线回退模式)"
};
println!("🔄 Provider: {provider_label}");
let mut registry = ToolRegistry::new();
registry.register(Arc::new(EchoTool)).unwrap();
registry.register(Arc::new(CalcTool)).unwrap();
registry.register(Arc::new(NoteTool { store: backend.clone() })).unwrap();
println!("🔧 注册工具: {:?}", registry.list_tools());
let bundle = Arc::new(AgentBuilder::new()
.provider(select_provider())
.tool_registry(Arc::new(registry))
.hook_executor(Arc::new(HookExecutor::new()))
.build().expect("RuntimeBundle 装配失败"));
let mut session = AgentSession::new(Arc::new(AssistantAgent), "e2e-1", bundle.clone());
println!("\n第 1 轮 用户: 帮我算 25 * 4");
let r1 = session.submit_turn("帮我算 25 * 4").await.expect("turn 1 失败");
println!(" → 回答: {}", r1.text());
println!("\n第 2 轮 用户: 记下来:结果是 100");
let r2 = session.submit_turn("记下来:结果是 100").await.expect("turn 2 失败");
println!(" → 回答: {}", r2.text());
println!("\n第 3 轮 用户: 我刚才算了什么?");
let r3 = session.submit_turn("我刚才算了什么?").await.expect("turn 3 失败");
println!(" → 回答: {}", r3.text());
let total = session.usage().total();
println!("\n📊 用量: prompt={}, completion={}, total={}",
total.prompt_tokens, total.completion_tokens, total.total_tokens);
println!("\n=== 持久化验证 ===");
// 显式释放所有对 backend 的 Arc 引用,确保 SqliteStore Connection 真正关闭。
// 释放顺序:session → bundle(间接持有 NoteTool → backend clone)→ backend 局部变量。
drop(session); // session.bundle Arc 计数 -1
drop(bundle); // bundle Arc 计数归零 → registry → NoteTool → backend clone Arc 计数 2→1
drop(backend); // backend 局部变量 Arc 计数 1→0 → SqliteStore::drop → Connection 自动 close
let backend2: Arc<dyn MemoryStore> =
Arc::new(SqliteStore::open(&db_path).expect("重开 SqliteStore 失败"));
let filter = MemoryFilter { prefix: Some("note:".into()), ..Default::default() };
let items = backend2.list(&filter).await.expect("list 失败");
println!("✓ 跨连接数据存活: 找到 {} 条 note", items.len());
assert!(!items.is_empty(), "持久化验证失败:重开后无数据");
for i in &items {
println!(" - {} = {}", i.id, i.content);
}
println!("\n✓ 端到端演示完成");
}
+20 -20
View File
@@ -38,14 +38,26 @@ async fn main() {
let ks = KnowledgeStore::new(store);
let pages = vec![
make_page("rust-1", "Rust 入门", "Rust 是一门系统级编程语言,注重安全性与并发。"),
make_page("python-1", "Python 简介", "Python 是一门动态类型的高级编程语言。"),
make_page(
"rust-1",
"Rust 入门",
"Rust 是一门系统级编程语言,注重安全性与并发。",
),
make_page(
"python-1",
"Python 简介",
"Python 是一门动态类型的高级编程语言。",
),
make_page(
"langgraph-1",
"LangGraph 框架",
"LangGraph 是 LangChain 的状态图扩展,用于构建多步 Agent。",
),
make_page("rust-async", "Rust 异步编程", "Rust 异步基于 tokio 与 futures 抽象。"),
make_page(
"rust-async",
"Rust 异步编程",
"Rust 异步基于 tokio 与 futures 抽象。",
),
];
for p in &pages {
ks.add_page(p.clone()).await.expect("保存页面失败");
@@ -62,14 +74,8 @@ async fn main() {
let result = retriever.retrieve("Rust 异步").await.unwrap();
println!("query: {}", result.query);
for item in &result.items {
println!(
" 命中: {} (score={:.3})",
item.page.title, item.score
);
assert!(
(0.0..=1.0).contains(&item.score),
"score 应在 [0, 1] 区间"
);
println!(" 命中: {} (score={:.3})", item.page.title, item.score);
assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
}
assert!(!result.items.is_empty(), "应至少命中一个页面");
@@ -85,14 +91,8 @@ async fn main() {
min_score: 0.5,
};
let retriever2 = MemoryRetriever::new(ks2, cfg);
let result = retriever2
.retrieve("完全不相关的火锅配方")
.await
.unwrap();
println!(
"无关 query → items.len = {} (期望 0)",
result.items.len()
);
let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
println!("无关 query → items.len = {} (期望 0)", result.items.len());
assert!(result.items.is_empty());
// 4. max_results 截断
@@ -146,4 +146,4 @@ async fn main() {
assert!(only_stop.items.is_empty(), "纯停用词 query 必须返回空结果");
println!("\n✓ knowledge_search_demo 完成");
}
}
+11 -7
View File
@@ -11,7 +11,7 @@
use agcore::llm::types::message::{ContentBlock, Message};
use agcore::prompt::{
validate_messages, PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext,
PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext, validate_messages,
};
fn message_text(msg: &Message) -> String {
@@ -27,16 +27,16 @@ fn message_text(msg: &Message) -> String {
})
.collect(),
Message::UserImage { .. } => "[image]".into(),
_ => String::new(),
}
}
fn main() {
// 1. PromptTemplate::compile + render —— 直接构造模板
println!("=== PromptTemplate::compile + render ===");
let tpl = PromptTemplate::compile(
"今日 {{location}} 天气:{{condition}},温度 {{temperature}}",
)
.expect("编译失败");
let tpl =
PromptTemplate::compile("今日 {{location}} 天气:{{condition}},温度 {{temperature}}")
.expect("编译失败");
let mut ctx = TemplateContext::new();
ctx.insert("location", "北京");
ctx.insert("condition", "");
@@ -58,7 +58,10 @@ fn main() {
.register("weather", "今日 {{location}}{{condition}}")
.expect("注册失败");
registry
.register("greet", "你好 {{name}}{{#if formal}} 见到您很荣幸。{{/if}}")
.register(
"greet",
"你好 {{name}}{{#if formal}} 见到您很荣幸。{{/if}}",
)
.expect("注册失败");
let mut ctx = TemplateContext::new();
@@ -88,6 +91,7 @@ fn main() {
Message::User { .. } | Message::UserImage { .. } => "user",
Message::Assistant { .. } => "assistant",
Message::ToolResult { .. } => "tool",
_ => "unknown",
};
println!("[{i}] {role}: {}", message_text(m));
}
@@ -105,4 +109,4 @@ fn main() {
}
println!("\n✓ prompt_composer 完成");
}
}
+60
View File
@@ -0,0 +1,60 @@
//! quick_start —— 30 行最小可运行示例,展示 Agent / BaseTool / Builder / Session 四层抽象。
//!
//! 运行:`cargo run --example quick_start`(离线,零配置)
use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider;
use agcore::llm::provider::LlmProvider;
use agcore::llm::types::{Usage, message::{ContentBlock, Message}, response_v2::{MessageResponse, StopReason}};
use agcore::tools::{BaseTool, ToolContext, ToolError, ToolRegistry};
use async_trait::async_trait;
use serde_json::{Value, json};
struct Greeter;
impl Agent for Greeter {
fn name(&self) -> &str { "greeter" }
fn system_prompt(&self) -> Option<&str> { Some("中文助手,先调用 echo 工具,再总结。") }
}
struct EchoTool;
#[async_trait]
impl BaseTool for EchoTool {
fn name(&self) -> &str { "echo" }
fn description(&self) -> &str { "回显文本" }
fn parameters(&self) -> Value {
json!({"type":"object","properties":{"text":{"type":"string"}},"required":["text"]})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let text = args.get("text").and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidArguments("text".into(), "需要 string 类型的 text 参数".into()))?;
Ok(json!({"echoed": format!("收到: {text}")}))
}
}
fn resp(content: Vec<ContentBlock>, stop: StopReason, u: (u32, u32)) -> MessageResponse {
MessageResponse { id: String::new(), model: "mock".into(), message: Message::Assistant { content },
usage: Usage::from_input_output(u.0, u.1), stop_reason: stop, extra: Default::default() }
}
#[tokio::main]
async fn main() {
let mut registry = ToolRegistry::new();
registry.register(Arc::new(EchoTool)).unwrap();
let provider: Arc<dyn LlmProvider> = Arc::new(MockProvider::new(vec![
resp(vec![ContentBlock::ToolUse { id: "c1".into(), name: "echo".into(),
input: json!({"text": "你好"}) }], StopReason::ToolUse, (5, 8)),
resp(vec![ContentBlock::Text { text: "EchoTool 已收到您的消息并完成回传。".into() }],
StopReason::Stop, (8, 16)),
]));
let bundle = Arc::new(AgentBuilder::new()
.provider(provider).tool_registry(Arc::new(registry))
.hook_executor(Arc::new(HookExecutor::new())).build().unwrap());
let mut session = AgentSession::new(Arc::new(Greeter), "qs", bundle);
let resp = session.submit_turn("你好").await.unwrap();
let text = resp.text();
println!("LLM: {text}");
assert!(text.contains("收到"), "响应应包含'收到'字样: {text}");
println!("\n✓ quick_start 完成");
}
+7 -6
View File
@@ -3,7 +3,7 @@ use std::env;
use agcore::init_tracing;
use agcore::llm::{
cycle::{CycleConfig, LlmCycle},
provider::{create_provider, ProviderConfig, ProviderType},
provider::{ProviderConfig, ProviderType, create_provider},
types::{message::ContentBlock, message::Message, response_v2::MessageResponse},
};
@@ -51,10 +51,11 @@ async fn main() {
base_url,
api_key,
model: model.clone(),
timeout_secs: 30,
max_retries: 3,
};
let provider = create_provider(provider_type, config)
.expect("创建 Provider 失败");
let provider = create_provider(provider_type, config).expect("创建 Provider 失败");
let cycle_config = CycleConfig {
model,
@@ -63,9 +64,9 @@ async fn main() {
..CycleConfig::default()
};
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![
Message::system("你是一个简洁的助手,对于任何问题都是用一句话回答。"),
]);
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![Message::system(
"你是一个简洁的助手,对于任何问题都是用一句话回答。",
)]);
println!("发送请求...");
+3 -5
View File
@@ -17,9 +17,9 @@ use std::sync::Arc;
use agcore::llm::cycle::{CycleConfig, LlmCycle};
use agcore::llm::mock::MockProvider;
use agcore::llm::provider::LlmProvider;
use agcore::llm::types::Usage;
use agcore::llm::types::message::{ContentBlock, Message};
use agcore::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
use agcore::llm::types::Usage;
use futures_util::StreamExt;
/// 构造预设的纯文本响应。
@@ -99,9 +99,7 @@ async fn main() {
// 上层 Agent 通过 `match` 或 `?` 处理 `AgentError::Llm(_)`。
println!("\n=== 阶段 2:错误路径(队列耗尽)===");
let mut cycle = LlmCycle::new_with_arc(dyn_provider, CycleConfig::default());
let result = cycle
.submit_stream("第二次提问".to_string(), vec![])
.await;
let result = cycle.submit_stream("第二次提问".to_string(), vec![]).await;
match result {
Ok(_) => panic!("阶段 2 必须失败(队列耗尽)"),
Err(e) => {
@@ -114,4 +112,4 @@ async fn main() {
}
println!("\n✓ streaming_events_demo 完成");
}
}
+25 -24
View File
@@ -8,27 +8,13 @@
//! 5. 错误路径:非法 JSON / 空 steps / 缺字段 → `AgentError::PlanParse`
//!
//! 运行:`cargo run --example task_agent_demo`
//!
//! ## 已知技术债(v0.2 迁移指南)
//!
//! 本示例使用 `#[deprecated]` 标记的旧 wire-format 类型:
//! - `ChatResponse`、`OpenaiChatMessage`、`FinishReason` —— `OpenaiChatProvider::chat_inner()`
//! 内部转换层仍在使用(参见 `docs/10a-phase0-types-and-trait.md` §2.5.1),
//! 故结构体定义保留。
//! - `StepStatus::Completed(ChatResponse)` —— 因为 `Step` 的"已完成"变体需携带
//! provider 响应,目前沿用旧的 `ChatResponse`。
//!
//! **触发迁移的条件**v0.2 引入 IR 层的 `StepResult` / 切换为 `MessageResponse`。
//! **迁移路径**:将本文件 `ChatResponse`/`OpenaiChatMessage`/`FinishReason` 替换为
//! `MessageResponse`/`Message`/`StopReason`,移除顶部 `#![allow(deprecated)]`。
//! 上层应用代码(`TaskAgent` 消费者)也可同步迁移。
#![allow(deprecated)]
use std::collections::HashMap;
use agcore::agent::{AgentError, JsonPlanParser, PlanParser, Step, StepStatus};
use agcore::llm::types::openai_message::OpenaiChatMessage;
use agcore::llm::types::shared::FinishReason;
use agcore::llm::types::{ChatResponse, Usage};
use agcore::llm::types::message::Message;
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
use agcore::llm::types::Usage;
#[tokio::main]
async fn main() {
@@ -67,21 +53,36 @@ async fn main() {
assert!(step.status.is_pending());
step.status = StepStatus::Running;
println!("Running: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
println!(
"Running: pending={}, terminal={}",
step.status.is_pending(),
step.status.is_terminal()
);
step.status = StepStatus::Completed(ChatResponse {
message: OpenaiChatMessage::assistant_text("天气:晴,22°C"),
step.status = StepStatus::Completed(MessageResponse {
id: String::new(),
model: "mock".into(),
message: Message::assistant("天气:晴,22°C"),
usage: Usage::from_input_output(5, 10),
stop_reason: Some(FinishReason::Stop),
stop_reason: StopReason::Stop,
extra: HashMap::new(),
});
println!("Completed: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
println!(
"Completed: pending={}, terminal={}",
step.status.is_pending(),
step.status.is_terminal()
);
assert!(step.status.is_terminal());
// 3. 失败路径
println!("\n=== Step 状态机:失败路径 ===");
let mut fail_step = Step::new(0, "调用天气 API");
fail_step.status = StepStatus::Failed(AgentError::Other("API 不可用".into()));
println!("Failed: pending={}, terminal={}", fail_step.status.is_pending(), fail_step.status.is_terminal());
println!(
"Failed: pending={}, terminal={}",
fail_step.status.is_pending(),
fail_step.status.is_terminal()
);
assert!(fail_step.status.is_terminal());
// 4. 跳过路径
+6 -1
View File
@@ -11,6 +11,7 @@
pub mod agent;
pub mod builder;
pub mod context;
pub mod error;
pub mod runtime;
pub mod session;
@@ -20,9 +21,13 @@ pub mod task;
// 重导出公共 API(按使用频度排序)
pub use agent::Agent;
pub use builder::AgentBuilder;
pub use context::{
ContextBudget, ContextSlot, DeriveStrategy, FocusedConfig, SlotConfig, SlotMeta, SlotMode,
SlotSource,
};
pub use error::AgentError;
pub use runtime::{AgentConfig, RuntimeBundle};
pub use session::AgentSession;
pub use session_memory::SessionMemory;
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
pub use task::JsonPlanParser;
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
+2 -4
View File
@@ -7,14 +7,12 @@
//! - **不绑定业务循环**`submit_turn` 在 `AgentSession` 上,不在 trait 上
use crate::agent::runtime::RuntimeBundle;
#[allow(deprecated)]
use crate::llm::types::ToolDefinition;
use crate::llm::types::tool::ToolDef;
/// Agent 角色抽象。
///
/// 实现此 trait 即可接入 Agent Runtime。典型实现是 struct 持有静态配置(name、system prompt 模板),
/// 也可以是基于配置动态生成的轻量实现。
#[allow(deprecated)]
pub trait Agent: Send + Sync {
/// 角色名(用于日志、调试、UI 展示)。
fn name(&self) -> &str;
@@ -26,7 +24,7 @@ pub trait Agent: Send + Sync {
///
/// **默认实现**:从 `bundle.tool_registry` 取全部工具(最常用模式)。
/// **子 trait / 具体实现可覆盖**:做白名单、过滤、按状态动态调整等。
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDefinition> {
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDef> {
bundle.tool_registry.definitions()
}
}
+8 -6
View File
@@ -92,15 +92,17 @@ impl AgentBuilder {
/// `AgentError::Config(...)`,提示调用 `.provider(...)` / `.tool_registry(...)` /
/// `.hook_executor(...)` 补齐。不 panic。
pub fn build(self) -> Result<RuntimeBundle, AgentError> {
let provider = self
.provider
.ok_or_else(|| AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into()))?;
let provider = self.provider.ok_or_else(|| {
AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into())
})?;
let tool_registry = self
.tool_registry
.ok_or_else(|| AgentError::Config("缺少 tool_registry,请先调用 .tool_registry(...)(即使是空 ToolRegistry 也需要传入)".into()))?;
let hook_executor = self
.hook_executor
.ok_or_else(|| AgentError::Config("缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into()))?;
let hook_executor = self.hook_executor.ok_or_else(|| {
AgentError::Config(
"缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into(),
)
})?;
let config = self.config.unwrap_or_default();
+900
View File
@@ -0,0 +1,900 @@
//! ContextSlot —— 多上下文槽位管理。
//!
//! 设计要点(参见 `docs/17-phase10-contextslot.md`):
//!
//! - **多上下文分区**:单个 session 内可创建/切换/派生多个独立消息上下文
//! - **三种模式**Full(完整历史)/ Focused(读时过滤)/ Readonly(禁止写入)
//! - **三种来源**New(全新)/ Derived(派生)/ Static(静态)
//! - **基于 MemoryStore trait 持久化**JSON blob 批次存储,每 slot 3-4 条 MemoryItem 记录
//! - **零新依赖方向**:放在 `agent/` 下利用已有的 `agent → memory` 依赖
use serde::{Deserialize, Serialize};
use time::OffsetDateTime;
use crate::agent::error::AgentError;
use crate::llm::types::message::Message;
use crate::memory::store::MemoryStore;
use crate::memory::types::{MemoryFilter, MemoryItem};
/// 上下文槽 —— 一段带策略配置的消息列表。
#[derive(Debug, Clone)]
pub struct ContextSlot {
/// 当前 slot 的唯一标识(同一个 session_id 内唯一)。
pub id: String,
/// 所属 session。
pub session_id: String,
/// 槽配置。
pub config: SlotConfig,
/// 消息列表(全量,Focused/Readonly 在读取时做策略过滤)。
pub messages: Vec<Message>,
/// 槽元数据。
pub meta: SlotMeta,
}
/// 槽配置。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SlotConfig {
/// 槽模式(Full / Focused / Readonly)。
pub mode: SlotMode,
/// 槽来源(New / Derived / Static)。
pub source: SlotSource,
/// 上下文预算(v0.2 纯数据结构,无消费逻辑)。
pub budget: ContextBudget,
/// 是否启用自动压缩(v0.2 保留字段,LlmCycle 内部自行判断)。
pub compact: bool,
}
impl Default for SlotConfig {
fn default() -> Self {
Self {
mode: SlotMode::Full,
source: SlotSource::New,
budget: ContextBudget::default(),
compact: true,
}
}
}
/// 槽模式。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[non_exhaustive]
pub enum SlotMode {
/// 完整对话历史(全部消息)。
Full,
/// 聚焦模式 —— 读取时按策略过滤,保持 LLM 注意力。
Focused(FocusedConfig),
/// 只读参考上下文 —— 禁止写入。
Readonly,
}
/// 聚焦模式配置。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FocusedConfig {
/// 是否保留 system prompt。
pub keep_system: bool,
/// 保留的最近消息条数(以消息条数而非对话轮次为单位,因为一轮对话可能包含多条 tool 消息)。
pub recent_messages: usize,
/// 摘要覆盖(v0.2 仅消费端:手动设置则注入,不自动生成)。
/// v0.3 将支持 Hook 驱动的自动摘要生成。
pub summary_override: Option<String>,
}
/// 槽来源。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[non_exhaustive]
pub enum SlotSource {
/// 全新空槽。
New,
/// 从父 slot 派生(记录 parent_id)。
Derived {
parent_id: String,
strategy: DeriveStrategy,
},
/// 预置静态消息(不持久化,随 session 生命周期存在)。
Static(Vec<Message>),
}
/// 派生策略。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum DeriveStrategy {
/// 完整复制父 slot 的消息。
Full,
/// 按聚焦策略复制父 slot 的消息。
Focused(FocusedConfig),
}
/// 上下文预算(v0.2 纯数据结构,无消费逻辑)。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ContextBudget {
/// system prompt 预算。
pub system: u32,
/// 对话历史预算。
pub history: u32,
/// 工具定义预算。
pub tools: u32,
/// 工具结果预算。
pub tool_results: u32,
/// 预留 buffer。
pub reserve: u32,
}
impl Default for ContextBudget {
fn default() -> Self {
Self {
system: 8_000,
history: 80_000,
tools: 10_000,
tool_results: 20_000,
reserve: 10_000,
}
}
}
impl ContextBudget {
/// 自动分配:按上下文窗口的固定比例分配预算。
/// v0.2 只做占位实现,v0.3 将根据实际 provider 的 context_window 计算。
pub fn auto() -> Self {
Self::default()
}
}
/// 槽元数据。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SlotMeta {
/// 父 slot id(仅 Derived 来源有值)。
pub parent_id: Option<String>,
/// 消息总数。
pub message_count: usize,
/// 总 token 估算值(由 add_messages 时累计,v0.2 为近似值)。
pub total_tokens: u32,
/// 创建时间(Unix 时间戳,秒)。
pub created_at: u64,
}
impl SlotMeta {
pub fn new() -> Self {
Self {
parent_id: None,
message_count: 0,
total_tokens: 0,
created_at: std::time::SystemTime::now()
.duration_since(std::time::SystemTime::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0),
}
}
}
impl Default for SlotMeta {
fn default() -> Self {
Self::new()
}
}
impl ContextSlot {
/// 持久化 key 前缀。
const KEY_DATA: &'static str = "slot_data";
const KEY_META: &'static str = "slot_meta";
const KEY_CONFIG: &'static str = "slot_config";
const KEY_REL: &'static str = "slot_rel";
/// 校验 id 不含冒号(避免破坏 key 格式与 list prefix 过滤)。
/// 失败时 panic —— 这是开发者错误而非用户错误。
fn assert_no_colon(id: &str, field: &str) {
if id.contains(':') {
panic!(
"{field} '{id}' contains ':' which would break key format. \
Use only letters, digits, hyphens and underscores."
);
}
}
pub(crate) fn data_key(session_id: &str, slot_id: &str) -> String {
Self::assert_no_colon(session_id, "session_id");
Self::assert_no_colon(slot_id, "slot_id");
format!("{}:{}:{}", Self::KEY_DATA, session_id, slot_id)
}
pub(crate) fn meta_key(session_id: &str, slot_id: &str) -> String {
Self::assert_no_colon(session_id, "session_id");
Self::assert_no_colon(slot_id, "slot_id");
format!("{}:{}:{}", Self::KEY_META, session_id, slot_id)
}
pub(crate) fn config_key(session_id: &str, slot_id: &str) -> String {
Self::assert_no_colon(session_id, "session_id");
Self::assert_no_colon(slot_id, "slot_id");
format!("{}:{}:{}", Self::KEY_CONFIG, session_id, slot_id)
}
pub(crate) fn rel_key(session_id: &str, child_id: &str) -> String {
Self::assert_no_colon(session_id, "session_id");
Self::assert_no_colon(child_id, "child_id");
format!("{}:{}:{}", Self::KEY_REL, session_id, child_id)
}
/// 构造 MemoryItem 的辅助函数。
fn make_item(key: String, content: String) -> MemoryItem {
MemoryItem {
id: key,
content,
metadata: serde_json::json!({}),
created_at: OffsetDateTime::now_utc(),
}
}
/// 创建一个新的空 ContextSlot(不持久化,仅内存构造)。
pub fn new(
session_id: impl Into<String>,
slot_id: impl Into<String>,
config: SlotConfig,
) -> Self {
Self {
id: slot_id.into(),
session_id: session_id.into(),
config,
messages: Vec::new(),
meta: SlotMeta::new(),
}
}
/// 保存 slot 数据到存储后端(全量写入,含 config)。
pub async fn save(&self, store: &dyn MemoryStore) -> Result<(), AgentError> {
let data = serde_json::to_string(&self.messages)
.map_err(|e| AgentError::Other(e.to_string()))?;
let meta = serde_json::to_string(&self.meta)
.map_err(|e| AgentError::Other(e.to_string()))?;
let config = serde_json::to_string(&self.config)
.map_err(|e| AgentError::Other(e.to_string()))?;
store
.save(Self::make_item(
Self::data_key(&self.session_id, &self.id),
data,
))
.await
.map_err(AgentError::Memory)?;
store
.save(Self::make_item(
Self::meta_key(&self.session_id, &self.id),
meta,
))
.await
.map_err(AgentError::Memory)?;
store
.save(Self::make_item(
Self::config_key(&self.session_id, &self.id),
config,
))
.await
.map_err(AgentError::Memory)?;
// 派生关系
if let SlotSource::Derived { parent_id, .. } = &self.config.source {
store
.save(Self::make_item(
Self::rel_key(&self.session_id, &self.id),
parent_id.clone(),
))
.await
.map_err(AgentError::Memory)?;
}
Ok(())
}
/// 从存储加载 slotconfig 从 `slot_config` key 自行恢复。
/// 若 config 记录不存在(旧版本升级场景),使用 `SlotConfig::default()`。
pub async fn load(
id: &str,
session_id: &str,
store: &dyn MemoryStore,
) -> Result<Option<Self>, AgentError> {
let meta_item = store
.get(&Self::meta_key(session_id, id))
.await
.map_err(AgentError::Memory)?;
let data_item = store
.get(&Self::data_key(session_id, id))
.await
.map_err(AgentError::Memory)?;
let config_item = store
.get(&Self::config_key(session_id, id))
.await
.map_err(AgentError::Memory)?;
match (meta_item, data_item) {
(Some(m), Some(d)) => {
let meta: SlotMeta = serde_json::from_str(&m.content)
.map_err(|e| AgentError::Other(e.to_string()))?;
let messages: Vec<Message> = serde_json::from_str(&d.content)
.map_err(|e| AgentError::Other(e.to_string()))?;
// config 从存储恢复;不存在则使用 default(兼容旧版本)
let config = match config_item {
Some(c) => serde_json::from_str(&c.content)
.map_err(|e| AgentError::Other(e.to_string()))?,
None => SlotConfig::default(),
};
Ok(Some(Self {
id: id.to_string(),
session_id: session_id.to_string(),
config,
messages,
meta,
}))
}
_ => Ok(None),
}
}
/// 列出某 session 下的所有 slot 元数据。
pub async fn list(
session_id: &str,
store: &dyn MemoryStore,
) -> Result<Vec<SlotMeta>, AgentError> {
let prefix_str = format!("{}:{}:", Self::KEY_META, session_id);
let filter = MemoryFilter {
prefix: Some(prefix_str),
..Default::default()
};
let items = store.list(&filter).await.map_err(AgentError::Memory)?;
let mut metas = Vec::new();
for item in items {
if let Ok(meta) = serde_json::from_str::<SlotMeta>(&item.content) {
metas.push(meta);
}
}
Ok(metas)
}
/// 删除 slot 的所有存储记录(slot_data + slot_meta + slot_config + slot_rel)。
pub async fn delete(
id: &str,
session_id: &str,
store: &dyn MemoryStore,
) -> Result<(), AgentError> {
store
.delete(&Self::data_key(session_id, id))
.await
.map_err(AgentError::Memory)?;
store
.delete(&Self::meta_key(session_id, id))
.await
.map_err(AgentError::Memory)?;
store
.delete(&Self::config_key(session_id, id))
.await
.map_err(AgentError::Memory)?;
// slot_rel 是 best-effort(仅 Derived 来源的 slot 才有此 key
let _ = store.delete(&Self::rel_key(session_id, id)).await;
Ok(())
}
/// 追加消息(Readonly 模式下返回 `SlotReadonly` 错误)。
/// Full / Focused 模式下允许追加。
pub fn append_messages(&mut self, new_messages: Vec<Message>) -> Result<(), AgentError> {
if matches!(self.config.mode, SlotMode::Readonly) {
return Err(AgentError::SlotReadonly(
"Readonly slot does not allow writes".into(),
));
}
let count = new_messages.len();
self.messages.extend(new_messages);
self.meta.message_count += count;
Ok(())
}
/// 按 FocusedConfig 过滤消息(静态辅助函数,被 `load_messages` 和 `derive_slot` 复用)。
///
/// 过滤逻辑:
/// 1. 保留第一条 system 消息(如果 `keep_system=true`
/// 2. 取最近 `recent_messages` 条非 system 消息(如果 `recent_messages > 0`
/// 3. 追加摘要消息(如果 `summary_override` 存在)
pub fn filter_focused(messages: &[Message], cfg: &FocusedConfig) -> Vec<Message> {
let mut result = Vec::new();
// 保留 system prompt
if cfg.keep_system
&& let Some(msg) = messages
.iter()
.find(|m| matches!(m, Message::System { .. }))
{
result.push(msg.clone());
}
// 处理 recent_messages=0 边界:上面已处理 system,下面仅取最近 N 条
if cfg.recent_messages > 0 {
let recent: Vec<&Message> = messages
.iter()
.filter(|m| !matches!(m, Message::System { .. }))
.collect();
let start = recent.len().saturating_sub(cfg.recent_messages);
for msg in recent.iter().skip(start) {
result.push((*msg).clone());
}
}
// 注入摘要
if let Some(summary) = &cfg.summary_override {
result.push(Message::system(format!("[上下文摘要] {}", summary)));
}
result
}
/// 返回消息列表。Focused 模式下按策略过滤(裁剪到最近 recent_messages 条)。
pub fn load_messages(&self) -> Vec<Message> {
match &self.config.mode {
SlotMode::Focused(cfg) => Self::filter_focused(&self.messages, cfg),
_ => self.messages.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::store::InMemoryStore;
fn make_store() -> std::sync::Arc<dyn MemoryStore> {
std::sync::Arc::new(InMemoryStore::new())
}
fn make_slot(id: &str, session: &str) -> ContextSlot {
ContextSlot::new(session, id, SlotConfig::default())
}
/// 提取 `Message` 的第一个 Text block 的内容(用于测试断言)。
/// 返回 None 表示该消息不含纯文本 block。
fn extract_text(msg: &Message) -> &str {
use crate::llm::types::message::ContentBlock;
let blocks = match msg {
Message::System { content }
| Message::User { content }
| Message::Assistant { content } => content,
Message::UserImage { .. } => return "",
Message::ToolResult { content, .. } => content,
};
for block in blocks {
if let ContentBlock::Text { text } = block {
return text;
}
}
""
}
// ===== 持久化 =====
#[tokio::test]
async fn slot_save_load_roundtrip() {
let store = make_store();
let mut slot = make_slot("default", "s1");
slot.append_messages(vec![Message::user_text("hi")]).unwrap();
slot.append_messages(vec![Message::assistant("hello")]).unwrap();
slot.save(&*store).await.unwrap();
let loaded = ContextSlot::load("default", "s1", &*store).await.unwrap();
let loaded = loaded.expect("slot should exist after save");
assert_eq!(loaded.id, "default");
assert_eq!(loaded.session_id, "s1");
assert_eq!(loaded.messages.len(), 2);
assert_eq!(loaded.meta.message_count, 2);
}
#[tokio::test]
async fn slot_session_isolation() {
let store = make_store();
let mut a = make_slot("main", "sA");
a.append_messages(vec![Message::user_text("only in A")])
.unwrap();
a.save(&*store).await.unwrap();
let mut b = make_slot("main", "sB");
b.append_messages(vec![Message::user_text("only in B")])
.unwrap();
b.save(&*store).await.unwrap();
let loaded_a = ContextSlot::load("main", "sA", &*store).await.unwrap().unwrap();
let loaded_b = ContextSlot::load("main", "sB", &*store).await.unwrap().unwrap();
assert_eq!(extract_text(&loaded_a.messages[0]), "only in A");
assert_eq!(extract_text(&loaded_b.messages[0]), "only in B");
}
#[tokio::test]
async fn slot_derived_parent_id_recorded() {
let store = make_store();
let slot = ContextSlot::new(
"s1",
"child",
SlotConfig {
mode: SlotMode::Full,
source: SlotSource::Derived {
parent_id: "default".to_string(),
strategy: DeriveStrategy::Full,
},
budget: ContextBudget::default(),
compact: true,
},
);
slot.save(&*store).await.unwrap();
// rel_key 直接读
let rel = store
.get(&ContextSlot::rel_key("s1", "child"))
.await
.unwrap()
.unwrap();
assert_eq!(rel.content, "default");
}
#[tokio::test]
async fn slot_readonly_rejects_write() {
let mut slot = make_slot("ro", "s1");
slot.config.mode = SlotMode::Readonly;
let result = slot.append_messages(vec![Message::user_text("nope")]);
assert!(matches!(result, Err(AgentError::SlotReadonly(_))));
assert!(slot.messages.is_empty());
}
#[tokio::test]
async fn slot_delete_then_load_none() {
let store = make_store();
let mut slot = make_slot("to_delete", "s1");
slot.append_messages(vec![Message::user_text("hi")]).unwrap();
slot.save(&*store).await.unwrap();
ContextSlot::delete("to_delete", "s1", &*store).await.unwrap();
let loaded = ContextSlot::load("to_delete", "s1", &*store).await.unwrap();
assert!(loaded.is_none());
}
#[tokio::test]
async fn slot_list_multiple() {
let store = make_store();
for id in ["alpha", "beta", "gamma"] {
let mut s = make_slot(id, "sX");
s.append_messages(vec![Message::user_text(id)]).unwrap();
s.save(&*store).await.unwrap();
}
// 不同 session 不该列出
let mut s2 = make_slot("alpha", "sY");
s2.append_messages(vec![Message::user_text("y")]).unwrap();
s2.save(&*store).await.unwrap();
let metas = ContextSlot::list("sX", &*store).await.unwrap();
assert_eq!(metas.len(), 3);
let metas_y = ContextSlot::list("sY", &*store).await.unwrap();
assert_eq!(metas_y.len(), 1);
}
// ===== Focused 模式 =====
#[tokio::test]
async fn slot_focused_recent_messages() {
let mut slot = make_slot("f", "s1");
slot.append_messages(vec![Message::system("sys")]).unwrap();
for i in 0..5 {
slot.append_messages(vec![Message::user_text(format!("u{i}"))])
.unwrap();
slot.append_messages(vec![Message::assistant(format!("a{i}"))])
.unwrap();
}
slot.config.mode = SlotMode::Focused(FocusedConfig {
keep_system: true,
recent_messages: 3,
summary_override: None,
});
let loaded = slot.load_messages();
// system + 最近 3 条 (assistant 4, user 4, assistant 5 实际是按 vec 顺序取最近 3 条非 system)
let has_sys = loaded.iter().any(|m| matches!(m, Message::System { .. }));
assert!(has_sys, "system 提示应保留");
// 最近 3 条非 system 应该是 a4, u4, a5 (按 messages 存储顺序的最后 3 条)
assert_eq!(loaded.len(), 1 + 3);
}
#[tokio::test]
async fn slot_focused_summary_override() {
let mut slot = make_slot("f", "s1");
slot.append_messages(vec![Message::user_text("u")]).unwrap();
slot.append_messages(vec![Message::assistant("a")]).unwrap();
slot.config.mode = SlotMode::Focused(FocusedConfig {
keep_system: false,
recent_messages: 100,
summary_override: Some("讨论了 X".to_string()),
});
let loaded = slot.load_messages();
// 2 条原始 + 1 条摘要 system = 3
assert_eq!(loaded.len(), 3);
// 最后一条是摘要
if let Message::System { content } = &loaded[2] {
let text = format!("{:?}", content);
assert!(text.contains("上下文摘要"));
assert!(text.contains("讨论了 X"));
} else {
panic!("最后一条应为 system 摘要");
}
}
#[tokio::test]
async fn slot_focused_zero_messages() {
let mut slot = make_slot("f", "s1");
slot.append_messages(vec![Message::system("sys")]).unwrap();
slot.append_messages(vec![Message::user_text("u")]).unwrap();
slot.config.mode = SlotMode::Focused(FocusedConfig {
keep_system: true,
recent_messages: 0,
summary_override: None,
});
let loaded = slot.load_messages();
// recent_messages=0 但 keep_system=true 应只含 system
assert_eq!(loaded.len(), 1);
assert!(matches!(loaded[0], Message::System { .. }));
}
// ===== 边界 =====
#[tokio::test]
async fn slot_empty_messages_roundtrip() {
let store = make_store();
let slot = make_slot("empty", "s1");
slot.save(&*store).await.unwrap();
let loaded = ContextSlot::load("empty", "s1", &*store).await.unwrap().unwrap();
assert!(loaded.messages.is_empty());
assert_eq!(loaded.meta.message_count, 0);
}
#[tokio::test]
async fn slot_save_on_readonly_side_effect() {
let store = make_store();
let mut slot = make_slot("ro", "s1");
slot.config.mode = SlotMode::Readonly;
// save 本身允许(只禁止 append
slot.save(&*store).await.unwrap();
let loaded = ContextSlot::load("ro", "s1", &*store).await.unwrap();
assert!(loaded.is_some());
}
// ===== 派生 (derive_slot 行为) =====
#[tokio::test]
async fn derive_full_copies_parent_messages() {
let mut parent = make_slot("p", "s1");
for i in 0..3 {
parent.append_messages(vec![Message::user_text(format!("u{i}"))])
.unwrap();
}
// 模拟 derive_slot 内部 Full 策略
let child_messages = parent.messages.clone();
let child = ContextSlot::new(
"s1",
"c",
SlotConfig {
mode: SlotMode::Full,
source: SlotSource::Derived {
parent_id: "p".to_string(),
strategy: DeriveStrategy::Full,
},
budget: ContextBudget::default(),
compact: true,
},
);
let mut child = child;
child.messages = child_messages;
let store = make_store();
child.save(&*store).await.unwrap();
let loaded = ContextSlot::load("c", "s1", &*store).await.unwrap().unwrap();
assert_eq!(loaded.messages.len(), 3);
assert!(matches!(loaded.config.source, SlotSource::Derived { .. }));
}
#[tokio::test]
async fn derive_focused_filters_parent_messages() {
let mut parent = make_slot("p", "s1");
parent.append_messages(vec![Message::system("sys")]).unwrap();
for i in 0..5 {
parent.append_messages(vec![Message::user_text(format!("u{i}"))])
.unwrap();
}
// 模拟 derive_slot 内部 Focused 策略:按 FocusedConfig 过滤
let cfg = FocusedConfig {
keep_system: true,
recent_messages: 2,
summary_override: None,
};
// 应用 load_messages 同样的过滤
let mut filtered = Vec::new();
if cfg.keep_system
&& let Some(m) = parent
.messages
.iter()
.find(|m| matches!(m, Message::System { .. }))
{
filtered.push(m.clone());
}
if cfg.recent_messages > 0 {
let recent: Vec<&Message> = parent
.messages
.iter()
.filter(|m| !matches!(m, Message::System { .. }))
.collect();
let start = recent.len().saturating_sub(cfg.recent_messages);
for m in recent.iter().skip(start) {
filtered.push((*m).clone());
}
}
assert_eq!(filtered.len(), 1 + 2); // system + 2 条
}
#[tokio::test]
async fn derived_slot_loadable_independently() {
let store = make_store();
let mut parent = make_slot("p", "s1");
parent.append_messages(vec![Message::user_text("u")]).unwrap();
parent.save(&*store).await.unwrap();
// 派生 child
let mut child = ContextSlot::new(
"s1",
"c",
SlotConfig {
mode: SlotMode::Full,
source: SlotSource::Derived {
parent_id: "p".to_string(),
strategy: DeriveStrategy::Full,
},
budget: ContextBudget::default(),
compact: true,
},
);
child.append_messages(vec![Message::user_text("derived msg")])
.unwrap();
child.save(&*store).await.unwrap();
// child 可独立加载
let loaded = ContextSlot::load("c", "s1", &*store).await.unwrap().unwrap();
assert_eq!(loaded.messages.len(), 1);
assert_eq!(extract_text(&loaded.messages[0]), "derived msg");
}
// ===== delete 保护 (AgentSession 层,但 ContextSlot.delete 不保护;逻辑测试在 session.rs) =====
#[tokio::test]
async fn slot_delete_cleans_all_records() {
let store = make_store();
let mut slot = ContextSlot::new(
"s1",
"x",
SlotConfig {
mode: SlotMode::Full,
source: SlotSource::Derived {
parent_id: "p".to_string(),
strategy: DeriveStrategy::Full,
},
budget: ContextBudget::default(),
compact: true,
},
);
slot.append_messages(vec![Message::user_text("u")]).unwrap();
slot.save(&*store).await.unwrap();
// 确认所有记录存在
assert!(store.get(&ContextSlot::data_key("s1", "x")).await.unwrap().is_some());
assert!(store.get(&ContextSlot::meta_key("s1", "x")).await.unwrap().is_some());
assert!(store.get(&ContextSlot::config_key("s1", "x")).await.unwrap().is_some());
assert!(store.get(&ContextSlot::rel_key("s1", "x")).await.unwrap().is_some());
ContextSlot::delete("x", "s1", &*store).await.unwrap();
// data/meta/config 已删
assert!(store.get(&ContextSlot::data_key("s1", "x")).await.unwrap().is_none());
assert!(store.get(&ContextSlot::meta_key("s1", "x")).await.unwrap().is_none());
assert!(store.get(&ContextSlot::config_key("s1", "x")).await.unwrap().is_none());
}
// ===== 基础类型测试 =====
#[test]
fn slot_meta_new_sets_zero_message_count() {
let m = SlotMeta::new();
assert_eq!(m.message_count, 0);
assert_eq!(m.total_tokens, 0);
assert!(m.parent_id.is_none());
}
#[test]
fn context_budget_default_sum_128k() {
let b = ContextBudget::default();
assert_eq!(b.system + b.history + b.tools + b.tool_results + b.reserve, 128_000);
}
#[test]
fn slot_config_default_is_full_new() {
let c = SlotConfig::default();
assert!(matches!(c.mode, SlotMode::Full));
assert!(matches!(c.source, SlotSource::New));
assert!(c.compact);
}
#[test]
fn focused_config_serializes_roundtrip() {
let cfg = FocusedConfig {
keep_system: true,
recent_messages: 5,
summary_override: Some("sum".into()),
};
let json = serde_json::to_string(&cfg).unwrap();
let back: FocusedConfig = serde_json::from_str(&json).unwrap();
assert_eq!(back.recent_messages, 5);
assert_eq!(back.summary_override.as_deref(), Some("sum"));
}
// ====== filter_focused 静态方法(被 load_messages 和 derive_slot 复用) ======
#[test]
fn filter_focused_keeps_system_and_recent() {
let mut messages = vec![Message::system("sys")];
for i in 0..5 {
messages.push(Message::user_text(format!("u{i}")));
messages.push(Message::assistant(format!("a{i}")));
}
let cfg = FocusedConfig {
keep_system: true,
recent_messages: 3,
summary_override: None,
};
let filtered = ContextSlot::filter_focused(&messages, &cfg);
// system + 3 条最近的非 system 消息
assert_eq!(filtered.len(), 1 + 3);
assert!(matches!(filtered[0], Message::System { .. }));
}
#[test]
fn filter_focused_injects_summary() {
let messages = vec![
Message::user_text("u"),
Message::assistant("a"),
];
let cfg = FocusedConfig {
keep_system: false,
recent_messages: 100,
summary_override: Some("讨论了 X".into()),
};
let filtered = ContextSlot::filter_focused(&messages, &cfg);
// 2 条原始 + 1 条摘要 system
assert_eq!(filtered.len(), 3);
if let Message::System { content } = &filtered[2] {
let text = format!("{:?}", content);
assert!(text.contains("上下文摘要"));
} else {
panic!("最后一条应为 system 摘要");
}
}
// ====== Colon 校验(key 格式保护) ======
#[test]
#[should_panic(expected = "session_id 's:1' contains ':'")]
fn key_constructor_rejects_colon_in_session_id() {
// 通过 make_slot 间接调用 slot.save 时会触发 data_key -> assert_no_colon
let store = make_store();
let slot = ContextSlot::new("s:1", "default", SlotConfig::default());
let _ = tokio_test_runtime(slot.save(&*store));
}
#[test]
#[should_panic(expected = "slot_id 'a:b' contains ':'")]
fn key_constructor_rejects_colon_in_slot_id() {
let store = make_store();
let slot = ContextSlot::new("s1", "a:b", SlotConfig::default());
let _ = tokio_test_runtime(slot.save(&*store));
}
/// 在同步测试中运行 future 的辅助函数。
fn tokio_test_runtime<F: std::future::Future>(f: F) -> F::Output {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap()
.block_on(f)
}
}
+54 -3
View File
@@ -18,6 +18,7 @@ use crate::tools::error::ToolError;
/// **不实现 `Clone`**:透传内层 `LlmError` / `MemoryError`,两者均未派生 `Clone`(保留
/// 完整错误信息,传递所有权)。如需在多 session 间共享错误状态,用 `Arc<AgentError>` 包装。
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum AgentError {
/// LLM 调用错误(透传 Phase 0)。
#[error("LLM 错误: {0}")]
@@ -35,6 +36,18 @@ pub enum AgentError {
#[error("Plan 解析错误: {0}")]
PlanParse(String),
/// Readonly slot 不允许写入(Phase 10 新增)。
#[error("Readonly slot 不允许写入: {0}")]
SlotReadonly(String),
/// Slot 不存在(Phase 10 新增)。
#[error("Slot '{0}' 不存在")]
SlotNotFound(String),
/// Slot 已存在(Phase 10 新增)。
#[error("Slot '{0}' 已存在")]
SlotAlreadyExists(String),
/// 钩子阻断操作(Agent 层特有)。
#[error("钩子阻断: {0}")]
HookBlocked(String),
@@ -59,6 +72,7 @@ impl AgentError {
/// - `Tool`:由内层 `is_recoverable()` 决定
/// - `HookBlocked` / `LimitExceeded`:不可恢复(需人工介入或终止循环)
/// - `Config` / `Other`:不可恢复
/// - `SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`:不可恢复(结构性错误)
pub fn is_recoverable(&self) -> bool {
match self {
Self::Llm(e) => matches!(
@@ -68,9 +82,13 @@ impl AgentError {
Self::Tool(e) => e.is_recoverable(),
Self::Memory(e) => e.is_recoverable(),
Self::PlanParse(_) => false,
Self::HookBlocked(_) | Self::LimitExceeded(_) | Self::Config(_) | Self::Other(_) => {
false
}
Self::SlotReadonly(_)
| Self::SlotNotFound(_)
| Self::SlotAlreadyExists(_)
| Self::HookBlocked(_)
| Self::LimitExceeded(_)
| Self::Config(_)
| Self::Other(_) => false,
}
}
}
@@ -180,4 +198,37 @@ mod tests {
let err = caller().unwrap_err();
assert!(matches!(err, AgentError::Memory(_)));
}
// ====== Phase 10: Slot 错误变体测试 ======
#[test]
fn slot_readonly_not_recoverable() {
assert!(!AgentError::SlotReadonly("readonly".into()).is_recoverable());
}
#[test]
fn slot_not_found_not_recoverable() {
assert!(!AgentError::SlotNotFound("missing".into()).is_recoverable());
}
#[test]
fn slot_already_exists_not_recoverable() {
assert!(!AgentError::SlotAlreadyExists("dup".into()).is_recoverable());
}
#[test]
fn slot_error_messages() {
assert_eq!(
format!("{}", AgentError::SlotReadonly("readonly".into())),
"Readonly slot 不允许写入: readonly"
);
assert_eq!(
format!("{}", AgentError::SlotNotFound("foo".into())),
"Slot 'foo' 不存在"
);
assert_eq!(
format!("{}", AgentError::SlotAlreadyExists("bar".into())),
"Slot 'bar' 已存在"
);
}
}
+1 -1
View File
@@ -16,8 +16,8 @@ use std::sync::Arc;
use std::time::Duration;
use crate::llm::compact::CompactConfig;
use crate::llm::provider::LlmProvider;
use crate::llm::hooks::HookExecutor;
use crate::llm::provider::LlmProvider;
use crate::memory::retriever::MemoryRetriever;
use crate::memory::store::MemoryStore;
use crate::tools::ToolRegistry;
+746 -102
View File
@@ -1,29 +1,41 @@
//! AgentSession —— 智能体"会话"实例。
//!
//! 设计要点(参见 `docs/7-agent-runtime.md` §3.2.3):
//! 设计要点(参见 `docs/7-agent-runtime.md` §3.2.3 与 `docs/17-phase10-contextslot.md`):
//!
//! - **会话 = 角色 + 状态**:绑定 `session_id` / `agent` / `bundle`,累计 `turn_index` 和 `cost_so_far`
//! - **多上下文分区**Phase 10):通过 `ContextSlot` 管理多个独立的消息上下文
//! - **最小 reference impl**`submit_turn` 演示"组装 LlmCycle → submit_with_tools → 累计 cost"的标准流程
//! - **不做业务循环**:多轮策略、错误重试、记忆回写由上层应用或具体 `TaskAgent` 决定
//! - **不持有 ConversationMemory**:上层可独立 new 一个 `ConversationMemory`,在合适的时机调 `add_message`
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::Arc;
use futures_core::Stream;
use crate::agent::agent::Agent;
use crate::agent::context::{
ContextSlot, DeriveStrategy, SlotConfig, SlotMode, SlotSource,
};
use crate::agent::error::AgentError;
use crate::agent::runtime::RuntimeBundle;
use crate::agent::session_memory::SessionMemory;
use crate::llm::cycle::{CostTracker, CycleConfig, LlmCycle};
use crate::llm::hooks::{HookContext, HookEvent};
use crate::llm::stream::StreamEvent;
use crate::llm::types::message::Message;
use crate::llm::types::response_v2::MessageResponse;
use crate::memory::store::InMemoryStore;
use crate::memory::store::{InMemoryStore, MemoryStore};
/// Agent 会话实例。
///
/// 同一 `Agent` 可被多个 `AgentSession` 复用(不同 session_id 互不干扰)。
/// `submit_turn` 一次只跑一轮 LLM 调用(含自动 tool 循环)。
///
/// **Phase 10 新增**:通过 `slots: HashMap<String, ContextSlot>` 管理多个独立的对话上下文。
/// `submit_turn` 默认写入当前活跃 slot`current_slot_id`)。
///
/// **不实现 `Clone`**session 持有累计 `turn_index` / `cost_so_far` / `session_memory`
/// 共享这些状态需要显式 sync 语义;如果上层需要并发访问,自己用 `Arc<Mutex<_>>` 包装。
pub struct AgentSession {
@@ -36,6 +48,10 @@ pub struct AgentSession {
cost_so_far: CostTracker,
/// 会话级记忆(Phase 4c 替换内联 HashMap)。
pub session_memory: SessionMemory,
/// Phase 10 新增:所有 slotid → ContextSlot)。
slots: HashMap<String, ContextSlot>,
/// Phase 10 新增:当前活跃 slot 的 id。
current_slot_id: String,
}
impl std::fmt::Debug for AgentSession {
@@ -46,6 +62,8 @@ impl std::fmt::Debug for AgentSession {
.field("turn_index", &self.turn_index)
.field("cost_so_far", &self.cost_so_far.total())
.field("session_memory", &"<SessionMemory>")
.field("slots", &self.slots.keys().collect::<Vec<_>>())
.field("current_slot_id", &self.current_slot_id)
.finish()
}
}
@@ -54,6 +72,14 @@ impl AgentSession {
/// 创建一个新的会话实例。
///
/// `agent` 与 `bundle` 共同决定 `submit_turn` 行为:system_prompt / 工具集 / LLM 后端均来自它们。
///
/// Phase 10 新增:自动创建 `"default"` slot,确保简单场景无感使用。
///
/// **注意**:`new()` 是同步函数,无法执行异步的 `ContextSlot::load()`。
/// 因此始终创建空的 default slot。若需要从存储恢复 session 历史,
/// 可在创建后调用 `switch_slot("default")` 尝试从存储加载。
/// v0.2 简化:`switch_slot` 在 HashMap 中已有 key 时不会重载——如需恢复,
/// 请在清空 `slots` 后调用 `switch_slot`,或等待 v0.3 的懒加载支持。)
pub fn new(
agent: Arc<dyn Agent>,
session_id: impl Into<String>,
@@ -65,6 +91,16 @@ impl AgentSession {
.clone()
.unwrap_or_else(|| Arc::new(InMemoryStore::new()));
let session_memory = SessionMemory::new(backend, &session_id_str);
// 自动创建 "default" slot
let default_slot = ContextSlot::new(
&session_id_str,
"default",
SlotConfig::default(),
);
let mut slots = HashMap::new();
slots.insert("default".to_string(), default_slot);
Self {
session_id: session_id_str,
agent,
@@ -72,6 +108,8 @@ impl AgentSession {
turn_index: 0,
cost_so_far: CostTracker::default(),
session_memory,
slots,
current_slot_id: "default".to_string(),
}
}
@@ -96,7 +134,9 @@ impl AgentSession {
key: impl Into<String>,
value: impl Into<String>,
) -> Result<(), AgentError> {
self.session_memory.set(&key.into(), &value.into()).await
self.session_memory
.set(&key.into(), &value.into())
.await
}
/// 读取一条会话级数据。
@@ -104,20 +144,158 @@ impl AgentSession {
self.session_memory.get(key).await
}
/// Phase 10: 当前 slot id。
pub fn current_slot_id(&self) -> &str {
&self.current_slot_id
}
/// Phase 10: 列出所有 slot 的不可变引用(按 id 顺序)。
pub fn slots(&self) -> impl Iterator<Item = (&String, &ContextSlot)> {
self.slots.iter()
}
/// Phase 10: 解析存储后端。
/// fallback 链:`session_memory_backend` → `memory_store` → `InMemoryStore`。
fn resolve_store(&self) -> Arc<dyn MemoryStore> {
self.bundle
.session_memory_backend
.clone()
.or_else(|| self.bundle.memory_store.clone())
.unwrap_or_else(|| Arc::new(InMemoryStore::new()))
}
// ====== Phase 10: Slot 管理方法 ======
/// 创建新 slot(config 可选,不传则使用默认值)。
pub async fn create_slot(
&mut self,
id: impl Into<String>,
config: Option<SlotConfig>,
) -> Result<(), AgentError> {
let id = id.into();
if self.slots.contains_key(&id) {
return Err(AgentError::SlotAlreadyExists(id));
}
let slot = ContextSlot::new(
&self.session_id,
&id,
config.unwrap_or_default(),
);
slot.save(&*self.resolve_store()).await?;
self.slots.insert(id, slot);
Ok(())
}
/// 切换到指定 slot。
/// - 如果 slot 已在内存中,直接切换 current_slot_id
/// - 如果不在内存中,尝试从存储加载(config 自动从 slot_config key 恢复)
/// - 存储中也不存在则返回 `SlotNotFound`
pub async fn switch_slot(&mut self, id: &str) -> Result<(), AgentError> {
if !self.slots.contains_key(id) {
let store = self.resolve_store();
match ContextSlot::load(id, &self.session_id, &*store).await? {
Some(slot) => {
self.slots.insert(id.to_string(), slot);
}
None => return Err(AgentError::SlotNotFound(id.to_string())),
}
}
self.current_slot_id = id.to_string();
Ok(())
}
/// 列出所有 slot id。
pub fn list_slots(&self) -> impl Iterator<Item = &String> {
self.slots.keys()
}
/// 从父 slot 派生新 slot(继承父 slot 的全量或聚焦消息)。
pub async fn derive_slot(
&mut self,
id: impl Into<String>,
parent_id: &str,
strategy: DeriveStrategy,
) -> Result<(), AgentError> {
let slot_id = id.into();
if self.slots.contains_key(&slot_id) {
return Err(AgentError::SlotAlreadyExists(slot_id));
}
let parent = self
.slots
.get(parent_id)
.ok_or_else(|| AgentError::SlotNotFound(parent_id.to_string()))?;
let parent_messages = parent.messages.clone();
let (messages, focused_cfg) = match &strategy {
DeriveStrategy::Full => (parent_messages, None),
DeriveStrategy::Focused(cfg) => {
let filtered = ContextSlot::filter_focused(&parent_messages, cfg);
(filtered, Some(cfg.clone()))
}
};
let mode = match focused_cfg {
Some(cfg) => SlotMode::Focused(cfg),
None => SlotMode::Full,
};
let slot = ContextSlot::new(
&self.session_id,
&slot_id,
SlotConfig {
mode,
source: SlotSource::Derived {
parent_id: parent_id.to_string(),
strategy,
},
budget: Default::default(),
compact: true,
},
);
let mut slot = slot;
slot.messages = messages;
slot.save(&*self.resolve_store()).await?;
self.slots.insert(slot_id, slot);
Ok(())
}
/// 删除一个 slot。
/// - 禁止删除 `"default"` slot
/// - 至少保留一个 slot
/// - 已删除后再 load 返回 None
pub async fn delete_slot(&mut self, id: &str) -> Result<(), AgentError> {
if id == "default" {
return Err(AgentError::Config("Cannot delete the 'default' slot".into()));
}
if self.slots.len() <= 1 {
return Err(AgentError::Config("Cannot delete the last slot".into()));
}
ContextSlot::delete(id, &self.session_id, &*self.resolve_store()).await?;
self.slots.remove(id);
if self.current_slot_id == id {
self.current_slot_id = "default".to_string();
}
Ok(())
}
// ====== 原始方法 ======
/// 提交一轮对话(含自动 tool 循环),返回 LLM 响应。
///
/// 流程
/// 1. 触发 `OnTurnStart` hook
/// 2. 组装 `LlmCycle`(注入 system_prompt / hook_executor / compact_config / 消息历史
/// 3. `submit_with_tools` 跑单轮对话
/// 4. 累计 `cost_so_far`
/// 5. 触发 `OnTurnEnd` hook
/// 6. `turn_index += 1`
/// Phase 10 改造
/// - 从当前 slot 加载历史(Focused 模式读时过滤)
/// - 提交完成后**增量追加**本轮新增消息到当前 slot(不覆盖,确保 Focused 语义不丢数据
///
/// **不做**
/// - 不持有 `ConversationMemory`(由上层独立 task 决定何时回写)
/// - 不做 Plan 拆解(Phase 4b 才加 `TaskAgent`
/// - 不做 session_data 持久化(Phase 4c 替换为 `SessionMemory`
/// 流程
/// 1. 检查当前 slot 不是 Readonly
/// 2. 触发 `OnTurnStart` hook
/// 3. 加载当前 slot 的历史消息
/// 4. 组装 `LlmCycle`(注入 system_prompt / compact_config / 历史)
/// 5. `submit_with_tools` 跑单轮对话
/// 6. 累计 `cost_so_far`
/// 7. **增量追加**本轮新增消息到当前 slot + 保存到 store
/// 8. 触发 `OnTurnEnd` hook
/// 9. `turn_index += 1`
pub async fn submit_turn(
&mut self,
user_input: impl Into<String>,
@@ -125,50 +303,196 @@ impl AgentSession {
let turn_index = self.turn_index;
let hook_executor = Arc::clone(&self.bundle.hook_executor);
// 0. Readonly slot 拒绝 submit_turn
{
let slot = self
.slots
.get(&self.current_slot_id)
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?;
if matches!(slot.config.mode, SlotMode::Readonly) {
return Err(AgentError::SlotReadonly(format!(
"Cannot submit turn on Readonly slot '{}'",
self.current_slot_id
)));
}
}
// 1. 触发 OnTurnStart hook
let start_ctx =
HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
let start_ctx = HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
hook_executor
.execute(HookEvent::OnTurnStart, &start_ctx)
.await;
// 2. 组装 LlmCycle —— 共享 bundle 中的 provider 句柄
// 工具列表从 agent.tool_definitions(bundle) 派生(默认 = bundle 全量);
// submit_with_tools 内部从 registry 自行取 definitions,此处仅消费以触发
// 子 trait 覆盖(白名单/过滤)的副作用。
// 2. 从当前 slot 加载历史消息
let history = self
.slots
.get(&self.current_slot_id)
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?
.load_messages();
// 3. 组装 LlmCycle
let _ = self.agent.tool_definitions(&self.bundle);
let mut cycle = LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default())
.with_messages(Vec::new());
// Phase 2 切换 system_prompt 字段为 Message::SystemFIX-D)。
// 若 agent 自带 system prompt,预置到 messages 列表头部。
let mut initial_messages: Vec<Message> = Vec::new();
let mut cycle =
LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default());
let mut messages_with_prompt = history;
if let Some(prompt) = self.agent.system_prompt() {
initial_messages.push(Message::system(prompt));
}
if !initial_messages.is_empty() {
cycle = cycle.with_messages(initial_messages);
messages_with_prompt.insert(0, Message::system(prompt));
}
let input_len = messages_with_prompt.len();
cycle = cycle.with_messages(messages_with_prompt);
if let Some(cfg) = self.bundle.config.compact_config.clone() {
cycle = cycle.with_compact_config(cfg);
}
// 3. 提交HookExecutor 不在这里传——内部 hook 由 LlmCycle 在 PreRequest/PostRequest 触发)
// 4. 提交
let response = cycle
.submit_with_tools(user_input.into(), &self.bundle.tool_registry)
.await?;
// 4. 累计 cost
// 5. 累计 cost
self.cost_so_far.add(&response.usage);
// 5. 触发 OnTurnEnd hook
// 6. 只将本轮新增消息追加到当前 slot(保留全量历史,确保 Focused 模式的"读时过滤"语义不丢失数据)
// cycle.messages() 包含 [system_prompt?, history..., user_input, tool_calls..., final_response]
// 新增消息 = cycle.messages()[input_len..](跳过 initial_messages,即跳过已被持久化的内容)
let new_messages: Vec<Message> = cycle
.messages()
.iter()
.skip(input_len)
.cloned()
.collect();
let store = self.resolve_store();
if let Some(slot) = self.slots.get_mut(&self.current_slot_id) {
slot.append_messages(new_messages)?;
slot.save(&*store).await?;
}
// 7. 触发 OnTurnEnd hook
let end_ctx = HookContext::new(HookEvent::OnTurnEnd).with_turn_index(turn_index);
hook_executor.execute(HookEvent::OnTurnEnd, &end_ctx).await;
// 6. turn_index 递增
// 8. turn_index 递增
self.turn_index += 1;
Ok(response)
}
/// 提交一轮对话(流式版本,含自动 tool 循环),返回 `StreamEvent` 流。
///
/// Phase 10 改造:
/// - 从当前 slot 加载历史(Focused 模式读时过滤)
/// - finalize_turn 需要传入本轮新增消息列表
///
/// **运行时要求**:内部委托 `submit_with_tools_stream`,需要 tokio 多线程运行时。
pub async fn submit_turn_stream(
&mut self,
user_input: impl Into<String>,
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, AgentError> {
let turn_index = self.turn_index;
let hook_executor = Arc::clone(&self.bundle.hook_executor);
// 0. Readonly 检查
{
let slot = self
.slots
.get(&self.current_slot_id)
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?;
if matches!(slot.config.mode, SlotMode::Readonly) {
return Err(AgentError::SlotReadonly(format!(
"Cannot submit turn stream on Readonly slot '{}'",
self.current_slot_id
)));
}
}
// 1. 触发 OnTurnStart hook
let start_ctx = HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
hook_executor
.execute(HookEvent::OnTurnStart, &start_ctx)
.await;
// 2. 触发子 trait 覆盖(白名单/过滤)的副作用
let _ = self.agent.tool_definitions(&self.bundle);
// 3. 从当前 slot 加载历史
let history = self
.slots
.get(&self.current_slot_id)
.ok_or_else(|| AgentError::SlotNotFound(self.current_slot_id.clone()))?
.load_messages();
// 4. 组装 LlmCycle
let mut cycle =
LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default());
let mut messages_with_prompt = history;
if let Some(prompt) = self.agent.system_prompt() {
messages_with_prompt.insert(0, Message::system(prompt));
}
cycle = cycle.with_messages(messages_with_prompt);
if let Some(cfg) = self.bundle.config.compact_config.clone() {
cycle = cycle.with_compact_config(cfg);
}
// 5. 调用流式工具循环
let stream = cycle
.submit_with_tools_stream(
user_input.into(),
Arc::clone(&self.bundle.tool_registry),
)
.await?;
// 6. turn_index 递增 —— 配合 finalize_turn 用 (turn_index - 1) 传递正确的 OnTurnEnd 序号
self.turn_index += 1;
// 注:hook_executor 不显式 drop,生命周期由 Arc 自动管理
let _ = hook_executor;
Ok(stream)
}
/// 完成一轮 turn:累计 cost + 触发 OnTurnEnd hook + 增量追加消息到当前 slot。
///
/// Phase 10 改造:
/// - 新增 `new_messages_from_cycle` 参数:流式场景下,本轮新增的消息列表
/// (由消费者在流消费完毕后从 `cycle.messages()[input_len..]` 获取并传入)
/// - 仅**增量追加**到当前 slot(不覆盖已有消息),与 submit_turn 行为一致
/// - 返回类型从 `()` 改为 `Result<(), AgentError>`,错误传播更清晰
///
/// 由消费者在收到 `MessageComplete.full_response` 后调用。
pub async fn finalize_turn(
&mut self,
response: &MessageResponse,
new_messages_from_cycle: Vec<Message>,
) -> Result<(), AgentError> {
self.cost_so_far.add(&response.usage);
// 防御性检查:current_slot_id 必须在 slots 中(与 submit_turn 行为一致)。
// 正常流程:submit_turn_stream 已注册 slotfinalize_turn 不应触发此分支。
if !self.slots.contains_key(&self.current_slot_id) {
return Err(AgentError::SlotNotFound(self.current_slot_id.clone()));
}
// 增量追加到当前 slot(仅 Full/Focused 模式允许,Readonly 阻断)
let store = self.resolve_store();
if let Some(slot) = self.slots.get_mut(&self.current_slot_id) {
if matches!(slot.config.mode, SlotMode::Readonly) {
return Err(AgentError::SlotReadonly(format!(
"Cannot finalize turn on Readonly slot '{}'",
self.current_slot_id
)));
}
slot.append_messages(new_messages_from_cycle)?;
slot.save(&*store).await?;
}
// 防御性 saturating_sub 防止误用 panic。
let end_ctx = HookContext::new(HookEvent::OnTurnEnd)
.with_turn_index(self.turn_index.saturating_sub(1));
self.bundle
.hook_executor
.execute(HookEvent::OnTurnEnd, &end_ctx)
.await;
Ok(())
}
}
#[cfg(test)]
@@ -223,9 +547,7 @@ mod tests {
id: String::new(),
model: String::new(),
message: Message::Assistant {
content: vec![ContentBlock::Text {
text: text.into(),
}],
content: vec![ContentBlock::Text { text: text.into() }],
},
usage: crate::llm::types::Usage::from_input_output(10, 5),
stop_reason: StopReason::Stop,
@@ -233,65 +555,7 @@ mod tests {
}
}
/// 烟雾测试 1AgentSession::submit_turn 跑通 mock provider。
#[tokio::test]
async fn submit_turn_runs_with_mock_provider() {
let provider = Arc::new(MockProvider::new(vec![assistant_text("hello back")]));
let agent = Arc::new(StubAgent {
name: "stub".into(),
prompt: Some("you are a test agent".into()),
});
let bundle = Arc::new(
AgentBuilder::new()
.provider(provider)
.tool_registry(Arc::new(ToolRegistry::new()))
.hook_executor(Arc::new(HookExecutor::new()))
.build()
.unwrap(),
);
let mut session = AgentSession::new(agent, "s1", bundle);
assert_eq!(session.turn_index(), 0);
let response = session.submit_turn("hi").await.unwrap();
assert_eq!(response.text(), "hello back");
assert_eq!(session.turn_index(), 1);
assert_eq!(session.usage().total().prompt_tokens, 10);
assert_eq!(session.usage().total().completion_tokens, 5);
}
/// 烟雾测试 2session_data 读写。
#[tokio::test]
async fn session_data_set_get() {
let provider = Arc::new(MockProvider::new(vec![]));
let agent = Arc::new(StubAgent {
name: "stub".into(),
prompt: None,
});
let bundle = Arc::new(
AgentBuilder::new()
.provider(provider)
.tool_registry(Arc::new(ToolRegistry::new()))
.hook_executor(Arc::new(HookExecutor::new()))
.build()
.unwrap(),
);
let mut session = AgentSession::new(agent, "s2", bundle);
assert!(session.get_session_data("k").await.unwrap().is_none());
session.set_session_data("k", "v").await.unwrap();
assert_eq!(session.get_session_data("k").await.unwrap(), Some("v".into()));
// 覆盖写
session.set_session_data("k", "v2").await.unwrap();
assert_eq!(
session.get_session_data("k").await.unwrap(),
Some("v2".into())
);
}
/// 烟雾测试 3submit_turn 触发 OnTurnStart / OnTurnEnd hook。
#[tokio::test]
async fn submit_turn_triggers_turn_hooks() {
fn build_session(provider_responses: Vec<MessageResponse>) -> (AgentSession, Arc<CountHook>, Arc<CountHook>) {
let mut hook_executor = HookExecutor::new();
let start_count = Arc::new(CountHook(AtomicU32::new(0)));
let end_count = Arc::new(CountHook(AtomicU32::new(0)));
@@ -304,13 +568,10 @@ mod tests {
Box::new(CountHookAdapter(end_count.clone())),
);
let provider = Arc::new(MockProvider::new(vec![
assistant_text("ok"),
assistant_text("ok 2"),
]));
let provider = Arc::new(MockProvider::new(provider_responses));
let agent = Arc::new(StubAgent {
name: "stub".into(),
prompt: None,
prompt: Some("you are a test agent".into()),
});
let bundle = Arc::new(
AgentBuilder::new()
@@ -320,7 +581,53 @@ mod tests {
.build()
.unwrap(),
);
let mut session = AgentSession::new(agent, "s3", bundle);
let session = AgentSession::new(agent, "test-session", bundle);
(session, start_count, end_count)
}
/// 烟雾测试 1AgentSession::submit_turn 跑通 mock provider(向后兼容)。
#[tokio::test]
async fn submit_turn_runs_with_mock_provider() {
let (mut session, start_count, end_count) = build_session(vec![assistant_text("hello back")]);
assert_eq!(session.turn_index(), 0);
let response = session.submit_turn("hi").await.unwrap();
assert_eq!(extract_text(&response.message), "hello back");
assert_eq!(session.turn_index(), 1);
assert_eq!(session.usage().total().prompt_tokens, 10);
assert_eq!(session.usage().total().completion_tokens, 5);
// hook 触发
assert_eq!(start_count.0.load(Ordering::SeqCst), 1);
assert_eq!(end_count.0.load(Ordering::SeqCst), 1);
}
/// 烟雾测试 2session_data 读写。
#[tokio::test]
async fn session_data_set_get() {
let (mut session, _, _) = build_session(vec![]);
assert!(session.get_session_data("k").await.unwrap().is_none());
session.set_session_data("k", "v").await.unwrap();
assert_eq!(
session.get_session_data("k").await.unwrap(),
Some("v".into())
);
// 覆盖写
session.set_session_data("k", "v2").await.unwrap();
assert_eq!(
session.get_session_data("k").await.unwrap(),
Some("v2".into())
);
}
/// 烟雾测试 3submit_turn 触发 OnTurnStart / OnTurnEnd hook。
#[tokio::test]
async fn submit_turn_triggers_turn_hooks() {
let (mut session, start_count, end_count) = build_session(vec![
assistant_text("ok"),
assistant_text("ok 2"),
]);
session.submit_turn("hi").await.unwrap();
assert_eq!(start_count.0.load(Ordering::SeqCst), 1);
@@ -330,4 +637,341 @@ mod tests {
assert_eq!(start_count.0.load(Ordering::SeqCst), 2);
assert_eq!(end_count.0.load(Ordering::SeqCst), 2);
}
}
/// 提取 Message 的第一个 Text block(测试辅助)。
fn extract_text(msg: &Message) -> &str {
use crate::llm::types::message::ContentBlock;
let blocks = match msg {
Message::System { content }
| Message::User { content }
| Message::Assistant { content } => content,
Message::UserImage { .. } => return "",
Message::ToolResult { content, .. } => content,
};
for block in blocks {
if let ContentBlock::Text { text } = block {
return text;
}
}
""
}
// ====== Phase 10 新增测试 ======
/// Phase 10: 默认 slot 自动创建。
#[tokio::test]
async fn default_slot_auto_created() {
let (session, _, _) = build_session(vec![]);
assert_eq!(session.current_slot_id(), "default");
let slots: Vec<_> = session.list_slots().collect();
assert_eq!(slots.len(), 1);
assert!(slots.contains(&&"default".to_string()));
}
/// Phase 10: submit_turn 写入当前 slot。
#[tokio::test]
async fn submit_turn_writes_to_current_slot() {
let (mut session, _, _) = build_session(vec![assistant_text("resp")]);
session.submit_turn("user input").await.unwrap();
// 检查 default slot 内存中的消息
let slot = session.slots.get("default").expect("default slot exists");
// submit_turn 增量追加的是 cycle.messages()[input_len..] 部分,
// 即 [user_input, tool_results?, final_response](不含 system_promptsystem 由 agent 提供)
assert!(slot.messages.len() >= 2, "应至少包含 user 和 assistant");
// 验证 user 输入和 assistant 响应都已写入
let has_user = slot.messages.iter().any(|m| extract_text(m) == "user input");
let has_resp = slot.messages.iter().any(|m| extract_text(m) == "resp");
assert!(has_user && has_resp, "slot 应包含 user input 和 assistant response");
}
/// Phase 10: create_slot 创建新 slot。
#[tokio::test]
async fn create_slot_basic() {
let (mut session, _, _) = build_session(vec![]);
session.create_slot("scratch", None).await.unwrap();
let slots: Vec<_> = session.list_slots().cloned().collect();
assert!(slots.contains(&"default".to_string()));
assert!(slots.contains(&"scratch".to_string()));
assert_eq!(slots.len(), 2);
}
/// Phase 10: create_slot 拒绝重复 id。
#[tokio::test]
async fn create_slot_rejects_duplicate() {
let (mut session, _, _) = build_session(vec![]);
session.create_slot("dup", None).await.unwrap();
let err = session.create_slot("dup", None).await.unwrap_err();
assert!(matches!(err, AgentError::SlotAlreadyExists(_)));
}
/// Phase 10: switch_slot 切换并保留各自消息。
#[tokio::test]
async fn switch_slot_isolates_messages() {
let (mut session, _, _) = build_session(vec![
assistant_text("resp a"),
assistant_text("resp b"),
assistant_text("resp c"),
]);
// 1. 在 default 中提交一次
session.submit_turn("msg in default").await.unwrap();
// 2. 创建 slot_a
session.create_slot("slot_a", None).await.unwrap();
session.switch_slot("slot_a").await.unwrap();
assert_eq!(session.current_slot_id(), "slot_a");
session.submit_turn("msg in slot_a").await.unwrap();
// 3. 检查 slot_a 的消息数
let slot_a = session.slots.get("slot_a").unwrap();
let slot_a_count = slot_a.messages.len();
assert!(slot_a_count >= 2, "slot_a 至少 2 条消息,实际 {}", slot_a_count);
// 4. 切回 default,验证 default 不包含 slot_a 的消息
session.switch_slot("default").await.unwrap();
let slot_default = session.slots.get("default").unwrap();
let default_count = slot_default.messages.len();
assert!(default_count >= 2);
// 验证 default 中没有 "msg in slot_a"
let default_has_a = slot_default
.messages
.iter()
.any(|m| extract_text(m) == "msg in slot_a");
assert!(!default_has_a, "default 不应包含 slot_a 的消息");
// 验证 slot_a 中没有 "msg in default"
let slot_a = session.slots.get("slot_a").unwrap();
let a_has_default = slot_a
.messages
.iter()
.any(|m| extract_text(m) == "msg in default");
assert!(!a_has_default, "slot_a 不应包含 default 的消息");
}
/// Phase 10: Readonly slot 拒绝写入。
#[tokio::test]
async fn readonly_slot_rejects_submit_turn() {
let (mut session, _, _) = build_session(vec![assistant_text("resp")]);
session
.create_slot(
"ro",
Some(SlotConfig {
mode: SlotMode::Readonly,
source: SlotSource::New,
budget: Default::default(),
compact: true,
}),
)
.await
.unwrap();
session.switch_slot("ro").await.unwrap();
let err = session.submit_turn("blocked").await.unwrap_err();
assert!(matches!(err, AgentError::SlotReadonly(_)));
}
/// Phase 10: delete_slot 删除非 default。
#[tokio::test]
async fn delete_slot_removes_non_default() {
let (mut session, _, _) = build_session(vec![]);
session.create_slot("to_delete", None).await.unwrap();
session.delete_slot("to_delete").await.unwrap();
let slots: Vec<_> = session.list_slots().cloned().collect();
assert!(!slots.contains(&"to_delete".to_string()));
assert_eq!(slots.len(), 1);
}
/// Phase 10: delete_slot 禁止删 default。
#[tokio::test]
async fn delete_slot_rejects_default() {
let (mut session, _, _) = build_session(vec![]);
let err = session.delete_slot("default").await.unwrap_err();
assert!(matches!(err, AgentError::Config(_)));
}
/// Phase 10: delete_slot 禁止删最后一个 slot。
#[tokio::test]
async fn delete_slot_rejects_last() {
let (mut session, _, _) = build_session(vec![]);
// 只有 default 一个 slot
let err = session.delete_slot("default").await.unwrap_err();
assert!(matches!(err, AgentError::Config(_)));
}
/// Phase 10: delete_slot 后 current 回退到 default。
#[tokio::test]
async fn delete_slot_falls_back_to_default() {
let (mut session, _, _) = build_session(vec![]);
session.create_slot("temp", None).await.unwrap();
session.switch_slot("temp").await.unwrap();
assert_eq!(session.current_slot_id(), "temp");
session.delete_slot("temp").await.unwrap();
assert_eq!(session.current_slot_id(), "default");
}
/// Phase 10: derive_slot Full 策略复制父 slot 消息。
#[tokio::test]
async fn derive_slot_full_copies_parent() {
let (mut session, _, _) = build_session(vec![assistant_text("resp")]);
session.submit_turn("parent msg").await.unwrap();
session
.derive_slot("child", "default", DeriveStrategy::Full)
.await
.unwrap();
let child = session.slots.get("child").unwrap();
assert!(matches!(child.config.source, SlotSource::Derived { .. }));
// child 应有 parent 的消息拷贝
let has_parent = child
.messages
.iter()
.any(|m| extract_text(m) == "parent msg");
assert!(has_parent);
}
/// Phase 10: derive_slot 拒绝重复 id。
#[tokio::test]
async fn derive_slot_rejects_duplicate() {
let (mut session, _, _) = build_session(vec![]);
session.create_slot("child", None).await.unwrap();
let err = session
.derive_slot("child", "default", DeriveStrategy::Full)
.await
.unwrap_err();
assert!(matches!(err, AgentError::SlotAlreadyExists(_)));
}
/// Phase 10: derive_slot 父 slot 不存在返回 SlotNotFound。
#[tokio::test]
async fn derive_slot_parent_not_found() {
let (mut session, _, _) = build_session(vec![]);
let err = session
.derive_slot("child", "nonexistent", DeriveStrategy::Full)
.await
.unwrap_err();
assert!(matches!(err, AgentError::SlotNotFound(_)));
}
/// Phase 10: switch_slot 加载不存在的 slot 返回 SlotNotFound。
#[tokio::test]
async fn switch_slot_not_found() {
let (mut session, _, _) = build_session(vec![]);
let err = session.switch_slot("missing").await.unwrap_err();
assert!(matches!(err, AgentError::SlotNotFound(_)));
}
/// Phase 10: slot 数据持久化到 storageswitch 时可恢复。
/// 使用 session_memory_backend 配置可验证持久化。
#[tokio::test]
async fn slot_persistence_roundtrip() {
// 创建一个共享的 InMemoryStore 作为后端
let backend = Arc::new(InMemoryStore::new());
let provider = Arc::new(MockProvider::new(vec![assistant_text("resp")]));
let agent = Arc::new(StubAgent {
name: "stub".into(),
prompt: None,
});
let bundle = Arc::new(
AgentBuilder::new()
.provider(provider)
.tool_registry(Arc::new(ToolRegistry::new()))
.hook_executor(Arc::new(HookExecutor::new()))
.session_memory_backend(backend.clone())
.build()
.unwrap(),
);
let mut session = AgentSession::new(agent, "persist-session", bundle);
session.create_slot("persist_test", None).await.unwrap();
session.switch_slot("persist_test").await.unwrap();
session.submit_turn("hi").await.unwrap();
// 验证 data/meta/config 三个 key 都已写入共享 backend
let stored_data = backend
.get(&ContextSlot::data_key("persist-session", "persist_test"))
.await
.unwrap();
assert!(stored_data.is_some(), "data 应已持久化");
let stored_meta = backend
.get(&ContextSlot::meta_key("persist-session", "persist_test"))
.await
.unwrap();
assert!(stored_meta.is_some(), "meta 应已持久化");
let stored_config = backend
.get(&ContextSlot::config_key("persist-session", "persist_test"))
.await
.unwrap();
assert!(stored_config.is_some(), "config 应已持久化");
}
/// Phase 10: 当只有 `memory_store`(无 `session_memory_backend`)时,resolve_store
/// 应 fallback 到 `memory_store`。
#[tokio::test]
async fn resolve_store_falls_back_to_memory_store() {
let backend = Arc::new(InMemoryStore::new());
let provider = Arc::new(MockProvider::new(vec![assistant_text("ok")]));
let agent = Arc::new(StubAgent {
name: "stub".into(),
prompt: None,
});
// 注意:这里只设置 memory_store,不设置 session_memory_backend
let bundle = Arc::new(
AgentBuilder::new()
.provider(provider)
.tool_registry(Arc::new(ToolRegistry::new()))
.hook_executor(Arc::new(HookExecutor::new()))
.memory_store(backend.clone())
.build()
.unwrap(),
);
let mut session = AgentSession::new(agent, "fb-session", bundle);
session.create_slot("fb_slot", None).await.unwrap();
// 验证 backend 中已存在 fb_slot 的数据
let stored = backend
.get(&ContextSlot::data_key("fb-session", "fb_slot"))
.await
.unwrap();
assert!(stored.is_some(), "memory_store fallback 应生效");
}
/// Phase 10: finalize_turn 在 current_slot 不存在时返回 SlotNotFound(与 submit_turn 一致)。
#[tokio::test]
async fn finalize_turn_slot_not_found() {
let (mut session, _, _) = build_session(vec![assistant_text("ok")]);
// 强制 current_slot_id 指向不存在的 slot(模拟异常状态)
session.current_slot_id = "ghost".to_string();
let response = assistant_text("ok");
let err = session.finalize_turn(&response, vec![]).await.unwrap_err();
assert!(matches!(err, AgentError::SlotNotFound(_)));
}
/// Phase 10: finalize_turn 在 Readonly slot 上返回 SlotReadonly。
#[tokio::test]
async fn finalize_turn_readonly_rejects() {
let (mut session, _, _) = build_session(vec![assistant_text("ok")]);
session
.create_slot(
"ro",
Some(SlotConfig {
mode: SlotMode::Readonly,
source: SlotSource::New,
budget: Default::default(),
compact: true,
}),
)
.await
.unwrap();
session.switch_slot("ro").await.unwrap();
let response = assistant_text("ok");
let err = session
.finalize_turn(&response, vec![Message::user_text("x")])
.await
.unwrap_err();
assert!(matches!(err, AgentError::SlotReadonly(_)));
}
}
+3 -11
View File
@@ -78,11 +78,7 @@ impl SessionMemory {
prefix: Some(format!("{}:", self.namespace)),
..Default::default()
};
let items = self
.store
.list(&filter)
.await
.map_err(AgentError::Memory)?;
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
let mut lines = Vec::with_capacity(items.len() + 2);
lines.push("<session-context>".to_string());
@@ -113,11 +109,7 @@ impl SessionMemory {
prefix: Some(format!("{}:", self.namespace)),
..Default::default()
};
let items = self
.store
.list(&filter)
.await
.map_err(AgentError::Memory)?;
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
for item in items {
self.store
@@ -181,4 +173,4 @@ mod tests {
assert!(mem_a.get("key").await.unwrap().is_none());
assert_eq!(mem_b.get("key").await.unwrap(), Some("val_b".into()));
}
}
}
+5 -11
View File
@@ -10,8 +10,7 @@
//! - 重试由上层新建 `Plan` 实现,`TaskAgent` 不做自动重试
use crate::agent::error::AgentError;
#[allow(deprecated)]
use crate::llm::types::ChatResponse;
use crate::llm::types::response_v2::MessageResponse;
use async_trait::async_trait;
@@ -56,14 +55,14 @@ impl Step {
/// 均未派生 `Clone`(保留原始错误信息,传递所有权而非克隆)。如需复制 `Plan`,
/// 只能 clone 处于 `Pending` / `Running` / `Completed` / `Skipped` 状态的步骤。
#[derive(Debug)]
#[allow(deprecated)]
#[non_exhaustive]
pub enum StepStatus {
/// 初始状态 —— 等待执行。
Pending,
/// 正在执行(`TaskAgent::execute_plan` 进入)。
Running,
/// 已完成(含 LLM 响应)。
Completed(ChatResponse),
Completed(MessageResponse),
/// 失败(含错误)。
Failed(AgentError),
/// 跳过(上层主动跳过)。
@@ -130,9 +129,7 @@ impl PlanParser for JsonPlanParser {
.collect::<Result<Vec<_>, AgentError>>()?;
if steps.is_empty() {
return Err(AgentError::PlanParse(
"Plan 至少需要一个步骤".into(),
));
return Err(AgentError::PlanParse("Plan 至少需要一个步骤".into()));
}
Ok(Plan {
@@ -203,10 +200,7 @@ mod tests {
let plan = Plan {
id: "p1".into(),
goal: "test goal".into(),
steps: vec![
Step::new(0, "first"),
Step::new(1, "second"),
],
steps: vec![Step::new(0, "first"), Step::new(1, "second")],
};
assert_eq!(plan.steps.len(), 2);
assert_eq!(plan.steps[0].index, 0);
+32 -16
View File
@@ -73,10 +73,7 @@ impl CompactState {
/// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。
pub fn estimate_message_tokens(messages: &[Message]) -> u32 {
messages
.iter()
.map(estimate_single_message_tokens)
.sum()
messages.iter().map(estimate_single_message_tokens).sum()
}
fn estimate_single_message_tokens(msg: &Message) -> u32 {
@@ -99,9 +96,7 @@ fn estimate_block_tokens(block: &ContentBlock) -> u32 {
match block {
ContentBlock::Text { text } => estimate_text_tokens(text),
ContentBlock::Thinking { text, .. } => estimate_text_tokens(text),
ContentBlock::ToolUse { input, .. } => {
estimate_text_tokens(&input.to_string())
}
ContentBlock::ToolUse { input, .. } => estimate_text_tokens(&input.to_string()),
ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
// ponytail: Image / Audio / File / Extension 在 IR 中固定估算。
// 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。
@@ -148,14 +143,25 @@ pub fn microcompact(messages: &mut [Message], keep_recent: usize) -> u32 {
// 第一遍:计算可释放 token(仅非错误 ToolResult
for msg in &messages[..prune_start] {
if matches!(msg, Message::ToolResult { is_error: false, .. }) {
if matches!(
msg,
Message::ToolResult {
is_error: false,
..
}
) {
freed_tokens += estimate_single_message_tokens(msg);
}
}
// 第二遍:替换内容(仅非错误 ToolResult
for msg in &mut messages[..prune_start] {
if let Message::ToolResult { content, is_error: false, .. } = msg {
if let Message::ToolResult {
content,
is_error: false,
..
} = msg
{
*content = vec![ContentBlock::Text {
text: "[pruned]".to_string(),
}];
@@ -177,13 +183,15 @@ mod tests {
fn estimate_message_tokens_handles_all_variants() {
let messages = vec![
Message::System {
content: vec![ContentBlock::Text {
text: "sys".into(),
}],
content: vec![ContentBlock::Text { text: "sys".into() }],
},
Message::user_text("hi"),
Message::assistant("ans"),
Message::user_image("b64", "image/png", crate::llm::types::shared::ImageDetail::Auto),
Message::user_image(
"b64",
"image/png",
crate::llm::types::shared::ImageDetail::Auto,
),
Message::tool_result("call_1", "tool res", false),
];
let tokens = estimate_message_tokens(&messages);
@@ -205,7 +213,10 @@ mod tests {
assert!(freed > 0);
assert_eq!(messages.len(), before_len); // 只改内容,不删消息
// 索引 1 是被压缩的 ToolResult
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
if let Message::ToolResult {
content, is_error, ..
} = &messages[1]
{
assert_eq!(content.len(), 1);
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
assert!(!is_error);
@@ -228,9 +239,14 @@ mod tests {
assert_eq!(freed, 0); // 错误 ToolResult 不计入
assert_eq!(messages.len(), before_len);
// 错误信息保留完整
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
if let Message::ToolResult {
content, is_error, ..
} = &messages[1]
{
assert!(is_error);
assert!(matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down")));
assert!(
matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down"))
);
} else {
panic!("expected ToolResult at index 1");
}
+29 -30
View File
@@ -8,11 +8,9 @@
use serde_json::Value;
use crate::llm::types::message::{ContentBlock, Message};
use crate::llm::types::openai_message::{
ContentField, OpenaiChatMessage, OpenaiContentPart,
};
use crate::llm::types::OpenaiToolCall;
use crate::llm::types::message::{ContentBlock, Message};
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage, OpenaiContentPart};
/// `OpenaiChatMessage` → IR `Message`。
///
@@ -24,11 +22,10 @@ use crate::llm::types::OpenaiToolCall;
/// - `Function`(已废弃)→ `Message::ToolResult``name` 作为 `tool_call_id` 兜底)
pub fn from_openai(msg: &OpenaiChatMessage) -> Message {
match msg {
OpenaiChatMessage::Developer { content, .. } | OpenaiChatMessage::System { content, .. } => {
Message::System {
content: content_to_blocks(content),
}
}
OpenaiChatMessage::Developer { content, .. }
| OpenaiChatMessage::System { content, .. } => Message::System {
content: content_to_blocks(content),
},
OpenaiChatMessage::User { content, .. } => Message::User {
content: content_to_blocks(content),
},
@@ -86,7 +83,11 @@ pub fn to_openai(msg: &Message) -> OpenaiChatMessage {
content: blocks_to_content(content),
name: None,
},
Message::UserImage { data, mime_type, detail } => {
Message::UserImage {
data,
mime_type,
detail,
} => {
// ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。
let mime = mime_type.clone();
let is_url = data.starts_with("http://") || data.starts_with("https://");
@@ -167,26 +168,25 @@ pub fn content_to_blocks(field: &ContentField) -> Vec<ContentBlock> {
ContentField::Array(parts) => parts
.iter()
.filter_map(|p| match p {
OpenaiContentPart::Text { text } => {
Some(ContentBlock::Text { text: text.clone() })
}
OpenaiContentPart::Refusal { refusal } => {
Some(ContentBlock::Text { text: refusal.clone() })
}
OpenaiContentPart::Text { text } => Some(ContentBlock::Text { text: text.clone() }),
OpenaiContentPart::Refusal { refusal } => Some(ContentBlock::Text {
text: refusal.clone(),
}),
OpenaiContentPart::Image { image_url, .. } => {
// ponytail: 简化处理 —— URL 直接通过,data URI 拆出
// data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。
let url = &image_url.url;
if let Some(rest) = url.strip_prefix("data:")
&& let Some((mime, b64)) = rest.split_once(";base64,") {
return Some(ContentBlock::Image {
source: crate::llm::types::message::ImageSource {
data: b64.to_string(),
mime_type: mime.to_string(),
is_url: false,
},
});
}
&& let Some((mime, b64)) = rest.split_once(";base64,")
{
return Some(ContentBlock::Image {
source: crate::llm::types::message::ImageSource {
data: b64.to_string(),
mime_type: mime.to_string(),
is_url: false,
},
});
}
Some(ContentBlock::Image {
source: crate::llm::types::message::ImageSource {
data: url.clone(),
@@ -263,7 +263,9 @@ mod tests {
match ir {
Message::System { content } => {
assert_eq!(content.len(), 1);
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful"));
assert!(
matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful")
);
}
_ => panic!("expected System variant"),
}
@@ -385,10 +387,7 @@ mod tests {
assert_eq!(parts.len(), 1);
match &parts[0] {
OpenaiContentPart::Image { image_url, .. } => {
assert_eq!(
image_url.url,
"data:image/png;base64,BASE64DATA"
);
assert_eq!(image_url.url, "data:image/png;base64,BASE64DATA");
}
_ => panic!("expected Image part"),
}
+828 -43
View File
File diff suppressed because it is too large Load Diff
+11 -4
View File
@@ -8,9 +8,12 @@ use std::time::Duration;
///
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查 API key、减少上下文)。
#[derive(thiserror::Error, Debug)]
#[non_exhaustive]
pub enum LlmError {
/// API 认证失败(API key 无效、过期或权限不足)。
#[error("LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确")]
#[error(
"LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确"
)]
Authentication(String),
/// 请求被限流,可选地附带重试等待时间。可重试。
@@ -18,7 +21,9 @@ pub enum LlmError {
RateLimit { retry_after: Option<Duration> },
/// HTTP 请求失败(网络错误或非 2xx 状态码),包含状态码与响应体。
#[error("LLM 请求失败(HTTP {status}: {body}。请检查 Provider 端点地址(base_url)和网络连通性")]
#[error(
"LLM 请求失败(HTTP {status}: {body}。请检查 Provider 端点地址(base_url)和网络连通性"
)]
Request { status: u16, body: String },
/// 请求超时。可重试。
@@ -30,10 +35,12 @@ pub enum LlmError {
Stream(String),
/// 上下文长度超出模型窗口限制。
#[error("LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compactionllm::compact")]
#[error(
"LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compactionllm::compact"
)]
ContextLength { actual: u32, limit: u32 },
/// 其他未分类的 LLM 调用失败。
#[error("LLM 调用失败: {0}")]
Other(String),
}
}
+2 -3
View File
@@ -7,6 +7,7 @@ use crate::llm::types::request_v2::MessageRequest;
/// 生命周期钩子事件点。
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum HookEvent {
/// LLM 请求发起之前(可阻断)。
PreRequest,
@@ -130,9 +131,7 @@ impl Default for HookExecutor {
impl HookExecutor {
/// 创建一个空的执行器。
pub fn new() -> Self {
Self {
hooks: Vec::new(),
}
Self { hooks: Vec::new() }
}
/// 注册一个钩子到指定事件点。
+11 -9
View File
@@ -97,8 +97,7 @@ impl LlmProvider for MockProvider {
async fn chat_stream(
&self,
_request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
let response = self.pop()?;
// 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。
let id = response.id.clone();
@@ -206,10 +205,7 @@ mod tests {
#[tokio::test]
async fn chat_returns_queued_response() {
let provider = MockProvider::new(vec![text_response("hello")]);
let resp = provider
.chat(MessageRequest::default())
.await
.unwrap();
let resp = provider.chat(MessageRequest::default()).await.unwrap();
assert_eq!(resp.text(), "hello");
assert_eq!(provider.remaining(), 0);
}
@@ -231,7 +227,10 @@ mod tests {
#[tokio::test]
async fn chat_stream_emits_text_delta_sequence() {
let provider = MockProvider::new(vec![text_response("hi")]);
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
let mut stream = provider
.chat_stream(MessageRequest::default())
.await
.unwrap();
let mut seen_start = false;
let mut seen_block_start = false;
@@ -283,7 +282,10 @@ mod tests {
extra: Default::default(),
};
let provider = MockProvider::new(vec![response]);
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
let mut stream = provider
.chat_stream(MessageRequest::default())
.await
.unwrap();
let mut saw_tool_args = false;
let mut saw_tool_end = false;
@@ -302,4 +304,4 @@ mod tests {
assert!(saw_tool_args);
assert!(saw_tool_end);
}
}
}
+456 -21
View File
@@ -1,11 +1,14 @@
pub mod anthropic;
pub mod ollama;
pub mod openai;
pub mod openai_compat;
pub mod registry;
use std::pin::Pin;
use std::time::Duration;
use futures_core::Stream;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use crate::llm::error::LlmError;
@@ -18,6 +21,7 @@ use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
/// 当前协议数量(5 种以内)完全可控,enum 的编译期安全检查优于运行时的 `HashMap::get()`。
/// 未来如果扩展到 15+ 种以上,再改为注册表模式。
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ProviderType {
/// OpenAI Chat Completions API(兼容 DeepSeek / Qwen 等 `/chat/completions` 端点)。
OpenaiChat,
@@ -29,6 +33,8 @@ pub enum ProviderType {
DeepSeek,
/// Qwen / 阿里云百炼(OpenAI-compatible `/chat/completions`)。
Qwen,
/// Ollama 本地推理(OpenAI-compatible `/chat/completions`,默认 `http://localhost:11434/v1`)。
Ollama,
}
impl std::str::FromStr for ProviderType {
@@ -41,47 +47,211 @@ impl std::str::FromStr for ProviderType {
"anthropic" | "claude" => Ok(ProviderType::Anthropic),
"deepseek" => Ok(ProviderType::DeepSeek),
"qwen" | "dashscope" | "tongyi" => Ok(ProviderType::Qwen),
"ollama" => Ok(ProviderType::Ollama),
_ => Err(format!("未知的 Provider 类型: {s}")),
}
}
}
/// Provider 构造参数 —— 通用 base_url + api_key + model。
/// Provider 构造参数 —— 通用 base_url + api_key + model + timeout / retry 配置
#[derive(Debug, Clone)]
pub struct ProviderConfig {
/// API base URL(如 `https://api.openai.com/v1`)。为空时由 Provider 选择默认值。
pub base_url: String,
/// API key。Ollama 等本地 Provider 可为空。
pub api_key: String,
/// 模型名(如 `gpt-4o` / `claude-sonnet-4-20250514`)。
pub model: String,
/// 请求超时秒数(默认 30)。应用于 Provider 的 HTTP Client 级别。
pub timeout_secs: u64,
/// 最大重试次数(默认 3)。
///
/// 当前此字段仅由 `from_env()` 采集,**实际重试逻辑由 `CycleConfig.retry.max_retries` 控制**。
/// 此处保留字段以与 Roadmap §Phase 5 Step 5.1 对齐;未来 Phase 6+ 可统一合并到 `CycleConfig`。
pub max_retries: u32,
}
impl Default for ProviderConfig {
fn default() -> Self {
Self {
base_url: String::new(),
api_key: String::new(),
model: String::new(),
timeout_secs: 30,
max_retries: 3,
}
}
}
impl ProviderConfig {
/// 从环境变量构造 `ProviderConfig`。
///
/// 必填变量:
/// - `{prefix}_BASE_URL`
/// - `{prefix}_API_KEY`
/// - `{prefix}_MODEL`
///
/// 可选变量(有默认值):
/// - `{prefix}_TIMEOUT_SECS`(默认 30,解析失败回退 30 并 warn)
/// - `{prefix}_MAX_RETRIES`(默认 3,解析失败回退 3 并 warn)
pub fn from_env(prefix: &str) -> Result<Self, String> {
let base_url = std::env::var(format!("{prefix}_BASE_URL"))
.map_err(|_| format!("{prefix}_BASE_URL 环境变量未设置"))?;
let api_key = std::env::var(format!("{prefix}_API_KEY"))
.map_err(|_| format!("{prefix}_API_KEY 环境变量未设置"))?;
let model = std::env::var(format!("{prefix}_MODEL"))
.map_err(|_| format!("{prefix}_MODEL 环境变量未设置"))?;
let timeout_secs = match std::env::var(format!("{prefix}_TIMEOUT_SECS")) {
Ok(v) => v.parse().unwrap_or_else(|_| {
tracing::warn!("{prefix}_TIMEOUT_SECS='{v}' 解析失败,使用默认值 30");
30
}),
Err(_) => 30,
};
let max_retries = match std::env::var(format!("{prefix}_MAX_RETRIES")) {
Ok(v) => v.parse().unwrap_or_else(|_| {
tracing::warn!("{prefix}_MAX_RETRIES='{v}' 解析失败,使用默认值 3");
3
}),
Err(_) => 3,
};
// ponytail: max_retries 当前仅采集,不传入 Provider。
// 实际重试由 CycleConfig.retry.max_retries 控制。
if max_retries != 3 {
tracing::warn!(
"ProviderConfig.max_retries={} 已采集但当前未生效;\
CycleConfig.retry.max_retries ",
max_retries,
);
}
Ok(Self {
base_url,
api_key,
model,
timeout_secs,
max_retries,
})
}
}
/// 构造带 timeout 的 `reqwest::Client`OpenAI-compatible 共享)。
fn build_client_with_timeout(timeout_secs: u64) -> Result<Client, LlmError> {
Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.build()
.map_err(|e| LlmError::Other(format!("创建 HTTP 客户端失败: {e}")))
}
/// 构造带 Anthropic 默认 headers + timeout 的 `reqwest::Client`。
///
/// Anthropic 由于需要保留 `x-api-key` / `anthropic-version` 默认 headers
/// 与 OpenAI-compatible 共享的 `build_client_with_timeout` 不同。
fn build_anthropic_client(api_key: &str, timeout_secs: u64) -> Result<Client, LlmError> {
use reqwest::header::{HeaderMap, HeaderValue};
let key_header = HeaderValue::from_str(api_key)
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
let version_header = HeaderValue::from_static("2023-06-01");
Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
headers.insert("anthropic-version", version_header);
headers
})
.build()
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))
}
/// Provider 工厂 —— exhaustive match 在编译期保证新 Provider 被注册。
///
/// `config.timeout_secs` 注入到 Provider 的 HTTP Client 超时配置。
/// 每个分支通过 `from_parts` (pub(crate)) 一次性构造,无冗余 client 创建。
pub fn create_provider(
provider_type: ProviderType,
config: ProviderConfig,
) -> Result<Box<dyn LlmProvider>, LlmError> {
match provider_type {
ProviderType::OpenaiChat => Ok(Box::new(openai::OpenaiChatProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::OpenaiChat => {
let client = build_client_with_timeout(config.timeout_secs)?;
Ok(Box::new(openai::OpenaiChatProvider(
openai::GenericOpenaiProvider::from_parts(
config.base_url,
config.api_key,
config.model,
"openai",
client,
Vec::new(),
config.timeout_secs,
),
)))
}
ProviderType::OpenaiResponse => Err(LlmError::Other(
"OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(),
)),
ProviderType::Anthropic => Ok(Box::new(anthropic::AnthropicProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::DeepSeek => Ok(Box::new(openai_compat::DeepSeekProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::Qwen => Ok(Box::new(openai_compat::QwenProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::Anthropic => {
let client = build_anthropic_client(&config.api_key, config.timeout_secs)?;
Ok(Box::new(anthropic::AnthropicProvider::from_parts(
config.base_url,
config.api_key,
config.model,
client,
config.timeout_secs,
)))
}
ProviderType::DeepSeek => {
let client = build_client_with_timeout(config.timeout_secs)?;
Ok(Box::new(openai_compat::DeepSeekProvider(
openai::GenericOpenaiProvider::from_parts(
config.base_url,
config.api_key,
config.model,
"deepseek",
client,
Vec::new(),
config.timeout_secs,
),
)))
}
ProviderType::Qwen => {
let client = build_client_with_timeout(config.timeout_secs)?;
Ok(Box::new(openai_compat::QwenProvider(
openai::GenericOpenaiProvider::from_parts(
config.base_url,
config.api_key,
config.model,
"qwen",
client,
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
config.timeout_secs,
),
)))
}
ProviderType::Ollama => {
let client = build_client_with_timeout(config.timeout_secs)?;
// ponytail: Ollama 默认 base_url 由 OllamaProvider 构造处理 —— 但 from_parts 不走
// OllamaProvider::new 的默认 URL 回退。这里保留 base_url(可能为空 → http://localhost:11434/v1)。
let base_url = if config.base_url.is_empty() {
"http://localhost:11434/v1".to_string()
} else {
config.base_url
};
Ok(Box::new(ollama::OllamaProvider(
openai::GenericOpenaiProvider::from_parts(
base_url,
config.api_key,
config.model,
"ollama",
client,
Vec::new(),
config.timeout_secs,
),
)))
}
}
}
@@ -142,3 +312,268 @@ pub trait LlmProvider: Send + Sync {
/// 返回 Provider 静态能力描述。
fn capabilities(&self) -> ProviderCapabilities;
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn provider_config_default_values() {
let config = ProviderConfig::default();
assert_eq!(config.timeout_secs, 30);
assert_eq!(config.max_retries, 3);
assert_eq!(config.base_url, "");
assert_eq!(config.api_key, "");
assert_eq!(config.model, "");
}
#[test]
fn provider_config_from_env_requires_all_three() {
// 使用 temp_env 移除所有相关变量,避免外部环境意外设置导致测试 flaky
temp_env::with_vars(
[
("TEST_PROVIDER_MISSING_BASE_URL", None::<&str>),
("TEST_PROVIDER_MISSING_API_KEY", None::<&str>),
("TEST_PROVIDER_MISSING_MODEL", None::<&str>),
("TEST_PROVIDER_MISSING_TIMEOUT_SECS", None::<&str>),
("TEST_PROVIDER_MISSING_MAX_RETRIES", None::<&str>),
],
|| {
let result = ProviderConfig::from_env("TEST_PROVIDER_MISSING");
assert!(result.is_err());
let msg = result.unwrap_err();
assert!(
msg.contains("TEST_PROVIDER_MISSING_BASE_URL"),
"error should mention missing var, got: {msg}"
);
},
);
}
#[test]
fn provider_config_from_env_uses_defaults_when_only_required_set() {
temp_env::with_vars(
[
("TEST_PROVIDER_BASE_URL", Some("http://localhost:11434/v1")),
("TEST_PROVIDER_API_KEY", Some("")),
("TEST_PROVIDER_MODEL", Some("llama3")),
],
|| {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
assert_eq!(config.base_url, "http://localhost:11434/v1");
assert_eq!(config.api_key, "");
assert_eq!(config.model, "llama3");
assert_eq!(config.timeout_secs, 30);
assert_eq!(config.max_retries, 3);
},
);
}
#[test]
fn provider_config_from_env_reads_custom_values() {
temp_env::with_vars(
[
("TEST_PROVIDER_BASE_URL", Some("http://x")),
("TEST_PROVIDER_API_KEY", Some("k")),
("TEST_PROVIDER_MODEL", Some("m")),
("TEST_PROVIDER_TIMEOUT_SECS", Some("60")),
("TEST_PROVIDER_MAX_RETRIES", Some("5")),
],
|| {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
assert_eq!(config.timeout_secs, 60);
assert_eq!(config.max_retries, 5);
},
);
}
#[test]
fn provider_config_from_env_falls_back_on_invalid_numbers() {
temp_env::with_vars(
[
("TEST_PROVIDER_BASE_URL", Some("http://x")),
("TEST_PROVIDER_API_KEY", Some("k")),
("TEST_PROVIDER_MODEL", Some("m")),
("TEST_PROVIDER_TIMEOUT_SECS", Some("not-a-number")),
("TEST_PROVIDER_MAX_RETRIES", Some("also-bad")),
],
|| {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
// 解析失败回退默认值
assert_eq!(config.timeout_secs, 30);
assert_eq!(config.max_retries, 3);
},
);
}
/// Timeout 传导集成测试:构造 `ProviderConfig` timeout=1s
/// `create_provider` 注入 1s 超时 client,请求一个故意延迟 3s 的 mock server
/// 验证返回 `LlmError::Timeout { duration: 1s }`。
#[tokio::test]
async fn create_provider_injects_timeout_into_openai_chat() {
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
// 故意延迟 3s 触发超时
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(3))
.set_body_json(json!({
"id": "x",
"object": "chat.completion",
"created": 0,
"model": "gpt-4o",
"choices": [],
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
})),
)
.mount(&server)
.await;
let provider = create_provider(
ProviderType::OpenaiChat,
ProviderConfig {
base_url: server.uri(),
api_key: "sk-test".into(),
model: "gpt-4o".into(),
timeout_secs: 1,
max_retries: 3,
},
)
.unwrap();
let err = provider
.chat(crate::llm::types::request_v2::MessageRequest {
model: "gpt-4o".into(),
messages: vec![],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Timeout { duration } => {
assert_eq!(duration, Duration::from_secs(1));
}
other => panic!("expected Timeout, got {other:?}"),
}
}
/// Timeout 传导验证:`create_provider` 生成的 DeepSeek Provider 也带 1s 超时,
/// 错误消息中的 duration 与 timeout_secs 一致(而非硬编码 120s)。
#[tokio::test]
async fn create_provider_injects_timeout_into_deepseek() {
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(3))
.set_body_json(json!({
"id": "x",
"object": "chat.completion",
"created": 0,
"model": "deepseek-chat",
"choices": [],
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
})),
)
.mount(&server)
.await;
let provider = create_provider(
ProviderType::DeepSeek,
ProviderConfig {
base_url: server.uri(),
api_key: "sk-test".into(),
model: "deepseek-chat".into(),
timeout_secs: 1,
max_retries: 3,
},
)
.unwrap();
let err = provider
.chat(crate::llm::types::request_v2::MessageRequest {
model: "deepseek-chat".into(),
messages: vec![],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Timeout { duration } => {
assert_eq!(duration, Duration::from_secs(1));
}
other => panic!("expected Timeout, got {other:?}"),
}
}
/// Timeout 传导验证:`create_provider` 生成的 Anthropic Provider 通过 `with_timeout`
/// 注入 1s 超时。
///
/// 与 OpenAI-compatible 路径不同,Anthropic 走 `AnthropicProvider::with_timeout()`
/// 重建底层 client(保留 default_headers),独立于 OpenAI-compatible 的 `build_client_with_timeout`。
/// 单独覆盖此路径以验证 `with_timeout` 不会因服务端延迟而返回硬编码 120s 的超时错误。
#[tokio::test]
async fn create_provider_injects_timeout_into_anthropic() {
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
// Anthropic Messages API 端点:`POST /v1/messages`
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(3))
.set_body_json(json!({
"id": "msg_timeout_test",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "ok"}],
"model": "claude-sonnet-4-20250514",
"stop_reason": "end_turn",
"usage": {"input_tokens": 1, "output_tokens": 1}
})),
)
.mount(&server)
.await;
let provider = create_provider(
ProviderType::Anthropic,
ProviderConfig {
base_url: server.uri(),
api_key: "sk-ant-test".into(),
model: "claude-sonnet-4-20250514".into(),
timeout_secs: 1,
max_retries: 3,
},
)
.unwrap();
let err = provider
.chat(crate::llm::types::request_v2::MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Timeout { duration } => {
assert_eq!(duration, Duration::from_secs(1));
}
other => panic!("expected Timeout, got {other:?}"),
}
}
}
+132 -55
View File
@@ -12,10 +12,10 @@ use async_trait::async_trait;
use bytes::Bytes;
use futures_core::Stream;
use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderValue};
use reqwest::Client;
use reqwest::header::{HeaderMap, HeaderValue};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use serde_json::{Value, json};
use tracing::{debug, error, info, warn};
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
@@ -39,16 +39,20 @@ pub struct AnthropicProvider {
#[allow(dead_code)]
api_key: String,
model: String,
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter
/// 因此单独存储以便错误消息与配置保持一致。
timeout_secs: u64,
}
impl AnthropicProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
let key_header = HeaderValue::from_str(&api_key)
.expect("Anthropic API key 包含无效的 HTTP 头部字符");
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let key_header =
HeaderValue::from_str(&api_key).expect("Anthropic API key 包含无效的 HTTP 头部字符");
let version_header = HeaderValue::from_static("2023-06-01");
let http_client = Client::builder()
.timeout(Duration::from_secs(120))
.timeout(Duration::from_secs(timeout_secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
@@ -67,14 +71,80 @@ impl AnthropicProvider {
},
api_key,
model,
timeout_secs,
}
}
/// ⚠️ 替换 HTTP Client**丢弃** `new()` 中设置的默认 headers`x-api-key` / `anthropic-version`)。
///
/// 调用此方法后,所有 Anthropic API 请求将以**无认证头**发送出去,预期会 401/403 失败。
/// 推荐改用 [`Self::with_timeout`],它会重建 client 并保留默认 headers。
///
/// 此方法仍保留以兼容调用方自定义 client 但不需要默认 headers 的极端场景。
#[deprecated(
since = "0.2.0",
note = "此方法会丢弃默认 headersx-api-key / anthropic-version),改为使用 `with_timeout` 或带 headers 的 `Client::builder()`"
)]
pub fn with_client(mut self, client: Client) -> Self {
self.http_client = client;
self
}
/// 替换 HTTP Client 的超时配置(重建底层 client,保留默认 headers)。
///
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后通过 `with_client` 注入的 Client
/// 将被替换。headers 构造逻辑与 `new()` 中的保持一致(`x-api-key` / `anthropic-version`)。
///
/// ponytail: 同值调用短路。当 `secs == self.timeout_secs` 时跳过 client 重建,
/// 避免 `create_provider` 路径 `new(timeout).with_timeout(timeout)` 的双重构造。
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
if secs == self.timeout_secs {
return Ok(self);
}
// ponytail: 重建 http_client 时保留已有默认 headersx-api-key / anthropic-version)。
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
let key_header = HeaderValue::from_str(&self.api_key)
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
let version_header = HeaderValue::from_static("2023-06-01");
self.http_client = Client::builder()
.timeout(Duration::from_secs(secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
headers.insert("anthropic-version", version_header);
headers
})
.build()
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
self.timeout_secs = secs;
Ok(self)
}
/// 一次性构造 —— `create_provider` 路径专用,避免 `new(...)` + `with_timeout(...)` 的双重 client 构造。
///
/// 调用方负责预先构造好符合 Anthropic 协议要求的 `http_client`(带正确的 `x-api-key` /
/// `anthropic-version` 默认 headers + 指定 timeout)。
pub(crate) fn from_parts(
base_url: String,
api_key: String,
model: String,
http_client: Client,
timeout_secs: u64,
) -> Self {
Self {
http_client,
base_url: if base_url.is_empty() {
"https://api.anthropic.com".to_string()
} else {
base_url
},
api_key,
model,
timeout_secs,
}
}
fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
}
@@ -104,12 +174,14 @@ impl AnthropicProvider {
Message::User { content } => {
api_messages.push(AnthropicMessage::user(content));
}
Message::UserImage { data, mime_type, detail } => {
Message::UserImage {
data,
mime_type,
detail,
} => {
// Anthropic image format: {type: "image", source: {type: "base64", media_type, data}}
let source = if data.starts_with("http://") || data.starts_with("https://") {
AnthropicImageSource::Url {
url: data.clone(),
}
AnthropicImageSource::Url { url: data.clone() }
} else {
AnthropicImageSource::Base64 {
media_type: mime_type.clone(),
@@ -194,7 +266,7 @@ impl AnthropicProvider {
.json(&body)
.send()
.await
.map_err(Self::map_reqwest_error)?;
.map_err(|e| self.map_reqwest_error(e))?;
let status = response.status();
if !status.is_success() {
@@ -215,8 +287,7 @@ impl AnthropicProvider {
async fn chat_stream_inner(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
let mut body = self.build_request_body(request)?;
body.stream = Some(true);
@@ -230,16 +301,16 @@ impl AnthropicProvider {
.json(&body)
.send()
.await
.map_err(Self::map_reqwest_error)?;
.map_err(|e| self.map_reqwest_error(e))?;
let status = response.status();
if !status.is_success() {
return Err(Self::handle_error_response(response).await);
}
let byte_stream = response.bytes_stream().map(|r| {
r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}")))
});
let byte_stream = response
.bytes_stream()
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}"))));
let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> =
Box::pin(byte_stream);
@@ -247,10 +318,10 @@ impl AnthropicProvider {
Ok(Box::pin(AnthropicSseStream::new(byte_stream)))
}
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
if e.is_timeout() {
LlmError::Timeout {
duration: Duration::from_secs(120),
duration: Duration::from_secs(self.timeout_secs),
}
} else if e.is_connect() {
LlmError::Other(format!("连接失败: {e}"))
@@ -291,13 +362,12 @@ impl AnthropicProvider {
blocks.push(ContentBlock::Text { text });
}
AnthropicContentBlockResp::ToolUse { id, name, input } => {
blocks.push(ContentBlock::ToolUse {
id,
name,
input,
});
blocks.push(ContentBlock::ToolUse { id, name, input });
}
AnthropicContentBlockResp::Thinking { thinking, signature } => {
AnthropicContentBlockResp::Thinking {
thinking,
signature,
} => {
blocks.push(ContentBlock::Thinking {
text: thinking,
signature,
@@ -336,8 +406,7 @@ impl LlmProvider for AnthropicProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.chat_stream_inner(request).await
}
@@ -418,7 +487,9 @@ impl AnthropicMessage {
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentPart {
Text { text: String },
Text {
text: String,
},
Image {
source: AnthropicImageSource,
},
@@ -437,13 +508,8 @@ enum AnthropicContentPart {
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicImageSource {
Base64 {
media_type: String,
data: String,
},
Url {
url: String,
},
Base64 { media_type: String, data: String },
Url { url: String },
}
fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> {
@@ -523,9 +589,18 @@ struct AnthropicUsage {
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentBlockResp {
Text { text: String },
ToolUse { id: String, name: String, input: Value },
Thinking { thinking: String, signature: Option<String> },
Text {
text: String,
},
ToolUse {
id: String,
name: String,
input: Value,
},
Thinking {
thinking: String,
signature: Option<String>,
},
}
// =============================================================================
@@ -662,18 +737,13 @@ impl AnthropicSseStream {
// 先把所有字段提前,避免 match 中 part-move
let block_type = match &content_block {
AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text,
AnthropicContentBlockStart::ToolUse { id, name } => {
ContentBlockType::ToolUse {
id: id.clone(),
name: name.clone(),
}
}
AnthropicContentBlockStart::ToolUse { id, name } => ContentBlockType::ToolUse {
id: id.clone(),
name: name.clone(),
},
AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking,
};
events.push(StreamEvent::ContentBlockStart {
index,
block_type,
});
events.push(StreamEvent::ContentBlockStart { index, block_type });
let builder = match content_block {
AnthropicContentBlockStart::Text { text } => {
crate::llm::types::response_v2::ContentBlockBuilder::Text(text)
@@ -742,7 +812,9 @@ impl AnthropicSseStream {
completion_tokens_details: None,
prompt_tokens_details: None,
};
events.push(StreamEvent::CostUpdate { usage: partial_usage });
events.push(StreamEvent::CostUpdate {
usage: partial_usage,
});
}
}
AnthropicSseEvent::MessageStop => {
@@ -753,7 +825,9 @@ impl AnthropicSseStream {
self.saw_terminal = true;
match self.partial.clone().finalize() {
Ok(full) => {
events.push(StreamEvent::MessageComplete { full_response: full });
events.push(StreamEvent::MessageComplete {
full_response: full,
});
}
Err(e) => {
events.push(StreamEvent::Error {
@@ -780,10 +854,7 @@ fn _unused_marker() {}
impl Stream for AnthropicSseStream {
type Item = Result<StreamEvent, LlmError>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
if let Some(data) = self.next_event_line() {
let mut events = self.handle_event_json(&data);
@@ -835,7 +906,12 @@ mod tests {
fn make_provider(base_url: String) -> AnthropicProvider {
// 跳过默认 header 注入:测试用自定义 base_url 直接 mock
AnthropicProvider::new(base_url, "sk-ant-test".into(), "claude-sonnet-4-20250514".into())
AnthropicProvider::new(
base_url,
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
30,
)
}
#[tokio::test]
@@ -979,6 +1055,7 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
"http://x".into(),
"k".into(),
"claude-sonnet-4-20250514".into(),
30,
)
.capabilities();
assert_eq!(caps.provider_name, "anthropic");
+72
View File
@@ -0,0 +1,72 @@
//! Ollama Provider —— OpenAI-compatible 协议的 newtype 包装,零 API key。
//!
//! 默认 base_url = `http://localhost:11434/v1`,空 api_key 也可工作。
//! 实现方式同 `DeepSeekProvider` / `QwenProvider`,共享 `GenericOpenaiProvider`
//! 的 HTTP/SSE/转换逻辑,仅配置不同。
use std::pin::Pin;
use async_trait::async_trait;
use futures_core::Stream;
use reqwest::Client;
use super::openai::GenericOpenaiProvider;
use super::{LlmProvider, ProviderCapabilities};
use crate::llm::error::LlmError;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
/// Ollama 本地 Provider —— OpenAI-compatible 协议的 newtype 包装。
///
/// Ollama 在 `localhost:11434` 暴露与 OpenAI 兼容的 `/v1/chat/completions`
/// 接口,因此完全复用 `GenericOpenaiProvider` 的实现。允许空 `api_key`。
pub struct OllamaProvider(pub GenericOpenaiProvider);
impl OllamaProvider {
/// 构造 Ollama Provider。
///
/// - `base_url` 为空时使用默认 `http://localhost:11434/v1`
/// - `api_key` 可为空字符串(Ollama 不校验)
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let url = if base_url.is_empty() {
"http://localhost:11434/v1".to_string()
} else {
base_url
};
Self(GenericOpenaiProvider::new_with_name(
url,
api_key,
model,
"ollama",
timeout_secs,
))
}
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
///
/// 与 `OpenaiChatProvider::with_client`、`DeepSeekProvider::with_client`、
/// `QwenProvider::with_client` 签名一致。
pub fn with_client(self, client: Client) -> Self {
Self(self.0.with_client(client))
}
}
#[async_trait]
impl LlmProvider for OllamaProvider {
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
self.0.chat(request).await
}
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
fn capabilities(&self) -> ProviderCapabilities {
let mut caps = self.0.capabilities();
caps.provider_name = "ollama";
caps
}
}
+93 -56
View File
@@ -52,31 +52,63 @@ pub struct GenericOpenaiProvider {
model: String,
provider_name: &'static str,
extra_headers: Vec<(String, String)>,
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter
/// 因此单独存储以便错误消息与配置保持一致。
timeout_secs: u64,
}
impl GenericOpenaiProvider {
/// 基础构造
pub fn new_with_name(
/// 一次性构造 —— `create_provider` 路径专用,避免 `new_with_name` + `with_client` 的双重 client 构造。
///
/// 调用方负责预先构造好带正确 timeout 的 `http_client`。`extra_headers` 与 `timeout_secs`
/// 一并设置字段,避免后续修改。
pub(crate) fn from_parts(
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
http_client: Client,
extra_headers: Vec<(String, String)>,
timeout_secs: u64,
) -> Self {
let http_client = Client::builder()
.timeout(Duration::from_secs(120))
.build()
.expect("创建 HTTP 客户端失败");
Self {
http_client,
base_url,
api_key,
model,
provider_name,
extra_headers: Vec::new(),
extra_headers,
timeout_secs,
}
}
/// 基础构造器。
///
/// `timeout_secs` 应用于 `reqwest::Client` 的请求超时配置。
/// 应由 `ProviderConfig::timeout_secs` 传入(调用方如不知道,可传 30)。
pub fn new_with_name(
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
timeout_secs: u64,
) -> Self {
let http_client = Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.build()
.expect("创建 HTTP 客户端失败");
Self::from_parts(
base_url,
api_key,
model,
provider_name,
http_client,
Vec::new(),
timeout_secs,
)
}
/// 带额外请求头的构造器(如 Qwen 需要 SSE 启用头)。
pub fn new_with_name_and_headers(
base_url: String,
@@ -84,10 +116,21 @@ impl GenericOpenaiProvider {
model: String,
provider_name: &'static str,
extra_headers: Vec<(String, String)>,
timeout_secs: u64,
) -> Self {
let mut base = Self::new_with_name(base_url, api_key, model, provider_name);
base.extra_headers = extra_headers;
base
let http_client = Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.build()
.expect("创建 HTTP 客户端失败");
Self::from_parts(
base_url,
api_key,
model,
provider_name,
http_client,
extra_headers,
timeout_secs,
)
}
pub fn with_client(mut self, client: Client) -> Self {
@@ -119,10 +162,10 @@ impl GenericOpenaiProvider {
Ok(builder.json(body))
}
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
if e.is_timeout() {
LlmError::Timeout {
duration: Duration::from_secs(120),
duration: Duration::from_secs(self.timeout_secs),
}
} else if e.is_connect() {
LlmError::Other(format!("连接失败: {}", e))
@@ -135,9 +178,7 @@ impl GenericOpenaiProvider {
///
/// ponytail: Qwen 等部分 OpenAI-compatible 提供方可能返回非标准 error body
/// (无法解析为 JSON),此处直接用 status code + 原始 body 兜底。
async fn handle_error_response(
response: reqwest::Response,
) -> LlmError {
async fn handle_error_response(response: reqwest::Response) -> LlmError {
let status = response.status().as_u16();
let body = response.text().await.unwrap_or_default();
@@ -148,10 +189,7 @@ impl GenericOpenaiProvider {
// 与 OpenAI 完全一致;DeepSeek/Qwen 通常遵循。
LlmError::RateLimit { retry_after: None }
}
_ if status >= 500 => LlmError::Request {
status,
body,
},
_ if status >= 500 => LlmError::Request { status, body },
_ if status == 400 && body.contains("context_length_exceeded") => {
LlmError::ContextLength {
actual: 0,
@@ -188,7 +226,7 @@ impl GenericOpenaiProvider {
Some(
tool_defs
.into_iter()
.map(|t| OpenaiTool::Function { function: t })
.map(|t| OpenaiTool::Function { function: t.into() })
.collect(),
)
};
@@ -200,7 +238,9 @@ impl GenericOpenaiProvider {
stop_sequences[0].clone(),
))
} else {
Some(crate::llm::types::shared::StopSequence::Multiple(stop_sequences))
Some(crate::llm::types::shared::StopSequence::Multiple(
stop_sequences,
))
};
let frequency_penalty = request.get_extra_opt("frequency_penalty");
@@ -245,9 +285,7 @@ impl GenericOpenaiProvider {
let stop_reason = match choice.finish_reason {
Some(FinishReason::Stop) => StopReason::Stop,
Some(FinishReason::Length) => StopReason::Length,
Some(FinishReason::ToolCalls) | Some(FinishReason::FunctionCall) => {
StopReason::ToolUse
}
Some(FinishReason::ToolCalls) | Some(FinishReason::FunctionCall) => StopReason::ToolUse,
Some(FinishReason::ContentFilter) => StopReason::ContentFilter,
Some(FinishReason::Other) | None => StopReason::Stop,
};
@@ -263,7 +301,10 @@ impl GenericOpenaiProvider {
}
/// 非流式 `chat()` 入口。
pub async fn chat_blocking(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
pub async fn chat_blocking(
&self,
request: MessageRequest,
) -> Result<MessageResponse, LlmError> {
let req = self.convert_request(request)?;
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
@@ -280,7 +321,7 @@ impl GenericOpenaiProvider {
.await
.map_err(|e| {
error!(error = %e, "请求失败");
Self::map_reqwest_error(e)
self.map_reqwest_error(e)
})?;
let status = response.status();
@@ -303,8 +344,7 @@ impl GenericOpenaiProvider {
pub async fn chat_stream_inner(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
let mut req = self.convert_request(request)?;
req.stream = Some(true);
req.stream_options = Some(StreamOptions {
@@ -322,7 +362,7 @@ impl GenericOpenaiProvider {
.await
.map_err(|e| {
error!(error = %e, "流式请求失败");
Self::map_reqwest_error(e)
self.map_reqwest_error(e)
})?;
let status = response.status();
@@ -331,9 +371,9 @@ impl GenericOpenaiProvider {
}
let byte_stream: std::pin::Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> = {
let s = response.bytes_stream().map(|r| {
r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e)))
});
let s = response
.bytes_stream()
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e))));
Box::pin(s)
};
@@ -368,8 +408,7 @@ impl LlmProvider for GenericOpenaiProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.chat_stream_inner(request).await
}
@@ -389,12 +428,13 @@ impl LlmProvider for GenericOpenaiProvider {
pub struct OpenaiChatProvider(pub GenericOpenaiProvider);
impl OpenaiChatProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
Self(GenericOpenaiProvider::new_with_name(
base_url,
api_key,
model,
"openai",
timeout_secs,
))
}
@@ -412,8 +452,7 @@ impl LlmProvider for OpenaiChatProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
@@ -436,7 +475,10 @@ enum BlockState {
/// 正在累积 text block。
InText { block_index: u32 },
/// 正在累积 tool_call block。
InTool { block_index: u32, tool_call_index: u32 },
InTool {
block_index: u32,
tool_call_index: u32,
},
/// 正在累积 refusal block。
InRefusal { block_index: u32 },
}
@@ -456,9 +498,7 @@ pub struct ChunkToEventStream {
}
impl ChunkToEventStream {
fn new(
chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>,
) -> Self {
fn new(chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>) -> Self {
Self {
chunks,
buffer: String::new(),
@@ -500,12 +540,7 @@ impl ChunkToEventStream {
let mut events = Vec::new();
// 元信息:MessageStart(仅在第一次见到 role=assistant 时)。
if self.partial.id.is_none()
&& chunk
.choices
.iter()
.any(|c| c.delta.role.is_some())
{
if self.partial.id.is_none() && chunk.choices.iter().any(|c| c.delta.role.is_some()) {
events.push(StreamEvent::MessageStart {
id: chunk.id.clone(),
model: chunk.model.clone(),
@@ -650,9 +685,7 @@ impl ChunkToEventStream {
self.partial.stop_reason = Some(match fr {
FinishReason::Stop => StopReason::Stop,
FinishReason::Length => StopReason::Length,
FinishReason::ToolCalls | FinishReason::FunctionCall => {
StopReason::ToolUse
}
FinishReason::ToolCalls | FinishReason::FunctionCall => StopReason::ToolUse,
FinishReason::ContentFilter => StopReason::ContentFilter,
FinishReason::Other => StopReason::Other,
});
@@ -692,7 +725,9 @@ impl ChunkToEventStream {
match self.partial.clone().finalize() {
Ok(full) => {
events.push(StreamEvent::MessageComplete { full_response: full });
events.push(StreamEvent::MessageComplete {
full_response: full,
});
}
Err(e) => {
events.push(StreamEvent::Error {
@@ -707,10 +742,7 @@ impl ChunkToEventStream {
impl Stream for ChunkToEventStream {
type Item = Result<StreamEvent, LlmError>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
// 先尝试从 buffer 取一行处理
if let Some(line) = self.next_line() {
@@ -814,6 +846,7 @@ mod tests {
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let response = provider
.chat_blocking(MessageRequest {
@@ -842,6 +875,7 @@ mod tests {
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let err = provider
.chat_blocking(MessageRequest {
@@ -868,6 +902,7 @@ mod tests {
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let err = provider
.chat_blocking(MessageRequest {
@@ -906,6 +941,7 @@ data: [DONE]\n\n";
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let mut stream = provider
.chat_stream_inner(MessageRequest {
@@ -974,6 +1010,7 @@ data: [DONE]\n\n";
"k".into(),
"gpt-4o".into(),
"openai",
30,
);
let ir = provider.convert_response(resp).unwrap();
assert_eq!(ir.stop_reason, StopReason::ToolUse);
+32 -39
View File
@@ -15,12 +15,12 @@ use std::pin::Pin;
use async_trait::async_trait;
use futures_core::Stream;
use super::openai::GenericOpenaiProvider;
use super::ProviderCapabilities;
use super::openai::GenericOpenaiProvider;
use crate::llm::error::LlmError;
use crate::llm::provider::LlmProvider;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
use crate::llm::provider::LlmProvider;
// =============================================================================
// DeepSeek
@@ -29,7 +29,7 @@ use crate::llm::provider::LlmProvider;
pub struct DeepSeekProvider(pub GenericOpenaiProvider);
impl DeepSeekProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let url = if base_url.is_empty() {
"https://api.deepseek.com".to_string()
} else {
@@ -40,26 +40,27 @@ impl DeepSeekProvider {
api_key,
model,
"deepseek",
timeout_secs,
))
}
}
impl DeepSeekProvider {
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
pub fn with_client(self, client: reqwest::Client) -> Self {
Self(self.0.with_client(client))
}
/// 测试中(带 mock_client)使用的构造器。
///
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
/// 无超时),不影响行为。
pub fn new_with_client(
base_url: String,
api_key: String,
model: String,
client: reqwest::Client,
) -> Self {
let url = if base_url.is_empty() {
"https://api.deepseek.com".to_string()
} else {
base_url
};
let mut inner = GenericOpenaiProvider::new_with_name(url, api_key, model, "deepseek");
inner.http_client = client;
Self(inner)
Self::new(base_url, api_key, model, 30).with_client(client)
}
}
@@ -72,8 +73,7 @@ impl LlmProvider for DeepSeekProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
@@ -91,7 +91,7 @@ impl LlmProvider for DeepSeekProvider {
pub struct QwenProvider(pub GenericOpenaiProvider);
impl QwenProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let url = if base_url.is_empty() {
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
} else {
@@ -104,31 +104,28 @@ impl QwenProvider {
model,
"qwen",
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
timeout_secs,
);
Self(inner)
}
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
pub fn with_client(self, client: reqwest::Client) -> Self {
Self(self.0.with_client(client))
}
/// 测试构造器。
///
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
/// 无超时),不影响行为。
pub fn new_with_client(
base_url: String,
api_key: String,
model: String,
client: reqwest::Client,
) -> Self {
let url = if base_url.is_empty() {
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
} else {
base_url
};
let mut inner = GenericOpenaiProvider::new_with_name_and_headers(
url,
api_key,
model,
"qwen",
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
);
inner.http_client = client;
Self(inner)
Self::new(base_url, api_key, model, 30).with_client(client)
}
}
@@ -141,8 +138,7 @@ impl LlmProvider for QwenProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
@@ -156,8 +152,8 @@ impl LlmProvider for QwenProvider {
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::message::Message as IrMessage;
use crate::llm::types::request_v2::MessageRequest;
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
@@ -182,11 +178,8 @@ mod tests {
.mount(&server)
.await;
let provider = DeepSeekProvider::new(
server.uri(),
"sk-test".into(),
"deepseek-chat".into(),
);
let provider =
DeepSeekProvider::new(server.uri(), "sk-test".into(), "deepseek-chat".into(), 30);
let response = provider
.chat(MessageRequest {
model: "deepseek-chat".into(),
@@ -219,7 +212,7 @@ mod tests {
.mount(&server)
.await;
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into());
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into(), 30);
let response = provider
.chat(MessageRequest {
model: "qwen-plus".into(),
+2 -4
View File
@@ -3,7 +3,7 @@
use std::collections::HashMap;
use crate::llm::error::LlmError;
use crate::llm::provider::{create_provider, LlmProvider, ProviderConfig, ProviderType};
use crate::llm::provider::{LlmProvider, ProviderConfig, ProviderType, create_provider};
/// Provider 注册表 —— 管理多个 LLM Provider 实例。
///
@@ -61,8 +61,6 @@ impl ProviderRegistry {
/// 获取默认 Provider。
pub fn get_default(&self) -> Option<&dyn LlmProvider> {
self.default_name
.as_ref()
.and_then(|name| self.get(name))
self.default_name.as_ref().and_then(|name| self.get(name))
}
}
+7 -8
View File
@@ -14,8 +14,8 @@ use std::pin::Pin;
use std::task::{Context, Poll};
use futures_core::stream::Stream;
use futures_util::future::poll_fn;
use futures_util::FutureExt;
use futures_util::future::poll_fn;
use serde_json::Value;
use crate::llm::error::LlmError;
@@ -95,9 +95,7 @@ impl Stream for ChunkToLegacyEventStream {
}
if let Some(usage) = &chunk.usage {
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate {
usage: *usage,
}));
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate { usage: *usage }));
}
Poll::Ready(None)
@@ -143,9 +141,7 @@ fn empty_message_response() -> MessageResponse {
MessageResponse {
id: String::new(),
model: String::new(),
message: Message::Assistant {
content: vec![],
},
message: Message::Assistant { content: vec![] },
usage: Usage::default(),
stop_reason: StopReason::Stop,
extra: HashMap::new(),
@@ -172,7 +168,10 @@ fn map_legacy_to_ir(legacy: LegacyStreamEvent) -> StreamEvent {
LegacyStreamEvent::AssistantTextDelta { text } => StreamEvent::TextDelta { text },
LegacyStreamEvent::ToolExecutionStarted { input, .. } => {
let arguments = serde_json::to_string(&input).unwrap_or_default();
StreamEvent::ToolCallArgumentsDelta { index: 0, arguments }
StreamEvent::ToolCallArgumentsDelta {
index: 0,
arguments,
}
}
LegacyStreamEvent::ToolExecutionCompleted { .. } => {
// 旧 ToolExecutionCompleted 不在 IR 流协议中——工具执行是消费方职责。
+9 -19
View File
@@ -20,15 +20,12 @@ use crate::llm::types::shared::ImageDetail;
/// 消费方 match 可直接区分文本和图片输入。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Message {
/// 系统提示(User & Assistant 之外的引导指令)。
System {
content: Vec<ContentBlock>,
},
System { content: Vec<ContentBlock> },
/// 用户输入。
User {
content: Vec<ContentBlock>,
},
User { content: Vec<ContentBlock> },
/// 用户的图片输入(快捷构造,免去构造 ContentBlock 的 boilerplate)。
UserImage {
data: String,
@@ -36,9 +33,7 @@ pub enum Message {
detail: ImageDetail,
},
/// Assistant 回复内容块(可能包含 text、thinking、tool_use 等多种 block 的混合)。
Assistant {
content: Vec<ContentBlock>,
},
Assistant { content: Vec<ContentBlock> },
/// 工具调用结果。
ToolResult {
tool_call_id: String,
@@ -103,6 +98,7 @@ impl Message {
/// block 的逃生舱。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ContentBlock {
/// 纯文本。
Text { text: String },
@@ -130,10 +126,7 @@ pub enum ContentBlock {
signature: Option<String>,
},
/// 逃生舱:Provider 特定 block 透传(OpenAI Response 内置工具等)。
Extension {
kind: String,
data: Value,
},
Extension { kind: String, data: Value },
}
/// 内容块类型标签 —— 用于 `StreamEvent::ContentBlockStart.block_type`。
@@ -141,6 +134,7 @@ pub enum ContentBlock {
/// 用途:在流式场景中,Provider 先下发 block 类型,再下发 block 内容增量。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ContentBlockType {
/// 文本块。
Text,
@@ -349,9 +343,7 @@ mod tests {
fn message_roundtrip_each_variant() {
let msgs = vec![
Message::System {
content: vec![ContentBlock::Text {
text: "sys".into(),
}],
content: vec![ContentBlock::Text { text: "sys".into() }],
},
Message::User {
content: vec![ContentBlock::Text {
@@ -376,9 +368,7 @@ mod tests {
},
Message::ToolResult {
tool_call_id: "call_1".into(),
content: vec![ContentBlock::Text {
text: "ok".into(),
}],
content: vec![ContentBlock::Text { text: "ok".into() }],
is_error: true,
},
];
+1 -5
View File
@@ -26,7 +26,7 @@ pub use shared::{
AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier,
StopSequence,
};
pub use tool::{FunctionCall, OpenaiToolCall, OpenaiToolDefinition};
pub use tool::{FunctionCall, OpenaiToolCall, ToolDef};
pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage};
// Re-export IR 内容块 / 消息类型供 `types::ContentBlock` 等历史路径消费。
@@ -96,7 +96,3 @@ impl From<ChatResponse> for OpenaiChatChunk {
}
}
}
/// 工具定义别名(无新类型冲突,保留)。
#[deprecated(since = "0.1.0", note = "ToolDefinition 仍直接对应 OpenAI wire-format;未来 v0.2 引入 IR 工具类型后会再次更新")]
pub type ToolDefinition = OpenaiToolDefinition;
+8 -5
View File
@@ -11,18 +11,21 @@ pub struct StreamOptions {
pub include_obfuscation: Option<bool>,
}
#[derive(Debug, Clone)]
#[derive(Default)]
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub enum ToolChoice {
#[default]
None,
Auto,
Required,
Named { name: String },
AllowedTools { tool_names: Vec<String> },
Named {
name: String,
},
AllowedTools {
tool_names: Vec<String>,
},
}
impl Serialize for ToolChoice {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
+48 -16
View File
@@ -10,14 +10,14 @@ use thiserror::Error;
use crate::llm::types::message::Message;
use crate::llm::types::request::ToolChoice;
use crate::llm::types::tool::OpenaiToolDefinition;
use crate::llm::types::tool::ToolDef;
/// Provider 无关的请求类型。
///
/// 设计要点:
/// - `system` 字段不存在;system 提示由调用方通过 `Message::System` 在 `messages` 中表达。
/// - `tools` / `tool_choice` 直接复用现有 `OpenaiToolDefinition` / `ToolChoice`
/// (10a §251 决策:先复用旧类型,Phase 2 切换为新 `ToolDefinition` 后再调整)
/// - `tools` 使用 Provider 无关的 `ToolDef` IR;各 Provider 适配层在 `convert_request`
/// 中转换为对应 wire format。`tool_choice` 复用现有 `ToolChoice`
/// - `extra` 作为逃生舱:Provider 特定字段(`web_search_options`、`previous_response_id` 等)
/// 通过 `extra.set_extra / get_extra` 传递,避免持续膨胀本结构体。
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
@@ -26,8 +26,8 @@ pub struct MessageRequest {
pub model: String,
/// 消息列表(包含 system / user / assistant / tool_result 等所有变体)。
pub messages: Vec<Message>,
/// 工具定义列表。
pub tools: Vec<OpenaiToolDefinition>,
/// 工具定义列表Provider 无关 IR
pub tools: Vec<ToolDef>,
/// 工具选择策略。
pub tool_choice: ToolChoice,
/// 最大输出 token 数。
@@ -127,14 +127,9 @@ mod tests {
#[test]
fn extra_set_and_get_roundtrip() {
let mut req = MessageRequest::default();
req.set_extra(
"previous_response_id",
"resp_abc123",
);
req.set_extra("previous_response_id", "resp_abc123");
let v: Option<String> = req
.get_extra("previous_response_id")
.expect("get_extra ok");
let v: Option<String> = req.get_extra("previous_response_id").expect("get_extra ok");
assert_eq!(v.as_deref(), Some("resp_abc123"));
let missing: Option<String> = req.get_extra("missing").expect("missing ok");
@@ -174,10 +169,7 @@ mod tests {
}
let opts: Options = req.get_extra_as().expect("get_extra_as ok");
assert_eq!(
opts.web_search_options.search_context_size,
"high"
);
assert_eq!(opts.web_search_options.search_context_size, "high");
assert_eq!(opts.user.as_deref(), Some("u_123"));
}
@@ -206,4 +198,44 @@ mod tests {
assert_eq!(decoded.stream, req.stream);
assert_eq!(decoded.extra.get("trace_id"), Some(&json!("t-1")));
}
#[test]
fn message_request_with_tools_roundtrip() {
// 验证 ToolDef 的 serde 属性与 OpenaiToolDefinition 一致:
// 同名字段(name/description/parameters)序列化结果应一致。
let params = json!({
"type": "object",
"properties": {"x": {"type": "number"}},
"required": ["x"],
});
let tool = super::ToolDef {
name: "add".to_string(),
description: Some("add two numbers".to_string()),
parameters: params.clone(),
};
let req = MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
tools: vec![tool],
tool_choice: ToolChoice::Auto,
max_tokens: None,
temperature: None,
top_p: None,
stop_sequences: vec![],
stream: false,
thinking: None,
extra: HashMap::new(),
};
let json = serde_json::to_string(&req).expect("serialize");
// 验证反序列化能还原所有字段(包括嵌套 parameters
let decoded: MessageRequest = serde_json::from_str(&json).expect("deserialize");
assert_eq!(decoded.tools.len(), 1);
assert_eq!(decoded.tools[0].name, "add");
assert_eq!(decoded.tools[0].description.as_deref(), Some("add two numbers"));
assert_eq!(decoded.tools[0].parameters, params);
// 验证序列化 JSON 不含 ToolDef 没有的字段(如 strict),保持 wire-format 兼容
assert!(!json.contains("strict"), "ToolDef 序列化不应包含 strict 字段");
}
}
+2 -6
View File
@@ -2,8 +2,8 @@ use crate::llm::types::openai_message::OpenaiChatMessage;
use crate::llm::types::shared::{FinishReason, ServiceTier};
use crate::llm::types::tool::OpenaiToolCall;
use crate::llm::types::usage::Usage;
use serde::{Deserialize, Serialize};
use crate::llm::types::{ContentField, OpenaiContentPart};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenLogprob {
@@ -135,11 +135,7 @@ impl From<OpenaiChatMessage> for Delta {
text.push_str(&t);
}
}
if text.is_empty() {
None
} else {
Some(text)
}
if text.is_empty() { None } else { Some(text) }
}
},
refusal: None,
+35 -23
View File
@@ -19,6 +19,7 @@ use crate::llm::types::usage::{CompletionTokensDetails, PromptTokensDetails, Usa
/// Phase 2 完成时统一收敛。
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum StopReason {
/// 自然停止。
Stop,
@@ -164,11 +165,15 @@ pub enum ContentBlockBuilder {
/// `thinking_signature`,最终通过 `finalize()` 回填到 `full_response` 的 `Thinking` block 中。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum StreamEvent {
/// 消息开始(元信息)。
MessageStart { id: String, model: String },
/// 内容块开始(告知块类型,携带 id/name for ToolUse)。
ContentBlockStart { index: u32, block_type: ContentBlockType },
ContentBlockStart {
index: u32,
block_type: ContentBlockType,
},
/// 内容块结束标记。
ContentBlockEnd { index: u32 },
/// 文本增量。
@@ -187,6 +192,24 @@ pub enum StreamEvent {
MessageComplete { full_response: MessageResponse },
/// 错误事件。
Error { message: String },
/// 工具开始执行 —— 在 `ToolCallEnd` 之后、`registry.invoke_all` 之前发出。
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
ToolExecutionStarted {
tool_name: String,
tool_call_id: String,
/// 工具参数(JSON 字符串形式),用于 UI 展示
arguments: String,
},
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
ToolExecutionCompleted {
tool_name: String,
tool_call_id: String,
/// 结果摘要(由 `CycleConfig.max_tool_result_bytes` 截断,默认 65536 字节/字符边界安全),
/// 用于 UI 反馈。完整结果已在内部 `messages` 中作为 `ToolResult` 回传给 LLM。
result_summary: String,
/// 是否执行出错
is_error: bool,
},
}
/// 流式响应累积状态。
@@ -320,9 +343,8 @@ impl PartialMessageResponse {
true
}
StreamEvent::ToolCallArgumentsDelta { index, arguments } => {
if let Some(ContentBlockBuilder::ToolUse {
arguments: buf, ..
}) = self.blocks.get_mut(index)
if let Some(ContentBlockBuilder::ToolUse { arguments: buf, .. }) =
self.blocks.get_mut(index)
{
buf.push_str(arguments);
}
@@ -348,6 +370,9 @@ impl PartialMessageResponse {
self.is_errored = true;
false
}
// 元事件:不参与内容块累积,不修改 partial 状态
//(Phase 9 —— 工具执行透明化,由 run_tool_loop 在工具前后插入)
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
}
}
@@ -360,9 +385,7 @@ impl PartialMessageResponse {
let mut content_blocks = Vec::with_capacity(self.blocks.len());
for (idx, builder) in self.blocks {
let block = Self::builder_to_block(idx, builder, self.thinking_signature.as_deref())
.map_err(|e| LlmError::Other(format!(
"partial 块 #{idx} finalize 失败: {e}"
)))?;
.map_err(|e| LlmError::Other(format!("partial 块 #{idx} finalize 失败: {e}")))?;
content_blocks.push(block);
}
@@ -743,10 +766,7 @@ mod tests {
Message::Assistant { content } => {
assert_eq!(content.len(), 2);
match (&content[0], &content[1]) {
(
ContentBlock::Text { text: t1 },
ContentBlock::Text { text: t2 },
) => {
(ContentBlock::Text { text: t1 }, ContentBlock::Text { text: t2 }) => {
assert_eq!(t1, "first");
assert_eq!(t2, "second");
}
@@ -772,9 +792,7 @@ mod tests {
index: 0,
block_type: ContentBlockType::Text,
},
StreamEvent::TextDelta {
text: "x".into(),
},
StreamEvent::TextDelta { text: "x".into() },
StreamEvent::ContentBlockEnd { index: 0 },
StreamEvent::MessageComplete {
full_response: empty_response(),
@@ -832,15 +850,9 @@ mod tests {
block_type: ContentBlockType::Text,
},
StreamEvent::ContentBlockEnd { index: 0 },
StreamEvent::TextDelta {
text: "t".into(),
},
StreamEvent::ThinkingDelta {
text: "p".into(),
},
StreamEvent::RefusalDelta {
text: "r".into(),
},
StreamEvent::TextDelta { text: "t".into() },
StreamEvent::ThinkingDelta { text: "p".into() },
StreamEvent::RefusalDelta { text: "r".into() },
StreamEvent::ToolCallArgumentsDelta {
index: 1,
arguments: "{\"x\":1}".into(),
+2
View File
@@ -13,6 +13,7 @@ pub enum Role {
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum FinishReason {
Stop,
Length,
@@ -67,6 +68,7 @@ pub enum StopSequence {
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "type")]
#[non_exhaustive]
pub enum ResponseFormat {
Text,
JsonObject,
+40
View File
@@ -1,6 +1,25 @@
use serde::{Deserialize, Serialize};
use serde_json::Value;
/// Provider 无关的工具定义 IRv0.2 引入,替换 `ToolDefinition` 别名)。
///
/// 字段最小化:仅承载跨 Provider 公共的概念(name、description、parameters)。
/// OpenAI 专属 `strict` 字段不在此表达,由 OpenAI 适配层通过
/// `MessageRequest.extra` 逃生舱在 `convert_request` 内补充。
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ToolDef {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default)]
pub parameters: Value,
}
/// 旧 OpenAI wire-format 工具定义(v0.2 降级为 `#[doc(hidden)]`)。
///
/// 由 `ToolDef` 替代;保留仅供 OpenAI 适配层消费 `ToolDef → OpenaiToolDefinition`
/// 转换与外部反序列化兼容路径使用,不作为公共 API。
#[doc(hidden)]
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct OpenaiToolDefinition {
pub name: String,
@@ -12,6 +31,27 @@ pub struct OpenaiToolDefinition {
pub strict: Option<bool>,
}
impl From<ToolDef> for OpenaiToolDefinition {
fn from(t: ToolDef) -> Self {
Self {
name: t.name,
description: t.description,
parameters: t.parameters,
strict: None,
}
}
}
impl From<OpenaiToolDefinition> for ToolDef {
fn from(t: OpenaiToolDefinition) -> Self {
Self {
name: t.name,
description: t.description,
parameters: t.parameters,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FunctionCall {
pub name: String,
+3 -3
View File
@@ -12,11 +12,11 @@ pub use conversation::{ConversationMemory, ConversationMemoryConfig};
pub use error::MemoryError;
pub use knowledge::KnowledgeStore;
pub use retriever::MemoryRetriever;
pub use store::{InMemoryStore, MemoryStore};
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
// 低频类型(配置/高级使用)
pub use conversation::MemoryStrategy;
pub use knowledge::{PageIndexEntry, KNOWLEDGE_PREFIX};
pub use retriever::{RetrieverConfig, RetrievalResult, ScoredItem};
pub use knowledge::{KNOWLEDGE_PREFIX, PageIndexEntry};
pub use retriever::{RetrievalResult, RetrieverConfig, ScoredItem};
pub use store::{EvictionConfig, EvictionPolicy};
pub use types::{KnowledgePage, MemoryFilter, MemoryItem};
+21 -14
View File
@@ -12,6 +12,7 @@ use crate::memory::types::MemoryItem;
/// 对话消息管理策略。
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[non_exhaustive]
pub enum MemoryStrategy {
/// 滑动窗口:达到上限时删除最旧消息。
SlidingWindow,
@@ -160,7 +161,12 @@ impl ConversationMemory {
}
fn make_message_id(&self, index: usize, now: &OffsetDateTime) -> String {
format!("{}{:010}_{}", self.session_prefix(), index, now.unix_timestamp_nanos())
format!(
"{}{:010}_{}",
self.session_prefix(),
index,
now.unix_timestamp_nanos()
)
}
async fn maybe_evict_and_compact(&mut self) {
@@ -175,15 +181,16 @@ impl ConversationMemory {
}
if let Some(ref compact_config) = self.config.compact_config
&& should_compact(&self.messages, compact_config, &self.compact_state) {
let keep_recent = compact_config.keep_recent;
let freed = microcompact(&mut self.messages, keep_recent);
if freed > 0 {
self.compact_state.record_success();
} else {
let _ = self.compact_state.record_failure();
}
&& should_compact(&self.messages, compact_config, &self.compact_state)
{
let keep_recent = compact_config.keep_recent;
let freed = microcompact(&mut self.messages, keep_recent);
if freed > 0 {
self.compact_state.record_success();
} else {
let _ = self.compact_state.record_failure();
}
}
}
}
@@ -196,7 +203,8 @@ mod tests {
#[tokio::test]
async fn add_and_get_history() {
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
let mut conv = ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
let mut conv =
ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
conv.add_message(Message::user_text("hello")).await.unwrap();
conv.add_message(Message::user_text("world")).await.unwrap();
assert_eq!(conv.len(), 2);
@@ -211,9 +219,7 @@ mod tests {
conv.add_message(Message::tool_result("call_1", "ok", false))
.await
.unwrap();
conv.add_message(Message::assistant("done"))
.await
.unwrap();
conv.add_message(Message::assistant("done")).await.unwrap();
let original = conv.get_history().to_vec();
assert_eq!(original.len(), 2);
@@ -263,7 +269,8 @@ mod tests {
#[tokio::test]
async fn clear_empties_messages() {
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
let mut conv = ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
let mut conv =
ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
conv.add_message(Message::user_text("hello")).await.unwrap();
assert!(!conv.is_empty());
conv.clear().await.unwrap();
+2 -1
View File
@@ -6,6 +6,7 @@ use thiserror::Error;
///
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查环境变量、重试)。
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum MemoryError {
/// 按 ID 未找到指定记忆条目。可重试——通常是 namespace 拼写错误或条目已被淘汰。
#[error("未找到记忆条目 '{0}',请检查 ID 或 namespace 是否正确")]
@@ -33,4 +34,4 @@ impl MemoryError {
pub fn is_recoverable(&self) -> bool {
matches!(self, Self::NotFound(_) | Self::RetrievalError(_))
}
}
}
+6 -3
View File
@@ -57,8 +57,8 @@ impl KnowledgeStore {
}
let now = OffsetDateTime::now_utc();
let id = format!("{KNOWLEDGE_PREFIX}{}", page.id);
let content = serde_json::to_string(&page)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
let content =
serde_json::to_string(&page).map_err(|e| MemoryError::Serialization(e.to_string()))?;
let item = MemoryItem {
id,
content,
@@ -128,7 +128,10 @@ impl KnowledgeStore {
.filter(|entry| {
entry.title.to_lowercase().contains(&needle)
|| entry.summary.to_lowercase().contains(&needle)
|| entry.tags.iter().any(|t| t.to_lowercase().contains(&needle))
|| entry
.tags
.iter()
.any(|t| t.to_lowercase().contains(&needle))
})
.map(|entry| entry.id.clone())
.collect()
+15 -8
View File
@@ -97,7 +97,11 @@ impl MemoryRetriever {
// 4. 过滤 → 排序 → 截取
items.retain(|i| i.score >= self.config.min_score);
items.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
items.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
items.truncate(self.config.max_results);
Ok(RetrievalResult {
@@ -159,12 +163,12 @@ fn char_bigrams(s: &str) -> Vec<String> {
fn default_stop_words() -> HashSet<String> {
[
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has",
"had", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall",
"can", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what",
"which", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor",
"so", "if", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at",
"by", "of", "as", "into", "through", "during", "before", "after", "above", "below",
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had",
"do", "does", "did", "will", "would", "should", "could", "may", "might", "shall", "can",
"this", "that", "these", "those", "it", "its", "they", "them", "their", "what", "which",
"who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor", "so", "if",
"then", "else", "with", "without", "for", "to", "from", "in", "on", "at", "by", "of", "as",
"into", "through", "during", "before", "after", "above", "below",
]
.iter()
.map(|s| s.to_string())
@@ -236,7 +240,10 @@ mod tests {
min_score: 0.99,
};
let retriever = MemoryRetriever::new(ks, config);
let result = retriever.retrieve("totally unrelated content").await.unwrap();
let result = retriever
.retrieve("totally unrelated content")
.await
.unwrap();
assert!(result.items.is_empty());
}
+7 -259
View File
@@ -1,14 +1,16 @@
//! MemoryStore 抽象接口与默认实现。
use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait;
use time::OffsetDateTime;
use crate::memory::error::MemoryError;
use crate::memory::types::{MemoryFilter, MemoryItem};
pub mod in_memory;
pub mod sqlite_store;
pub use in_memory::InMemoryStore;
pub use sqlite_store::SqliteStore;
/// 底层记忆存储抽象接口。
///
/// 下游可实现此 trait 以对接持久化后端(JSON 文件、SQLite、Redis 等)。
@@ -32,6 +34,7 @@ pub trait MemoryStore: Send + Sync {
/// 淘汰策略。
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum EvictionPolicy {
/// 不淘汰(默认)。
None,
@@ -57,258 +60,3 @@ impl Default for EvictionConfig {
}
}
}
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
pub struct InMemoryStore {
items: Mutex<HashMap<String, MemoryItem>>,
eviction: EvictionConfig,
/// 自上次淘汰检查以来的写入次数。
writes_since_check: Mutex<usize>,
}
impl InMemoryStore {
/// 创建一个无淘汰策略的 InMemoryStore。
pub fn new() -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction: EvictionConfig::default(),
writes_since_check: Mutex::new(0),
}
}
/// 创建一个带淘汰配置的 InMemoryStore。
pub fn with_eviction(eviction: EvictionConfig) -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction,
writes_since_check: Mutex::new(0),
}
}
fn maybe_evict(&self) {
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
let should_check = {
let mut counter = self.writes_since_check.lock().unwrap();
*counter += 1;
if *counter >= self.eviction.check_interval {
*counter = 0;
true
} else {
false
}
};
if !should_check {
return;
}
let policy = self.eviction.policy.clone();
match policy {
EvictionPolicy::None => {}
EvictionPolicy::Ttl { ttl_secs } => {
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
let mut items = self.items.lock().unwrap();
items.retain(|_, v| v.created_at > cutoff);
}
EvictionPolicy::Capacity { max_items } => {
let mut items = self.items.lock().unwrap();
if items.len() > max_items {
let mut vec: Vec<_> = items.drain().collect();
// O(n) 部分排序:保留 created_at 最大的 max_items 个
vec.select_nth_unstable_by(max_items, |a, b| {
b.1.created_at.cmp(&a.1.created_at)
});
vec.truncate(max_items);
*items = vec.into_iter().collect();
}
}
}
}
}
impl Default for InMemoryStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl MemoryStore for InMemoryStore {
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
{
let mut items = self.items.lock().unwrap();
items.insert(item.id.clone(), item);
}
self.maybe_evict();
Ok(())
}
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
Ok(items.get(id).cloned())
}
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
let mut items = self.items.lock().unwrap();
items.remove(id);
Ok(())
}
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
let mut result: Vec<MemoryItem> = items
.values()
.filter(|v| match &filter.prefix {
Some(p) => v.id.starts_with(p),
None => true,
})
.filter(|v| match filter.since {
Some(t) => v.created_at > t,
None => true,
})
.cloned()
.collect();
// 按 created_at 升序排列(最旧在前)
result.sort_by_key(|v| v.created_at);
// 应用 offset
if let Some(offset) = filter.offset {
if offset < result.len() {
result.drain(..offset);
} else {
result.clear();
}
}
// 应用 limit
if let Some(limit) = filter.limit {
result.truncate(limit);
}
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use time::OffsetDateTime;
fn make_item(id: &str) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({}),
created_at: OffsetDateTime::now_utc(),
}
}
#[tokio::test]
async fn save_get_delete_list() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
let got = store.get("a").await.unwrap();
assert!(got.is_some());
assert_eq!(got.unwrap().id, "a");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
store.delete("a").await.unwrap();
assert!(store.get("a").await.unwrap().is_none());
}
#[tokio::test]
async fn save_is_upsert() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
let mut item = make_item("a");
item.content = "updated".to_string();
store.save(item).await.unwrap();
let got = store.get("a").await.unwrap().unwrap();
assert_eq!(got.content, "updated");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn list_with_prefix_and_limit() {
let store = InMemoryStore::new();
store.save(make_item("foo_a")).await.unwrap();
store.save(make_item("foo_b")).await.unwrap();
store.save(make_item("bar_a")).await.unwrap();
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
limit: Some(1),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn capacity_eviction() {
// 强制每次写入都检查
let eviction = EvictionConfig {
policy: EvictionPolicy::Capacity { max_items: 2 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
// 第一条和第二条共存
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
// 第三条写入触发淘汰:a 或 b 之一被淘汰
store.save(make_item("c")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
// 留下的应该是 b 和 c(最新的两个)
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"b"));
assert!(ids.contains(&"c"));
}
#[tokio::test]
async fn ttl_eviction() {
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
// 改用小 TTL + 睡眠:先 save asleepsave b 时 a 已过期被淘汰。
let eviction = EvictionConfig {
policy: EvictionPolicy::Ttl { ttl_secs: 1 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
store.save(make_item("a")).await.unwrap();
// 等待超过 1 秒
std::thread::sleep(std::time::Duration::from_millis(1100));
// 触发淘汰:a 已超过 ttl_secs=1,应被淘汰
store.save(make_item("b")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
// 由于 ttl_secs=1,且 b 刚写入,可能刚好处于临界值。
// 我们只断言 list 不包含 "a" 即可。
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(
!ids.contains(&"a"),
"expected 'a' to be evicted, but found in {ids:?}"
);
}
#[tokio::test]
async fn none_policy_no_eviction() {
let eviction = EvictionConfig {
policy: EvictionPolicy::None,
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
for i in 0..100 {
store.save(make_item(&format!("item_{i}"))).await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
}
}
+266
View File
@@ -0,0 +1,266 @@
//! 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait;
use time::OffsetDateTime;
use crate::memory::error::MemoryError;
use crate::memory::store::{EvictionConfig, EvictionPolicy, MemoryStore};
use crate::memory::types::{MemoryFilter, MemoryItem};
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
pub struct InMemoryStore {
items: Mutex<HashMap<String, MemoryItem>>,
eviction: EvictionConfig,
/// 自上次淘汰检查以来的写入次数。
writes_since_check: Mutex<usize>,
}
impl InMemoryStore {
/// 创建一个无淘汰策略的 InMemoryStore。
pub fn new() -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction: EvictionConfig::default(),
writes_since_check: Mutex::new(0),
}
}
/// 创建一个带淘汰配置的 InMemoryStore。
pub fn with_eviction(eviction: EvictionConfig) -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction,
writes_since_check: Mutex::new(0),
}
}
fn maybe_evict(&self) {
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
let should_check = {
let mut counter = self.writes_since_check.lock().unwrap();
*counter += 1;
if *counter >= self.eviction.check_interval {
*counter = 0;
true
} else {
false
}
};
if !should_check {
return;
}
let policy = self.eviction.policy.clone();
match policy {
EvictionPolicy::None => {}
EvictionPolicy::Ttl { ttl_secs } => {
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
let mut items = self.items.lock().unwrap();
items.retain(|_, v| v.created_at > cutoff);
}
EvictionPolicy::Capacity { max_items } => {
let mut items = self.items.lock().unwrap();
if items.len() > max_items {
let mut vec: Vec<_> = items.drain().collect();
// O(n) 部分排序:保留 created_at 最大的 max_items 个
vec.select_nth_unstable_by(max_items, |a, b| {
b.1.created_at.cmp(&a.1.created_at)
});
vec.truncate(max_items);
*items = vec.into_iter().collect();
}
}
}
}
}
impl Default for InMemoryStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl MemoryStore for InMemoryStore {
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
{
let mut items = self.items.lock().unwrap();
items.insert(item.id.clone(), item);
}
self.maybe_evict();
Ok(())
}
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
Ok(items.get(id).cloned())
}
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
let mut items = self.items.lock().unwrap();
items.remove(id);
Ok(())
}
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
let mut result: Vec<MemoryItem> = items
.values()
.filter(|v| match &filter.prefix {
Some(p) => v.id.starts_with(p),
None => true,
})
.filter(|v| match filter.since {
Some(t) => v.created_at > t,
None => true,
})
.cloned()
.collect();
// 按 created_at 升序排列(最旧在前)
result.sort_by_key(|v| v.created_at);
// 应用 offset
if let Some(offset) = filter.offset {
if offset < result.len() {
result.drain(..offset);
} else {
result.clear();
}
}
// 应用 limit
if let Some(limit) = filter.limit {
result.truncate(limit);
}
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use time::OffsetDateTime;
fn make_item(id: &str) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({}),
created_at: OffsetDateTime::now_utc(),
}
}
#[tokio::test]
async fn save_get_delete_list() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
let got = store.get("a").await.unwrap();
assert!(got.is_some());
assert_eq!(got.unwrap().id, "a");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
store.delete("a").await.unwrap();
assert!(store.get("a").await.unwrap().is_none());
}
#[tokio::test]
async fn save_is_upsert() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
let mut item = make_item("a");
item.content = "updated".to_string();
store.save(item).await.unwrap();
let got = store.get("a").await.unwrap().unwrap();
assert_eq!(got.content, "updated");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn list_with_prefix_and_limit() {
let store = InMemoryStore::new();
store.save(make_item("foo_a")).await.unwrap();
store.save(make_item("foo_b")).await.unwrap();
store.save(make_item("bar_a")).await.unwrap();
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
limit: Some(1),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn capacity_eviction() {
// 强制每次写入都检查
let eviction = EvictionConfig {
policy: EvictionPolicy::Capacity { max_items: 2 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
// 第一条和第二条共存
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
// 第三条写入触发淘汰:a 或 b 之一被淘汰
store.save(make_item("c")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
// 留下的应该是 b 和 c(最新的两个)
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"b"));
assert!(ids.contains(&"c"));
}
#[tokio::test]
async fn ttl_eviction() {
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
// 改用小 TTL + 睡眠:先 save asleepsave b 时 a 已过期被淘汰。
let eviction = EvictionConfig {
policy: EvictionPolicy::Ttl { ttl_secs: 1 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
store.save(make_item("a")).await.unwrap();
// 等待超过 1 秒
std::thread::sleep(std::time::Duration::from_millis(1100));
// 触发淘汰:a 已超过 ttl_secs=1,应被淘汰
store.save(make_item("b")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
// 由于 ttl_secs=1,且 b 刚写入,可能刚好处于临界值。
// 我们只断言 list 不包含 "a" 即可。
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(
!ids.contains(&"a"),
"expected 'a' to be evicted, but found in {ids:?}"
);
}
#[tokio::test]
async fn none_policy_no_eviction() {
let eviction = EvictionConfig {
policy: EvictionPolicy::None,
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
for i in 0..100 {
store.save(make_item(&format!("item_{i}"))).await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
}
}
+545
View File
@@ -0,0 +1,545 @@
//! SqliteStore —— 基于 rusqlite 的持久化 MemoryStore 实现。
//!
//! 单进程独享、写入串行化(WAL + Mutex),适合本地 Agent 长期持久化场景。
use std::path::Path;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use rusqlite::{params, params_from_iter, Connection, ErrorCode};
use time::format_description::well_known::Rfc3339;
use time::OffsetDateTime;
use tracing::{debug, error, instrument, warn};
use crate::memory::error::MemoryError;
use crate::memory::store::MemoryStore;
use crate::memory::types::{MemoryFilter, MemoryItem};
const INITIAL_USER_VERSION: i64 = 1;
const BUSY_TIMEOUT_MS: i64 = 5000;
const WAL_AUTOCHECKPOINT_PAGES: i64 = 1000;
/// SQLite 持久化后端的 MemoryStore 实现。
///
/// 设计要点:
/// - 单进程独享:`Arc<Mutex<Connection>>` 串行化所有 IO
/// - WAL 模式 + `synchronous=NORMAL` 兼顾崩溃安全与吞吐
/// - `created_at` 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
/// - 所有 IO 通过 `tokio::task::spawn_blocking` 卸载到阻塞线程池
pub struct SqliteStore {
conn: Arc<Mutex<Connection>>,
}
impl SqliteStore {
/// 打开或创建一个 SQLite 数据库。
///
/// - `path = ":memory:"` 使用内存数据库(测试场景)
/// - 其他路径:自动创建父目录;文件已存在则附加打开
/// - 启动时执行 `migrate()`,失败立即返回错误
#[instrument(skip(path), fields(path = %path.as_ref().display()))]
pub fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
let path_ref = path.as_ref();
let path_str = path_ref.to_string_lossy();
let conn = if path_str == ":memory:" {
Connection::open_in_memory()
} else {
if let Some(parent) = path_ref.parent()
&& !parent.as_os_str().is_empty()
{
std::fs::create_dir_all(parent).map_err(|e| {
MemoryError::Storage(format!(
"创建数据库父目录失败 ({}): {}",
parent.display(),
e
))
})?;
}
Connection::open(path_ref)
}
.map_err(|e| map_sqlite_error(e, "打开数据库"))?;
migrate(&conn)?;
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
})
}
}
#[async_trait]
impl MemoryStore for SqliteStore {
#[instrument(skip(self, item), fields(id = %item.id))]
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
let conn = Arc::clone(&self.conn);
let created_at_str = item
.created_at
.to_offset(time::UtcOffset::UTC)
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format created_at: {e}")))?;
let metadata_str = serde_json::to_string(&item.metadata)
.map_err(|e| MemoryError::Serialization(format!("serialize metadata: {e}")))?;
let id = item.id;
let content = item.content;
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
conn.execute(
"INSERT INTO memory_items (id, content, metadata, created_at) \
VALUES (?1, ?2, ?3, ?4) \
ON CONFLICT(id) DO UPDATE SET \
content=excluded.content, \
metadata=excluded.metadata, \
created_at=excluded.created_at",
params![id, content, metadata_str, created_at_str],
)
.map_err(|e| map_sqlite_error(e, "保存记忆"))?;
Ok(())
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, id))]
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
let conn = Arc::clone(&self.conn);
let id_owned = id.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<MemoryItem>, MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
let mut stmt = conn
.prepare("SELECT id, content, metadata, created_at FROM memory_items WHERE id = ?1")
.map_err(|e| map_sqlite_error(e, "prepare get"))?;
let mut rows = stmt
.query_map(params![id_owned], row_to_item)
.map_err(|e| map_sqlite_error(e, "query get"))?;
match rows.next() {
None => Ok(None),
Some(row) => row
.map(Some)
.map_err(|e| map_sqlite_error(e, "decode row")),
}
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, id))]
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
let conn = Arc::clone(&self.conn);
let id_owned = id.to_string();
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
conn.execute(
"DELETE FROM memory_items WHERE id = ?1",
params![id_owned],
)
.map_err(|e| map_sqlite_error(e, "delete"))?;
Ok(())
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, filter))]
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
let mut sql = String::from(
"SELECT id, content, metadata, created_at FROM memory_items WHERE 1=1",
);
let mut param_values: Vec<String> = Vec::new();
let mut ph_idx = 0usize;
if filter.prefix.is_some() {
ph_idx += 1;
sql.push_str(&format!(" AND id LIKE ?{ph_idx} || '%'"));
}
if filter.since.is_some() {
ph_idx += 1;
sql.push_str(&format!(" AND created_at > ?{ph_idx}"));
}
// ORDER BY created_at ASC(按时间升序,最旧在前)
sql.push_str(" ORDER BY created_at ASC");
let limit_sql: String = match (filter.limit, filter.offset) {
(Some(_), Some(_)) => {
ph_idx += 1;
let limit_p = ph_idx;
ph_idx += 1;
let offset_p = ph_idx;
format!(" LIMIT ?{limit_p} OFFSET ?{offset_p}")
}
(Some(_), None) => {
ph_idx += 1;
let limit_p = ph_idx;
format!(" LIMIT ?{limit_p}")
}
(None, Some(_)) => {
// SQLite 中 LIMIT -1 表示无限制
ph_idx += 1;
let offset_p = ph_idx;
format!(" LIMIT -1 OFFSET ?{offset_p}")
}
(None, None) => String::new(),
};
sql.push_str(&limit_sql);
if let Some(p) = &filter.prefix {
param_values.push(p.clone());
}
if let Some(t) = filter.since {
let s = t
.to_offset(time::UtcOffset::UTC)
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format since: {e}")))?;
param_values.push(s);
}
if let Some(l) = filter.limit {
param_values.push(l.to_string());
}
if let Some(o) = filter.offset {
param_values.push(o.to_string());
}
let conn = Arc::clone(&self.conn);
let sql_owned = sql;
let param_values_owned = param_values;
tokio::task::spawn_blocking(move || -> Result<Vec<MemoryItem>, MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
let mut stmt = conn
.prepare(&sql_owned)
.map_err(|e| map_sqlite_error(e, "list prepare"))?;
let params_iter: Vec<&dyn rusqlite::ToSql> = param_values_owned
.iter()
.map(|s| s as &dyn rusqlite::ToSql)
.collect();
let rows = stmt
.query_map(params_from_iter(params_iter), row_to_item)
.map_err(|e| map_sqlite_error(e, "list query"))?;
let mut result = Vec::new();
for row in rows {
result.push(row.map_err(|e| map_sqlite_error(e, "list row"))?);
}
debug!(count = result.len(), "SqliteStore::list 完成");
Ok(result)
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
}
fn row_to_item(row: &rusqlite::Row<'_>) -> Result<MemoryItem, rusqlite::Error> {
let id: String = row.get(0)?;
let content: String = row.get(1)?;
let metadata_str: String = row.get(2)?;
let created_at_str: String = row.get(3)?;
let metadata: serde_json::Value = serde_json::from_str(&metadata_str).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(2, rusqlite::types::Type::Text, Box::new(e))
})?;
let created_at = OffsetDateTime::parse(&created_at_str, &Rfc3339).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(3, rusqlite::types::Type::Text, Box::new(e))
})?;
Ok(MemoryItem {
id,
content,
metadata,
created_at: created_at.to_offset(time::UtcOffset::UTC),
})
}
fn migrate(conn: &Connection) -> Result<(), MemoryError> {
conn.pragma_update(None, "journal_mode", "WAL")
.map_err(|e| map_sqlite_error(e, "PRAGMA journal_mode"))?;
conn.pragma_update(None, "synchronous", "NORMAL")
.map_err(|e| map_sqlite_error(e, "PRAGMA synchronous"))?;
conn.execute_batch(&format!("PRAGMA busy_timeout = {BUSY_TIMEOUT_MS};"))
.map_err(|e| map_sqlite_error(e, "PRAGMA busy_timeout"))?;
conn.execute_batch(&format!(
"PRAGMA wal_autocheckpoint = {WAL_AUTOCHECKPOINT_PAGES};"
))
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_autocheckpoint"))?;
let version: i64 = conn
.query_row("PRAGMA user_version", [], |row| row.get(0))
.map_err(|e| map_sqlite_error(e, "PRAGMA user_version"))?;
if version < INITIAL_USER_VERSION {
conn.execute_batch(
"CREATE TABLE IF NOT EXISTS memory_items (
id TEXT PRIMARY KEY,
content TEXT NOT NULL,
metadata TEXT NOT NULL DEFAULT '{}',
created_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_memory_items_created_at
ON memory_items(created_at);
PRAGMA user_version = 1;",
)
.map_err(|e| map_sqlite_error(e, "create schema v1"))?;
}
let check_result: String = conn
.query_row("PRAGMA quick_check", [], |row| row.get(0))
.map_err(|e| map_sqlite_error(e, "PRAGMA quick_check"))?;
if check_result != "ok" {
error!(result = %check_result, "数据库文件 quick_check 失败");
return Err(MemoryError::Storage(format!(
"数据库文件损坏: {check_result}"
)));
}
conn.execute_batch("PRAGMA wal_checkpoint(TRUNCATE);")
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_checkpoint"))?;
Ok(())
}
fn map_sqlite_error(e: rusqlite::Error, ctx: &str) -> MemoryError {
match &e {
rusqlite::Error::SqliteFailure(err, _) => match err.code {
ErrorCode::ConstraintViolation => MemoryError::InvalidInput(format!("{ctx}: {e}")),
ErrorCode::DatabaseBusy | ErrorCode::DatabaseLocked => {
warn!("SQLite 忙: {e}");
MemoryError::Storage(format!("{ctx}: {e}"))
}
_ => MemoryError::Storage(format!("{ctx}: {e}")),
},
rusqlite::Error::InvalidQuery
| rusqlite::Error::InvalidParameterName(_)
| rusqlite::Error::InvalidColumnIndex(_)
| rusqlite::Error::InvalidColumnName(_) => {
MemoryError::InvalidInput(format!("{ctx}: {e}"))
}
rusqlite::Error::FromSqlConversionFailure(_, _, _)
| rusqlite::Error::ToSqlConversionFailure(_) => {
MemoryError::Serialization(format!("{ctx}: {e}"))
}
_ => MemoryError::Storage(format!("{ctx}: {e}")),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::store::InMemoryStore;
use std::sync::Arc;
use tempfile::TempDir;
use time::OffsetDateTime;
fn make_item(id: &str) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({"id_key": id}),
created_at: OffsetDateTime::now_utc(),
}
}
fn make_item_at(id: &str, when: OffsetDateTime) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({}),
created_at: when,
}
}
#[tokio::test]
async fn crud_basic() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
let got_a = store.get("a").await.unwrap();
assert!(got_a.is_some());
assert_eq!(got_a.unwrap().id, "a");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
store.delete("a").await.unwrap();
assert!(store.get("a").await.unwrap().is_none());
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
assert_eq!(list[0].id, "b");
}
#[tokio::test]
async fn save_is_upsert() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("a")).await.unwrap();
let mut item = make_item("a");
item.content = "updated".to_string();
item.metadata = serde_json::json!({"rev": 2});
let original_created_at = item.created_at;
store.save(item).await.unwrap();
let got = store.get("a").await.unwrap().unwrap();
assert_eq!(got.content, "updated");
assert_eq!(got.metadata["rev"], serde_json::json!(2));
// created_at 保持调用方传入值
assert_eq!(got.created_at, original_created_at);
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn list_with_prefix() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("foo_a")).await.unwrap();
store.save(make_item("foo_b")).await.unwrap();
store.save(make_item("bar_a")).await.unwrap();
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"foo_a"));
assert!(ids.contains(&"foo_b"));
assert!(!ids.contains(&"bar_a"));
}
#[tokio::test]
async fn list_with_since_filter() {
let store = SqliteStore::open(":memory:").unwrap();
let t0 = OffsetDateTime::now_utc();
store
.save(make_item_at("early", t0 - time::Duration::seconds(60)))
.await
.unwrap();
store.save(make_item_at("middle", t0)).await.unwrap();
store
.save(make_item_at("late", t0 + time::Duration::seconds(60)))
.await
.unwrap();
let filter = MemoryFilter {
since: Some(t0 - time::Duration::seconds(1)),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert_eq!(list.len(), 2);
assert!(ids.contains(&"middle"));
assert!(ids.contains(&"late"));
assert!(!ids.contains(&"early"));
}
#[tokio::test]
async fn list_with_offset_and_limit() {
let store = SqliteStore::open(":memory:").unwrap();
// 写入 5 条时间递增的记录
let base = OffsetDateTime::now_utc() - time::Duration::seconds(5);
for i in 0..5 {
let mut item = make_item(&format!("item_{i}"));
item.created_at = base + time::Duration::seconds(i);
store.save(item).await.unwrap();
}
// offset=1, limit=2 -> item_1, item_2
let filter = MemoryFilter {
offset: Some(1),
limit: Some(2),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
assert_eq!(list[0].id, "item_1");
assert_eq!(list[1].id, "item_2");
}
#[tokio::test]
async fn concurrent_writers_no_data_loss() {
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
let mut handles = Vec::new();
for w in 0..10 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
for i in 0..10 {
let id = format!("w{w}_i{i}");
s.save(make_item(&id)).await.unwrap();
}
}));
}
for h in handles {
h.await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
// 验证所有 id 唯一
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
ids.sort();
ids.dedup();
assert_eq!(ids.len(), 100);
}
#[tokio::test]
async fn persistence_round_trip() {
let dir = TempDir::new().unwrap();
let path = dir.path().join("memory.db");
// 阶段 1:写入 3 条
{
let store = SqliteStore::open(&path).unwrap();
store.save(make_item("alpha")).await.unwrap();
store.save(make_item("beta")).await.unwrap();
store.save(make_item("gamma")).await.unwrap();
}
// 阶段 2:重新打开,验证数据完整
{
let store = SqliteStore::open(&path).unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 3);
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"alpha"));
assert!(ids.contains(&"beta"));
assert!(ids.contains(&"gamma"));
// 单条读回
let got = store.get("beta").await.unwrap().unwrap();
assert_eq!(got.content, "content-beta");
}
}
#[tokio::test]
async fn open_invalid_path_returns_error() {
// 路径指向已存在的目录而非文件,open 应失败
let dir = TempDir::new().unwrap();
match SqliteStore::open(dir.path()) {
Err(MemoryError::Storage(_)) => {}
Err(other) => panic!("expected Storage error, got {other:?}"),
Ok(_) => panic!("expected error when opening a directory as database"),
}
}
#[tokio::test]
async fn trait_object_compatibility() {
// ponytail: 回归验证 SqliteStore 可作为 Arc<dyn MemoryStore> 与 InMemoryStore 互换
// 所有现有消费者(Conversation / Knowledge / Retriever / SessionMemory)均通过 trait object 引用,
// 此测试确保 trait 接口契约在 SqliteStore 上同样成立。
let sqlite: Arc<dyn MemoryStore> =
Arc::new(SqliteStore::open(":memory:").unwrap());
let in_mem: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let stores: Vec<Arc<dyn MemoryStore>> = vec![Arc::clone(&sqlite), Arc::clone(&in_mem)];
for store in &stores {
store.save(make_item("x")).await.unwrap();
let got = store.get("x").await.unwrap();
assert_eq!(got.unwrap().id, "x");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
store.delete("x").await.unwrap();
assert!(store.get("x").await.unwrap().is_none());
}
}
}
+2 -2
View File
@@ -1,7 +1,7 @@
pub mod composer;
pub mod error;
pub mod template;
pub mod composer;
pub use composer::{PromptComposer, validate_messages};
pub use error::PromptError;
pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue};
pub use composer::{validate_messages, PromptComposer};
+17 -15
View File
@@ -48,7 +48,11 @@ impl PromptComposer {
/// 添加一条 Tool 消息(工具执行结果回传)。
pub fn tool(mut self, tool_call_id: impl Into<String>, content: impl Into<String>) -> Self {
self.push_message(Message::tool_result(tool_call_id.into(), content.into(), false));
self.push_message(Message::tool_result(
tool_call_id.into(),
content.into(),
false,
));
self
}
@@ -133,11 +137,7 @@ impl PromptComposer {
}
/// 添加一条含指定 ContentBlock 的 Tool 消息。
pub fn tool_content(
mut self,
tool_call_id: impl Into<String>,
block: ContentBlock,
) -> Self {
pub fn tool_content(mut self, tool_call_id: impl Into<String>, block: ContentBlock) -> Self {
self.push_message(Message::ToolResult {
tool_call_id: tool_call_id.into(),
content: vec![block],
@@ -187,9 +187,7 @@ impl PromptComposer {
/// 验证消息序列是否符合 LLM API 要求(Tool 消息必须紧跟含 tool_calls 的 Assistant)。
pub fn validate_messages(messages: &[Message]) -> Result<(), PromptError> {
if messages.is_empty() {
return Err(PromptError::InvalidSequence(
"消息列表不能为空".to_string(),
));
return Err(PromptError::InvalidSequence("消息列表不能为空".to_string()));
}
let mut last_tool_call_ids: Vec<String> = Vec::new();
@@ -297,7 +295,8 @@ mod tests {
#[test]
fn test_template_if() {
let tpl = PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
let tpl =
PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
let mut ctx = TemplateContext::new();
ctx.insert("name", "Bob");
@@ -312,11 +311,14 @@ mod tests {
fn test_template_each() {
let tpl = PromptTemplate::compile("Items: {{#each items}}{{item}}, {{/each}}").unwrap();
let mut ctx = TemplateContext::new();
ctx.insert("items", TemplateValue::Array(vec![
TemplateValue::String("a".to_string()),
TemplateValue::String("b".to_string()),
TemplateValue::String("c".to_string()),
]));
ctx.insert(
"items",
TemplateValue::Array(vec![
TemplateValue::String("a".to_string()),
TemplateValue::String("b".to_string()),
TemplateValue::String("c".to_string()),
]),
);
let result = tpl.render(&ctx).unwrap();
assert_eq!(result, "Items: a, b, c, ");
+7 -2
View File
@@ -1,6 +1,7 @@
use thiserror::Error;
#[derive(Error, Debug)]
#[non_exhaustive]
pub enum PromptError {
#[error("模板解析错误: {0}。请检查模板语法({{var}} / {{#if}} / {{#each}}")]
Parse(String),
@@ -8,7 +9,9 @@ pub enum PromptError {
#[error("渲染错误: 变量 '{0}' 未找到。请在 TemplateContext 中插入该变量")]
VariableNotFound(String),
#[error("渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板")]
#[error(
"渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板"
)]
PartialNotFound(String),
#[error("渲染错误: '{0}' 不是数组,无法遍历。请确认传入的是数组或先判空")]
@@ -20,7 +23,9 @@ pub enum PromptError {
#[error("渲染错误: {0}")]
Render(String),
#[error("消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)")]
#[error(
"消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)"
)]
InvalidSequence(String),
#[error("文件读取错误: {0}。请检查模板文件路径与权限")]
+21 -34
View File
@@ -1,6 +1,6 @@
use serde_json::Value;
use std::collections::HashMap;
use std::fmt;
use serde_json::Value;
use crate::prompt::error::PromptError;
@@ -140,7 +140,9 @@ fn json_to_template_value(v: &Value) -> Result<TemplateValue, PromptError> {
#[derive(Debug, Clone)]
enum Fragment {
Literal(String),
Variable { name: String },
Variable {
name: String,
},
If {
condition: String,
body: Vec<Fragment>,
@@ -223,8 +225,7 @@ fn compile_fragments(template: &str) -> Result<Vec<Fragment>, PromptError> {
let tag = tag_content.trim();
if let Some(rest) = tag.strip_prefix("#if ") {
let (body, else_body, new_i) =
parse_block(template, i, "if")?;
let (body, else_body, new_i) = parse_block(template, i, "if")?;
let condition = rest.trim().to_string();
fragments.push(Fragment::If {
condition,
@@ -331,10 +332,7 @@ fn parse_block(
Err(PromptError::Parse(format!("未闭合的 {{#{}}}", kind)))
}
fn parse_each_block(
template: &str,
start: usize,
) -> Result<(Vec<Fragment>, usize), PromptError> {
fn parse_each_block(template: &str, start: usize) -> Result<(Vec<Fragment>, usize), PromptError> {
let bytes = template.as_bytes();
let len = bytes.len();
let mut depth = 1u32;
@@ -368,9 +366,7 @@ fn parse_each_block(
}
}
Err(PromptError::Parse(
"未闭合的 {{#each}} 块".to_string(),
))
Err(PromptError::Parse("未闭合的 {{#each}} 块".to_string()))
}
fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), PromptError> {
@@ -395,9 +391,7 @@ fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), Prom
}
}
Err(PromptError::Parse(
"未闭合的 {{#raw}} 块".to_string(),
))
Err(PromptError::Parse("未闭合的 {{#raw}} 块".to_string()))
}
// ===== Renderer =====
@@ -418,33 +412,28 @@ fn render_fragments(
Fragment::Literal(text) => {
output.push_str(text);
}
Fragment::Variable { name } => {
match ctx.get(name) {
Some(val) => {
output.push_str(&format!("{}", val));
}
None => {
return Err(PromptError::VariableNotFound(name.clone()));
}
Fragment::Variable { name } => match ctx.get(name) {
Some(val) => {
output.push_str(&format!("{}", val));
}
}
None => {
return Err(PromptError::VariableNotFound(name.clone()));
}
},
Fragment::If {
condition,
body,
else_body,
} => {
let truthy = ctx
.get(condition)
.map(|v| v.is_truthy())
.unwrap_or(false);
let truthy = ctx.get(condition).map(|v| v.is_truthy()).unwrap_or(false);
let target = if truthy { body } else { else_body };
render_fragments(target, ctx, partials, output, depth + 1)?;
}
Fragment::Each { variable, body } => {
let arr = match ctx.get(variable) {
Some(val) => val.as_array().ok_or_else(|| {
PromptError::NotAnArray(variable.clone())
})?,
Some(val) => val
.as_array()
.ok_or_else(|| PromptError::NotAnArray(variable.clone()))?,
None => {
return Err(PromptError::VariableNotFound(variable.clone()));
}
@@ -504,10 +493,8 @@ impl PromptTemplateRegistry {
/// 延迟编译注册:只存储原始字符串,首次渲染时编译。
pub fn register_lazy(&mut self, name: &str, template: &str) {
self.templates.insert(
name.to_string(),
StoredTemplate::Raw(template.to_string()),
);
self.templates
.insert(name.to_string(), StoredTemplate::Raw(template.to_string()));
}
/// 从文件读取并编译注册。
+10 -3
View File
@@ -4,9 +4,12 @@ use std::sync::Arc;
/// 工具调用过程中可能发生的所有错误。
#[derive(thiserror::Error, Debug, Clone)]
#[non_exhaustive]
pub enum ToolError {
/// 工具未注册。不可恢复——需调用方先 `registry.register(...)`。
#[error("工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写")]
#[error(
"工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写"
)]
NotFound(String),
/// 工具执行失败(可恢复——文本回传 LLM 由其决定重试或放弃)。
@@ -14,11 +17,15 @@ pub enum ToolError {
ExecutionFailed(String, String),
/// 工具参数无效(可恢复——文本回传 LLM)。
#[error("工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema")]
#[error(
"工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema"
)]
InvalidArguments(String, String),
/// 权限被拒绝(不可恢复——终止循环)。
#[error("权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过")]
#[error(
"权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过"
)]
PermissionDenied(String, String),
/// MCP 协议错误(不可恢复)。
+17 -37
View File
@@ -9,19 +9,18 @@
use std::collections::HashMap;
use std::process::Stdio;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use serde_json::{Value, json};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::{oneshot, Mutex};
use tokio::sync::{Mutex, oneshot};
#[allow(deprecated)]
use crate::llm::types::ToolDefinition;
use crate::llm::types::tool::ToolDef;
use crate::tools::base::{BaseTool, ToolContext, ToolRef};
use crate::tools::error::ToolError;
@@ -136,7 +135,6 @@ impl std::fmt::Debug for McpClient {
}
}
#[allow(deprecated)]
impl McpClient {
/// 创建一个 MCP 客户端。
pub fn new(server_name: impl Into<String>, transport: McpTransport) -> Self {
@@ -226,9 +224,7 @@ impl McpClient {
"version": env!("CARGO_PKG_VERSION")
}
});
let _response = self
.send_request("initialize", Some(init_params))
.await?;
let _response = self.send_request("initialize", Some(init_params)).await?;
// 发送 initialized 通知(无 id
self.send_notification("notifications/initialized", Some(json!({})))
@@ -239,7 +235,7 @@ impl McpClient {
}
/// 列出服务器支持的工具(调用 `tools/list`)。
pub async fn list_tools(&mut self) -> Result<Vec<ToolDefinition>, ToolError> {
pub async fn list_tools(&mut self) -> Result<Vec<ToolDef>, ToolError> {
if !self.is_initialized() {
return Err(ToolError::McpNotInitialized(self.server_name.clone()));
}
@@ -274,11 +270,10 @@ impl McpClient {
description: description.clone(),
input_schema: input_schema.clone(),
});
defs.push(ToolDefinition {
defs.push(ToolDef {
name,
description,
parameters: input_schema,
strict: None,
});
}
Ok(defs)
@@ -337,11 +332,7 @@ impl McpClient {
if let Some(state) = self.process.take() {
let mut state = state.lock().await;
// 优雅等待 5 秒
let graceful = tokio::time::timeout(
Duration::from_secs(5),
state.child.wait(),
)
.await;
let graceful = tokio::time::timeout(Duration::from_secs(5), state.child.wait()).await;
if graceful.is_err() {
// 超时则强杀
let _ = state.child.kill().await;
@@ -372,11 +363,7 @@ impl McpClient {
tools
}
async fn send_request(
&self,
method: &str,
params: Option<Value>,
) -> Result<Value, ToolError> {
async fn send_request(&self, method: &str, params: Option<Value>) -> Result<Value, ToolError> {
let state_arc = self
.process
.as_ref()
@@ -412,9 +399,11 @@ impl McpClient {
.write_all(b"\n")
.await
.map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?;
state.stdin.flush().await.map_err(|e| {
ToolError::McpError(format!("flush stdin 失败: {e}"))
})?;
state
.stdin
.flush()
.await
.map_err(|e| ToolError::McpError(format!("flush stdin 失败: {e}")))?;
}
// 等待响应(带超时)
@@ -471,10 +460,7 @@ impl McpClient {
}
/// 持续读取 stdout,将响应分发到对应的 oneshot sender。
async fn read_loop(
mut reader: BufReader<ChildStdout>,
state: Arc<Mutex<ChildProcessState>>,
) {
async fn read_loop(mut reader: BufReader<ChildStdout>, state: Arc<Mutex<ChildProcessState>>) {
let mut line = String::new();
loop {
line.clear();
@@ -542,7 +528,6 @@ enum McpClientHandle {
}
#[async_trait]
#[allow(deprecated)]
impl BaseTool for McpToolAdapter {
fn name(&self) -> &str {
&self.name
@@ -556,11 +541,7 @@ impl BaseTool for McpToolAdapter {
self.parameters.clone()
}
async fn execute(
&self,
_args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
// 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。
// 实际生产中应持有 Arc<McpClient> 并通过 mcp.call_tool() 执行。
// 这里返回错误,提示需要通过其他方式调用 MCP 工具。
@@ -617,8 +598,7 @@ mod tests {
#[test]
fn test_jsonrpc_response_parse_error() {
let s =
r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
let s = r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
let resp: JsonRpcResponse = serde_json::from_str(s).unwrap();
assert_eq!(resp.id, 1);
assert!(resp.result.is_none());
+21 -18
View File
@@ -148,9 +148,7 @@ mod tests {
#[test]
fn test_default_config_denies_delete() {
let checker = PermissionChecker::new(PermissionConfig::default());
assert!(checker
.check("rm_file", &p(Permission::Delete))
.is_err());
assert!(checker.check("rm_file", &p(Permission::Delete)).is_err());
}
#[test]
@@ -246,12 +244,16 @@ mod tests {
allow_unspecified: false,
};
let checker = PermissionChecker::new(cfg);
assert!(checker
.check("t", &[Permission::Custom("db:read".into())])
.is_ok());
assert!(checker
.check("t", &[Permission::Custom("db:write".into())])
.is_err());
assert!(
checker
.check("t", &[Permission::Custom("db:read".into())])
.is_ok()
);
assert!(
checker
.check("t", &[Permission::Custom("db:write".into())])
.is_err()
);
}
#[test]
@@ -262,12 +264,11 @@ mod tests {
allow_unspecified: false,
};
let checker = PermissionChecker::new(cfg);
assert!(checker
.check(
"t",
&[Permission::Read, Permission::Network]
)
.is_ok());
assert!(
checker
.check("t", &[Permission::Read, Permission::Network])
.is_ok()
);
}
#[test]
@@ -279,8 +280,10 @@ mod tests {
};
let checker = PermissionChecker::new(cfg);
// 任一权限不在白名单则拒绝
assert!(checker
.check("t", &[Permission::Read, Permission::Write])
.is_err());
assert!(
checker
.check("t", &[Permission::Read, Permission::Write])
.is_err()
);
}
}
+10 -10
View File
@@ -7,8 +7,7 @@ use std::time::Duration;
use futures::future::join_all;
use serde_json::Value;
#[allow(deprecated)]
use crate::llm::types::ToolDefinition;
use crate::llm::types::tool::ToolDef;
use crate::tools::base::{ToolContext, ToolRef};
use crate::tools::error::ToolError;
use crate::tools::permission::PermissionChecker;
@@ -71,7 +70,6 @@ impl std::fmt::Debug for ToolRegistry {
}
}
#[allow(deprecated)]
impl ToolRegistry {
/// 创建一个新的工具注册表。
pub fn new() -> Self {
@@ -127,16 +125,15 @@ impl ToolRegistry {
self.inner.tools.keys().cloned().collect()
}
/// 获取所有工具的 `ToolDefinition` 列表(用于传递给 LLM)。
pub fn definitions(&self) -> Vec<ToolDefinition> {
/// 获取所有工具的 `ToolDef` 列表(用于传递给 LLM)。
pub fn definitions(&self) -> Vec<ToolDef> {
self.inner
.tools
.values()
.map(|tool| ToolDefinition {
.map(|tool| ToolDef {
name: tool.name().to_string(),
description: Some(tool.description().to_string()),
parameters: tool.parameters(),
strict: None,
})
.collect()
}
@@ -348,7 +345,10 @@ mod tests {
async fn test_invoke_success() {
let mut reg = ToolRegistry::new();
reg.register(Arc::new(AddTool { base: 100 })).unwrap();
let result = reg.invoke("call_1", "add", json!({ "n": 5 })).await.unwrap();
let result = reg
.invoke("call_1", "add", json!({ "n": 5 }))
.await
.unwrap();
let value = result.output.unwrap();
assert_eq!(value["result"], 105);
assert_eq!(result.tool_call_id, "call_1");
@@ -372,8 +372,8 @@ mod tests {
#[tokio::test]
async fn test_invoke_with_permission_denied() {
let mut reg = ToolRegistry::new()
.with_permission_checker(PermissionChecker::new(Default::default()));
let mut reg =
ToolRegistry::new().with_permission_checker(PermissionChecker::new(Default::default()));
reg.register(Arc::new(ShellTool)).unwrap();
let result = reg.invoke("call_z", "shell", json!({})).await;
assert!(matches!(result, Err(ToolError::PermissionDenied(_, _))));