22 Commits
Author SHA1 Message Date
徐涛 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
63 changed files with 3924 additions and 955 deletions
+12
View File
@@ -198,6 +198,18 @@ pub use vector_store::VectorStore;
5. **风险评估** - 潜在风险、缓解措施
6. **验收标准** - 可验证的完成条件
### 进度同步规范 (docs/roadmap.md)
完成一项实施后,必须检查 `docs/roadmap.md` 是否存在对应内容;若存在,必须同步标记为完成:
- **Step / Phase 状态行**:对应 Step 加 ✅ 标记;Phase 章节末尾「状态」行从 ⏳ 改为 ✅ Phase X 全部交付物已完成
- **里程碑表**:更新对应里程碑状态从 ⏳ 改为 ✅ + 完成日期
- **依赖关系图(Mermaid**:节点 `class``pending` / `core` 改为 `done`,必要时更新节点摘要
- **文末「已完成 / 进行中阶段」列表**:追加一行 `- ✅ Phase X — 一句话要点`
- **顶部「当前状态」**:补充新完成 Phase,更新「下一步」指向
参考案例:2026-07-05 完成 Phase 7 SqliteStore 时同步更新 6 处(顶部状态 / Phase 章节 / 依赖图 / M3 / 下一步行动 / 已完成列表)。
---
## 项目特定规则
+61
View File
@@ -2,6 +2,67 @@
本项目所有重要变更均记录于此文件。格式参考 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/)。
## [0.2.0-rc.1] - 2026-07-05
v0.2.0 候选发布。Phase 5-7 三大 P0 全部交付完成,API 稳定性扫尾,新增 2 个面向新用户的集成示例。
### Added
**Phase 5 — 热身准备**
- `ProviderConfig::from_env(prefix)`:从 `{prefix}_BASE_URL` / `{prefix}_API_KEY` / `{prefix}_MODEL` / `{prefix}_TIMEOUT_SECS` / `{prefix}_MAX_RETRIES` 环境变量构造配置
- `ProviderConfig::timeout_secs` / `max_retries` 字段(默认 30 / 3
- `OllamaProvider`:本地推理 ProviderOpenAI-compatible`http://localhost:11434/v1` 默认端点)
- `ProviderType::Ollama` 变体 + `FromStr` 解析
**Phase 6 — ToolDef IR 正式化**
- `ToolDef` 结构体(name / description / parameters)替代已废弃的 `OpenaiToolDefinition`
- `MessageRequest.tools` 切换为 `Vec<ToolDef>`
- `OpenaiToolDefinition` 降级为 `#[doc(hidden)]`,仅供 OpenAI 适配层内部消费
**Phase 7 — SqliteStore 持久化**
- `SqliteStore``MemoryStore` 的 SQLite 后端实现,基于 `rusqlite 0.32` bundled
- WAL 模式 + `synchronous=NORMAL` + `busy_timeout=5s` 兼顾崩溃安全与吞吐
- `Arc<Mutex<Connection>>` + `spawn_blocking` 卸载 IO10×10 并发写入无 race
- `MemoryStore::save / get / delete / list` CRUD + prefix / since / offset+limit 过滤
- 进程重启数据不丢的 round-trip 验证
**Phase 8 — MVP 集成出口**
- `examples/quick_start`:30 行最小可运行示例(MockProvider + EchoTool + submit_turn
- `examples/end_to_end`3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证
### Changed
- **API 稳定性护栏**:14 个公开枚举标记 `#[non_exhaustive]`,覆盖 P0 核心 IR`Message` / `ContentBlock` / `ContentBlockType` / `StreamEvent` / `HookEvent`)、P0 Error`AgentError` / `LlmError` / `ToolError` / `MemoryError` / `PromptError`)、P1 其他(`MemoryStrategy` / `StepStatus` / `ToolChoice` / `ResponseFormat`
- **`StepStatus::Completed` 字段类型**:从废弃的 `ChatResponse` 切换为 IR 层 `MessageResponse`(同时清理 `task_agent_demo.rs``ChatResponse` / `OpenaiChatMessage` / `FinishReason` 三处废弃类型引用)
### Non-exhaustive 清单
为防止未来新增变体时下游 exhaustive match 静默失效,14 个枚举追加 `#[non_exhaustive]`
| 优先级 | 枚举 |
|--------|------|
| P0 核心 IR | `Message`, `ContentBlock`, `ContentBlockType`, `StreamEvent`, `HookEvent` |
| P0 Error | `AgentError`, `LlmError`, `ToolError`, `MemoryError`, `PromptError` |
| P1 其他 | `MemoryStrategy`, `StepStatus`, `ToolChoice`, `ResponseFormat` |
明确不加:内部 wire-format`OpenaiChatMessage` 等)/ 语义已收敛(`Role` / `ServiceTier` / `Modality` / `ImageDetail` / `AudioFormat` / `StopSequence`/ 使用面窄(`Permission` / `McpTransport` 等)。
### Deprecated
(继承自 0.1.0,无新增)`ChatResponse` / `ToolDefinition` 保持 `#[deprecated]` 标记。
### Fixed
- 修复 `StepStatus::Completed(ChatResponse)` 字段类型与 IR 体系不一致问题(已完成迁移)
### Migration Guide (v0.1 → v0.2.0-rc.1)
1. **枚举 match**14 个 `#[non_exhaustive]` 枚举在 crate 外必须使用 `_ =>` 通配分支
2. **`StepStatus::Completed`**:字段类型从 `ChatResponse` 切换为 `MessageResponse`,需做字段映射(参考 `docs/15-phase8-mvp-integration.md` §3.1.2
3. **`ToolDefinition``ToolDef`**Phase 6 已彻底替换 `#[deprecated]` 别名,需全局重命名
---
## [0.1.0] - 2026-07-04
首个公开版本。涵盖 Phase 0-4c 的全部核心能力、Provider IR 重构、LlmCycle 简化,以及面向用户的 7 个离线示例。
+5 -2
View File
@@ -1,6 +1,6 @@
[package]
name = "agcore"
version = "0.1.0"
version = "0.2.0-rc.1"
edition = "2024"
[dependencies]
@@ -19,8 +19,11 @@ futures-core = "0.3"
bytes = "1"
async-stream = "0.3"
tokio-util = { version = "0.7", features = ["rt"] }
time = { version = "0.3", features = ["serde"] }
time = { version = "0.3", features = ["serde", "parsing", "formatting", "macros"] }
rusqlite = { version = "0.32", features = ["bundled"] }
[dev-dependencies]
dotenvy = "0.15.7"
wiremock = "0.6"
temp-env = "0.3"
tempfile = "3"
+5 -2
View File
@@ -26,7 +26,7 @@ AG Core 不是 Agent 产品,而是 Agent 的**底层依赖库**:上层应用
```toml
[dependencies]
agcore = "0.1"
agcore = "0.2"
tokio = { version = "1", features = ["macros", "rt-multi-thread"] }
```
@@ -110,10 +110,12 @@ let provider = create_provider(
).expect("创建 Provider 失败");
```
更多端到端示例见 [`examples/`](./examples/) 目录(共 7 个,全部可 `cargo run --example <name>`):
更多端到端示例见 [`examples/`](./examples/) 目录(共 10 个,全部可 `cargo run --example <name>`):
| 示例 | 说明 |
|------|------|
| `quick_start` | **30 行最小示例**MockProvider + EchoTool + submit_turn,新用户 5 分钟上手 |
| `end_to_end` | **完整集成示例**3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证 |
| `agent_session_demo` | Agent + 会话 + SessionMemory 完整链路(MockProvider 离线) |
| `custom_tool` | 自定义工具注册、单次 / 并行调用、权限检查 |
| `prompt_composer` | 提示词模板与组合器(纯离线) |
@@ -121,6 +123,7 @@ let provider = create_provider(
| `conversation_memory_demo` | 对话记忆滑动窗口与隔离 |
| `knowledge_search_demo` | 知识页面关键词检索 |
| `streaming_events_demo` | LLM 流式响应事件消费(含错误路径) |
| `simple_visit` | 真实 LLM 调用(OpenAI / Anthropic,设置 `OPENAI_*` / `ANTHROPIC_*` 环境变量) |
## 核心模块
+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 个示例 |
+88 -36
View File
@@ -1,13 +1,13 @@
# AG Core Roadmap
> 定稿日期:2026-05-11
> 最后更新:2026-07-04
> 最后更新:2026-07-05
## 愿景
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
**当前状态**v0.1.0 已发布(2026-07-04)。Phase 0-4c 全部完成,Provider IR 重构 + LlmCycle 简化 + 7 个离线示例已交付。v0.2.0 已细分为 8 个增量 PhasePhase 5-12),本周开发启动
**当前状态**v0.1.0 已发布(2026-07-04)。Phase 0-8 全部完成,v0.2.0-rc.1 已打标签。Provider IR 重构 + LlmCycle 简化 + 10 个离线示例(含 `quick_start` 30 行最小示例与 `end_to_end` 完整集成示例)+ SqliteStore 持久化 + 14 个公开枚举 `#[non_exhaustive]` 护栏 + `StepStatus` IR 迁移已交付。下一步进入 Phase 9(流式体验增强 / `submit_turn_stream`
---
@@ -334,13 +334,23 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
| 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 警告 |
| **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 全部交付物已完成
---
@@ -352,11 +362,11 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
| 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` 全绿 |
| **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:新类型存在但旧代码照常编译
@@ -366,6 +376,19 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
**依赖**:无(仅与 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 持久化
@@ -376,8 +399,8 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
| 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 |
| **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)
@@ -386,6 +409,22 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
**依赖**`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 候选)
@@ -394,14 +433,22 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
| Step | 内容 | 验证标准 |
|------|------|---------|
| **8.1** | API 稳定性扫尾:`#[deprecated]` 整理 + CHANGELOG v0.2 + 公开类型回顾 | 人工 review + `cargo doc` warning |
| **8.2** | Quick Start 示例(30`main.rs`):MockProvider + EchoTool + 一次 `submit_turn` | `cargo run --example quick_start` exit 0 |
| **8.3** | 端到端示例:SqliteStore + Ollama/OpenAI(from_env) + 自定义 Tool + 轮对话 | `cargo run --example end_to_end`Mock fallback,无需 API key|
| **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` 标签**
**Phase 8 全部完成**。**已`v0.2.0-rc.1` 标签**。
**实际新增**2026-07-056 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`57 行)+ `end_to_end.rs`237 行)
**依赖**Phase 5ProviderConfig from_env+ Phase 6ToolDef+ Phase 7SqliteStore
**优先级**P0
**状态**:✅ Phase 8 全部交付物已完成
---
@@ -470,10 +517,10 @@ pub struct ContextBudget { system, history, tools, tool_results, reserve }
```mermaid
graph BT
P5["Phase 5<br/>热身准备"]:::warmup
P6["Phase 6<br/>ToolDefinition IR"]:::core
P7["Phase 7<br/>SqliteStore"]:::core
P8["Phase 8<br/>MVP 出口 (rc.1)"]:::mvp
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["Phase 9<br/>流式体验增强"]:::p1
P10["Phase 10<br/>ContextSlot"]:::p1
P11["Phase 11<br/>测试与检索"]:::p1
@@ -490,6 +537,7 @@ graph BT
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
@@ -501,16 +549,16 @@ graph BT
### 关键里程碑
| 里程碑 | Phase 完成条件 | 可验证指标 |
|--------|---------------|-----------|
| **M1** | Phase 5 | 热身三项完成:`from_env()` 可用 / Ollama 类型存在 / `#[non_exhaustive]` 就位 |
| **M2** | Phase 6 | `ToolDef` 全量切换,`cargo test --all-targets` 全绿 |
| **M3** | Phase 7 | SqliteStore CRUD + 并发测试通过,进程重启数据不丢 |
| **M4** | **Phase 8 (rc.1)** | P0 五项全部交付,`cargo run --example quick_start` 跑通 |
| **M5** | Phase 9 | `submit_turn_stream` 流式事件序列验证通过 |
| **M6** | Phase 10 | ContextSlot 创建/切换/派生集成测试通过 |
| **M7** | Phase 11 | wiremock + 并发测试补强,测试总量 200+ |
| **M8** | Phase 12(可选) | P2 功能按需交付 |
| 里程碑 | 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` 流式事件序列验证通过 | ⏳ |
| **M6** | Phase 10 | ContextSlot 创建/切换/派生集成测试通过 | ⏳ |
| **M7** | Phase 11 | wiremock + 并发测试补强,测试总量 200+ | ⏳ |
| **M8** | Phase 12(可选) | P2 功能按需交付 | ⏳ |
---
@@ -553,10 +601,10 @@ graph BT
## 下一步行动
1. **Phase 5 启动**ProviderConfig from_env + Ollama Provider + #[non_exhaustive] 前置,三个 Step 并行推进
2. **Phase 6 方案准备**:ToolDef 结构体定义 + 兼容转换,出实施笔记(实施时直接走代码评审)
3. **示例先行**每完成一个 Phase 立即更新对应示例,验证通过后再合入
4. **里程碑追踪**Phase 8MVP 出口)为 v0.2.0-rc.1 节点,逐 Phase 验收
1. **Phase 9 启动**流式体验增强(`AgentSession::submit_turn_stream` 流式事件序列),P1 功能
2. **示例先行**:每完成一个 Phase 立即更新对应示例,验证通过后再合入
3. **里程碑追踪** Phase 8MVP 出口,v0.2.0-rc.1 已打标签)为节点,逐 Phase 验收
4. **v0.2.0 正式版**Phase 8-11 全部完成后,去掉 rc 后缀打 `v0.2.0` 正式版
**已完成 / 进行中阶段**
- ✅ Phase 0 Foundation — 全部交付物已完成
@@ -566,9 +614,13 @@ graph BT
- ✅ Phase 4a Core Glue — 全部交付物已完成
- ✅ Phase 4b Task Execution — 全部交付物已完成
- ✅ Phase 4c Session Memory — 全部交付物已完成
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen 适配
- ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy
- ✅ Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
- ✅ Phase 7 SqliteStore — `rusqlite 0.32` + WAL 模式 + `Arc<Mutex<Connection>>` + `spawn_blocking``memory/store.rs``store/{in_memory,sqlite_store}.rs` 模块化;9 个内联测试覆盖 CRUD/upsert/过滤/10×10 并发/持久化 round-trip`InMemoryStore ↔ SqliteStore` trait-box 互换兼容
-**Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)``Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 57 行 + `end_to_end` 237 行),10 个离线示例全部 exit 0**v0.2.0-rc.1 标签已打**
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen/Ollama 适配
- ✅ LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、7 个离线示例、README + 错误消息友好化、CHANGELOG 初始化
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、8 个离线示例(含 `simple_visit`、README + 错误消息友好化、CHANGELOG 初始化
- 📋 **v0.2 规划细化完成** — 8 个增量 PhasePhase 5-12),17 个可验证 Step,覆盖 P0-P2 全部 12 项功能 + ContextSlot
---
+7 -8
View File
@@ -15,9 +15,9 @@ use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider;
use agcore::llm::types::Usage;
use agcore::llm::types::message::{ContentBlock, Message};
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
use agcore::llm::types::Usage;
use agcore::tools::ToolRegistry;
/// 计算器角色 Agent。
@@ -72,7 +72,10 @@ async fn main() {
// 4. 提交第一轮
println!("=== 提交第 1 轮 ===");
let resp = session.submit_turn("1+1=?").await.expect("submit_turn 失败");
let resp = session
.submit_turn("1+1=?")
.await
.expect("submit_turn 失败");
println!("LLM: {}", resp.text());
session
.set_session_data("last_q", "1+1=?")
@@ -107,11 +110,7 @@ async fn main() {
// 8. 跨 session 数据隔离验证
println!("=== 数据隔离验证 ===");
let other = AgentSession::new(
Arc::new(CalculatorAgent),
"other-session",
bundle,
);
let other = AgentSession::new(Arc::new(CalculatorAgent), "other-session", bundle);
assert!(
other.get_session_data("last_q").await.unwrap().is_none(),
"新会话不应看到旧 session 的 last_q"
@@ -127,4 +126,4 @@ async fn main() {
);
println!("\n✓ agent_session_demo 完成");
}
}
+9 -17
View File
@@ -29,6 +29,7 @@ fn message_text(msg: &Message) -> &str {
.next()
.unwrap_or(""),
Message::UserImage { .. } => "[image]",
_ => "",
}
}
@@ -80,11 +81,8 @@ async fn main() {
// 3. 多角色混合 + clear
println!("\n=== 多角色写入 + clear ===");
let store3 = Arc::new(InMemoryStore::new());
let mut memory3 = ConversationMemory::new(
store3,
"session-3",
ConversationMemoryConfig::default(),
);
let mut memory3 =
ConversationMemory::new(store3, "session-3", ConversationMemoryConfig::default());
memory3
.add_message(Message::user_text("你好"))
.await
@@ -98,7 +96,9 @@ async fn main() {
.await
.unwrap();
memory3
.add_message(Message::assistant("我无法查询实时天气,但你可以查看天气应用。"))
.add_message(Message::assistant(
"我无法查询实时天气,但你可以查看天气应用。",
))
.await
.unwrap();
println!(
@@ -119,16 +119,8 @@ async fn main() {
// 4. Session 隔离
println!("\n=== Session 隔离(共用 InMemoryStore===");
let store4 = Arc::new(InMemoryStore::new());
let mut a = ConversationMemory::new(
store4.clone(),
"s-a",
ConversationMemoryConfig::default(),
);
let mut b = ConversationMemory::new(
store4.clone(),
"s-b",
ConversationMemoryConfig::default(),
);
let mut a = ConversationMemory::new(store4.clone(), "s-a", ConversationMemoryConfig::default());
let mut b = ConversationMemory::new(store4.clone(), "s-b", ConversationMemoryConfig::default());
a.add_message(Message::user_text("A 的消息")).await.unwrap();
b.add_message(Message::user_text("B 的消息")).await.unwrap();
println!(
@@ -140,4 +132,4 @@ async fn main() {
assert_eq!(b.len(), 1);
println!("\n✓ conversation_memory_demo 完成");
}
}
+11 -16
View File
@@ -17,7 +17,7 @@ use agcore::tools::{
ToolRegistry,
};
use async_trait::async_trait;
use serde_json::{json, Value};
use serde_json::{Value, json};
/// 天气查询工具 —— 模拟根据城市返回天气数据。
struct WeatherTool;
@@ -42,11 +42,7 @@ impl BaseTool for WeatherTool {
fn required_permissions(&self) -> Vec<Permission> {
vec![Permission::Network]
}
async fn execute(
&self,
args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
async fn execute(&self, args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
let city = args["city"].as_str().unwrap_or("未知");
// 模拟查询:根据城市名给出不同温度
let (temperature, condition) = match city {
@@ -84,11 +80,7 @@ impl BaseTool for DeleteFileTool {
fn required_permissions(&self) -> Vec<Permission> {
vec![Permission::Delete]
}
async fn execute(
&self,
_args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
Ok(json!({"deleted": true}))
}
}
@@ -138,9 +130,8 @@ async fn main() {
// 5. 权限检查:默认 PermissionConfig 黑名单含 Delete
println!("\n=== 权限检查(默认 PermissionConfigdenied = [Delete, Shell]===");
let mut registry_with_checker = ToolRegistry::new().with_permission_checker(PermissionChecker::new(
PermissionConfig::default(),
));
let mut registry_with_checker = ToolRegistry::new()
.with_permission_checker(PermissionChecker::new(PermissionConfig::default()));
registry_with_checker
.register(Arc::new(WeatherTool) as ToolRef)
.unwrap();
@@ -155,7 +146,11 @@ async fn main() {
.unwrap();
println!(
"get_weather 权限检查: {}",
if r.output.is_ok() { "通过 ✓" } else { "阻断 ✗" }
if r.output.is_ok() {
"通过 ✓"
} else {
"阻断 ✗"
}
);
// delete_file 声明 Delete → 在 denied 列表 → 阻断
@@ -166,4 +161,4 @@ async fn main() {
println!("delete_file 权限检查: 阻断 ✗ ({err})");
println!("\n✓ custom_tool 完成");
}
}
+247
View File
@@ -0,0 +1,247 @@
//! end_to_end —— 3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证。
//!
//! 运行:`cargo run --example end_to_end`(离线,零配置)
//!
//! ## 真实 LLM Provider 切换
//!
//! 设置环境变量即可使用真实 LLM Provider
//! - `AG_LLM_BASE_URL` —— API 端点(如 `https://api.openai.com/v1`
//! - `AG_LLM_API_KEY` —— API key
//! - `AG_LLM_MODEL` —— 模型名(如 `gpt-4o-mini`
//! - `AG_LLM_PROVIDER`(可选)—— Provider 类型,默认 OpenaiChatOpenAI / DeepSeek / Qwen / Ollama
//!
//! 未设置上述变量时自动降级为 MockProvider,零配置可运行。
use std::env;
use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider;
use agcore::llm::provider::{create_provider, LlmProvider, ProviderConfig, ProviderType};
use agcore::llm::types::{Usage, message::{ContentBlock, Message}, response_v2::{MessageResponse, StopReason}};
use agcore::memory::store::{MemoryStore, SqliteStore};
use agcore::memory::types::{MemoryFilter, MemoryItem};
use agcore::tools::{BaseTool, ToolContext, ToolError, ToolRegistry};
use async_trait::async_trait;
use serde_json::{Value, json};
use tempfile::TempDir;
use time::OffsetDateTime;
// === Agent ===
struct AssistantAgent;
impl Agent for AssistantAgent {
fn name(&self) -> &str { "end-to-end assistant" }
fn system_prompt(&self) -> Option<&str> { Some("简洁助手,必要时调用工具完成任务。") }
}
// === Tools ===
struct EchoTool;
#[async_trait]
impl BaseTool for EchoTool {
fn name(&self) -> &str { "echo" }
fn description(&self) -> &str { "回显输入文本" }
fn parameters(&self) -> Value {
json!({"type":"object","properties":{"text":{"type":"string"}},"required":["text"]})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let text = args.get("text").and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidArguments("text".into(), "需要 string 类型的 text 参数".into()))?;
Ok(json!({"echoed": format!("收到: {text}")}))
}
}
/// 四则运算:'a op b' 格式(ponytail: 基础 +-*/ 不引入 rhai 依赖)。
struct CalcTool;
#[async_trait]
impl BaseTool for CalcTool {
fn name(&self) -> &str { "calc" }
fn description(&self) -> &str { "四则运算:'a op b' 格式,op ∈ {+, -, *, /}" }
fn parameters(&self) -> Value {
json!({"type":"object","properties":{"expr":{"type":"string"}},"required":["expr"]})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let expr = args["expr"].as_str().unwrap_or("");
let parts: Vec<&str> = expr.split_whitespace().collect();
if parts.len() != 3 {
return Err(ToolError::InvalidArguments("expr".into(), "需要 'a op b' 三段式".into()));
}
let a: i64 = parts[0].parse().map_err(|_| ToolError::InvalidArguments("expr".into(), format!("无法解析 '{}'", parts[0])))?;
let b: i64 = parts[2].parse().map_err(|_| ToolError::InvalidArguments("expr".into(), format!("无法解析 '{}'", parts[2])))?;
let result = match parts[1] {
"+" => a + b,
"-" => a - b,
"*" => a * b,
"/" => a.checked_div(b).ok_or_else(|| {
ToolError::InvalidArguments("expr".into(), "除数不能为 0".into())
})?,
op => return Err(ToolError::InvalidArguments("expr".into(), format!("不支持的运算符: {op}"))),
};
Ok(json!({"result": result}))
}
}
/// 通过 MemoryStore trait 读写笔记:直接持有 Arc<dyn MemoryStore>
/// 绕开 AgentSession 封装(NoteTool 在 tool.execute 中直接操作 store)。
/// 关键前缀 "note:" 用于 list 过滤。
struct NoteTool { store: Arc<dyn MemoryStore> }
impl NoteTool { const PREFIX: &'static str = "note:"; }
#[async_trait]
impl BaseTool for NoteTool {
fn name(&self) -> &str { "note" }
fn description(&self) -> &str { "笔记 save/query: save(key, content) / query()" }
fn parameters(&self) -> Value {
json!({
"type":"object",
"properties":{
"action":{"type":"string","enum":["save","query"]},
"key":{"type":"string"},
"content":{"type":"string"}
},
"required":["action"]
})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let action = args["action"].as_str().unwrap_or("");
match action {
"save" => {
let key = args["key"].as_str().unwrap_or("");
let content = args["content"].as_str().unwrap_or("");
let item = MemoryItem {
id: format!("{}{}", Self::PREFIX, key),
content: content.to_string(),
metadata: json!({}),
created_at: OffsetDateTime::now_utc(),
};
self.store.save(item).await
.map_err(|e| ToolError::ExecutionFailed("note".into(), e.to_string()))?;
Ok(json!({"saved": key}))
}
"query" => {
let filter = MemoryFilter { prefix: Some(Self::PREFIX.into()), ..Default::default() };
let items = self.store.list(&filter).await
.map_err(|e| ToolError::ExecutionFailed("note".into(), e.to_string()))?;
let notes: Vec<String> = items.into_iter().map(|i| i.content).collect();
Ok(json!({"notes": notes}))
}
_ => Err(ToolError::InvalidArguments("action".into(), format!("未知 action: {action}"))),
}
}
}
// === Mock response helper ===
fn resp(content: Vec<ContentBlock>, stop: StopReason, u: (u32, u32)) -> MessageResponse {
MessageResponse { id: String::new(), model: "mock".into(),
message: Message::Assistant { content },
usage: Usage::from_input_output(u.0, u.1),
stop_reason: stop, extra: Default::default() }
}
fn mock_responses() -> Vec<MessageResponse> {
vec![
// 第 1 轮:calc(25 * 4) → tool_result(100) → 文本回答
resp(vec![ContentBlock::ToolUse { id: "t1".into(), name: "calc".into(),
input: json!({"expr": "25 * 4"}) }], StopReason::ToolUse, (5, 8)),
resp(vec![ContentBlock::Text { text: "25 * 4 = 100".into() }], StopReason::Stop, (8, 12)),
// 第 2 轮:note(save, last_calc, "100") → tool_result(saved) → 文本回答
resp(vec![ContentBlock::ToolUse { id: "t2".into(), name: "note".into(),
input: json!({"action": "save", "key": "last_calc", "content": "100"}) }],
StopReason::ToolUse, (10, 14)),
resp(vec![ContentBlock::Text { text: "已记录:last_calc = 100".into() }], StopReason::Stop, (12, 16)),
// 第 3 轮:note(query) → tool_result([100]) → 文本回答
resp(vec![ContentBlock::ToolUse { id: "t3".into(), name: "note".into(),
input: json!({"action": "query"}) }], StopReason::ToolUse, (8, 8)),
resp(vec![ContentBlock::Text { text: "您刚才的计算结果是 100".into() }], StopReason::Stop, (10, 14)),
// 后续冗余响应(防止队列耗尽报错)
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
resp(vec![ContentBlock::Text { text: "done".into() }], StopReason::Stop, (1, 1)),
]
}
// === Provider selection ===
fn select_provider() -> Arc<dyn LlmProvider> {
if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
let cfg = ProviderConfig::from_env("AG_LLM").expect("AG_LLM_* 环境变量解析失败");
let provider_type = env::var("AG_LLM_PROVIDER").ok()
.and_then(|s| s.parse::<ProviderType>().ok())
.unwrap_or(ProviderType::OpenaiChat);
Arc::from(create_provider(provider_type, cfg).expect("Provider 创建失败"))
} else {
let found: Vec<&str> = ["AG_LLM_BASE_URL", "AG_LLM_API_KEY", "AG_LLM_MODEL"]
.iter().filter(|k| env::var(k).is_ok()).copied().collect();
eprintln!("AG_LLM_* 环境变量不完整(检测到: {:?}),回退到 MockProvider", found);
Arc::new(MockProvider::new(mock_responses()))
}
}
// === Main ===
#[tokio::main]
async fn main() {
println!("=== agcore 端到端演示 ===");
let dir = TempDir::new().expect("TempDir 创建失败");
let db_path = dir.path().join("agcore.db");
let backend: Arc<dyn MemoryStore> =
Arc::new(SqliteStore::open(&db_path).expect("SqliteStore 打开失败"));
println!("💾 SqliteStore: {}", db_path.display());
let provider_label = if env::var("AG_LLM_BASE_URL").is_ok() && env::var("AG_LLM_API_KEY").is_ok() {
"真实 LLM Provider"
} else {
"MockProvider (离线回退模式)"
};
println!("🔄 Provider: {provider_label}");
let mut registry = ToolRegistry::new();
registry.register(Arc::new(EchoTool)).unwrap();
registry.register(Arc::new(CalcTool)).unwrap();
registry.register(Arc::new(NoteTool { store: backend.clone() })).unwrap();
println!("🔧 注册工具: {:?}", registry.list_tools());
let bundle = Arc::new(AgentBuilder::new()
.provider(select_provider())
.tool_registry(Arc::new(registry))
.hook_executor(Arc::new(HookExecutor::new()))
.build().expect("RuntimeBundle 装配失败"));
let mut session = AgentSession::new(Arc::new(AssistantAgent), "e2e-1", bundle.clone());
println!("\n第 1 轮 用户: 帮我算 25 * 4");
let r1 = session.submit_turn("帮我算 25 * 4").await.expect("turn 1 失败");
println!(" → 回答: {}", r1.text());
println!("\n第 2 轮 用户: 记下来:结果是 100");
let r2 = session.submit_turn("记下来:结果是 100").await.expect("turn 2 失败");
println!(" → 回答: {}", r2.text());
println!("\n第 3 轮 用户: 我刚才算了什么?");
let r3 = session.submit_turn("我刚才算了什么?").await.expect("turn 3 失败");
println!(" → 回答: {}", r3.text());
let total = session.usage().total();
println!("\n📊 用量: prompt={}, completion={}, total={}",
total.prompt_tokens, total.completion_tokens, total.total_tokens);
println!("\n=== 持久化验证 ===");
// 显式释放所有对 backend 的 Arc 引用,确保 SqliteStore Connection 真正关闭。
// 释放顺序:session → bundle(间接持有 NoteTool → backend clone)→ backend 局部变量。
drop(session); // session.bundle Arc 计数 -1
drop(bundle); // bundle Arc 计数归零 → registry → NoteTool → backend clone Arc 计数 2→1
drop(backend); // backend 局部变量 Arc 计数 1→0 → SqliteStore::drop → Connection 自动 close
let backend2: Arc<dyn MemoryStore> =
Arc::new(SqliteStore::open(&db_path).expect("重开 SqliteStore 失败"));
let filter = MemoryFilter { prefix: Some("note:".into()), ..Default::default() };
let items = backend2.list(&filter).await.expect("list 失败");
println!("✓ 跨连接数据存活: 找到 {} 条 note", items.len());
assert!(!items.is_empty(), "持久化验证失败:重开后无数据");
for i in &items {
println!(" - {} = {}", i.id, i.content);
}
println!("\n✓ 端到端演示完成");
}
+20 -20
View File
@@ -38,14 +38,26 @@ async fn main() {
let ks = KnowledgeStore::new(store);
let pages = vec![
make_page("rust-1", "Rust 入门", "Rust 是一门系统级编程语言,注重安全性与并发。"),
make_page("python-1", "Python 简介", "Python 是一门动态类型的高级编程语言。"),
make_page(
"rust-1",
"Rust 入门",
"Rust 是一门系统级编程语言,注重安全性与并发。",
),
make_page(
"python-1",
"Python 简介",
"Python 是一门动态类型的高级编程语言。",
),
make_page(
"langgraph-1",
"LangGraph 框架",
"LangGraph 是 LangChain 的状态图扩展,用于构建多步 Agent。",
),
make_page("rust-async", "Rust 异步编程", "Rust 异步基于 tokio 与 futures 抽象。"),
make_page(
"rust-async",
"Rust 异步编程",
"Rust 异步基于 tokio 与 futures 抽象。",
),
];
for p in &pages {
ks.add_page(p.clone()).await.expect("保存页面失败");
@@ -62,14 +74,8 @@ async fn main() {
let result = retriever.retrieve("Rust 异步").await.unwrap();
println!("query: {}", result.query);
for item in &result.items {
println!(
" 命中: {} (score={:.3})",
item.page.title, item.score
);
assert!(
(0.0..=1.0).contains(&item.score),
"score 应在 [0, 1] 区间"
);
println!(" 命中: {} (score={:.3})", item.page.title, item.score);
assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
}
assert!(!result.items.is_empty(), "应至少命中一个页面");
@@ -85,14 +91,8 @@ async fn main() {
min_score: 0.5,
};
let retriever2 = MemoryRetriever::new(ks2, cfg);
let result = retriever2
.retrieve("完全不相关的火锅配方")
.await
.unwrap();
println!(
"无关 query → items.len = {} (期望 0)",
result.items.len()
);
let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
println!("无关 query → items.len = {} (期望 0)", result.items.len());
assert!(result.items.is_empty());
// 4. max_results 截断
@@ -146,4 +146,4 @@ async fn main() {
assert!(only_stop.items.is_empty(), "纯停用词 query 必须返回空结果");
println!("\n✓ knowledge_search_demo 完成");
}
}
+11 -7
View File
@@ -11,7 +11,7 @@
use agcore::llm::types::message::{ContentBlock, Message};
use agcore::prompt::{
validate_messages, PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext,
PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext, validate_messages,
};
fn message_text(msg: &Message) -> String {
@@ -27,16 +27,16 @@ fn message_text(msg: &Message) -> String {
})
.collect(),
Message::UserImage { .. } => "[image]".into(),
_ => String::new(),
}
}
fn main() {
// 1. PromptTemplate::compile + render —— 直接构造模板
println!("=== PromptTemplate::compile + render ===");
let tpl = PromptTemplate::compile(
"今日 {{location}} 天气:{{condition}},温度 {{temperature}}",
)
.expect("编译失败");
let tpl =
PromptTemplate::compile("今日 {{location}} 天气:{{condition}},温度 {{temperature}}")
.expect("编译失败");
let mut ctx = TemplateContext::new();
ctx.insert("location", "北京");
ctx.insert("condition", "");
@@ -58,7 +58,10 @@ fn main() {
.register("weather", "今日 {{location}}{{condition}}")
.expect("注册失败");
registry
.register("greet", "你好 {{name}}{{#if formal}} 见到您很荣幸。{{/if}}")
.register(
"greet",
"你好 {{name}}{{#if formal}} 见到您很荣幸。{{/if}}",
)
.expect("注册失败");
let mut ctx = TemplateContext::new();
@@ -88,6 +91,7 @@ fn main() {
Message::User { .. } | Message::UserImage { .. } => "user",
Message::Assistant { .. } => "assistant",
Message::ToolResult { .. } => "tool",
_ => "unknown",
};
println!("[{i}] {role}: {}", message_text(m));
}
@@ -105,4 +109,4 @@ fn main() {
}
println!("\n✓ prompt_composer 完成");
}
}
+60
View File
@@ -0,0 +1,60 @@
//! quick_start —— 30 行最小可运行示例,展示 Agent / BaseTool / Builder / Session 四层抽象。
//!
//! 运行:`cargo run --example quick_start`(离线,零配置)
use std::sync::Arc;
use agcore::agent::{Agent, AgentBuilder, AgentSession};
use agcore::llm::hooks::HookExecutor;
use agcore::llm::mock::MockProvider;
use agcore::llm::provider::LlmProvider;
use agcore::llm::types::{Usage, message::{ContentBlock, Message}, response_v2::{MessageResponse, StopReason}};
use agcore::tools::{BaseTool, ToolContext, ToolError, ToolRegistry};
use async_trait::async_trait;
use serde_json::{Value, json};
struct Greeter;
impl Agent for Greeter {
fn name(&self) -> &str { "greeter" }
fn system_prompt(&self) -> Option<&str> { Some("中文助手,先调用 echo 工具,再总结。") }
}
struct EchoTool;
#[async_trait]
impl BaseTool for EchoTool {
fn name(&self) -> &str { "echo" }
fn description(&self) -> &str { "回显文本" }
fn parameters(&self) -> Value {
json!({"type":"object","properties":{"text":{"type":"string"}},"required":["text"]})
}
async fn execute(&self, args: Value, _: &ToolContext<'_>) -> Result<Value, ToolError> {
let text = args.get("text").and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidArguments("text".into(), "需要 string 类型的 text 参数".into()))?;
Ok(json!({"echoed": format!("收到: {text}")}))
}
}
fn resp(content: Vec<ContentBlock>, stop: StopReason, u: (u32, u32)) -> MessageResponse {
MessageResponse { id: String::new(), model: "mock".into(), message: Message::Assistant { content },
usage: Usage::from_input_output(u.0, u.1), stop_reason: stop, extra: Default::default() }
}
#[tokio::main]
async fn main() {
let mut registry = ToolRegistry::new();
registry.register(Arc::new(EchoTool)).unwrap();
let provider: Arc<dyn LlmProvider> = Arc::new(MockProvider::new(vec![
resp(vec![ContentBlock::ToolUse { id: "c1".into(), name: "echo".into(),
input: json!({"text": "你好"}) }], StopReason::ToolUse, (5, 8)),
resp(vec![ContentBlock::Text { text: "EchoTool 已收到您的消息并完成回传。".into() }],
StopReason::Stop, (8, 16)),
]));
let bundle = Arc::new(AgentBuilder::new()
.provider(provider).tool_registry(Arc::new(registry))
.hook_executor(Arc::new(HookExecutor::new())).build().unwrap());
let mut session = AgentSession::new(Arc::new(Greeter), "qs", bundle);
let resp = session.submit_turn("你好").await.unwrap();
let text = resp.text();
println!("LLM: {text}");
assert!(text.contains("收到"), "响应应包含'收到'字样: {text}");
println!("\n✓ quick_start 完成");
}
+7 -6
View File
@@ -3,7 +3,7 @@ use std::env;
use agcore::init_tracing;
use agcore::llm::{
cycle::{CycleConfig, LlmCycle},
provider::{create_provider, ProviderConfig, ProviderType},
provider::{ProviderConfig, ProviderType, create_provider},
types::{message::ContentBlock, message::Message, response_v2::MessageResponse},
};
@@ -51,10 +51,11 @@ async fn main() {
base_url,
api_key,
model: model.clone(),
timeout_secs: 30,
max_retries: 3,
};
let provider = create_provider(provider_type, config)
.expect("创建 Provider 失败");
let provider = create_provider(provider_type, config).expect("创建 Provider 失败");
let cycle_config = CycleConfig {
model,
@@ -63,9 +64,9 @@ async fn main() {
..CycleConfig::default()
};
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![
Message::system("你是一个简洁的助手,对于任何问题都是用一句话回答。"),
]);
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![Message::system(
"你是一个简洁的助手,对于任何问题都是用一句话回答。",
)]);
println!("发送请求...");
+3 -5
View File
@@ -17,9 +17,9 @@ use std::sync::Arc;
use agcore::llm::cycle::{CycleConfig, LlmCycle};
use agcore::llm::mock::MockProvider;
use agcore::llm::provider::LlmProvider;
use agcore::llm::types::Usage;
use agcore::llm::types::message::{ContentBlock, Message};
use agcore::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
use agcore::llm::types::Usage;
use futures_util::StreamExt;
/// 构造预设的纯文本响应。
@@ -99,9 +99,7 @@ async fn main() {
// 上层 Agent 通过 `match` 或 `?` 处理 `AgentError::Llm(_)`。
println!("\n=== 阶段 2:错误路径(队列耗尽)===");
let mut cycle = LlmCycle::new_with_arc(dyn_provider, CycleConfig::default());
let result = cycle
.submit_stream("第二次提问".to_string(), vec![])
.await;
let result = cycle.submit_stream("第二次提问".to_string(), vec![]).await;
match result {
Ok(_) => panic!("阶段 2 必须失败(队列耗尽)"),
Err(e) => {
@@ -114,4 +112,4 @@ async fn main() {
}
println!("\n✓ streaming_events_demo 完成");
}
}
+25 -24
View File
@@ -8,27 +8,13 @@
//! 5. 错误路径:非法 JSON / 空 steps / 缺字段 → `AgentError::PlanParse`
//!
//! 运行:`cargo run --example task_agent_demo`
//!
//! ## 已知技术债(v0.2 迁移指南)
//!
//! 本示例使用 `#[deprecated]` 标记的旧 wire-format 类型:
//! - `ChatResponse`、`OpenaiChatMessage`、`FinishReason` —— `OpenaiChatProvider::chat_inner()`
//! 内部转换层仍在使用(参见 `docs/10a-phase0-types-and-trait.md` §2.5.1),
//! 故结构体定义保留。
//! - `StepStatus::Completed(ChatResponse)` —— 因为 `Step` 的"已完成"变体需携带
//! provider 响应,目前沿用旧的 `ChatResponse`。
//!
//! **触发迁移的条件**v0.2 引入 IR 层的 `StepResult` / 切换为 `MessageResponse`。
//! **迁移路径**:将本文件 `ChatResponse`/`OpenaiChatMessage`/`FinishReason` 替换为
//! `MessageResponse`/`Message`/`StopReason`,移除顶部 `#![allow(deprecated)]`。
//! 上层应用代码(`TaskAgent` 消费者)也可同步迁移。
#![allow(deprecated)]
use std::collections::HashMap;
use agcore::agent::{AgentError, JsonPlanParser, PlanParser, Step, StepStatus};
use agcore::llm::types::openai_message::OpenaiChatMessage;
use agcore::llm::types::shared::FinishReason;
use agcore::llm::types::{ChatResponse, Usage};
use agcore::llm::types::message::Message;
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
use agcore::llm::types::Usage;
#[tokio::main]
async fn main() {
@@ -67,21 +53,36 @@ async fn main() {
assert!(step.status.is_pending());
step.status = StepStatus::Running;
println!("Running: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
println!(
"Running: pending={}, terminal={}",
step.status.is_pending(),
step.status.is_terminal()
);
step.status = StepStatus::Completed(ChatResponse {
message: OpenaiChatMessage::assistant_text("天气:晴,22°C"),
step.status = StepStatus::Completed(MessageResponse {
id: String::new(),
model: "mock".into(),
message: Message::assistant("天气:晴,22°C"),
usage: Usage::from_input_output(5, 10),
stop_reason: Some(FinishReason::Stop),
stop_reason: StopReason::Stop,
extra: HashMap::new(),
});
println!("Completed: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
println!(
"Completed: pending={}, terminal={}",
step.status.is_pending(),
step.status.is_terminal()
);
assert!(step.status.is_terminal());
// 3. 失败路径
println!("\n=== Step 状态机:失败路径 ===");
let mut fail_step = Step::new(0, "调用天气 API");
fail_step.status = StepStatus::Failed(AgentError::Other("API 不可用".into()));
println!("Failed: pending={}, terminal={}", fail_step.status.is_pending(), fail_step.status.is_terminal());
println!(
"Failed: pending={}, terminal={}",
fail_step.status.is_pending(),
fail_step.status.is_terminal()
);
assert!(fail_step.status.is_terminal());
// 4. 跳过路径
+1 -1
View File
@@ -24,5 +24,5 @@ pub use error::AgentError;
pub use runtime::{AgentConfig, RuntimeBundle};
pub use session::AgentSession;
pub use session_memory::SessionMemory;
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
pub use task::JsonPlanParser;
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
+2 -4
View File
@@ -7,14 +7,12 @@
//! - **不绑定业务循环**`submit_turn` 在 `AgentSession` 上,不在 trait 上
use crate::agent::runtime::RuntimeBundle;
#[allow(deprecated)]
use crate::llm::types::ToolDefinition;
use crate::llm::types::tool::ToolDef;
/// Agent 角色抽象。
///
/// 实现此 trait 即可接入 Agent Runtime。典型实现是 struct 持有静态配置(name、system prompt 模板),
/// 也可以是基于配置动态生成的轻量实现。
#[allow(deprecated)]
pub trait Agent: Send + Sync {
/// 角色名(用于日志、调试、UI 展示)。
fn name(&self) -> &str;
@@ -26,7 +24,7 @@ pub trait Agent: Send + Sync {
///
/// **默认实现**:从 `bundle.tool_registry` 取全部工具(最常用模式)。
/// **子 trait / 具体实现可覆盖**:做白名单、过滤、按状态动态调整等。
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDefinition> {
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDef> {
bundle.tool_registry.definitions()
}
}
+8 -6
View File
@@ -92,15 +92,17 @@ impl AgentBuilder {
/// `AgentError::Config(...)`,提示调用 `.provider(...)` / `.tool_registry(...)` /
/// `.hook_executor(...)` 补齐。不 panic。
pub fn build(self) -> Result<RuntimeBundle, AgentError> {
let provider = self
.provider
.ok_or_else(|| AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into()))?;
let provider = self.provider.ok_or_else(|| {
AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into())
})?;
let tool_registry = self
.tool_registry
.ok_or_else(|| AgentError::Config("缺少 tool_registry,请先调用 .tool_registry(...)(即使是空 ToolRegistry 也需要传入)".into()))?;
let hook_executor = self
.hook_executor
.ok_or_else(|| AgentError::Config("缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into()))?;
let hook_executor = self.hook_executor.ok_or_else(|| {
AgentError::Config(
"缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into(),
)
})?;
let config = self.config.unwrap_or_default();
+1
View File
@@ -18,6 +18,7 @@ use crate::tools::error::ToolError;
/// **不实现 `Clone`**:透传内层 `LlmError` / `MemoryError`,两者均未派生 `Clone`(保留
/// 完整错误信息,传递所有权)。如需在多 session 间共享错误状态,用 `Arc<AgentError>` 包装。
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum AgentError {
/// LLM 调用错误(透传 Phase 0)。
#[error("LLM 错误: {0}")]
+1 -1
View File
@@ -16,8 +16,8 @@ use std::sync::Arc;
use std::time::Duration;
use crate::llm::compact::CompactConfig;
use crate::llm::provider::LlmProvider;
use crate::llm::hooks::HookExecutor;
use crate::llm::provider::LlmProvider;
use crate::memory::retriever::MemoryRetriever;
use crate::memory::store::MemoryStore;
use crate::tools::ToolRegistry;
+9 -8
View File
@@ -126,8 +126,7 @@ impl AgentSession {
let hook_executor = Arc::clone(&self.bundle.hook_executor);
// 1. 触发 OnTurnStart hook
let start_ctx =
HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
let start_ctx = HookContext::new(HookEvent::OnTurnStart).with_turn_index(turn_index);
hook_executor
.execute(HookEvent::OnTurnStart, &start_ctx)
.await;
@@ -137,8 +136,9 @@ impl AgentSession {
// submit_with_tools 内部从 registry 自行取 definitions,此处仅消费以触发
// 子 trait 覆盖(白名单/过滤)的副作用。
let _ = self.agent.tool_definitions(&self.bundle);
let mut cycle = LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default())
.with_messages(Vec::new());
let mut cycle =
LlmCycle::new_with_arc(Arc::clone(&self.bundle.provider), CycleConfig::default())
.with_messages(Vec::new());
// Phase 2 切换 system_prompt 字段为 Message::SystemFIX-D)。
// 若 agent 自带 system prompt,预置到 messages 列表头部。
let mut initial_messages: Vec<Message> = Vec::new();
@@ -223,9 +223,7 @@ mod tests {
id: String::new(),
model: String::new(),
message: Message::Assistant {
content: vec![ContentBlock::Text {
text: text.into(),
}],
content: vec![ContentBlock::Text { text: text.into() }],
},
usage: crate::llm::types::Usage::from_input_output(10, 5),
stop_reason: StopReason::Stop,
@@ -280,7 +278,10 @@ mod tests {
assert!(session.get_session_data("k").await.unwrap().is_none());
session.set_session_data("k", "v").await.unwrap();
assert_eq!(session.get_session_data("k").await.unwrap(), Some("v".into()));
assert_eq!(
session.get_session_data("k").await.unwrap(),
Some("v".into())
);
// 覆盖写
session.set_session_data("k", "v2").await.unwrap();
assert_eq!(
+3 -11
View File
@@ -78,11 +78,7 @@ impl SessionMemory {
prefix: Some(format!("{}:", self.namespace)),
..Default::default()
};
let items = self
.store
.list(&filter)
.await
.map_err(AgentError::Memory)?;
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
let mut lines = Vec::with_capacity(items.len() + 2);
lines.push("<session-context>".to_string());
@@ -113,11 +109,7 @@ impl SessionMemory {
prefix: Some(format!("{}:", self.namespace)),
..Default::default()
};
let items = self
.store
.list(&filter)
.await
.map_err(AgentError::Memory)?;
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
for item in items {
self.store
@@ -181,4 +173,4 @@ mod tests {
assert!(mem_a.get("key").await.unwrap().is_none());
assert_eq!(mem_b.get("key").await.unwrap(), Some("val_b".into()));
}
}
}
+5 -11
View File
@@ -10,8 +10,7 @@
//! - 重试由上层新建 `Plan` 实现,`TaskAgent` 不做自动重试
use crate::agent::error::AgentError;
#[allow(deprecated)]
use crate::llm::types::ChatResponse;
use crate::llm::types::response_v2::MessageResponse;
use async_trait::async_trait;
@@ -56,14 +55,14 @@ impl Step {
/// 均未派生 `Clone`(保留原始错误信息,传递所有权而非克隆)。如需复制 `Plan`,
/// 只能 clone 处于 `Pending` / `Running` / `Completed` / `Skipped` 状态的步骤。
#[derive(Debug)]
#[allow(deprecated)]
#[non_exhaustive]
pub enum StepStatus {
/// 初始状态 —— 等待执行。
Pending,
/// 正在执行(`TaskAgent::execute_plan` 进入)。
Running,
/// 已完成(含 LLM 响应)。
Completed(ChatResponse),
Completed(MessageResponse),
/// 失败(含错误)。
Failed(AgentError),
/// 跳过(上层主动跳过)。
@@ -130,9 +129,7 @@ impl PlanParser for JsonPlanParser {
.collect::<Result<Vec<_>, AgentError>>()?;
if steps.is_empty() {
return Err(AgentError::PlanParse(
"Plan 至少需要一个步骤".into(),
));
return Err(AgentError::PlanParse("Plan 至少需要一个步骤".into()));
}
Ok(Plan {
@@ -203,10 +200,7 @@ mod tests {
let plan = Plan {
id: "p1".into(),
goal: "test goal".into(),
steps: vec![
Step::new(0, "first"),
Step::new(1, "second"),
],
steps: vec![Step::new(0, "first"), Step::new(1, "second")],
};
assert_eq!(plan.steps.len(), 2);
assert_eq!(plan.steps[0].index, 0);
+32 -16
View File
@@ -73,10 +73,7 @@ impl CompactState {
/// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。
pub fn estimate_message_tokens(messages: &[Message]) -> u32 {
messages
.iter()
.map(estimate_single_message_tokens)
.sum()
messages.iter().map(estimate_single_message_tokens).sum()
}
fn estimate_single_message_tokens(msg: &Message) -> u32 {
@@ -99,9 +96,7 @@ fn estimate_block_tokens(block: &ContentBlock) -> u32 {
match block {
ContentBlock::Text { text } => estimate_text_tokens(text),
ContentBlock::Thinking { text, .. } => estimate_text_tokens(text),
ContentBlock::ToolUse { input, .. } => {
estimate_text_tokens(&input.to_string())
}
ContentBlock::ToolUse { input, .. } => estimate_text_tokens(&input.to_string()),
ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
// ponytail: Image / Audio / File / Extension 在 IR 中固定估算。
// 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。
@@ -148,14 +143,25 @@ pub fn microcompact(messages: &mut [Message], keep_recent: usize) -> u32 {
// 第一遍:计算可释放 token(仅非错误 ToolResult
for msg in &messages[..prune_start] {
if matches!(msg, Message::ToolResult { is_error: false, .. }) {
if matches!(
msg,
Message::ToolResult {
is_error: false,
..
}
) {
freed_tokens += estimate_single_message_tokens(msg);
}
}
// 第二遍:替换内容(仅非错误 ToolResult
for msg in &mut messages[..prune_start] {
if let Message::ToolResult { content, is_error: false, .. } = msg {
if let Message::ToolResult {
content,
is_error: false,
..
} = msg
{
*content = vec![ContentBlock::Text {
text: "[pruned]".to_string(),
}];
@@ -177,13 +183,15 @@ mod tests {
fn estimate_message_tokens_handles_all_variants() {
let messages = vec![
Message::System {
content: vec![ContentBlock::Text {
text: "sys".into(),
}],
content: vec![ContentBlock::Text { text: "sys".into() }],
},
Message::user_text("hi"),
Message::assistant("ans"),
Message::user_image("b64", "image/png", crate::llm::types::shared::ImageDetail::Auto),
Message::user_image(
"b64",
"image/png",
crate::llm::types::shared::ImageDetail::Auto,
),
Message::tool_result("call_1", "tool res", false),
];
let tokens = estimate_message_tokens(&messages);
@@ -205,7 +213,10 @@ mod tests {
assert!(freed > 0);
assert_eq!(messages.len(), before_len); // 只改内容,不删消息
// 索引 1 是被压缩的 ToolResult
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
if let Message::ToolResult {
content, is_error, ..
} = &messages[1]
{
assert_eq!(content.len(), 1);
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
assert!(!is_error);
@@ -228,9 +239,14 @@ mod tests {
assert_eq!(freed, 0); // 错误 ToolResult 不计入
assert_eq!(messages.len(), before_len);
// 错误信息保留完整
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
if let Message::ToolResult {
content, is_error, ..
} = &messages[1]
{
assert!(is_error);
assert!(matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down")));
assert!(
matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down"))
);
} else {
panic!("expected ToolResult at index 1");
}
+29 -30
View File
@@ -8,11 +8,9 @@
use serde_json::Value;
use crate::llm::types::message::{ContentBlock, Message};
use crate::llm::types::openai_message::{
ContentField, OpenaiChatMessage, OpenaiContentPart,
};
use crate::llm::types::OpenaiToolCall;
use crate::llm::types::message::{ContentBlock, Message};
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage, OpenaiContentPart};
/// `OpenaiChatMessage` → IR `Message`。
///
@@ -24,11 +22,10 @@ use crate::llm::types::OpenaiToolCall;
/// - `Function`(已废弃)→ `Message::ToolResult``name` 作为 `tool_call_id` 兜底)
pub fn from_openai(msg: &OpenaiChatMessage) -> Message {
match msg {
OpenaiChatMessage::Developer { content, .. } | OpenaiChatMessage::System { content, .. } => {
Message::System {
content: content_to_blocks(content),
}
}
OpenaiChatMessage::Developer { content, .. }
| OpenaiChatMessage::System { content, .. } => Message::System {
content: content_to_blocks(content),
},
OpenaiChatMessage::User { content, .. } => Message::User {
content: content_to_blocks(content),
},
@@ -86,7 +83,11 @@ pub fn to_openai(msg: &Message) -> OpenaiChatMessage {
content: blocks_to_content(content),
name: None,
},
Message::UserImage { data, mime_type, detail } => {
Message::UserImage {
data,
mime_type,
detail,
} => {
// ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。
let mime = mime_type.clone();
let is_url = data.starts_with("http://") || data.starts_with("https://");
@@ -167,26 +168,25 @@ pub fn content_to_blocks(field: &ContentField) -> Vec<ContentBlock> {
ContentField::Array(parts) => parts
.iter()
.filter_map(|p| match p {
OpenaiContentPart::Text { text } => {
Some(ContentBlock::Text { text: text.clone() })
}
OpenaiContentPart::Refusal { refusal } => {
Some(ContentBlock::Text { text: refusal.clone() })
}
OpenaiContentPart::Text { text } => Some(ContentBlock::Text { text: text.clone() }),
OpenaiContentPart::Refusal { refusal } => Some(ContentBlock::Text {
text: refusal.clone(),
}),
OpenaiContentPart::Image { image_url, .. } => {
// ponytail: 简化处理 —— URL 直接通过,data URI 拆出
// data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。
let url = &image_url.url;
if let Some(rest) = url.strip_prefix("data:")
&& let Some((mime, b64)) = rest.split_once(";base64,") {
return Some(ContentBlock::Image {
source: crate::llm::types::message::ImageSource {
data: b64.to_string(),
mime_type: mime.to_string(),
is_url: false,
},
});
}
&& let Some((mime, b64)) = rest.split_once(";base64,")
{
return Some(ContentBlock::Image {
source: crate::llm::types::message::ImageSource {
data: b64.to_string(),
mime_type: mime.to_string(),
is_url: false,
},
});
}
Some(ContentBlock::Image {
source: crate::llm::types::message::ImageSource {
data: url.clone(),
@@ -263,7 +263,9 @@ mod tests {
match ir {
Message::System { content } => {
assert_eq!(content.len(), 1);
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful"));
assert!(
matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful")
);
}
_ => panic!("expected System variant"),
}
@@ -385,10 +387,7 @@ mod tests {
assert_eq!(parts.len(), 1);
match &parts[0] {
OpenaiContentPart::Image { image_url, .. } => {
assert_eq!(
image_url.url,
"data:image/png;base64,BASE64DATA"
);
assert_eq!(image_url.url, "data:image/png;base64,BASE64DATA");
}
_ => panic!("expected Image part"),
}
+45 -42
View File
@@ -9,10 +9,10 @@ pub use usage::{CostTracker, Usage};
use std::pin::Pin;
use std::sync::Arc;
use futures_core::stream::Stream;
use async_stream::stream;
use futures_core::stream::Stream;
use crate::llm::compact::{should_compact, microcompact, CompactConfig, CompactState};
use crate::llm::compact::{CompactConfig, CompactState, microcompact, should_compact};
use crate::llm::cycle::retry::should_retry;
use crate::llm::error::LlmError;
use crate::llm::hooks::{HookContext, HookExecutor};
@@ -21,8 +21,8 @@ use crate::llm::stream::StreamEvent;
use crate::llm::types::message::{ContentBlock, Message};
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StopReason};
#[allow(deprecated)]
use crate::llm::types::{ToolChoice, ToolDefinition};
use crate::llm::types::tool::ToolDef;
use crate::llm::types::ToolChoice;
/// LLM 调用周期配置。
pub struct CycleConfig {
@@ -113,8 +113,12 @@ impl LlmCycle {
note = "请改用 Message::system_text() + with_messages()"
)]
pub fn with_system_prompt(mut self, prompt: String) -> Self {
self.messages
.insert(0, Message::System { content: vec![ContentBlock::Text { text: prompt }] });
self.messages.insert(
0,
Message::System {
content: vec![ContentBlock::Text { text: prompt }],
},
);
self
}
@@ -175,7 +179,7 @@ impl LlmCycle {
pub async fn submit_messages(
&mut self,
messages: Vec<Message>,
tools: Vec<ToolDefinition>,
tools: Vec<ToolDef>,
) -> Result<MessageResponse, LlmError> {
let request = MessageRequest {
model: self.config.model.clone(),
@@ -188,8 +192,8 @@ impl LlmCycle {
};
if let Some(ref executor) = self.hook_executor {
let ctx = HookContext::new(crate::llm::hooks::HookEvent::PreRequest)
.with_request(&request);
let ctx =
HookContext::new(crate::llm::hooks::HookEvent::PreRequest).with_request(&request);
let results = executor
.execute(crate::llm::hooks::HookEvent::PreRequest, &ctx)
.await;
@@ -218,7 +222,8 @@ impl LlmCycle {
}
Err(e) => {
if let Some(ref executor) = self.hook_executor {
let ctx = HookContext::new(crate::llm::hooks::HookEvent::OnError).with_error(&e);
let ctx =
HookContext::new(crate::llm::hooks::HookEvent::OnError).with_error(&e);
executor
.execute(crate::llm::hooks::HookEvent::OnError, &ctx)
.await;
@@ -232,7 +237,7 @@ impl LlmCycle {
pub async fn submit(
&mut self,
prompt: String,
tools: Vec<ToolDefinition>,
tools: Vec<ToolDef>,
) -> Result<MessageResponse, LlmError> {
self.messages.push(Message::user_text(prompt));
@@ -342,7 +347,7 @@ impl LlmCycle {
pub async fn submit_stream(
&mut self,
prompt: String,
tools: Vec<ToolDefinition>,
tools: Vec<ToolDef>,
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, LlmError> {
self.messages.push(Message::user_text(prompt));
@@ -359,8 +364,8 @@ impl LlmCycle {
// PreRequest hook
if let Some(ref executor) = self.hook_executor {
let ctx = HookContext::new(crate::llm::hooks::HookEvent::PreRequest)
.with_request(&request);
let ctx =
HookContext::new(crate::llm::hooks::HookEvent::PreRequest).with_request(&request);
let results = executor
.execute(crate::llm::hooks::HookEvent::PreRequest, &ctx)
.await;
@@ -424,7 +429,7 @@ impl LlmCycle {
}))
}
fn build_request(&self, tools: &[ToolDefinition]) -> MessageRequest {
fn build_request(&self, tools: &[ToolDef]) -> MessageRequest {
// ponytail: Phase 2 简化 —— 直接 clone self.messages,无任何转换 / system prompt 注入。
// 系统消息如需存在,由调用方通过 `with_messages()` 自行管理。
MessageRequest {
@@ -443,7 +448,7 @@ impl LlmCycle {
/// 用于 `submit_with_tools()` 的多轮 tool 循环。
async fn submit_request(
&mut self,
tools: &[ToolDefinition],
tools: &[ToolDef],
) -> Result<MessageResponse, LlmError> {
let mut attempts = 0;
@@ -496,8 +501,8 @@ impl LlmCycle {
}
Err(e) => {
if let Some(ref executor) = self.hook_executor {
let ctx = HookContext::new(crate::llm::hooks::HookEvent::OnError)
.with_error(&e);
let ctx =
HookContext::new(crate::llm::hooks::HookEvent::OnError).with_error(&e);
executor
.execute(crate::llm::hooks::HookEvent::OnError, &ctx)
.await;
@@ -593,11 +598,8 @@ impl LlmCycle {
// 真实 tool_call_id 而非 tool_name 充当 —— 这条 FIX-A 修复与 Phase 2 消息切换
// 同步生效。
// ponytail: Phase 2 直接存储 Message::ToolResultis_error 由 ToolInvocation.output 推断。
self.messages.push(Message::tool_result(
result.tool_call_id,
content,
is_error,
));
self.messages
.push(Message::tool_result(result.tool_call_id, content, is_error));
}
// 每轮工具执行后触发 compaction
@@ -641,9 +643,7 @@ fn has_tool_calls_in_response(response: &MessageResponse) -> bool {
/// ponytail: 当前 Phase 0 实现,`arguments_json_string` 内含 JSON 序列化的 input。
/// 消费方在调用 `registry.invoke_all()` 时反序列化一次。该小段冗余序列化
/// 在 Phase 2 切换为 `Vec<Message>` 后可整体消除。
fn extract_tool_calls_from_response(
response: &MessageResponse,
) -> Vec<(String, String, String)> {
fn extract_tool_calls_from_response(response: &MessageResponse) -> Vec<(String, String, String)> {
let mut out = Vec::new();
if let Message::Assistant { content } = &response.message {
for block in content {
@@ -665,7 +665,11 @@ fn truncate_tool_result(s: &str, max_bytes: usize) -> String {
while end > 0 && !s.is_char_boundary(end) {
end -= 1;
}
format!("{}\n\n[... truncated, original size: {} bytes ...]", &s[..end], s.len())
format!(
"{}\n\n[... truncated, original size: {} bytes ...]",
&s[..end],
s.len()
)
}
#[cfg(test)]
@@ -675,7 +679,7 @@ mod tests {
use crate::tools::{BaseTool, ToolRegistry};
use async_trait::async_trait;
use futures_core::Stream;
use serde_json::{json, Value};
use serde_json::{Value, json};
use std::pin::Pin;
/// 模拟 Provider —— 预定义响应序列,按调用顺序返回。
@@ -729,9 +733,7 @@ mod tests {
id: String::new(),
model: String::new(),
message: Message::Assistant {
content: vec![ContentBlock::Text {
text: text.into(),
}],
content: vec![ContentBlock::Text { text: text.into() }],
},
usage: empty_usage(),
stop_reason: StopReason::Stop,
@@ -755,7 +757,9 @@ mod tests {
MessageResponse {
id: String::new(),
model: String::new(),
message: Message::Assistant { content: tool_blocks },
message: Message::Assistant {
content: tool_blocks,
},
usage: empty_usage(),
stop_reason: StopReason::ToolUse,
extra: std::collections::HashMap::new(),
@@ -810,16 +814,13 @@ mod tests {
let messages = cycle.messages();
assert_eq!(messages.len(), 4);
assert!(matches!(messages[0], Message::User { .. }));
assert!(matches!(
messages[1],
Message::Assistant {
content: _,
}
));
assert!(matches!(messages[1], Message::Assistant { content: _ }));
if let Message::Assistant { content } = &messages[1] {
assert!(content
.iter()
.any(|b| matches!(b, ContentBlock::ToolUse { .. })));
assert!(
content
.iter()
.any(|b| matches!(b, ContentBlock::ToolUse { .. }))
);
}
assert!(matches!(
messages[2],
@@ -875,7 +876,9 @@ mod tests {
registry.register(std::sync::Arc::new(AddTool)).unwrap();
let result = cycle.submit_with_tools("test".to_string(), &registry).await;
assert!(matches!(result, Err(LlmError::Other(msg)) if msg.contains("达到最大工具循环轮次")));
assert!(
matches!(result, Err(LlmError::Other(msg)) if msg.contains("达到最大工具循环轮次"))
);
}
#[tokio::test]
+11 -4
View File
@@ -8,9 +8,12 @@ use std::time::Duration;
///
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查 API key、减少上下文)。
#[derive(thiserror::Error, Debug)]
#[non_exhaustive]
pub enum LlmError {
/// API 认证失败(API key 无效、过期或权限不足)。
#[error("LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确")]
#[error(
"LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确"
)]
Authentication(String),
/// 请求被限流,可选地附带重试等待时间。可重试。
@@ -18,7 +21,9 @@ pub enum LlmError {
RateLimit { retry_after: Option<Duration> },
/// HTTP 请求失败(网络错误或非 2xx 状态码),包含状态码与响应体。
#[error("LLM 请求失败(HTTP {status}: {body}。请检查 Provider 端点地址(base_url)和网络连通性")]
#[error(
"LLM 请求失败(HTTP {status}: {body}。请检查 Provider 端点地址(base_url)和网络连通性"
)]
Request { status: u16, body: String },
/// 请求超时。可重试。
@@ -30,10 +35,12 @@ pub enum LlmError {
Stream(String),
/// 上下文长度超出模型窗口限制。
#[error("LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compactionllm::compact")]
#[error(
"LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compactionllm::compact"
)]
ContextLength { actual: u32, limit: u32 },
/// 其他未分类的 LLM 调用失败。
#[error("LLM 调用失败: {0}")]
Other(String),
}
}
+2 -3
View File
@@ -7,6 +7,7 @@ use crate::llm::types::request_v2::MessageRequest;
/// 生命周期钩子事件点。
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum HookEvent {
/// LLM 请求发起之前(可阻断)。
PreRequest,
@@ -130,9 +131,7 @@ impl Default for HookExecutor {
impl HookExecutor {
/// 创建一个空的执行器。
pub fn new() -> Self {
Self {
hooks: Vec::new(),
}
Self { hooks: Vec::new() }
}
/// 注册一个钩子到指定事件点。
+11 -9
View File
@@ -97,8 +97,7 @@ impl LlmProvider for MockProvider {
async fn chat_stream(
&self,
_request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
let response = self.pop()?;
// 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。
let id = response.id.clone();
@@ -206,10 +205,7 @@ mod tests {
#[tokio::test]
async fn chat_returns_queued_response() {
let provider = MockProvider::new(vec![text_response("hello")]);
let resp = provider
.chat(MessageRequest::default())
.await
.unwrap();
let resp = provider.chat(MessageRequest::default()).await.unwrap();
assert_eq!(resp.text(), "hello");
assert_eq!(provider.remaining(), 0);
}
@@ -231,7 +227,10 @@ mod tests {
#[tokio::test]
async fn chat_stream_emits_text_delta_sequence() {
let provider = MockProvider::new(vec![text_response("hi")]);
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
let mut stream = provider
.chat_stream(MessageRequest::default())
.await
.unwrap();
let mut seen_start = false;
let mut seen_block_start = false;
@@ -283,7 +282,10 @@ mod tests {
extra: Default::default(),
};
let provider = MockProvider::new(vec![response]);
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
let mut stream = provider
.chat_stream(MessageRequest::default())
.await
.unwrap();
let mut saw_tool_args = false;
let mut saw_tool_end = false;
@@ -302,4 +304,4 @@ mod tests {
assert!(saw_tool_args);
assert!(saw_tool_end);
}
}
}
+456 -21
View File
@@ -1,11 +1,14 @@
pub mod anthropic;
pub mod ollama;
pub mod openai;
pub mod openai_compat;
pub mod registry;
use std::pin::Pin;
use std::time::Duration;
use futures_core::Stream;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use crate::llm::error::LlmError;
@@ -18,6 +21,7 @@ use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
/// 当前协议数量(5 种以内)完全可控,enum 的编译期安全检查优于运行时的 `HashMap::get()`。
/// 未来如果扩展到 15+ 种以上,再改为注册表模式。
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ProviderType {
/// OpenAI Chat Completions API(兼容 DeepSeek / Qwen 等 `/chat/completions` 端点)。
OpenaiChat,
@@ -29,6 +33,8 @@ pub enum ProviderType {
DeepSeek,
/// Qwen / 阿里云百炼(OpenAI-compatible `/chat/completions`)。
Qwen,
/// Ollama 本地推理(OpenAI-compatible `/chat/completions`,默认 `http://localhost:11434/v1`)。
Ollama,
}
impl std::str::FromStr for ProviderType {
@@ -41,47 +47,211 @@ impl std::str::FromStr for ProviderType {
"anthropic" | "claude" => Ok(ProviderType::Anthropic),
"deepseek" => Ok(ProviderType::DeepSeek),
"qwen" | "dashscope" | "tongyi" => Ok(ProviderType::Qwen),
"ollama" => Ok(ProviderType::Ollama),
_ => Err(format!("未知的 Provider 类型: {s}")),
}
}
}
/// Provider 构造参数 —— 通用 base_url + api_key + model。
/// Provider 构造参数 —— 通用 base_url + api_key + model + timeout / retry 配置
#[derive(Debug, Clone)]
pub struct ProviderConfig {
/// API base URL(如 `https://api.openai.com/v1`)。为空时由 Provider 选择默认值。
pub base_url: String,
/// API key。Ollama 等本地 Provider 可为空。
pub api_key: String,
/// 模型名(如 `gpt-4o` / `claude-sonnet-4-20250514`)。
pub model: String,
/// 请求超时秒数(默认 30)。应用于 Provider 的 HTTP Client 级别。
pub timeout_secs: u64,
/// 最大重试次数(默认 3)。
///
/// 当前此字段仅由 `from_env()` 采集,**实际重试逻辑由 `CycleConfig.retry.max_retries` 控制**。
/// 此处保留字段以与 Roadmap §Phase 5 Step 5.1 对齐;未来 Phase 6+ 可统一合并到 `CycleConfig`。
pub max_retries: u32,
}
impl Default for ProviderConfig {
fn default() -> Self {
Self {
base_url: String::new(),
api_key: String::new(),
model: String::new(),
timeout_secs: 30,
max_retries: 3,
}
}
}
impl ProviderConfig {
/// 从环境变量构造 `ProviderConfig`。
///
/// 必填变量:
/// - `{prefix}_BASE_URL`
/// - `{prefix}_API_KEY`
/// - `{prefix}_MODEL`
///
/// 可选变量(有默认值):
/// - `{prefix}_TIMEOUT_SECS`(默认 30,解析失败回退 30 并 warn)
/// - `{prefix}_MAX_RETRIES`(默认 3,解析失败回退 3 并 warn)
pub fn from_env(prefix: &str) -> Result<Self, String> {
let base_url = std::env::var(format!("{prefix}_BASE_URL"))
.map_err(|_| format!("{prefix}_BASE_URL 环境变量未设置"))?;
let api_key = std::env::var(format!("{prefix}_API_KEY"))
.map_err(|_| format!("{prefix}_API_KEY 环境变量未设置"))?;
let model = std::env::var(format!("{prefix}_MODEL"))
.map_err(|_| format!("{prefix}_MODEL 环境变量未设置"))?;
let timeout_secs = match std::env::var(format!("{prefix}_TIMEOUT_SECS")) {
Ok(v) => v.parse().unwrap_or_else(|_| {
tracing::warn!("{prefix}_TIMEOUT_SECS='{v}' 解析失败,使用默认值 30");
30
}),
Err(_) => 30,
};
let max_retries = match std::env::var(format!("{prefix}_MAX_RETRIES")) {
Ok(v) => v.parse().unwrap_or_else(|_| {
tracing::warn!("{prefix}_MAX_RETRIES='{v}' 解析失败,使用默认值 3");
3
}),
Err(_) => 3,
};
// ponytail: max_retries 当前仅采集,不传入 Provider。
// 实际重试由 CycleConfig.retry.max_retries 控制。
if max_retries != 3 {
tracing::warn!(
"ProviderConfig.max_retries={} 已采集但当前未生效;\
CycleConfig.retry.max_retries ",
max_retries,
);
}
Ok(Self {
base_url,
api_key,
model,
timeout_secs,
max_retries,
})
}
}
/// 构造带 timeout 的 `reqwest::Client`OpenAI-compatible 共享)。
fn build_client_with_timeout(timeout_secs: u64) -> Result<Client, LlmError> {
Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.build()
.map_err(|e| LlmError::Other(format!("创建 HTTP 客户端失败: {e}")))
}
/// 构造带 Anthropic 默认 headers + timeout 的 `reqwest::Client`。
///
/// Anthropic 由于需要保留 `x-api-key` / `anthropic-version` 默认 headers
/// 与 OpenAI-compatible 共享的 `build_client_with_timeout` 不同。
fn build_anthropic_client(api_key: &str, timeout_secs: u64) -> Result<Client, LlmError> {
use reqwest::header::{HeaderMap, HeaderValue};
let key_header = HeaderValue::from_str(api_key)
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
let version_header = HeaderValue::from_static("2023-06-01");
Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
headers.insert("anthropic-version", version_header);
headers
})
.build()
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))
}
/// Provider 工厂 —— exhaustive match 在编译期保证新 Provider 被注册。
///
/// `config.timeout_secs` 注入到 Provider 的 HTTP Client 超时配置。
/// 每个分支通过 `from_parts` (pub(crate)) 一次性构造,无冗余 client 创建。
pub fn create_provider(
provider_type: ProviderType,
config: ProviderConfig,
) -> Result<Box<dyn LlmProvider>, LlmError> {
match provider_type {
ProviderType::OpenaiChat => Ok(Box::new(openai::OpenaiChatProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::OpenaiChat => {
let client = build_client_with_timeout(config.timeout_secs)?;
Ok(Box::new(openai::OpenaiChatProvider(
openai::GenericOpenaiProvider::from_parts(
config.base_url,
config.api_key,
config.model,
"openai",
client,
Vec::new(),
config.timeout_secs,
),
)))
}
ProviderType::OpenaiResponse => Err(LlmError::Other(
"OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(),
)),
ProviderType::Anthropic => Ok(Box::new(anthropic::AnthropicProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::DeepSeek => Ok(Box::new(openai_compat::DeepSeekProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::Qwen => Ok(Box::new(openai_compat::QwenProvider::new(
config.base_url,
config.api_key,
config.model,
))),
ProviderType::Anthropic => {
let client = build_anthropic_client(&config.api_key, config.timeout_secs)?;
Ok(Box::new(anthropic::AnthropicProvider::from_parts(
config.base_url,
config.api_key,
config.model,
client,
config.timeout_secs,
)))
}
ProviderType::DeepSeek => {
let client = build_client_with_timeout(config.timeout_secs)?;
Ok(Box::new(openai_compat::DeepSeekProvider(
openai::GenericOpenaiProvider::from_parts(
config.base_url,
config.api_key,
config.model,
"deepseek",
client,
Vec::new(),
config.timeout_secs,
),
)))
}
ProviderType::Qwen => {
let client = build_client_with_timeout(config.timeout_secs)?;
Ok(Box::new(openai_compat::QwenProvider(
openai::GenericOpenaiProvider::from_parts(
config.base_url,
config.api_key,
config.model,
"qwen",
client,
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
config.timeout_secs,
),
)))
}
ProviderType::Ollama => {
let client = build_client_with_timeout(config.timeout_secs)?;
// ponytail: Ollama 默认 base_url 由 OllamaProvider 构造处理 —— 但 from_parts 不走
// OllamaProvider::new 的默认 URL 回退。这里保留 base_url(可能为空 → http://localhost:11434/v1)。
let base_url = if config.base_url.is_empty() {
"http://localhost:11434/v1".to_string()
} else {
config.base_url
};
Ok(Box::new(ollama::OllamaProvider(
openai::GenericOpenaiProvider::from_parts(
base_url,
config.api_key,
config.model,
"ollama",
client,
Vec::new(),
config.timeout_secs,
),
)))
}
}
}
@@ -142,3 +312,268 @@ pub trait LlmProvider: Send + Sync {
/// 返回 Provider 静态能力描述。
fn capabilities(&self) -> ProviderCapabilities;
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn provider_config_default_values() {
let config = ProviderConfig::default();
assert_eq!(config.timeout_secs, 30);
assert_eq!(config.max_retries, 3);
assert_eq!(config.base_url, "");
assert_eq!(config.api_key, "");
assert_eq!(config.model, "");
}
#[test]
fn provider_config_from_env_requires_all_three() {
// 使用 temp_env 移除所有相关变量,避免外部环境意外设置导致测试 flaky
temp_env::with_vars(
[
("TEST_PROVIDER_MISSING_BASE_URL", None::<&str>),
("TEST_PROVIDER_MISSING_API_KEY", None::<&str>),
("TEST_PROVIDER_MISSING_MODEL", None::<&str>),
("TEST_PROVIDER_MISSING_TIMEOUT_SECS", None::<&str>),
("TEST_PROVIDER_MISSING_MAX_RETRIES", None::<&str>),
],
|| {
let result = ProviderConfig::from_env("TEST_PROVIDER_MISSING");
assert!(result.is_err());
let msg = result.unwrap_err();
assert!(
msg.contains("TEST_PROVIDER_MISSING_BASE_URL"),
"error should mention missing var, got: {msg}"
);
},
);
}
#[test]
fn provider_config_from_env_uses_defaults_when_only_required_set() {
temp_env::with_vars(
[
("TEST_PROVIDER_BASE_URL", Some("http://localhost:11434/v1")),
("TEST_PROVIDER_API_KEY", Some("")),
("TEST_PROVIDER_MODEL", Some("llama3")),
],
|| {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
assert_eq!(config.base_url, "http://localhost:11434/v1");
assert_eq!(config.api_key, "");
assert_eq!(config.model, "llama3");
assert_eq!(config.timeout_secs, 30);
assert_eq!(config.max_retries, 3);
},
);
}
#[test]
fn provider_config_from_env_reads_custom_values() {
temp_env::with_vars(
[
("TEST_PROVIDER_BASE_URL", Some("http://x")),
("TEST_PROVIDER_API_KEY", Some("k")),
("TEST_PROVIDER_MODEL", Some("m")),
("TEST_PROVIDER_TIMEOUT_SECS", Some("60")),
("TEST_PROVIDER_MAX_RETRIES", Some("5")),
],
|| {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
assert_eq!(config.timeout_secs, 60);
assert_eq!(config.max_retries, 5);
},
);
}
#[test]
fn provider_config_from_env_falls_back_on_invalid_numbers() {
temp_env::with_vars(
[
("TEST_PROVIDER_BASE_URL", Some("http://x")),
("TEST_PROVIDER_API_KEY", Some("k")),
("TEST_PROVIDER_MODEL", Some("m")),
("TEST_PROVIDER_TIMEOUT_SECS", Some("not-a-number")),
("TEST_PROVIDER_MAX_RETRIES", Some("also-bad")),
],
|| {
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
// 解析失败回退默认值
assert_eq!(config.timeout_secs, 30);
assert_eq!(config.max_retries, 3);
},
);
}
/// Timeout 传导集成测试:构造 `ProviderConfig` timeout=1s
/// `create_provider` 注入 1s 超时 client,请求一个故意延迟 3s 的 mock server
/// 验证返回 `LlmError::Timeout { duration: 1s }`。
#[tokio::test]
async fn create_provider_injects_timeout_into_openai_chat() {
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
// 故意延迟 3s 触发超时
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(3))
.set_body_json(json!({
"id": "x",
"object": "chat.completion",
"created": 0,
"model": "gpt-4o",
"choices": [],
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
})),
)
.mount(&server)
.await;
let provider = create_provider(
ProviderType::OpenaiChat,
ProviderConfig {
base_url: server.uri(),
api_key: "sk-test".into(),
model: "gpt-4o".into(),
timeout_secs: 1,
max_retries: 3,
},
)
.unwrap();
let err = provider
.chat(crate::llm::types::request_v2::MessageRequest {
model: "gpt-4o".into(),
messages: vec![],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Timeout { duration } => {
assert_eq!(duration, Duration::from_secs(1));
}
other => panic!("expected Timeout, got {other:?}"),
}
}
/// Timeout 传导验证:`create_provider` 生成的 DeepSeek Provider 也带 1s 超时,
/// 错误消息中的 duration 与 timeout_secs 一致(而非硬编码 120s)。
#[tokio::test]
async fn create_provider_injects_timeout_into_deepseek() {
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(3))
.set_body_json(json!({
"id": "x",
"object": "chat.completion",
"created": 0,
"model": "deepseek-chat",
"choices": [],
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
})),
)
.mount(&server)
.await;
let provider = create_provider(
ProviderType::DeepSeek,
ProviderConfig {
base_url: server.uri(),
api_key: "sk-test".into(),
model: "deepseek-chat".into(),
timeout_secs: 1,
max_retries: 3,
},
)
.unwrap();
let err = provider
.chat(crate::llm::types::request_v2::MessageRequest {
model: "deepseek-chat".into(),
messages: vec![],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Timeout { duration } => {
assert_eq!(duration, Duration::from_secs(1));
}
other => panic!("expected Timeout, got {other:?}"),
}
}
/// Timeout 传导验证:`create_provider` 生成的 Anthropic Provider 通过 `with_timeout`
/// 注入 1s 超时。
///
/// 与 OpenAI-compatible 路径不同,Anthropic 走 `AnthropicProvider::with_timeout()`
/// 重建底层 client(保留 default_headers),独立于 OpenAI-compatible 的 `build_client_with_timeout`。
/// 单独覆盖此路径以验证 `with_timeout` 不会因服务端延迟而返回硬编码 120s 的超时错误。
#[tokio::test]
async fn create_provider_injects_timeout_into_anthropic() {
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
// Anthropic Messages API 端点:`POST /v1/messages`
Mock::given(method("POST"))
.and(path("/v1/messages"))
.respond_with(
ResponseTemplate::new(200)
.set_delay(Duration::from_secs(3))
.set_body_json(json!({
"id": "msg_timeout_test",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "ok"}],
"model": "claude-sonnet-4-20250514",
"stop_reason": "end_turn",
"usage": {"input_tokens": 1, "output_tokens": 1}
})),
)
.mount(&server)
.await;
let provider = create_provider(
ProviderType::Anthropic,
ProviderConfig {
base_url: server.uri(),
api_key: "sk-ant-test".into(),
model: "claude-sonnet-4-20250514".into(),
timeout_secs: 1,
max_retries: 3,
},
)
.unwrap();
let err = provider
.chat(crate::llm::types::request_v2::MessageRequest {
model: "claude-sonnet-4-20250514".into(),
messages: vec![],
..Default::default()
})
.await
.unwrap_err();
match err {
LlmError::Timeout { duration } => {
assert_eq!(duration, Duration::from_secs(1));
}
other => panic!("expected Timeout, got {other:?}"),
}
}
}
+132 -55
View File
@@ -12,10 +12,10 @@ use async_trait::async_trait;
use bytes::Bytes;
use futures_core::Stream;
use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderValue};
use reqwest::Client;
use reqwest::header::{HeaderMap, HeaderValue};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use serde_json::{Value, json};
use tracing::{debug, error, info, warn};
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
@@ -39,16 +39,20 @@ pub struct AnthropicProvider {
#[allow(dead_code)]
api_key: String,
model: String,
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter
/// 因此单独存储以便错误消息与配置保持一致。
timeout_secs: u64,
}
impl AnthropicProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
let key_header = HeaderValue::from_str(&api_key)
.expect("Anthropic API key 包含无效的 HTTP 头部字符");
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let key_header =
HeaderValue::from_str(&api_key).expect("Anthropic API key 包含无效的 HTTP 头部字符");
let version_header = HeaderValue::from_static("2023-06-01");
let http_client = Client::builder()
.timeout(Duration::from_secs(120))
.timeout(Duration::from_secs(timeout_secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
@@ -67,14 +71,80 @@ impl AnthropicProvider {
},
api_key,
model,
timeout_secs,
}
}
/// ⚠️ 替换 HTTP Client**丢弃** `new()` 中设置的默认 headers`x-api-key` / `anthropic-version`)。
///
/// 调用此方法后,所有 Anthropic API 请求将以**无认证头**发送出去,预期会 401/403 失败。
/// 推荐改用 [`Self::with_timeout`],它会重建 client 并保留默认 headers。
///
/// 此方法仍保留以兼容调用方自定义 client 但不需要默认 headers 的极端场景。
#[deprecated(
since = "0.2.0",
note = "此方法会丢弃默认 headersx-api-key / anthropic-version),改为使用 `with_timeout` 或带 headers 的 `Client::builder()`"
)]
pub fn with_client(mut self, client: Client) -> Self {
self.http_client = client;
self
}
/// 替换 HTTP Client 的超时配置(重建底层 client,保留默认 headers)。
///
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后通过 `with_client` 注入的 Client
/// 将被替换。headers 构造逻辑与 `new()` 中的保持一致(`x-api-key` / `anthropic-version`)。
///
/// ponytail: 同值调用短路。当 `secs == self.timeout_secs` 时跳过 client 重建,
/// 避免 `create_provider` 路径 `new(timeout).with_timeout(timeout)` 的双重构造。
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
if secs == self.timeout_secs {
return Ok(self);
}
// ponytail: 重建 http_client 时保留已有默认 headersx-api-key / anthropic-version)。
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
let key_header = HeaderValue::from_str(&self.api_key)
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
let version_header = HeaderValue::from_static("2023-06-01");
self.http_client = Client::builder()
.timeout(Duration::from_secs(secs))
.default_headers({
let mut headers = HeaderMap::new();
headers.insert("x-api-key", key_header);
headers.insert("anthropic-version", version_header);
headers
})
.build()
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
self.timeout_secs = secs;
Ok(self)
}
/// 一次性构造 —— `create_provider` 路径专用,避免 `new(...)` + `with_timeout(...)` 的双重 client 构造。
///
/// 调用方负责预先构造好符合 Anthropic 协议要求的 `http_client`(带正确的 `x-api-key` /
/// `anthropic-version` 默认 headers + 指定 timeout)。
pub(crate) fn from_parts(
base_url: String,
api_key: String,
model: String,
http_client: Client,
timeout_secs: u64,
) -> Self {
Self {
http_client,
base_url: if base_url.is_empty() {
"https://api.anthropic.com".to_string()
} else {
base_url
},
api_key,
model,
timeout_secs,
}
}
fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
}
@@ -104,12 +174,14 @@ impl AnthropicProvider {
Message::User { content } => {
api_messages.push(AnthropicMessage::user(content));
}
Message::UserImage { data, mime_type, detail } => {
Message::UserImage {
data,
mime_type,
detail,
} => {
// Anthropic image format: {type: "image", source: {type: "base64", media_type, data}}
let source = if data.starts_with("http://") || data.starts_with("https://") {
AnthropicImageSource::Url {
url: data.clone(),
}
AnthropicImageSource::Url { url: data.clone() }
} else {
AnthropicImageSource::Base64 {
media_type: mime_type.clone(),
@@ -194,7 +266,7 @@ impl AnthropicProvider {
.json(&body)
.send()
.await
.map_err(Self::map_reqwest_error)?;
.map_err(|e| self.map_reqwest_error(e))?;
let status = response.status();
if !status.is_success() {
@@ -215,8 +287,7 @@ impl AnthropicProvider {
async fn chat_stream_inner(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
let mut body = self.build_request_body(request)?;
body.stream = Some(true);
@@ -230,16 +301,16 @@ impl AnthropicProvider {
.json(&body)
.send()
.await
.map_err(Self::map_reqwest_error)?;
.map_err(|e| self.map_reqwest_error(e))?;
let status = response.status();
if !status.is_success() {
return Err(Self::handle_error_response(response).await);
}
let byte_stream = response.bytes_stream().map(|r| {
r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}")))
});
let byte_stream = response
.bytes_stream()
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}"))));
let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> =
Box::pin(byte_stream);
@@ -247,10 +318,10 @@ impl AnthropicProvider {
Ok(Box::pin(AnthropicSseStream::new(byte_stream)))
}
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
if e.is_timeout() {
LlmError::Timeout {
duration: Duration::from_secs(120),
duration: Duration::from_secs(self.timeout_secs),
}
} else if e.is_connect() {
LlmError::Other(format!("连接失败: {e}"))
@@ -291,13 +362,12 @@ impl AnthropicProvider {
blocks.push(ContentBlock::Text { text });
}
AnthropicContentBlockResp::ToolUse { id, name, input } => {
blocks.push(ContentBlock::ToolUse {
id,
name,
input,
});
blocks.push(ContentBlock::ToolUse { id, name, input });
}
AnthropicContentBlockResp::Thinking { thinking, signature } => {
AnthropicContentBlockResp::Thinking {
thinking,
signature,
} => {
blocks.push(ContentBlock::Thinking {
text: thinking,
signature,
@@ -336,8 +406,7 @@ impl LlmProvider for AnthropicProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.chat_stream_inner(request).await
}
@@ -418,7 +487,9 @@ impl AnthropicMessage {
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentPart {
Text { text: String },
Text {
text: String,
},
Image {
source: AnthropicImageSource,
},
@@ -437,13 +508,8 @@ enum AnthropicContentPart {
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicImageSource {
Base64 {
media_type: String,
data: String,
},
Url {
url: String,
},
Base64 { media_type: String, data: String },
Url { url: String },
}
fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> {
@@ -523,9 +589,18 @@ struct AnthropicUsage {
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum AnthropicContentBlockResp {
Text { text: String },
ToolUse { id: String, name: String, input: Value },
Thinking { thinking: String, signature: Option<String> },
Text {
text: String,
},
ToolUse {
id: String,
name: String,
input: Value,
},
Thinking {
thinking: String,
signature: Option<String>,
},
}
// =============================================================================
@@ -662,18 +737,13 @@ impl AnthropicSseStream {
// 先把所有字段提前,避免 match 中 part-move
let block_type = match &content_block {
AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text,
AnthropicContentBlockStart::ToolUse { id, name } => {
ContentBlockType::ToolUse {
id: id.clone(),
name: name.clone(),
}
}
AnthropicContentBlockStart::ToolUse { id, name } => ContentBlockType::ToolUse {
id: id.clone(),
name: name.clone(),
},
AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking,
};
events.push(StreamEvent::ContentBlockStart {
index,
block_type,
});
events.push(StreamEvent::ContentBlockStart { index, block_type });
let builder = match content_block {
AnthropicContentBlockStart::Text { text } => {
crate::llm::types::response_v2::ContentBlockBuilder::Text(text)
@@ -742,7 +812,9 @@ impl AnthropicSseStream {
completion_tokens_details: None,
prompt_tokens_details: None,
};
events.push(StreamEvent::CostUpdate { usage: partial_usage });
events.push(StreamEvent::CostUpdate {
usage: partial_usage,
});
}
}
AnthropicSseEvent::MessageStop => {
@@ -753,7 +825,9 @@ impl AnthropicSseStream {
self.saw_terminal = true;
match self.partial.clone().finalize() {
Ok(full) => {
events.push(StreamEvent::MessageComplete { full_response: full });
events.push(StreamEvent::MessageComplete {
full_response: full,
});
}
Err(e) => {
events.push(StreamEvent::Error {
@@ -780,10 +854,7 @@ fn _unused_marker() {}
impl Stream for AnthropicSseStream {
type Item = Result<StreamEvent, LlmError>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
if let Some(data) = self.next_event_line() {
let mut events = self.handle_event_json(&data);
@@ -835,7 +906,12 @@ mod tests {
fn make_provider(base_url: String) -> AnthropicProvider {
// 跳过默认 header 注入:测试用自定义 base_url 直接 mock
AnthropicProvider::new(base_url, "sk-ant-test".into(), "claude-sonnet-4-20250514".into())
AnthropicProvider::new(
base_url,
"sk-ant-test".into(),
"claude-sonnet-4-20250514".into(),
30,
)
}
#[tokio::test]
@@ -979,6 +1055,7 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
"http://x".into(),
"k".into(),
"claude-sonnet-4-20250514".into(),
30,
)
.capabilities();
assert_eq!(caps.provider_name, "anthropic");
+72
View File
@@ -0,0 +1,72 @@
//! Ollama Provider —— OpenAI-compatible 协议的 newtype 包装,零 API key。
//!
//! 默认 base_url = `http://localhost:11434/v1`,空 api_key 也可工作。
//! 实现方式同 `DeepSeekProvider` / `QwenProvider`,共享 `GenericOpenaiProvider`
//! 的 HTTP/SSE/转换逻辑,仅配置不同。
use std::pin::Pin;
use async_trait::async_trait;
use futures_core::Stream;
use reqwest::Client;
use super::openai::GenericOpenaiProvider;
use super::{LlmProvider, ProviderCapabilities};
use crate::llm::error::LlmError;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
/// Ollama 本地 Provider —— OpenAI-compatible 协议的 newtype 包装。
///
/// Ollama 在 `localhost:11434` 暴露与 OpenAI 兼容的 `/v1/chat/completions`
/// 接口,因此完全复用 `GenericOpenaiProvider` 的实现。允许空 `api_key`。
pub struct OllamaProvider(pub GenericOpenaiProvider);
impl OllamaProvider {
/// 构造 Ollama Provider。
///
/// - `base_url` 为空时使用默认 `http://localhost:11434/v1`
/// - `api_key` 可为空字符串(Ollama 不校验)
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let url = if base_url.is_empty() {
"http://localhost:11434/v1".to_string()
} else {
base_url
};
Self(GenericOpenaiProvider::new_with_name(
url,
api_key,
model,
"ollama",
timeout_secs,
))
}
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
///
/// 与 `OpenaiChatProvider::with_client`、`DeepSeekProvider::with_client`、
/// `QwenProvider::with_client` 签名一致。
pub fn with_client(self, client: Client) -> Self {
Self(self.0.with_client(client))
}
}
#[async_trait]
impl LlmProvider for OllamaProvider {
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
self.0.chat(request).await
}
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
fn capabilities(&self) -> ProviderCapabilities {
let mut caps = self.0.capabilities();
caps.provider_name = "ollama";
caps
}
}
+93 -56
View File
@@ -52,31 +52,63 @@ pub struct GenericOpenaiProvider {
model: String,
provider_name: &'static str,
extra_headers: Vec<(String, String)>,
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter
/// 因此单独存储以便错误消息与配置保持一致。
timeout_secs: u64,
}
impl GenericOpenaiProvider {
/// 基础构造
pub fn new_with_name(
/// 一次性构造 —— `create_provider` 路径专用,避免 `new_with_name` + `with_client` 的双重 client 构造。
///
/// 调用方负责预先构造好带正确 timeout 的 `http_client`。`extra_headers` 与 `timeout_secs`
/// 一并设置字段,避免后续修改。
pub(crate) fn from_parts(
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
http_client: Client,
extra_headers: Vec<(String, String)>,
timeout_secs: u64,
) -> Self {
let http_client = Client::builder()
.timeout(Duration::from_secs(120))
.build()
.expect("创建 HTTP 客户端失败");
Self {
http_client,
base_url,
api_key,
model,
provider_name,
extra_headers: Vec::new(),
extra_headers,
timeout_secs,
}
}
/// 基础构造器。
///
/// `timeout_secs` 应用于 `reqwest::Client` 的请求超时配置。
/// 应由 `ProviderConfig::timeout_secs` 传入(调用方如不知道,可传 30)。
pub fn new_with_name(
base_url: String,
api_key: String,
model: String,
provider_name: &'static str,
timeout_secs: u64,
) -> Self {
let http_client = Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.build()
.expect("创建 HTTP 客户端失败");
Self::from_parts(
base_url,
api_key,
model,
provider_name,
http_client,
Vec::new(),
timeout_secs,
)
}
/// 带额外请求头的构造器(如 Qwen 需要 SSE 启用头)。
pub fn new_with_name_and_headers(
base_url: String,
@@ -84,10 +116,21 @@ impl GenericOpenaiProvider {
model: String,
provider_name: &'static str,
extra_headers: Vec<(String, String)>,
timeout_secs: u64,
) -> Self {
let mut base = Self::new_with_name(base_url, api_key, model, provider_name);
base.extra_headers = extra_headers;
base
let http_client = Client::builder()
.timeout(Duration::from_secs(timeout_secs))
.build()
.expect("创建 HTTP 客户端失败");
Self::from_parts(
base_url,
api_key,
model,
provider_name,
http_client,
extra_headers,
timeout_secs,
)
}
pub fn with_client(mut self, client: Client) -> Self {
@@ -119,10 +162,10 @@ impl GenericOpenaiProvider {
Ok(builder.json(body))
}
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
if e.is_timeout() {
LlmError::Timeout {
duration: Duration::from_secs(120),
duration: Duration::from_secs(self.timeout_secs),
}
} else if e.is_connect() {
LlmError::Other(format!("连接失败: {}", e))
@@ -135,9 +178,7 @@ impl GenericOpenaiProvider {
///
/// ponytail: Qwen 等部分 OpenAI-compatible 提供方可能返回非标准 error body
/// (无法解析为 JSON),此处直接用 status code + 原始 body 兜底。
async fn handle_error_response(
response: reqwest::Response,
) -> LlmError {
async fn handle_error_response(response: reqwest::Response) -> LlmError {
let status = response.status().as_u16();
let body = response.text().await.unwrap_or_default();
@@ -148,10 +189,7 @@ impl GenericOpenaiProvider {
// 与 OpenAI 完全一致;DeepSeek/Qwen 通常遵循。
LlmError::RateLimit { retry_after: None }
}
_ if status >= 500 => LlmError::Request {
status,
body,
},
_ if status >= 500 => LlmError::Request { status, body },
_ if status == 400 && body.contains("context_length_exceeded") => {
LlmError::ContextLength {
actual: 0,
@@ -188,7 +226,7 @@ impl GenericOpenaiProvider {
Some(
tool_defs
.into_iter()
.map(|t| OpenaiTool::Function { function: t })
.map(|t| OpenaiTool::Function { function: t.into() })
.collect(),
)
};
@@ -200,7 +238,9 @@ impl GenericOpenaiProvider {
stop_sequences[0].clone(),
))
} else {
Some(crate::llm::types::shared::StopSequence::Multiple(stop_sequences))
Some(crate::llm::types::shared::StopSequence::Multiple(
stop_sequences,
))
};
let frequency_penalty = request.get_extra_opt("frequency_penalty");
@@ -245,9 +285,7 @@ impl GenericOpenaiProvider {
let stop_reason = match choice.finish_reason {
Some(FinishReason::Stop) => StopReason::Stop,
Some(FinishReason::Length) => StopReason::Length,
Some(FinishReason::ToolCalls) | Some(FinishReason::FunctionCall) => {
StopReason::ToolUse
}
Some(FinishReason::ToolCalls) | Some(FinishReason::FunctionCall) => StopReason::ToolUse,
Some(FinishReason::ContentFilter) => StopReason::ContentFilter,
Some(FinishReason::Other) | None => StopReason::Stop,
};
@@ -263,7 +301,10 @@ impl GenericOpenaiProvider {
}
/// 非流式 `chat()` 入口。
pub async fn chat_blocking(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
pub async fn chat_blocking(
&self,
request: MessageRequest,
) -> Result<MessageResponse, LlmError> {
let req = self.convert_request(request)?;
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
@@ -280,7 +321,7 @@ impl GenericOpenaiProvider {
.await
.map_err(|e| {
error!(error = %e, "请求失败");
Self::map_reqwest_error(e)
self.map_reqwest_error(e)
})?;
let status = response.status();
@@ -303,8 +344,7 @@ impl GenericOpenaiProvider {
pub async fn chat_stream_inner(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
let mut req = self.convert_request(request)?;
req.stream = Some(true);
req.stream_options = Some(StreamOptions {
@@ -322,7 +362,7 @@ impl GenericOpenaiProvider {
.await
.map_err(|e| {
error!(error = %e, "流式请求失败");
Self::map_reqwest_error(e)
self.map_reqwest_error(e)
})?;
let status = response.status();
@@ -331,9 +371,9 @@ impl GenericOpenaiProvider {
}
let byte_stream: std::pin::Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> = {
let s = response.bytes_stream().map(|r| {
r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e)))
});
let s = response
.bytes_stream()
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {}", e))));
Box::pin(s)
};
@@ -368,8 +408,7 @@ impl LlmProvider for GenericOpenaiProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.chat_stream_inner(request).await
}
@@ -389,12 +428,13 @@ impl LlmProvider for GenericOpenaiProvider {
pub struct OpenaiChatProvider(pub GenericOpenaiProvider);
impl OpenaiChatProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
Self(GenericOpenaiProvider::new_with_name(
base_url,
api_key,
model,
"openai",
timeout_secs,
))
}
@@ -412,8 +452,7 @@ impl LlmProvider for OpenaiChatProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
@@ -436,7 +475,10 @@ enum BlockState {
/// 正在累积 text block。
InText { block_index: u32 },
/// 正在累积 tool_call block。
InTool { block_index: u32, tool_call_index: u32 },
InTool {
block_index: u32,
tool_call_index: u32,
},
/// 正在累积 refusal block。
InRefusal { block_index: u32 },
}
@@ -456,9 +498,7 @@ pub struct ChunkToEventStream {
}
impl ChunkToEventStream {
fn new(
chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>,
) -> Self {
fn new(chunks: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>>) -> Self {
Self {
chunks,
buffer: String::new(),
@@ -500,12 +540,7 @@ impl ChunkToEventStream {
let mut events = Vec::new();
// 元信息:MessageStart(仅在第一次见到 role=assistant 时)。
if self.partial.id.is_none()
&& chunk
.choices
.iter()
.any(|c| c.delta.role.is_some())
{
if self.partial.id.is_none() && chunk.choices.iter().any(|c| c.delta.role.is_some()) {
events.push(StreamEvent::MessageStart {
id: chunk.id.clone(),
model: chunk.model.clone(),
@@ -650,9 +685,7 @@ impl ChunkToEventStream {
self.partial.stop_reason = Some(match fr {
FinishReason::Stop => StopReason::Stop,
FinishReason::Length => StopReason::Length,
FinishReason::ToolCalls | FinishReason::FunctionCall => {
StopReason::ToolUse
}
FinishReason::ToolCalls | FinishReason::FunctionCall => StopReason::ToolUse,
FinishReason::ContentFilter => StopReason::ContentFilter,
FinishReason::Other => StopReason::Other,
});
@@ -692,7 +725,9 @@ impl ChunkToEventStream {
match self.partial.clone().finalize() {
Ok(full) => {
events.push(StreamEvent::MessageComplete { full_response: full });
events.push(StreamEvent::MessageComplete {
full_response: full,
});
}
Err(e) => {
events.push(StreamEvent::Error {
@@ -707,10 +742,7 @@ impl ChunkToEventStream {
impl Stream for ChunkToEventStream {
type Item = Result<StreamEvent, LlmError>;
fn poll_next(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
// 先尝试从 buffer 取一行处理
if let Some(line) = self.next_line() {
@@ -814,6 +846,7 @@ mod tests {
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let response = provider
.chat_blocking(MessageRequest {
@@ -842,6 +875,7 @@ mod tests {
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let err = provider
.chat_blocking(MessageRequest {
@@ -868,6 +902,7 @@ mod tests {
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let err = provider
.chat_blocking(MessageRequest {
@@ -906,6 +941,7 @@ data: [DONE]\n\n";
"sk-test".into(),
"gpt-4o".into(),
"openai",
30,
);
let mut stream = provider
.chat_stream_inner(MessageRequest {
@@ -974,6 +1010,7 @@ data: [DONE]\n\n";
"k".into(),
"gpt-4o".into(),
"openai",
30,
);
let ir = provider.convert_response(resp).unwrap();
assert_eq!(ir.stop_reason, StopReason::ToolUse);
+32 -39
View File
@@ -15,12 +15,12 @@ use std::pin::Pin;
use async_trait::async_trait;
use futures_core::Stream;
use super::openai::GenericOpenaiProvider;
use super::ProviderCapabilities;
use super::openai::GenericOpenaiProvider;
use crate::llm::error::LlmError;
use crate::llm::provider::LlmProvider;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
use crate::llm::provider::LlmProvider;
// =============================================================================
// DeepSeek
@@ -29,7 +29,7 @@ use crate::llm::provider::LlmProvider;
pub struct DeepSeekProvider(pub GenericOpenaiProvider);
impl DeepSeekProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let url = if base_url.is_empty() {
"https://api.deepseek.com".to_string()
} else {
@@ -40,26 +40,27 @@ impl DeepSeekProvider {
api_key,
model,
"deepseek",
timeout_secs,
))
}
}
impl DeepSeekProvider {
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
pub fn with_client(self, client: reqwest::Client) -> Self {
Self(self.0.with_client(client))
}
/// 测试中(带 mock_client)使用的构造器。
///
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
/// 无超时),不影响行为。
pub fn new_with_client(
base_url: String,
api_key: String,
model: String,
client: reqwest::Client,
) -> Self {
let url = if base_url.is_empty() {
"https://api.deepseek.com".to_string()
} else {
base_url
};
let mut inner = GenericOpenaiProvider::new_with_name(url, api_key, model, "deepseek");
inner.http_client = client;
Self(inner)
Self::new(base_url, api_key, model, 30).with_client(client)
}
}
@@ -72,8 +73,7 @@ impl LlmProvider for DeepSeekProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
@@ -91,7 +91,7 @@ impl LlmProvider for DeepSeekProvider {
pub struct QwenProvider(pub GenericOpenaiProvider);
impl QwenProvider {
pub fn new(base_url: String, api_key: String, model: String) -> Self {
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
let url = if base_url.is_empty() {
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
} else {
@@ -104,31 +104,28 @@ impl QwenProvider {
model,
"qwen",
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
timeout_secs,
);
Self(inner)
}
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
pub fn with_client(self, client: reqwest::Client) -> Self {
Self(self.0.with_client(client))
}
/// 测试构造器。
///
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
/// 无超时),不影响行为。
pub fn new_with_client(
base_url: String,
api_key: String,
model: String,
client: reqwest::Client,
) -> Self {
let url = if base_url.is_empty() {
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
} else {
base_url
};
let mut inner = GenericOpenaiProvider::new_with_name_and_headers(
url,
api_key,
model,
"qwen",
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
);
inner.http_client = client;
Self(inner)
Self::new(base_url, api_key, model, 30).with_client(client)
}
}
@@ -141,8 +138,7 @@ impl LlmProvider for QwenProvider {
async fn chat_stream(
&self,
request: MessageRequest,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
{
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
self.0.chat_stream(request).await
}
@@ -156,8 +152,8 @@ impl LlmProvider for QwenProvider {
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::types::request_v2::MessageRequest;
use crate::llm::types::message::Message as IrMessage;
use crate::llm::types::request_v2::MessageRequest;
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
@@ -182,11 +178,8 @@ mod tests {
.mount(&server)
.await;
let provider = DeepSeekProvider::new(
server.uri(),
"sk-test".into(),
"deepseek-chat".into(),
);
let provider =
DeepSeekProvider::new(server.uri(), "sk-test".into(), "deepseek-chat".into(), 30);
let response = provider
.chat(MessageRequest {
model: "deepseek-chat".into(),
@@ -219,7 +212,7 @@ mod tests {
.mount(&server)
.await;
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into());
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into(), 30);
let response = provider
.chat(MessageRequest {
model: "qwen-plus".into(),
+2 -4
View File
@@ -3,7 +3,7 @@
use std::collections::HashMap;
use crate::llm::error::LlmError;
use crate::llm::provider::{create_provider, LlmProvider, ProviderConfig, ProviderType};
use crate::llm::provider::{LlmProvider, ProviderConfig, ProviderType, create_provider};
/// Provider 注册表 —— 管理多个 LLM Provider 实例。
///
@@ -61,8 +61,6 @@ impl ProviderRegistry {
/// 获取默认 Provider。
pub fn get_default(&self) -> Option<&dyn LlmProvider> {
self.default_name
.as_ref()
.and_then(|name| self.get(name))
self.default_name.as_ref().and_then(|name| self.get(name))
}
}
+7 -8
View File
@@ -14,8 +14,8 @@ use std::pin::Pin;
use std::task::{Context, Poll};
use futures_core::stream::Stream;
use futures_util::future::poll_fn;
use futures_util::FutureExt;
use futures_util::future::poll_fn;
use serde_json::Value;
use crate::llm::error::LlmError;
@@ -95,9 +95,7 @@ impl Stream for ChunkToLegacyEventStream {
}
if let Some(usage) = &chunk.usage {
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate {
usage: *usage,
}));
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate { usage: *usage }));
}
Poll::Ready(None)
@@ -143,9 +141,7 @@ fn empty_message_response() -> MessageResponse {
MessageResponse {
id: String::new(),
model: String::new(),
message: Message::Assistant {
content: vec![],
},
message: Message::Assistant { content: vec![] },
usage: Usage::default(),
stop_reason: StopReason::Stop,
extra: HashMap::new(),
@@ -172,7 +168,10 @@ fn map_legacy_to_ir(legacy: LegacyStreamEvent) -> StreamEvent {
LegacyStreamEvent::AssistantTextDelta { text } => StreamEvent::TextDelta { text },
LegacyStreamEvent::ToolExecutionStarted { input, .. } => {
let arguments = serde_json::to_string(&input).unwrap_or_default();
StreamEvent::ToolCallArgumentsDelta { index: 0, arguments }
StreamEvent::ToolCallArgumentsDelta {
index: 0,
arguments,
}
}
LegacyStreamEvent::ToolExecutionCompleted { .. } => {
// 旧 ToolExecutionCompleted 不在 IR 流协议中——工具执行是消费方职责。
+9 -19
View File
@@ -20,15 +20,12 @@ use crate::llm::types::shared::ImageDetail;
/// 消费方 match 可直接区分文本和图片输入。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Message {
/// 系统提示(User & Assistant 之外的引导指令)。
System {
content: Vec<ContentBlock>,
},
System { content: Vec<ContentBlock> },
/// 用户输入。
User {
content: Vec<ContentBlock>,
},
User { content: Vec<ContentBlock> },
/// 用户的图片输入(快捷构造,免去构造 ContentBlock 的 boilerplate)。
UserImage {
data: String,
@@ -36,9 +33,7 @@ pub enum Message {
detail: ImageDetail,
},
/// Assistant 回复内容块(可能包含 text、thinking、tool_use 等多种 block 的混合)。
Assistant {
content: Vec<ContentBlock>,
},
Assistant { content: Vec<ContentBlock> },
/// 工具调用结果。
ToolResult {
tool_call_id: String,
@@ -103,6 +98,7 @@ impl Message {
/// block 的逃生舱。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ContentBlock {
/// 纯文本。
Text { text: String },
@@ -130,10 +126,7 @@ pub enum ContentBlock {
signature: Option<String>,
},
/// 逃生舱:Provider 特定 block 透传(OpenAI Response 内置工具等)。
Extension {
kind: String,
data: Value,
},
Extension { kind: String, data: Value },
}
/// 内容块类型标签 —— 用于 `StreamEvent::ContentBlockStart.block_type`。
@@ -141,6 +134,7 @@ pub enum ContentBlock {
/// 用途:在流式场景中,Provider 先下发 block 类型,再下发 block 内容增量。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ContentBlockType {
/// 文本块。
Text,
@@ -349,9 +343,7 @@ mod tests {
fn message_roundtrip_each_variant() {
let msgs = vec![
Message::System {
content: vec![ContentBlock::Text {
text: "sys".into(),
}],
content: vec![ContentBlock::Text { text: "sys".into() }],
},
Message::User {
content: vec![ContentBlock::Text {
@@ -376,9 +368,7 @@ mod tests {
},
Message::ToolResult {
tool_call_id: "call_1".into(),
content: vec![ContentBlock::Text {
text: "ok".into(),
}],
content: vec![ContentBlock::Text { text: "ok".into() }],
is_error: true,
},
];
+1 -5
View File
@@ -26,7 +26,7 @@ pub use shared::{
AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier,
StopSequence,
};
pub use tool::{FunctionCall, OpenaiToolCall, OpenaiToolDefinition};
pub use tool::{FunctionCall, OpenaiToolCall, ToolDef};
pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage};
// Re-export IR 内容块 / 消息类型供 `types::ContentBlock` 等历史路径消费。
@@ -96,7 +96,3 @@ impl From<ChatResponse> for OpenaiChatChunk {
}
}
}
/// 工具定义别名(无新类型冲突,保留)。
#[deprecated(since = "0.1.0", note = "ToolDefinition 仍直接对应 OpenAI wire-format;未来 v0.2 引入 IR 工具类型后会再次更新")]
pub type ToolDefinition = OpenaiToolDefinition;
+8 -5
View File
@@ -11,18 +11,21 @@ pub struct StreamOptions {
pub include_obfuscation: Option<bool>,
}
#[derive(Debug, Clone)]
#[derive(Default)]
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub enum ToolChoice {
#[default]
None,
Auto,
Required,
Named { name: String },
AllowedTools { tool_names: Vec<String> },
Named {
name: String,
},
AllowedTools {
tool_names: Vec<String>,
},
}
impl Serialize for ToolChoice {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
+48 -16
View File
@@ -10,14 +10,14 @@ use thiserror::Error;
use crate::llm::types::message::Message;
use crate::llm::types::request::ToolChoice;
use crate::llm::types::tool::OpenaiToolDefinition;
use crate::llm::types::tool::ToolDef;
/// Provider 无关的请求类型。
///
/// 设计要点:
/// - `system` 字段不存在;system 提示由调用方通过 `Message::System` 在 `messages` 中表达。
/// - `tools` / `tool_choice` 直接复用现有 `OpenaiToolDefinition` / `ToolChoice`
/// (10a §251 决策:先复用旧类型,Phase 2 切换为新 `ToolDefinition` 后再调整)
/// - `tools` 使用 Provider 无关的 `ToolDef` IR;各 Provider 适配层在 `convert_request`
/// 中转换为对应 wire format。`tool_choice` 复用现有 `ToolChoice`
/// - `extra` 作为逃生舱:Provider 特定字段(`web_search_options`、`previous_response_id` 等)
/// 通过 `extra.set_extra / get_extra` 传递,避免持续膨胀本结构体。
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
@@ -26,8 +26,8 @@ pub struct MessageRequest {
pub model: String,
/// 消息列表(包含 system / user / assistant / tool_result 等所有变体)。
pub messages: Vec<Message>,
/// 工具定义列表。
pub tools: Vec<OpenaiToolDefinition>,
/// 工具定义列表Provider 无关 IR
pub tools: Vec<ToolDef>,
/// 工具选择策略。
pub tool_choice: ToolChoice,
/// 最大输出 token 数。
@@ -127,14 +127,9 @@ mod tests {
#[test]
fn extra_set_and_get_roundtrip() {
let mut req = MessageRequest::default();
req.set_extra(
"previous_response_id",
"resp_abc123",
);
req.set_extra("previous_response_id", "resp_abc123");
let v: Option<String> = req
.get_extra("previous_response_id")
.expect("get_extra ok");
let v: Option<String> = req.get_extra("previous_response_id").expect("get_extra ok");
assert_eq!(v.as_deref(), Some("resp_abc123"));
let missing: Option<String> = req.get_extra("missing").expect("missing ok");
@@ -174,10 +169,7 @@ mod tests {
}
let opts: Options = req.get_extra_as().expect("get_extra_as ok");
assert_eq!(
opts.web_search_options.search_context_size,
"high"
);
assert_eq!(opts.web_search_options.search_context_size, "high");
assert_eq!(opts.user.as_deref(), Some("u_123"));
}
@@ -206,4 +198,44 @@ mod tests {
assert_eq!(decoded.stream, req.stream);
assert_eq!(decoded.extra.get("trace_id"), Some(&json!("t-1")));
}
#[test]
fn message_request_with_tools_roundtrip() {
// 验证 ToolDef 的 serde 属性与 OpenaiToolDefinition 一致:
// 同名字段(name/description/parameters)序列化结果应一致。
let params = json!({
"type": "object",
"properties": {"x": {"type": "number"}},
"required": ["x"],
});
let tool = super::ToolDef {
name: "add".to_string(),
description: Some("add two numbers".to_string()),
parameters: params.clone(),
};
let req = MessageRequest {
model: "gpt-4o".into(),
messages: vec![Message::user_text("hi")],
tools: vec![tool],
tool_choice: ToolChoice::Auto,
max_tokens: None,
temperature: None,
top_p: None,
stop_sequences: vec![],
stream: false,
thinking: None,
extra: HashMap::new(),
};
let json = serde_json::to_string(&req).expect("serialize");
// 验证反序列化能还原所有字段(包括嵌套 parameters
let decoded: MessageRequest = serde_json::from_str(&json).expect("deserialize");
assert_eq!(decoded.tools.len(), 1);
assert_eq!(decoded.tools[0].name, "add");
assert_eq!(decoded.tools[0].description.as_deref(), Some("add two numbers"));
assert_eq!(decoded.tools[0].parameters, params);
// 验证序列化 JSON 不含 ToolDef 没有的字段(如 strict),保持 wire-format 兼容
assert!(!json.contains("strict"), "ToolDef 序列化不应包含 strict 字段");
}
}
+2 -6
View File
@@ -2,8 +2,8 @@ use crate::llm::types::openai_message::OpenaiChatMessage;
use crate::llm::types::shared::{FinishReason, ServiceTier};
use crate::llm::types::tool::OpenaiToolCall;
use crate::llm::types::usage::Usage;
use serde::{Deserialize, Serialize};
use crate::llm::types::{ContentField, OpenaiContentPart};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenLogprob {
@@ -135,11 +135,7 @@ impl From<OpenaiChatMessage> for Delta {
text.push_str(&t);
}
}
if text.is_empty() {
None
} else {
Some(text)
}
if text.is_empty() { None } else { Some(text) }
}
},
refusal: None,
+14 -23
View File
@@ -19,6 +19,7 @@ use crate::llm::types::usage::{CompletionTokensDetails, PromptTokensDetails, Usa
/// Phase 2 完成时统一收敛。
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum StopReason {
/// 自然停止。
Stop,
@@ -164,11 +165,15 @@ pub enum ContentBlockBuilder {
/// `thinking_signature`,最终通过 `finalize()` 回填到 `full_response` 的 `Thinking` block 中。
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum StreamEvent {
/// 消息开始(元信息)。
MessageStart { id: String, model: String },
/// 内容块开始(告知块类型,携带 id/name for ToolUse)。
ContentBlockStart { index: u32, block_type: ContentBlockType },
ContentBlockStart {
index: u32,
block_type: ContentBlockType,
},
/// 内容块结束标记。
ContentBlockEnd { index: u32 },
/// 文本增量。
@@ -320,9 +325,8 @@ impl PartialMessageResponse {
true
}
StreamEvent::ToolCallArgumentsDelta { index, arguments } => {
if let Some(ContentBlockBuilder::ToolUse {
arguments: buf, ..
}) = self.blocks.get_mut(index)
if let Some(ContentBlockBuilder::ToolUse { arguments: buf, .. }) =
self.blocks.get_mut(index)
{
buf.push_str(arguments);
}
@@ -360,9 +364,7 @@ impl PartialMessageResponse {
let mut content_blocks = Vec::with_capacity(self.blocks.len());
for (idx, builder) in self.blocks {
let block = Self::builder_to_block(idx, builder, self.thinking_signature.as_deref())
.map_err(|e| LlmError::Other(format!(
"partial 块 #{idx} finalize 失败: {e}"
)))?;
.map_err(|e| LlmError::Other(format!("partial 块 #{idx} finalize 失败: {e}")))?;
content_blocks.push(block);
}
@@ -743,10 +745,7 @@ mod tests {
Message::Assistant { content } => {
assert_eq!(content.len(), 2);
match (&content[0], &content[1]) {
(
ContentBlock::Text { text: t1 },
ContentBlock::Text { text: t2 },
) => {
(ContentBlock::Text { text: t1 }, ContentBlock::Text { text: t2 }) => {
assert_eq!(t1, "first");
assert_eq!(t2, "second");
}
@@ -772,9 +771,7 @@ mod tests {
index: 0,
block_type: ContentBlockType::Text,
},
StreamEvent::TextDelta {
text: "x".into(),
},
StreamEvent::TextDelta { text: "x".into() },
StreamEvent::ContentBlockEnd { index: 0 },
StreamEvent::MessageComplete {
full_response: empty_response(),
@@ -832,15 +829,9 @@ mod tests {
block_type: ContentBlockType::Text,
},
StreamEvent::ContentBlockEnd { index: 0 },
StreamEvent::TextDelta {
text: "t".into(),
},
StreamEvent::ThinkingDelta {
text: "p".into(),
},
StreamEvent::RefusalDelta {
text: "r".into(),
},
StreamEvent::TextDelta { text: "t".into() },
StreamEvent::ThinkingDelta { text: "p".into() },
StreamEvent::RefusalDelta { text: "r".into() },
StreamEvent::ToolCallArgumentsDelta {
index: 1,
arguments: "{\"x\":1}".into(),
+2
View File
@@ -13,6 +13,7 @@ pub enum Role {
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum FinishReason {
Stop,
Length,
@@ -67,6 +68,7 @@ pub enum StopSequence {
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case", tag = "type")]
#[non_exhaustive]
pub enum ResponseFormat {
Text,
JsonObject,
+40
View File
@@ -1,6 +1,25 @@
use serde::{Deserialize, Serialize};
use serde_json::Value;
/// Provider 无关的工具定义 IRv0.2 引入,替换 `ToolDefinition` 别名)。
///
/// 字段最小化:仅承载跨 Provider 公共的概念(name、description、parameters)。
/// OpenAI 专属 `strict` 字段不在此表达,由 OpenAI 适配层通过
/// `MessageRequest.extra` 逃生舱在 `convert_request` 内补充。
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ToolDef {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default)]
pub parameters: Value,
}
/// 旧 OpenAI wire-format 工具定义(v0.2 降级为 `#[doc(hidden)]`)。
///
/// 由 `ToolDef` 替代;保留仅供 OpenAI 适配层消费 `ToolDef → OpenaiToolDefinition`
/// 转换与外部反序列化兼容路径使用,不作为公共 API。
#[doc(hidden)]
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct OpenaiToolDefinition {
pub name: String,
@@ -12,6 +31,27 @@ pub struct OpenaiToolDefinition {
pub strict: Option<bool>,
}
impl From<ToolDef> for OpenaiToolDefinition {
fn from(t: ToolDef) -> Self {
Self {
name: t.name,
description: t.description,
parameters: t.parameters,
strict: None,
}
}
}
impl From<OpenaiToolDefinition> for ToolDef {
fn from(t: OpenaiToolDefinition) -> Self {
Self {
name: t.name,
description: t.description,
parameters: t.parameters,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FunctionCall {
pub name: String,
+3 -3
View File
@@ -12,11 +12,11 @@ pub use conversation::{ConversationMemory, ConversationMemoryConfig};
pub use error::MemoryError;
pub use knowledge::KnowledgeStore;
pub use retriever::MemoryRetriever;
pub use store::{InMemoryStore, MemoryStore};
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
// 低频类型(配置/高级使用)
pub use conversation::MemoryStrategy;
pub use knowledge::{PageIndexEntry, KNOWLEDGE_PREFIX};
pub use retriever::{RetrieverConfig, RetrievalResult, ScoredItem};
pub use knowledge::{KNOWLEDGE_PREFIX, PageIndexEntry};
pub use retriever::{RetrievalResult, RetrieverConfig, ScoredItem};
pub use store::{EvictionConfig, EvictionPolicy};
pub use types::{KnowledgePage, MemoryFilter, MemoryItem};
+21 -14
View File
@@ -12,6 +12,7 @@ use crate::memory::types::MemoryItem;
/// 对话消息管理策略。
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[non_exhaustive]
pub enum MemoryStrategy {
/// 滑动窗口:达到上限时删除最旧消息。
SlidingWindow,
@@ -160,7 +161,12 @@ impl ConversationMemory {
}
fn make_message_id(&self, index: usize, now: &OffsetDateTime) -> String {
format!("{}{:010}_{}", self.session_prefix(), index, now.unix_timestamp_nanos())
format!(
"{}{:010}_{}",
self.session_prefix(),
index,
now.unix_timestamp_nanos()
)
}
async fn maybe_evict_and_compact(&mut self) {
@@ -175,15 +181,16 @@ impl ConversationMemory {
}
if let Some(ref compact_config) = self.config.compact_config
&& should_compact(&self.messages, compact_config, &self.compact_state) {
let keep_recent = compact_config.keep_recent;
let freed = microcompact(&mut self.messages, keep_recent);
if freed > 0 {
self.compact_state.record_success();
} else {
let _ = self.compact_state.record_failure();
}
&& should_compact(&self.messages, compact_config, &self.compact_state)
{
let keep_recent = compact_config.keep_recent;
let freed = microcompact(&mut self.messages, keep_recent);
if freed > 0 {
self.compact_state.record_success();
} else {
let _ = self.compact_state.record_failure();
}
}
}
}
@@ -196,7 +203,8 @@ mod tests {
#[tokio::test]
async fn add_and_get_history() {
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
let mut conv = ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
let mut conv =
ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
conv.add_message(Message::user_text("hello")).await.unwrap();
conv.add_message(Message::user_text("world")).await.unwrap();
assert_eq!(conv.len(), 2);
@@ -211,9 +219,7 @@ mod tests {
conv.add_message(Message::tool_result("call_1", "ok", false))
.await
.unwrap();
conv.add_message(Message::assistant("done"))
.await
.unwrap();
conv.add_message(Message::assistant("done")).await.unwrap();
let original = conv.get_history().to_vec();
assert_eq!(original.len(), 2);
@@ -263,7 +269,8 @@ mod tests {
#[tokio::test]
async fn clear_empties_messages() {
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
let mut conv = ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
let mut conv =
ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
conv.add_message(Message::user_text("hello")).await.unwrap();
assert!(!conv.is_empty());
conv.clear().await.unwrap();
+2 -1
View File
@@ -6,6 +6,7 @@ use thiserror::Error;
///
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查环境变量、重试)。
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum MemoryError {
/// 按 ID 未找到指定记忆条目。可重试——通常是 namespace 拼写错误或条目已被淘汰。
#[error("未找到记忆条目 '{0}',请检查 ID 或 namespace 是否正确")]
@@ -33,4 +34,4 @@ impl MemoryError {
pub fn is_recoverable(&self) -> bool {
matches!(self, Self::NotFound(_) | Self::RetrievalError(_))
}
}
}
+6 -3
View File
@@ -57,8 +57,8 @@ impl KnowledgeStore {
}
let now = OffsetDateTime::now_utc();
let id = format!("{KNOWLEDGE_PREFIX}{}", page.id);
let content = serde_json::to_string(&page)
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
let content =
serde_json::to_string(&page).map_err(|e| MemoryError::Serialization(e.to_string()))?;
let item = MemoryItem {
id,
content,
@@ -128,7 +128,10 @@ impl KnowledgeStore {
.filter(|entry| {
entry.title.to_lowercase().contains(&needle)
|| entry.summary.to_lowercase().contains(&needle)
|| entry.tags.iter().any(|t| t.to_lowercase().contains(&needle))
|| entry
.tags
.iter()
.any(|t| t.to_lowercase().contains(&needle))
})
.map(|entry| entry.id.clone())
.collect()
+15 -8
View File
@@ -97,7 +97,11 @@ impl MemoryRetriever {
// 4. 过滤 → 排序 → 截取
items.retain(|i| i.score >= self.config.min_score);
items.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
items.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
items.truncate(self.config.max_results);
Ok(RetrievalResult {
@@ -159,12 +163,12 @@ fn char_bigrams(s: &str) -> Vec<String> {
fn default_stop_words() -> HashSet<String> {
[
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has",
"had", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall",
"can", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what",
"which", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor",
"so", "if", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at",
"by", "of", "as", "into", "through", "during", "before", "after", "above", "below",
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had",
"do", "does", "did", "will", "would", "should", "could", "may", "might", "shall", "can",
"this", "that", "these", "those", "it", "its", "they", "them", "their", "what", "which",
"who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor", "so", "if",
"then", "else", "with", "without", "for", "to", "from", "in", "on", "at", "by", "of", "as",
"into", "through", "during", "before", "after", "above", "below",
]
.iter()
.map(|s| s.to_string())
@@ -236,7 +240,10 @@ mod tests {
min_score: 0.99,
};
let retriever = MemoryRetriever::new(ks, config);
let result = retriever.retrieve("totally unrelated content").await.unwrap();
let result = retriever
.retrieve("totally unrelated content")
.await
.unwrap();
assert!(result.items.is_empty());
}
+7 -259
View File
@@ -1,14 +1,16 @@
//! MemoryStore 抽象接口与默认实现。
use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait;
use time::OffsetDateTime;
use crate::memory::error::MemoryError;
use crate::memory::types::{MemoryFilter, MemoryItem};
pub mod in_memory;
pub mod sqlite_store;
pub use in_memory::InMemoryStore;
pub use sqlite_store::SqliteStore;
/// 底层记忆存储抽象接口。
///
/// 下游可实现此 trait 以对接持久化后端(JSON 文件、SQLite、Redis 等)。
@@ -32,6 +34,7 @@ pub trait MemoryStore: Send + Sync {
/// 淘汰策略。
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum EvictionPolicy {
/// 不淘汰(默认)。
None,
@@ -57,258 +60,3 @@ impl Default for EvictionConfig {
}
}
}
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
pub struct InMemoryStore {
items: Mutex<HashMap<String, MemoryItem>>,
eviction: EvictionConfig,
/// 自上次淘汰检查以来的写入次数。
writes_since_check: Mutex<usize>,
}
impl InMemoryStore {
/// 创建一个无淘汰策略的 InMemoryStore。
pub fn new() -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction: EvictionConfig::default(),
writes_since_check: Mutex::new(0),
}
}
/// 创建一个带淘汰配置的 InMemoryStore。
pub fn with_eviction(eviction: EvictionConfig) -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction,
writes_since_check: Mutex::new(0),
}
}
fn maybe_evict(&self) {
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
let should_check = {
let mut counter = self.writes_since_check.lock().unwrap();
*counter += 1;
if *counter >= self.eviction.check_interval {
*counter = 0;
true
} else {
false
}
};
if !should_check {
return;
}
let policy = self.eviction.policy.clone();
match policy {
EvictionPolicy::None => {}
EvictionPolicy::Ttl { ttl_secs } => {
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
let mut items = self.items.lock().unwrap();
items.retain(|_, v| v.created_at > cutoff);
}
EvictionPolicy::Capacity { max_items } => {
let mut items = self.items.lock().unwrap();
if items.len() > max_items {
let mut vec: Vec<_> = items.drain().collect();
// O(n) 部分排序:保留 created_at 最大的 max_items 个
vec.select_nth_unstable_by(max_items, |a, b| {
b.1.created_at.cmp(&a.1.created_at)
});
vec.truncate(max_items);
*items = vec.into_iter().collect();
}
}
}
}
}
impl Default for InMemoryStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl MemoryStore for InMemoryStore {
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
{
let mut items = self.items.lock().unwrap();
items.insert(item.id.clone(), item);
}
self.maybe_evict();
Ok(())
}
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
Ok(items.get(id).cloned())
}
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
let mut items = self.items.lock().unwrap();
items.remove(id);
Ok(())
}
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
let mut result: Vec<MemoryItem> = items
.values()
.filter(|v| match &filter.prefix {
Some(p) => v.id.starts_with(p),
None => true,
})
.filter(|v| match filter.since {
Some(t) => v.created_at > t,
None => true,
})
.cloned()
.collect();
// 按 created_at 升序排列(最旧在前)
result.sort_by_key(|v| v.created_at);
// 应用 offset
if let Some(offset) = filter.offset {
if offset < result.len() {
result.drain(..offset);
} else {
result.clear();
}
}
// 应用 limit
if let Some(limit) = filter.limit {
result.truncate(limit);
}
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use time::OffsetDateTime;
fn make_item(id: &str) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({}),
created_at: OffsetDateTime::now_utc(),
}
}
#[tokio::test]
async fn save_get_delete_list() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
let got = store.get("a").await.unwrap();
assert!(got.is_some());
assert_eq!(got.unwrap().id, "a");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
store.delete("a").await.unwrap();
assert!(store.get("a").await.unwrap().is_none());
}
#[tokio::test]
async fn save_is_upsert() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
let mut item = make_item("a");
item.content = "updated".to_string();
store.save(item).await.unwrap();
let got = store.get("a").await.unwrap().unwrap();
assert_eq!(got.content, "updated");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn list_with_prefix_and_limit() {
let store = InMemoryStore::new();
store.save(make_item("foo_a")).await.unwrap();
store.save(make_item("foo_b")).await.unwrap();
store.save(make_item("bar_a")).await.unwrap();
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
limit: Some(1),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn capacity_eviction() {
// 强制每次写入都检查
let eviction = EvictionConfig {
policy: EvictionPolicy::Capacity { max_items: 2 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
// 第一条和第二条共存
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
// 第三条写入触发淘汰:a 或 b 之一被淘汰
store.save(make_item("c")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
// 留下的应该是 b 和 c(最新的两个)
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"b"));
assert!(ids.contains(&"c"));
}
#[tokio::test]
async fn ttl_eviction() {
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
// 改用小 TTL + 睡眠:先 save asleepsave b 时 a 已过期被淘汰。
let eviction = EvictionConfig {
policy: EvictionPolicy::Ttl { ttl_secs: 1 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
store.save(make_item("a")).await.unwrap();
// 等待超过 1 秒
std::thread::sleep(std::time::Duration::from_millis(1100));
// 触发淘汰:a 已超过 ttl_secs=1,应被淘汰
store.save(make_item("b")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
// 由于 ttl_secs=1,且 b 刚写入,可能刚好处于临界值。
// 我们只断言 list 不包含 "a" 即可。
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(
!ids.contains(&"a"),
"expected 'a' to be evicted, but found in {ids:?}"
);
}
#[tokio::test]
async fn none_policy_no_eviction() {
let eviction = EvictionConfig {
policy: EvictionPolicy::None,
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
for i in 0..100 {
store.save(make_item(&format!("item_{i}"))).await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
}
}
+266
View File
@@ -0,0 +1,266 @@
//! 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait;
use time::OffsetDateTime;
use crate::memory::error::MemoryError;
use crate::memory::store::{EvictionConfig, EvictionPolicy, MemoryStore};
use crate::memory::types::{MemoryFilter, MemoryItem};
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
pub struct InMemoryStore {
items: Mutex<HashMap<String, MemoryItem>>,
eviction: EvictionConfig,
/// 自上次淘汰检查以来的写入次数。
writes_since_check: Mutex<usize>,
}
impl InMemoryStore {
/// 创建一个无淘汰策略的 InMemoryStore。
pub fn new() -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction: EvictionConfig::default(),
writes_since_check: Mutex::new(0),
}
}
/// 创建一个带淘汰配置的 InMemoryStore。
pub fn with_eviction(eviction: EvictionConfig) -> Self {
Self {
items: Mutex::new(HashMap::new()),
eviction,
writes_since_check: Mutex::new(0),
}
}
fn maybe_evict(&self) {
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
let should_check = {
let mut counter = self.writes_since_check.lock().unwrap();
*counter += 1;
if *counter >= self.eviction.check_interval {
*counter = 0;
true
} else {
false
}
};
if !should_check {
return;
}
let policy = self.eviction.policy.clone();
match policy {
EvictionPolicy::None => {}
EvictionPolicy::Ttl { ttl_secs } => {
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
let mut items = self.items.lock().unwrap();
items.retain(|_, v| v.created_at > cutoff);
}
EvictionPolicy::Capacity { max_items } => {
let mut items = self.items.lock().unwrap();
if items.len() > max_items {
let mut vec: Vec<_> = items.drain().collect();
// O(n) 部分排序:保留 created_at 最大的 max_items 个
vec.select_nth_unstable_by(max_items, |a, b| {
b.1.created_at.cmp(&a.1.created_at)
});
vec.truncate(max_items);
*items = vec.into_iter().collect();
}
}
}
}
}
impl Default for InMemoryStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl MemoryStore for InMemoryStore {
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
{
let mut items = self.items.lock().unwrap();
items.insert(item.id.clone(), item);
}
self.maybe_evict();
Ok(())
}
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
Ok(items.get(id).cloned())
}
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
let mut items = self.items.lock().unwrap();
items.remove(id);
Ok(())
}
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
let items = self.items.lock().unwrap();
let mut result: Vec<MemoryItem> = items
.values()
.filter(|v| match &filter.prefix {
Some(p) => v.id.starts_with(p),
None => true,
})
.filter(|v| match filter.since {
Some(t) => v.created_at > t,
None => true,
})
.cloned()
.collect();
// 按 created_at 升序排列(最旧在前)
result.sort_by_key(|v| v.created_at);
// 应用 offset
if let Some(offset) = filter.offset {
if offset < result.len() {
result.drain(..offset);
} else {
result.clear();
}
}
// 应用 limit
if let Some(limit) = filter.limit {
result.truncate(limit);
}
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use time::OffsetDateTime;
fn make_item(id: &str) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({}),
created_at: OffsetDateTime::now_utc(),
}
}
#[tokio::test]
async fn save_get_delete_list() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
let got = store.get("a").await.unwrap();
assert!(got.is_some());
assert_eq!(got.unwrap().id, "a");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
store.delete("a").await.unwrap();
assert!(store.get("a").await.unwrap().is_none());
}
#[tokio::test]
async fn save_is_upsert() {
let store = InMemoryStore::new();
store.save(make_item("a")).await.unwrap();
let mut item = make_item("a");
item.content = "updated".to_string();
store.save(item).await.unwrap();
let got = store.get("a").await.unwrap().unwrap();
assert_eq!(got.content, "updated");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn list_with_prefix_and_limit() {
let store = InMemoryStore::new();
store.save(make_item("foo_a")).await.unwrap();
store.save(make_item("foo_b")).await.unwrap();
store.save(make_item("bar_a")).await.unwrap();
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
limit: Some(1),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn capacity_eviction() {
// 强制每次写入都检查
let eviction = EvictionConfig {
policy: EvictionPolicy::Capacity { max_items: 2 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
// 第一条和第二条共存
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
// 第三条写入触发淘汰:a 或 b 之一被淘汰
store.save(make_item("c")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
// 留下的应该是 b 和 c(最新的两个)
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"b"));
assert!(ids.contains(&"c"));
}
#[tokio::test]
async fn ttl_eviction() {
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
// 改用小 TTL + 睡眠:先 save asleepsave b 时 a 已过期被淘汰。
let eviction = EvictionConfig {
policy: EvictionPolicy::Ttl { ttl_secs: 1 },
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
store.save(make_item("a")).await.unwrap();
// 等待超过 1 秒
std::thread::sleep(std::time::Duration::from_millis(1100));
// 触发淘汰:a 已超过 ttl_secs=1,应被淘汰
store.save(make_item("b")).await.unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
// 由于 ttl_secs=1,且 b 刚写入,可能刚好处于临界值。
// 我们只断言 list 不包含 "a" 即可。
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(
!ids.contains(&"a"),
"expected 'a' to be evicted, but found in {ids:?}"
);
}
#[tokio::test]
async fn none_policy_no_eviction() {
let eviction = EvictionConfig {
policy: EvictionPolicy::None,
check_interval: 1,
};
let store = InMemoryStore::with_eviction(eviction);
for i in 0..100 {
store.save(make_item(&format!("item_{i}"))).await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
}
}
+545
View File
@@ -0,0 +1,545 @@
//! SqliteStore —— 基于 rusqlite 的持久化 MemoryStore 实现。
//!
//! 单进程独享、写入串行化(WAL + Mutex),适合本地 Agent 长期持久化场景。
use std::path::Path;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use rusqlite::{params, params_from_iter, Connection, ErrorCode};
use time::format_description::well_known::Rfc3339;
use time::OffsetDateTime;
use tracing::{debug, error, instrument, warn};
use crate::memory::error::MemoryError;
use crate::memory::store::MemoryStore;
use crate::memory::types::{MemoryFilter, MemoryItem};
const INITIAL_USER_VERSION: i64 = 1;
const BUSY_TIMEOUT_MS: i64 = 5000;
const WAL_AUTOCHECKPOINT_PAGES: i64 = 1000;
/// SQLite 持久化后端的 MemoryStore 实现。
///
/// 设计要点:
/// - 单进程独享:`Arc<Mutex<Connection>>` 串行化所有 IO
/// - WAL 模式 + `synchronous=NORMAL` 兼顾崩溃安全与吞吐
/// - `created_at` 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
/// - 所有 IO 通过 `tokio::task::spawn_blocking` 卸载到阻塞线程池
pub struct SqliteStore {
conn: Arc<Mutex<Connection>>,
}
impl SqliteStore {
/// 打开或创建一个 SQLite 数据库。
///
/// - `path = ":memory:"` 使用内存数据库(测试场景)
/// - 其他路径:自动创建父目录;文件已存在则附加打开
/// - 启动时执行 `migrate()`,失败立即返回错误
#[instrument(skip(path), fields(path = %path.as_ref().display()))]
pub fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
let path_ref = path.as_ref();
let path_str = path_ref.to_string_lossy();
let conn = if path_str == ":memory:" {
Connection::open_in_memory()
} else {
if let Some(parent) = path_ref.parent()
&& !parent.as_os_str().is_empty()
{
std::fs::create_dir_all(parent).map_err(|e| {
MemoryError::Storage(format!(
"创建数据库父目录失败 ({}): {}",
parent.display(),
e
))
})?;
}
Connection::open(path_ref)
}
.map_err(|e| map_sqlite_error(e, "打开数据库"))?;
migrate(&conn)?;
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
})
}
}
#[async_trait]
impl MemoryStore for SqliteStore {
#[instrument(skip(self, item), fields(id = %item.id))]
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
let conn = Arc::clone(&self.conn);
let created_at_str = item
.created_at
.to_offset(time::UtcOffset::UTC)
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format created_at: {e}")))?;
let metadata_str = serde_json::to_string(&item.metadata)
.map_err(|e| MemoryError::Serialization(format!("serialize metadata: {e}")))?;
let id = item.id;
let content = item.content;
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
conn.execute(
"INSERT INTO memory_items (id, content, metadata, created_at) \
VALUES (?1, ?2, ?3, ?4) \
ON CONFLICT(id) DO UPDATE SET \
content=excluded.content, \
metadata=excluded.metadata, \
created_at=excluded.created_at",
params![id, content, metadata_str, created_at_str],
)
.map_err(|e| map_sqlite_error(e, "保存记忆"))?;
Ok(())
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, id))]
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
let conn = Arc::clone(&self.conn);
let id_owned = id.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<MemoryItem>, MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
let mut stmt = conn
.prepare("SELECT id, content, metadata, created_at FROM memory_items WHERE id = ?1")
.map_err(|e| map_sqlite_error(e, "prepare get"))?;
let mut rows = stmt
.query_map(params![id_owned], row_to_item)
.map_err(|e| map_sqlite_error(e, "query get"))?;
match rows.next() {
None => Ok(None),
Some(row) => row
.map(Some)
.map_err(|e| map_sqlite_error(e, "decode row")),
}
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, id))]
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
let conn = Arc::clone(&self.conn);
let id_owned = id.to_string();
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
conn.execute(
"DELETE FROM memory_items WHERE id = ?1",
params![id_owned],
)
.map_err(|e| map_sqlite_error(e, "delete"))?;
Ok(())
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, filter))]
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
let mut sql = String::from(
"SELECT id, content, metadata, created_at FROM memory_items WHERE 1=1",
);
let mut param_values: Vec<String> = Vec::new();
let mut ph_idx = 0usize;
if filter.prefix.is_some() {
ph_idx += 1;
sql.push_str(&format!(" AND id LIKE ?{ph_idx} || '%'"));
}
if filter.since.is_some() {
ph_idx += 1;
sql.push_str(&format!(" AND created_at > ?{ph_idx}"));
}
// ORDER BY created_at ASC(按时间升序,最旧在前)
sql.push_str(" ORDER BY created_at ASC");
let limit_sql: String = match (filter.limit, filter.offset) {
(Some(_), Some(_)) => {
ph_idx += 1;
let limit_p = ph_idx;
ph_idx += 1;
let offset_p = ph_idx;
format!(" LIMIT ?{limit_p} OFFSET ?{offset_p}")
}
(Some(_), None) => {
ph_idx += 1;
let limit_p = ph_idx;
format!(" LIMIT ?{limit_p}")
}
(None, Some(_)) => {
// SQLite 中 LIMIT -1 表示无限制
ph_idx += 1;
let offset_p = ph_idx;
format!(" LIMIT -1 OFFSET ?{offset_p}")
}
(None, None) => String::new(),
};
sql.push_str(&limit_sql);
if let Some(p) = &filter.prefix {
param_values.push(p.clone());
}
if let Some(t) = filter.since {
let s = t
.to_offset(time::UtcOffset::UTC)
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format since: {e}")))?;
param_values.push(s);
}
if let Some(l) = filter.limit {
param_values.push(l.to_string());
}
if let Some(o) = filter.offset {
param_values.push(o.to_string());
}
let conn = Arc::clone(&self.conn);
let sql_owned = sql;
let param_values_owned = param_values;
tokio::task::spawn_blocking(move || -> Result<Vec<MemoryItem>, MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
let mut stmt = conn
.prepare(&sql_owned)
.map_err(|e| map_sqlite_error(e, "list prepare"))?;
let params_iter: Vec<&dyn rusqlite::ToSql> = param_values_owned
.iter()
.map(|s| s as &dyn rusqlite::ToSql)
.collect();
let rows = stmt
.query_map(params_from_iter(params_iter), row_to_item)
.map_err(|e| map_sqlite_error(e, "list query"))?;
let mut result = Vec::new();
for row in rows {
result.push(row.map_err(|e| map_sqlite_error(e, "list row"))?);
}
debug!(count = result.len(), "SqliteStore::list 完成");
Ok(result)
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
}
fn row_to_item(row: &rusqlite::Row<'_>) -> Result<MemoryItem, rusqlite::Error> {
let id: String = row.get(0)?;
let content: String = row.get(1)?;
let metadata_str: String = row.get(2)?;
let created_at_str: String = row.get(3)?;
let metadata: serde_json::Value = serde_json::from_str(&metadata_str).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(2, rusqlite::types::Type::Text, Box::new(e))
})?;
let created_at = OffsetDateTime::parse(&created_at_str, &Rfc3339).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(3, rusqlite::types::Type::Text, Box::new(e))
})?;
Ok(MemoryItem {
id,
content,
metadata,
created_at: created_at.to_offset(time::UtcOffset::UTC),
})
}
fn migrate(conn: &Connection) -> Result<(), MemoryError> {
conn.pragma_update(None, "journal_mode", "WAL")
.map_err(|e| map_sqlite_error(e, "PRAGMA journal_mode"))?;
conn.pragma_update(None, "synchronous", "NORMAL")
.map_err(|e| map_sqlite_error(e, "PRAGMA synchronous"))?;
conn.execute_batch(&format!("PRAGMA busy_timeout = {BUSY_TIMEOUT_MS};"))
.map_err(|e| map_sqlite_error(e, "PRAGMA busy_timeout"))?;
conn.execute_batch(&format!(
"PRAGMA wal_autocheckpoint = {WAL_AUTOCHECKPOINT_PAGES};"
))
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_autocheckpoint"))?;
let version: i64 = conn
.query_row("PRAGMA user_version", [], |row| row.get(0))
.map_err(|e| map_sqlite_error(e, "PRAGMA user_version"))?;
if version < INITIAL_USER_VERSION {
conn.execute_batch(
"CREATE TABLE IF NOT EXISTS memory_items (
id TEXT PRIMARY KEY,
content TEXT NOT NULL,
metadata TEXT NOT NULL DEFAULT '{}',
created_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_memory_items_created_at
ON memory_items(created_at);
PRAGMA user_version = 1;",
)
.map_err(|e| map_sqlite_error(e, "create schema v1"))?;
}
let check_result: String = conn
.query_row("PRAGMA quick_check", [], |row| row.get(0))
.map_err(|e| map_sqlite_error(e, "PRAGMA quick_check"))?;
if check_result != "ok" {
error!(result = %check_result, "数据库文件 quick_check 失败");
return Err(MemoryError::Storage(format!(
"数据库文件损坏: {check_result}"
)));
}
conn.execute_batch("PRAGMA wal_checkpoint(TRUNCATE);")
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_checkpoint"))?;
Ok(())
}
fn map_sqlite_error(e: rusqlite::Error, ctx: &str) -> MemoryError {
match &e {
rusqlite::Error::SqliteFailure(err, _) => match err.code {
ErrorCode::ConstraintViolation => MemoryError::InvalidInput(format!("{ctx}: {e}")),
ErrorCode::DatabaseBusy | ErrorCode::DatabaseLocked => {
warn!("SQLite 忙: {e}");
MemoryError::Storage(format!("{ctx}: {e}"))
}
_ => MemoryError::Storage(format!("{ctx}: {e}")),
},
rusqlite::Error::InvalidQuery
| rusqlite::Error::InvalidParameterName(_)
| rusqlite::Error::InvalidColumnIndex(_)
| rusqlite::Error::InvalidColumnName(_) => {
MemoryError::InvalidInput(format!("{ctx}: {e}"))
}
rusqlite::Error::FromSqlConversionFailure(_, _, _)
| rusqlite::Error::ToSqlConversionFailure(_) => {
MemoryError::Serialization(format!("{ctx}: {e}"))
}
_ => MemoryError::Storage(format!("{ctx}: {e}")),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::store::InMemoryStore;
use std::sync::Arc;
use tempfile::TempDir;
use time::OffsetDateTime;
fn make_item(id: &str) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({"id_key": id}),
created_at: OffsetDateTime::now_utc(),
}
}
fn make_item_at(id: &str, when: OffsetDateTime) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({}),
created_at: when,
}
}
#[tokio::test]
async fn crud_basic() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
let got_a = store.get("a").await.unwrap();
assert!(got_a.is_some());
assert_eq!(got_a.unwrap().id, "a");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
store.delete("a").await.unwrap();
assert!(store.get("a").await.unwrap().is_none());
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
assert_eq!(list[0].id, "b");
}
#[tokio::test]
async fn save_is_upsert() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("a")).await.unwrap();
let mut item = make_item("a");
item.content = "updated".to_string();
item.metadata = serde_json::json!({"rev": 2});
let original_created_at = item.created_at;
store.save(item).await.unwrap();
let got = store.get("a").await.unwrap().unwrap();
assert_eq!(got.content, "updated");
assert_eq!(got.metadata["rev"], serde_json::json!(2));
// created_at 保持调用方传入值
assert_eq!(got.created_at, original_created_at);
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn list_with_prefix() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("foo_a")).await.unwrap();
store.save(make_item("foo_b")).await.unwrap();
store.save(make_item("bar_a")).await.unwrap();
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"foo_a"));
assert!(ids.contains(&"foo_b"));
assert!(!ids.contains(&"bar_a"));
}
#[tokio::test]
async fn list_with_since_filter() {
let store = SqliteStore::open(":memory:").unwrap();
let t0 = OffsetDateTime::now_utc();
store
.save(make_item_at("early", t0 - time::Duration::seconds(60)))
.await
.unwrap();
store.save(make_item_at("middle", t0)).await.unwrap();
store
.save(make_item_at("late", t0 + time::Duration::seconds(60)))
.await
.unwrap();
let filter = MemoryFilter {
since: Some(t0 - time::Duration::seconds(1)),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert_eq!(list.len(), 2);
assert!(ids.contains(&"middle"));
assert!(ids.contains(&"late"));
assert!(!ids.contains(&"early"));
}
#[tokio::test]
async fn list_with_offset_and_limit() {
let store = SqliteStore::open(":memory:").unwrap();
// 写入 5 条时间递增的记录
let base = OffsetDateTime::now_utc() - time::Duration::seconds(5);
for i in 0..5 {
let mut item = make_item(&format!("item_{i}"));
item.created_at = base + time::Duration::seconds(i);
store.save(item).await.unwrap();
}
// offset=1, limit=2 -> item_1, item_2
let filter = MemoryFilter {
offset: Some(1),
limit: Some(2),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
assert_eq!(list[0].id, "item_1");
assert_eq!(list[1].id, "item_2");
}
#[tokio::test]
async fn concurrent_writers_no_data_loss() {
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
let mut handles = Vec::new();
for w in 0..10 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
for i in 0..10 {
let id = format!("w{w}_i{i}");
s.save(make_item(&id)).await.unwrap();
}
}));
}
for h in handles {
h.await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
// 验证所有 id 唯一
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
ids.sort();
ids.dedup();
assert_eq!(ids.len(), 100);
}
#[tokio::test]
async fn persistence_round_trip() {
let dir = TempDir::new().unwrap();
let path = dir.path().join("memory.db");
// 阶段 1:写入 3 条
{
let store = SqliteStore::open(&path).unwrap();
store.save(make_item("alpha")).await.unwrap();
store.save(make_item("beta")).await.unwrap();
store.save(make_item("gamma")).await.unwrap();
}
// 阶段 2:重新打开,验证数据完整
{
let store = SqliteStore::open(&path).unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 3);
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"alpha"));
assert!(ids.contains(&"beta"));
assert!(ids.contains(&"gamma"));
// 单条读回
let got = store.get("beta").await.unwrap().unwrap();
assert_eq!(got.content, "content-beta");
}
}
#[tokio::test]
async fn open_invalid_path_returns_error() {
// 路径指向已存在的目录而非文件,open 应失败
let dir = TempDir::new().unwrap();
match SqliteStore::open(dir.path()) {
Err(MemoryError::Storage(_)) => {}
Err(other) => panic!("expected Storage error, got {other:?}"),
Ok(_) => panic!("expected error when opening a directory as database"),
}
}
#[tokio::test]
async fn trait_object_compatibility() {
// ponytail: 回归验证 SqliteStore 可作为 Arc<dyn MemoryStore> 与 InMemoryStore 互换
// 所有现有消费者(Conversation / Knowledge / Retriever / SessionMemory)均通过 trait object 引用,
// 此测试确保 trait 接口契约在 SqliteStore 上同样成立。
let sqlite: Arc<dyn MemoryStore> =
Arc::new(SqliteStore::open(":memory:").unwrap());
let in_mem: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let stores: Vec<Arc<dyn MemoryStore>> = vec![Arc::clone(&sqlite), Arc::clone(&in_mem)];
for store in &stores {
store.save(make_item("x")).await.unwrap();
let got = store.get("x").await.unwrap();
assert_eq!(got.unwrap().id, "x");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
store.delete("x").await.unwrap();
assert!(store.get("x").await.unwrap().is_none());
}
}
}
+2 -2
View File
@@ -1,7 +1,7 @@
pub mod composer;
pub mod error;
pub mod template;
pub mod composer;
pub use composer::{PromptComposer, validate_messages};
pub use error::PromptError;
pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue};
pub use composer::{validate_messages, PromptComposer};
+17 -15
View File
@@ -48,7 +48,11 @@ impl PromptComposer {
/// 添加一条 Tool 消息(工具执行结果回传)。
pub fn tool(mut self, tool_call_id: impl Into<String>, content: impl Into<String>) -> Self {
self.push_message(Message::tool_result(tool_call_id.into(), content.into(), false));
self.push_message(Message::tool_result(
tool_call_id.into(),
content.into(),
false,
));
self
}
@@ -133,11 +137,7 @@ impl PromptComposer {
}
/// 添加一条含指定 ContentBlock 的 Tool 消息。
pub fn tool_content(
mut self,
tool_call_id: impl Into<String>,
block: ContentBlock,
) -> Self {
pub fn tool_content(mut self, tool_call_id: impl Into<String>, block: ContentBlock) -> Self {
self.push_message(Message::ToolResult {
tool_call_id: tool_call_id.into(),
content: vec![block],
@@ -187,9 +187,7 @@ impl PromptComposer {
/// 验证消息序列是否符合 LLM API 要求(Tool 消息必须紧跟含 tool_calls 的 Assistant)。
pub fn validate_messages(messages: &[Message]) -> Result<(), PromptError> {
if messages.is_empty() {
return Err(PromptError::InvalidSequence(
"消息列表不能为空".to_string(),
));
return Err(PromptError::InvalidSequence("消息列表不能为空".to_string()));
}
let mut last_tool_call_ids: Vec<String> = Vec::new();
@@ -297,7 +295,8 @@ mod tests {
#[test]
fn test_template_if() {
let tpl = PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
let tpl =
PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
let mut ctx = TemplateContext::new();
ctx.insert("name", "Bob");
@@ -312,11 +311,14 @@ mod tests {
fn test_template_each() {
let tpl = PromptTemplate::compile("Items: {{#each items}}{{item}}, {{/each}}").unwrap();
let mut ctx = TemplateContext::new();
ctx.insert("items", TemplateValue::Array(vec![
TemplateValue::String("a".to_string()),
TemplateValue::String("b".to_string()),
TemplateValue::String("c".to_string()),
]));
ctx.insert(
"items",
TemplateValue::Array(vec![
TemplateValue::String("a".to_string()),
TemplateValue::String("b".to_string()),
TemplateValue::String("c".to_string()),
]),
);
let result = tpl.render(&ctx).unwrap();
assert_eq!(result, "Items: a, b, c, ");
+7 -2
View File
@@ -1,6 +1,7 @@
use thiserror::Error;
#[derive(Error, Debug)]
#[non_exhaustive]
pub enum PromptError {
#[error("模板解析错误: {0}。请检查模板语法({{var}} / {{#if}} / {{#each}}")]
Parse(String),
@@ -8,7 +9,9 @@ pub enum PromptError {
#[error("渲染错误: 变量 '{0}' 未找到。请在 TemplateContext 中插入该变量")]
VariableNotFound(String),
#[error("渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板")]
#[error(
"渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板"
)]
PartialNotFound(String),
#[error("渲染错误: '{0}' 不是数组,无法遍历。请确认传入的是数组或先判空")]
@@ -20,7 +23,9 @@ pub enum PromptError {
#[error("渲染错误: {0}")]
Render(String),
#[error("消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)")]
#[error(
"消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)"
)]
InvalidSequence(String),
#[error("文件读取错误: {0}。请检查模板文件路径与权限")]
+21 -34
View File
@@ -1,6 +1,6 @@
use serde_json::Value;
use std::collections::HashMap;
use std::fmt;
use serde_json::Value;
use crate::prompt::error::PromptError;
@@ -140,7 +140,9 @@ fn json_to_template_value(v: &Value) -> Result<TemplateValue, PromptError> {
#[derive(Debug, Clone)]
enum Fragment {
Literal(String),
Variable { name: String },
Variable {
name: String,
},
If {
condition: String,
body: Vec<Fragment>,
@@ -223,8 +225,7 @@ fn compile_fragments(template: &str) -> Result<Vec<Fragment>, PromptError> {
let tag = tag_content.trim();
if let Some(rest) = tag.strip_prefix("#if ") {
let (body, else_body, new_i) =
parse_block(template, i, "if")?;
let (body, else_body, new_i) = parse_block(template, i, "if")?;
let condition = rest.trim().to_string();
fragments.push(Fragment::If {
condition,
@@ -331,10 +332,7 @@ fn parse_block(
Err(PromptError::Parse(format!("未闭合的 {{#{}}}", kind)))
}
fn parse_each_block(
template: &str,
start: usize,
) -> Result<(Vec<Fragment>, usize), PromptError> {
fn parse_each_block(template: &str, start: usize) -> Result<(Vec<Fragment>, usize), PromptError> {
let bytes = template.as_bytes();
let len = bytes.len();
let mut depth = 1u32;
@@ -368,9 +366,7 @@ fn parse_each_block(
}
}
Err(PromptError::Parse(
"未闭合的 {{#each}} 块".to_string(),
))
Err(PromptError::Parse("未闭合的 {{#each}} 块".to_string()))
}
fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), PromptError> {
@@ -395,9 +391,7 @@ fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), Prom
}
}
Err(PromptError::Parse(
"未闭合的 {{#raw}} 块".to_string(),
))
Err(PromptError::Parse("未闭合的 {{#raw}} 块".to_string()))
}
// ===== Renderer =====
@@ -418,33 +412,28 @@ fn render_fragments(
Fragment::Literal(text) => {
output.push_str(text);
}
Fragment::Variable { name } => {
match ctx.get(name) {
Some(val) => {
output.push_str(&format!("{}", val));
}
None => {
return Err(PromptError::VariableNotFound(name.clone()));
}
Fragment::Variable { name } => match ctx.get(name) {
Some(val) => {
output.push_str(&format!("{}", val));
}
}
None => {
return Err(PromptError::VariableNotFound(name.clone()));
}
},
Fragment::If {
condition,
body,
else_body,
} => {
let truthy = ctx
.get(condition)
.map(|v| v.is_truthy())
.unwrap_or(false);
let truthy = ctx.get(condition).map(|v| v.is_truthy()).unwrap_or(false);
let target = if truthy { body } else { else_body };
render_fragments(target, ctx, partials, output, depth + 1)?;
}
Fragment::Each { variable, body } => {
let arr = match ctx.get(variable) {
Some(val) => val.as_array().ok_or_else(|| {
PromptError::NotAnArray(variable.clone())
})?,
Some(val) => val
.as_array()
.ok_or_else(|| PromptError::NotAnArray(variable.clone()))?,
None => {
return Err(PromptError::VariableNotFound(variable.clone()));
}
@@ -504,10 +493,8 @@ impl PromptTemplateRegistry {
/// 延迟编译注册:只存储原始字符串,首次渲染时编译。
pub fn register_lazy(&mut self, name: &str, template: &str) {
self.templates.insert(
name.to_string(),
StoredTemplate::Raw(template.to_string()),
);
self.templates
.insert(name.to_string(), StoredTemplate::Raw(template.to_string()));
}
/// 从文件读取并编译注册。
+10 -3
View File
@@ -4,9 +4,12 @@ use std::sync::Arc;
/// 工具调用过程中可能发生的所有错误。
#[derive(thiserror::Error, Debug, Clone)]
#[non_exhaustive]
pub enum ToolError {
/// 工具未注册。不可恢复——需调用方先 `registry.register(...)`。
#[error("工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写")]
#[error(
"工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写"
)]
NotFound(String),
/// 工具执行失败(可恢复——文本回传 LLM 由其决定重试或放弃)。
@@ -14,11 +17,15 @@ pub enum ToolError {
ExecutionFailed(String, String),
/// 工具参数无效(可恢复——文本回传 LLM)。
#[error("工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema")]
#[error(
"工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema"
)]
InvalidArguments(String, String),
/// 权限被拒绝(不可恢复——终止循环)。
#[error("权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过")]
#[error(
"权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过"
)]
PermissionDenied(String, String),
/// MCP 协议错误(不可恢复)。
+17 -37
View File
@@ -9,19 +9,18 @@
use std::collections::HashMap;
use std::process::Stdio;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use serde_json::{Value, json};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::{oneshot, Mutex};
use tokio::sync::{Mutex, oneshot};
#[allow(deprecated)]
use crate::llm::types::ToolDefinition;
use crate::llm::types::tool::ToolDef;
use crate::tools::base::{BaseTool, ToolContext, ToolRef};
use crate::tools::error::ToolError;
@@ -136,7 +135,6 @@ impl std::fmt::Debug for McpClient {
}
}
#[allow(deprecated)]
impl McpClient {
/// 创建一个 MCP 客户端。
pub fn new(server_name: impl Into<String>, transport: McpTransport) -> Self {
@@ -226,9 +224,7 @@ impl McpClient {
"version": env!("CARGO_PKG_VERSION")
}
});
let _response = self
.send_request("initialize", Some(init_params))
.await?;
let _response = self.send_request("initialize", Some(init_params)).await?;
// 发送 initialized 通知(无 id
self.send_notification("notifications/initialized", Some(json!({})))
@@ -239,7 +235,7 @@ impl McpClient {
}
/// 列出服务器支持的工具(调用 `tools/list`)。
pub async fn list_tools(&mut self) -> Result<Vec<ToolDefinition>, ToolError> {
pub async fn list_tools(&mut self) -> Result<Vec<ToolDef>, ToolError> {
if !self.is_initialized() {
return Err(ToolError::McpNotInitialized(self.server_name.clone()));
}
@@ -274,11 +270,10 @@ impl McpClient {
description: description.clone(),
input_schema: input_schema.clone(),
});
defs.push(ToolDefinition {
defs.push(ToolDef {
name,
description,
parameters: input_schema,
strict: None,
});
}
Ok(defs)
@@ -337,11 +332,7 @@ impl McpClient {
if let Some(state) = self.process.take() {
let mut state = state.lock().await;
// 优雅等待 5 秒
let graceful = tokio::time::timeout(
Duration::from_secs(5),
state.child.wait(),
)
.await;
let graceful = tokio::time::timeout(Duration::from_secs(5), state.child.wait()).await;
if graceful.is_err() {
// 超时则强杀
let _ = state.child.kill().await;
@@ -372,11 +363,7 @@ impl McpClient {
tools
}
async fn send_request(
&self,
method: &str,
params: Option<Value>,
) -> Result<Value, ToolError> {
async fn send_request(&self, method: &str, params: Option<Value>) -> Result<Value, ToolError> {
let state_arc = self
.process
.as_ref()
@@ -412,9 +399,11 @@ impl McpClient {
.write_all(b"\n")
.await
.map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?;
state.stdin.flush().await.map_err(|e| {
ToolError::McpError(format!("flush stdin 失败: {e}"))
})?;
state
.stdin
.flush()
.await
.map_err(|e| ToolError::McpError(format!("flush stdin 失败: {e}")))?;
}
// 等待响应(带超时)
@@ -471,10 +460,7 @@ impl McpClient {
}
/// 持续读取 stdout,将响应分发到对应的 oneshot sender。
async fn read_loop(
mut reader: BufReader<ChildStdout>,
state: Arc<Mutex<ChildProcessState>>,
) {
async fn read_loop(mut reader: BufReader<ChildStdout>, state: Arc<Mutex<ChildProcessState>>) {
let mut line = String::new();
loop {
line.clear();
@@ -542,7 +528,6 @@ enum McpClientHandle {
}
#[async_trait]
#[allow(deprecated)]
impl BaseTool for McpToolAdapter {
fn name(&self) -> &str {
&self.name
@@ -556,11 +541,7 @@ impl BaseTool for McpToolAdapter {
self.parameters.clone()
}
async fn execute(
&self,
_args: Value,
_ctx: &ToolContext<'_>,
) -> Result<Value, ToolError> {
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
// 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。
// 实际生产中应持有 Arc<McpClient> 并通过 mcp.call_tool() 执行。
// 这里返回错误,提示需要通过其他方式调用 MCP 工具。
@@ -617,8 +598,7 @@ mod tests {
#[test]
fn test_jsonrpc_response_parse_error() {
let s =
r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
let s = r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
let resp: JsonRpcResponse = serde_json::from_str(s).unwrap();
assert_eq!(resp.id, 1);
assert!(resp.result.is_none());
+21 -18
View File
@@ -148,9 +148,7 @@ mod tests {
#[test]
fn test_default_config_denies_delete() {
let checker = PermissionChecker::new(PermissionConfig::default());
assert!(checker
.check("rm_file", &p(Permission::Delete))
.is_err());
assert!(checker.check("rm_file", &p(Permission::Delete)).is_err());
}
#[test]
@@ -246,12 +244,16 @@ mod tests {
allow_unspecified: false,
};
let checker = PermissionChecker::new(cfg);
assert!(checker
.check("t", &[Permission::Custom("db:read".into())])
.is_ok());
assert!(checker
.check("t", &[Permission::Custom("db:write".into())])
.is_err());
assert!(
checker
.check("t", &[Permission::Custom("db:read".into())])
.is_ok()
);
assert!(
checker
.check("t", &[Permission::Custom("db:write".into())])
.is_err()
);
}
#[test]
@@ -262,12 +264,11 @@ mod tests {
allow_unspecified: false,
};
let checker = PermissionChecker::new(cfg);
assert!(checker
.check(
"t",
&[Permission::Read, Permission::Network]
)
.is_ok());
assert!(
checker
.check("t", &[Permission::Read, Permission::Network])
.is_ok()
);
}
#[test]
@@ -279,8 +280,10 @@ mod tests {
};
let checker = PermissionChecker::new(cfg);
// 任一权限不在白名单则拒绝
assert!(checker
.check("t", &[Permission::Read, Permission::Write])
.is_err());
assert!(
checker
.check("t", &[Permission::Read, Permission::Write])
.is_err()
);
}
}
+10 -10
View File
@@ -7,8 +7,7 @@ use std::time::Duration;
use futures::future::join_all;
use serde_json::Value;
#[allow(deprecated)]
use crate::llm::types::ToolDefinition;
use crate::llm::types::tool::ToolDef;
use crate::tools::base::{ToolContext, ToolRef};
use crate::tools::error::ToolError;
use crate::tools::permission::PermissionChecker;
@@ -71,7 +70,6 @@ impl std::fmt::Debug for ToolRegistry {
}
}
#[allow(deprecated)]
impl ToolRegistry {
/// 创建一个新的工具注册表。
pub fn new() -> Self {
@@ -127,16 +125,15 @@ impl ToolRegistry {
self.inner.tools.keys().cloned().collect()
}
/// 获取所有工具的 `ToolDefinition` 列表(用于传递给 LLM)。
pub fn definitions(&self) -> Vec<ToolDefinition> {
/// 获取所有工具的 `ToolDef` 列表(用于传递给 LLM)。
pub fn definitions(&self) -> Vec<ToolDef> {
self.inner
.tools
.values()
.map(|tool| ToolDefinition {
.map(|tool| ToolDef {
name: tool.name().to_string(),
description: Some(tool.description().to_string()),
parameters: tool.parameters(),
strict: None,
})
.collect()
}
@@ -348,7 +345,10 @@ mod tests {
async fn test_invoke_success() {
let mut reg = ToolRegistry::new();
reg.register(Arc::new(AddTool { base: 100 })).unwrap();
let result = reg.invoke("call_1", "add", json!({ "n": 5 })).await.unwrap();
let result = reg
.invoke("call_1", "add", json!({ "n": 5 }))
.await
.unwrap();
let value = result.output.unwrap();
assert_eq!(value["result"], 105);
assert_eq!(result.tool_call_id, "call_1");
@@ -372,8 +372,8 @@ mod tests {
#[tokio::test]
async fn test_invoke_with_permission_denied() {
let mut reg = ToolRegistry::new()
.with_permission_checker(PermissionChecker::new(Default::default()));
let mut reg =
ToolRegistry::new().with_permission_checker(PermissionChecker::new(Default::default()));
reg.register(Arc::new(ShellTool)).unwrap();
let result = reg.invoke("call_z", "shell", json!({})).await;
assert!(matches!(result, Err(ToolError::PermissionDenied(_, _))));