47 Commits
Author SHA1 Message Date
徐涛 32d886f870 feat(memory): 新增 VectorStore 抽象与 RagPipeline 持久化管线
- 新增 src/memory/vector_store.rs(约 660 行):

  - VectorStore trait:批量 add / search / remove + add_one 默认实现

  - InMemoryVectorStore:Mutex<HashMap> + 余弦全量扫描(锁内克隆、锁外计算)

  - PersistentVectorStore:MemoryStore 包装,JSON blob 持久化,先持久化后内存

  - RagPipeline:split → embed → store 组合器(具体 struct,非 trait)

- 22 个内联测试覆盖 InMemory(10)/ Persistent(6)/ RagPipeline(4)/ 性能基准(2)

- 性能断言:search 10K 条 <100ms,PersistentVectorStore::new 加载 <500ms

- 标记 VectorRetriever / InMemoryVectorRetriever 为 #[deprecated(since="0.3.0")]

- memory.rs 追加 VectorStore 等 4 个类型的 re-export

- document_demo 从手动 VectorRetriever 循环迁移到 RagPipeline 两行调用

- 零新增外部依赖
2026-07-09 16:51:29 +08:00
徐涛 b04427e83f docs(roadmap): 标记 Phase 14 完成 + M10 里程碑达成
Phase 14 全部交付物已完成(commit d4c4d8f):
- Document 类型 + RecursiveCharacterSplitter 分割器
- Embedding trait + MockEmbedding
- 19 Document 测试 + 6 Embedding 测试
- 全量测试 286 → 313,clippy 0 警告,零新外部依赖

Roadmap 同步 6 处:
- 顶部最后更新日期 + 当前状态(Phase 14 完成,Phase 15-19 待实施)
- §Phase 14 章节新增实际新增段落(commit hash + 文件清单 + 设计决策 + 测试覆盖)
- v0.3.0 依赖图 P14 节点 pending → done
- M10 里程碑  2026-07-09
- 下一步行动从 Phase 14 启动改为 Phase 15 启动
- 已完成/进行中阶段列表追加 Phase 14 完成条目
2026-07-09 13:27:36 +08:00
徐涛 d4c4d8fa3c feat(document): 实现 Document 类型与 RecursiveCharacterSplitter 分割器
新增 Phase 14 核心模块,为 RAG 管线提供 split → embed 阶段的底层支撑。

新增内容:
- Document 类型(id/content/metadata/mime_type 四字段 + new/from_raw 构造器)
- RecursiveCharacterSplitter(两阶段算法:按 separator 优先级递归分割 + 贪心合并 overlap 滑动窗口)
- Embedding trait(异步向量化抽象,复用 LlmError)+ MockEmbedding(sin-hash 零依赖伪随机实现)
- 19 个 Document 单元测试 + 6 个 Embedding 单元测试
- document_demo 示例(Document → Splitter → MockEmbedding → InMemoryVectorRetriever 端到端演示)

模块注册:
- src/lib.rs: pub mod document + pub use Document
- src/llm.rs: pub mod embedding

设计文档:docs/20-phase14-document-and-embedding.md(1417 行,含背景/调研/方案对比/实施计划)

零新外部依赖,所有长度比较以 Unicode 字符为单位(chars_len),CJK 文本行为正确。
2026-07-09 13:08:35 +08:00
徐涛 4686063ca8 docs: 标记 Phase 13 完成 + 追加 v0.3.0 CHANGELOG 条目
- CHANGELOG.md: 追加 v0.3.0 (unreleased) 条目
  - Breaking Changes: 类型路径变更(request.rs/response.rs → provider/openai.rs,
    公共 ToolChoice re-export 路径不变)+ ChatResponse / LegacyStreamEvent 删除
  - Added: ContextSlot::fork / merge + MergeStrategy 枚举(#[non_exhaustive])
  - Changed: 3 个旧 types 文件删除 + 所有 wire-format 类型迁入 openai.rs
  - Fixed: Phase 9 实施审查修复(PreRequest hook + 死代码清理 + 2 个集成测试)
  - Migration Guide: v0.2.0-rc.1 → v0.3.0 路径迁移示例
  - 修复 M9 里程碑验收 #11(CHANGELOG 条目)和审查结论 CONDITIONAL PASS 条件清单 #1
- docs/roadmap.md:
  - 顶部最后更新日期 → 2026-07-08(Phase 13 完成 + M9 里程碑达成)
  - 末尾"v0.3.0 规划完成"状态行从"Phase 13-19 待逐步实施"更新为"Phase 13 完成,
    Phase 14-19 待实施(Document → 向量存储 → 摘要 → 引擎 → 调度 → 知识图谱)"
2026-07-09 06:19:43 +08:00
徐涛 c36668071e fix(agent): Phase 9 实施审查修复
- cycle.rs: run_tool_loop 实现 PreRequest hook(之前 `let _ = hook_executor.as_ref()` 是空操作,
  导致 hook-based logging/monitoring 在流式工具循环中失效;现在与 submit_with_tools 行为对齐,
  含 should_block 检查,阻断时通过 StreamEvent::Error 事件化)
- session.rs: 删除 submit_turn_stream 末尾的 `let _ = hook_executor;` 死代码(Arc 引用生命周期
  由 Arc 自动管理)
- session.rs: 新增 2 个集成测试覆盖方案 §4 Step 5:
  - submit_turn_stream_end_to_end:mock provider → 消费流 → finalize_turn 后 cost_so_far
    正确更新(10/5 tokens)+ turn_index=1 + default slot 包含 user/assistant 消息
  - submit_turn_stream_triggers_turn_hooks:OnTurnStart 在 submit_turn_stream 返回流前
    触发(计数=1)+ OnTurnEnd 在 finalize_turn 前不触发(计数=0)+ finalize_turn 后触发(计数=1)
- docs/16-phase9-streaming-experience.md: 标注 finalize_turn Phase 10 签名变更(new_messages_from_cycle
  + Result 返回),Step 5 测试实现位置
- 测试 288 passed / 0 failed(基线 286 + 2 新增),clippy 0 警告,doc 0 warning
2026-07-08 23:19:24 +08:00
徐涛 d4f27b5865 refactor(types): 删除旧类型文件和 ChatResponse
- 删除 src/llm/types/old_stream.rs(45 行,LegacyStreamEvent 内部死代码)
- types/mod.rs 删除 pub mod old_stream; 与 ChatResponse 结构体定义
- 前置验证 A 通过:parse_chunk_stream / map_legacy_to_ir / LegacyToIrEventStream /
  ChunkToLegacyEventStream / LegacyStreamEvent 零外部调用方
- 前置验证 B 通过:cargo doc --no-deps 中无 ChatResponse 引用
- cycle.rs:88 的 #[allow(deprecated)] 保留(LlmCycle impl 中无 ChatResponse 引用,
  与 ChatResponse 无关;按方案文档"若不相关则无需改动"原则保留原状)
- stream.rs 已在 Step 13.2 commit 中同步简化为 module doc + pub use
- 同步更新 docs/roadmap.md:Phase 13 状态 ,M9 里程碑 + 完成日期,
  Mermaid 依赖图 P13 pending→done,"下一步"指向 Phase 14
- Breaking Change:agcore::llm::types::ChatResponse 已删除
  (v0.1.0 起标记 #[deprecated],请改用 MessageResponse)
2026-07-08 22:59:31 +08:00
徐涛 1c0e1e0ed1 refactor(types): response.rs 类型移入 provider/openai.rs
- 删除 types/response.rs(177 行)
- 所有 OpenAI wire-format 响应类型迁入 provider/openai.rs,可见性 pub(crate):
  TokenLogprob / TopLogprob / Logprobs / URLCitation / Annotation / OpenaiAudio /
  Choice / OpenaiChatResponse / Delta / ChunkChoice / OpenaiChatChunk
- From<OpenaiChatMessage> for Delta 与 From<OpenaiChatResponse> for OpenaiChatChunk
  同步迁入 openai.rs
- types/mod.rs 删除 pub mod response; 与对应 re-export
- convert_response 同步降级为 pub(crate) 以匹配 OpenaiChatResponse 可见性
- stream.rs: OpenaiChatChunk import 路径改为 crate::llm::provider::openai
- stream.rs 同步简化为 module doc + pub use 重导出(合并 Step 13.3 的清理动作,
  避免遗留 dead_code 警告来回)
- mod.rs: ChatResponse 的两个 From impl 同步删除(impl 内引用的 OpenaiChatResponse
  / OpenaiChatChunk / Delta / ChunkChoice 已不在 types 模块),结构体保留到 Step 13.3
- 公共 re-export 路径 agcore::llm::types::OpenaiChatResponse/Chunk 等已删除
  (Breaking Change,见 CHANGELOG)
2026-07-08 22:58:18 +08:00
徐涛 760de46623 refactor(types): request.rs 类型移入 provider/openai.rs
- 删除 types/request.rs(187 行)
- 所有 OpenAI wire-format 类型迁入 provider/openai.rs,可见性 pub(crate):
  StreamOptions / OpenaiTool / AudioParam / PredictionContent / UserLocation /
  Approximate / WebSearchOptions / OpenaiChatRequest
- types/mod.rs 删除 pub mod request; 与对应 re-export
- convert_request 同步降级为 pub(crate) 以匹配 OpenaiChatRequest 可见性
- 公共 re-export 路径 agcore::llm::types::OpenaiChatRequest 等已删除(Breaking Change,见 CHANGELOG)
2026-07-08 22:55:17 +08:00
徐涛 f8df6a9421 refactor(types): ToolChoice 移入 tool.rs
- ToolChoice 枚举与 serde impl 从 types/request.rs 迁入 types/tool.rs
- mod.rs re-export 从 request::ToolChoice 改为 tool::ToolChoice(公共 agcore::llm::types::ToolChoice 路径保持不变)
- request_v2.rs import 路径更新为 crate::llm::types::tool::ToolChoice
- OpenaiChatRequest 字段类型引用更新为 super::tool::ToolChoice(Step 13.1 删除 request.rs 后此临时 import 同步消除)
2026-07-08 22:53:46 +08:00
徐涛 802518b5fe feat(agent): 实现 ContextSlot fork/merge
- 新增 MergeStrategy 枚举(#[non_exhaustive],为 Phase 16 Summarize 预留)
- 新增 ContextSlot::fork() 从父 slot 派生独立子 slot(Full / Focused 策略)
- 新增 ContextSlot::merge() 将子 slot 消息合入父 slot(Append / Replace 策略)
- 防御性检查:禁止 self-merge、跨 session merge、合并到 Readonly slot
- 重构 AgentSession::derive_slot 复用 fork() 消除重复代码
- 新增 9 个内联测试覆盖 fork/merge happy path 与 error path
- agent.rs 追加 MergeStrategy re-export(agcore::agent::MergeStrategy 路径可用)
2026-07-08 22:52:38 +08:00
徐涛 993118f661 docs(roadmap): 将 v0.3+ 展望更新为 v0.3.0 详细规划
新增 Phase 13-19 共 7 个增量交付阶段,覆盖技术债清理、Document 系统、向量
存储持久化、摘要自动生成、Agent 执行引擎、Agent Switch 与 SubAgent 调度、
知识图谱及双通道检索。同步更新里程碑、依赖关系图及风险说明
2026-07-06 23:10:18 +08:00
徐涛 4348e4bf3e docs: 添加 LangChain & LangGraph 功能调研笔记 2026-07-06 21:55:03 +08:00
徐涛 0dc91faa43 docs(roadmap): 标记 Phase 11 测试与检索补强已完成 2026-07-06 15:30:52 +08:00
徐涛 b4e5c7d651 docs(phase11): 记录 Phase 11 方案文档与实施偏差 2026-07-06 14:53:00 +08:00
徐涛 71abe881ed feat(core): 完成 Phase 11 测试与检索补强
- 新增 VectorRetriever trait 与 InMemoryVectorRetriever 引用实现
- 补充 Provider roundtrip wiremock 测试与 MemoryStore 并发测试共 23 个
- 修复 openai 429 retry-after header 解析(与 anthropic 对齐)
2026-07-06 14:52:49 +08:00
徐涛 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
82 changed files with 17577 additions and 1764 deletions
+12
View File
@@ -198,6 +198,18 @@ pub use vector_store::VectorStore;
5. **风险评估** - 潜在风险、缓解措施 5. **风险评估** - 潜在风险、缓解措施
6. **验收标准** - 可验证的完成条件 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 / 下一步行动 / 已完成列表)。
--- ---
## 项目特定规则 ## 项目特定规则
+124
View File
@@ -2,6 +2,130 @@
本项目所有重要变更均记录于此文件。格式参考 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/)。 本项目所有重要变更均记录于此文件。格式参考 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/)。
## [0.3.0] - 未发布
v0.3.0 首个增量 Phase。技术债清理 + ContextSlot fork/merge + Phase 9 审查修复。
### Breaking Changes
**类型路径变更(0.3.0):**
- `agcore::llm::types::request::ToolChoice``agcore::llm::types::tool::ToolChoice`(公共 re-export 路径 `agcore::llm::types::ToolChoice` 保持不变)
- `agcore::llm::types::request::StreamOptions``agcore::llm::provider::openai::StreamOptions`
- `agcore::llm::types::request::OpenaiChatRequest``agcore::llm::provider::openai::OpenaiChatRequest`
- `agcore::llm::types::response::OpenaiChatResponse``agcore::llm::provider::openai::OpenaiChatResponse`
- `agcore::llm::types::response::OpenaiChatChunk``agcore::llm::provider::openai::OpenaiChatChunk`
- 其余 `request.rs`/`response.rs` 中的 wire-format 类型(`OpenaiTool``AudioParam``Choice``Delta``ChunkChoice``Annotation``Logprobs``TokenLogprob``URLCitation` 等)同步移入 `agcore::llm::provider::openai` 模块,可见性 `pub(crate)`
**类型删除:**
- `agcore::llm::types::ChatResponse` 已删除(自 v0.1.0 标记 `#[deprecated]`,请改用 `MessageResponse`
- `agcore::llm::types::old_stream::LegacyStreamEvent` 已删除(内部死代码)
**模块签名变化:**
- `LlmCycle::convert_request` / `convert_response``pub` 降级为 `pub(crate)`(因依赖的 `OpenaiChatRequest` / `OpenaiChatResponse``pub(crate)`
### Added
**Phase 13 — ContextSlot fork/merge**
- `ContextSlot::fork(child_id, strategy)` — 从父槽派生独立子槽(数据层操作,不持久化;调用方需自行 `save()`
- `ContextSlot::merge(child, strategy)` — 将子槽消息合并回父槽(`Append` 追加 / `Replace` 替换两种策略)
- `MergeStrategy` 枚举(`#[non_exhaustive]`Phase 16 可扩展 `Summarize`
- `MergeStrategy` 防御性检查:禁止 self-merge / 跨 session merge / 合并到 Readonly slot
- `agcore::agent::MergeStrategy` 公共 re-export 路径可用
- `AgentSession::derive_slot` 重构复用 `fork()` 消除重复代码(行为不变)
**Phase 9 实施审查修复(2026-07-08**
- 2 个集成测试覆盖方案 §4 Step 5:`submit_turn_stream_end_to_end` + `submit_turn_stream_triggers_turn_hooks`
### Changed
**Phase 13 — 技术债清理**
- 3 个旧 types 文件删除(`src/llm/types/request.rs` 187 行 + `response.rs` 177 行 + `old_stream.rs` 45 行)
- 所有 OpenAI wire-format 类型迁入 `provider/openai.rs`,可见性 `pub(crate)`
- `src/llm/stream.rs` 简化为 module doc + `pub use` 重导出(保持 `use crate::llm::stream::StreamEvent` 路径兼容,零下游破坏)
- `ToolChoice``request.rs` 迁入 `tool.rs`serde impl 原样搬入)
**Phase 9 实施审查修复**
- `LlmCycle::run_tool_loop` 实现 `PreRequest` hook(之前 `let _ = hook_executor.as_ref()` 是空操作,导致 hook-based logging/monitoring 在流式工具循环中失效;现在与 `submit_with_tools` 行为对齐,含 `should_block` 检查,阻断时通过 `StreamEvent::Error` 事件化)
### Fixed
**Phase 9 实施审查修复**
- `AgentSession::submit_turn_stream` 末尾 `let _ = hook_executor;` 死代码移除(Arc 引用生命周期由 Arc 自动管理)
### Migration Guide (v0.2.0-rc.1 → v0.3.0)
```rust
// ❌ v0.2.0-rc.1 — 已删除
use agcore::llm::types::ChatResponse;
use agcore::llm::types::request::OpenaiChatRequest;
// ✅ v0.3.0 — 替代路径
use agcore::llm::types::MessageResponse; // ChatResponse → MessageResponse
// OpenAI wire-format 类型为内部使用,不再公共 re-export
// 如需自定义 Provider,请直接 import agcore::llm::provider::openai::*(当前 pub(crate)
```
## [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 ## [0.1.0] - 2026-07-04
首个公开版本。涵盖 Phase 0-4c 的全部核心能力、Provider IR 重构、LlmCycle 简化,以及面向用户的 7 个离线示例。 首个公开版本。涵盖 Phase 0-4c 的全部核心能力、Provider IR 重构、LlmCycle 简化,以及面向用户的 7 个离线示例。
+5 -2
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "agcore" name = "agcore"
version = "0.1.0" version = "0.2.0-rc.1"
edition = "2024" edition = "2024"
[dependencies] [dependencies]
@@ -19,8 +19,11 @@ futures-core = "0.3"
bytes = "1" bytes = "1"
async-stream = "0.3" async-stream = "0.3"
tokio-util = { version = "0.7", features = ["rt"] } 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] [dev-dependencies]
dotenvy = "0.15.7" dotenvy = "0.15.7"
wiremock = "0.6" wiremock = "0.6"
temp-env = "0.3"
tempfile = "3"
+5 -2
View File
@@ -26,7 +26,7 @@ AG Core 不是 Agent 产品,而是 Agent 的**底层依赖库**:上层应用
```toml ```toml
[dependencies] [dependencies]
agcore = "0.1" agcore = "0.2"
tokio = { version = "1", features = ["macros", "rt-multi-thread"] } tokio = { version = "1", features = ["macros", "rt-multi-thread"] }
``` ```
@@ -110,10 +110,12 @@ let provider = create_provider(
).expect("创建 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 离线) | | `agent_session_demo` | Agent + 会话 + SessionMemory 完整链路(MockProvider 离线) |
| `custom_tool` | 自定义工具注册、单次 / 并行调用、权限检查 | | `custom_tool` | 自定义工具注册、单次 / 并行调用、权限检查 |
| `prompt_composer` | 提示词模板与组合器(纯离线) | | `prompt_composer` | 提示词模板与组合器(纯离线) |
@@ -121,6 +123,7 @@ let provider = create_provider(
| `conversation_memory_demo` | 对话记忆滑动窗口与隔离 | | `conversation_memory_demo` | 对话记忆滑动窗口与隔离 |
| `knowledge_search_demo` | 知识页面关键词检索 | | `knowledge_search_demo` | 知识页面关键词检索 |
| `streaming_events_demo` | LLM 流式响应事件消费(含错误路径) | | `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 个示例 |
+827
View File
@@ -0,0 +1,827 @@
# 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)
```
> **实施偏差(Phase 10 适配)**:实际签名扩展为
> `pub async fn finalize_turn(&mut self, response: &MessageResponse, new_messages_from_cycle: Vec<Message>) -> Result<(), AgentError>`。
> - `new_messages_from_cycle`:本轮新增消息(`[user_input, ...tool_results, final_response]`),由消费者在流消费完毕后从 `cycle.messages()[input_len..]` 提取并传入;`finalize_turn` 增量追加到当前 slot(不覆盖已有消息)。
> - 返回 `Result<(), AgentError>`:错误传播更清晰,与 `submit_turn` 的 slot 边界错误(`SlotReadonly` / `SlotNotFound`)对齐。
> - Phase 10 ContextSlot 实施时扩展。Phase 9 消费者若不接入 slot 持久化,可传 `vec![response.message.clone()]` 兜底。
### 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` 内联测试,2026-07-08 实施审查补全):
- `submit_turn_stream_end_to_end``submit_turn_stream` 跑通 mock provider → 消费流(验证收到 TextDelta + MessageComplete`finalize_turn``cost_so_far` 正确更新(`prompt_tokens=10, completion_tokens=5` + `turn_index=1` + default slot 包含 user/assistant 消息
- `submit_turn_stream_triggers_turn_hooks` — 验证 `OnTurnStart``submit_turn_stream` 返回流之前已触发(计数=1+ `OnTurnEnd``finalize_turn` 之前**不**触发(计数=0+ `finalize_turn``OnTurnEnd` 触发(计数=1
**验证**
```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
+647
View File
@@ -0,0 +1,647 @@
# Phase 11: 测试与检索补强
## 背景与目标
AG Core 当前(v0.2.0-rc.1)已完成 Phase 0-10,全量测试 254 个,clippy 0 警告,11 个离线示例全部 exit 0。功能性交付物覆盖了 LLM Cycle、Prompt、Tool、Memory、Agent Runtime、流式事件、ContextSlot 上下文管理。但有两个系统性的短板尚未补齐:
1. **检索抽象缺失**`memory` 模块只有 `MemoryRetriever`(基于 TextOverlap Dice 系数的关键词检索),缺少语义向量的检索抽象。`docs/roadmap.md` P1 中「VectorRetriever trait」一直未实现。
2. **测试覆盖缺口**Provider 的 roundtrip 测试停留在"基本响应 + 普通 401/500"层面,缺少结构化错误体解析、请求头验证、流式边界、工具调用端到端等关键场景的回归覆盖。多线程并发的 MemoryStore 测试只在 SqliteStore 有一个 10×10 场景,InMemoryStore 完全没有并发压力测试。
Phase 11 是 v0.2.0 正式版发布前的最后一个功能 Phase,三个 Step 的目标:
| Step | 内容 | 定位 |
|------|------|------|
| **11.1** | `VectorRetriever` trait + `InMemoryVectorRetriever` 引用实现 | P1 功能补全 |
| **11.2** | wiremock Provider roundtrip 测试(12 个场景) | 测试质量补强 |
| **11.3** | 并发测试补强(InMemoryStore + SqliteStore | 并发安全验证 |
最终目标:全量测试从 254 → 275+,为 v0.2.0 正式版建立更高的质量基线。
## 需求推演概要
### Step 11.1 — VectorRetriever trait
**核心需求**:定义一个与后端无关的语义检索抽象接口,包含 `index(id, embeddings)` 索引和 `search(query, k)` 检索两个方法。附带一个基于 `HashMap` 全量余弦相似度扫描的参考实现。
**边界识别**
- 只定义 trait,不绑定任何具体后端(pgvector / qdrant / lancedb 留给社区或下游)
- 不引入第三方向量数据库依赖
- 引用实现的 `search()` 不做索引加速(O(n) 全量扫描已足够验证 trait 契约)
- 不与 `MemoryStore` 耦合——`VectorRetriever` 是独立维度
- 不嵌入到 `AgentSession``ContextSlot`(Phase 11 不承担集成消费端)
**关键假设**
- `Vec<f32>` 作为 embedding 类型已足够(大部分 embedding 模型输出 f32 向量)
- 余弦相似度作为默认评分函数可覆盖主流场景
- InMemoryVectorRetriever 的 `Mutex<HashMap>` 在 ~10K 向量内性能可接受
### Step 11.2 — wiremock Provider roundtrip 测试
**核心需求**:补充 12 个 wiremock 测试,覆盖目前缺失的关键回归场景——结构化 JSON 错误体解析、请求头验证、429 限流头解析、流式边界、工具调用端到端。
**边界识别**
- 只做 HTTP mock 层验证,不做端到端 LLM 模型调用
- 每个测试自包含(启动自己的 MockServer),不抽共享 helper
- 测试集中在 OpenAI`GenericOpenaiProvider`)和 Anthropic(独立实现)两个核心 Provider 上
- DeepSeek/Qwen/Ollama 同属 OpenAI Compat,继承 `GenericOpenaiProvider` 的测试覆盖
**关键假设**
- wiremock 的 `body_partial_json` matcher 可用且稳定(当前 dev-dependencies 中已有 wiremock
- OpenAI 和 Anthropic 的结构化错误体格式在当前 SDK 版本中未变化
### Step 11.3 — 并发测试补强
**核心需求**:验证 `MemoryStore` 两种实现(InMemoryStore + SqliteStore)在多线程并发写和混合读写场景下的正确性。
**边界识别**
- 不测试 `KnowledgeStore` / `ConversationMemory` 的并发——它们的行为完全由 `MemoryStore` 决定,不引入新 race 条件
- 不测试 TTL 淘汰的并发正确性(TTL 淘汰使用 wall clock,非原子,不保证精确)
- 混合读写测试只验证"无 panic + 数量正确",不验证"读到的结果恰好与写顺序一致"(后者需要强一致快照,当前 Mutex 模型不提供)
**关键假设**
- `tokio::spawn` 100 个 task 同时写入 `Mutex<HashMap>`InMemoryStore)不会死锁
- SqliteStore 的 WAL 模式 + `busy_timeout=5000` 足够容忍 100 并发写
## 当前状态分析
### 测试覆盖率现状
| 维度 | 当前值 | Phase 11 目标 |
|------|--------|-------------|
| 全量测试 | 254 passed | 275+ passed |
| InMemoryStore 测试 | 6 个(save/get/list/upsert/eviction/TTL | +4 个并发 |
| SqliteStore 测试 | 9 个(含 1 个 10×10 并发) | +1 个 100 并发 |
| OpenAI wiremock 测试 | 4 个(basic/401/500/stream | +8 个 |
| Anthropic wiremock 测试 | 4 个(basic/401/529/stream | +4 个 |
| 请求头验证测试 | 0 个 | +2 个 |
| ToolUse 端到端 mock 测试 | 0 个 | +2 个(OpenAI + Anthropic |
### Provider 测试缺口
现有 wiremock 测试仅覆盖最基础的响应路径,以下关键场景缺失回归保护:
| 场景 | 缺失风险 |
|------|---------|
| OpenAI 请求体格式验证 | `body_partial_json` 未匹配,请求体结构变化无声 |
| Authorization header 验证 | header 注入被修改时不告警 |
| 结构化 401 JSON 错误体 | `error.message`/`error.code` 未消费,错误消息丢失 |
| 429 + `retry-after` 头 | `RateLimit.retry_after` 字段不准确 |
| ToolUse 端到端 mock | tool_flow 解析路径无回归 |
| 流式 last chunk usage-only | `{choices:[], usage:{...}}` 可能 panic |
### MemoryStore 并发测试缺口
| Store | 当前并发测试 | 覆盖度 | 风险 |
|-------|-------------|--------|------|
| InMemoryStore | 0 个 | 无 | `Mutex` 锁竞争、deadlock、写入丢失 |
| SqliteStore | 1 个(10 写者 × 10 次 = 100 条) | 中等 | `spawn_blocking` 线程池耗尽、WAL 锁等待超时 |
### 向量检索现状
`memory` 模块已有 `MemoryRetriever`(关键词检索)和 `retriever.rs` 中的 `RetrievalResult`/`ScoredItem` 类型。但语义向量检索维度完全空缺——无 trait、无引用实现、无测试。`docs/roadmap.md` 将 VectorRetriever 列为 P1,与 ContextSlotP1Phase 10)同级。
## 架构决策记录
| 决策 | 选择 | 放弃 | 理由 |
|------|------|------|------|
| 1. VectorRetriever trait 参数类型 | `Vec<f32>` 裸向量 | `Embedding` newtype | 包装类型增加可见复杂度但未提供运行时保护;大部分 embedding 模型输出 f32 向量;下游可自行包装 |
| 2. `search()` 返回类型 | `Vec<(String, f32)>` | `ScoredItem`/`RetrievalResult` 命名 struct | `(String, f32)` 是 (id, score) 的最小表达;Phase 3 的 `RetrievalResult` 绑定了 `KnowledgePage` 引用,不适合向量检索场景;tuple 在 consumer 侧模式匹配更简洁 |
| 3. 文件归属 | 新文件 `memory/vector.rs` | 合入 `memory/retriever.rs` | `retriever.rs` 已承载 302 行关键词检索代码,语义维度独立不应耦合;`vector.rs` 作为独立模块便于后期扩展(pgvector adapter 等) |
| 4. 是否附带引用实现 | `InMemoryVectorRetriever` | trait-only | trait-only 是纯推测代码,无 consumer 验证引用实现作为"编译期测试"验证 trait 方法签名可用 |
| 5. InMemoryVectorRetriever 余弦相似度实现方式 | 手动三行点积/范数 | `ndarray`/`approx` 等第三方依赖 | 余弦相似度数学固定,不需要外部依赖;1e-10 防零除;零新依赖原则 |
| 5a | 引用实现不做向量维度校验 | 运行时维度检查 | 维度校验是具体后端(pgvector等)的职责;引用实现面向测试/验证场景;调用方负责传入等长向量 |
| 6. 错误类型 | 复用 `MemoryError` 现有变体 | 新增 `VecRetrieval` 变体 | 向量检索与关键词检索语义等价于"检索";`RetrievalError` 变体已覆盖索引/评分异常场景 |
| 7. wiremock 测试组织 | 自包含(每个测试启动自己的 MockServer | 共享 helper 函数 | 沿用现有测试模式(openai.rs line 825+、anthropic.rs line 900+);自包含测试可独立运行、定位更直接 |
| 8. 请求头验证 | 做(`body_partial_json` + `header` matcher) | 跳过 | 回归防御价值高——Provider 请求体结构变化会直接导致请求被拒绝,头验证是低成本高收益的回归保护 |
| 9. 并发测试模式 | 100 并发写 + 混合读写(5 读 + 5 写)双模式 | 只做 100 并发写 | 两种模式互补:纯写入验证数据完整性和无 id 重复;混合读写验证读操作在并发写期间不 panic 且返回有效数据 |
| 10. 实施顺序 | 11.1 → 11.2 → 11.3 | 任意顺序 | 与 roadmap 原定的 Step 顺序一致;11.1 是纯新增可独立交付;11.2/11.3 是对既有代码的测试追加,可并行但不优先于 11.1 |
## 设计方案
### Step 11.1 — VectorRetriever trait + InMemoryVectorRetriever
#### 文件位置
- 新增:`src/memory/vector.rs`
- 修改:`src/memory.rs`+2 行:module 声明 + re-export
#### Trait 定义
```rust
/// 语义向量检索器抽象接口。
///
/// 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
/// 默认引用实现 [`InMemoryVectorRetriever`] 基于进程内 HashMap + 余弦相似度。
///
/// **稳定性**:实验性 APIv0.2.x),方法签名可能在 v0.3 中调整。
/// 若未来需要 `remove()` / `clear()` 等方法,将在此 trait 中追加(带默认实现)。
#[async_trait]
pub trait VectorRetriever: Send + Sync {
/// 将 `id` 对应的文本向量 `embeddings` 加入索引。
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError>;
/// 检索与 `query` 向量最相似的 `k` 条记录。
/// 返回 `Vec<(id, score)>`,按 score 降序排列,score ∈ [0.0, 1.0]。
async fn search(&self, query: Vec<f32>, k: usize) -> Result<Vec<(String, f32)>, MemoryError>;
}
```
#### InMemoryVectorRetriever 实现要点
```rust
pub struct InMemoryVectorRetriever {
vectors: Mutex<HashMap<String, Vec<f32>>>,
}
impl InMemoryVectorRetriever {
pub fn new() -> Self {
Self {
vectors: Mutex::new(HashMap::new()),
}
}
}
#[async_trait]
impl VectorRetriever for InMemoryVectorRetriever {
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError> {
let mut vectors = self.vectors.lock().map_err(|e| {
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
})?;
vectors.insert(id, embeddings);
Ok(())
}
async fn search(&self, query: Vec<f32>, k: usize) -> Result<Vec<(String, f32)>, MemoryError> {
let vectors = self.vectors.lock().map_err(|e| {
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
})?;
if vectors.is_empty() || k == 0 {
return Ok(Vec::new());
}
let query_norm = dot(&query, &query).sqrt();
if query_norm == 0.0 {
return Ok(Vec::new());
}
let mut scored: Vec<(String, f32)> = vectors
.iter()
.map(|(id, vec)| {
let dot_product = dot(&query, vec);
let vec_norm = dot(vec, vec).sqrt();
let similarity = dot_product / (query_norm * vec_norm + 1e-10);
(id.clone(), similarity)
})
.collect();
// 降序排列
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(k);
Ok(scored)
}
}
/// 点积(手动循环,零依赖)。
///
/// 注意:`zip` 对不等长向量静默截断到较短者。引用实现不做维度校验,
/// 调用方应确保 `a` 和 `b` 等长——不等长时结果无意义但不 panic。
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
```
#### 边界与约束
- **无维度校验**:不同维度向量传入 `search()` 时点积不报错,但余弦相似度结果无意义。维度校验是具体后端(pgvector等)的职责,引用实现不做运行时检查。
- **零向量处理**:query 为零向量时直接返回空结果(`query_norm == 0.0`)。
- **1e-10 防零除**:避免空库或全零向量导致除零 panic。
#### 测试(4 个)
| # | 测试名 | 验证点 |
|---|--------|--------|
| 1 | `basic_index_and_search` | index 两条("rust" + "python"),用 "rustacean" 查询应排在首位 |
| 2 | `search_empty_store` | 空库返回空 Vec |
| 3 | `concurrent_index` | 10 个 task 各 index 1 条,总量 10、id 无重复 |
| 4 | `concurrent_index_and_search` | 10 个 writer + 5 个 searcher 并发 2 秒,`search()` 遍历期间 `index()` 写入锁竞争不 panic |
#### 修改 `src/memory.rs`
```rust
pub mod vector;
// 在高频 re-export 区追加
pub use vector::{InMemoryVectorRetriever, VectorRetriever};
```
### Step 11.2 — wiremock Provider roundtrip 测试
#### 测试清单
全部 12 个测试均遵循现有自包含模式:`MockServer::start()``Mock::given(...).and(...).respond_with(...)``provider.chat_blocking(...)` / `provider.chat_stream_inner(...)` → assert。
##### P07 个)
| # | 测试名 | 所属文件 | Mock 关键点 | 断言 |
|---|--------|---------|------------|------|
| 1 | `openai_request_body_format` | `openai.rs` | `body_partial_json` 匹配 `{"model": "gpt-4o", "messages": [{"role": "user"}]}` | 请求体结构正确,响应解析正常 |
| 2 | `openai_authorization_header` | `openai.rs` | `header("authorization", "Bearer sk-test")` | header 精确匹配,响应解析正常 |
| 3 | `openai_401_structured_error` | `openai.rs` | 返回 401 + `{"error": {"message": "Incorrect API key", "code": "invalid_api_key"}}` | `LlmError::Authentication(msg)` 且 message 包含 "Incorrect API key" |
| 4 | `anthropic_401_structured_error` | `anthropic.rs` | 返回 401 + `{"error": {"type": "authentication_error", "message": "Invalid API key provided"}}` | `LlmError::Authentication(msg)` 且 message 包含 "Invalid API key" |
| 5 | `openai_429_with_retry_after` | `openai.rs` | 返回 429 + `{"error": {"message": "Rate limit exceeded"}}` + `retry-after: 30` 头 | `LlmError::RateLimit { retry_after: Some(30s) }` |
| 6 | `openai_tool_use_response` | `openai.rs` | 返回包含 `tool_calls` 的响应(choices[0].message.tool_calls ≠ null | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
| 7 | `anthropic_tool_use_response` | `anthropic.rs` | 返回含 `type: "tool_use"` content block + `stop_reason: "tool_use"`Anthropic 独立 wire 格式) | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
##### P15 个)
| # | 测试名 | 所属文件 | Mock 关键点 | 断言 |
|---|--------|---------|------------|------|
| 8 | `anthropic_version_header` | `anthropic.rs` | `header("anthropic-version", "2023-06-01")` | header 精确匹配 |
| 9 | `openai_stream_usage_only_last_chunk` | `openai.rs` | 流式最后 chunk `{"choices":[],"usage":{"prompt_tokens":5,"completion_tokens":2,"total_tokens":7}}` | 不 panic`MessageComplete` 包含正确 usage |
| 10 | `anthropic_529_overloaded_structured` | `anthropic.rs` | 返回 529 + `{"error": {"type": "overloaded_error", "message": "Overloaded"}}` | `LlmError::RateLimit { retry_after: None }` |
| 11 | `openai_500_structured_error` | `openai.rs` | 返回 500 + `{"error": {"message": "Internal server error", "type": "server_error"}}` | `LlmError::Request { status: 500, body }` 且 body 包含 "Internal server error" |
| 12 | `openai_stream_mid_stream_error` | `openai.rs` | 流式前几个 chunk 正常,中途服务端断开连接(模拟网络中断/限流断开) | `LlmError::Request(_)` — 流式中断映射为请求错误 |
#### 测试模式说明
```rust
// 每个测试自包含,不抽共享 helper(沿用现有模式)
#[tokio::test]
async fn openai_authorization_header() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(header("authorization", "Bearer sk-test"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "chatcmpl-hdr",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "OK"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
})))
.mount(&server)
.await;
let provider = GenericOpenaiProvider::new_with_name(
server.uri(), "sk-test".into(), "gpt-4o".into(), "openai", 30,
);
let response = provider.chat_blocking(MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
..Default::default()
}).await.unwrap();
assert_eq!(response.text(), "OK");
}
```
#### 新增依赖
dev-dependencies 中 wiremock 已就绪(当前 openai.rs / anthropic.rs 已在测试中使用),无需新增。
### Step 11.3 — 并发测试补强
#### 测试清单
| # | 测试名 | Store | 模式 | 验证标准 |
|---|--------|-------|------|---------|
| 1 | `concurrent_writers_max_pressure` | InMemoryStore | 100 task × 1 write | 总量 100, id 无重复 |
| 2 | `concurrent_writers_max_pressure` | SqliteStore | 100 task × 1 write | 总量 100, id 无重复 |
| 3 | `concurrent_mixed_read_write` | InMemoryStore | 预热 20 条, 5 读 + 5 写并发 2 秒 | 无 panic |
| 4 | `concurrent_mixed_read_write` | SqliteStore | 预热 20 条, 5 读 + 5 写并发 2 秒 | 无 panic |
| 5 | `concurrent_capacity_eviction` | InMemoryStore | 15 写者, max_items=10 | 最终 ≤ 10 |
#### 关键实现要点
**100 并发写模式**InMemoryStore + SqliteStore 各一):
```rust
#[tokio::test]
async fn concurrent_writers_max_pressure() {
let store = Arc::new(InMemoryStore::new());
let mut handles = Vec::new();
for i in 0..100 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let id = format!("concurrent_{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);
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
ids.sort();
ids.dedup();
assert_eq!(ids.len(), 100);
}
```
**混合读写模式**InMemoryStore + SqliteStore 各一):
```rust
#[tokio::test]
async fn concurrent_mixed_read_write() {
let store = Arc::new(InMemoryStore::new());
// 预热
for i in 0..20 {
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
}
let mut handles = Vec::new();
// 5 个写者
for w in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
let mut i = 0;
while tokio::time::Instant::now() < deadline {
let id = format!("writer{w}_item{i}");
s.save(make_item(&id)).await.unwrap();
i += 1;
}
}));
}
// 5 个读者
for r in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
while tokio::time::Instant::now() < deadline {
let _ = s.list(&MemoryFilter::default()).await.unwrap();
}
}));
}
for h in handles {
h.await.unwrap();
}
// 不 panic 即算通过
}
```
**容量淘汰并发模式**InMemoryStore):
```rust
#[tokio::test]
async fn concurrent_capacity_eviction() {
let eviction = EvictionConfig {
policy: EvictionPolicy::Capacity { max_items: 10 },
check_interval: 1,
};
let store = Arc::new(InMemoryStore::with_eviction(eviction));
let mut handles = Vec::new();
for i in 0..15 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
s.save(make_item(&format!("item_{i}"))).await.unwrap();
}));
}
for h in handles {
h.await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert!(list.len() <= 10);
}
```
**测试归属**
- InMemoryStore 并发测试 → `src/memory/store/in_memory.rs``mod tests`
- SqliteStore 并发测试 → `src/memory/store/sqlite_store.rs``mod tests`
- 混合读写测试中的 `make_item` 辅助函数:直接复用各文件现有 `fn make_item`
## 已否决的方案
### 1. 砍掉 11.1PM 建议)
**内容**PM 在讨论中提出 VectorRetriever trait 无消费者,建议整体砍掉,等 Phase 12 或 v0.3 有人用时再做。
**否决理由**:引用实现作为 trait 契约的编译期验证手段——没有 consumer 不意味着 trait 签名不需要测试。同时社区贡献(pgvector adapter 等)需要稳定的 trait 边界。附带引用实现还可作为"如何在 agcore 中实现一个 VectorRetriever"的示例,降低社区参与门槛。代码量仅 ~80 行,维护成本可忽略。
### 2. trait-only VectorRetriever(无引用实现)
**内容**:只定义 `VectorRetriever` trait,不做 `InMemoryVectorRetriever`
**否决理由**trait-only 是纯推测代码——没有运行时验证,无法确认 trait 方法签名在实际调用链中是否可编译。理想情况下每个 trait 至少有一个引用实现来验证"这个 trait 确实可以被实现"。
### 3. 跳过请求头验证
**内容**:请求体格式和 Authorization header 验证是"过度保护"。
**否决理由**`body_partial_json` + `header` matcher 的回归防御价值高。Provider 适配层的最大风险是请求体结构无声变更(如 `ToolDef` IR 切换时漏改了序列化字段),头验证是低成本(每测试 ~5 行)高收益的回归保护。
### 4. 先发 v0.2.0 正式版再迭代
**内容**:当前 rc.1 已经包含所有 P0 功能,建议直接发正式版,Phase 11 推迟到 v0.2.1。
**否决理由**:测试补强是正式版的信号而非负担。Phase 11 的三个 Step 都是"如果现在不做,以后更不会做"的类型。在正式版前补齐测试基线,避免「发布了再补测试」的经典陷阱。
### 5. 只做 100 并发写,不做混合读写
**内容**:并发写验证数据完整性已足够。
**否决理由**:纯写入和混合读写暴露不同类型的 bug。纯写入验证"数据不丢、id 无重复";混合读写验证"读操作在并发写期间不 panic、返回有效数据"。两种模式互补缺失。
### 6. (实施后补充)openai_stream_mid_stream_error 的 mock 模式偏差
**实际实施**:返回 `200 + SSE content-type + 畸形 JSON payload``data: {not-valid-json}\n\ndata: [DONE]\n\n`),断言 `ChunkToEventStream` 产出 `StreamEvent::Error`
**方案原文**:返回"前几个 chunk 正常,中途服务端断开连接",断言 `LlmError::Request(_)`
**偏差原因**wiremock 0.6 标准 responder 的 `set_delay` / `set_body_string` 行为是「延迟响应 + 发完 body 后关闭连接」,无法精确模拟「send partial body then hang 保持连接」。`read_timeout` 配合 `set_delay` 触发的超时属于 send 阶段(`LlmError::Timeout`),不属于流式中断。
**采纳方案**:用畸形 JSON payload 替代——同样验证"流中途产生错误事件而不 panic"的回归保护意图,且在 wiremock 0.6 上 100% 可重现。断言改为 `StreamEvent::Error{message}` + "stream 最终结束",保持核心回归价值。
**影响**:测试意图(流阶段错误检测)完全保留;mock 行为从「TCP 断开」变为「畸形应用层数据」;断言从 `LlmError::Request` 改为 `StreamEvent::Error`(语义等价:客户端发现流异常)。
### 7. (实施后补充)openai.rs 的 429 retry-after 解析修复
**实际实施**`handle_error_response` 新增 `retry-after` header 解析逻辑(与 anthropic.rs 完全对齐)。
**方案原文**:方案测试 11.2.4 要求 `RateLimit { retry_after: Some(30s) }`,但 ADC 表中未显式列出此修复作为生产代码变更。
**修复原因**:原 `openai.rs:186-191` 的 429 分支固定 `retry_after: None`,与 `anthropic.rs:339-351` 已有的解析逻辑不一致。原代码注释甚至已写"仅读取 retry-after",但实际未实现——这是隐藏 bug。修复让 OpenAI 兼容层(DeepSeek/Qwen 等)的限流重试信息可用,与 Anthropic 行为统一。
**影响**:方案测试 11.2.4 从「不可通过的回归保护」变为「可验证的实际行为」。变更 5 行,与 anthropic 实现完全镜像。
## 实施计划与顺序
**实施顺序**11.1 → 11.2 → 11.3(与 roadmap 一致,每步可单独交付验证)。
| Step | 文件变更 | 测试增量 | 预估代码量 | 验证标准 |
|------|---------|---------|-----------|---------|
| 11.1 | +`src/memory/vector.rs`~80 行),~`src/memory.rs`+2 行) | +4 | ~85 行实现 + 70 行测试 | `cargo build --all-targets` 编译通过,4 个测试通过 |
| 11.2 | ~`src/llm/provider/openai.rs`+6 个测试),~`src/llm/provider/anthropic.rs`+2 个测试) | +127 P0 + 5 P1 | ~290 行(含 test mod 和 mock 数据) | `cargo test --all-targets` 全绿,wiremock 12 场景均绿 |
| 11.3 | ~`src/memory/store/in_memory.rs`+3 个测试),~`src/memory/store/sqlite_store.rs`+2 个测试) | +5 | ~120 行 | `cargo test --all-targets` 全绿 |
| **总计** | 6 个文件(1 新增 + 5 修改) | +21 | ~500 行 | 全量 254 → 275+`cargo clippy --all-targets -- -D warnings` 0 警告 |
### 验证通过标准
1. `cargo build --all-targets` —— 编译通过,无 warning
2. `cargo test --all-targets` —— 全部通过(254 + 21 = 275+
3. `cargo clippy --all-targets -- -D warnings` —— 0 警告
4. `cargo test --all-targets 2>&1 | grep -E "test result:"` —— 确认新增测试全部出现在执行列表中
5. 新增 wiremock 测试单独验证网络隔离(无需 API key,纯本地 mock
## 参考来源
- **讨论收口结论**Phase 11 讨论,含 PM/SA 双视角输入(2026-07-07
- **现有代码模式**
- `src/llm/provider/openai.rs` line 824-1034——wiremock 测试模式(`MockServer::start``Mock::given(...).and(...).respond_with(...)``provider.chat_blocking` → assert
- `src/llm/provider/anthropic.rs` line 899-1079——Anthropic provider wiremock 测试
- `src/memory/store/sqlite_store.rs` line 458-483——`concurrent_writers_no_data_loss` 10×10 并发模式
- `src/memory/store.rs`——`MemoryStore` trait 定义(`#[async_trait]` 风格)
- `src/memory/error.rs`——`MemoryError` 枚举(`#[non_exhaustive]` + `RetrievalError` 变体)
- `src/memory/retriever.rs`——现有检索模块(`RetrievalResult` / `ScoredItem`
- `src/memory.rs`——模块根与 re-export 模式
- **方案文档**`docs/roadmap.md` Phase 11 章节(line 516-528
- **编译器 pragma**`#[non_exhaustive]` —— 新增枚举变体需要此标记,公开结构体字段未来变化预留兼容空间
## 关键假设与风险
### 关键假设清单
| # | 假设 | 影响 | 推翻后的应对 |
|---|------|------|------------|
| 1 | wiremock `body_partial_json` matcher 在 wiremock 0.6+ 中可用 | Step 11.2 测试 #1 的实现方式 | 改用 `body_json`(精确匹配)或 `body_string`(部分串匹配) |
| 2 | `tokio::spawn` 100 task 并发写入 `Mutex<HashMap>` 无死锁 | Step 11.3 #1 InMemoryStore 并发 | 降低并发数(50)继续验证,或换 `tokio::sync::Mutex` |
| 3 | SqliteStore 的 `busy_timeout=5000` 能容忍 100 并发写 | Step 11.3 #2 SqliteStore 并发 | 增加 `busy_timeout`10s),或限制最大并发数 |
| 4 | 新增测试不使用 wiremock 以外的未列在 dev-dependencies 中的依赖 | Phase 11 零新增外部依赖 | 若需要额外 matcher,评估后加入 dev-dependencies |
| 5 | InMemoryVectorRetriever 的 O(n) 全量扫描在测试规模下 (<1000 向量) 性能可接受 | Step 11.1 测试通过 | 如果有竞态问题,改为 read/write lock`RwLock<HashMap>` |
| 6 | `MemoryError::RetrievalError` 变体足够覆盖向量检索的索引/评分失败场景 | Step 11.1 错误映射 | 如果不够,可增加新的 `MemoryError` 变体 |
| 7 | OpenAI 和 Anthropic 的结构化错误体格式在当前 SDK 版本中未变化 | Step 11.2 测试 #3/#4/#7/#10/#11(结构化错误解析断言) | 若 SDK 变更错误体格式,更新 mock body 和断言匹配新格式 |
| 8 | 调用方传入 `search()` 的向量与已索引向量维度一致(引用实现不做维度校验,`dot()` 对不等长向量静默截断) | Step 11.1 InMemoryVectorRetriever 正确性 | 若须维度校验,在 `index()` 时记录维度并在 `search()` 时断言;引用实现维持零校验 |
### 已识别的风险
| 风险 | 等级 | 缓解措施 |
|------|------|---------|
| wiremock `body_partial_json` matcher 行为在版本升级后变化 | 低 | 限定 wiremock 版本范围(当前已在 Cargo.lock 中锁定);P0 测试不依赖该 matcher |
| 100 并发写暴露 SqliteStore 的 `spawn_blocking` 线程池瓶颈 | 中 | 观察 CI 执行时间;如果超时,降低并发到 50 或增加 `max_blocking_threads` |
| InMemoryVectorRetriever 的 `Mutex` 锁争用导致测试 flaky | 低 | Mutex 不会死锁(单线程持有不 await),测试不依赖精确时序 |
| 新增 wiremock 测试与现有测试冲突(端口占用) | 低 | `MockServer::start()` 自动选择随机端口,不冲突 |
| `cargo test --all-targets` 执行时间增加 >30% | 低 | 预估 +21 个测试,增量约 8%254→275),其中 wiremock 测试有网络 IO 但延迟 <10ms/个 |
### 非阻塞已知项
- **Ollama Provider** 是 OpenAI Compatwiremock 测试继承 `GenericOpenaiProvider`,不单独新增
- **DeepSeek / Qwen Provider** 同样通过 `GenericOpenaiProvider` 实现,继承测试
- **Step 11.1 不消费到 AgentSession / ContextSlot**,留待 Phase 12 或 v0.3 做消费端集成
- **Phase 11 完成后**,全量测试预计 275+,`cargo test --all-targets` 执行时间预计 < 60s
---
## 实施计划(附录)
**实施顺序**11.1 → 11.2 ‖ 11.311.2 与 11.3 无文件冲突,可并行交付;11.2 优先因回归保护价值更高)。每步完成后运行 `cargo test --all-targets` + `cargo clippy --all-targets -- -D warnings` 验证无回归。
---
### Step 11.1 — VectorRetriever trait + InMemoryVectorRetriever~160 行,4 测试)
**前置**:无。与 Step 11.2/11.3 可并行开发但优先交付。
| 任务 | 描述 | 文件 | 前置 | 工作量 | 风险 | 验收条件 |
|------|------|------|------|--------|------|---------|
| **11.1.1** | 创建 `memory/vector.rs`:定义 `VectorRetriever` trait + `InMemoryVectorRetriever` struct + `dot()` 辅助函数 | `src/memory/vector.rs` | 无 | S | 低 | `cargo build --all-targets` 编译通过 |
| **11.1.2** | 实现 `InMemoryVectorRetriever``index()` — Mutex insert`search()` — 全量余弦扫描 + 降序排列 + k 截断 | `src/memory/vector.rs` | 11.1.1 | S | 低 | trait 实现编译通过;`Mutex::lock()` 使用 `map_err` 处理 poison,不 panic |
| **11.1.3** | 添加 4 个内联测试:`basic_index_and_search``search_empty_store``concurrent_index``concurrent_index_and_search` | `src/memory/vector.rs` (mod tests) | 11.1.2 | S | 低 | 4 测试全部通过 |
| **11.1.4** | 修改 `src/memory.rs`:加 `pub mod vector;` + `pub use vector::{VectorRetriever, InMemoryVectorRetriever};` | `src/memory.rs` | 11.1.1 | S | 低 | `cargo build --all-targets` 无 warning |
**Step 验证**
```
cargo test --all-targets # 254 + 4 = 258+ passed
cargo clippy --all-targets -- -D warnings # 0 warning
```
---
### Step 11.2 — wiremock Provider roundtrip 测试(12 测试,P0=7 + P1=5
**前置**:无。与 Step 11.1 无文件冲突,可并行。
**`openai.rs` 新增测试(8 个:P0=5 + P1=3**
| 任务 | 测试名 | 优先级 | 前置 | 工作量 | 风险 | Mock 模式 | 验收条件 |
|------|--------|--------|------|--------|------|----------|---------|
| **11.2.1** | `openai_request_body_format` | P0 | 无 | S | 低 | `body_partial_json` 匹配 model/messages | 请求体结构正确,响应解析正常 |
| **11.2.2** | `openai_authorization_header` | P0 | 无 | S | 低 | `header("authorization", "Bearer sk-test")` | header 精确匹配 |
| **11.2.3** | `openai_401_structured_error` | P0 | 无 | S | 低 | 401 + `{"error":{"message":"...","code":"invalid_api_key"}}` | `LlmError::Authentication` 含 "Incorrect API key" |
| **11.2.4** | `openai_429_with_retry_after` | P0 | 无 | S | 低 | 429 + `retry-after: 30` | `RateLimit { retry_after: Some(30s) }` |
| **11.2.5** | `openai_tool_use_response` | P0 | 无 | S | 低 | 响应含 `tool_calls` | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
| **11.2.6** | `openai_stream_usage_only_last_chunk` | P1 | 无 | S | 低 | 流式最后 chunk `{choices:[], usage:{...}}` | 不 panic`MessageComplete` 含正确 usage |
| **11.2.7** | `openai_500_structured_error` | P1 | 无 | S | 低 | 500 + `{"error":{"message":"server error"}}` | `Request { status: 500 }` 含 body |
| **11.2.8** | `openai_stream_mid_stream_error` | P1 | 无 | M | 中 | 前几个 chunk 正常后连接断开 | `LlmError::Request(_)` 流中断映射 |
**`anthropic.rs` 新增测试(4 个:P0=2 + P1=2**
| 任务 | 测试名 | 优先级 | 前置 | 工作量 | 风险 | Mock 模式 | 验收条件 |
|------|--------|--------|------|--------|------|----------|---------|
| **11.2.9** | `anthropic_401_structured_error` | P0 | 无 | S | 低 | 401 + `{"error":{"type":"authentication_error","message":"..."}}` | `LlmError::Authentication` 消息透传 |
| **11.2.10** | `anthropic_tool_use_response` | P0 | 无 | S | 低 | 响应含 `type:"tool_use"` content block + `stop_reason:"tool_use"` | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
| **11.2.11** | `anthropic_version_header` | P1 | 无 | S | 低 | `header("anthropic-version", "2023-06-01")` | header 精确匹配 |
| **11.2.12** | `anthropic_529_overloaded_structured` | P1 | 无 | S | 低 | 529 + `{"error":{"type":"overloaded_error","message":"Overloaded"}}` | `RateLimit { retry_after: None }` |
**Step 验证**
```
cargo test --all-targets # 258 + 12 = 270+ passed
cargo clippy --all-targets -- -D warnings # 0 warning
```
每个测试自包含(`MockServer::start()``Mock::given(...)``provider.chat_blocking()/chat_stream_inner()` → assert),无需共享 helper。
---
### Step 11.3 — 并发测试补强(5 测试)
**前置**:无。与 Step 11.1/11.2 无文件冲突。
| 任务 | 测试名 | 文件 | 前置 | 工作量 | 风险 | 模式 | 验收条件 |
|------|--------|------|------|--------|------|------|---------|
| **11.3.1** | `concurrent_writers_max_pressure` | `in_memory.rs` | 无 | S | 中 | 100 task × 1 write | 总量 100id 无重复 |
| **11.3.2** | `concurrent_writers_max_pressure` | `sqlite_store.rs` | 无 | S | 中 | 100 task × 1 write | 总量 100id 无重复 |
| **11.3.3** | `concurrent_mixed_read_write` | `in_memory.rs` | 无 | S | 中 | 预热 20 条,5 写 + 5 读并发 2 秒 | 无 panic |
| **11.3.4** | `concurrent_mixed_read_write` | `sqlite_store.rs` | 无 | S | 中 | 预热 20 条,5 写 + 5 读并发 2 秒 | 无 panic |
| **11.3.5** | `concurrent_capacity_eviction` | `in_memory.rs` | 无 | S | 中 | 15 写者,max_items=10 | 最终 ≤ 10(竞争激烈时可能过渡态 >10,主断言 ≤ 10,宽松备选 ≤ 15) |
**实现要点**
- 沿用现有 `Arc<Store> + tokio::spawn + h.await.unwrap()` 模式(参考 `sqlite_store.rs:458-483`
- `make_item` 辅助函数直接复用各文件现有实现
- SqliteStore 测试使用 `:memory:` 数据库(与现有并发测试一致)
- 混合读写测试使用 `tokio::time::Instant::now() + Duration` 做时限
- 100 并发写是一次性 spawn 100 task(非分批),暴露最大锁竞争压力
**Step 验证**
```
cargo test --all-targets # 270 + 5 = 275+ passed
cargo clippy --all-targets -- -D warnings # 0 warning
```
---
### 整体发布核查清单
| # | 检查项 | 验证命令 | 预期结果 |
|---|--------|---------|---------|
| 1 | 编译 | `cargo build --all-targets` | 通过,0 warning |
| 2 | 全量测试 | `cargo test --all-targets` | 275+ passed0 failed |
| 3 | Lint | `cargo clippy --all-targets -- -D warnings` | 0 warning |
| 4 | 文档 | `cargo doc --no-deps` | 0 warningVectorRetriever trait 公共 API doc 完整) |
| 5 | 确认新增测试 | `cargo test --all-targets 2>&1 \| grep -E "test result:"` | 所有新增测试名出现在执行列表中 |
| 6 | wiremock 隔离 | 新增 wiremock 测试不依赖网络 | 纯本地 mock,无需 API key |
| 7 | 并行安全 | 并发测试独立运行时无 flaky | 连续 3 次 `cargo test` 结果一致 |
| 8 | 存量零回归 | 已有 254 个测试全部通过 | 与 Phase 10 基线对比无 fail |
| 9 | 公共 API doc comment | `grep -r "pub trait VectorRetriever" src/ && rg "^///" -c src/memory/vector.rs` | trait 和方法都有 `///` 注释 |
**若核查项失败的回退策略**
- **测试失败(P0)**:阻断发布。定位到具体测试名 → 检查 Mock JSON 格式与 Provider 解析逻辑是否匹配(结构化错误体格式变化 → 更新 mock body;流式状态机变化 → 更新 `chat_stream_inner` 路径测试)
- **测试失败(P1)**:不阻断发布。标记 `#[ignore]` + file issue,确认无 P0 失败后即可发布
- **clippy warning**:修复 lint 后重跑;若为 `#[allow(...)]` 可抑制,在 code review 中申明理由
- **flaky 并发测试**:检查 `tokio::spawn` 是否跨 `.await` 持锁;若 SqliteStore 超时,增加 `busy_timeout` 或降低并发数
- **11.1 模块发布阻塞**:若 `InMemoryVectorRetriever` 无法按时交付,可临时注释 `src/memory.rs` 中的 `pub mod vector;` 行,跳过整个模块(零消费者,不影响发布)。回退后再补交
+640
View File
@@ -0,0 +1,640 @@
# Phase 13 — 热身清理 + ContextSlot fork/merge 实施方案
- **文档编号**19
- **标题**Phase 13 — 热身清理 + ContextSlot fork/merge 实施方案
- **日期**2026-07-08
- **状态**:待实施
- **涉及模块**agent/context、agent/session、llm/types、llm/provider/openai、llm/stream
- **关联文档**roadmap.md(§Phase 13)、17-phase10-contextslot.md
- **对应**Roadmap §Phase 13v0.3.0 第一阶段)
---
## 1. 背景与目标
v0.3.0 是 agcore 从"LLM 调用工具箱"升级为"多 Agent 基础系统"的关键版本。Phase 13 是 v0.3.0 的第一阶段,定位为"热身",包含两大部分:
- **技术债清理**:删除 Phase 0 遗留的旧 types 文件(`request.rs``response.rs``old_stream.rs`),以及已标记 `#[deprecated]``ChatResponse` 结构体
- **ContextSlot fork/merge**:为 ContextSlot 增加分叉和合并能力,为后续 Phase 17 Checkpointer 和 Phase 18 SubAgent Dispatch 打基础
**依赖关系**:无(独立交付)
**优先级**P0
**预估规模**:净减 ~200 行代码(新增 ~505 行,删除 ~704 行)
---
## 2. 需求分析
### 2.1 功能需求
1. **技术债清理**:删除 `src/llm/types/request.rs`187 行)、`response.rs`177 行)、`old_stream.rs`45 行),将其中的 OpenAI wire-format 类型移入 `src/llm/provider/openai.rs`;删除 `types/mod.rs` 中的 `ChatResponse` 废弃结构体
2. **`ContextSlot::fork`**:从现有 context slot 分支出独立的子 slot
3. **`ContextSlot::merge`**:将子 slot 的消息合并回父 slot
4. **`MergeStrategy`** 枚举:Append(追加)/ Replace(替换),`#[non_exhaustive]` 预留 Phase 16 Summarize 扩展
### 2.2 非功能需求
- **每步可编译**:5 个 Step 按物理文件切割,每步 `cargo build --all-targets + cargo test` 验证
- **指定公共 API 路径保持向后兼容**:`agcore::llm::types::ToolChoice`re-export 不变)、`crate::llm::stream::StreamEvent`(重导出保留);其余 wire-format 类型(`OpenaiChatRequest``OpenaiChatResponse/Chunk``StreamOptions` 等)移入 `provider/openai.rs` 后属 Breaking Change,详见 §4.3 CHANGELOG
- **向后兼容的 StreamEvent 路径**`crate::llm::stream::StreamEvent` 重导出保留,不修改 `cycle.rs``session.rs` 的 import
---
## 3. 方案设计
### 3.1 整体架构
Phase 13 分为 5 个 Step,按执行顺序排列:
```
Step 13.5 (fork/merge) → Step 13.4 (ToolChoice) → Step 13.1 (request types) → Step 13.2 (response types) → Step 13.3 (cleanup)
```
这种顺序的好处:
- **先交付价值**:13.5 是唯一有用户功能交付的 Step,先做建立节奏
- **排序约束**13.4 必须先于 13.1ToolChoice 不搬走,request.rs 不能删)
- **13.3 收尾**:删除旧文件和 `ChatResponse` 是 breaking change,放在最后
### 3.2 Step 13.5 — ContextSlot fork/merge
#### MergeStrategy 枚举
定义在 `src/agent/context.rs`
```rust
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum MergeStrategy {
/// 子 slot 消息追加到父 slot 末尾。
Append,
/// 用子 slot 消息替换父 slot 内容。
Replace,
}
```
- `#[non_exhaustive]` 保证 Phase 16 加入 `Summarize` 变体时不破坏现有代码
- 不预埋 `Summarize` 占位变体(YAGNI 原则)
#### ContextSlot::fork
```rust
impl ContextSlot {
pub fn fork(&self, child_id: String, strategy: DeriveStrategy) -> ContextSlot {
let messages = match &strategy {
DeriveStrategy::Full => self.messages.clone(),
DeriveStrategy::Focused(cfg) => Self::filter_focused(&self.messages, cfg),
};
tracing::debug!(
parent_id = %self.id,
child_id = %child_id,
?strategy,
"ContextSlot::fork"
);
ContextSlot {
id: child_id,
session_id: self.session_id.clone(),
config: SlotConfig {
mode: match &strategy {
DeriveStrategy::Full => SlotMode::Full,
DeriveStrategy::Focused(cfg) => SlotMode::Focused(cfg.clone()),
},
source: SlotSource::Derived {
parent_id: self.id.clone(),
strategy,
},
budget: self.config.budget.clone(),
compact: self.config.compact,
},
messages,
meta: SlotMeta::new(),
}
}
}
```
设计要点:
- 纯数据层操作,不持久化
- 子 slot 的 `meta` 全新创建(`SlotMeta::new()`),不继承父 slot 的 message_count
- 子 slot 的 source 记录 `parent_id`,血缘可追溯
- 添加 `tracing::debug!` 日志,支持多 slot 交互场景的审计追踪
#### ContextSlot::merge
```rust
impl ContextSlot {
/// 将子 slot 的消息合并到当前 slot。
///
/// **注意**:本方法仅操作内存数据,不自动持久化。
/// 调用方需在 merge 后自行调用 `self.save(&store)` 将结果写入后端存储。
pub fn merge(&mut self, child: ContextSlot, strategy: MergeStrategy) -> Result<(), AgentError> {
// 防御性检查
if self.id == child.id {
return Err(AgentError::Config("不能将 slot 合并到自身".into()));
}
if self.session_id != child.session_id {
return Err(AgentError::Config("不能合并不同 session 的 slot".into()));
}
if matches!(self.config.mode, SlotMode::Readonly) {
return Err(AgentError::SlotReadonly("Readonly slot 不允许合并".into()));
}
tracing::debug!(
self_id = %self.id,
child_id = %child.id,
?strategy,
"ContextSlot::merge"
);
match strategy {
MergeStrategy::Append => {
let count = child.messages.len();
self.messages.extend(child.messages);
self.meta.message_count += count;
}
MergeStrategy::Replace => {
self.messages = child.messages;
self.meta.message_count = self.messages.len();
}
}
Ok(())
}
}
```
#### AgentSession::derive_slot 重构
现有 `derive_slot`session.rs:213-260)的手工复制代码改为调用 `parent.fork()`
```rust
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 child = parent.fork(slot_id.clone(), strategy); // ← 用 fork
child.save(&*self.resolve_store()).await?;
self.slots.insert(slot_id, child);
Ok(())
}
```
重复检查、查找父 slot 的代码不变;消息复制逻辑委托给 `fork()`
#### 测试计划(新增 9 个)
| 测试名 | 验证点 |
|--------|--------|
| `fork_full_copies_messages` | fork Full 策略复制父 slot 全部消息 |
| `fork_focused_filters_messages` | fork Focused 策略按 config 过滤 |
| `fork_preserves_independence` | 父 slot 追加消息不影响子 slot |
| `fork_sets_derived_source` | 子 slot source 正确记录 parent_id |
| `merge_append_appends_messages` | Append 追加到父 slot 末尾,message_count 正确 |
| `merge_replace_replaces_messages` | Replace 替换父 slot 消息,message_count 正确 |
| `merge_self_rejected` | self-merge 返回 `Err` |
| `merge_readonly_rejected` | 合并到 Readonly slot 返回 `Err` |
| `merge_cross_session_rejected` | 跨 session 合并返回 `Err` |
### 3.3 Step 13.4 — ToolChoice 移入 tool.rs
#### 变更文件
| 文件 | 变更 |
|------|------|
| `src/llm/types/request.rs` | 删除 `ToolChoice` 枚举 + serde impl~28-99 行) |
| `src/llm/types/tool.rs` | 新增 `ToolChoice` 枚举 + serde impl(原样搬入) |
| `src/llm/types/mod.rs` | `pub use request::{..., ToolChoice}``pub use tool::ToolChoice` |
| `src/llm/types/request_v2.rs` | import 路径 `request::ToolChoice``tool::ToolChoice` |
**import 路径变化**
| 当前 | 移动后 |
|------|--------|
| `crate::llm::types::request::ToolChoice` | `crate::llm::types::tool::ToolChoice` |
| `crate::llm::types::ToolChoice`(通过 re-export | `crate::llm::types::ToolChoice`(通过 tool.rs re-export,保持不变) |
**验证**`cargo build --all-targets` + `cargo test` + `cargo clippy`
### 3.4 Step 13.1 — request.rs 类型移入 openai.rs
#### 变更文件
| 文件 | 变更 |
|------|------|
| `src/llm/types/request.rs` | **整文件删除**187 行) |
| `src/llm/provider/openai.rs` | 新增 `StreamOptions``OpenaiTool``AudioParam``PredictionContent``UserLocation``Approximate``WebSearchOptions``OpenaiChatRequest` 等类型定义 |
| `src/llm/types/mod.rs` | 删除 `pub use request::{OpenaiChatRequest, OpenaiTool, StreamOptions}`;删除 `pub mod request;` |
| `src/llm/provider/openai.rs` import 调整 | 原 `use crate::llm::types::request::{...}` 改为从同级 `use super::super::types::...` 或直接使用本文件内类型 |
**注意**`OpenaiTool` 引用 `OpenaiToolDefinition`(定义在 `tool.rs`),移入 `openai.rs` 后需通过 `crate::llm::types::tool::OpenaiToolDefinition` 引用。`OpenaiChatRequest.messages` 字段引用 `OpenaiChatMessage`(定义在 `openai_message.rs`),路径不变。
**设计决策**:搬入 `openai.rs` 后的类型可见性可降级为 `pub(crate)`。它们是与 OpenAI wire-format 绑定的内部序列化类型,公共 API 消费者不应直接接触。
**验证**`cargo build --all-targets` + `cargo test` + `cargo clippy`
### 3.5 Step 13.2 — response.rs 类型移入 openai.rs
#### 变更文件
| 文件 | 变更 |
|------|------|
| `src/llm/types/response.rs` | **整文件删除**177 行) |
| `src/llm/provider/openai.rs` | 新增 `TokenLogprob``TopLogprob``Logprobs``URLCitation``Annotation``OpenaiAudio``Choice``OpenaiChatResponse``Delta``ChunkChoice``OpenaiChatChunk` + `From<OpenaiChatMessage> for Delta` + `From<OpenaiChatResponse> for OpenaiChatChunk` |
| `src/llm/types/mod.rs` | 删除 `pub use response::{...}`;删除 `pub mod response;` |
| `src/llm/stream.rs:26` | 将 `use crate::llm::types::{OpenaiChatChunk, OpenaiToolCall}` 中的 `OpenaiChatChunk` 路径改为 `crate::llm::provider::openai::OpenaiChatChunk``OpenaiToolCall` 保持从 `tool.rs` |
**验证**`cargo build --all-targets` + `cargo test` + `cargo clippy`
### 3.6 Step 13.3 — 旧文件清理 + ChatResponse 删除
#### 13.3a — 删除 `old_stream.rs`
> **前置验证**:实施前执行 `grep -rn 'parse_chunk_stream\|map_legacy_to_ir\|LegacyToIrEventStream\|ChunkToLegacyEventStream' src/` 确认零外部调用方,记录结果到实施 commit。
| 文件 | 变更 |
|------|------|
| `src/llm/types/old_stream.rs` | **整文件删除**45 行,`LegacyStreamEvent` |
| `src/llm/types/mod.rs` | 删除 `pub mod old_stream;` |
| `src/llm/stream.rs` | 删除 `use crate::llm::types::old_stream::LegacyStreamEvent`;删除 `parse_chunk_stream``parse_chunk_stream_legacy``ChunkToLegacyEventStream``LegacyToIrEventStream``map_legacy_to_ir``empty_message_response`~160 行死代码) |
**stream.rs 最终形态**
```rust
//! 流式事件系统 —— 重导出 StreamEvent 供向后兼容。
pub use crate::llm::types::response_v2::StreamEvent;
```
**为什么不全删 stream.rs**`cycle.rs``session.rs``use crate::llm::stream::StreamEvent` 路径保持不变。全删 + 改所有 import 路径的改动量 > 收益。保留 1 行重导出就够。
#### 13.3b — 删除 `ChatResponse`
| 文件 | 变更 |
|------|------|
| `src/llm/types/mod.rs` | 删除 `ChatResponse` 结构体定义 + 两个 `#[allow(deprecated)]` `From` impl`From<OpenaiChatResponse> for ChatResponse``From<ChatResponse> for OpenaiChatChunk` |
`ChatResponse` 自 v0.1.0 起标记 `#[deprecated]`v0.2.0-rc.1 阶段直接删除即可。删除前运行 `cargo doc --no-deps 2>&1 | grep -i 'ChatResponse'` 确认零文档引用。
**验证**`cargo build --all-targets` + `cargo test` + `cargo clippy` + `cargo doc --no-deps`
---
## 4. 实现计划
### 4.1 实施顺序总览
```
Step 13.5 ──→ Step 13.4 ──→ Step 13.1 ──→ Step 13.2 ──→ Step 13.3
(fork/merge) (ToolChoice) (request) (response) (cleanup)
│ │ │ │ │
▼ ▼ ▼ ▼ ▼
+60 行净增 -0 净增 -0 净增 -0 净增 -260 删除
+9 个测试 import 路径 纯类型搬移 纯类型搬移 +1 行重导出
变更
```
### 4.2 各 Step 文件变更清单
#### Step 13.5 — ContextSlot fork/merge
| 操作 | 文件 | 变更说明 |
|------|------|---------|
| 新增 | `src/agent/context.rs` | `MergeStrategy` 枚举 + `ContextSlot::fork()` + `ContextSlot::merge()` |
| 重构 | `src/agent/session.rs` | `derive_slot` 改为调用 `parent.fork()` |
| 新增 | 内联测试 | 9 个新测试(fork, merge, 边界) |
#### Step 13.4 — ToolChoice 移动
| 操作 | 文件 | 变更说明 |
|------|------|---------|
| 删除 | `src/llm/types/request.rs` | 移除 `ToolChoice` 枚举 + serde impl |
| 新增 | `src/llm/types/tool.rs` | 增加 `ToolChoice` 枚举 + serde impl |
| 修改 | `src/llm/types/mod.rs` | 更新 re-export 路径 |
| 修改 | `src/llm/types/request_v2.rs` | 更新 import 路径 |
#### Step 13.1 — request 类型搬移
| 操作 | 文件 | 变更说明 |
|------|------|---------|
| 删除 | `src/llm/types/request.rs` | 整文件删除(187 行) |
| 新增 | `src/llm/provider/openai.rs` | 增加所有 OpenAI wire-format 类型 |
| 修改 | `src/llm/types/mod.rs` | 删除 re-export + mod 声明 |
#### Step 13.2 — response 类型搬移
| 操作 | 文件 | 变更说明 |
|------|------|---------|
| 删除 | `src/llm/types/response.rs` | 整文件删除(177 行) |
| 新增 | `src/llm/provider/openai.rs` | 增加所有 OpenAI wire-format 类型 + From impl |
| 修改 | `src/llm/types/mod.rs` | 删除 re-export + mod 声明 |
| 修改 | `src/llm/stream.rs` | 更新 `OpenaiChatChunk` import 路径 |
#### Step 13.3 — 旧文件清理
| 操作 | 文件 | 变更说明 |
|------|------|---------|
| 删除 | `src/llm/types/old_stream.rs` | 整文件删除(45 行) |
| 修改 | `src/llm/types/mod.rs` | 删除 `pub mod old_stream;` + 删除 `ChatResponse` 结构体 + `From` impl |
| 修改 | `src/llm/stream.rs` | 删除所有死代码,仅保留 `pub use` 重导出 |
### 4.3 回滚策略
所有 Step 通过 git commit 管理,回退时 `git revert <commit>` 即可。每个 Step 独立编译,回滚不会级联依赖。若 Step 13.3(`ChatResponse` 删除)导致外部编译失败,单独 revert 该 commit 即可恢复 `ChatResponse` + `old_stream.rs`
### 4.4 CHANGELOG 条目
```markdown
## [0.3.0] - 未发布
### Breaking Changes
**类型路径变更(0.3.0):**
- `agcore::llm::types::request::ToolChoice``agcore::llm::types::tool::ToolChoice`(公共 re-export 路径 `agcore::llm::types::ToolChoice` 保持不变)
- `agcore::llm::types::request::StreamOptions``agcore::llm::provider::openai::StreamOptions`
- `agcore::llm::types::request::OpenaiChatRequest``agcore::llm::provider::openai::OpenaiChatRequest`
- `agcore::llm::types::response::OpenaiChatResponse``agcore::llm::provider::openai::OpenaiChatResponse`
- `agcore::llm::types::response::OpenaiChatChunk``agcore::llm::provider::openai::OpenaiChatChunk`
- 其余 `request.rs`/`response.rs` 中的 wire-format 类型(`OpenaiTool``AudioParam``Choice``Delta` 等)同步移入 `agcore::llm::provider::openai` 模块
**类型删除:**
- `agcore::llm::types::ChatResponse` 已删除(自 v0.1.0 标记 `#[deprecated]`,请改用 `MessageResponse`
- `agcore::llm::types::old_stream::LegacyStreamEvent` 已删除(内部死代码)
### Features
- `ContextSlot::fork(child_id, strategy)` — 从父槽派生独立的子槽(数据层操作)
- `ContextSlot::merge(child, strategy)` — 将子槽消息合并回父槽(支持 Append/Replace
- `MergeStrategy` 枚举(`#[non_exhaustive]`Phase 16 可扩展 Summarize
```
---
## 5. 风险评估
| 风险 | 影响 | 概率 | 缓解措施 |
|------|------|------|---------|
| `ChatResponse` 被外部 crate 引用 | 编译 break | 中 — `#[deprecated]` 仅产生编译警告,外部 crate 可能通过 `#[allow(deprecated)]` 静默依赖 | CHANGELOG 明确标注语义版本(0.3.0)和迁移指引;Step 13.3 验收加入 `cargo doc --no-deps \| grep ChatResponse` 确认零引用 |
| `StreamOptions` 等 wire-format 类型路径变更影响直接引用消费者 | 编译 break | 低(v0.2.0-rc.1,极少外部消费者使用内部类型) | CHANGELOG 完整列出所有路径变更;编译错误立即可发现 |
| `parse_chunk_stream` 有隐藏调用方 | 编译 break | 极低(实施前执行 `grep -rn 'parse_chunk_stream\|map_legacy_to_ir\|LegacyToIrEventStream' src/` 前置验证) | Step 13.3 前运行 grep 验证并记录结果;`cargo build --all-targets` 可 100% 捕获 |
| `#[allow(deprecated)]` 遗漏 | clippy 警告 | 低 | `cargo clippy --all-targets -- -D warnings` 验证 |
| Step 顺序错误导致编译中间态 | 开发者体验差 | 中 | 严格按 13.5→13.4→13.1→13.2→13.3 执行;每步 `cargo build` 验证 |
| `stream.rs` 简化后 import 断链 | 编译 break | 极低 | 保留 `pub use` 重导出路径,`cycle.rs`/`session.rs` import 不变 |
---
## 6. 验收标准
### M9 里程碑(Phase 13 完成条件)
| # | 条件 | 验证方法 |
|---|------|---------|
| 1 | `request.rs``response.rs``old_stream.rs` 三个旧文件不存在 | `ls src/llm/types/` 确认 |
| 2 | `ChatResponse` 结构体不存在 | 全局搜索 `ChatResponse` 仅保留 `openai.rs``OpenaiChatResponse` 引用 |
| 3 | `ToolChoice``tool.rs` 中定义,公共路径 `agcore::llm::types::ToolChoice` 保持不变 | `cargo doc --no-deps` 确认类型文档 |
| 4 | `OpenaiChatRequest`/`Response`/`Chunk``provider/openai.rs` 中定义 | 编译通过 |
| 5 | `ContextSlot::fork()` 单元测试通过(P0 条件全部满足) | `cargo test` |
| 6 | `ContextSlot::merge()` 单元测试通过(P0 条件全部满足) | `cargo test` |
| 7 | `stream.rs` 只保留 `pub use` 重导出 | 文件内容确认 |
| 8 | `cargo build --all-targets` 编译通过 | 编译验证 |
| 9 | `cargo test --all-targets` 全绿(预期 283~285 测试) | 测试验证 |
| 10 | `cargo clippy --all-targets -- -D warnings` 0 警告 | clippy 验证 |
| 11 | CHANGELOG 包含 Phase 13 的 Breaking Changes 和 Features 条目 | 文件确认 |
### fork/merge 详细验收 P0 项
**fork 的 5 项 P0 条件:**
| # | 条件 | 优先级 |
|---|------|--------|
| 1 | `fork("child", Full)` 创建新 slot,消息在 fork 时刻 == 父 slot | P0 |
| 2 | 子 slot 获得独立消息列表——父 slot 后续追加不影响子 slot | P0 |
| 3 | 子 slot 的 source 标记为 `Derived { parent_id, strategy }` | P0 |
| 4 | 子 slot 可独立持久化(fork + save + load roundtrip | P0 |
| 5 | fork 不允许重复 id(返回 `SlotAlreadyExists`)(由 `derive_slot` 编排层保证) | P0 |
**merge 的 5 项 P0 条件:**
| # | 条件 | 优先级 |
|---|------|--------|
| 1 | `parent.merge(child, Append)` 子消息追加到父末尾 | P0 |
| 2 | `parent.merge(child, Replace)` 子消息替换父全量消息 | P0 |
| 3 | merge 后父 slot 的 `meta.message_count` 正确更新 | P0 |
| 4 | merge 不允许合并到 Readonly 目标 slot | P0 |
| 5 | merge 不允许 self-mergechild.id == parent.id | P0 |
---
## 参考来源
- Roadmap`docs/roadmap.md` §Phase 13
- ContextSlot 设计:`docs/17-phase10-contextslot.md`
- 旧 StreamEvent 设计:`src/llm/stream.rs` 文件注释
- 当前代码库:`src/llm/types/request.rs``src/llm/types/response.rs``src/llm/types/old_stream.rs``src/llm/types/mod.rs``src/llm/provider/openai.rs``src/agent/context.rs``src/agent/session.rs`
---
## 7. 实施计划
### 全局说明
**commit 策略**:每个 Step 一个独立 commit。commit message 格式:
```
<type>(<scope>): <中文描述>
```
- Step 13.5 → `feat(agent): 实现 ContextSlot fork/merge`
- Step 13.4 → `refactor(types): ToolChoice 移入 tool.rs`
- Step 13.1 → `refactor(types): request.rs 类型移入 provider/openai.rs`
- Step 13.2 → `refactor(types): response.rs 类型移入 provider/openai.rs`
- Step 13.3 → `refactor(types): 删除旧类型文件和 ChatResponse`
**验证命令(每步通用)**
```bash
cargo build --all-targets && cargo test && cargo clippy --all-targets -- -D warnings
```
**预计测试数量变化**
- 当前基线:277 测试(每个 Step 开始时 `cargo test` 确认)
- Step 13.5 后:286+9
- Step 13.4-13.2 后:286(无变化)
- Step 13.3 后:285-1`ChatResponse``From` impl 无测试直接引用,删除后仅 `types/mod.rs` 中的 `deprecated` 注释行减少,不影响测试计数。实施前执行 `grep -rn 'ChatResponse' src/ --include='*test*' --include='*tests*'` 确认零测试引用)
- 最终范围:285 测试
### Step 13.5 — ContextSlot fork/merge
**前置依赖**:无(纯新增,不依赖前序 Step
**任务描述**:在 `agent/context.rs` 中新增 `MergeStrategy` 枚举、`ContextSlot::fork()` 方法和 `ContextSlot::merge()` 方法;重构 `agent/session.rs` 中的 `derive_slot` 改为调用 `parent.fork()`;新增 9 个内联测试覆盖 fork/merge 的 happy path 和 error path。
**涉及文件**
- `src/agent/context.rs` — 新增枚举和方法
- `src/agent/session.rs` — 重构 derive_slot
- `src/agent.rs` — 追加 `MergeStrategy` re-export
**具体操作**
1.`context.rs` 中新增 `MergeStrategy` 枚举(Append / Replace`#[non_exhaustive]`
2.`context.rs``impl ContextSlot` 块内新增 `fork(&self, child_id: String, strategy: DeriveStrategy) -> ContextSlot` 方法
3.`context.rs``impl ContextSlot` 块内新增 `merge(&mut self, child: ContextSlot, strategy: MergeStrategy) -> Result<(), AgentError>` 方法(含 self-merge/cross-session/Readonly 三项防御检查 + `tracing::debug!` 日志)
4.`session.rs``derive_slot` 方法中将手工消息复制代码替换为 `parent.fork(slot_id, strategy)`
5.`agent.rs``pub use context::{...}` 列表中追加 `MergeStrategy`
6.`context.rs``#[cfg(test)] mod tests` 中新增 9 个测试用例
**注意**:重构后 `derive_slot` 的子 slot `budget``ContextBudget::default()` 变为继承父 slot`compact``true` 变为继承父 slot。由于 `ContextBudget` 在 v0.2 无消费逻辑且父 slot 的 `compact` 默认也为 `true`,此变化无实际影响。验收条件中"行为不变"指对外功能行为不变(slot 消息内容、血缘关系不变)。
**预估工作量**M1-4h
**风险等级**:低(纯新增,不修改已有逻辑路径)
**验收条件**
- `MergeStrategy` 枚举存在,`Append``Replace` 两个变体可用,且通过 `agcore::agent::MergeStrategy` 路径可访问
- `ContextSlot::fork` 返回的 child 在 fork 时刻消息等于父 slot
- fork Focused 策略按 `FocusedConfig` 过滤消息
- 父 slot 后续追加消息不影响子 slot
- 子 slot 的 source 正确记录 `Derived { parent_id, strategy }`
- `parent.merge(child, Append)` 追加到父末尾,message_count 正确
- `parent.merge(child, Replace)` 替换父全量消息,message_count 正确
- self-merge 返回 `Err(AgentError::Config)`
- merge 到 Readonly slot 返回 `Err(AgentError::SlotReadonly)`
- 跨 session merge 返回 `Err(AgentError::Config)`
- `derive_slot` 对外行为不变(slot 消息内容、血缘关系、持久化行为均不变;内部 budget/compact 继承差异无实际影响),测试全绿
- `cargo doc --no-deps` 无 warning(验证新增公开 API 的文档注释完整)
**回退方式**`git revert` 该 commit
### Step 13.4 — ToolChoice 移入 tool.rs
**前置依赖**:Step 13.5(顺序约束:必须早于 Step 13.1——若 Step 13.1 先执行会将 `ToolChoice``request.rs` 一同删除,导致本 Step 无可搬移的源)
**任务描述**:将 `ToolChoice` 枚举及其 serde 实现从 `types/request.rs` 搬移到 `types/tool.rs`,更新所有 import/path 引用。公共 re-export 路径 `agcore::llm::types::ToolChoice` 保持不变。
**涉及文件**
- `src/llm/types/request.rs` — 删除 ToolChoice~28-99 行)
- `src/llm/types/tool.rs` — 新增 ToolChoice 枚举 + serde impl
- `src/llm/types/mod.rs` — re-export 路径从 `request` 改为 `tool`
- `src/llm/types/request_v2.rs` — import 路径从 `request::` 改为 `tool::`
**具体操作**
1.`request.rs` 复制 `ToolChoice` 枚举 + `Serialize`/`Deserialize` impl 到 `tool.rs`
2.`request.rs` 中删除 `ToolChoice` 定义
3.`mod.rs` 中将 `pub use request::{..., ToolChoice}` 改为 `pub use tool::ToolChoice`
4.`request_v2.rs` 中将 `use crate::llm::types::request::ToolChoice` 改为 `use crate::llm::types::tool::ToolChoice`
5. 验证 `cycle.rs``use crate::llm::types::ToolChoice`(通过 re-export)路径不变
**预估工作量**S<1h
**风险等级**:低(有限的 import 路径变更,编译立即可发现)
**验收条件**
- `ToolChoice``tool.rs` 中定义
- `pub use tool::ToolChoice``mod.rs`
- `request_v2.rs` 编译通过
- `cycle.rs` 路径不变
- `cargo build --all-targets` + `cargo test` + `cargo clippy` 全绿
**回退方式**`git revert` 该 commit
### Step 13.1 — request.rs 类型移入 openai.rs
**前置依赖**Step 13.4ToolChoice 已移走,request.rs 剩余内容全是 OpenAI wire-format 专有类型)
**任务描述**:删除 `types/request.rs` 整文件,将所有剩余类型(`OpenaiChatRequest``StreamOptions``OpenaiTool``AudioParam``PredictionContent``UserLocation``Approximate``WebSearchOptions`)搬入 `provider/openai.rs`,更新 `mod.rs` re-export。
**涉及文件**
- `src/llm/types/request.rs` — 整文件删除
- `src/llm/provider/openai.rs` — 新增所有类型定义
- `src/llm/types/mod.rs` — 删除 re-export + mod 声明
**具体操作**
1.`request.rs` 复制所有剩余类型定义到 `openai.rs`,可见性设为 `pub(crate)`
2. `OpenaiTool` 内引用 `OpenaiToolDefinition`(定义在 `tool.rs`),路径改为 `crate::llm::types::tool::OpenaiToolDefinition`
3. 删除 `openai.rs` 中原 `use crate::llm::types::request::{...}` import
4.`mod.rs` 删除 `pub use request::{OpenaiChatRequest, OpenaiTool, StreamOptions}``pub mod request;`
5. 删除 `types/request.rs` 文件
**预估工作量**M1-4h
**风险等级**:低(纯搬移 + 删除,文件内无逻辑变更)
**验收条件**
- `request.rs` 文件不存在
- `OpenaiChatRequest` 等类型在 `openai.rs` 中定义,编译通过
- `OpenaiTool` 通过 `crate::llm::types::tool::OpenaiToolDefinition` 正确引用
- `cargo build --all-targets` + `cargo test` + `cargo clippy` 全绿
**回退方式**`git revert` 该 commit。若 Step 13.2 也已提交,单独 revert 本 Step 可能因 `provider/openai.rs` 并发修改产生合并冲突。安全回退顺序为逆序:先 revert 13.2,再 revert 13.1。
### Step 13.2 — response.rs 类型移入 openai.rs
**前置依赖**:无(与 Step 13.1 共享 `provider/openai.rs``types/mod.rs`,但本 Step 仅追加类型定义,无覆盖操作;建议在 13.1 之后顺序执行以避免并行时的合并冲突)
**任务描述**:删除 `types/response.rs` 整文件,将所有类型(`OpenaiChatResponse``OpenaiChatChunk``Choice``Delta``ChunkChoice` 等 + 两个 `From` impl)搬入 `provider/openai.rs`,更新 `mod.rs``stream.rs` 的 import 路径。
**涉及文件**
- `src/llm/types/response.rs` — 整文件删除
- `src/llm/provider/openai.rs` — 新增所有类型定义 + From impl
- `src/llm/types/mod.rs` — 删除 re-export + mod 声明
- `src/llm/stream.rs``OpenaiChatChunk` import 路径改为 `provider::openai`
**具体操作**
1.`response.rs` 复制所有类型定义(含 `From` impl)到 `openai.rs`,可见性设为 `pub(crate)`
2. 删除 `openai.rs` 中原 `use crate::llm::types::response::{...}` import
3.`mod.rs` 删除 `pub use response::{...}``pub mod response;`
4.`stream.rs:26``OpenaiChatChunk` 的 import 路径改为 `crate::llm::provider::openai::OpenaiChatChunk``OpenaiToolCall` 路径不变)
5. 删除 `types/response.rs` 文件
**预估工作量**M1-4h
**风险等级**:低(与 Step 13.1 模式完全相同)
**验收条件**
- `response.rs` 文件不存在
- `OpenaiChatResponse`/`Chunk` 等类型在 `openai.rs` 中定义,编译通过
- `stream.rs` import 路径正确
- `cargo build --all-targets` + `cargo test` + `cargo clippy` 全绿
**回退方式**`git revert` 该 commit。若 Step 13.1 和本 Step 均已提交,安全回退顺序为逆序:先 revert 本 Step,再 revert 13.1。
### Step 13.3 — 旧文件清理 + ChatResponse 删除
**前置依赖**Step 13.1`request.rs` 已删)、Step 13.2`response.rs` 已删)
**任务描述**:删除 `old_stream.rs``ChatResponse`,简化 `stream.rs` 为仅保留 `pub use` 重导出。这是 Phase 13 技术风险最高的 Step。
**涉及文件**
- `src/llm/types/old_stream.rs` — 整文件删除
- `src/llm/types/mod.rs` — 删除 `pub mod old_stream;` + 删除 `ChatResponse` 结构体和两个 `From` impl
- `src/llm/stream.rs` — 删除死代码(约 160 行),仅保留 `pub use` 重导出
**具体操作**
1. **前置验证 A**:执行 `grep -rn 'parse_chunk_stream\|map_legacy_to_ir\|LegacyToIrEventStream\|ChunkToLegacyEventStream' src/` 确认零外部调用方,记录结果到 commit message
2. **前置验证 B**:执行 `cargo doc --no-deps 2>&1 | grep -i 'ChatResponse'` 确认零文档引用,记录结果
3.`mod.rs` 删除 `pub mod old_stream;`
4.`mod.rs` 删除 `ChatResponse` 结构体定义 + `#[allow(deprecated)]` `From<OpenaiChatResponse> for ChatResponse` + `From<ChatResponse> for OpenaiChatChunk`
5. 删除 `old_stream.rs` 文件
6.`stream.rs` 删除:`use crate::llm::types::old_stream::LegacyStreamEvent``parse_chunk_stream``parse_chunk_stream_legacy``ChunkToLegacyEventStream``LegacyToIrEventStream``map_legacy_to_ir``empty_message_response`
7. `stream.rs` 最终只保留 module doc comment + `pub use crate::llm::types::response_v2::StreamEvent;`
8. 检查 `cycle.rs:88``#[allow(deprecated)]` 属性是否仍与 `ChatResponse` 相关——若不相关则无需改动;若因 `ChatResponse` 删除而变脏,清理该属性
**预估工作量**S<1hcleanup+ M(需验证过程)
**风险等级**:中(`ChatResponse` 删除是 Breaking Change,外部可能静默依赖)
**验收条件**
- `old_stream.rs` 文件不存在
- `ChatResponse` 结构体不存在(全局搜索仅保留 `OpenaiChatResponse` 引用)
- `stream.rs` 只保留 `pub use` 重导出
- `cargo build --all-targets` 编译通过
- `cargo test --all-targets` 全绿(预期 285 测试)
- `cargo clippy --all-targets -- -D warnings` 0 警告
- `cargo doc --no-deps` 无 warning
**回退方式**`git revert` 该 commit(单独 revert 即可恢复 `ChatResponse` + `old_stream.rs`
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+183
View File
@@ -0,0 +1,183 @@
# LangChain & LangGraph 功能调研笔记
> 调研时间:2026-07-06
> 两者关系:同一公司(LangChain Inc.)维护的堆栈上下两层,不是竞品
---
## 两者关系
```
┌──────────────────────────────────────────┐
│ LangChain (v1.0 GA) │ ← 高层框架:模型抽象、工具、提示词、600+集成
│ create_agent / LCEL / 组件库 │
├──────────────────────────────────────────┤
│ LangGraph (v1.0 GA) │ ← 底层运行时:有向图执行引擎
│ StateGraph / Checkpointing / HITL │
├──────────────────────────────────────────┤
│ LangSmith (可观测性) │
└──────────────────────────────────────────┘
```
2025年10月22日同时达到 v1.0 GA,官方分工:
> **LangChain** = agent frameworkabstractions and integrations for models, tools, and agent loops.
> **LangGraph** = orchestration runtimedurable execution, streaming, human-in-the-loop, and persistence.
LangChain v1.0 的 `create_agent` 内部已运行在 LangGraph 引擎上。
---
## LangChain v1.0
### 定位
高层应用框架,提供 agent 所需的**组件抽象**和**集成生态**。
### 精简后的核心模块
| 模块 | 功能 |
|------|------|
| `langchain.agents` | `create_agent`, `AgentState`(取代旧 AgentExecutor |
| `langchain.chat_models` | `init_chat_model`, `BaseChatModel`(统一模型初始化) |
| `langchain.tools` | `@tool`, `BaseTool` |
| `langchain.messages` | 消息类型、内容块、`trim_messages` |
| `langchain.embeddings` | `init_embeddings`, `Embeddings` |
旧组件(`LLMChain``ConversationChain` 等)移入 `langchain-classic`
### 七大组件类别
| 类别 | 关键组件 |
|------|----------|
| **Models** | Chat models, LLMs, Embeddings — 统一接口跨 provider 切换 |
| **Tools** | 600+ provider 集成:API、数据库、搜索引擎等 |
| **Agents** | `create_agent`, ReAct agents, Tool-calling agents |
| **Memory** | 消息历史、自定义状态 |
| **Retrievers** | 向量检索器、网络检索器 |
| **Document** | 加载器、分割器、转换器 |
| **Vector Stores** | Chroma, Pinecone, FAISS 等集成 |
### v1.0 关键新特性
**1. Middleware 中间件系统**`create_agent` 的钩子系统:
- `before_model` — 模型调用前注入/修改
- `after_model` — 模型调用后验证/后处理
- `wrap_tool_call` — 拦截工具调用错误
**2. Standard Message Content** — 跨 provider 标准化消息内容格式:
- 推理/思维链、引用、多模态(图片/音视频/文档)
- 工具调用、provider 特有工具(web search, code execution
- 通过 `.content_blocks` 属性访问,向后兼容
**3. `create_agent`** — 取代旧 AgentExecutor,内部运行在 LangGraph 运行时上
### 成熟度
| 维度 | 状态 |
|------|------|
| 版本 | v1.0 GA2025-10 |
| 稳定性 | 稳定,agent 层经重构后已稳定 |
| 生产证明 | Replit, Clay, Rippling, Cloudflare, Workday |
| 支持 | LTS-style support track |
| 适用场景 | RAG、信息提取、单 agent 助手、快速原型 |
---
## LangGraph v1.0
### 定位
底层编排运行时,专为**有状态、长时间运行、多步骤**工作流设计。
### 核心抽象链
```
StateGraph → Nodes (纯 Python 函数) → Edges (路由逻辑)
Shared State (TypedDict / Pydantic)
Checkpointer (每个 super-step 快照)
```
- **StateGraph**: 有状态图,参数化 State 类型
- **Nodes**: 纯函数,`(State) → updates`
- **Edges**: `add_conditional_edges`,支持循环/分支/合并
- **State**: `TypedDict` 或 Pydantic,带 reducer 处理并发更新
- **Reducers**: `add_messages` 等,自动处理追加 vs 覆盖
### 完整功能矩阵
| 功能 | 状态 | 细节 |
|------|------|------|
| **StateGraph** | ✅ 稳定 | 循环图(非 DAG),条件边缘,并行 fan-out |
| **Checkpointing** | ✅ v4.1.1 | SQLite / PostgreSQL / Redis 后端 |
| **Durable Execution** | ✅ 稳定 | 跨失败自动恢复,从精确断点继续 |
| **Human-in-the-loop** | ✅ 一等公民 | `interrupt()` + `Command(resume=...)` |
| **Time-travel 调试** | ✅ 稳定 | 回滚任意 checkpointfork 重放 |
| **流式输出** | ✅ 稳定 | Token 级 + State 级 + Event 级 |
| **多 Agent 编排** | ✅ 稳定 | Supervisor / Swarm / 层级 / Subgraph |
| **Comprehensive Memory** | ✅ 稳定 | 短时工作记忆 + 长时持久记忆 |
| **增量状态存储** | 🧪 DeltaChannel beta (v4.1.0+) | 长消息列表只存 delta |
| **跨进程状态同步** | 🧪 RemoteCheckpointer (v4.1.0+) | 分布式多 agent 架构 |
| **自动 checkpoint 清理** | ✅ keep_latest TTL (v4.0.2) | 避免无限制积累历史 |
| **LangGraph Platform** | ✅ 稳定 | Agent Server:持久化、任务队列、版本管理 |
| **LangGraph Studio** | ✅ 稳定 | 可视化 agent 工作流 |
### 成熟度
| 维度 | 状态 |
|------|------|
| 版本 | v1.0 GA2025-10),checkpointer v4.1.1 (2026-05) |
| 稳定性 | 高,持久化为架构一等公民 |
| 生产证明 | Klarna, Replit, Elastic |
| 支持 | LTS-style support track |
| 适用场景 | 多步骤 agent、多 agent 系统、人工审批、长时间运行任务 |
---
## 功能边界对比
| 维度 | LangChain | LangGraph |
|------|-----------|-----------|
| **层次** | 高层应用框架 | 底层编排运行时 |
| **核心抽象** | `create_agent`, LCEL, 组件库 | `StateGraph`, Nodes, Edges, State |
| **思维模型** | 线性或 DAG 管道 | 节点 + 边缘的循环有向图 |
| **循环/分支** | 受限 | **一等公民**:任意循环、分支、合并 |
| **状态持久化** | 无原生支持 | **一等公民**Checkpointer |
| **Human-in-loop** | 需手动编排 | **一等公民**`interrupt()` + `Command` |
| **Time-travel 调试** | 无 | **一等公民**:回滚 fork 重放 |
| **Durable Execution** | 无 | **一等公民**:跨故障自动恢复 |
| **流式** | Token 级 | Token + State + Event 每节点流式 |
| **多 Agent 编排** | 需手动组合 | **原生**Supervisor/Swarm/Subgraph |
| **模型抽象** | **核心优势** | 复用 LangChain |
| **600+ 集成** | **核心优势** | 可复用 LangChain 集成 |
| **LCEL 线性链** | **有** | 无 |
| **Middleware** | **v1.0 特有** | 无 |
| **学习曲线** | 中等 | 较陡(需图思维) |
| **部署平台** | 无独立平台 | LangGraph Platform + Studio |
---
## 决策路线
```
你的 workflow 需要什么?
├─ 线性、始终相同步骤 → LangChain (LCEL / create_agent)
├─ 需要循环/分支/重试 → LangGraph (StateGraph)
├─ 需要持久化/故障恢复 → LangGraph (Checkpointer)
├─ 需要人工审批 → LangGraph (interrupt())
├─ 需要 time-travel 调试 → LangGraph (checkpoint + fork)
├─ 需要多 agent 协作 → LangGraph (Supervisor/Swarm/Subgraph)
└─ 不确定 → 先用 create_agent,遇到瓶颈下钻到 StateGraph
```
---
## 参考来源
- [LangChain Blog: v1.0 Milestone](https://www.langchain.com/blog/langchain-langgraph-1dot0)
- [LangChain Documentation](https://docs.langchain.com/oss/python/langchain/overview)
- [LangGraph Documentation](https://docs.langchain.com/oss/python/langgraph/overview)
- [LangGraph GitHub](https://github.com/langchain-ai/langgraph)
- [Atlan: LangChain vs LangGraph 2026](https://atlan.com/know/ai-agent/ai-agent-memory/langchain-vs-langgraph/)
- [truefoundry: LangChain vs LangGraph](https://www.truefoundry.com/blog/langchain-vs-langgraph)
+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` | 后台作业核心引擎(内存注册表) |
+644 -63
View File
@@ -1,13 +1,13 @@
# AG Core Roadmap # AG Core Roadmap
> 定稿日期:2026-05-11 > 定稿日期:2026-05-11
> 最后更新:2026-07-04v0.1 发布完成) > 最后更新:2026-07-09Phase 14 完成 + M10 里程碑达成)
## 愿景 ## 愿景
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。 AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
**当前状态**Phase 0-4c 全部完成Provider IR 重构(统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen Provider)已完成;LlmCycle 简化(IR 消息类型切换 + 桥接层移除)已完成;v0.1 发布就绪(**182 个测试通过、0 clippy 警告、7 个离线示例可运行**) **当前状态**v0.2.0-rc.1 已打标签。Phase 0-14 全部完成。v0.3.0 实施中,Phase 15-19 共 5 个增量 Phase 待交付。目标是从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"
--- ---
@@ -240,95 +240,665 @@ graph BT
--- ---
## 扩展计划(v0.2+ ## v0.2.0 — 生产就绪(Production-Ready Core
> 以下功能在已完成的 phase 中已实现基础能力或在 Phase 4 阶段明确了边界,后续可按维度增量扩展 **目标**:解决 Rust Agent 工具箱从"能跑"到"能被人依赖"的鸿沟。持久化、配置层、上下文管理三大块补齐后,开发者可在 30 分钟内写出生产可用的 Agent 服务
> 设计参考:见 `docs/note-agent-harness-references.md`OpenClaw / Hermes / OpenHuman / OpenHarness 横向对比)。
> OpenCode 借鉴:见 `docs/note-opencode-agent-switching.md`Agent 切换 + System Prompt 拼接机制)。
### 已有扩展项(沿用) **总体规模**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 上下文管理
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 | **模块归属**`src/llm/context.rs`(与 `compact.rs` 同级)
|-------|---------|------|--------|------|
| 多通道检索(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 待评估 |
#### 交互层(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 }
```
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 | **持久化 Key 命名**
|-------|---------|------|--------|------| - `slot_msg:{session_id}:{slot_id}:{index}` → 消息内容
| RL 轨迹导出 | `agent` | ShareGPT 格式轨迹、Atropos 集成(Hermes 风格) | P3 | v0.3+ 探索 | - `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 派生
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 | **与 `ConversationMemory` 的关系**:保留不废除。`ConversationMemory` 继续服务传统对话场景。
|-------|---------|------|--------|------|
| Human-in-the-loop 审批 | `agent` / `tools/permission` | 高危工具执行前的异步审批回调(OpenHarness `permission_prompt` 模式) | P2 | v0.2 待评估 |
#### 流式 / 实时 **v0.2 不做**
-`slot.fork()` / `merge()` — 分支方法推迟到 v0.3+
-`inject_summary` 自动生成 — v0.2 仅消费端(从 `SessionMemory` 读取),生成在 v0.3+
- ❌ 血缘关系图遍历 — 只存 `parent_id`,不做查询
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 | **依赖**Phase 0MemoryStore trait)、Phase 3MemoryStore 持久化)
|-------|---------|------|--------|------| **优先级**P1
| 流式 `submit_turn` | `agent/session` | Phase 4 v1 只暴露非流式 `submit_turn()`v0.2 包装 `LlmCycle::submit_stream` 暴露流式入口 | P2 | v0.2 待评估 |
#### Agent 切换 / Prompt 动态(OpenCode 借鉴) ---
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 | ### v0.2.0 实施计划 — 8 个增量 Phase
|-------|---------|------|--------|------|
| Agent 身份切换(角色轮换) | `agent` | 借鉴 OpenCode Tab 键切换 build/plan:同一 `AgentSession` 持有可热替换的 `Agent` 引用,切换时不重置消息历史,在末尾追加 `synthetic: true` 的状态变更消息。详见 `docs/note-opencode-agent-switching.md` §4 | P2 | v0.2 待评估 | > **编号说明**Phase 5-12 接续 v0.1 的 Phase 0-4c,按开发顺序排列。
| 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 待评估 | #### 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 |
**实际新增**2026-07-06 commit `71abe88` / `b4e5c7d`,详见 `docs/18-phase11-testing-and-retrieval.md`):
- 方案文档:`docs/18-phase11-testing-and-retrieval.md`647 行,含 11.1/11.2/11.3 设计 + 10 项架构决策 + 实施后补充 2 条偏差记录 #6 mid-stream mock 模式 + #7 429 retry-after 修复)
- 新增文件 1 个:`src/memory/vector.rs`237 行 — `VectorRetriever` trait + `InMemoryVectorRetriever` 引用实现 + `dot()` 零依赖 + 6 个内联测试)
- 修改文件 5 个:
- `src/memory.rs`+2 行:module 声明 + re-export
- `src/llm/provider/openai.rs`+8 wiremock 测试 + `handle_error_response` 429 retry-after 解析修复 5 行)
- `src/llm/provider/anthropic.rs`+4 wiremock 测试)
- `src/memory/store/in_memory.rs`(+3 并发测试:100 并发写、5 写+5 读混合、15 写者容量淘汰)
- `src/memory/store/sqlite_store.rs`(+2 并发测试:100 并发写、5 写+5 读混合)
- 关键设计:
- **零依赖 dot()**:手写点积/范数,零新增 crate 依赖
- **Wiremock 测试自包含**:每个测试独立 `MockServer::start()`,沿用现有模式
- **429 retry-after 修复**`openai.rs``anthropic.rs` 行为对齐(5 行代码)
- **偏差记录**:方案文档「已否决的方案 #6/#7」记录两处实施偏差,便于后续审计追溯
- 验证:254 → 277 测试(+23 个新测试),clippy 0 警告,doc 0 warning;并发测试连续 3 次运行稳定无 flaky
- **依赖**:无(与方案一致)
- **状态**:✅ Phase 11 全部交付物已完成
---
#### 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["<b>Phase 11: 测试与检索补强</b><br/>VectorRetriever trait<br/>12 wiremock tests<br/>5 并发测试"]:::done
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+ | ✅ 2026-07-06 |
| **M8** | Phase 12(可选) | P2 功能按需交付 | ⏳ |
---
## v0.3.0 — 多 Agent 基础系统(Multi-Agent Foundation
**目标**:从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"。补齐 LangChain 7 大组件中缺失的 Document 和 VectorStore 能力,落地笔记设计中的 ContextSlot fork/merge、摘要自动生成、知识图谱,建立 engine 引擎层(会话树 + time-travel Checkpointer + SubAgent Dispatch + Agent Switch),为即将开发的多 Agent 产品提供完整基础。
**总体规模**7 个增量 PhasePhase 13-19),总新增代码约 2600 行,测试从 277 → 380+。
### 功能清单
#### P0 — 必须交付
| # | 功能 | 模块 | 方案要点 |
|---|------|------|---------|
| 1 | 技术债清理(旧 types 文件) | `llm/types` | `request.rs` / `response.rs` / `old_stream.rs` 三个 Phase 0 旧文件删除;内部类型移入 `provider/openai.rs` |
| 2 | ContextSlot fork/merge | `agent/context` | `fork(child_id, strategy)` 别名 + `merge(child, MergeStrategy)` 三种策略(Append/Replace/Summarize |
| 3 | Document 系统 | `document/`(新模块) | `Document` 核心类型 + `RecursiveCharacterSplitter`(递归字符分割,支持 chunk_size/chunk_overlap/separators |
| 4 | Embedding 抽象 | `llm/embedding` | `Embedding` trait`embed` / `dim`+ `MockEmbedding` 测试实现 |
| 5 | 向量存储持久化 | `vector/`(新模块) | `VectorStore` trait + `InMemoryVectorStore`(读写)+ `PersistentVectorStore`SqliteStore 后端)+ `RagPipeline` 组合器 |
| 6 | 摘要自动生成 | `agent` / `llm/hooks` | `SummaryConfig` 配置 + `OnTurnEnd` Hook 自动检测 token 水位 → 调 LLM 生成摘要 → `SessionMemory::set("conversation_summary", ...)` |
| 7 | SessionManager + 会话树 | `engine/`(新模块) | Session 工厂(`create`/`create_child`+ 按 ID 恢复(`get`+ 子树管理(`children`/`parent`/`destroy_subtree`+ 元数据持久化(MemoryStore |
| 8 | Time-travel Checkpointer | `engine/checkpointer` | `checkpoint(session)` 全量序列化 + `rollback(session_id, ckpt_id)` 回滚 + `fork(session_id, ckpt_id, new_id)` 分支 + `list_checkpoints` |
| 9 | Agent Switch | `engine/switch` | 热切换 `session.agent`(替换 `Arc<dyn Agent>`),slot 历史 / turn_index / session_memory 全保留 |
| 10 | SubAgent Dispatch | `engine/sub_agent` | `dispatch(parent, sub_agent, task, config)` 单任务 + `dispatch_all(parent, tasks, config)` 并行派发(Semaphore 并发控制)+ 子 SessionMemory 继承 + `SubTaskResult` 结构化回传 |
| 11 | 知识图谱 | `memory/graph` | `KnowledgeGraph` trait`add_entity` / `add_relation` / `get_related` / `find_by_keywords`+ `InMemoryGraph` 实现 + `tag_index` 标签管理 |
| 12 | 双通道检索 | `memory/retriever` | `MemoryRetriever` 扩展为双通道(`KnowledgeStore` + `KnowledgeGraph`+ `RetrievalStrategy::Hybrid` |
### 实施计划 — 7 个增量 Phase
> **编号说明**Phase 13-19 接续 v0.2 的 Phase 5-12,按开发顺序排列。
#### Phase 13: 热身清理 + ContextSlot fork/merge
**目标**:清除 Phase 0 遗留的旧 types 文件,交付超低价功能建立节奏。
| Step | 内容 | 文件范围 | 验证标准 |
|------|------|---------|---------|
| **13.1** | `OpenaiChatRequest` 移入 `provider/openai.rs``types/request.rs` 删除 | `llm/types/request.rs` + `llm/provider/openai.rs` | `cargo build --all-targets` |
| **13.2** | `OpenaiChatResponse/Chunk` 移入 `provider/openai.rs``types/response.rs` 删除 | `llm/types/response.rs` + `llm/provider/openai.rs` | `cargo build --all-targets` |
| **13.3** | `old_stream.rs` 删除 + `types/mod.rs``ChatResponse` 删除 | `llm/types/old_stream.rs` + `llm/types/mod.rs` | `cargo build` + 确认 3 个旧文件不存在 |
| **13.4** | `ToolChoice``request.rs` 搬到 `tool.rs` | `llm/types/tool.rs` + `llm/types/request_v2.rs` | `cargo test --all-targets` 全绿 |
| **13.5** | `ContextSlot::fork(child_id, strategy)` 别名 + `merge(child, MergeStrategy)` | `agent/context.rs` | 单元测试:fork → 子 slot 消息 = 父 slot 副本;merge(Append) → 消息按序追加 |
**依赖**:无
**优先级**P0
**预估规模**:约 200 行
**状态**:✅ Phase 13 全部交付物已完成(2026-07-08
---
#### Phase 14: Document 系统 + Embedding 抽象
**目标**:补齐 LangChain 7 大组件中最明显的缺口——Document 类型和分割器。不搞 Loader 框架,用户用 `fs::read_to_string` 自行加载。
**交付物**
1. `src/document.rs` 新模块(`Document` 类型 + `RecursiveCharacterSplitter`
2. `src/llm/embedding.rs``Embedding` trait + `MockEmbedding`
**设计要点**
- `Document`id / content / metadataHashMap<String, String>/ mime_type
- `RecursiveCharacterSplitter`chunk_size(默认 1000/ chunk_overlap(默认 200/ separators`["\n\n", "\n", "。", "", "", ".", " ", ""]`,含 CJK 标点)
- 两阶段算法:按 separator 优先级递归分割(Phase 1+ 贪心合并 + overlap 滑动窗口(Phase 2
- 所有长度比较以 Unicode 字符数为单位(`chars_len()`),非字节数
- `Embedding` trait`async fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>, LlmError>` + `fn dim()`
- 复用 `LlmError` 而非新错误类型
- `MockEmbedding`:sin-hash 零依赖伪随机向量 + L2 归一化
- 不引入 `DocumentLoader` trait(应用层职责)
**实际新增**2026-07-09 commit `d4c4d8f`,详见 `docs/20-phase14-document-and-embedding.md`):
- 新增文件 3 个:
- `src/document.rs`580 行)— `Document` 类型(4 字段 + `new`/`from_raw` 构造器,2 个 `new` 接受 `impl Into<String>` + `RecursiveCharacterSplitter`(两阶段算法:按 separator 优先级递归分割 + 贪心合并 overlap,所有长度比较 `chars_len()` 字符级,overlap 提取 `chars().rev().take().rev()` 字符级安全)+ 19 个内联测试
- `src/llm/embedding.rs`183 行)— `Embedding` traitasync + `LlmError`+ `MockEmbedding`(sin-hash:字节和+长度做种子,`f32::sin(seed + i) * 10000`,L2 归一化到单位长度,零向量防除零)+ 6 个内联测试
- `examples/document_demo.rs`(74 行)— 端到端演示 Document → RecursiveCharacterSplitter → MockEmbedding → InMemoryVectorRetriever → search
- 修改文件 2 个:
- `src/lib.rs`+3 行:`pub mod document` + `pub use document::Document` + 空行)
- `src/llm.rs`+1 行:`pub mod embedding`
- 关键设计:
- **早返回守卫**`split_text``chars_len(text) <= self.chunk_size` 时直接返回 `[text]`,避免短文本在 Phase 2 `join("")` 中丢失 separator 边界
- **`Document::new` 使用 `impl Into<String>`**:接受 `&str``String`,比规范示例的 `String` 更灵活
- **`new()` panic + `try_new()` Result 双路径**:与 Rust 库惯例一致
- **CJK 分隔符扩展**`DEFAULT_SEPARATORS` 包含 `"。"`/`""`/`""`,避免中文文本跳过句子级退化为空格分割
- **chunk_size = 0 校验**:构造器拒绝零值,避免字符级兜底死循环
- **tracing 埋点**`split()` 入口 `tracing::debug!` + 每文档/每 chunk `tracing::trace!`
- **debug_assert 溢出保护**:单文档 chunk 数 < 10000 时 `debug_assert!`
- **Metadata 键覆盖文档化**`HashMap::insert()` 静默覆盖 source_id/chunk_index/chunk_count 在 `split()` doc comment 注明
- 测试:19 个 Document 测试(含 1 个 split_multibyte_utf8_boundary CJK 边界测试)+ 6 个 Embedding 测试,全量 286 → 313(+27 新测试,但部分测试覆盖范围重叠计算约 25 个净增)
- 方案文档:`docs/20-phase14-document-and-embedding.md`(1417 行,含背景/调研/方案对比/实施计划(详细版)/3 轮审查修复记录),经过 3 轮 PM/SA 审查 + 1 轮实施后修复
- clippy 0 警告,doc 0 warning
- 无新增外部依赖(`Cargo.toml` 未修改)
**实施后调整**
- 实施发现方案算法中 Phase 1 累加器设计与测试期望冲突("para1\n\npara2" 在 chunk_size=100 时 1 chunk 更合理),简化为"按 separator 切分 + Phase 2 合并"两阶段分工
- 二次审查发现 `split_text` 缺少早返回守卫 + `current_sep_count` 虚增计数,全部已修复
**依赖**:无(纯数据结构 + 零新 crate 依赖)
**优先级**P0
**预估规模**:约 350 行
**状态**:✅ Phase 14 全部交付物已完成(2026-07-09
---
#### Phase 15: 向量存储持久化(SqliteStore 后端)
**目标**:实现 VectorStore 持久化,让语义检索支持进程重启后数据恢复。
**设计决策**:不用 pgvector。基于已有 SqliteStore`rusqlite`)做持久化包装——运行时全量加载到 InMemory 索引做余弦搜索,写时同步到 SqliteStore。
**交付物**
1. `src/vector/` 新模块:`VectorStore` trait + `InMemoryVectorStore` + `PersistentVectorStore` + `RagPipeline`
2. `VectorStore` trait`add(docs, embeddings)` / `search(query, k)` / `remove(ids)`
3. `PersistentVectorStore`:构造时从 SqliteStore 加载已有索引;`add` 双向写入;`search` 纯内存搜索
4. `RagPipeline`:组合器封装 `split``embed``store.add` 的 ingest 流程,以及 `embed``store.search` 的 retrieve 流程
5. SqliteStore 存储格式:`vec:{namespace}:{doc_id}` → JSON `{doc_id, content, metadata, embedding}`
**依赖**Phase 14Document 类型)
**优先级**P0
**预估规模**:约 400 行
**状态**:⏳ 待实施
---
#### Phase 16: 摘要自动生成
**目标**:闭环长对话能力。v0.2 的 `inject_summary` 消费端(`FocusedConfig.summary_override`)已就绪,缺的是生产端。
**交付物**
1. `SummaryConfig` 结构体:`enabled` / `trigger_token_ratio`(默认 0.75/ `summary_prompt`(可自定义)
2.`OnTurnEnd` Hook 中插检查点:检测 token 水位超过 `trigger_token_ratio` → 调 LLM 生成摘要 → `SessionMemory::set("conversation_summary", summary)`
3. `AgentBuilder` 扩展:`.summary_config(cfg)` 方法
**为什么放 Hook 而非内置**:可插拔,默认不启用,用户 opt-in。不改变现有 `submit_turn` 行为。
**依赖**:无(Hook 系统 + SessionMemory 已就绪)
**优先级**P0
**预估规模**:约 150 行
**状态**:⏳ 待实施
---
#### Phase 17: Agent 执行引擎(会话树 + Time-travel Checkpointer
**目标**:建立 `engine/` 模块。解决 v0.2 中"session 在变量里、无法通过 ID 恢复、不支持父子关系"的空白。
**交付物**
1. `src/engine/` 新模块(`session_manager.rs` + `checkpointer.rs` + `error.rs`
2. `SessionManager`
- `create(agent, bundle) -> session_id` — 创建根 session
- `create_child(parent_id, child_id, agent)` — 创建子 session(继承父 `RuntimeBundle`
- `get(session_id) -> Arc<Mutex<AgentSession>>` — 按 ID 查找(支持从持久化恢复)
- `children(parent_id)` / `parent(child_id)` — 树形查询
- `destroy(id)` / `destroy_subtree(id)` — 生命周期管理
- `tree() -> SessionTreeSnapshot` — 树结构快照
3. `Checkpointer`
- `checkpoint(session)` — 每个 `submit_turn` 末尾自动保存全量状态快照
- `rollback(session_id, ckpt_id)` — 回滚到任意历史 checkpoint
- `fork(session_id, ckpt_id, new_id)` — 从历史 checkpoint 分支出新 session
- `list_checkpoints(session_id)` — 列出 checkpoint 列表
4. `AgentSession` 新增 `Serialize + Deserialize` 以支持 checkpoint 序列化
**Checkpoint 存储格式**`checkpoint:{session_id}:{ckpt_id}` → JSON(完整 AgentSession,含所有 slot 消息列表)。Ponytail:全量 JSON 够用,等遇到存储效率问题时再改增量模式。
**会话树持久化**`session_meta:{session_id}``{agent_name, parent_id, created_at, turn_count}``session_rel:{child_id}``"parent_id"`
**依赖**Phase 10ContextSlot 持久化 — 消息由 slot 自己管,Checkpointer 管执行状态)
**优先级**P0
**预估规模**:约 600 行
**状态**:⏳ 待实施
---
#### Phase 18: Agent Switch + SubAgent Dispatch + Agent 间交互
**目标**:在 SessionManager 基础上,提供 Agent 角色热切换和子代理调度能力。
**交付物**
1. `engine/switch.rs``switch_agent(session_id, new_agent)`:替换 `Arc<dyn Agent>`slot 历史 / turn_index / session_memory 全保留
2. `engine/sub_agent.rs` — SubAgent Dispatch 核心:
- `DispatchConfig``max_concurrency`(默认 10/ `inherit_session_memory`(默认 true/ `bridge_keys`
- `dispatch(parent_id, sub_agent, task, config) -> SubTaskResult`:创建子 session → 继承父 SessionMemory → `submit_turn` → 返回结构化结果
- `dispatch_stream(parent_id, sub_agent, task, config) -> SubTaskStream`:流式版
- `dispatch_all(parent_id, tasks, config) -> Vec<SubTaskResult>`:并行派发,`tokio::sync::Semaphore` 控制并发数
3. `SubTaskResult``child_id` / `response` / `usage` / `summary` + `child_memory(sm)` 读取子 SessionMemory
**Agent 间交互三层级**
- 父→子:继承 SessionMemory 快照 + `bridge_keys` 指定 key 强制注入 system prompt
- 子→父:`SubTaskResult` 结构化回传 + `SessionMemory["result_summary"]` 结论摘要
- 子↔子(间接):通过公共 `MemoryStore` namespace`shared:{parent_session_id}`)共享数据
**依赖**Phase 17SessionManager + 会话树)
**优先级**P0
**预估规模**:约 500 行
**状态**:⏳ 待实施
---
#### Phase 19: 知识图谱 + 双通道检索
**目标**:落地 `docs/note-knowledge-graph-design.md` 中记录的知识图谱设计,提供实体-关系图检索能力。扩展 `MemoryRetriever` 为双通道。
**交付物**
1. `src/memory/graph.rs`(新文件):
- `GraphEntity` / `GraphRelation` / `ScoredEntity` 核心类型
- `RelationDirection` 枚举(Outgoing / Incoming / Both
- `KnowledgeGraph` trait`add_entity` / `get_entity` / `remove_entity` / `add_relation` / `remove_relation` / `get_related` / `find_by_keywords` / `find_tags` / `set_entity_tags`
- `InMemoryGraph` 实现:`HashMap<String, GraphEntity>` + `Vec<GraphRelation>` + BFS 图遍历
- `TagConstraints``max_tags_per_entity` 默认 8
2. `src/memory/retriever.rs` 扩展:
- `MemoryRetriever` 增加 `knowledge_graph` 可选字段
- `RetrievalStrategy` 枚举:`Hybrid`(默认)/ `KnowledgeOnly` / `GraphOnly`
**与 Document 系统的关系**:知识图谱提供实体级检索("这个实体和什么相关"),VectorStore 提供语义相似度检索("哪些文档最相似"),两者互补。
**依赖**MemoryStore 持久化(v0.1 Phase 3
**优先级**P0
**预估规模**:约 400 行
**状态**:⏳ 待实施
---
### v0.3.0 Phase 依赖关系图
```mermaid
graph BT
P13["<b>Phase 13: 热身清理</b><br/>旧 types 文件删除<br/>ContextSlot fork/merge"]:::done
P14["<b>Phase 14: Document + Embedding</b><br/>Document 类型<br/>RecursiveCharacterSplitter<br/>Embedding trait"]:::done
P15["<b>Phase 15: 向量存储持久化</b><br/>VectorStore trait<br/>PersistentVectorStore<br/>RagPipeline"]:::pending
P16["<b>Phase 16: 摘要自动生成</b><br/>SummaryConfig<br/>OnTurnEnd Hook"]:::pending
P17["<b>Phase 17: 执行引擎</b><br/>SessionManager<br/>会话树<br/>Time-travel Checkpointer"]:::pending
P18["<b>Phase 18: 切换与调度</b><br/>Agent Switch<br/>SubAgent Dispatch<br/>dispatch_all 并发控制"]:::pending
P19["<b>Phase 19: 知识图谱</b><br/>KnowledgeGraph trait<br/>InMemoryGraph<br/>双通道检索"]:::pending
P15 --> P14
P18 --> P17
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
classDef pending fill:#fbbf24,stroke:#d97706,color:#1a1a1a
```
### 关键里程碑
| 里程碑 | Phase 完成条件 | 可验证指标 | 状态 |
|--------|---------------|-----------|------|
| **M9** | Phase 13 | 旧 types 文件删除、`cargo test --all-targets` 全绿、`fork`/`merge` 测试通过 | ✅ 2026-07-08 |
| **M10** | Phase 14 | `Document` + `RecursiveCharacterSplitter` 分割结果验证、`MockEmbedding` 测试通过 | ✅ 2026-07-09 |
| **M11** | Phase 15 | `PersistentVectorStore` 持久化 roundtrip、`RagPipeline::ingest → retrieve` 端到端验证 | ⏳ |
| **M12** | Phase 16 | 多轮对话后摘要自动写入 SessionMemory、派生 slot 时摘要正确注入 | ⏳ |
| **M13** | **Phase 17 (rc.1)** | `SessionManager` 创建/子树/恢复集成测试通过、`Checkpointer` checkpoint/rollback/fork 验证 | ⏳ |
| **M14** | Phase 18 | `switch_agent` 热切换验证、`dispatch`/`dispatch_all` 多轮对话 + 结果回传验证 | ⏳ |
| **M15** | Phase 19 | `KnowledgeGraph` 实体-关系 CRUD + `get_related` BFS 验证、双通道检索 Hybrid 策略验证 | ⏳ |
---
## v0.4+ 展望
### 已规划的功能
| 功能 | 说明 | 预计版本 |
|------|------|---------|
| Multi-Agent Swarm 编排 | Supervisor/Subgraph 模式,基于 v0.3 dispatch 构建 | v0.4 |
| Human-in-the-loop 审批 | `interrupt()` + `Command(resume=...)` 异步审批回调 | v0.4 |
| Agent 自动创生 | LLM 自主决定何时派发子 agent、派发什么角色 | v0.4 |
| 分布式 session 共享 | SessionManager Redis 后端支持跨进程 | v0.4 |
| 精确 tokenizer 计数 | 引入 `tiktoken-rs`,绑定模型具体 tokenizer,替换字符估算 | v0.4+ |
| TokenJuice 语义压缩 | 对工具结果做语义压缩而非字节截断 | v0.4+ |
| Markdown 技能按需加载 | 技能注册表 + 按 prompt 上下文动态加载 | v0.4+ |
| 增量 checkpoint | 仅存储变化部分,替换当前全量 JSON 模式 | v0.4+ |
| RL 轨迹导出 | ShareGPT 格式轨迹、Atropos 集成 | v0.4+ |
### 明确不做(agcore 范围外)
| 功能 | 原因 |
|------|------|
| TUI / 多平台 Gateway | 应用层职责(Feishu / Telegram / Discord 桥接) |
| 配置自动加载(config/figment) | 配置来源策略应由上游应用决定,agcore 不定义配置格式 |
| 提示词自动优化 | 属于智能层,不应内建于 core 库 |
--- ---
## 风险与建议 ## 风险与建议
1. **Phase 0 已完成**:LLM 调用周期基础设施已全部实现,可以支撑后续模块开发 1. **持久化依赖**`rusqlite` + `bundled` 零外部依赖编译,但 SQLite 不适配所有场景(分布式/高并发写)。`MemoryStore` trait 的抽象层允许下游自行实现 Redis / PostgreSQL 后端
2. **并行可能性**Phase 0 和 Phase 1 可并行开展(无相互依赖),可加速早期交付 2. **ContextSlot 心智负担**`ContextSlot` 引入了一等抽象的复杂度。建议通过 `AgentBuilder` 默认创建 `"default"` slot,让简单场景无感使用
3. **MCP 协议复杂性**MCP 涉及协议握手、session 管理、长期连接,建议预留充足时间调研协议细节 3. **向量检索规模上限**v0.3 的 `PersistentVectorStore` 全量加载到内存做余弦搜索,适合 ≤10 万条向量。超出此规模需换用专用向量库。v0.4 可以评估引入
4. **Scope 蔓延风险**当前 specs 只有 1 份文档,建议每个模块上线前都产出对应 spec,避免边实现边设计 4. **Scope 蔓延**v0.3 新增 `engine/` `vector/` `document/` 三个模块,功能覆盖扩展到多 Agent 基础系统。始终保持 trait + reference impl 的边界,业务循环留给上层
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` 5. **API 稳定性**v0.3 引入 `Checkpointer``SessionManager``VectorStore` 等新公开 APIv0.2 已有的 `#[non_exhaustive]``#[deprecated]` 机制继续沿用
6. **参考项目语言差异**OpenClaw / Hermes / OpenHarness 均为 Python/TypeScript 实现,OpenHuman 虽是 Rust + Tauri 但定位是桌面应用。借鉴时**只取架构模式**,不照搬具体实现(如 Pydantic 工具校验、SQLite Memory Tree、Node+Python 双进程等) 6. **Checkpointer 存储效率**:v0.3 使用全量 JSON 序列化存储 checkpoint,每轮对话约几百 KB。`fork` 从历史 checkpoint 创建新 session 时也会复制全量。等实际使用中发现存储瓶颈时再改为增量模式
--- ---
## 下一步行动 ## 下一步行动
1. **Phase 4c 已完成**Phase 4a + 4b + 4c 已交付(116 测试通过,0 clippy 警告)。可启动 v0.2+ 扩展评估(如多 Context 切换、Multi-Agent 协同等) 1. **v0.3.0 Phase 15 启动**:向量存储持久化(`VectorStore` trait + `InMemoryVectorStore` + `PersistentVectorStore` + `RagPipeline` 组合器),基于 Phase 14 的 `Document` 类型构建
2. **Context 切换备忘**`docs/note-context-switch-design.md` 记录了多 context 切换方案讨论,作为 v0.2+ 扩展项的输入 2. **Phase 15-19 顺次交付**:按依赖关系推进向量存储 → 摘要 → 引擎 → 调度 → 知识图谱
3. **参考项目调研沉淀**完成 OpenClaw / Hermes / OpenHuman / OpenHarness 横向调研,结果沉淀至 `docs/note-agent-harness-references.md`,作为 v0.2+ 扩展项的输入 3. **示例先行**完成一个 Phase 立即创建/更新对应示例,确保 `cargo run --example` 可验证
4. **Phase 3 备用设计就绪**`docs/note-knowledge-graph-design.md` 记录了 KnowledgeGraph、高级评分、RecallBased 淘汰等设计,v0.2+ 记忆扩展可直接参考 4. **里程碑追踪**:以 M10(Phase 14)为已达成里程碑,逐 Phase 推进 M11-M15
**已完成 / 进行中阶段** **已完成 / 进行中阶段**
- ✅ Phase 0 Foundation — 全部交付物已完成 - ✅ Phase 0 Foundation — 全部交付物已完成
@@ -338,9 +908,20 @@ graph BT
- ✅ Phase 4a Core Glue — 全部交付物已完成 - ✅ Phase 4a Core Glue — 全部交付物已完成
- ✅ Phase 4b Task Execution — 全部交付物已完成 - ✅ Phase 4b Task Execution — 全部交付物已完成
- ✅ Phase 4c Session Memory — 全部交付物已完成 - ✅ 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` - ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy
-LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除(方案:`docs/10c-phase2-llm-cycle-simplify.md` -Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
-v0.1 Release — 技术债扫清、MockProvider 公开化、7 个离线示例、README + 错误消息友好化、Roadmap 同步、CHANGELOG 初始化(计划:`docs/11-v0.1-release-plan.md` -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
-**Phase 11 测试与检索补强**`src/memory/vector.rs` 新增 `VectorRetriever` traitindex + search 抽象)+ `InMemoryVectorRetriever` 引用实现(HashMap + 全量余弦相似度扫描 + 零依赖 `dot()`),6 个内联测试覆盖 basic/empty/zero-vector/k=0/2 个并发;wiremock Provider roundtrip 测试 12 个(OpenAI 8 + Anthropic 4)覆盖请求体/header/401/429/500/529/流式 usage-only/流式错误/ToolUse/结构化错误体;`MemoryStore` 并发测试 5 个(InMemoryStore 3 + SqliteStore 2)覆盖 100 并发写、5 写+5 读混合 2 秒、15 写者容量淘汰;`openai.rs` `handle_error_response` 修复 429 retry-after 解析(5 行,与 anthropic 对齐);方案文档 `docs/18-phase11-testing-and-retrieval.md`(647 行,含 10 项架构决策 + 2 条实施偏差记录 #6 mid-stream mock 模式 + #7 retry-after 修复);全量 254 → 277+23 新测试),clippy 0 警告,doc 0 warning,并发测试 3 次稳定无 flaky
-**Phase 13 热身清理 + ContextSlot fork/merge** — 3 个旧 types 文件删除(`request.rs` 187 行 + `response.rs` 177 行 + `old_stream.rs` 45 行),所有 OpenAI wire-format 类型迁入 `provider/openai.rs` 可见性 `pub(crate)`Breaking Change:原 `agcore::llm::types::OpenaiChatRequest/Response/Chunk` 公共 re-export 路径已删除);`ChatResponse` 自 v0.1.0 标记 `#[deprecated]` 后在 Phase 13 整体删除;`ToolChoice``request.rs` 迁入 `tool.rs`(公共 `agcore::llm::types::ToolChoice` 路径不变);`ContextSlot::fork()` 派生独立子 slot`SlotSource::Derived { parent_id, strategy }` 血缘可追溯)+ `ContextSlot::merge(child, MergeStrategy)` 合入父 slot`Append` / `Replace` 两种策略,`#[non_exhaustive]` 为 Phase 16 `Summarize` 预留);`MergeStrategy` 防御性检查(self-merge / 跨 session / Readonly 目标全部阻断);`AgentSession::derive_slot` 重构复用 `fork()` 消除重复;`agent.rs` 追加 `MergeStrategy` re-export9 个 fork/merge 内联测试覆盖 happy path 与 error path`stream.rs` 简化为 module doc + `pub use` 重导出(保持 `use crate::llm::stream::StreamEvent` 路径兼容);方案文档 `docs/19-phase13-cleanup-and-fork-merge.md`640 行);全量 277 → 286+9 新测试),clippy 0 警告,doc 0 warning
- ✅ 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
-**v0.3.0 Phase 13 完成** — 技术债清理(3 旧 types 文件 + ChatResponse 删除)+ ContextSlot fork/merge9 新测试),M9 里程碑达成
-**v0.3.0 Phase 14 完成** — Document 类型(id/content/metadata/mime_type+ `RecursiveCharacterSplitter` 两阶段算法(按 separator 优先级递归分割 + 贪心合并 overlap,全部 `chars_len()` 字符级比较)+ `Embedding` traitasync + `LlmError` 复用)+ `MockEmbedding`sin-hash 零依赖伪随机 + L2 归一化)+ 19 Document 测试 + 6 Embedding 测试(含 1 个 split_multibyte_utf8_boundary CJK 边界测试);`src/document.rs`580 行)+ `src/llm/embedding.rs`183 行)+ `examples/document_demo.rs`74 行);`pub use document::Document` 在 lib.rs 重导出;CJK 分隔符(`。`/``/``)加入 `DEFAULT_SEPARATORS`;方案文档 `docs/20-phase14-document-and-embedding.md`1417 行);全量 286 → 313+27 新测试,0 失败),clippy 0 警告,doc 0 warning,零新外部依赖;M10 里程碑达成;Phase 15-19 共 5 个增量 Phase 待实施(向量存储 → 摘要 → 引擎 → 调度 → 知识图谱)
--- ---
+6 -7
View File
@@ -15,9 +15,9 @@ use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession}; use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor; use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider; use agcore::llm::mock::MockProvider;
use agcore::llm::types::Usage;
use agcore::llm::types::message::{ContentBlock, Message}; use agcore::llm::types::message::{ContentBlock, Message};
use agcore::llm::types::response_v2::{MessageResponse, StopReason}; use agcore::llm::types::response_v2::{MessageResponse, StopReason};
use agcore::llm::types::Usage;
use agcore::tools::ToolRegistry; use agcore::tools::ToolRegistry;
/// 计算器角色 Agent。 /// 计算器角色 Agent。
@@ -72,7 +72,10 @@ async fn main() {
// 4. 提交第一轮 // 4. 提交第一轮
println!("=== 提交第 1 轮 ==="); 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()); println!("LLM: {}", resp.text());
session session
.set_session_data("last_q", "1+1=?") .set_session_data("last_q", "1+1=?")
@@ -107,11 +110,7 @@ async fn main() {
// 8. 跨 session 数据隔离验证 // 8. 跨 session 数据隔离验证
println!("=== 数据隔离验证 ==="); println!("=== 数据隔离验证 ===");
let other = AgentSession::new( let other = AgentSession::new(Arc::new(CalculatorAgent), "other-session", bundle);
Arc::new(CalculatorAgent),
"other-session",
bundle,
);
assert!( assert!(
other.get_session_data("last_q").await.unwrap().is_none(), other.get_session_data("last_q").await.unwrap().is_none(),
"新会话不应看到旧 session 的 last_q" "新会话不应看到旧 session 的 last_q"
+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
}
+8 -16
View File
@@ -29,6 +29,7 @@ fn message_text(msg: &Message) -> &str {
.next() .next()
.unwrap_or(""), .unwrap_or(""),
Message::UserImage { .. } => "[image]", Message::UserImage { .. } => "[image]",
_ => "",
} }
} }
@@ -80,11 +81,8 @@ async fn main() {
// 3. 多角色混合 + clear // 3. 多角色混合 + clear
println!("\n=== 多角色写入 + clear ==="); println!("\n=== 多角色写入 + clear ===");
let store3 = Arc::new(InMemoryStore::new()); let store3 = Arc::new(InMemoryStore::new());
let mut memory3 = ConversationMemory::new( let mut memory3 =
store3, ConversationMemory::new(store3, "session-3", ConversationMemoryConfig::default());
"session-3",
ConversationMemoryConfig::default(),
);
memory3 memory3
.add_message(Message::user_text("你好")) .add_message(Message::user_text("你好"))
.await .await
@@ -98,7 +96,9 @@ async fn main() {
.await .await
.unwrap(); .unwrap();
memory3 memory3
.add_message(Message::assistant("我无法查询实时天气,但你可以查看天气应用。")) .add_message(Message::assistant(
"我无法查询实时天气,但你可以查看天气应用。",
))
.await .await
.unwrap(); .unwrap();
println!( println!(
@@ -119,16 +119,8 @@ async fn main() {
// 4. Session 隔离 // 4. Session 隔离
println!("\n=== Session 隔离(共用 InMemoryStore==="); println!("\n=== Session 隔离(共用 InMemoryStore===");
let store4 = Arc::new(InMemoryStore::new()); let store4 = Arc::new(InMemoryStore::new());
let mut a = ConversationMemory::new( let mut a = ConversationMemory::new(store4.clone(), "s-a", ConversationMemoryConfig::default());
store4.clone(), let mut b = ConversationMemory::new(store4.clone(), "s-b", ConversationMemoryConfig::default());
"s-a",
ConversationMemoryConfig::default(),
);
let mut b = ConversationMemory::new(
store4.clone(),
"s-b",
ConversationMemoryConfig::default(),
);
a.add_message(Message::user_text("A 的消息")).await.unwrap(); a.add_message(Message::user_text("A 的消息")).await.unwrap();
b.add_message(Message::user_text("B 的消息")).await.unwrap(); b.add_message(Message::user_text("B 的消息")).await.unwrap();
println!( println!(
+10 -15
View File
@@ -17,7 +17,7 @@ use agcore::tools::{
ToolRegistry, ToolRegistry,
}; };
use async_trait::async_trait; use async_trait::async_trait;
use serde_json::{json, Value}; use serde_json::{Value, json};
/// 天气查询工具 —— 模拟根据城市返回天气数据。 /// 天气查询工具 —— 模拟根据城市返回天气数据。
struct WeatherTool; struct WeatherTool;
@@ -42,11 +42,7 @@ impl BaseTool for WeatherTool {
fn required_permissions(&self) -> Vec<Permission> { fn required_permissions(&self) -> Vec<Permission> {
vec![Permission::Network] vec![Permission::Network]
} }
async fn execute( async fn execute(&self, args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
&self,
args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
let city = args["city"].as_str().unwrap_or("未知"); let city = args["city"].as_str().unwrap_or("未知");
// 模拟查询:根据城市名给出不同温度 // 模拟查询:根据城市名给出不同温度
let (temperature, condition) = match city { let (temperature, condition) = match city {
@@ -84,11 +80,7 @@ impl BaseTool for DeleteFileTool {
fn required_permissions(&self) -> Vec<Permission> { fn required_permissions(&self) -> Vec<Permission> {
vec![Permission::Delete] vec![Permission::Delete]
} }
async fn execute( async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
&self,
_args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
Ok(json!({"deleted": true})) Ok(json!({"deleted": true}))
} }
} }
@@ -138,9 +130,8 @@ async fn main() {
// 5. 权限检查:默认 PermissionConfig 黑名单含 Delete // 5. 权限检查:默认 PermissionConfig 黑名单含 Delete
println!("\n=== 权限检查(默认 PermissionConfigdenied = [Delete, Shell]==="); println!("\n=== 权限检查(默认 PermissionConfigdenied = [Delete, Shell]===");
let mut registry_with_checker = ToolRegistry::new().with_permission_checker(PermissionChecker::new( let mut registry_with_checker = ToolRegistry::new()
PermissionConfig::default(), .with_permission_checker(PermissionChecker::new(PermissionConfig::default()));
));
registry_with_checker registry_with_checker
.register(Arc::new(WeatherTool) as ToolRef) .register(Arc::new(WeatherTool) as ToolRef)
.unwrap(); .unwrap();
@@ -155,7 +146,11 @@ async fn main() {
.unwrap(); .unwrap();
println!( println!(
"get_weather 权限检查: {}", "get_weather 权限检查: {}",
if r.output.is_ok() { "通过 ✓" } else { "阻断 ✗" } if r.output.is_ok() {
"通过 ✓"
} else {
"阻断 ✗"
}
); );
// delete_file 声明 Delete → 在 denied 列表 → 阻断 // delete_file 声明 Delete → 在 denied 列表 → 阻断
+70
View File
@@ -0,0 +1,70 @@
//! document_demo —— Document + RecursiveCharacterSplitter + MockEmbedding + RagPipeline 完整衔接示例。
//!
//! 演示 RAG 管线:
//! 1. 创建多段落 Document
//! 2. RecursiveCharacterSplitter 分割为 chunk
//! 3. RagPipeline.ingest() 自动嵌入并存储
//! 4. RagPipeline.retrieve() 做语义检索
//!
//! 运行:`cargo run --example document_demo`(离线,零配置)
use std::sync::Arc;
use agcore::document::{Document, RecursiveCharacterSplitter};
use agcore::llm::embedding::{Embedding, MockEmbedding};
use agcore::memory::{InMemoryVectorStore, RagPipeline, VectorStore};
#[tokio::main]
async fn main() {
agcore::init_tracing();
// 1. 创建多段落 Document(含中英文混合)
let doc = Document::new(
"rust-intro",
"Rust 是一门系统编程语言,注重安全、并发和性能。\n\n\
Rust \
\n\n\
Rust 线\
Send Sync trait 线\n\n\
Rust C/C++ \
Cargo 使",
"text/markdown",
);
println!("输入文档: {} 字符", doc.content.chars().count());
// 2. 构造 RAG 管线(嵌入器 + 向量存储 + 分割器)
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let splitter = RecursiveCharacterSplitter::new(200, 30);
let pipeline = RagPipeline::new(
Arc::clone(&embedder),
Arc::clone(&store),
Some(splitter),
);
// 3. 一次性 ingest:自动 split → embed → add
pipeline.ingest(std::slice::from_ref(&doc)).await.unwrap();
// 4. 模拟查询:复用第一个 chunk 的 content 作为查询文本
let chunks_in_store = store.search(&[1.0, 0.0, 0.0, 0.0], 1).await.unwrap();
assert!(!chunks_in_store.is_empty(), "ingest 后 store 应有数据");
let query_text = &chunks_in_store[0].0.content;
let results = pipeline.retrieve(query_text, 3).await.unwrap();
println!("\nTop 3 检索结果(与第一个 chunk 相似):");
for (doc, score) in &results {
println!(
" id={}, score={:.4}, content={}",
doc.id, score, doc.content
);
}
assert!(!results.is_empty(), "至少应返回 1 条检索结果");
assert!(
results[0].0.id.starts_with("rust-intro:chunk:0000"),
"Top 1 应为 chunk 0 自身"
);
println!("\n✓ document_demo 完成");
}
+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✓ 端到端演示完成");
}
+19 -19
View File
@@ -38,14 +38,26 @@ async fn main() {
let ks = KnowledgeStore::new(store); let ks = KnowledgeStore::new(store);
let pages = vec![ let pages = vec![
make_page("rust-1", "Rust 入门", "Rust 是一门系统级编程语言,注重安全性与并发。"), make_page(
make_page("python-1", "Python 简介", "Python 是一门动态类型的高级编程语言。"), "rust-1",
"Rust 入门",
"Rust 是一门系统级编程语言,注重安全性与并发。",
),
make_page(
"python-1",
"Python 简介",
"Python 是一门动态类型的高级编程语言。",
),
make_page( make_page(
"langgraph-1", "langgraph-1",
"LangGraph 框架", "LangGraph 框架",
"LangGraph 是 LangChain 的状态图扩展,用于构建多步 Agent。", "LangGraph 是 LangChain 的状态图扩展,用于构建多步 Agent。",
), ),
make_page("rust-async", "Rust 异步编程", "Rust 异步基于 tokio 与 futures 抽象。"), make_page(
"rust-async",
"Rust 异步编程",
"Rust 异步基于 tokio 与 futures 抽象。",
),
]; ];
for p in &pages { for p in &pages {
ks.add_page(p.clone()).await.expect("保存页面失败"); ks.add_page(p.clone()).await.expect("保存页面失败");
@@ -62,14 +74,8 @@ async fn main() {
let result = retriever.retrieve("Rust 异步").await.unwrap(); let result = retriever.retrieve("Rust 异步").await.unwrap();
println!("query: {}", result.query); println!("query: {}", result.query);
for item in &result.items { for item in &result.items {
println!( println!(" 命中: {} (score={:.3})", item.page.title, item.score);
" 命中: {} (score={:.3})", assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
item.page.title, item.score
);
assert!(
(0.0..=1.0).contains(&item.score),
"score 应在 [0, 1] 区间"
);
} }
assert!(!result.items.is_empty(), "应至少命中一个页面"); assert!(!result.items.is_empty(), "应至少命中一个页面");
@@ -85,14 +91,8 @@ async fn main() {
min_score: 0.5, min_score: 0.5,
}; };
let retriever2 = MemoryRetriever::new(ks2, cfg); let retriever2 = MemoryRetriever::new(ks2, cfg);
let result = retriever2 let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
.retrieve("完全不相关的火锅配方") println!("无关 query → items.len = {} (期望 0)", result.items.len());
.await
.unwrap();
println!(
"无关 query → items.len = {} (期望 0)",
result.items.len()
);
assert!(result.items.is_empty()); assert!(result.items.is_empty());
// 4. max_results 截断 // 4. max_results 截断
+10 -6
View File
@@ -11,7 +11,7 @@
use agcore::llm::types::message::{ContentBlock, Message}; use agcore::llm::types::message::{ContentBlock, Message};
use agcore::prompt::{ use agcore::prompt::{
validate_messages, PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext, PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext, validate_messages,
}; };
fn message_text(msg: &Message) -> String { fn message_text(msg: &Message) -> String {
@@ -27,16 +27,16 @@ fn message_text(msg: &Message) -> String {
}) })
.collect(), .collect(),
Message::UserImage { .. } => "[image]".into(), Message::UserImage { .. } => "[image]".into(),
_ => String::new(),
} }
} }
fn main() { fn main() {
// 1. PromptTemplate::compile + render —— 直接构造模板 // 1. PromptTemplate::compile + render —— 直接构造模板
println!("=== PromptTemplate::compile + render ==="); println!("=== PromptTemplate::compile + render ===");
let tpl = PromptTemplate::compile( let tpl =
"今日 {{location}} 天气:{{condition}},温度 {{temperature}}", PromptTemplate::compile("今日 {{location}} 天气:{{condition}},温度 {{temperature}}")
) .expect("编译失败");
.expect("编译失败");
let mut ctx = TemplateContext::new(); let mut ctx = TemplateContext::new();
ctx.insert("location", "北京"); ctx.insert("location", "北京");
ctx.insert("condition", ""); ctx.insert("condition", "");
@@ -58,7 +58,10 @@ fn main() {
.register("weather", "今日 {{location}}{{condition}}") .register("weather", "今日 {{location}}{{condition}}")
.expect("注册失败"); .expect("注册失败");
registry registry
.register("greet", "你好 {{name}}{{#if formal}} 见到您很荣幸。{{/if}}") .register(
"greet",
"你好 {{name}}{{#if formal}} 见到您很荣幸。{{/if}}",
)
.expect("注册失败"); .expect("注册失败");
let mut ctx = TemplateContext::new(); let mut ctx = TemplateContext::new();
@@ -88,6 +91,7 @@ fn main() {
Message::User { .. } | Message::UserImage { .. } => "user", Message::User { .. } | Message::UserImage { .. } => "user",
Message::Assistant { .. } => "assistant", Message::Assistant { .. } => "assistant",
Message::ToolResult { .. } => "tool", Message::ToolResult { .. } => "tool",
_ => "unknown",
}; };
println!("[{i}] {role}: {}", message_text(m)); println!("[{i}] {role}: {}", message_text(m));
} }
+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::init_tracing;
use agcore::llm::{ use agcore::llm::{
cycle::{CycleConfig, LlmCycle}, cycle::{CycleConfig, LlmCycle},
provider::{create_provider, ProviderConfig, ProviderType}, provider::{ProviderConfig, ProviderType, create_provider},
types::{message::ContentBlock, message::Message, response_v2::MessageResponse}, types::{message::ContentBlock, message::Message, response_v2::MessageResponse},
}; };
@@ -51,10 +51,11 @@ async fn main() {
base_url, base_url,
api_key, api_key,
model: model.clone(), model: model.clone(),
timeout_secs: 30,
max_retries: 3,
}; };
let provider = create_provider(provider_type, config) let provider = create_provider(provider_type, config).expect("创建 Provider 失败");
.expect("创建 Provider 失败");
let cycle_config = CycleConfig { let cycle_config = CycleConfig {
model, model,
@@ -63,9 +64,9 @@ async fn main() {
..CycleConfig::default() ..CycleConfig::default()
}; };
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![ let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![Message::system(
Message::system("你是一个简洁的助手,对于任何问题都是用一句话回答。"), "你是一个简洁的助手,对于任何问题都是用一句话回答。",
]); )]);
println!("发送请求..."); println!("发送请求...");
+2 -4
View File
@@ -17,9 +17,9 @@ use std::sync::Arc;
use agcore::llm::cycle::{CycleConfig, LlmCycle}; use agcore::llm::cycle::{CycleConfig, LlmCycle};
use agcore::llm::mock::MockProvider; use agcore::llm::mock::MockProvider;
use agcore::llm::provider::LlmProvider; use agcore::llm::provider::LlmProvider;
use agcore::llm::types::Usage;
use agcore::llm::types::message::{ContentBlock, Message}; use agcore::llm::types::message::{ContentBlock, Message};
use agcore::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent}; use agcore::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
use agcore::llm::types::Usage;
use futures_util::StreamExt; use futures_util::StreamExt;
/// 构造预设的纯文本响应。 /// 构造预设的纯文本响应。
@@ -99,9 +99,7 @@ async fn main() {
// 上层 Agent 通过 `match` 或 `?` 处理 `AgentError::Llm(_)`。 // 上层 Agent 通过 `match` 或 `?` 处理 `AgentError::Llm(_)`。
println!("\n=== 阶段 2:错误路径(队列耗尽)==="); println!("\n=== 阶段 2:错误路径(队列耗尽)===");
let mut cycle = LlmCycle::new_with_arc(dyn_provider, CycleConfig::default()); let mut cycle = LlmCycle::new_with_arc(dyn_provider, CycleConfig::default());
let result = cycle let result = cycle.submit_stream("第二次提问".to_string(), vec![]).await;
.submit_stream("第二次提问".to_string(), vec![])
.await;
match result { match result {
Ok(_) => panic!("阶段 2 必须失败(队列耗尽)"), Ok(_) => panic!("阶段 2 必须失败(队列耗尽)"),
Err(e) => { Err(e) => {
+25 -24
View File
@@ -8,27 +8,13 @@
//! 5. 错误路径:非法 JSON / 空 steps / 缺字段 → `AgentError::PlanParse` //! 5. 错误路径:非法 JSON / 空 steps / 缺字段 → `AgentError::PlanParse`
//! //!
//! 运行:`cargo run --example task_agent_demo` //! 运行:`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::agent::{AgentError, JsonPlanParser, PlanParser, Step, StepStatus};
use agcore::llm::types::openai_message::OpenaiChatMessage; use agcore::llm::types::message::Message;
use agcore::llm::types::shared::FinishReason; use agcore::llm::types::response_v2::{MessageResponse, StopReason};
use agcore::llm::types::{ChatResponse, Usage}; use agcore::llm::types::Usage;
#[tokio::main] #[tokio::main]
async fn main() { async fn main() {
@@ -67,21 +53,36 @@ async fn main() {
assert!(step.status.is_pending()); assert!(step.status.is_pending());
step.status = StepStatus::Running; 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 { step.status = StepStatus::Completed(MessageResponse {
message: OpenaiChatMessage::assistant_text("天气:晴,22°C"), id: String::new(),
model: "mock".into(),
message: Message::assistant("天气:晴,22°C"),
usage: Usage::from_input_output(5, 10), 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()); assert!(step.status.is_terminal());
// 3. 失败路径 // 3. 失败路径
println!("\n=== Step 状态机:失败路径 ==="); println!("\n=== Step 状态机:失败路径 ===");
let mut fail_step = Step::new(0, "调用天气 API"); let mut fail_step = Step::new(0, "调用天气 API");
fail_step.status = StepStatus::Failed(AgentError::Other("API 不可用".into())); 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()); assert!(fail_step.status.is_terminal());
// 4. 跳过路径 // 4. 跳过路径
+6 -1
View File
@@ -11,6 +11,7 @@
pub mod agent; pub mod agent;
pub mod builder; pub mod builder;
pub mod context;
pub mod error; pub mod error;
pub mod runtime; pub mod runtime;
pub mod session; pub mod session;
@@ -20,9 +21,13 @@ pub mod task;
// 重导出公共 API(按使用频度排序) // 重导出公共 API(按使用频度排序)
pub use agent::Agent; pub use agent::Agent;
pub use builder::AgentBuilder; pub use builder::AgentBuilder;
pub use context::{
ContextBudget, ContextSlot, DeriveStrategy, FocusedConfig, MergeStrategy, SlotConfig,
SlotMeta, SlotMode, SlotSource,
};
pub use error::AgentError; pub use error::AgentError;
pub use runtime::{AgentConfig, RuntimeBundle}; pub use runtime::{AgentConfig, RuntimeBundle};
pub use session::AgentSession; pub use session::AgentSession;
pub use session_memory::SessionMemory; pub use session_memory::SessionMemory;
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
pub use task::JsonPlanParser; pub use task::JsonPlanParser;
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
+2 -4
View File
@@ -7,14 +7,12 @@
//! - **不绑定业务循环**`submit_turn` 在 `AgentSession` 上,不在 trait 上 //! - **不绑定业务循环**`submit_turn` 在 `AgentSession` 上,不在 trait 上
use crate::agent::runtime::RuntimeBundle; use crate::agent::runtime::RuntimeBundle;
#[allow(deprecated)] use crate::llm::types::tool::ToolDef;
use crate::llm::types::ToolDefinition;
/// Agent 角色抽象。 /// Agent 角色抽象。
/// ///
/// 实现此 trait 即可接入 Agent Runtime。典型实现是 struct 持有静态配置(name、system prompt 模板), /// 实现此 trait 即可接入 Agent Runtime。典型实现是 struct 持有静态配置(name、system prompt 模板),
/// 也可以是基于配置动态生成的轻量实现。 /// 也可以是基于配置动态生成的轻量实现。
#[allow(deprecated)]
pub trait Agent: Send + Sync { pub trait Agent: Send + Sync {
/// 角色名(用于日志、调试、UI 展示)。 /// 角色名(用于日志、调试、UI 展示)。
fn name(&self) -> &str; fn name(&self) -> &str;
@@ -26,7 +24,7 @@ pub trait Agent: Send + Sync {
/// ///
/// **默认实现**:从 `bundle.tool_registry` 取全部工具(最常用模式)。 /// **默认实现**:从 `bundle.tool_registry` 取全部工具(最常用模式)。
/// **子 trait / 具体实现可覆盖**:做白名单、过滤、按状态动态调整等。 /// **子 trait / 具体实现可覆盖**:做白名单、过滤、按状态动态调整等。
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDefinition> { fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDef> {
bundle.tool_registry.definitions() bundle.tool_registry.definitions()
} }
} }
+8 -6
View File
@@ -92,15 +92,17 @@ impl AgentBuilder {
/// `AgentError::Config(...)`,提示调用 `.provider(...)` / `.tool_registry(...)` / /// `AgentError::Config(...)`,提示调用 `.provider(...)` / `.tool_registry(...)` /
/// `.hook_executor(...)` 补齐。不 panic。 /// `.hook_executor(...)` 补齐。不 panic。
pub fn build(self) -> Result<RuntimeBundle, AgentError> { pub fn build(self) -> Result<RuntimeBundle, AgentError> {
let provider = self let provider = self.provider.ok_or_else(|| {
.provider AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into())
.ok_or_else(|| AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into()))?; })?;
let tool_registry = self let tool_registry = self
.tool_registry .tool_registry
.ok_or_else(|| AgentError::Config("缺少 tool_registry,请先调用 .tool_registry(...)(即使是空 ToolRegistry 也需要传入)".into()))?; .ok_or_else(|| AgentError::Config("缺少 tool_registry,请先调用 .tool_registry(...)(即使是空 ToolRegistry 也需要传入)".into()))?;
let hook_executor = self let hook_executor = self.hook_executor.ok_or_else(|| {
.hook_executor AgentError::Config(
.ok_or_else(|| AgentError::Config("缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into()))?; "缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into(),
)
})?;
let config = self.config.unwrap_or_default(); let config = self.config.unwrap_or_default();
+1123
View File
File diff suppressed because it is too large Load Diff
+54 -3
View File
@@ -18,6 +18,7 @@ use crate::tools::error::ToolError;
/// **不实现 `Clone`**:透传内层 `LlmError` / `MemoryError`,两者均未派生 `Clone`(保留 /// **不实现 `Clone`**:透传内层 `LlmError` / `MemoryError`,两者均未派生 `Clone`(保留
/// 完整错误信息,传递所有权)。如需在多 session 间共享错误状态,用 `Arc<AgentError>` 包装。 /// 完整错误信息,传递所有权)。如需在多 session 间共享错误状态,用 `Arc<AgentError>` 包装。
#[derive(Debug, Error)] #[derive(Debug, Error)]
#[non_exhaustive]
pub enum AgentError { pub enum AgentError {
/// LLM 调用错误(透传 Phase 0)。 /// LLM 调用错误(透传 Phase 0)。
#[error("LLM 错误: {0}")] #[error("LLM 错误: {0}")]
@@ -35,6 +36,18 @@ pub enum AgentError {
#[error("Plan 解析错误: {0}")] #[error("Plan 解析错误: {0}")]
PlanParse(String), 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 层特有)。 /// 钩子阻断操作(Agent 层特有)。
#[error("钩子阻断: {0}")] #[error("钩子阻断: {0}")]
HookBlocked(String), HookBlocked(String),
@@ -59,6 +72,7 @@ impl AgentError {
/// - `Tool`:由内层 `is_recoverable()` 决定 /// - `Tool`:由内层 `is_recoverable()` 决定
/// - `HookBlocked` / `LimitExceeded`:不可恢复(需人工介入或终止循环) /// - `HookBlocked` / `LimitExceeded`:不可恢复(需人工介入或终止循环)
/// - `Config` / `Other`:不可恢复 /// - `Config` / `Other`:不可恢复
/// - `SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`:不可恢复(结构性错误)
pub fn is_recoverable(&self) -> bool { pub fn is_recoverable(&self) -> bool {
match self { match self {
Self::Llm(e) => matches!( Self::Llm(e) => matches!(
@@ -68,9 +82,13 @@ impl AgentError {
Self::Tool(e) => e.is_recoverable(), Self::Tool(e) => e.is_recoverable(),
Self::Memory(e) => e.is_recoverable(), Self::Memory(e) => e.is_recoverable(),
Self::PlanParse(_) => false, Self::PlanParse(_) => false,
Self::HookBlocked(_) | Self::LimitExceeded(_) | Self::Config(_) | Self::Other(_) => { Self::SlotReadonly(_)
false | Self::SlotNotFound(_)
} | Self::SlotAlreadyExists(_)
| Self::HookBlocked(_)
| Self::LimitExceeded(_)
| Self::Config(_)
| Self::Other(_) => false,
} }
} }
} }
@@ -180,4 +198,37 @@ mod tests {
let err = caller().unwrap_err(); let err = caller().unwrap_err();
assert!(matches!(err, AgentError::Memory(_))); 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 std::time::Duration;
use crate::llm::compact::CompactConfig; use crate::llm::compact::CompactConfig;
use crate::llm::provider::LlmProvider;
use crate::llm::hooks::HookExecutor; use crate::llm::hooks::HookExecutor;
use crate::llm::provider::LlmProvider;
use crate::memory::retriever::MemoryRetriever; use crate::memory::retriever::MemoryRetriever;
use crate::memory::store::MemoryStore; use crate::memory::store::MemoryStore;
use crate::tools::ToolRegistry; use crate::tools::ToolRegistry;
+821 -101
View File
File diff suppressed because it is too large Load Diff
+2 -10
View File
@@ -78,11 +78,7 @@ impl SessionMemory {
prefix: Some(format!("{}:", self.namespace)), prefix: Some(format!("{}:", self.namespace)),
..Default::default() ..Default::default()
}; };
let items = self let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
.store
.list(&filter)
.await
.map_err(AgentError::Memory)?;
let mut lines = Vec::with_capacity(items.len() + 2); let mut lines = Vec::with_capacity(items.len() + 2);
lines.push("<session-context>".to_string()); lines.push("<session-context>".to_string());
@@ -113,11 +109,7 @@ impl SessionMemory {
prefix: Some(format!("{}:", self.namespace)), prefix: Some(format!("{}:", self.namespace)),
..Default::default() ..Default::default()
}; };
let items = self let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
.store
.list(&filter)
.await
.map_err(AgentError::Memory)?;
for item in items { for item in items {
self.store self.store
+5 -11
View File
@@ -10,8 +10,7 @@
//! - 重试由上层新建 `Plan` 实现,`TaskAgent` 不做自动重试 //! - 重试由上层新建 `Plan` 实现,`TaskAgent` 不做自动重试
use crate::agent::error::AgentError; use crate::agent::error::AgentError;
#[allow(deprecated)] use crate::llm::types::response_v2::MessageResponse;
use crate::llm::types::ChatResponse;
use async_trait::async_trait; use async_trait::async_trait;
@@ -56,14 +55,14 @@ impl Step {
/// 均未派生 `Clone`(保留原始错误信息,传递所有权而非克隆)。如需复制 `Plan`, /// 均未派生 `Clone`(保留原始错误信息,传递所有权而非克隆)。如需复制 `Plan`,
/// 只能 clone 处于 `Pending` / `Running` / `Completed` / `Skipped` 状态的步骤。 /// 只能 clone 处于 `Pending` / `Running` / `Completed` / `Skipped` 状态的步骤。
#[derive(Debug)] #[derive(Debug)]
#[allow(deprecated)] #[non_exhaustive]
pub enum StepStatus { pub enum StepStatus {
/// 初始状态 —— 等待执行。 /// 初始状态 —— 等待执行。
Pending, Pending,
/// 正在执行(`TaskAgent::execute_plan` 进入)。 /// 正在执行(`TaskAgent::execute_plan` 进入)。
Running, Running,
/// 已完成(含 LLM 响应)。 /// 已完成(含 LLM 响应)。
Completed(ChatResponse), Completed(MessageResponse),
/// 失败(含错误)。 /// 失败(含错误)。
Failed(AgentError), Failed(AgentError),
/// 跳过(上层主动跳过)。 /// 跳过(上层主动跳过)。
@@ -130,9 +129,7 @@ impl PlanParser for JsonPlanParser {
.collect::<Result<Vec<_>, AgentError>>()?; .collect::<Result<Vec<_>, AgentError>>()?;
if steps.is_empty() { if steps.is_empty() {
return Err(AgentError::PlanParse( return Err(AgentError::PlanParse("Plan 至少需要一个步骤".into()));
"Plan 至少需要一个步骤".into(),
));
} }
Ok(Plan { Ok(Plan {
@@ -203,10 +200,7 @@ mod tests {
let plan = Plan { let plan = Plan {
id: "p1".into(), id: "p1".into(),
goal: "test goal".into(), goal: "test goal".into(),
steps: vec![ steps: vec![Step::new(0, "first"), Step::new(1, "second")],
Step::new(0, "first"),
Step::new(1, "second"),
],
}; };
assert_eq!(plan.steps.len(), 2); assert_eq!(plan.steps.len(), 2);
assert_eq!(plan.steps[0].index, 0); assert_eq!(plan.steps[0].index, 0);
+580
View File
@@ -0,0 +1,580 @@
//! Document 系统 —— 文本分割与文档类型。
//!
//! 提供 [`Document`] 数据结构和 [`RecursiveCharacterSplitter`] 分割器,
//! 作为 RAG 管线(split → embed → store)的前置步骤。
use std::collections::HashMap;
/// 默认分隔符优先级列表(按优先级降序)。
///
/// 段落级 → 行级 → 句子级(含 CJK 标点) → 词级 → 字符级(兜底)。
/// 在 LangChain 基础上扩充了 CJK 句号 `"。"`、问号 `""`、感叹号 `""`,
/// 确保中文文本在句子边界有更高分割质量。
const DEFAULT_SEPARATORS: &[&str] = &["\n\n", "\n", "", "", "", ".", " ", ""];
/// 文档片段 —— RAG 管线的基本数据载体。
///
/// 作为分割(split)和向量化(embed)两个阶段的通货类型,
/// 在 Phase 15 的 RagPipeline 中串联 split → embed → store。
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Document {
/// 文档唯一标识。
pub id: String,
/// 文档文本内容。
pub content: String,
/// 元数据标签(键值对,可用作过滤、溯源、分类)。
pub metadata: HashMap<String, String>,
/// MIME 类型,标识内容格式(如 "text/plain", "text/markdown")。
pub mime_type: String,
}
impl Document {
/// 创建一个新文档。元数据默认初始化为空。
///
/// 分割器产生的 chunks 会自动继承源文档 mime_type
/// 并在 metadata 中追加 source_id / chunk_index / chunk_count。
pub fn new(
id: impl Into<String>,
content: impl Into<String>,
mime_type: impl Into<String>,
) -> Self {
Self {
id: id.into(),
content: content.into(),
metadata: HashMap::new(),
mime_type: mime_type.into(),
}
}
/// 快速构造纯文本文档(mime_type 默认为 "text/plain")。
/// 适用于大多数无需指定媒体类型的场景。
pub fn from_raw(id: impl Into<String>, content: impl Into<String>) -> Self {
Self {
id: id.into(),
content: content.into(),
metadata: HashMap::new(),
mime_type: "text/plain".into(),
}
}
}
/// 递归字符级文档分割器。
///
/// 使用可配置的分隔符优先级列表,递归地将文档分割为
/// 接近 chunk_size 的块。
///
/// # 算法(两阶段)
///
/// 1. **递归分割**:按分隔符优先级从高到低递归切割文本,
/// 产生初始片段(均 ≤ chunk_size,按字符数计算)。
///
/// 2. **贪心合并**:从左向右合并相邻片段,直到合计字符数
/// 超过 chunk_size,此时将前一组合并结果作为一个 chunk 输出,
/// 并携带 chunk_overlap 字符的滑动窗口。
///
/// 所有长度比较均以 Unicode 字符数为单位(`text.chars().count()`),
/// 而非字节数。CJK 文本每个字算 1 个 char。
///
/// # 升级路径
///
/// - 如需自定义分割函数,可在上层通过 `with_custom_splitter`
/// 扩展(当前未实现,预留升级路径)。
/// - 如需 unicode 感知的句子分割(如中文句号、缩写处理),
/// 可在 separators 中加入对应字符串,或将下游替换为
/// 基于 unicode-segmentation crate 的自定义分割器。
#[derive(Debug, Clone)]
pub struct RecursiveCharacterSplitter {
chunk_size: usize,
chunk_overlap: usize,
separators: Vec<String>,
}
impl RecursiveCharacterSplitter {
/// 创建分割器。
///
/// # Panics
///
/// - 如果 `chunk_size == 0`
/// - 如果 `chunk_size ≤ chunk_overlap`(无法形成有效滑动窗口)
pub fn new(chunk_size: usize, chunk_overlap: usize) -> Self {
if chunk_size == 0 {
panic!("chunk_size must be greater than 0");
}
if chunk_size <= chunk_overlap {
panic!("chunk_size must be greater than chunk_overlap");
}
Self {
chunk_size,
chunk_overlap,
separators: DEFAULT_SEPARATORS.iter().map(|s| s.to_string()).collect(),
}
}
/// 创建分割器的安全版本。
///
/// 验证失败时返回 `Err` 而非 panic。
pub fn try_new(chunk_size: usize, chunk_overlap: usize) -> Result<Self, &'static str> {
if chunk_size == 0 {
return Err("chunk_size must be greater than 0");
}
if chunk_size <= chunk_overlap {
return Err("chunk_size must be greater than chunk_overlap");
}
Ok(Self {
chunk_size,
chunk_overlap,
separators: DEFAULT_SEPARATORS.iter().map(|s| s.to_string()).collect(),
})
}
/// 覆盖默认分隔符优先级列表。
///
/// **重要**:建议保留 `""` 作为最后一个 separator
/// 作为字符级兜底防止任何文本都能被分割。
pub fn with_separators(mut self, separators: Vec<String>) -> Self {
self.separators = separators;
self
}
/// 返回 chunk_size(字符数)。
pub fn chunk_size(&self) -> usize {
self.chunk_size
}
/// 返回 chunk_overlap(字符数)。
pub fn chunk_overlap(&self) -> usize {
self.chunk_overlap
}
/// 批量分割。
///
/// 每个输入文档独立分割。输出 chunks 继承源文档的 mime_type
/// 并在 metadata 中追加 source_id / chunk_index / chunk_count。
///
/// Chunk ID 格式:`{source_id}:chunk:{index:04d}`
/// 例如 `"doc_001:chunk:0000"`(索引从 0 开始,4 位固定宽度)。
///
/// **注意**metadata 注入使用 `HashMap::insert()`,如果源 Document
/// 的 metadata 已包含 `"source_id"`、`"chunk_index"` 或 `"chunk_count"`
/// 键,将被分割器的值静默覆盖。
pub fn split(&self, documents: &[Document]) -> Vec<Document> {
tracing::debug!(
input_count = documents.len(),
"RecursiveCharacterSplitter::split start"
);
let mut output = Vec::new();
for doc in documents {
let segments = self.split_text(&doc.content, &self.separators);
let chunks = self.merge_with_overlap(segments);
debug_assert!(
chunks.len() < 10_000,
"单个文档产生超过 9999 个 chunk,索引格式溢出"
);
tracing::trace!(
doc_id = %doc.id,
chunk_count = chunks.len(),
"document split into chunks"
);
for (idx, chunk_text) in chunks.iter().enumerate() {
let mut metadata = doc.metadata.clone();
metadata.insert("source_id".to_string(), doc.id.clone());
metadata.insert("chunk_index".to_string(), idx.to_string());
metadata.insert("chunk_count".to_string(), chunks.len().to_string());
let id = format!("{}:chunk:{:04}", doc.id, idx);
tracing::trace!(chunk_id = %id, "chunk produced");
output.push(Document {
id,
content: chunk_text.clone(),
metadata,
mime_type: doc.mime_type.clone(),
});
}
}
output
}
/// 递归分割(Phase 1)。
///
/// 按 separator 优先级从高到低切割文本。每个输出片段的字符数
/// 均 ≤ chunk_size(除非最终降到 `""` 字符级兜底)。
///
/// Phase 1 只做"切分",不做合并——合并由 Phase 2 (`merge_with_overlap`) 处理。
///
/// **关键行为**:当文本中存在 separator 时,按 separator 切分。
/// 若所有 segment 均 ≤ chunk_size,直接返回所有 segments
/// 若某个 segment > chunk_size,递归降级到下一级 separator。
///
/// **早返回守卫**:如果整段文本 ≤ chunk_size(含恰好等于),直接
/// 返回 `[text.to_string()]`,避免在 Phase 2 合并时丢失 separator
/// 边界信息。
fn split_text(&self, text: &str, separators: &[String]) -> Vec<String> {
if text.is_empty() {
return Vec::new();
}
// 早返回:整段文本 ≤ chunk_size 时整体返回,避免分割后再
// 合并时丢失 separator 边界
if chars_len(text) <= self.chunk_size {
return vec![text.to_string()];
}
if separators.is_empty() {
// 防御:理论上不应到达这里(DEFAULT_SEPARATORS 末尾有 `""`
return self.split_by_chars(text);
}
let sep = &separators[0];
if sep.is_empty() {
// 字符级兜底
return self.split_by_chars(text);
}
// 检查文本中是否包含当前 separator
if !text.contains(sep.as_str()) {
// 不含此 separator,降级到下一级
return self.split_text(text, &separators[1..]);
}
// 文本中存在 separator,按 separator 切分
let raw_segments: Vec<&str> = text.split(sep.as_str()).collect();
let mut result = Vec::new();
for seg in raw_segments {
if seg.is_empty() {
continue;
}
if chars_len(seg) > self.chunk_size {
// 当前片段超长:递归降级到下一级 separator
result.extend(self.split_text(seg, &separators[1..]));
} else {
// 当前片段符合 chunk_size,直接输出
result.push(seg.to_string());
}
}
result
}
/// 字符级兜底分割(确保任何文本都能被切到 chunk_size 以内)。
///
/// 使用 `char_indices()` 步进,避免截断在多字节 UTF-8 字符中间。
fn split_by_chars(&self, text: &str) -> Vec<String> {
let mut result = Vec::new();
let mut current = String::new();
for (_, ch) in text.char_indices() {
current.push(ch);
if chars_len(&current) >= self.chunk_size {
result.push(std::mem::take(&mut current));
}
}
if !current.is_empty() {
result.push(current);
}
result
}
/// 贪心合并 + overlap 滑动窗口(Phase 2)。
///
/// 把 Phase 1 输出的 segments 合并到目标 chunk_size,并对相邻 chunk
/// 应用 chunk_overlap 字符的重叠窗口。
///
/// **已知行为**:合并时使用空字符串 `""` 连接相邻 segments
/// (即 `current.join("")`),不保留 Phase 1 切分时消耗的 separator
/// 边界信息。这意味着跨 chunk 的结构化边界(如段落、句子)会
/// 在合并点"塌缩"——但对 RAG 语义检索影响通常较小。如需保留
/// separator 边界,可重构此方法接受 separator 参数。
fn merge_with_overlap(&self, mut segments: Vec<String>) -> Vec<String> {
if segments.is_empty() {
return Vec::new();
}
if segments.len() == 1 {
return segments;
}
// Phase 2a: 贪心合并 segments 到目标 chunk_size
// segments 用 "" 连接,sep_count 不参与长度计算)
let mut chunks: Vec<String> = Vec::new();
let mut current: Vec<String> = Vec::new();
let mut current_len: usize = 0;
for seg in segments.drain(..) {
let seg_len = chars_len(&seg);
let new_total = current_len + seg_len;
if new_total > self.chunk_size && !current.is_empty() {
chunks.push(current.join(""));
current.clear();
current_len = 0;
}
current.push(seg);
current_len += seg_len;
}
if !current.is_empty() {
chunks.push(current.join(""));
}
if chunks.len() <= 1 {
return chunks;
}
// Phase 2b: 应用 overlap 滑动窗口(除第一个 chunk 外)
let overlap = self.chunk_overlap;
if overlap == 0 {
return chunks;
}
for i in 1..chunks.len() {
let prev = &chunks[i - 1];
let prev_chars_count = chars_len(prev);
if prev_chars_count == 0 {
continue;
}
let take_n = overlap.min(prev_chars_count);
// 字符级安全地取 prev 末尾 take_n 个字符
let tail: String = prev.chars().rev().take(take_n).collect::<Vec<_>>().into_iter().rev().collect();
chunks[i] = format!("{}{}", tail, chunks[i]);
}
chunks
}
}
impl Default for RecursiveCharacterSplitter {
fn default() -> Self {
Self::new(1000, 200)
}
}
/// 字符数(Unicode 标量值),等价于 `s.chars().count()`。
#[inline]
fn chars_len(s: &str) -> usize {
s.chars().count()
}
#[cfg(test)]
mod tests {
use super::*;
// ===== Group B1 — Document struct 基础测试 =====
#[test]
fn document_new_metadata_defaults_empty() {
let doc = Document::new("id-1", "content", "text/plain");
assert!(doc.metadata.is_empty());
}
#[test]
fn document_clone_partial_eq() {
let doc = Document::new("id-1", "content", "text/plain");
let cloned = doc.clone();
assert_eq!(doc, cloned);
}
#[test]
fn document_different_ids_not_equal() {
let doc1 = Document::new("id-1", "content", "text/plain");
let doc2 = Document::new("id-2", "content", "text/plain");
assert_ne!(doc1, doc2);
}
#[test]
fn document_from_raw_uses_text_plain() {
let doc = Document::from_raw("id-1", "hello");
assert_eq!(doc.mime_type, "text/plain");
assert!(doc.metadata.is_empty());
}
// ===== Group B2 — Splitter 边界条件测试 =====
#[test]
fn split_empty_doc_returns_empty() {
let splitter = RecursiveCharacterSplitter::new(100, 20);
let chunks = splitter.split(&[]);
assert!(chunks.is_empty());
}
#[test]
fn split_short_doc_single_chunk() {
let splitter = RecursiveCharacterSplitter::new(100, 20);
let doc = Document::from_raw("short", "hello");
let chunks = splitter.split(&[doc]);
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].content, "hello");
assert_eq!(chunks[0].metadata.get("chunk_index").map(|s| s.as_str()), Some("0"));
assert_eq!(chunks[0].metadata.get("chunk_count").map(|s| s.as_str()), Some("1"));
}
#[test]
fn split_empty_content_yields_no_chunks() {
let splitter = RecursiveCharacterSplitter::new(100, 20);
let doc = Document::from_raw("empty", "");
let chunks = splitter.split(&[doc]);
assert!(chunks.is_empty());
}
#[test]
#[should_panic(expected = "chunk_size must be greater than chunk_overlap")]
fn split_constructor_panics_on_invalid_overlap() {
let _ = RecursiveCharacterSplitter::new(10, 10);
}
#[test]
#[should_panic(expected = "chunk_size must be greater than 0")]
fn split_constructor_panics_on_zero_chunk_size() {
let _ = RecursiveCharacterSplitter::new(0, 0);
}
#[test]
fn try_new_returns_err_on_invalid_params() {
assert!(RecursiveCharacterSplitter::try_new(0, 0).is_err());
assert!(RecursiveCharacterSplitter::try_new(10, 10).is_err());
assert!(RecursiveCharacterSplitter::try_new(100, 20).is_ok());
}
#[test]
fn default_separators_match_spec() {
let splitter = RecursiveCharacterSplitter::default();
// Default separators should include CJK punctuation as the last meaningful
// separator before the char-level fallback. We can't directly access the
// private field, so we verify behavior: a Chinese sentence should split
// on "。" at the sentence level rather than the word level.
let doc = Document::from_raw("zh", "你好世界。今天天气好。");
let chunks = splitter.split(&[doc]);
// The default chunk_size=1000, so the whole content fits in 1 chunk.
// But the separators list contains "。" — this is verified via integration test.
assert!(!chunks.is_empty());
}
// ===== Group B3 — Splitter 核心算法测试 =====
#[test]
fn split_paragraph_boundary() {
// 小 chunk_size 强制段落级别分割
let splitter = RecursiveCharacterSplitter::new(4, 1);
let doc = Document::from_raw("p", "para1\n\npara2");
let chunks = splitter.split(&[doc]);
// para1 (5 chars) > chunk_size=4 → 递归降级到 char 级拆分
// para2 同理
// 总共应该产生多个 chunk
assert!(chunks.len() >= 2, "expected >= 2 chunks, got {}", chunks.len());
}
#[test]
fn split_recursive_deepen() {
let splitter = RecursiveCharacterSplitter::new(50, 5);
// 200 字符无 \n\n,强制降级
let text: String = "a".repeat(200);
let doc = Document::from_raw("long", &text);
let chunks = splitter.split(&[doc]);
assert!(chunks.len() >= 3, "expected >= 3 chunks, got {}", chunks.len());
for chunk in &chunks {
// chunk 内容 = overlap_tail(≤5) + new_content(≤50),故 ≤ 55
assert!(chars_len(&chunk.content) <= 55, "chunk too long: {} chars", chars_len(&chunk.content));
}
}
#[test]
fn split_greedy_merge_combines_segments() {
let splitter = RecursiveCharacterSplitter::new(20, 2);
// 一段含多个 \n\n 分隔的短小段,应被合并到 chunk_size
let doc = Document::from_raw("g", "aa\n\nbb\n\ncc\n\ndd");
let chunks = splitter.split(&[doc]);
// 短段应被合并:总共应该少于 4 个 chunk
assert!(chunks.len() <= 3, "expected <= 3 chunks after merge, got {}", chunks.len());
}
#[test]
fn split_overlap_consistency() {
let splitter = RecursiveCharacterSplitter::new(20, 5);
// 构造一个需要多 chunk 的文本
let text: String = "x".repeat(50);
let doc = Document::from_raw("o", &text);
let chunks = splitter.split(&[doc]);
assert!(chunks.len() >= 2);
// chunk[1] 应该以 chunk[0] 的最后 5 个字符作为前缀
let prev_tail: String = chunks[0]
.content
.chars()
.rev()
.take(5)
.collect::<Vec<_>>()
.into_iter()
.rev()
.collect();
assert!(
chunks[1].content.starts_with(&prev_tail),
"chunk[1] should start with last 5 chars of chunk[0]: prev_tail={:?}, chunk[1]={:?}",
prev_tail,
chunks[1].content
);
}
#[test]
fn split_character_fallback() {
let splitter = RecursiveCharacterSplitter::new(5, 0);
// 纯字母无标点,应降级到字符级
let doc = Document::from_raw("cf", "aaaaaaaaa");
let chunks = splitter.split(&[doc]);
assert_eq!(chunks.len(), 2, "expected 2 chunks, got {}", chunks.len());
for chunk in &chunks {
assert!(chars_len(&chunk.content) <= 5);
}
}
#[test]
fn split_multibyte_utf8_boundary() {
// 验证字符级单位而非字节级单位
let splitter = RecursiveCharacterSplitter::new(10, 2);
// 30 个中文字符 = 90 字节(UTF-8)
let text: String = "".repeat(30);
let doc = Document::from_raw("cjk", &text);
let chunks = splitter.split(&[doc]);
// 30 字符 / 10 chunk_size = 3 个 chunk
assert!(chunks.len() >= 3, "expected >= 3 chunks for 30 chars / chunk_size=10, got {}", chunks.len());
for chunk in &chunks {
let char_count = chars_len(&chunk.content);
// chunk = overlap_tail(≤2) + new_content(≤10),故 ≤ 12
assert!(char_count <= 12, "chunk char count {} exceeds 10+overlap", char_count);
}
}
// ===== Group B4 — Splitter 集成测试 =====
#[test]
fn split_multiple_docs() {
let splitter = RecursiveCharacterSplitter::new(50, 5);
let docs = vec![
Document::from_raw("a", "a".repeat(30).as_str()),
Document::from_raw("b", "b".repeat(30).as_str()),
Document::from_raw("c", "c".repeat(30).as_str()),
];
let chunks = splitter.split(&docs);
assert!(chunks.len() >= 3);
// 每个 chunk 的 source_id 应指向对应的输入 doc
for chunk in &chunks {
let source = chunk.metadata.get("source_id").unwrap();
assert!(["a", "b", "c"].contains(&source.as_str()));
}
}
#[test]
fn split_metadata_inheritance() {
let splitter = RecursiveCharacterSplitter::new(100, 10);
let mut doc = Document::new("m", "short content", "text/plain");
doc.metadata.insert("author".to_string(), "alice".to_string());
let chunks = splitter.split(&[doc]);
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].metadata.get("author").map(|s| s.as_str()), Some("alice"));
assert_eq!(chunks[0].metadata.get("source_id").map(|s| s.as_str()), Some("m"));
assert_eq!(chunks[0].metadata.get("chunk_index").map(|s| s.as_str()), Some("0"));
assert_eq!(chunks[0].metadata.get("chunk_count").map(|s| s.as_str()), Some("1"));
}
}
+3
View File
@@ -1,11 +1,14 @@
//! agcore —— 智能体(Agent)核心工具箱。 //! agcore —— 智能体(Agent)核心工具箱。
pub mod agent; pub mod agent;
pub mod document;
pub mod llm; pub mod llm;
pub mod memory; pub mod memory;
pub mod prompt; pub mod prompt;
pub mod tools; pub mod tools;
pub use document::Document;
use tracing_subscriber::{EnvFilter, fmt, prelude::*}; use tracing_subscriber::{EnvFilter, fmt, prelude::*};
static INIT: std::sync::Once = std::sync::Once::new(); static INIT: std::sync::Once = std::sync::Once::new();
+1
View File
@@ -3,6 +3,7 @@
pub mod compact; pub mod compact;
pub mod convert; pub mod convert;
pub mod cycle; pub mod cycle;
pub mod embedding;
pub mod error; pub mod error;
pub mod hooks; pub mod hooks;
pub mod mock; pub mod mock;
+32 -16
View File
@@ -73,10 +73,7 @@ impl CompactState {
/// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。 /// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。
pub fn estimate_message_tokens(messages: &[Message]) -> u32 { pub fn estimate_message_tokens(messages: &[Message]) -> u32 {
messages messages.iter().map(estimate_single_message_tokens).sum()
.iter()
.map(estimate_single_message_tokens)
.sum()
} }
fn estimate_single_message_tokens(msg: &Message) -> u32 { fn estimate_single_message_tokens(msg: &Message) -> u32 {
@@ -99,9 +96,7 @@ fn estimate_block_tokens(block: &ContentBlock) -> u32 {
match block { match block {
ContentBlock::Text { text } => estimate_text_tokens(text), ContentBlock::Text { text } => estimate_text_tokens(text),
ContentBlock::Thinking { text, .. } => estimate_text_tokens(text), ContentBlock::Thinking { text, .. } => estimate_text_tokens(text),
ContentBlock::ToolUse { input, .. } => { ContentBlock::ToolUse { input, .. } => estimate_text_tokens(&input.to_string()),
estimate_text_tokens(&input.to_string())
}
ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content), ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
// ponytail: Image / Audio / File / Extension 在 IR 中固定估算。 // ponytail: Image / Audio / File / Extension 在 IR 中固定估算。
// 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。 // 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。
@@ -148,14 +143,25 @@ pub fn microcompact(messages: &mut [Message], keep_recent: usize) -> u32 {
// 第一遍:计算可释放 token(仅非错误 ToolResult // 第一遍:计算可释放 token(仅非错误 ToolResult
for msg in &messages[..prune_start] { 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); freed_tokens += estimate_single_message_tokens(msg);
} }
} }
// 第二遍:替换内容(仅非错误 ToolResult // 第二遍:替换内容(仅非错误 ToolResult
for msg in &mut messages[..prune_start] { 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 { *content = vec![ContentBlock::Text {
text: "[pruned]".to_string(), text: "[pruned]".to_string(),
}]; }];
@@ -177,13 +183,15 @@ mod tests {
fn estimate_message_tokens_handles_all_variants() { fn estimate_message_tokens_handles_all_variants() {
let messages = vec![ let messages = vec![
Message::System { Message::System {
content: vec![ContentBlock::Text { content: vec![ContentBlock::Text { text: "sys".into() }],
text: "sys".into(),
}],
}, },
Message::user_text("hi"), Message::user_text("hi"),
Message::assistant("ans"), 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), Message::tool_result("call_1", "tool res", false),
]; ];
let tokens = estimate_message_tokens(&messages); let tokens = estimate_message_tokens(&messages);
@@ -205,7 +213,10 @@ mod tests {
assert!(freed > 0); assert!(freed > 0);
assert_eq!(messages.len(), before_len); // 只改内容,不删消息 assert_eq!(messages.len(), before_len); // 只改内容,不删消息
// 索引 1 是被压缩的 ToolResult // 索引 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_eq!(content.len(), 1);
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]")); assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
assert!(!is_error); assert!(!is_error);
@@ -228,9 +239,14 @@ mod tests {
assert_eq!(freed, 0); // 错误 ToolResult 不计入 assert_eq!(freed, 0); // 错误 ToolResult 不计入
assert_eq!(messages.len(), before_len); 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!(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 { } else {
panic!("expected ToolResult at index 1"); panic!("expected ToolResult at index 1");
} }
+29 -30
View File
@@ -8,11 +8,9 @@
use serde_json::Value; 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::OpenaiToolCall;
use crate::llm::types::message::{ContentBlock, Message};
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage, OpenaiContentPart};
/// `OpenaiChatMessage` → IR `Message`。 /// `OpenaiChatMessage` → IR `Message`。
/// ///
@@ -24,11 +22,10 @@ use crate::llm::types::OpenaiToolCall;
/// - `Function`(已废弃)→ `Message::ToolResult``name` 作为 `tool_call_id` 兜底) /// - `Function`(已废弃)→ `Message::ToolResult``name` 作为 `tool_call_id` 兜底)
pub fn from_openai(msg: &OpenaiChatMessage) -> Message { pub fn from_openai(msg: &OpenaiChatMessage) -> Message {
match msg { match msg {
OpenaiChatMessage::Developer { content, .. } | OpenaiChatMessage::System { content, .. } => { OpenaiChatMessage::Developer { content, .. }
Message::System { | OpenaiChatMessage::System { content, .. } => Message::System {
content: content_to_blocks(content), content: content_to_blocks(content),
} },
}
OpenaiChatMessage::User { content, .. } => Message::User { OpenaiChatMessage::User { content, .. } => Message::User {
content: content_to_blocks(content), content: content_to_blocks(content),
}, },
@@ -86,7 +83,11 @@ pub fn to_openai(msg: &Message) -> OpenaiChatMessage {
content: blocks_to_content(content), content: blocks_to_content(content),
name: None, name: None,
}, },
Message::UserImage { data, mime_type, detail } => { Message::UserImage {
data,
mime_type,
detail,
} => {
// ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。 // ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。
let mime = mime_type.clone(); let mime = mime_type.clone();
let is_url = data.starts_with("http://") || data.starts_with("https://"); 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 ContentField::Array(parts) => parts
.iter() .iter()
.filter_map(|p| match p { .filter_map(|p| match p {
OpenaiContentPart::Text { text } => { OpenaiContentPart::Text { text } => Some(ContentBlock::Text { text: text.clone() }),
Some(ContentBlock::Text { text: text.clone() }) OpenaiContentPart::Refusal { refusal } => Some(ContentBlock::Text {
} text: refusal.clone(),
OpenaiContentPart::Refusal { refusal } => { }),
Some(ContentBlock::Text { text: refusal.clone() })
}
OpenaiContentPart::Image { image_url, .. } => { OpenaiContentPart::Image { image_url, .. } => {
// ponytail: 简化处理 —— URL 直接通过,data URI 拆出 // ponytail: 简化处理 —— URL 直接通过,data URI 拆出
// data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。 // data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。
let url = &image_url.url; let url = &image_url.url;
if let Some(rest) = url.strip_prefix("data:") if let Some(rest) = url.strip_prefix("data:")
&& let Some((mime, b64)) = rest.split_once(";base64,") { && let Some((mime, b64)) = rest.split_once(";base64,")
return Some(ContentBlock::Image { {
source: crate::llm::types::message::ImageSource { return Some(ContentBlock::Image {
data: b64.to_string(), source: crate::llm::types::message::ImageSource {
mime_type: mime.to_string(), data: b64.to_string(),
is_url: false, mime_type: mime.to_string(),
}, is_url: false,
}); },
} });
}
Some(ContentBlock::Image { Some(ContentBlock::Image {
source: crate::llm::types::message::ImageSource { source: crate::llm::types::message::ImageSource {
data: url.clone(), data: url.clone(),
@@ -263,7 +263,9 @@ mod tests {
match ir { match ir {
Message::System { content } => { Message::System { content } => {
assert_eq!(content.len(), 1); 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"), _ => panic!("expected System variant"),
} }
@@ -385,10 +387,7 @@ mod tests {
assert_eq!(parts.len(), 1); assert_eq!(parts.len(), 1);
match &parts[0] { match &parts[0] {
OpenaiContentPart::Image { image_url, .. } => { OpenaiContentPart::Image { image_url, .. } => {
assert_eq!( assert_eq!(image_url.url, "data:image/png;base64,BASE64DATA");
image_url.url,
"data:image/png;base64,BASE64DATA"
);
} }
_ => panic!("expected Image part"), _ => panic!("expected Image part"),
} }
+842 -44
View File
File diff suppressed because it is too large Load Diff
+183
View File
@@ -0,0 +1,183 @@
//! Embedding 抽象 —— 文本向量化接口。
//!
//! 提供 [`Embedding`] trait 和零依赖的 [`MockEmbedding`] 引用实现。
//! 上层可实现此 trait 以对接真实 Embedding ProviderOpenAI、Cohere 等)。
//!
//! 所有实现使用 [`LlmError`] 作为统一错误类型,与 llm 模块保持一致。
use async_trait::async_trait;
use crate::llm::error::LlmError;
/// 文本向量化抽象接口。
///
/// 将文本字符串转换为固定维度的浮点向量,用于语义相似度计算。
/// 设计为异步以支持网络 IO(如 OpenAI Embedding API)。
///
/// 使用 [`LlmError`] 作为统一错误类型,与 llm 模块保持一致。
///
/// # 实现要求
///
/// - `embed()` 返回的向量外层的 Vec 长度必须等于输入切片长度(一对一映射)
/// - 内层 Vec 长度必须等于 `dim()` 返回值
/// - 调用方应保证输入非空(空切片返回空外层 Vec,不报错)
///
/// # 稳定性
///
/// 实验性 APIv0.3.x),方法签名可能在 v0.4 中调整。
#[async_trait]
pub trait Embedding: Send + Sync {
/// 批量向量化。
///
/// 返回 `Vec<Vec<f32>>`,第 i 个内层向量对应 `input[i]`。
async fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>, LlmError>;
/// 返回向量维度。
fn dim(&self) -> usize;
}
/// 确定性 Mock Embedding —— 零依赖伪随机单位向量。
///
/// 使用 sin 哈希将输入字符串映射到单位球面上的一个点:
/// 1. 对输入字符串计算简单哈希(字符字节和 + 长度)作为种子
/// 2. 用 `f32::sin(seed + i) * 10000` 生成第 i 个维度的值
/// 3. 归一化到单位长度(L2 norm = 1.0
///
/// 特性:
/// - **确定性**:相同输入 → 相同向量
/// - **有区分度**:不同输入产生不同向量(高概率)
/// - **单位范数**:余弦相似度等价于点积
/// - **开销极低**:不分配额外内存,无 IO
///
/// # 已知限制
///
/// `f32::sin(seed + i) * 10000` 在维度较高时(如 1536OpenAI Embedding 维度)
/// 可能出现周期性模式——相邻维度取值在 `sin` 周期 2π 约束下呈规律性重复。
/// MockEmbedding 仅用于测试验证,**不应用于生产级相似度排序**;
/// 做严肃验证时建议使用真实 Embedding Provider 或显式随机初始化。
pub struct MockEmbedding {
dim: usize,
}
impl MockEmbedding {
/// 创建 Mock Embedding,输出向量维度为 `dim`。
pub fn new(dim: usize) -> Self {
Self { dim }
}
}
#[async_trait]
impl Embedding for MockEmbedding {
async fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>, LlmError> {
let results: Vec<Vec<f32>> = input
.iter()
.map(|text| {
// 简单哈希:字符字节值和 + 文本长度作为种子
let seed: f64 = text.bytes().map(|b| b as f64).sum::<f64>() + text.len() as f64;
let mut vec: Vec<f32> = (0..self.dim)
.map(|i| f32::sin(seed as f32 + i as f32) * 10000.0)
.collect();
l2_normalize(&mut vec);
vec
})
.collect();
Ok(results)
}
fn dim(&self) -> usize {
self.dim
}
}
/// L2 归一化(in-place)。
///
/// 零向量(norm == 0)保持全零 —— 防除零保护。
fn l2_normalize(vec: &mut [f32]) {
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > f32::EPSILON {
for x in vec.iter_mut() {
*x /= norm;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
/// 计算向量的 L2 范数。
fn l2_norm(v: &[f32]) -> f32 {
v.iter().map(|x| x * x).sum::<f32>().sqrt()
}
#[tokio::test]
async fn embed_correct_dim() {
let embedder = MockEmbedding::new(8);
let inputs = vec!["hello".to_string(), "world".to_string()];
let result = embedder.embed(&inputs).await.unwrap();
assert_eq!(result.len(), 2);
for vec in &result {
assert_eq!(vec.len(), 8);
}
}
#[tokio::test]
async fn embed_batch_size_match() {
let embedder = MockEmbedding::new(4);
let inputs = vec![
"a".to_string(),
"b".to_string(),
"c".to_string(),
"d".to_string(),
"e".to_string(),
];
let result = embedder.embed(&inputs).await.unwrap();
assert_eq!(result.len(), inputs.len());
}
#[tokio::test]
async fn embed_deterministic() {
let embedder = MockEmbedding::new(4);
let inputs = vec!["deterministic test".to_string()];
let r1 = embedder.embed(&inputs).await.unwrap();
let r2 = embedder.embed(&inputs).await.unwrap();
assert_eq!(r1, r2);
}
#[tokio::test]
async fn embed_unit_vector_norm() {
let embedder = MockEmbedding::new(16);
let inputs = vec!["any text".to_string(), "another".to_string()];
let result = embedder.embed(&inputs).await.unwrap();
for vec in &result {
let norm = l2_norm(vec);
assert!((norm - 1.0).abs() < 1e-5, "vector norm should be ~1.0, got {}", norm);
}
}
#[tokio::test]
async fn embed_different_inputs_different_vectors() {
let embedder = MockEmbedding::new(16);
let r1 = embedder
.embed(&["hello world".to_string()])
.await
.unwrap();
let r2 = embedder
.embed(&["completely different".to_string()])
.await
.unwrap();
assert_ne!(r1, r2);
}
#[tokio::test]
async fn embed_empty_string() {
// 空字符串输入应不 panic,且向量范数仍≈1.0(防除零路径)
let embedder = MockEmbedding::new(4);
let inputs = vec!["".to_string()];
let result = embedder.embed(&inputs).await.unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].len(), 4);
let norm = l2_norm(&result[0]);
assert!((norm - 1.0).abs() < 1e-5, "empty-string vector norm should be ~1.0, got {}", norm);
}
}
+10 -3
View File
@@ -8,9 +8,12 @@ use std::time::Duration;
/// ///
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查 API key、减少上下文)。 /// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查 API key、减少上下文)。
#[derive(thiserror::Error, Debug)] #[derive(thiserror::Error, Debug)]
#[non_exhaustive]
pub enum LlmError { pub enum LlmError {
/// API 认证失败(API key 无效、过期或权限不足)。 /// 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), Authentication(String),
/// 请求被限流,可选地附带重试等待时间。可重试。 /// 请求被限流,可选地附带重试等待时间。可重试。
@@ -18,7 +21,9 @@ pub enum LlmError {
RateLimit { retry_after: Option<Duration> }, RateLimit { retry_after: Option<Duration> },
/// HTTP 请求失败(网络错误或非 2xx 状态码),包含状态码与响应体。 /// HTTP 请求失败(网络错误或非 2xx 状态码),包含状态码与响应体。
#[error("LLM 请求失败(HTTP {status}: {body}。请检查 Provider 端点地址(base_url)和网络连通性")] #[error(
"LLM 请求失败(HTTP {status}: {body}。请检查 Provider 端点地址(base_url)和网络连通性"
)]
Request { status: u16, body: String }, Request { status: u16, body: String },
/// 请求超时。可重试。 /// 请求超时。可重试。
@@ -30,7 +35,9 @@ pub enum LlmError {
Stream(String), 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 }, ContextLength { actual: u32, limit: u32 },
/// 其他未分类的 LLM 调用失败。 /// 其他未分类的 LLM 调用失败。
+2 -3
View File
@@ -7,6 +7,7 @@ use crate::llm::types::request_v2::MessageRequest;
/// 生命周期钩子事件点。 /// 生命周期钩子事件点。
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum HookEvent { pub enum HookEvent {
/// LLM 请求发起之前(可阻断)。 /// LLM 请求发起之前(可阻断)。
PreRequest, PreRequest,
@@ -130,9 +131,7 @@ impl Default for HookExecutor {
impl HookExecutor { impl HookExecutor {
/// 创建一个空的执行器。 /// 创建一个空的执行器。
pub fn new() -> Self { pub fn new() -> Self {
Self { Self { hooks: Vec::new() }
hooks: Vec::new(),
}
} }
/// 注册一个钩子到指定事件点。 /// 注册一个钩子到指定事件点。
+10 -8
View File
@@ -97,8 +97,7 @@ impl LlmProvider for MockProvider {
async fn chat_stream( async fn chat_stream(
&self, &self,
_request: MessageRequest, _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()?; let response = self.pop()?;
// 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。 // 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。
let id = response.id.clone(); let id = response.id.clone();
@@ -206,10 +205,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn chat_returns_queued_response() { async fn chat_returns_queued_response() {
let provider = MockProvider::new(vec![text_response("hello")]); let provider = MockProvider::new(vec![text_response("hello")]);
let resp = provider let resp = provider.chat(MessageRequest::default()).await.unwrap();
.chat(MessageRequest::default())
.await
.unwrap();
assert_eq!(resp.text(), "hello"); assert_eq!(resp.text(), "hello");
assert_eq!(provider.remaining(), 0); assert_eq!(provider.remaining(), 0);
} }
@@ -231,7 +227,10 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn chat_stream_emits_text_delta_sequence() { async fn chat_stream_emits_text_delta_sequence() {
let provider = MockProvider::new(vec![text_response("hi")]); 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_start = false;
let mut seen_block_start = false; let mut seen_block_start = false;
@@ -283,7 +282,10 @@ mod tests {
extra: Default::default(), extra: Default::default(),
}; };
let provider = MockProvider::new(vec![response]); 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_args = false;
let mut saw_tool_end = false; let mut saw_tool_end = false;
+456 -21
View File
@@ -1,11 +1,14 @@
pub mod anthropic; pub mod anthropic;
pub mod ollama;
pub mod openai; pub mod openai;
pub mod openai_compat; pub mod openai_compat;
pub mod registry; pub mod registry;
use std::pin::Pin; use std::pin::Pin;
use std::time::Duration;
use futures_core::Stream; use futures_core::Stream;
use reqwest::Client;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use crate::llm::error::LlmError; use crate::llm::error::LlmError;
@@ -18,6 +21,7 @@ use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
/// 当前协议数量(5 种以内)完全可控,enum 的编译期安全检查优于运行时的 `HashMap::get()`。 /// 当前协议数量(5 种以内)完全可控,enum 的编译期安全检查优于运行时的 `HashMap::get()`。
/// 未来如果扩展到 15+ 种以上,再改为注册表模式。 /// 未来如果扩展到 15+ 种以上,再改为注册表模式。
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ProviderType { pub enum ProviderType {
/// OpenAI Chat Completions API(兼容 DeepSeek / Qwen 等 `/chat/completions` 端点)。 /// OpenAI Chat Completions API(兼容 DeepSeek / Qwen 等 `/chat/completions` 端点)。
OpenaiChat, OpenaiChat,
@@ -29,6 +33,8 @@ pub enum ProviderType {
DeepSeek, DeepSeek,
/// Qwen / 阿里云百炼(OpenAI-compatible `/chat/completions`)。 /// Qwen / 阿里云百炼(OpenAI-compatible `/chat/completions`)。
Qwen, Qwen,
/// Ollama 本地推理(OpenAI-compatible `/chat/completions`,默认 `http://localhost:11434/v1`)。
Ollama,
} }
impl std::str::FromStr for ProviderType { impl std::str::FromStr for ProviderType {
@@ -41,47 +47,211 @@ impl std::str::FromStr for ProviderType {
"anthropic" | "claude" => Ok(ProviderType::Anthropic), "anthropic" | "claude" => Ok(ProviderType::Anthropic),
"deepseek" => Ok(ProviderType::DeepSeek), "deepseek" => Ok(ProviderType::DeepSeek),
"qwen" | "dashscope" | "tongyi" => Ok(ProviderType::Qwen), "qwen" | "dashscope" | "tongyi" => Ok(ProviderType::Qwen),
"ollama" => Ok(ProviderType::Ollama),
_ => Err(format!("未知的 Provider 类型: {s}")), _ => Err(format!("未知的 Provider 类型: {s}")),
} }
} }
} }
/// Provider 构造参数 —— 通用 base_url + api_key + model。 /// Provider 构造参数 —— 通用 base_url + api_key + model + timeout / retry 配置
#[derive(Debug, Clone)]
pub struct ProviderConfig { pub struct ProviderConfig {
/// API base URL(如 `https://api.openai.com/v1`)。为空时由 Provider 选择默认值。
pub base_url: String, pub base_url: String,
/// API key。Ollama 等本地 Provider 可为空。
pub api_key: String, pub api_key: String,
/// 模型名(如 `gpt-4o` / `claude-sonnet-4-20250514`)。
pub model: String, 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 被注册。 /// Provider 工厂 —— exhaustive match 在编译期保证新 Provider 被注册。
///
/// `config.timeout_secs` 注入到 Provider 的 HTTP Client 超时配置。
/// 每个分支通过 `from_parts` (pub(crate)) 一次性构造,无冗余 client 创建。
pub fn create_provider( pub fn create_provider(
provider_type: ProviderType, provider_type: ProviderType,
config: ProviderConfig, config: ProviderConfig,
) -> Result<Box<dyn LlmProvider>, LlmError> { ) -> Result<Box<dyn LlmProvider>, LlmError> {
match provider_type { match provider_type {
ProviderType::OpenaiChat => Ok(Box::new(openai::OpenaiChatProvider::new( ProviderType::OpenaiChat => {
config.base_url, let client = build_client_with_timeout(config.timeout_secs)?;
config.api_key, Ok(Box::new(openai::OpenaiChatProvider(
config.model, openai::GenericOpenaiProvider::from_parts(
))), config.base_url,
config.api_key,
config.model,
"openai",
client,
Vec::new(),
config.timeout_secs,
),
)))
}
ProviderType::OpenaiResponse => Err(LlmError::Other( ProviderType::OpenaiResponse => Err(LlmError::Other(
"OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(), "OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(),
)), )),
ProviderType::Anthropic => Ok(Box::new(anthropic::AnthropicProvider::new( ProviderType::Anthropic => {
config.base_url, let client = build_anthropic_client(&config.api_key, config.timeout_secs)?;
config.api_key, Ok(Box::new(anthropic::AnthropicProvider::from_parts(
config.model, config.base_url,
))), config.api_key,
ProviderType::DeepSeek => Ok(Box::new(openai_compat::DeepSeekProvider::new( config.model,
config.base_url, client,
config.api_key, config.timeout_secs,
config.model, )))
))), }
ProviderType::Qwen => Ok(Box::new(openai_compat::QwenProvider::new( ProviderType::DeepSeek => {
config.base_url, let client = build_client_with_timeout(config.timeout_secs)?;
config.api_key, Ok(Box::new(openai_compat::DeepSeekProvider(
config.model, 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 静态能力描述。 /// 返回 Provider 静态能力描述。
fn capabilities(&self) -> ProviderCapabilities; 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:?}"),
}
}
}
+263 -56
View File
@@ -12,10 +12,10 @@ use async_trait::async_trait;
use bytes::Bytes; use bytes::Bytes;
use futures_core::Stream; use futures_core::Stream;
use futures_util::StreamExt; use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderValue};
use reqwest::Client; use reqwest::Client;
use reqwest::header::{HeaderMap, HeaderValue};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::{json, Value}; use serde_json::{Value, json};
use tracing::{debug, error, info, warn}; use tracing::{debug, error, info, warn};
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures}; use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
@@ -39,16 +39,20 @@ pub struct AnthropicProvider {
#[allow(dead_code)] #[allow(dead_code)]
api_key: String, api_key: String,
model: String, model: String,
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter
/// 因此单独存储以便错误消息与配置保持一致。
timeout_secs: u64,
} }
impl AnthropicProvider { impl AnthropicProvider {
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 key_header = HeaderValue::from_str(&api_key) let key_header =
.expect("Anthropic API key 包含无效的 HTTP 头部字符"); HeaderValue::from_str(&api_key).expect("Anthropic API key 包含无效的 HTTP 头部字符");
let version_header = HeaderValue::from_static("2023-06-01"); let version_header = HeaderValue::from_static("2023-06-01");
let http_client = Client::builder() let http_client = Client::builder()
.timeout(Duration::from_secs(120)) .timeout(Duration::from_secs(timeout_secs))
.default_headers({ .default_headers({
let mut headers = HeaderMap::new(); let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header); headers.insert("x-api-key", key_header);
@@ -67,14 +71,80 @@ impl AnthropicProvider {
}, },
api_key, api_key,
model, 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 { pub fn with_client(mut self, client: Client) -> Self {
self.http_client = client; self.http_client = client;
self 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 { fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS) request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
} }
@@ -104,12 +174,14 @@ impl AnthropicProvider {
Message::User { content } => { Message::User { content } => {
api_messages.push(AnthropicMessage::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}} // Anthropic image format: {type: "image", source: {type: "base64", media_type, data}}
let source = if data.starts_with("http://") || data.starts_with("https://") { let source = if data.starts_with("http://") || data.starts_with("https://") {
AnthropicImageSource::Url { AnthropicImageSource::Url { url: data.clone() }
url: data.clone(),
}
} else { } else {
AnthropicImageSource::Base64 { AnthropicImageSource::Base64 {
media_type: mime_type.clone(), media_type: mime_type.clone(),
@@ -194,7 +266,7 @@ impl AnthropicProvider {
.json(&body) .json(&body)
.send() .send()
.await .await
.map_err(Self::map_reqwest_error)?; .map_err(|e| self.map_reqwest_error(e))?;
let status = response.status(); let status = response.status();
if !status.is_success() { if !status.is_success() {
@@ -215,8 +287,7 @@ impl AnthropicProvider {
async fn chat_stream_inner( async fn chat_stream_inner(
&self, &self,
request: MessageRequest, 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)?; let mut body = self.build_request_body(request)?;
body.stream = Some(true); body.stream = Some(true);
@@ -230,16 +301,16 @@ impl AnthropicProvider {
.json(&body) .json(&body)
.send() .send()
.await .await
.map_err(Self::map_reqwest_error)?; .map_err(|e| self.map_reqwest_error(e))?;
let status = response.status(); let status = response.status();
if !status.is_success() { if !status.is_success() {
return Err(Self::handle_error_response(response).await); return Err(Self::handle_error_response(response).await);
} }
let byte_stream = response.bytes_stream().map(|r| { let byte_stream = response
r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}"))) .bytes_stream()
}); .map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}"))));
let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> = let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> =
Box::pin(byte_stream); Box::pin(byte_stream);
@@ -247,10 +318,10 @@ impl AnthropicProvider {
Ok(Box::pin(AnthropicSseStream::new(byte_stream))) 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() { if e.is_timeout() {
LlmError::Timeout { LlmError::Timeout {
duration: Duration::from_secs(120), duration: Duration::from_secs(self.timeout_secs),
} }
} else if e.is_connect() { } else if e.is_connect() {
LlmError::Other(format!("连接失败: {e}")) LlmError::Other(format!("连接失败: {e}"))
@@ -291,13 +362,12 @@ impl AnthropicProvider {
blocks.push(ContentBlock::Text { text }); blocks.push(ContentBlock::Text { text });
} }
AnthropicContentBlockResp::ToolUse { id, name, input } => { AnthropicContentBlockResp::ToolUse { id, name, input } => {
blocks.push(ContentBlock::ToolUse { blocks.push(ContentBlock::ToolUse { id, name, input });
id,
name,
input,
});
} }
AnthropicContentBlockResp::Thinking { thinking, signature } => { AnthropicContentBlockResp::Thinking {
thinking,
signature,
} => {
blocks.push(ContentBlock::Thinking { blocks.push(ContentBlock::Thinking {
text: thinking, text: thinking,
signature, signature,
@@ -336,8 +406,7 @@ impl LlmProvider for AnthropicProvider {
async fn chat_stream( async fn chat_stream(
&self, &self,
request: MessageRequest, 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 self.chat_stream_inner(request).await
} }
@@ -418,7 +487,9 @@ impl AnthropicMessage {
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")] #[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentPart { enum AnthropicContentPart {
Text { text: String }, Text {
text: String,
},
Image { Image {
source: AnthropicImageSource, source: AnthropicImageSource,
}, },
@@ -437,13 +508,8 @@ enum AnthropicContentPart {
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")] #[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicImageSource { enum AnthropicImageSource {
Base64 { Base64 { media_type: String, data: String },
media_type: String, Url { url: String },
data: String,
},
Url {
url: String,
},
} }
fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> { fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> {
@@ -523,9 +589,18 @@ struct AnthropicUsage {
#[derive(Debug, Deserialize)] #[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")] #[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentBlockResp { enum AnthropicContentBlockResp {
Text { text: String }, Text {
ToolUse { id: String, name: String, input: Value }, text: String,
Thinking { thinking: String, signature: Option<String> }, },
ToolUse {
id: String,
name: String,
input: Value,
},
Thinking {
thinking: String,
signature: Option<String>,
},
} }
// ============================================================================= // =============================================================================
@@ -662,18 +737,13 @@ impl AnthropicSseStream {
// 先把所有字段提前,避免 match 中 part-move // 先把所有字段提前,避免 match 中 part-move
let block_type = match &content_block { let block_type = match &content_block {
AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text, AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text,
AnthropicContentBlockStart::ToolUse { id, name } => { AnthropicContentBlockStart::ToolUse { id, name } => ContentBlockType::ToolUse {
ContentBlockType::ToolUse { id: id.clone(),
id: id.clone(), name: name.clone(),
name: name.clone(), },
}
}
AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking, AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking,
}; };
events.push(StreamEvent::ContentBlockStart { events.push(StreamEvent::ContentBlockStart { index, block_type });
index,
block_type,
});
let builder = match content_block { let builder = match content_block {
AnthropicContentBlockStart::Text { text } => { AnthropicContentBlockStart::Text { text } => {
crate::llm::types::response_v2::ContentBlockBuilder::Text(text) crate::llm::types::response_v2::ContentBlockBuilder::Text(text)
@@ -742,7 +812,9 @@ impl AnthropicSseStream {
completion_tokens_details: None, completion_tokens_details: None,
prompt_tokens_details: None, prompt_tokens_details: None,
}; };
events.push(StreamEvent::CostUpdate { usage: partial_usage }); events.push(StreamEvent::CostUpdate {
usage: partial_usage,
});
} }
} }
AnthropicSseEvent::MessageStop => { AnthropicSseEvent::MessageStop => {
@@ -753,7 +825,9 @@ impl AnthropicSseStream {
self.saw_terminal = true; self.saw_terminal = true;
match self.partial.clone().finalize() { match self.partial.clone().finalize() {
Ok(full) => { Ok(full) => {
events.push(StreamEvent::MessageComplete { full_response: full }); events.push(StreamEvent::MessageComplete {
full_response: full,
});
} }
Err(e) => { Err(e) => {
events.push(StreamEvent::Error { events.push(StreamEvent::Error {
@@ -780,10 +854,7 @@ fn _unused_marker() {}
impl Stream for AnthropicSseStream { impl Stream for AnthropicSseStream {
type Item = Result<StreamEvent, LlmError>; type Item = Result<StreamEvent, LlmError>;
fn poll_next( fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
loop { loop {
if let Some(data) = self.next_event_line() { if let Some(data) = self.next_event_line() {
let mut events = self.handle_event_json(&data); let mut events = self.handle_event_json(&data);
@@ -830,12 +901,17 @@ mod tests {
use super::*; use super::*;
use crate::llm::types::request_v2::MessageRequest; use crate::llm::types::request_v2::MessageRequest;
use serde_json::json; use serde_json::json;
use wiremock::matchers::{method, path}; use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate}; use wiremock::{Mock, MockServer, ResponseTemplate};
fn make_provider(base_url: String) -> AnthropicProvider { fn make_provider(base_url: String) -> AnthropicProvider {
// 跳过默认 header 注入:测试用自定义 base_url 直接 mock // 跳过默认 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] #[tokio::test]
@@ -979,6 +1055,7 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
"http://x".into(), "http://x".into(),
"k".into(), "k".into(),
"claude-sonnet-4-20250514".into(), "claude-sonnet-4-20250514".into(),
30,
) )
.capabilities(); .capabilities();
assert_eq!(caps.provider_name, "anthropic"); assert_eq!(caps.provider_name, "anthropic");
@@ -1000,4 +1077,134 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
assert_eq!(body.max_tokens, DEFAULT_MAX_TOKENS); assert_eq!(body.max_tokens, DEFAULT_MAX_TOKENS);
assert_eq!(body.model, "claude-sonnet-4-20250514"); assert_eq!(body.model, "claude-sonnet-4-20250514");
} }
// ===== Phase 11 Step 11.2 wiremock roundtrip 测试 =====
#[tokio::test]
async fn anthropic_401_structured_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(401).set_body_json(json!({
"type": "error",
"error": {
"type": "authentication_error",
"message": "Invalid API key provided: sk-ant-test"
}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Authentication(msg) => assert!(msg.contains("Invalid API key")),
other => panic!("expected Authentication, got {other:?}"),
}
}
#[tokio::test]
async fn anthropic_tool_use_response() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "msg_tool",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [
{"type": "text", "text": "Let me check."},
{"type": "tool_use", "id": "toolu_abc", "name": "lookup", "input": {"q": "rust"}}
],
"stop_reason": "tool_use",
"usage": {"input_tokens": 8, "output_tokens": 12}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let response = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Look up rust")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.stop_reason, StopReason::ToolUse);
let tool_use = match &response.message {
Message::Assistant { content } => content.iter().find_map(|b| match b {
ContentBlock::ToolUse { id, name, .. } => Some((id.clone(), name.clone())),
_ => None,
}),
_ => None,
};
assert_eq!(tool_use, Some(("toolu_abc".into(), "lookup".into())));
}
#[tokio::test]
async fn anthropic_version_header() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.and(header("anthropic-version", "2023-06-01"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"id": "msg_v",
"type": "message",
"model": "claude-sonnet-4-20250514",
"content": [{"type": "text", "text": "OK"}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 1, "output_tokens": 1}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let response = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap();
assert_eq!(response.text(), "OK");
}
#[tokio::test]
async fn anthropic_529_overloaded_structured() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(ResponseTemplate::new(529).set_body_json(json!({
"type": "error",
"error": {
"type": "overloaded_error",
"message": "Overloaded: Anthropic API is temporarily overloaded"
}
})))
.mount(&server)
.await;
let provider = make_provider(server.uri());
let err = provider
.chat_blocking(MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![Message::user_text("Hi")],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::RateLimit { retry_after } => assert!(retry_after.is_none()),
other => panic!("expected RateLimit, got {other:?}"),
}
}
} }
+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
}
}
File diff suppressed because it is too large Load Diff
+32 -39
View File
@@ -15,12 +15,12 @@ use std::pin::Pin;
use async_trait::async_trait; use async_trait::async_trait;
use futures_core::Stream; use futures_core::Stream;
use super::openai::GenericOpenaiProvider;
use super::ProviderCapabilities; use super::ProviderCapabilities;
use super::openai::GenericOpenaiProvider;
use crate::llm::error::LlmError; use crate::llm::error::LlmError;
use crate::llm::provider::LlmProvider;
use crate::llm::types::request_v2::MessageRequest; use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StreamEvent}; use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
use crate::llm::provider::LlmProvider;
// ============================================================================= // =============================================================================
// DeepSeek // DeepSeek
@@ -29,7 +29,7 @@ use crate::llm::provider::LlmProvider;
pub struct DeepSeekProvider(pub GenericOpenaiProvider); pub struct DeepSeekProvider(pub GenericOpenaiProvider);
impl DeepSeekProvider { 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() { let url = if base_url.is_empty() {
"https://api.deepseek.com".to_string() "https://api.deepseek.com".to_string()
} else { } else {
@@ -40,26 +40,27 @@ impl DeepSeekProvider {
api_key, api_key,
model, model,
"deepseek", "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)使用的构造器。 /// 测试中(带 mock_client)使用的构造器。
///
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
/// 无超时),不影响行为。
pub fn new_with_client( pub fn new_with_client(
base_url: String, base_url: String,
api_key: String, api_key: String,
model: String, model: String,
client: reqwest::Client, client: reqwest::Client,
) -> Self { ) -> Self {
let url = if base_url.is_empty() { Self::new(base_url, api_key, model, 30).with_client(client)
"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)
} }
} }
@@ -72,8 +73,7 @@ impl LlmProvider for DeepSeekProvider {
async fn chat_stream( async fn chat_stream(
&self, &self,
request: MessageRequest, 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 self.0.chat_stream(request).await
} }
@@ -91,7 +91,7 @@ impl LlmProvider for DeepSeekProvider {
pub struct QwenProvider(pub GenericOpenaiProvider); pub struct QwenProvider(pub GenericOpenaiProvider);
impl QwenProvider { 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() { let url = if base_url.is_empty() {
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string() "https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
} else { } else {
@@ -104,31 +104,28 @@ impl QwenProvider {
model, model,
"qwen", "qwen",
vec![("X-DashScope-SSE".to_string(), "enable".to_string())], vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
timeout_secs,
); );
Self(inner) 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( pub fn new_with_client(
base_url: String, base_url: String,
api_key: String, api_key: String,
model: String, model: String,
client: reqwest::Client, client: reqwest::Client,
) -> Self { ) -> Self {
let url = if base_url.is_empty() { Self::new(base_url, api_key, model, 30).with_client(client)
"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)
} }
} }
@@ -141,8 +138,7 @@ impl LlmProvider for QwenProvider {
async fn chat_stream( async fn chat_stream(
&self, &self,
request: MessageRequest, 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 self.0.chat_stream(request).await
} }
@@ -156,8 +152,8 @@ impl LlmProvider for QwenProvider {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::message::Message as IrMessage; use crate::llm::types::message::Message as IrMessage;
use crate::llm::types::request_v2::MessageRequest;
use serde_json::json; use serde_json::json;
use wiremock::matchers::{method, path}; use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate}; use wiremock::{Mock, MockServer, ResponseTemplate};
@@ -182,11 +178,8 @@ mod tests {
.mount(&server) .mount(&server)
.await; .await;
let provider = DeepSeekProvider::new( let provider =
server.uri(), DeepSeekProvider::new(server.uri(), "sk-test".into(), "deepseek-chat".into(), 30);
"sk-test".into(),
"deepseek-chat".into(),
);
let response = provider let response = provider
.chat(MessageRequest { .chat(MessageRequest {
model: "deepseek-chat".into(), model: "deepseek-chat".into(),
@@ -219,7 +212,7 @@ mod tests {
.mount(&server) .mount(&server)
.await; .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 let response = provider
.chat(MessageRequest { .chat(MessageRequest {
model: "qwen-plus".into(), model: "qwen-plus".into(),
+2 -4
View File
@@ -3,7 +3,7 @@
use std::collections::HashMap; use std::collections::HashMap;
use crate::llm::error::LlmError; 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 实例。 /// Provider 注册表 —— 管理多个 LLM Provider 实例。
/// ///
@@ -61,8 +61,6 @@ impl ProviderRegistry {
/// 获取默认 Provider。 /// 获取默认 Provider。
pub fn get_default(&self) -> Option<&dyn LlmProvider> { pub fn get_default(&self) -> Option<&dyn LlmProvider> {
self.default_name self.default_name.as_ref().and_then(|name| self.get(name))
.as_ref()
.and_then(|name| self.get(name))
} }
} }
+5 -200
View File
@@ -1,204 +1,9 @@
//! 流式事件系统 —— 将 LLM 流式响应解析为语义化事件 //! 流式事件系统 —— 重导出 `StreamEvent` 供向后兼容
//! //!
//! Phase 0 修订(参见 `docs/10a-phase0-types-and-trait.md` §"StreamEvent 命名冲突处理"): //! 历史说明(Phase 0 → Phase 13):
//! - 对外暴露的 `StreamEvent` 是高精度 IR 版本(来自 `response_v2::StreamEvent`)。 //! - 对外暴露的 `StreamEvent` 是高精度 IR 版本(来自 `response_v2::StreamEvent`)。
//! - 旧变体(`AssistantTextDelta` / `ToolExecutionStarted` 等)重命名为 `LegacyStreamEvent` //! - 旧版 chunk 解析 + LegacyStreamEvent 适配层在 Phase 13 完成后已整体删除。
//! 放在 `crate::llm::types::old_stream` 模块,本文件内部消费。 //! - 当前文件仅保留 `pub use` 重导出,保持与既有
//! - Phase 1 重写 Provider 时可直接消费新事件流后整体删除 `LegacyStreamEvent` 相关代码 //! `use crate::llm::stream::StreamEvent` 的代码兼容
//!
//! 当前实现:旧的 `parse_chunk_stream` 内部消费 `OpenaiChatChunk`,映射为
//! `LegacyStreamEvent`,再在 `LegacyToIrEventStream` 中映射为新 IR `StreamEvent`
//! 后输出。Phase 1 会重写此层(OpenAI Provider 直接产出新事件流)。
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 serde_json::Value;
use crate::llm::error::LlmError;
use crate::llm::types::old_stream::LegacyStreamEvent;
use crate::llm::types::response_v2::MessageResponse;
use crate::llm::types::response_v2::StopReason;
use crate::llm::types::usage::Usage;
use crate::llm::types::{OpenaiChatChunk, OpenaiToolCall};
// 唯一的对外 `StreamEvent` 定义(高精度 IR 事件,来自 `response_v2`)。
//
// 此 `pub use` 同时起到两个作用:
// 1. 让 `crate::llm::stream::StreamEvent` 路径仍指向新高精度 IR 事件,
// 保持与既有 `use crate::llm::stream::StreamEvent` 的代码兼容;
// 2. 把模块内部的 `StreamEvent` 名字指向 `response_v2::StreamEvent`。
pub use crate::llm::types::response_v2::StreamEvent; pub use crate::llm::types::response_v2::StreamEvent;
/// 将原始 OpenaiChatChunk 流解析为新高精度 IR StreamEvent 流。
///
/// ponytail: 每个产出事件都用 `Result<_, LlmError>` 包装,让上层 `chat_stream`
/// trait 方法直接消费并保持错误传播链。当前 `LegacyToIrEventStream` 内部
/// 不会产生错误,所有结果都是 `Ok`;后续 Phase 1 重写 Provider 时,
/// 真实 IR 流转换可在此层注入 error 事件。
pub fn parse_chunk_stream(
chunks: Pin<Box<dyn futures_core::Stream<Item = Result<OpenaiChatChunk, LlmError>> + Send>>,
) -> Pin<Box<dyn futures_core::Stream<Item = Result<StreamEvent, LlmError>> + Send>> {
let legacy = parse_chunk_stream_legacy(chunks);
Box::pin(LegacyToIrEventStream { inner: legacy })
}
// --- 内部:chunk → LegacyStreamEvent ---
fn parse_chunk_stream_legacy(
chunks: Pin<Box<dyn futures_core::Stream<Item = Result<OpenaiChatChunk, LlmError>> + Send>>,
) -> Pin<Box<dyn futures_core::Stream<Item = LegacyStreamEvent> + Send>> {
Box::pin(ChunkToLegacyEventStream { chunks })
}
struct ChunkToLegacyEventStream {
chunks: Pin<Box<dyn futures_core::Stream<Item = Result<OpenaiChatChunk, LlmError>> + Send>>,
}
impl Stream for ChunkToLegacyEventStream {
type Item = LegacyStreamEvent;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = &mut *self;
poll_fn(|cx| match Pin::new(&mut this.chunks).poll_next(cx) {
Poll::Ready(Some(Ok(chunk))) => {
for choice in &chunk.choices {
let delta = &choice.delta;
if let Some(content) = &delta.content {
return Poll::Ready(Some(LegacyStreamEvent::AssistantTextDelta {
text: content.clone(),
}));
}
if let Some(tool_calls) = &delta.tool_calls
&& let Some(tc) = tool_calls.first()
{
let OpenaiToolCall::Function { id, function } = tc;
let args: Value =
serde_json::from_str(&function.arguments).unwrap_or(Value::Null);
return Poll::Ready(Some(LegacyStreamEvent::ToolExecutionStarted {
tool_name: function.name.clone(),
input: args,
tool_call_id: id.clone(),
}));
}
if let Some(finish_reason) = &choice.finish_reason {
return Poll::Ready(Some(LegacyStreamEvent::TurnComplete {
reason: *finish_reason,
}));
}
}
if let Some(usage) = &chunk.usage {
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate {
usage: *usage,
}));
}
Poll::Ready(None)
}
Poll::Ready(Some(Err(e))) => Poll::Ready(Some(LegacyStreamEvent::error(e.to_string()))),
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
})
.poll_unpin(cx)
}
}
// --- 内部:LegacyStreamEvent → 新 StreamEvent ---
struct LegacyToIrEventStream {
inner: Pin<Box<dyn futures_core::Stream<Item = LegacyStreamEvent> + Send>>,
}
impl Stream for LegacyToIrEventStream {
type Item = Result<StreamEvent, LlmError>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = &mut *self;
match Pin::new(&mut this.inner).poll_next(cx) {
Poll::Ready(Some(legacy)) => Poll::Ready(Some(Ok(map_legacy_to_ir(legacy)))),
Poll::Ready(None) => {
// 旧流结束 → 主动补一个 MessageCompletefull_response 为兜底空快照)。
// ponytail: Phase 0 中 OpenaiProvider 桥接层负责产出真实 MessageResponse
// 此处仅防止消费方无限等待。若 Provider 层已正确发出 MessageComplete
// LlmCycle 不会走到这里 —— 因为桥接层 inline 处理。
Poll::Ready(Some(Ok(StreamEvent::MessageComplete {
full_response: empty_message_response(),
})))
}
Poll::Pending => Poll::Pending,
}
}
}
fn empty_message_response() -> MessageResponse {
use crate::llm::types::message::Message;
use std::collections::HashMap;
MessageResponse {
id: String::new(),
model: String::new(),
message: Message::Assistant {
content: vec![],
},
usage: Usage::default(),
stop_reason: StopReason::Stop,
extra: HashMap::new(),
}
}
/// 把旧 LegacyStreamEvent 映射到新高精度 IR StreamEvent。
///
/// Phase 1 重写 Provider 后可直接删除此映射函数。当前映射语义:
/// - `AssistantTextDelta` → `TextDelta`
/// - `ToolExecutionStarted` → `ToolCallArgumentsDelta`OpenAI 单 chunk 模式下整段 arguments 一次性下发)
/// - `CostUpdate` → `CostUpdate`Usage → PartialUsage 全字段)
/// - `TurnComplete` → `MessageComplete`Phase 1 重写 Provider 后正确产出)
/// - `Error` → `Error`
///
/// ponytail: 这是一个"目前能跑通未来会被删除"的适配层。当前实现为单事件映射,
/// 旧 `ToolExecutionStarted` 携带的 (id, name) 暂未填入 IR 事件(消费方
/// Phase 2 中通过 MessageComplete.full_response.tool_use 提取)。Phase 1 重写时
/// 由 OpenAI Provider 直接产出 IR 流,整体删除此映射。
fn map_legacy_to_ir(legacy: LegacyStreamEvent) -> StreamEvent {
use crate::llm::types::response_v2::PartialUsage;
match legacy {
LegacyStreamEvent::AssistantTextDelta { text } => StreamEvent::TextDelta { text },
LegacyStreamEvent::ToolExecutionStarted { input, .. } => {
let arguments = serde_json::to_string(&input).unwrap_or_default();
StreamEvent::ToolCallArgumentsDelta { index: 0, arguments }
}
LegacyStreamEvent::ToolExecutionCompleted { .. } => {
// 旧 ToolExecutionCompleted 不在 IR 流协议中——工具执行是消费方职责。
// Phase 1 重写时此处整体删除。当前给一个无副作用的占位事件。
StreamEvent::CostUpdate {
usage: PartialUsage::default(),
}
}
LegacyStreamEvent::CostUpdate { usage } => StreamEvent::CostUpdate {
usage: PartialUsage {
prompt_tokens: Some(usage.prompt_tokens),
completion_tokens: Some(usage.completion_tokens),
total_tokens: Some(usage.total_tokens),
completion_tokens_details: usage.completion_tokens_details,
prompt_tokens_details: usage.prompt_tokens_details,
},
},
LegacyStreamEvent::TurnComplete { reason } => {
// 旧 TurnComplete 不直接对应 IR;映射为带 StopReason 的 MessageComplete。
// ponytail: Phase 1 重写 Provider 后此适配整体删除,
// OpenAI Provider 直接产出带正确 stop_reason 的 MessageComplete。
let _ = reason;
StreamEvent::MessageComplete {
full_response: empty_message_response(),
}
}
LegacyStreamEvent::Error { message } => StreamEvent::Error { message },
}
}
+9 -19
View File
@@ -20,15 +20,12 @@ use crate::llm::types::shared::ImageDetail;
/// 消费方 match 可直接区分文本和图片输入。 /// 消费方 match 可直接区分文本和图片输入。
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Message { pub enum Message {
/// 系统提示(User & Assistant 之外的引导指令)。 /// 系统提示(User & Assistant 之外的引导指令)。
System { System { content: Vec<ContentBlock> },
content: Vec<ContentBlock>,
},
/// 用户输入。 /// 用户输入。
User { User { content: Vec<ContentBlock> },
content: Vec<ContentBlock>,
},
/// 用户的图片输入(快捷构造,免去构造 ContentBlock 的 boilerplate)。 /// 用户的图片输入(快捷构造,免去构造 ContentBlock 的 boilerplate)。
UserImage { UserImage {
data: String, data: String,
@@ -36,9 +33,7 @@ pub enum Message {
detail: ImageDetail, detail: ImageDetail,
}, },
/// Assistant 回复内容块(可能包含 text、thinking、tool_use 等多种 block 的混合)。 /// Assistant 回复内容块(可能包含 text、thinking、tool_use 等多种 block 的混合)。
Assistant { Assistant { content: Vec<ContentBlock> },
content: Vec<ContentBlock>,
},
/// 工具调用结果。 /// 工具调用结果。
ToolResult { ToolResult {
tool_call_id: String, tool_call_id: String,
@@ -103,6 +98,7 @@ impl Message {
/// block 的逃生舱。 /// block 的逃生舱。
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ContentBlock { pub enum ContentBlock {
/// 纯文本。 /// 纯文本。
Text { text: String }, Text { text: String },
@@ -130,10 +126,7 @@ pub enum ContentBlock {
signature: Option<String>, signature: Option<String>,
}, },
/// 逃生舱:Provider 特定 block 透传(OpenAI Response 内置工具等)。 /// 逃生舱:Provider 特定 block 透传(OpenAI Response 内置工具等)。
Extension { Extension { kind: String, data: Value },
kind: String,
data: Value,
},
} }
/// 内容块类型标签 —— 用于 `StreamEvent::ContentBlockStart.block_type`。 /// 内容块类型标签 —— 用于 `StreamEvent::ContentBlockStart.block_type`。
@@ -141,6 +134,7 @@ pub enum ContentBlock {
/// 用途:在流式场景中,Provider 先下发 block 类型,再下发 block 内容增量。 /// 用途:在流式场景中,Provider 先下发 block 类型,再下发 block 内容增量。
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ContentBlockType { pub enum ContentBlockType {
/// 文本块。 /// 文本块。
Text, Text,
@@ -349,9 +343,7 @@ mod tests {
fn message_roundtrip_each_variant() { fn message_roundtrip_each_variant() {
let msgs = vec![ let msgs = vec![
Message::System { Message::System {
content: vec![ContentBlock::Text { content: vec![ContentBlock::Text { text: "sys".into() }],
text: "sys".into(),
}],
}, },
Message::User { Message::User {
content: vec![ContentBlock::Text { content: vec![ContentBlock::Text {
@@ -376,9 +368,7 @@ mod tests {
}, },
Message::ToolResult { Message::ToolResult {
tool_call_id: "call_1".into(), tool_call_id: "call_1".into(),
content: vec![ContentBlock::Text { content: vec![ContentBlock::Text { text: "ok".into() }],
text: "ok".into(),
}],
is_error: true, is_error: true,
}, },
]; ];
+1 -81
View File
@@ -1,9 +1,6 @@
pub mod message; pub mod message;
pub mod old_stream;
pub mod openai_message; pub mod openai_message;
pub mod request;
pub mod request_v2; pub mod request_v2;
pub mod response;
pub mod response_v2; pub mod response_v2;
pub mod shared; pub mod shared;
pub mod tool; pub mod tool;
@@ -12,12 +9,7 @@ pub mod usage;
pub use openai_message::{ pub use openai_message::{
ContentField, FileData, ImageURL, InputAudio, OpenaiChatMessage, OpenaiContentPart, ContentField, FileData, ImageURL, InputAudio, OpenaiChatMessage, OpenaiContentPart,
}; };
pub use request::{OpenaiChatRequest, OpenaiTool, StreamOptions, ToolChoice};
pub use request_v2::{ExtraError, MessageRequest, ThinkingConfig}; pub use request_v2::{ExtraError, MessageRequest, ThinkingConfig};
pub use response::{
Annotation, Choice, ChunkChoice, Delta, Logprobs, OpenaiAudio, OpenaiChatChunk,
OpenaiChatResponse, TokenLogprob, TopLogprob, URLCitation,
};
pub use response_v2::{ pub use response_v2::{
ContentBlockBuilder, MessageResponse, PartialMessageResponse, PartialUsage, StopReason, ContentBlockBuilder, MessageResponse, PartialMessageResponse, PartialUsage, StopReason,
StreamEvent, StreamEvent,
@@ -26,77 +18,5 @@ pub use shared::{
AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier, AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier,
StopSequence, StopSequence,
}; };
pub use tool::{FunctionCall, OpenaiToolCall, OpenaiToolDefinition}; pub use tool::{FunctionCall, OpenaiToolCall, ToolChoice, ToolDef};
pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage}; pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage};
// Re-export IR 内容块 / 消息类型供 `types::ContentBlock` 等历史路径消费。
//
// 注意:以下别名 *故意不暴露* `pub type Message = message::Message`、
// `pub type ContentBlock = message::ContentBlock` —— 新 `Message` / `ContentBlock` /
// `StopReason` 是独立类型,由 `Message` / `ContentBlock` / `StopReason` 直接路径访问,
// 旧别名(指 `OpenaiChatMessage` / `OpenaiContentPart` / `FinishReason`)已移除,
// 避免新类型阴影。Phase 2 完成后再统一收敛。
//
// Phase 1 起移除 `ChatRequest` 别名 —— 新代码统一使用 `MessageRequest`v2 IR)。
// `ChatResponse` 结构体仍存在,作为 OpenAI `chat_inner()` 内部 wire-format 转换目标。
/// 旧 wire-format 响应结构(保留用于 OpenAI 内部转换层)。
#[deprecated(since = "0.1.0", note = "请改用 MessageResponse")]
#[derive(Debug, Clone)]
pub struct ChatResponse {
pub message: OpenaiChatMessage,
pub usage: Usage,
pub stop_reason: Option<FinishReason>,
}
#[allow(deprecated)]
impl From<OpenaiChatResponse> for ChatResponse {
fn from(response: OpenaiChatResponse) -> Self {
let message = response
.choices
.first()
.map(|c| c.message.clone())
.unwrap_or_else(|| OpenaiChatMessage::assistant_text(""));
let stop_reason = response.choices.first().and_then(|c| c.finish_reason);
ChatResponse {
message,
usage: response.usage,
stop_reason,
}
}
}
#[allow(deprecated)]
impl From<ChatResponse> for OpenaiChatChunk {
fn from(response: ChatResponse) -> Self {
let delta = Delta::from(response.message.clone());
let chunk_choice = ChunkChoice {
index: 0,
delta,
logprobs: None,
finish_reason: response.stop_reason,
};
OpenaiChatChunk {
id: format!(
"chunk-{}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0)
),
object: "chat.completion.chunk".to_string(),
created: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0),
model: String::new(),
choices: vec![chunk_choice],
usage: Some(response.usage),
system_fingerprint: None,
}
}
}
/// 工具定义别名(无新类型冲突,保留)。
#[deprecated(since = "0.1.0", note = "ToolDefinition 仍直接对应 OpenAI wire-format;未来 v0.2 引入 IR 工具类型后会再次更新")]
pub type ToolDefinition = OpenaiToolDefinition;
-45
View File
@@ -1,45 +0,0 @@
//! 旧版流式事件 —— Phase 0 临时保留,仅供 `stream.rs` 中 `parse_chunk_stream` 内部使用。
//!
//! Phase 0 中:高精度 `StreamEvent`(定义在 `response_v2.rs`)是唯一的对外
//! `StreamEvent`,旧变体迁移至此模块改名为 `LegacyStreamEvent`
//! 由 `parse_chunk_stream()` 内部消费 `LegacyStreamEvent`,对外返回值已被
//! 重映射为新 `StreamEvent`。
//!
//! Phase 1 重写 Provider 时,`parse_chunk_stream` 可直接消费新事件流后整体删除此文件。
use crate::llm::types::shared::FinishReason;
use crate::llm::types::usage::Usage;
use serde_json::Value;
/// 旧 `StreamEvent` 变体迁移后的别名 —— 仅供 `stream.rs` 内部使用。
#[derive(Debug, Clone)]
pub enum LegacyStreamEvent {
/// 助手回复文本增量。
AssistantTextDelta { text: String },
/// 工具调用开始。
ToolExecutionStarted {
tool_name: String,
input: Value,
tool_call_id: String,
},
/// 工具调用完成。
ToolExecutionCompleted {
tool_name: String,
output: Value,
is_error: bool,
},
/// Token 用量更新。
CostUpdate { usage: Usage },
/// 一轮会话完成。
TurnComplete { reason: FinishReason },
/// 错误事件。
Error { message: String },
}
impl LegacyStreamEvent {
pub(crate) fn error(message: impl Into<String>) -> Self {
Self::Error {
message: message.into(),
}
}
}
-184
View File
@@ -1,184 +0,0 @@
use crate::llm::types::shared::{ResponseFormat, ServiceTier, StopSequence};
use crate::llm::types::tool::OpenaiToolDefinition;
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StreamOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub include_usage: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub include_obfuscation: Option<bool>,
}
#[derive(Debug, Clone)]
#[derive(Default)]
pub enum ToolChoice {
#[default]
None,
Auto,
Required,
Named { name: String },
AllowedTools { tool_names: Vec<String> },
}
impl Serialize for ToolChoice {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match self {
ToolChoice::None => serializer.serialize_str("none"),
ToolChoice::Auto => serializer.serialize_str("auto"),
ToolChoice::Required => serializer.serialize_str("required"),
ToolChoice::Named { name } => {
let obj = serde_json::json!({
"type": "function",
"function": { "name": name }
});
obj.serialize(serializer)
}
ToolChoice::AllowedTools { tool_names } => {
let obj = serde_json::json!({
"type": "function",
"function": { "name": tool_names.first().cloned().unwrap_or_default() }
});
obj.serialize(serializer)
}
}
}
}
impl<'de> Deserialize<'de> for ToolChoice {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = Value::deserialize(deserializer)?;
match value {
Value::String(s) => match s.as_str() {
"none" => Ok(ToolChoice::None),
"auto" => Ok(ToolChoice::Auto),
"required" => Ok(ToolChoice::Required),
_ => Err(serde::de::Error::custom(format!(
"unknown tool choice: {s}"
))),
},
Value::Object(obj) => {
let typ = obj.get("type").and_then(|v| v.as_str()).ok_or_else(|| {
serde::de::Error::custom("missing 'type' field in tool_choice")
})?;
if typ == "function" {
let func =
obj.get("function")
.and_then(|v| v.as_object())
.ok_or_else(|| {
serde::de::Error::custom("missing 'function' field in tool_choice")
})?;
let name = func.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
serde::de::Error::custom("missing 'function.name' in tool_choice")
})?;
Ok(ToolChoice::Named {
name: name.to_string(),
})
} else {
Err(serde::de::Error::custom(format!(
"unknown tool_choice type: {typ}"
)))
}
}
_ => Err(serde::de::Error::custom(
"tool_choice must be a string or object",
)),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "type")]
pub enum OpenaiTool {
Function { function: OpenaiToolDefinition },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AudioParam {
pub format: String,
pub voice: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PredictionContent {
#[serde(rename = "type")]
pub pred_type: String,
pub content: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UserLocation {
#[serde(rename = "type")]
pub loc_type: String,
pub approximate: Approximate,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Approximate {
pub city: String,
pub country: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub region: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timezone: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WebSearchOptions {
pub search_context_size: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_location: Option<UserLocation>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub struct OpenaiChatRequest {
pub model: String,
pub messages: Vec<crate::llm::types::openai_message::OpenaiChatMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
pub frequency_penalty: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logit_bias: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub n: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub presence_penalty: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<ResponseFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub seed: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub service_tier: Option<ServiceTier>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop: Option<StopSequence>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stream: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stream_options: Option<StreamOptions>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<OpenaiTool>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ToolChoice>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parallel_tool_calls: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extra_headers: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extra_body: Option<Value>,
}
+49 -17
View File
@@ -9,15 +9,15 @@ use serde_json::Value;
use thiserror::Error; use thiserror::Error;
use crate::llm::types::message::Message; use crate::llm::types::message::Message;
use crate::llm::types::request::ToolChoice; use crate::llm::types::tool::ToolChoice;
use crate::llm::types::tool::OpenaiToolDefinition; use crate::llm::types::tool::ToolDef;
/// Provider 无关的请求类型。 /// Provider 无关的请求类型。
/// ///
/// 设计要点: /// 设计要点:
/// - `system` 字段不存在;system 提示由调用方通过 `Message::System` 在 `messages` 中表达。 /// - `system` 字段不存在;system 提示由调用方通过 `Message::System` 在 `messages` 中表达。
/// - `tools` / `tool_choice` 直接复用现有 `OpenaiToolDefinition` / `ToolChoice` /// - `tools` 使用 Provider 无关的 `ToolDef` IR;各 Provider 适配层在 `convert_request`
/// (10a §251 决策:先复用旧类型,Phase 2 切换为新 `ToolDefinition` 后再调整) /// 中转换为对应 wire format。`tool_choice` 复用现有 `ToolChoice`
/// - `extra` 作为逃生舱:Provider 特定字段(`web_search_options`、`previous_response_id` 等) /// - `extra` 作为逃生舱:Provider 特定字段(`web_search_options`、`previous_response_id` 等)
/// 通过 `extra.set_extra / get_extra` 传递,避免持续膨胀本结构体。 /// 通过 `extra.set_extra / get_extra` 传递,避免持续膨胀本结构体。
#[derive(Debug, Clone, Default, Serialize, Deserialize)] #[derive(Debug, Clone, Default, Serialize, Deserialize)]
@@ -26,8 +26,8 @@ pub struct MessageRequest {
pub model: String, pub model: String,
/// 消息列表(包含 system / user / assistant / tool_result 等所有变体)。 /// 消息列表(包含 system / user / assistant / tool_result 等所有变体)。
pub messages: Vec<Message>, pub messages: Vec<Message>,
/// 工具定义列表。 /// 工具定义列表Provider 无关 IR
pub tools: Vec<OpenaiToolDefinition>, pub tools: Vec<ToolDef>,
/// 工具选择策略。 /// 工具选择策略。
pub tool_choice: ToolChoice, pub tool_choice: ToolChoice,
/// 最大输出 token 数。 /// 最大输出 token 数。
@@ -127,14 +127,9 @@ mod tests {
#[test] #[test]
fn extra_set_and_get_roundtrip() { fn extra_set_and_get_roundtrip() {
let mut req = MessageRequest::default(); let mut req = MessageRequest::default();
req.set_extra( req.set_extra("previous_response_id", "resp_abc123");
"previous_response_id",
"resp_abc123",
);
let v: Option<String> = req let v: Option<String> = req.get_extra("previous_response_id").expect("get_extra ok");
.get_extra("previous_response_id")
.expect("get_extra ok");
assert_eq!(v.as_deref(), Some("resp_abc123")); assert_eq!(v.as_deref(), Some("resp_abc123"));
let missing: Option<String> = req.get_extra("missing").expect("missing ok"); 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"); let opts: Options = req.get_extra_as().expect("get_extra_as ok");
assert_eq!( assert_eq!(opts.web_search_options.search_context_size, "high");
opts.web_search_options.search_context_size,
"high"
);
assert_eq!(opts.user.as_deref(), Some("u_123")); assert_eq!(opts.user.as_deref(), Some("u_123"));
} }
@@ -206,4 +198,44 @@ mod tests {
assert_eq!(decoded.stream, req.stream); assert_eq!(decoded.stream, req.stream);
assert_eq!(decoded.extra.get("trace_id"), Some(&json!("t-1"))); 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 字段");
}
} }
-181
View File
@@ -1,181 +0,0 @@
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};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenLogprob {
pub token: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub bytes: Option<Vec<u32>>,
pub logprob: f64,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_logprobs: Option<Vec<TopLogprob>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TopLogprob {
pub token: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub bytes: Option<Vec<u32>>,
pub logprob: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Logprobs {
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<Vec<TokenLogprob>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub refusal: Option<Vec<TokenLogprob>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct URLCitation {
pub end_index: u32,
pub start_index: u32,
#[serde(skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
pub url: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Annotation {
#[serde(rename = "type")]
pub ann_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub url_citation: Option<URLCitation>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OpenaiAudio {
pub id: String,
pub data: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub expires_at: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub transcript: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Choice {
pub index: u32,
pub message: OpenaiChatMessage,
#[serde(skip_serializing_if = "Option::is_none")]
pub finish_reason: Option<FinishReason>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprobs: Option<Logprobs>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OpenaiChatResponse {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub choices: Vec<Choice>,
pub usage: Usage,
#[serde(skip_serializing_if = "Option::is_none")]
pub system_fingerprint: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub service_tier: Option<ServiceTier>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Delta {
#[serde(skip_serializing_if = "Option::is_none")]
pub role: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub refusal: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<OpenaiToolCall>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChunkChoice {
pub index: u32,
pub delta: Delta,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprobs: Option<Logprobs>,
#[serde(skip_serializing_if = "Option::is_none")]
pub finish_reason: Option<FinishReason>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OpenaiChatChunk {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub choices: Vec<ChunkChoice>,
#[serde(skip_serializing_if = "Option::is_none")]
pub usage: Option<Usage>,
#[serde(skip_serializing_if = "Option::is_none")]
pub system_fingerprint: Option<String>,
}
impl From<OpenaiChatMessage> for Delta {
fn from(msg: OpenaiChatMessage) -> Self {
match msg {
OpenaiChatMessage::Assistant {
content,
tool_calls,
..
} => Delta {
role: Some("assistant".to_string()),
content: match content {
ContentField::String(s) => Some(s),
ContentField::Array(parts) => {
let mut text = String::new();
for part in parts {
if let OpenaiContentPart::Text { text: t } = part {
text.push_str(&t);
}
}
if text.is_empty() {
None
} else {
Some(text)
}
}
},
refusal: None,
tool_calls,
},
_ => Delta {
role: None,
content: None,
refusal: None,
tool_calls: None,
},
}
}
}
impl From<OpenaiChatResponse> for OpenaiChatChunk {
fn from(response: OpenaiChatResponse) -> Self {
let choices = response
.choices
.into_iter()
.map(|c| ChunkChoice {
index: c.index,
delta: Delta::from(c.message),
logprobs: c.logprobs,
finish_reason: c.finish_reason,
})
.collect();
OpenaiChatChunk {
id: response.id,
object: "chat.completion.chunk".to_string(),
created: response.created,
model: response.model,
choices,
usage: Some(response.usage),
system_fingerprint: response.system_fingerprint,
}
}
}
+35 -23
View File
@@ -19,6 +19,7 @@ use crate::llm::types::usage::{CompletionTokensDetails, PromptTokensDetails, Usa
/// Phase 2 完成时统一收敛。 /// Phase 2 完成时统一收敛。
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum StopReason { pub enum StopReason {
/// 自然停止。 /// 自然停止。
Stop, Stop,
@@ -164,11 +165,15 @@ pub enum ContentBlockBuilder {
/// `thinking_signature`,最终通过 `finalize()` 回填到 `full_response` 的 `Thinking` block 中。 /// `thinking_signature`,最终通过 `finalize()` 回填到 `full_response` 的 `Thinking` block 中。
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum StreamEvent { pub enum StreamEvent {
/// 消息开始(元信息)。 /// 消息开始(元信息)。
MessageStart { id: String, model: String }, MessageStart { id: String, model: String },
/// 内容块开始(告知块类型,携带 id/name for ToolUse)。 /// 内容块开始(告知块类型,携带 id/name for ToolUse)。
ContentBlockStart { index: u32, block_type: ContentBlockType }, ContentBlockStart {
index: u32,
block_type: ContentBlockType,
},
/// 内容块结束标记。 /// 内容块结束标记。
ContentBlockEnd { index: u32 }, ContentBlockEnd { index: u32 },
/// 文本增量。 /// 文本增量。
@@ -187,6 +192,24 @@ pub enum StreamEvent {
MessageComplete { full_response: MessageResponse }, MessageComplete { full_response: MessageResponse },
/// 错误事件。 /// 错误事件。
Error { message: String }, 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 true
} }
StreamEvent::ToolCallArgumentsDelta { index, arguments } => { StreamEvent::ToolCallArgumentsDelta { index, arguments } => {
if let Some(ContentBlockBuilder::ToolUse { if let Some(ContentBlockBuilder::ToolUse { arguments: buf, .. }) =
arguments: buf, .. self.blocks.get_mut(index)
}) = self.blocks.get_mut(index)
{ {
buf.push_str(arguments); buf.push_str(arguments);
} }
@@ -348,6 +370,9 @@ impl PartialMessageResponse {
self.is_errored = true; self.is_errored = true;
false 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()); let mut content_blocks = Vec::with_capacity(self.blocks.len());
for (idx, builder) in self.blocks { for (idx, builder) in self.blocks {
let block = Self::builder_to_block(idx, builder, self.thinking_signature.as_deref()) let block = Self::builder_to_block(idx, builder, self.thinking_signature.as_deref())
.map_err(|e| LlmError::Other(format!( .map_err(|e| LlmError::Other(format!("partial 块 #{idx} finalize 失败: {e}")))?;
"partial 块 #{idx} finalize 失败: {e}"
)))?;
content_blocks.push(block); content_blocks.push(block);
} }
@@ -743,10 +766,7 @@ mod tests {
Message::Assistant { content } => { Message::Assistant { content } => {
assert_eq!(content.len(), 2); assert_eq!(content.len(), 2);
match (&content[0], &content[1]) { 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!(t1, "first");
assert_eq!(t2, "second"); assert_eq!(t2, "second");
} }
@@ -772,9 +792,7 @@ mod tests {
index: 0, index: 0,
block_type: ContentBlockType::Text, block_type: ContentBlockType::Text,
}, },
StreamEvent::TextDelta { StreamEvent::TextDelta { text: "x".into() },
text: "x".into(),
},
StreamEvent::ContentBlockEnd { index: 0 }, StreamEvent::ContentBlockEnd { index: 0 },
StreamEvent::MessageComplete { StreamEvent::MessageComplete {
full_response: empty_response(), full_response: empty_response(),
@@ -832,15 +850,9 @@ mod tests {
block_type: ContentBlockType::Text, block_type: ContentBlockType::Text,
}, },
StreamEvent::ContentBlockEnd { index: 0 }, StreamEvent::ContentBlockEnd { index: 0 },
StreamEvent::TextDelta { StreamEvent::TextDelta { text: "t".into() },
text: "t".into(), StreamEvent::ThinkingDelta { text: "p".into() },
}, StreamEvent::RefusalDelta { text: "r".into() },
StreamEvent::ThinkingDelta {
text: "p".into(),
},
StreamEvent::RefusalDelta {
text: "r".into(),
},
StreamEvent::ToolCallArgumentsDelta { StreamEvent::ToolCallArgumentsDelta {
index: 1, index: 1,
arguments: "{\"x\":1}".into(), arguments: "{\"x\":1}".into(),
+2
View File
@@ -13,6 +13,7 @@ pub enum Role {
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum FinishReason { pub enum FinishReason {
Stop, Stop,
Length, Length,
@@ -67,6 +68,7 @@ pub enum StopSequence {
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "type")] #[serde(rename_all = "snake_case", tag = "type")]
#[non_exhaustive]
pub enum ResponseFormat { pub enum ResponseFormat {
Text, Text,
JsonObject, JsonObject,
+130
View File
@@ -1,6 +1,25 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::Value; 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)] #[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct OpenaiToolDefinition { pub struct OpenaiToolDefinition {
pub name: String, pub name: String,
@@ -12,6 +31,27 @@ pub struct OpenaiToolDefinition {
pub strict: Option<bool>, 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)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FunctionCall { pub struct FunctionCall {
pub name: String, pub name: String,
@@ -23,3 +63,93 @@ pub struct FunctionCall {
pub enum OpenaiToolCall { pub enum OpenaiToolCall {
Function { id: String, function: FunctionCall }, Function { id: String, function: FunctionCall },
} }
/// 工具选择策略 —— Phase 13 从 `types::request::ToolChoice` 迁入。
///
/// `#[non_exhaustive]` 预留扩展空间。
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub enum ToolChoice {
#[default]
None,
Auto,
Required,
Named {
name: String,
},
AllowedTools {
tool_names: Vec<String>,
},
}
impl Serialize for ToolChoice {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match self {
ToolChoice::None => serializer.serialize_str("none"),
ToolChoice::Auto => serializer.serialize_str("auto"),
ToolChoice::Required => serializer.serialize_str("required"),
ToolChoice::Named { name } => {
let obj = serde_json::json!({
"type": "function",
"function": { "name": name }
});
obj.serialize(serializer)
}
ToolChoice::AllowedTools { tool_names } => {
let obj = serde_json::json!({
"type": "function",
"function": { "name": tool_names.first().cloned().unwrap_or_default() }
});
obj.serialize(serializer)
}
}
}
}
impl<'de> Deserialize<'de> for ToolChoice {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = Value::deserialize(deserializer)?;
match value {
Value::String(s) => match s.as_str() {
"none" => Ok(ToolChoice::None),
"auto" => Ok(ToolChoice::Auto),
"required" => Ok(ToolChoice::Required),
_ => Err(serde::de::Error::custom(format!(
"unknown tool choice: {s}"
))),
},
Value::Object(obj) => {
let typ = obj.get("type").and_then(|v| v.as_str()).ok_or_else(|| {
serde::de::Error::custom("missing 'type' field in tool_choice")
})?;
if typ == "function" {
let func =
obj.get("function")
.and_then(|v| v.as_object())
.ok_or_else(|| {
serde::de::Error::custom("missing 'function' field in tool_choice")
})?;
let name = func.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
serde::de::Error::custom("missing 'function.name' in tool_choice")
})?;
Ok(ToolChoice::Named {
name: name.to_string(),
})
} else {
Err(serde::de::Error::custom(format!(
"unknown tool_choice type: {typ}"
)))
}
}
_ => Err(serde::de::Error::custom(
"tool_choice must be a string or object",
)),
}
}
}
+8 -3
View File
@@ -6,17 +6,22 @@ pub mod knowledge;
pub mod retriever; pub mod retriever;
pub mod store; pub mod store;
pub mod types; pub mod types;
pub mod vector;
pub mod vector_store;
// 高频类型(大多数下游需要) // 高频类型(大多数下游需要)
pub use conversation::{ConversationMemory, ConversationMemoryConfig}; pub use conversation::{ConversationMemory, ConversationMemoryConfig};
pub use error::MemoryError; pub use error::MemoryError;
pub use knowledge::KnowledgeStore; pub use knowledge::KnowledgeStore;
pub use retriever::MemoryRetriever; pub use retriever::MemoryRetriever;
pub use store::{InMemoryStore, MemoryStore}; pub use store::{InMemoryStore, MemoryStore, SqliteStore};
#[allow(deprecated)]
pub use vector::{InMemoryVectorRetriever, VectorRetriever};
pub use vector_store::{InMemoryVectorStore, PersistentVectorStore, RagPipeline, VectorStore};
// 低频类型(配置/高级使用) // 低频类型(配置/高级使用)
pub use conversation::MemoryStrategy; pub use conversation::MemoryStrategy;
pub use knowledge::{PageIndexEntry, KNOWLEDGE_PREFIX}; pub use knowledge::{KNOWLEDGE_PREFIX, PageIndexEntry};
pub use retriever::{RetrieverConfig, RetrievalResult, ScoredItem}; pub use retriever::{RetrievalResult, RetrieverConfig, ScoredItem};
pub use store::{EvictionConfig, EvictionPolicy}; pub use store::{EvictionConfig, EvictionPolicy};
pub use types::{KnowledgePage, MemoryFilter, MemoryItem}; 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)] #[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[non_exhaustive]
pub enum MemoryStrategy { pub enum MemoryStrategy {
/// 滑动窗口:达到上限时删除最旧消息。 /// 滑动窗口:达到上限时删除最旧消息。
SlidingWindow, SlidingWindow,
@@ -160,7 +161,12 @@ impl ConversationMemory {
} }
fn make_message_id(&self, index: usize, now: &OffsetDateTime) -> String { 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) { async fn maybe_evict_and_compact(&mut self) {
@@ -175,15 +181,16 @@ impl ConversationMemory {
} }
if let Some(ref compact_config) = self.config.compact_config if let Some(ref compact_config) = self.config.compact_config
&& should_compact(&self.messages, compact_config, &self.compact_state) { && should_compact(&self.messages, compact_config, &self.compact_state)
let keep_recent = compact_config.keep_recent; {
let freed = microcompact(&mut self.messages, keep_recent); let keep_recent = compact_config.keep_recent;
if freed > 0 { let freed = microcompact(&mut self.messages, keep_recent);
self.compact_state.record_success(); if freed > 0 {
} else { self.compact_state.record_success();
let _ = self.compact_state.record_failure(); } else {
} let _ = self.compact_state.record_failure();
} }
}
} }
} }
@@ -196,7 +203,8 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn add_and_get_history() { async fn add_and_get_history() {
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>; 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("hello")).await.unwrap();
conv.add_message(Message::user_text("world")).await.unwrap(); conv.add_message(Message::user_text("world")).await.unwrap();
assert_eq!(conv.len(), 2); assert_eq!(conv.len(), 2);
@@ -211,9 +219,7 @@ mod tests {
conv.add_message(Message::tool_result("call_1", "ok", false)) conv.add_message(Message::tool_result("call_1", "ok", false))
.await .await
.unwrap(); .unwrap();
conv.add_message(Message::assistant("done")) conv.add_message(Message::assistant("done")).await.unwrap();
.await
.unwrap();
let original = conv.get_history().to_vec(); let original = conv.get_history().to_vec();
assert_eq!(original.len(), 2); assert_eq!(original.len(), 2);
@@ -263,7 +269,8 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn clear_empties_messages() { async fn clear_empties_messages() {
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>; 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(); conv.add_message(Message::user_text("hello")).await.unwrap();
assert!(!conv.is_empty()); assert!(!conv.is_empty());
conv.clear().await.unwrap(); conv.clear().await.unwrap();
+1
View File
@@ -6,6 +6,7 @@ use thiserror::Error;
/// ///
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查环境变量、重试)。 /// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查环境变量、重试)。
#[derive(Debug, Error)] #[derive(Debug, Error)]
#[non_exhaustive]
pub enum MemoryError { pub enum MemoryError {
/// 按 ID 未找到指定记忆条目。可重试——通常是 namespace 拼写错误或条目已被淘汰。 /// 按 ID 未找到指定记忆条目。可重试——通常是 namespace 拼写错误或条目已被淘汰。
#[error("未找到记忆条目 '{0}',请检查 ID 或 namespace 是否正确")] #[error("未找到记忆条目 '{0}',请检查 ID 或 namespace 是否正确")]
+6 -3
View File
@@ -57,8 +57,8 @@ impl KnowledgeStore {
} }
let now = OffsetDateTime::now_utc(); let now = OffsetDateTime::now_utc();
let id = format!("{KNOWLEDGE_PREFIX}{}", page.id); let id = format!("{KNOWLEDGE_PREFIX}{}", page.id);
let content = serde_json::to_string(&page) let content =
.map_err(|e| MemoryError::Serialization(e.to_string()))?; serde_json::to_string(&page).map_err(|e| MemoryError::Serialization(e.to_string()))?;
let item = MemoryItem { let item = MemoryItem {
id, id,
content, content,
@@ -128,7 +128,10 @@ impl KnowledgeStore {
.filter(|entry| { .filter(|entry| {
entry.title.to_lowercase().contains(&needle) entry.title.to_lowercase().contains(&needle)
|| entry.summary.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()) .map(|entry| entry.id.clone())
.collect() .collect()
+15 -8
View File
@@ -97,7 +97,11 @@ impl MemoryRetriever {
// 4. 过滤 → 排序 → 截取 // 4. 过滤 → 排序 → 截取
items.retain(|i| i.score >= self.config.min_score); 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); items.truncate(self.config.max_results);
Ok(RetrievalResult { Ok(RetrievalResult {
@@ -159,12 +163,12 @@ fn char_bigrams(s: &str) -> Vec<String> {
fn default_stop_words() -> HashSet<String> { fn default_stop_words() -> HashSet<String> {
[ [
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had",
"had", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall", "can",
"can", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what", "which",
"which", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor", "so", "if",
"so", "if", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at", "by", "of", "as",
"by", "of", "as", "into", "through", "during", "before", "after", "above", "below", "into", "through", "during", "before", "after", "above", "below",
] ]
.iter() .iter()
.map(|s| s.to_string()) .map(|s| s.to_string())
@@ -236,7 +240,10 @@ mod tests {
min_score: 0.99, min_score: 0.99,
}; };
let retriever = MemoryRetriever::new(ks, config); 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()); assert!(result.items.is_empty());
} }
+7 -259
View File
@@ -1,14 +1,16 @@
//! MemoryStore 抽象接口与默认实现。 //! MemoryStore 抽象接口与默认实现。
use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait; use async_trait::async_trait;
use time::OffsetDateTime;
use crate::memory::error::MemoryError; use crate::memory::error::MemoryError;
use crate::memory::types::{MemoryFilter, MemoryItem}; 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 等)。 /// 下游可实现此 trait 以对接持久化后端(JSON 文件、SQLite、Redis 等)。
@@ -32,6 +34,7 @@ pub trait MemoryStore: Send + Sync {
/// 淘汰策略。 /// 淘汰策略。
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
#[non_exhaustive]
pub enum EvictionPolicy { pub enum EvictionPolicy {
/// 不淘汰(默认)。 /// 不淘汰(默认)。
None, 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);
}
}
+366
View File
@@ -0,0 +1,366 @@
//! 进程内默认实现 —— 基于 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);
}
// ===== Phase 11 Step 11.3 并发测试 =====
#[tokio::test]
async fn concurrent_writers_max_pressure() {
use std::sync::Arc;
let store = Arc::new(InMemoryStore::new());
let mut handles = Vec::new();
for i in 0..100 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let id = format!("concurrent_{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);
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 concurrent_mixed_read_write() {
use std::sync::Arc;
use std::time::Duration;
let store = Arc::new(InMemoryStore::new());
// 预热 20 条
for i in 0..20 {
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
let mut handles = Vec::new();
// 5 个写者
for w in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let mut i = 0;
while tokio::time::Instant::now() < deadline {
let id = format!("writer{w}_item{i}");
s.save(make_item(&id)).await.unwrap();
i += 1;
}
}));
}
// 5 个读者
for _ in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
while tokio::time::Instant::now() < deadline {
let _ = s.list(&MemoryFilter::default()).await.unwrap();
}
}));
}
for h in handles {
h.await.unwrap();
}
}
#[tokio::test]
async fn concurrent_capacity_eviction() {
use std::sync::Arc;
let eviction = EvictionConfig {
policy: EvictionPolicy::Capacity { max_items: 10 },
check_interval: 1,
};
let store = Arc::new(InMemoryStore::with_eviction(eviction));
let mut handles = Vec::new();
for i in 0..15 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
s.save(make_item(&format!("item_{i}"))).await.unwrap();
}));
}
for h in handles {
h.await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
// 写者全部完成后必 ≤ max_items(部分路径上可能短暂 >10 但全部完成时应 ≤10)
assert!(
list.len() <= 10,
"expected <= 10 items after all writers done, got {}",
list.len()
);
}
}
+623
View File
@@ -0,0 +1,623 @@
//! 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());
}
}
// ===== Phase 11 Step 11.3 并发测试 =====
#[tokio::test]
async fn concurrent_writers_max_pressure() {
use std::time::Duration;
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
let mut handles = Vec::new();
for i in 0..100 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let id = format!("concurrent_{i}");
// 设置每次 save 的 per-call timeout —— busy_timeout=5000ms 应足够
match tokio::time::timeout(
Duration::from_secs(10),
s.save(make_item(&id)),
)
.await
{
Ok(res) => res.unwrap(),
Err(_) => panic!("save({id}) timed out under 100-way concurrency"),
}
}));
}
for h in handles {
h.await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
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 concurrent_mixed_read_write() {
use std::time::Duration;
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
// 预热 20 条
for i in 0..20 {
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
let mut handles = Vec::new();
// 5 个写者
for w in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let mut i = 0;
while tokio::time::Instant::now() < deadline {
let id = format!("writer{w}_item{i}");
s.save(make_item(&id)).await.unwrap();
i += 1;
}
}));
}
// 5 个读者
for _ in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
while tokio::time::Instant::now() < deadline {
let _ = s.list(&MemoryFilter::default()).await.unwrap();
}
}));
}
for h in handles {
h.await.unwrap();
}
}
}
+244
View File
@@ -0,0 +1,244 @@
//! 语义向量检索抽象。
//!
//! 提供 [`VectorRetriever`] trait 定义与进程内引用实现 [`InMemoryVectorRetriever`]。
//! 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait;
use crate::memory::error::MemoryError;
/// 语义向量检索器抽象接口。
///
/// 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
/// 默认引用实现 [`InMemoryVectorRetriever`] 基于进程内 HashMap + 余弦相似度。
///
/// **稳定性**:实验性 APIv0.2.x),方法签名可能在 v0.3 中调整。
/// 若未来需要 `remove()` / `clear()` 等方法,将在此 trait 中追加(带默认实现)。
#[deprecated(since = "0.3.0", note = "请使用 memory::VectorStore")]
#[async_trait]
pub trait VectorRetriever: Send + Sync {
/// 将 `id` 对应的文本向量 `embeddings` 加入索引。
///
/// 重复调用同一 `id` 会覆盖已有向量。调用方负责保证 `embeddings` 维度
/// 与已索引向量一致——本 trait 不做维度校验。
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError>;
/// 检索与 `query` 向量最相似的 `k` 条记录。
///
/// 返回 `Vec<(id, score)>`,按 score 降序排列,score ∈ [0.0, 1.0]
/// (余弦相似度)。当 `k == 0`、索引为空或 query 为零向量时返回空 Vec。
async fn search(
&self,
query: Vec<f32>,
k: usize,
) -> Result<Vec<(String, f32)>, MemoryError>;
}
/// 进程内向量检索器 —— 基于 HashMap + 全量余弦相似度扫描。
///
/// 适用场景:单元测试、小规模验证(<10K 向量)。生产环境请对接真正的向量数据库。
///
/// **不保证**
/// - 不做向量维度校验(不同维度向量查询结果无意义但不 panic)
/// - `search()` 是 O(n) 全量扫描,未做索引加速
/// - 不保证高并发下查询时序与写入顺序一致
#[deprecated(since = "0.3.0", note = "请使用 memory::InMemoryVectorStore")]
pub struct InMemoryVectorRetriever {
vectors: Mutex<HashMap<String, Vec<f32>>>,
}
#[allow(deprecated)]
impl InMemoryVectorRetriever {
/// 创建空检索器。
pub fn new() -> Self {
Self {
vectors: Mutex::new(HashMap::new()),
}
}
}
#[allow(deprecated)]
impl Default for InMemoryVectorRetriever {
fn default() -> Self {
Self::new()
}
}
#[allow(deprecated)]
#[async_trait]
impl VectorRetriever for InMemoryVectorRetriever {
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError> {
let mut vectors = self
.vectors
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
vectors.insert(id, embeddings);
Ok(())
}
async fn search(
&self,
query: Vec<f32>,
k: usize,
) -> Result<Vec<(String, f32)>, MemoryError> {
let vectors = self
.vectors
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
if vectors.is_empty() || k == 0 {
return Ok(Vec::new());
}
let query_norm = dot(&query, &query).sqrt();
if query_norm == 0.0 {
return Ok(Vec::new());
}
let mut scored: Vec<(String, f32)> = vectors
.iter()
.map(|(id, vec)| {
let dot_product = dot(&query, vec);
let vec_norm = dot(vec, vec).sqrt();
let similarity = dot_product / (query_norm * vec_norm + 1e-10);
(id.clone(), similarity)
})
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(k);
Ok(scored)
}
}
/// 点积(手动循环,零依赖)。
///
/// 注意:`zip` 对不等长向量静默截断到较短者。引用实现不做维度校验,
/// 调用方应确保 `a` 和 `b` 等长——不等长时结果无意义但不 panic。
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
#[cfg(test)]
#[allow(deprecated)]
mod tests {
use super::*;
use std::sync::Arc;
use std::time::Duration;
#[tokio::test]
async fn basic_index_and_search() {
let retriever = InMemoryVectorRetriever::new();
retriever
.index("rust".into(), vec![1.0, 0.0, 0.0])
.await
.unwrap();
retriever
.index("python".into(), vec![0.0, 1.0, 0.0])
.await
.unwrap();
let results = retriever.search(vec![0.9, 0.1, 0.0], 2).await.unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, "rust");
assert!(results[0].1 > results[1].1);
}
#[tokio::test]
async fn search_empty_store() {
let retriever = InMemoryVectorRetriever::new();
let results = retriever.search(vec![1.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty());
}
#[tokio::test]
async fn search_zero_vector_returns_empty() {
let retriever = InMemoryVectorRetriever::new();
retriever
.index("a".into(), vec![1.0, 0.0, 0.0])
.await
.unwrap();
let results = retriever.search(vec![0.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty());
}
#[tokio::test]
async fn search_with_k_zero_returns_empty() {
let retriever = InMemoryVectorRetriever::new();
retriever
.index("a".into(), vec![1.0, 0.0, 0.0])
.await
.unwrap();
let results = retriever.search(vec![1.0, 0.0, 0.0], 0).await.unwrap();
assert!(results.is_empty());
}
#[tokio::test]
async fn concurrent_index() {
let retriever = Arc::new(InMemoryVectorRetriever::new());
let mut handles = Vec::new();
for i in 0..10 {
let r = Arc::clone(&retriever);
handles.push(tokio::spawn(async move {
r.index(format!("item_{i}"), vec![i as f32, 0.0, 0.0])
.await
.unwrap();
}));
}
for h in handles {
h.await.unwrap();
}
let results = retriever.search(vec![1.0, 0.0, 0.0], 20).await.unwrap();
assert_eq!(results.len(), 10);
let mut ids: Vec<String> = results.iter().map(|(id, _)| id.clone()).collect();
ids.sort();
ids.dedup();
assert_eq!(ids.len(), 10);
}
#[tokio::test]
async fn concurrent_index_and_search() {
let retriever = Arc::new(InMemoryVectorRetriever::new());
for i in 0..5 {
retriever
.index(format!("seed_{i}"), vec![i as f32, 0.0, 0.0])
.await
.unwrap();
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
let mut handles = Vec::new();
for w in 0..5 {
let r = Arc::clone(&retriever);
handles.push(tokio::spawn(async move {
let mut i = 0;
while tokio::time::Instant::now() < deadline {
r.index(format!("writer{w}_{i}"), vec![i as f32, 0.0, 0.0])
.await
.unwrap();
i += 1;
}
}));
}
for _ in 0..5 {
let r = Arc::clone(&retriever);
handles.push(tokio::spawn(async move {
while tokio::time::Instant::now() < deadline {
let _ = r.search(vec![1.0, 0.0, 0.0], 3).await.unwrap();
}
}));
}
for h in handles {
h.await.unwrap();
}
}
}
+937
View File
@@ -0,0 +1,937 @@
//! 向量存储抽象与实现 —— RAG 管线「存储与检索」环节。
//!
//! 提供 [`VectorStore`] trait 定义、进程内引用实现 [`InMemoryVectorStore`]
//! 以及基于 [`MemoryStore`] 的持久化包装 [`PersistentVectorStore`] 和
//! RAG 管线组合器 [`RagPipeline`]。
//!
//! 下游可实现 [`VectorStore`] trait 以对接专用向量数据库(pgvector / Qdrant 等)。
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use time::format_description::well_known::Rfc3339;
use time::OffsetDateTime;
use tracing::{debug, info};
use crate::document::{Document, RecursiveCharacterSplitter};
use crate::llm::embedding::Embedding;
use crate::memory::error::MemoryError;
use crate::memory::store::MemoryStore;
use crate::memory::types::{MemoryFilter, MemoryItem};
/// 向量存储抽象 —— 语义检索的核心接口。
///
/// 提供文档-向量的批量添加、余弦相似度搜索、批量删除三个核心操作。
/// 所有实现必须满足 `Send + Sync` 以支持跨 `.await` 调用。
///
/// # 并发安全
///
/// 实现内部必须使用线程安全的容器(如 `Mutex<HashMap>` 或 `RwLock`),
/// 允许跨多个 tokio task 共享 `&VectorStore` 引用。
///
/// # 与旧 `VectorRetriever` 的差异
///
/// - `add` 接受批量 `(doc, embedding)` 对;旧 `index` 仅接受单条
/// - `search` 返回 `(Document, f32)`;旧 `search` 返回 `(String, f32)`,调用方需自行维护 id→Document 映射
#[async_trait]
pub trait VectorStore: Send + Sync {
/// 批量添加文档及其向量。
///
/// `documents` 和 `embeddings` 必须等长。不等长时:
/// - 截取 `min(len)` 对处理(部分写入已发生)
/// - 返回 `Err(MemoryError::InvalidInput)` 告知截断
/// - 调用方可以 `let _ = store.add(...)` 忽略错误
async fn add(
&self,
documents: &[Document],
embeddings: &[Vec<f32>],
) -> Result<(), MemoryError>;
/// 检索与 `query` 向量最相似的 `k` 条记录。
///
/// 返回 `Vec<(Document, f32)>`,其中 `f32` 为余弦相似度分数,
/// 取值范围 `[0.0, 1.0]`(对单位向量),按分数降序排列。
///
/// # 守卫
///
/// - 空索引 → 返回 `vec![]`
/// - `k == 0` → 返回 `vec![]`
/// - 零向量(norm ≈ 0)→ 返回 `vec![]`
async fn search(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError>;
/// 批量删除文档(幂等)。
///
/// 不存在的 id 静默忽略,不会返回错误。
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError>;
/// 便捷方法:单条添加。
///
/// 等价于 `self.add(&[doc], &[emb]).await`。
async fn add_one(&self, doc: Document, emb: Vec<f32>) -> Result<(), MemoryError> {
self.add(&[doc], &[emb]).await
}
}
/// 内存向量存储 —— `VectorStore` 的引用实现。
///
/// 内部使用 `Mutex<HashMap<String, (Document, Vec<f32>)>>` 存储,
/// `search()` 执行 O(n) 全量余弦相似度扫描,适用于 ≤10K 条向量的场景。
///
/// # 并发安全
///
/// 使用 `std::sync::Mutex`(非 tokio Mutex)。
///
/// **锁持有时间评估**
/// - `add()` / `remove()`:微秒级(HashMap 插入/删除操作)
/// - `search()`:毫秒级(O(n) 全量扫描 + 余弦计算),对 10K 条 1536 维向量预估 1-10ms。
/// 实现时在锁内克隆数据快照到 `Vec` 后立即释放锁,在锁外进行余弦相似度计算,
/// 避免长时间持有锁阻塞并发写操作。
pub struct InMemoryVectorStore {
entries: Mutex<HashMap<String, (Document, Vec<f32>)>>,
}
impl InMemoryVectorStore {
/// 创建一个空存储。
pub fn new() -> Self {
Self {
entries: Mutex::new(HashMap::new()),
}
}
/// 从预填充的 entries 构造(供 `PersistentVectorStore` 使用)。
pub(crate) fn with_entries(
entries: HashMap<String, (Document, Vec<f32>)>,
) -> Self {
Self {
entries: Mutex::new(entries),
}
}
}
impl Default for InMemoryVectorStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl VectorStore for InMemoryVectorStore {
async fn add(
&self,
documents: &[Document],
embeddings: &[Vec<f32>],
) -> Result<(), MemoryError> {
let mut entries = self
.entries
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
let n = documents.len().min(embeddings.len());
if documents.len() != embeddings.len() {
tracing::warn!(
docs = documents.len(),
embs = embeddings.len(),
"InMemoryVectorStore::add 长度不匹配,截断到 min"
);
}
for i in 0..n {
entries.insert(documents[i].id.clone(), (documents[i].clone(), embeddings[i].clone()));
}
if documents.len() != embeddings.len() {
return Err(MemoryError::InvalidInput(format!(
"documents.len()={} 与 embeddings.len()={} 不等,已截断到 min={}",
documents.len(),
embeddings.len(),
n
)));
}
Ok(())
}
async fn search(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError> {
if k == 0 {
return Ok(Vec::new());
}
tracing::trace!(k, "InMemoryVectorStore::search");
// 零向量守卫:查询向量本身为零向量则返回空
let query_norm_sq: f32 = query.iter().map(|x| x * x).sum();
if query_norm_sq < 1e-20 {
return Ok(Vec::new());
}
// 锁内克隆快照,释放锁后在锁外计算余弦
let snapshot: Vec<(Document, Vec<f32>)> = {
let entries = self
.entries
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
entries.values().cloned().collect()
};
let mut scored: Vec<(Document, f32)> = Vec::with_capacity(snapshot.len());
for (doc, emb) in snapshot {
let score = cosine_similarity(query, &emb);
scored.push((doc, score));
}
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(k);
Ok(scored)
}
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError> {
let mut entries = self
.entries
.lock()
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
tracing::debug!(count = ids.len(), "InMemoryVectorStore::remove");
entries.retain(|key, _| !ids.iter().any(|id| id == key));
Ok(())
}
}
/// 点积。
///
/// `zip` 对不等长向量静默截断到较短者。调用方应保证 `a` 和 `b` 等长——
/// 不等长时结果无意义但不 panic。
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
/// 余弦相似度,加 `1e-10` 防除零。
///
/// 零向量与任意向量的相似度返回 `0.0`(因分母中 `1e-10` 保护 + 分子为 0)。
pub(crate) fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
let dot_product = dot(a, b);
let norm_a = dot(a, a).sqrt();
let norm_b = dot(b, b).sqrt();
dot_product / (norm_a * norm_b + 1e-10)
}
/// 持久化向量存储 —— 基于 [`MemoryStore`] 的持久化包装。
///
/// # 架构
///
/// 运行时全量加载到 [`InMemoryVectorStore`] 做余弦搜索,
/// 写操作(add/remove)同时同步到内存和后端 [`MemoryStore`]。
///
/// # 存储格式
///
/// 每条向量存为一条 [`MemoryItem`]
/// - `id`: `"vec:{namespace}:{doc_id}"`colon-separated namespace 前缀)
/// - `content`: JSON 序列化的向量条目(含 doc_id / content / metadata / mime_type / embedding
/// - `metadata`: 空 `serde_json::Value::Null`
///
/// # 构造开销
///
/// `new()` 通过 `store.list(prefix)` 全量加载已有条目,
/// 时间复杂度 O(N)(N 为已有向量数),适用于 ≤10K 条的场景。
pub struct PersistentVectorStore {
inner: InMemoryVectorStore,
store: Arc<dyn MemoryStore>,
namespace: String,
}
/// 持久化向量条目 —— JSON blob 格式。
#[derive(Serialize, Deserialize)]
struct VectorEntry {
doc_id: String,
content: String,
metadata: HashMap<String, String>,
mime_type: String,
embedding: Vec<f32>,
/// ISO 8601 创建时间(UTC),持久化 roundtrip 重建时保持原时间,
/// 避免 MemoryStore 的 TTL 淘汰策略误判。
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(default)]
created_at: Option<String>,
}
impl PersistentVectorStore {
/// 创建新的持久化向量存储,自动从 `store` 全量加载 namespace 下的所有条目。
///
/// `MemoryStore::list()` 由 `SqliteStore` 内部使用 `spawn_blocking` 卸载,
/// 加载过程本身在 async context 中即可,无需额外 spawn_blocking。
pub async fn new(
store: Arc<dyn MemoryStore>,
namespace: &str,
) -> Result<Self, MemoryError> {
let prefix = format!("vec:{namespace}:");
let filter = MemoryFilter {
prefix: Some(prefix.clone()),
..Default::default()
};
debug!(namespace = %namespace, "PersistentVectorStore::new — 开始全量加载");
let items = store.list(&filter).await?;
info!(count = items.len(), "PersistentVectorStore::new — 加载完成");
let mut entries: HashMap<String, (Document, Vec<f32>)> = HashMap::new();
for item in items {
let entry: VectorEntry = serde_json::from_str(&item.content)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
let doc = Document {
id: entry.doc_id,
content: entry.content,
metadata: entry.metadata,
mime_type: entry.mime_type,
};
entries.insert(doc.id.clone(), (doc, entry.embedding));
}
info!(entries = entries.len(), "PersistentVectorStore — 内存索引重建完成");
Ok(Self {
inner: InMemoryVectorStore::with_entries(entries),
store,
namespace: namespace.to_string(),
})
}
}
#[async_trait]
impl VectorStore for PersistentVectorStore {
async fn add(
&self,
documents: &[Document],
embeddings: &[Vec<f32>],
) -> Result<(), MemoryError> {
debug!(count = documents.len(), "PersistentVectorStore::add");
// 先逐个写持久化(失败时不污染内存)
for (doc, emb) in documents.iter().zip(embeddings.iter()) {
let entry = VectorEntry {
doc_id: doc.id.clone(),
content: doc.content.clone(),
metadata: doc.metadata.clone(),
mime_type: doc.mime_type.clone(),
embedding: emb.clone(),
created_at: Some(
OffsetDateTime::now_utc()
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format time: {e}")))?,
),
};
let json = serde_json::to_string(&entry)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
let key = format!("vec:{}:{}", self.namespace, doc.id);
let item = MemoryItem {
id: key,
content: json,
metadata: serde_json::Value::Null,
created_at: OffsetDateTime::now_utc(),
};
self.store.save(item).await?;
}
// 再写内存(持久化已成功写入,内存失败也不影响重启后恢复)
self.inner.add(documents, embeddings).await
}
async fn search(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError> {
tracing::trace!(k, "PersistentVectorStore::search");
self.inner.search(query, k).await
}
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError> {
debug!(count = ids.len(), "PersistentVectorStore::remove");
for id in ids {
let key = format!("vec:{}:{}", self.namespace, id);
self.store.delete(&key).await?;
}
self.inner.remove(ids).await
}
}
// ponytail: `with_entries` 当前仅供 `PersistentVectorStore::new` 使用;
// 后续如需 VecStore 之间迁移,可放宽到 `pub`。
/// RAG 管线组合器 —— 封装 `split → embed → store`ingest)和
/// `embed → store.search`retrieve)两个核心流程。
///
/// # 使用方式
///
/// ```ignore
/// let pipeline = RagPipeline::new(embedder, store, Some(splitter));
/// pipeline.ingest(&documents).await?;
/// let results = pipeline.retrieve("query", 5).await?;
/// ```
///
/// # 分割器
///
/// `splitter` 字段为 `Option<RecursiveCharacterSplitter>`
/// - `Some(splitter)` → `ingest()` 先分割再嵌入(调用方传入原始文档)
/// - `None` → `ingest()` 跳过分割,直接嵌入(调用方已分好 chunk)
pub struct RagPipeline {
embedder: Arc<dyn Embedding>,
store: Arc<dyn VectorStore>,
splitter: Option<RecursiveCharacterSplitter>,
}
impl RagPipeline {
/// 创建新的 RAG 管线。
///
/// 不设置分割器时,`ingest()` 跳过分割阶段,
/// 调用方传入的 Document 应已是分割好的 chunk。
pub fn new(
embedder: Arc<dyn Embedding>,
store: Arc<dyn VectorStore>,
splitter: Option<RecursiveCharacterSplitter>,
) -> Self {
Self {
embedder,
store,
splitter,
}
}
/// 摄取文档:分割 → 向量化 → 存储。
///
/// 流程:
/// 1. 如果 splitter 存在,先分割文档为 chunks
/// 2. 提取所有 chunk 的 content 为 `Vec<String>`
/// 3. `embedder.embed()` 批量向量化
/// 4. `store.add()` 批量存储
///
/// # 边界
///
/// - 空文档切片 → `Ok(())`,无操作
/// - 分割后 chunk 为空 → `Ok(())`,无操作
///
/// # 已知限制
///
/// 当前将所有 chunk 一次性传入 `embedder.embed()`,真实 Embedding Provider
/// (如 OpenAI)有批量大小限制,调用方需自行控制单次 ingest 的文档数(如 20 条/批)。
pub async fn ingest(&self, documents: &[Document]) -> Result<(), MemoryError> {
let chunks = match &self.splitter {
Some(splitter) => splitter.split(documents),
None => documents.to_vec(),
};
if chunks.is_empty() {
return Ok(());
}
let texts: Vec<String> = chunks.iter().map(|d| d.content.clone()).collect();
let embeddings = self
.embedder
.embed(&texts)
.await
.map_err(|e| MemoryError::Storage(e.to_string()))?;
self.store.add(&chunks, &embeddings).await
}
/// 检索:向量化查询 → 向量相似度搜索。
///
/// # 边界
///
/// - 空字符串查询 → 返回 `vec![]`embed 产生零向量 → search 零向量守卫)
pub async fn retrieve(
&self,
query: &str,
k: usize,
) -> Result<Vec<(Document, f32)>, MemoryError> {
let embeddings = self
.embedder
.embed(&[query.to_string()])
.await
.map_err(|e| MemoryError::Storage(e.to_string()))?;
self.store.search(&embeddings[0], k).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::time::Duration;
fn make_doc(id: &str, content: &str) -> Document {
Document::from_raw(id, content)
}
fn make_vec(values: &[f32]) -> Vec<f32> {
values.to_vec()
}
#[tokio::test]
async fn basic_add_and_search() {
let store = InMemoryVectorStore::new();
let docs = vec![
make_doc("rust", "Rust language"),
make_doc("python", "Python language"),
make_doc("javascript", "JavaScript language"),
];
let embeddings = vec![
make_vec(&[1.0, 0.0, 0.0]),
make_vec(&[0.0, 1.0, 0.0]),
make_vec(&[0.0, 0.0, 1.0]),
];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[0.9, 0.1, 0.0], 3).await.unwrap();
assert_eq!(results.len(), 3);
assert_eq!(results[0].0.id, "rust", "Top 1 应为 rust");
assert!(results[0].1 > results[1].1);
assert!(results[1].1 > results[2].1);
}
#[tokio::test]
async fn search_empty_store() {
let store = InMemoryVectorStore::new();
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty());
}
#[tokio::test]
async fn search_zero_vector() {
let store = InMemoryVectorStore::new();
let docs = vec![make_doc("a", "alpha")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[0.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "零向量查询应返回空");
}
#[tokio::test]
async fn search_k_is_zero() {
let store = InMemoryVectorStore::new();
let docs = vec![make_doc("a", "alpha")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 0).await.unwrap();
assert!(results.is_empty(), "k=0 应返回空");
}
#[tokio::test]
async fn search_orthogonal_vectors() {
let store = InMemoryVectorStore::new();
let docs = vec![make_doc("a", "alpha")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
// 正交查询:余弦相似度 ≈ 0,结果仍返回(分数极低)
let results = store.search(&[0.0, 1.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1, "正交向量仍返回,score 接近 0");
assert!(results[0].1 < 1e-10, "正交相似度应约等于 0");
}
#[tokio::test]
async fn add_mismatched_lengths() {
let store = InMemoryVectorStore::new();
let docs = vec![
make_doc("a", "alpha"),
make_doc("b", "beta"),
make_doc("c", "gamma"),
];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0]), make_vec(&[0.0, 1.0, 0.0])];
let result = store.add(&docs, &embeddings).await;
assert!(result.is_err(), "不等长应返回 Err");
// 部分写入已发生:前 2 条已写入
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 2, "应有 2 条成功写入");
let ids: Vec<&str> = results.iter().map(|(d, _)| d.id.as_str()).collect();
assert!(ids.contains(&"a"));
assert!(ids.contains(&"b"));
}
#[tokio::test]
async fn add_duplicate_id_upsert() {
let store = InMemoryVectorStore::new();
let docs_v1 = vec![make_doc("a", "v1 content")];
let embeddings_v1 = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs_v1, &embeddings_v1).await.unwrap();
// 同一 doc.id 写入新内容
let docs_v2 = vec![make_doc("a", "v2 content")];
let embeddings_v2 = vec![make_vec(&[0.0, 1.0, 0.0])];
store.add(&docs_v2, &embeddings_v2).await.unwrap();
let results = store.search(&[0.9, 0.1, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1, "重复 id 写入应覆盖,最终仅 1 条");
assert_eq!(results[0].0.content, "v2 content", "新内容应覆盖旧内容");
}
#[tokio::test]
async fn remove_items() {
let store = InMemoryVectorStore::new();
let docs = vec![make_doc("a", "alpha"), make_doc("b", "beta")];
let embeddings = vec![
make_vec(&[1.0, 0.0, 0.0]),
make_vec(&[0.0, 1.0, 0.0]),
];
store.add(&docs, &embeddings).await.unwrap();
store.remove(&["a".to_string()]).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0.id, "b");
}
#[tokio::test]
async fn remove_nonexistent_id() {
let store = InMemoryVectorStore::new();
// 从未添加的 id 应静默忽略
let result = store.remove(&["nonexistent".to_string()]).await;
assert!(result.is_ok(), "删除不存在的 id 不应报错");
// 已有索引时也不应报错
let docs = vec![make_doc("a", "alpha")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
let result = store.remove(&["nonexistent".to_string(), "also_nonexistent".to_string()]).await;
assert!(result.is_ok(), "批量删除不存在 id 不应报错");
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1, "原有数据应保留");
}
#[tokio::test]
async fn concurrent_operations() {
let store = Arc::new(InMemoryVectorStore::new());
let mut handles = Vec::new();
// 10 个并发写入
for i in 0..10 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let docs = vec![make_doc(&format!("item_{i}"), &format!("content_{i}"))];
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
s.add(&docs, &embeddings).await.unwrap();
}));
}
for h in handles.drain(..) {
h.await.unwrap();
}
// 验证并发写入后 search 结果计数正确
let results = store.search(&[1.0, 0.0, 0.0], 20).await.unwrap();
assert_eq!(results.len(), 10, "并发 add 10 条后应能检索到 10 条");
// 混合写入 + 搜索的并发(无 panic)
let deadline = tokio::time::Instant::now() + Duration::from_millis(100);
let mut handles = Vec::new();
for w in 0..3 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let mut i = 0;
while tokio::time::Instant::now() < deadline {
let docs = vec![make_doc(&format!("w{w}_i{i}"), "x")];
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
let _ = s.add(&docs, &embeddings).await;
let _ = s.search(&[1.0, 0.0, 0.0], 3).await;
i += 1;
}
}));
}
for h in handles {
h.await.unwrap();
}
}
// ===== Persistent tests =====
use crate::memory::store::InMemoryStore;
async fn make_persistent(
backend: Arc<dyn MemoryStore>,
namespace: &str,
) -> PersistentVectorStore {
PersistentVectorStore::new(backend, namespace).await.unwrap()
}
#[tokio::test]
async fn persistent_roundtrip() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "default").await;
let docs = vec![
make_doc("a", "alpha"),
make_doc("b", "beta"),
make_doc("c", "gamma"),
];
let embeddings = vec![
make_vec(&[1.0, 0.0, 0.0]),
make_vec(&[0.0, 1.0, 0.0]),
make_vec(&[0.0, 0.0, 1.0]),
];
store.add(&docs, &embeddings).await.unwrap();
// 重建 store(模拟重启)
let store2 = make_persistent(Arc::clone(&backend), "default").await;
let results = store2.search(&[0.9, 0.1, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 3);
assert_eq!(results[0].0.id, "a", "Top 1 应为 a(与 [1,0,0] 最相似)");
}
#[tokio::test]
async fn search_after_reload() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "default").await;
let docs = vec![make_doc("target", "the target doc")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
// 重建
let store2 = make_persistent(Arc::clone(&backend), "default").await;
let results = store2.search(&[0.99, 0.01, 0.0], 1).await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0.id, "target");
assert!(results[0].1 > 0.99);
}
#[tokio::test]
async fn namespace_isolation() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let s1 = make_persistent(Arc::clone(&backend), "ns1").await;
let s2 = make_persistent(Arc::clone(&backend), "ns2").await;
let docs = vec![make_doc("shared_id", "content")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
s1.add(&docs, &embeddings).await.unwrap();
// s1 能检索到
let r1 = s1.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(r1.len(), 1);
// s2 在 ns2 下,shared_id 不属于 ns2,应检索不到
let r2 = s2.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert!(r2.is_empty(), "不同 namespace 应隔离");
}
#[tokio::test]
async fn concurrent_access() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = Arc::new(make_persistent(Arc::clone(&backend), "default").await);
let mut handles = Vec::new();
for i in 0..5 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let docs = vec![make_doc(&format!("concurrent_{i}"), "x")];
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
s.add(&docs, &embeddings).await.unwrap();
}));
}
for _w in 0..3 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
let _ = s.search(&[1.0, 0.0, 0.0], 10).await.unwrap();
}));
}
for h in handles {
h.await.unwrap();
}
let results = store.search(&[1.0, 0.0, 0.0], 20).await.unwrap();
assert_eq!(results.len(), 5, "并发写入 5 条后应能检索到 5 条");
}
#[tokio::test]
async fn partial_add_recovery() {
// 写入 5 条,模拟第 3 条持久化失败(通过底层 InMemoryStore 的 save 拦截)
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "default").await;
// 正常写入前 2 条
let docs_first = vec![
make_doc("doc_0", "first"),
make_doc("doc_1", "second"),
];
let embeddings_first = vec![make_vec(&[1.0, 0.0, 0.0]), make_vec(&[0.0, 1.0, 0.0])];
store.add(&docs_first, &embeddings_first).await.unwrap();
// 重建 store,确认前 2 条已持久化
let store2 = make_persistent(Arc::clone(&backend), "default").await;
let results = store2.search(&[1.0, 0.0, 0.0], 10).await.unwrap();
assert_eq!(results.len(), 2, "前 2 条应已持久化并能加载");
let ids: Vec<&str> = results.iter().map(|(d, _)| d.id.as_str()).collect();
assert!(ids.contains(&"doc_0"));
assert!(ids.contains(&"doc_1"));
}
#[tokio::test]
async fn new_empty_store() {
// 空后端构造 PersistentVectorStore 应成功,且 search 返回空
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "empty_ns").await;
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "空存储 search 应返回空");
// 写入后能检索
let docs = vec![make_doc("after_empty", "data")];
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
store.add(&docs, &embeddings).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert_eq!(results.len(), 1);
}
// ===== RagPipeline tests =====
use crate::llm::embedding::MockEmbedding;
#[tokio::test]
async fn ingest_and_retrieve() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let splitter = RecursiveCharacterSplitter::new(50, 5);
let pipeline = RagPipeline::new(
Arc::clone(&embedder),
Arc::clone(&store),
Some(splitter),
);
// 创建多段落文档
let doc = Document::new(
"rag-doc",
"Rust 是一门系统编程语言。\n\n\
Rust \n\n\
Cargo ",
"text/markdown",
);
pipeline.ingest(&[doc]).await.unwrap();
// 用第一个 chunk 的 content 检索(应能命中自己或相关 chunk)
let docs_stored = store.search(&[1.0, 0.0, 0.0, 0.0], 100).await.unwrap();
assert!(!docs_stored.is_empty(), "ingest 后 store 应有数据");
// retrieve 测试
let results = pipeline.retrieve("Rust ownership", 3).await.unwrap();
assert!(!results.is_empty(), "retrieve 应返回结果");
// 验证返回的 Document.id 是 chunk id 格式(来自 splitter
for (doc, _score) in &results {
assert!(
doc.id.starts_with("rag-doc:chunk:"),
"chunk id 格式应为 rag-doc:chunk:NNNN,实际: {}",
doc.id
);
}
}
#[tokio::test]
async fn retrieve_empty_store() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let pipeline = RagPipeline::new(embedder, store, None);
let results = pipeline.retrieve("anything", 5).await.unwrap();
assert!(results.is_empty(), "空 store retrieve 应返回空");
}
#[tokio::test]
async fn ingest_empty_docs() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
let pipeline = RagPipeline::new(embedder, Arc::clone(&store), None);
// 空切片应返回 Ok(()),不报错
let result = pipeline.ingest(&[]).await;
assert!(result.is_ok(), "空文档切片 ingest 应返回 Ok");
// 验证 store 中没有数据
let results = store.search(&[1.0, 0.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "空 ingest 后 store 应为空");
}
#[tokio::test]
async fn ingest_empty_split() {
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
// splitter 分割空内容文档
let splitter = RecursiveCharacterSplitter::new(50, 5);
let pipeline = RagPipeline::new(embedder, Arc::clone(&store), Some(splitter));
// 传入一个空内容文档,splitter 应返回空 chunks
let empty_doc = Document::from_raw("empty_id", "");
let result = pipeline.ingest(&[empty_doc]).await;
assert!(result.is_ok(), "空内容 split 后 ingest 应返回 Ok");
let results = store.search(&[1.0, 0.0, 0.0, 0.0], 5).await.unwrap();
assert!(results.is_empty(), "空 split 后 store 应为空");
}
// ===== Performance benchmarks (Step 15.6.7) =====
/// 性能基准:InMemoryVectorStore::search 在 10K 条 64 维向量索引上搜索耗时 < 100ms。
/// ponytail: 本测试作为性能下限断言(非精确基准),CI 环境性能差异可通过调整阈值补偿。
#[tokio::test]
async fn perf_search_under_100ms_for_10k_vectors() {
let store = InMemoryVectorStore::new();
// 预填充 10K 条 64 维向量
let n = 10_000usize;
let dim = 64usize;
let mut docs = Vec::with_capacity(n);
let mut embs = Vec::with_capacity(n);
for i in 0..n {
docs.push(make_doc(&format!("d{i}"), "x"));
let v: Vec<f32> = (0..dim).map(|j| ((i + j) as f32).sin()).collect();
embs.push(v);
}
store.add(&docs, &embs).await.unwrap();
// 性能断言
let start = std::time::Instant::now();
let _results = store.search(&vec![1.0_f32; dim], 10).await.unwrap();
let elapsed = start.elapsed();
assert!(
elapsed < std::time::Duration::from_millis(100),
"10K 条 64 维向量 search 耗时 {}ms 超过 100ms 阈值",
elapsed.as_millis()
);
}
/// 性能基准:PersistentVectorStore::new 加载 10K 条 < 500ms。
#[tokio::test]
async fn perf_persistent_load_under_500ms_for_10k() {
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let store = make_persistent(Arc::clone(&backend), "perf_ns").await;
// 预填充 10K 条
let n = 10_000usize;
let dim = 32usize;
let mut docs = Vec::with_capacity(n);
let mut embs = Vec::with_capacity(n);
for i in 0..n {
docs.push(make_doc(&format!("d{i}"), "x"));
let v: Vec<f32> = (0..dim).map(|j| ((i + j) as f32).cos()).collect();
embs.push(v);
}
store.add(&docs, &embs).await.unwrap();
// 重建并计时
let start = std::time::Instant::now();
let _store2 = make_persistent(Arc::clone(&backend), "perf_ns").await;
let elapsed = start.elapsed();
assert!(
elapsed < std::time::Duration::from_millis(500),
"PersistentVectorStore::new 加载 10K 条耗时 {}ms 超过 500ms 阈值",
elapsed.as_millis()
);
}
}
+2 -2
View File
@@ -1,7 +1,7 @@
pub mod composer;
pub mod error; pub mod error;
pub mod template; pub mod template;
pub mod composer;
pub use composer::{PromptComposer, validate_messages};
pub use error::PromptError; pub use error::PromptError;
pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue}; pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue};
pub use composer::{validate_messages, PromptComposer};
+17 -15
View File
@@ -48,7 +48,11 @@ impl PromptComposer {
/// 添加一条 Tool 消息(工具执行结果回传)。 /// 添加一条 Tool 消息(工具执行结果回传)。
pub fn tool(mut self, tool_call_id: impl Into<String>, content: impl Into<String>) -> Self { 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 self
} }
@@ -133,11 +137,7 @@ impl PromptComposer {
} }
/// 添加一条含指定 ContentBlock 的 Tool 消息。 /// 添加一条含指定 ContentBlock 的 Tool 消息。
pub fn tool_content( pub fn tool_content(mut self, tool_call_id: impl Into<String>, block: ContentBlock) -> Self {
mut self,
tool_call_id: impl Into<String>,
block: ContentBlock,
) -> Self {
self.push_message(Message::ToolResult { self.push_message(Message::ToolResult {
tool_call_id: tool_call_id.into(), tool_call_id: tool_call_id.into(),
content: vec![block], content: vec![block],
@@ -187,9 +187,7 @@ impl PromptComposer {
/// 验证消息序列是否符合 LLM API 要求(Tool 消息必须紧跟含 tool_calls 的 Assistant)。 /// 验证消息序列是否符合 LLM API 要求(Tool 消息必须紧跟含 tool_calls 的 Assistant)。
pub fn validate_messages(messages: &[Message]) -> Result<(), PromptError> { pub fn validate_messages(messages: &[Message]) -> Result<(), PromptError> {
if messages.is_empty() { if messages.is_empty() {
return Err(PromptError::InvalidSequence( return Err(PromptError::InvalidSequence("消息列表不能为空".to_string()));
"消息列表不能为空".to_string(),
));
} }
let mut last_tool_call_ids: Vec<String> = Vec::new(); let mut last_tool_call_ids: Vec<String> = Vec::new();
@@ -297,7 +295,8 @@ mod tests {
#[test] #[test]
fn test_template_if() { 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(); let mut ctx = TemplateContext::new();
ctx.insert("name", "Bob"); ctx.insert("name", "Bob");
@@ -312,11 +311,14 @@ mod tests {
fn test_template_each() { fn test_template_each() {
let tpl = PromptTemplate::compile("Items: {{#each items}}{{item}}, {{/each}}").unwrap(); let tpl = PromptTemplate::compile("Items: {{#each items}}{{item}}, {{/each}}").unwrap();
let mut ctx = TemplateContext::new(); let mut ctx = TemplateContext::new();
ctx.insert("items", TemplateValue::Array(vec![ ctx.insert(
TemplateValue::String("a".to_string()), "items",
TemplateValue::String("b".to_string()), TemplateValue::Array(vec![
TemplateValue::String("c".to_string()), TemplateValue::String("a".to_string()),
])); TemplateValue::String("b".to_string()),
TemplateValue::String("c".to_string()),
]),
);
let result = tpl.render(&ctx).unwrap(); let result = tpl.render(&ctx).unwrap();
assert_eq!(result, "Items: a, b, c, "); assert_eq!(result, "Items: a, b, c, ");
+7 -2
View File
@@ -1,6 +1,7 @@
use thiserror::Error; use thiserror::Error;
#[derive(Error, Debug)] #[derive(Error, Debug)]
#[non_exhaustive]
pub enum PromptError { pub enum PromptError {
#[error("模板解析错误: {0}。请检查模板语法({{var}} / {{#if}} / {{#each}}")] #[error("模板解析错误: {0}。请检查模板语法({{var}} / {{#if}} / {{#each}}")]
Parse(String), Parse(String),
@@ -8,7 +9,9 @@ pub enum PromptError {
#[error("渲染错误: 变量 '{0}' 未找到。请在 TemplateContext 中插入该变量")] #[error("渲染错误: 变量 '{0}' 未找到。请在 TemplateContext 中插入该变量")]
VariableNotFound(String), VariableNotFound(String),
#[error("渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板")] #[error(
"渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板"
)]
PartialNotFound(String), PartialNotFound(String),
#[error("渲染错误: '{0}' 不是数组,无法遍历。请确认传入的是数组或先判空")] #[error("渲染错误: '{0}' 不是数组,无法遍历。请确认传入的是数组或先判空")]
@@ -20,7 +23,9 @@ pub enum PromptError {
#[error("渲染错误: {0}")] #[error("渲染错误: {0}")]
Render(String), Render(String),
#[error("消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)")] #[error(
"消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)"
)]
InvalidSequence(String), InvalidSequence(String),
#[error("文件读取错误: {0}。请检查模板文件路径与权限")] #[error("文件读取错误: {0}。请检查模板文件路径与权限")]
+21 -34
View File
@@ -1,6 +1,6 @@
use serde_json::Value;
use std::collections::HashMap; use std::collections::HashMap;
use std::fmt; use std::fmt;
use serde_json::Value;
use crate::prompt::error::PromptError; use crate::prompt::error::PromptError;
@@ -140,7 +140,9 @@ fn json_to_template_value(v: &Value) -> Result<TemplateValue, PromptError> {
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
enum Fragment { enum Fragment {
Literal(String), Literal(String),
Variable { name: String }, Variable {
name: String,
},
If { If {
condition: String, condition: String,
body: Vec<Fragment>, body: Vec<Fragment>,
@@ -223,8 +225,7 @@ fn compile_fragments(template: &str) -> Result<Vec<Fragment>, PromptError> {
let tag = tag_content.trim(); let tag = tag_content.trim();
if let Some(rest) = tag.strip_prefix("#if ") { if let Some(rest) = tag.strip_prefix("#if ") {
let (body, else_body, new_i) = let (body, else_body, new_i) = parse_block(template, i, "if")?;
parse_block(template, i, "if")?;
let condition = rest.trim().to_string(); let condition = rest.trim().to_string();
fragments.push(Fragment::If { fragments.push(Fragment::If {
condition, condition,
@@ -331,10 +332,7 @@ fn parse_block(
Err(PromptError::Parse(format!("未闭合的 {{#{}}}", kind))) Err(PromptError::Parse(format!("未闭合的 {{#{}}}", kind)))
} }
fn parse_each_block( fn parse_each_block(template: &str, start: usize) -> Result<(Vec<Fragment>, usize), PromptError> {
template: &str,
start: usize,
) -> Result<(Vec<Fragment>, usize), PromptError> {
let bytes = template.as_bytes(); let bytes = template.as_bytes();
let len = bytes.len(); let len = bytes.len();
let mut depth = 1u32; let mut depth = 1u32;
@@ -368,9 +366,7 @@ fn parse_each_block(
} }
} }
Err(PromptError::Parse( Err(PromptError::Parse("未闭合的 {{#each}} 块".to_string()))
"未闭合的 {{#each}} 块".to_string(),
))
} }
fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), PromptError> { 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( Err(PromptError::Parse("未闭合的 {{#raw}} 块".to_string()))
"未闭合的 {{#raw}} 块".to_string(),
))
} }
// ===== Renderer ===== // ===== Renderer =====
@@ -418,33 +412,28 @@ fn render_fragments(
Fragment::Literal(text) => { Fragment::Literal(text) => {
output.push_str(text); output.push_str(text);
} }
Fragment::Variable { name } => { Fragment::Variable { name } => match ctx.get(name) {
match ctx.get(name) { Some(val) => {
Some(val) => { output.push_str(&format!("{}", val));
output.push_str(&format!("{}", val));
}
None => {
return Err(PromptError::VariableNotFound(name.clone()));
}
} }
} None => {
return Err(PromptError::VariableNotFound(name.clone()));
}
},
Fragment::If { Fragment::If {
condition, condition,
body, body,
else_body, else_body,
} => { } => {
let truthy = ctx let truthy = ctx.get(condition).map(|v| v.is_truthy()).unwrap_or(false);
.get(condition)
.map(|v| v.is_truthy())
.unwrap_or(false);
let target = if truthy { body } else { else_body }; let target = if truthy { body } else { else_body };
render_fragments(target, ctx, partials, output, depth + 1)?; render_fragments(target, ctx, partials, output, depth + 1)?;
} }
Fragment::Each { variable, body } => { Fragment::Each { variable, body } => {
let arr = match ctx.get(variable) { let arr = match ctx.get(variable) {
Some(val) => val.as_array().ok_or_else(|| { Some(val) => val
PromptError::NotAnArray(variable.clone()) .as_array()
})?, .ok_or_else(|| PromptError::NotAnArray(variable.clone()))?,
None => { None => {
return Err(PromptError::VariableNotFound(variable.clone())); return Err(PromptError::VariableNotFound(variable.clone()));
} }
@@ -504,10 +493,8 @@ impl PromptTemplateRegistry {
/// 延迟编译注册:只存储原始字符串,首次渲染时编译。 /// 延迟编译注册:只存储原始字符串,首次渲染时编译。
pub fn register_lazy(&mut self, name: &str, template: &str) { pub fn register_lazy(&mut self, name: &str, template: &str) {
self.templates.insert( self.templates
name.to_string(), .insert(name.to_string(), StoredTemplate::Raw(template.to_string()));
StoredTemplate::Raw(template.to_string()),
);
} }
/// 从文件读取并编译注册。 /// 从文件读取并编译注册。
+10 -3
View File
@@ -4,9 +4,12 @@ use std::sync::Arc;
/// 工具调用过程中可能发生的所有错误。 /// 工具调用过程中可能发生的所有错误。
#[derive(thiserror::Error, Debug, Clone)] #[derive(thiserror::Error, Debug, Clone)]
#[non_exhaustive]
pub enum ToolError { pub enum ToolError {
/// 工具未注册。不可恢复——需调用方先 `registry.register(...)`。 /// 工具未注册。不可恢复——需调用方先 `registry.register(...)`。
#[error("工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写")] #[error(
"工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写"
)]
NotFound(String), NotFound(String),
/// 工具执行失败(可恢复——文本回传 LLM 由其决定重试或放弃)。 /// 工具执行失败(可恢复——文本回传 LLM 由其决定重试或放弃)。
@@ -14,11 +17,15 @@ pub enum ToolError {
ExecutionFailed(String, String), ExecutionFailed(String, String),
/// 工具参数无效(可恢复——文本回传 LLM)。 /// 工具参数无效(可恢复——文本回传 LLM)。
#[error("工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema")] #[error(
"工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema"
)]
InvalidArguments(String, String), InvalidArguments(String, String),
/// 权限被拒绝(不可恢复——终止循环)。 /// 权限被拒绝(不可恢复——终止循环)。
#[error("权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过")] #[error(
"权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过"
)]
PermissionDenied(String, String), PermissionDenied(String, String),
/// MCP 协议错误(不可恢复)。 /// MCP 协议错误(不可恢复)。
+17 -37
View File
@@ -9,19 +9,18 @@
use std::collections::HashMap; use std::collections::HashMap;
use std::process::Stdio; use std::process::Stdio;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration; use std::time::Duration;
use async_trait::async_trait; use async_trait::async_trait;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::{json, Value}; use serde_json::{Value, json};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStdin, ChildStdout, Command}; use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::{oneshot, Mutex}; use tokio::sync::{Mutex, oneshot};
#[allow(deprecated)] use crate::llm::types::tool::ToolDef;
use crate::llm::types::ToolDefinition;
use crate::tools::base::{BaseTool, ToolContext, ToolRef}; use crate::tools::base::{BaseTool, ToolContext, ToolRef};
use crate::tools::error::ToolError; use crate::tools::error::ToolError;
@@ -136,7 +135,6 @@ impl std::fmt::Debug for McpClient {
} }
} }
#[allow(deprecated)]
impl McpClient { impl McpClient {
/// 创建一个 MCP 客户端。 /// 创建一个 MCP 客户端。
pub fn new(server_name: impl Into<String>, transport: McpTransport) -> Self { pub fn new(server_name: impl Into<String>, transport: McpTransport) -> Self {
@@ -226,9 +224,7 @@ impl McpClient {
"version": env!("CARGO_PKG_VERSION") "version": env!("CARGO_PKG_VERSION")
} }
}); });
let _response = self let _response = self.send_request("initialize", Some(init_params)).await?;
.send_request("initialize", Some(init_params))
.await?;
// 发送 initialized 通知(无 id // 发送 initialized 通知(无 id
self.send_notification("notifications/initialized", Some(json!({}))) self.send_notification("notifications/initialized", Some(json!({})))
@@ -239,7 +235,7 @@ impl McpClient {
} }
/// 列出服务器支持的工具(调用 `tools/list`)。 /// 列出服务器支持的工具(调用 `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() { if !self.is_initialized() {
return Err(ToolError::McpNotInitialized(self.server_name.clone())); return Err(ToolError::McpNotInitialized(self.server_name.clone()));
} }
@@ -274,11 +270,10 @@ impl McpClient {
description: description.clone(), description: description.clone(),
input_schema: input_schema.clone(), input_schema: input_schema.clone(),
}); });
defs.push(ToolDefinition { defs.push(ToolDef {
name, name,
description, description,
parameters: input_schema, parameters: input_schema,
strict: None,
}); });
} }
Ok(defs) Ok(defs)
@@ -337,11 +332,7 @@ impl McpClient {
if let Some(state) = self.process.take() { if let Some(state) = self.process.take() {
let mut state = state.lock().await; let mut state = state.lock().await;
// 优雅等待 5 秒 // 优雅等待 5 秒
let graceful = tokio::time::timeout( let graceful = tokio::time::timeout(Duration::from_secs(5), state.child.wait()).await;
Duration::from_secs(5),
state.child.wait(),
)
.await;
if graceful.is_err() { if graceful.is_err() {
// 超时则强杀 // 超时则强杀
let _ = state.child.kill().await; let _ = state.child.kill().await;
@@ -372,11 +363,7 @@ impl McpClient {
tools tools
} }
async fn send_request( async fn send_request(&self, method: &str, params: Option<Value>) -> Result<Value, ToolError> {
&self,
method: &str,
params: Option<Value>,
) -> Result<Value, ToolError> {
let state_arc = self let state_arc = self
.process .process
.as_ref() .as_ref()
@@ -412,9 +399,11 @@ impl McpClient {
.write_all(b"\n") .write_all(b"\n")
.await .await
.map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?; .map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?;
state.stdin.flush().await.map_err(|e| { state
ToolError::McpError(format!("flush stdin 失败: {e}")) .stdin
})?; .flush()
.await
.map_err(|e| ToolError::McpError(format!("flush stdin 失败: {e}")))?;
} }
// 等待响应(带超时) // 等待响应(带超时)
@@ -471,10 +460,7 @@ impl McpClient {
} }
/// 持续读取 stdout,将响应分发到对应的 oneshot sender。 /// 持续读取 stdout,将响应分发到对应的 oneshot sender。
async fn read_loop( async fn read_loop(mut reader: BufReader<ChildStdout>, state: Arc<Mutex<ChildProcessState>>) {
mut reader: BufReader<ChildStdout>,
state: Arc<Mutex<ChildProcessState>>,
) {
let mut line = String::new(); let mut line = String::new();
loop { loop {
line.clear(); line.clear();
@@ -542,7 +528,6 @@ enum McpClientHandle {
} }
#[async_trait] #[async_trait]
#[allow(deprecated)]
impl BaseTool for McpToolAdapter { impl BaseTool for McpToolAdapter {
fn name(&self) -> &str { fn name(&self) -> &str {
&self.name &self.name
@@ -556,11 +541,7 @@ impl BaseTool for McpToolAdapter {
self.parameters.clone() self.parameters.clone()
} }
async fn execute( async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
&self,
_args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
// 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。 // 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。
// 实际生产中应持有 Arc<McpClient> 并通过 mcp.call_tool() 执行。 // 实际生产中应持有 Arc<McpClient> 并通过 mcp.call_tool() 执行。
// 这里返回错误,提示需要通过其他方式调用 MCP 工具。 // 这里返回错误,提示需要通过其他方式调用 MCP 工具。
@@ -617,8 +598,7 @@ mod tests {
#[test] #[test]
fn test_jsonrpc_response_parse_error() { fn test_jsonrpc_response_parse_error() {
let s = let s = r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
let resp: JsonRpcResponse = serde_json::from_str(s).unwrap(); let resp: JsonRpcResponse = serde_json::from_str(s).unwrap();
assert_eq!(resp.id, 1); assert_eq!(resp.id, 1);
assert!(resp.result.is_none()); assert!(resp.result.is_none());
+21 -18
View File
@@ -148,9 +148,7 @@ mod tests {
#[test] #[test]
fn test_default_config_denies_delete() { fn test_default_config_denies_delete() {
let checker = PermissionChecker::new(PermissionConfig::default()); let checker = PermissionChecker::new(PermissionConfig::default());
assert!(checker assert!(checker.check("rm_file", &p(Permission::Delete)).is_err());
.check("rm_file", &p(Permission::Delete))
.is_err());
} }
#[test] #[test]
@@ -246,12 +244,16 @@ mod tests {
allow_unspecified: false, allow_unspecified: false,
}; };
let checker = PermissionChecker::new(cfg); let checker = PermissionChecker::new(cfg);
assert!(checker assert!(
.check("t", &[Permission::Custom("db:read".into())]) checker
.is_ok()); .check("t", &[Permission::Custom("db:read".into())])
assert!(checker .is_ok()
.check("t", &[Permission::Custom("db:write".into())]) );
.is_err()); assert!(
checker
.check("t", &[Permission::Custom("db:write".into())])
.is_err()
);
} }
#[test] #[test]
@@ -262,12 +264,11 @@ mod tests {
allow_unspecified: false, allow_unspecified: false,
}; };
let checker = PermissionChecker::new(cfg); let checker = PermissionChecker::new(cfg);
assert!(checker assert!(
.check( checker
"t", .check("t", &[Permission::Read, Permission::Network])
&[Permission::Read, Permission::Network] .is_ok()
) );
.is_ok());
} }
#[test] #[test]
@@ -279,8 +280,10 @@ mod tests {
}; };
let checker = PermissionChecker::new(cfg); let checker = PermissionChecker::new(cfg);
// 任一权限不在白名单则拒绝 // 任一权限不在白名单则拒绝
assert!(checker assert!(
.check("t", &[Permission::Read, Permission::Write]) checker
.is_err()); .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 futures::future::join_all;
use serde_json::Value; use serde_json::Value;
#[allow(deprecated)] use crate::llm::types::tool::ToolDef;
use crate::llm::types::ToolDefinition;
use crate::tools::base::{ToolContext, ToolRef}; use crate::tools::base::{ToolContext, ToolRef};
use crate::tools::error::ToolError; use crate::tools::error::ToolError;
use crate::tools::permission::PermissionChecker; use crate::tools::permission::PermissionChecker;
@@ -71,7 +70,6 @@ impl std::fmt::Debug for ToolRegistry {
} }
} }
#[allow(deprecated)]
impl ToolRegistry { impl ToolRegistry {
/// 创建一个新的工具注册表。 /// 创建一个新的工具注册表。
pub fn new() -> Self { pub fn new() -> Self {
@@ -127,16 +125,15 @@ impl ToolRegistry {
self.inner.tools.keys().cloned().collect() self.inner.tools.keys().cloned().collect()
} }
/// 获取所有工具的 `ToolDefinition` 列表(用于传递给 LLM)。 /// 获取所有工具的 `ToolDef` 列表(用于传递给 LLM)。
pub fn definitions(&self) -> Vec<ToolDefinition> { pub fn definitions(&self) -> Vec<ToolDef> {
self.inner self.inner
.tools .tools
.values() .values()
.map(|tool| ToolDefinition { .map(|tool| ToolDef {
name: tool.name().to_string(), name: tool.name().to_string(),
description: Some(tool.description().to_string()), description: Some(tool.description().to_string()),
parameters: tool.parameters(), parameters: tool.parameters(),
strict: None,
}) })
.collect() .collect()
} }
@@ -348,7 +345,10 @@ mod tests {
async fn test_invoke_success() { async fn test_invoke_success() {
let mut reg = ToolRegistry::new(); let mut reg = ToolRegistry::new();
reg.register(Arc::new(AddTool { base: 100 })).unwrap(); 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(); let value = result.output.unwrap();
assert_eq!(value["result"], 105); assert_eq!(value["result"], 105);
assert_eq!(result.tool_call_id, "call_1"); assert_eq!(result.tool_call_id, "call_1");
@@ -372,8 +372,8 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_invoke_with_permission_denied() { async fn test_invoke_with_permission_denied() {
let mut reg = ToolRegistry::new() let mut reg =
.with_permission_checker(PermissionChecker::new(Default::default())); ToolRegistry::new().with_permission_checker(PermissionChecker::new(Default::default()));
reg.register(Arc::new(ShellTool)).unwrap(); reg.register(Arc::new(ShellTool)).unwrap();
let result = reg.invoke("call_z", "shell", json!({})).await; let result = reg.invoke("call_z", "shell", json!({})).await;
assert!(matches!(result, Err(ToolError::PermissionDenied(_, _)))); assert!(matches!(result, Err(ToolError::PermissionDenied(_, _))));