Compare commits
47
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
32d886f870 | ||
|
|
b04427e83f | ||
|
|
d4c4d8fa3c | ||
|
|
4686063ca8 | ||
|
|
c36668071e | ||
|
|
d4f27b5865 | ||
|
|
1c0e1e0ed1 | ||
|
|
760de46623 | ||
|
|
f8df6a9421 | ||
|
|
802518b5fe | ||
|
|
993118f661 | ||
|
|
4348e4bf3e | ||
|
|
0dc91faa43 | ||
|
|
b4e5c7d651 | ||
|
|
71abe881ed | ||
|
|
2af92cd554 | ||
|
|
635942248b | ||
|
|
fe51961202 | ||
|
|
212cfcc916 | ||
|
|
88d00ac927 | ||
|
|
358e971094 | ||
|
|
85b92ae9d4 | ||
|
|
57b2fbaaed | ||
|
|
2c8e31919d | ||
|
|
c6651c9b75 | ||
|
|
e636e16820 | ||
|
|
1c89d23ba2 | ||
|
|
5b4343a051 | ||
|
|
6e1182e64c | ||
|
|
b8f4fe0fe3 | ||
|
|
7574f9c24c | ||
|
|
c82af60f81 | ||
|
|
c8a91f6eaf | ||
|
|
13edacd775 | ||
|
|
821cea8e60 | ||
|
|
517ef7db32 | ||
|
|
4cf5918b9c | ||
|
|
9da9b83167 | ||
|
|
b1875192fd | ||
|
|
d3067e2f53 | ||
|
|
5648b1d217 | ||
|
|
98dfe6c1ed | ||
|
|
9e476e79bb | ||
|
|
3bd135ec98 | ||
|
|
76f3235ed7 | ||
|
|
6315f2d008 | ||
|
|
fba78f5f33 |
@@ -198,6 +198,18 @@ pub use vector_store::VectorStore;
|
|||||||
5. **风险评估** - 潜在风险、缓解措施
|
5. **风险评估** - 潜在风险、缓解措施
|
||||||
6. **验收标准** - 可验证的完成条件
|
6. **验收标准** - 可验证的完成条件
|
||||||
|
|
||||||
|
### 进度同步规范 (docs/roadmap.md)
|
||||||
|
|
||||||
|
完成一项实施后,必须检查 `docs/roadmap.md` 是否存在对应内容;若存在,必须同步标记为完成:
|
||||||
|
|
||||||
|
- **Step / Phase 状态行**:对应 Step 加 ✅ 标记;Phase 章节末尾「状态」行从 ⏳ 改为 ✅ Phase X 全部交付物已完成
|
||||||
|
- **里程碑表**:更新对应里程碑状态从 ⏳ 改为 ✅ + 完成日期
|
||||||
|
- **依赖关系图(Mermaid)**:节点 `class` 从 `pending` / `core` 改为 `done`,必要时更新节点摘要
|
||||||
|
- **文末「已完成 / 进行中阶段」列表**:追加一行 `- ✅ Phase X — 一句话要点`
|
||||||
|
- **顶部「当前状态」**:补充新完成 Phase,更新「下一步」指向
|
||||||
|
|
||||||
|
参考案例:2026-07-05 完成 Phase 7 SqliteStore 时同步更新 6 处(顶部状态 / Phase 章节 / 依赖图 / M3 / 下一步行动 / 已完成列表)。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 项目特定规则
|
## 项目特定规则
|
||||||
|
|||||||
+124
@@ -2,6 +2,130 @@
|
|||||||
|
|
||||||
本项目所有重要变更均记录于此文件。格式参考 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/)。
|
本项目所有重要变更均记录于此文件。格式参考 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/)。
|
||||||
|
|
||||||
|
## [0.3.0] - 未发布
|
||||||
|
|
||||||
|
v0.3.0 首个增量 Phase。技术债清理 + ContextSlot fork/merge + Phase 9 审查修复。
|
||||||
|
|
||||||
|
### Breaking Changes
|
||||||
|
|
||||||
|
**类型路径变更(0.3.0):**
|
||||||
|
- `agcore::llm::types::request::ToolChoice` → `agcore::llm::types::tool::ToolChoice`(公共 re-export 路径 `agcore::llm::types::ToolChoice` 保持不变)
|
||||||
|
- `agcore::llm::types::request::StreamOptions` → `agcore::llm::provider::openai::StreamOptions`
|
||||||
|
- `agcore::llm::types::request::OpenaiChatRequest` → `agcore::llm::provider::openai::OpenaiChatRequest`
|
||||||
|
- `agcore::llm::types::response::OpenaiChatResponse` → `agcore::llm::provider::openai::OpenaiChatResponse`
|
||||||
|
- `agcore::llm::types::response::OpenaiChatChunk` → `agcore::llm::provider::openai::OpenaiChatChunk`
|
||||||
|
- 其余 `request.rs`/`response.rs` 中的 wire-format 类型(`OpenaiTool`、`AudioParam`、`Choice`、`Delta`、`ChunkChoice`、`Annotation`、`Logprobs`、`TokenLogprob`、`URLCitation` 等)同步移入 `agcore::llm::provider::openai` 模块,可见性 `pub(crate)`
|
||||||
|
|
||||||
|
**类型删除:**
|
||||||
|
- `agcore::llm::types::ChatResponse` 已删除(自 v0.1.0 标记 `#[deprecated]`,请改用 `MessageResponse`)
|
||||||
|
- `agcore::llm::types::old_stream::LegacyStreamEvent` 已删除(内部死代码)
|
||||||
|
|
||||||
|
**模块签名变化:**
|
||||||
|
- `LlmCycle::convert_request` / `convert_response` 由 `pub` 降级为 `pub(crate)`(因依赖的 `OpenaiChatRequest` / `OpenaiChatResponse` 已 `pub(crate)`)
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
**Phase 13 — ContextSlot fork/merge**
|
||||||
|
- `ContextSlot::fork(child_id, strategy)` — 从父槽派生独立子槽(数据层操作,不持久化;调用方需自行 `save()`)
|
||||||
|
- `ContextSlot::merge(child, strategy)` — 将子槽消息合并回父槽(`Append` 追加 / `Replace` 替换两种策略)
|
||||||
|
- `MergeStrategy` 枚举(`#[non_exhaustive]`,Phase 16 可扩展 `Summarize`)
|
||||||
|
- `MergeStrategy` 防御性检查:禁止 self-merge / 跨 session merge / 合并到 Readonly slot
|
||||||
|
- `agcore::agent::MergeStrategy` 公共 re-export 路径可用
|
||||||
|
- `AgentSession::derive_slot` 重构复用 `fork()` 消除重复代码(行为不变)
|
||||||
|
|
||||||
|
**Phase 9 实施审查修复(2026-07-08)**
|
||||||
|
- 2 个集成测试覆盖方案 §4 Step 5:`submit_turn_stream_end_to_end` + `submit_turn_stream_triggers_turn_hooks`
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
**Phase 13 — 技术债清理**
|
||||||
|
- 3 个旧 types 文件删除(`src/llm/types/request.rs` 187 行 + `response.rs` 177 行 + `old_stream.rs` 45 行)
|
||||||
|
- 所有 OpenAI wire-format 类型迁入 `provider/openai.rs`,可见性 `pub(crate)`
|
||||||
|
- `src/llm/stream.rs` 简化为 module doc + `pub use` 重导出(保持 `use crate::llm::stream::StreamEvent` 路径兼容,零下游破坏)
|
||||||
|
- `ToolChoice` 从 `request.rs` 迁入 `tool.rs`(serde impl 原样搬入)
|
||||||
|
|
||||||
|
**Phase 9 实施审查修复**
|
||||||
|
- `LlmCycle::run_tool_loop` 实现 `PreRequest` hook(之前 `let _ = hook_executor.as_ref()` 是空操作,导致 hook-based logging/monitoring 在流式工具循环中失效;现在与 `submit_with_tools` 行为对齐,含 `should_block` 检查,阻断时通过 `StreamEvent::Error` 事件化)
|
||||||
|
|
||||||
|
### Fixed
|
||||||
|
|
||||||
|
**Phase 9 实施审查修复**
|
||||||
|
- `AgentSession::submit_turn_stream` 末尾 `let _ = hook_executor;` 死代码移除(Arc 引用生命周期由 Arc 自动管理)
|
||||||
|
|
||||||
|
### Migration Guide (v0.2.0-rc.1 → v0.3.0)
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// ❌ v0.2.0-rc.1 — 已删除
|
||||||
|
use agcore::llm::types::ChatResponse;
|
||||||
|
use agcore::llm::types::request::OpenaiChatRequest;
|
||||||
|
|
||||||
|
// ✅ v0.3.0 — 替代路径
|
||||||
|
use agcore::llm::types::MessageResponse; // ChatResponse → MessageResponse
|
||||||
|
// OpenAI wire-format 类型为内部使用,不再公共 re-export
|
||||||
|
// 如需自定义 Provider,请直接 import agcore::llm::provider::openai::*(当前 pub(crate))
|
||||||
|
```
|
||||||
|
|
||||||
|
## [0.2.0-rc.1] - 2026-07-05
|
||||||
|
|
||||||
|
v0.2.0 候选发布。Phase 5-7 三大 P0 全部交付完成,API 稳定性扫尾,新增 2 个面向新用户的集成示例。
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
**Phase 5 — 热身准备**
|
||||||
|
- `ProviderConfig::from_env(prefix)`:从 `{prefix}_BASE_URL` / `{prefix}_API_KEY` / `{prefix}_MODEL` / `{prefix}_TIMEOUT_SECS` / `{prefix}_MAX_RETRIES` 环境变量构造配置
|
||||||
|
- `ProviderConfig::timeout_secs` / `max_retries` 字段(默认 30 / 3)
|
||||||
|
- `OllamaProvider`:本地推理 Provider(OpenAI-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` 卸载 IO;10×10 并发写入无 race
|
||||||
|
- `MemoryStore::save / get / delete / list` CRUD + prefix / since / offset+limit 过滤
|
||||||
|
- 进程重启数据不丢的 round-trip 验证
|
||||||
|
|
||||||
|
**Phase 8 — MVP 集成出口**
|
||||||
|
- `examples/quick_start`:30 行最小可运行示例(MockProvider + EchoTool + submit_turn)
|
||||||
|
- `examples/end_to_end`:3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
- **API 稳定性护栏**:14 个公开枚举标记 `#[non_exhaustive]`,覆盖 P0 核心 IR(`Message` / `ContentBlock` / `ContentBlockType` / `StreamEvent` / `HookEvent`)、P0 Error(`AgentError` / `LlmError` / `ToolError` / `MemoryError` / `PromptError`)、P1 其他(`MemoryStrategy` / `StepStatus` / `ToolChoice` / `ResponseFormat`)
|
||||||
|
- **`StepStatus::Completed` 字段类型**:从废弃的 `ChatResponse` 切换为 IR 层 `MessageResponse`(同时清理 `task_agent_demo.rs` 的 `ChatResponse` / `OpenaiChatMessage` / `FinishReason` 三处废弃类型引用)
|
||||||
|
|
||||||
|
### Non-exhaustive 清单
|
||||||
|
|
||||||
|
为防止未来新增变体时下游 exhaustive match 静默失效,14 个枚举追加 `#[non_exhaustive]`:
|
||||||
|
|
||||||
|
| 优先级 | 枚举 |
|
||||||
|
|--------|------|
|
||||||
|
| P0 核心 IR | `Message`, `ContentBlock`, `ContentBlockType`, `StreamEvent`, `HookEvent` |
|
||||||
|
| P0 Error | `AgentError`, `LlmError`, `ToolError`, `MemoryError`, `PromptError` |
|
||||||
|
| P1 其他 | `MemoryStrategy`, `StepStatus`, `ToolChoice`, `ResponseFormat` |
|
||||||
|
|
||||||
|
明确不加:内部 wire-format(`OpenaiChatMessage` 等)/ 语义已收敛(`Role` / `ServiceTier` / `Modality` / `ImageDetail` / `AudioFormat` / `StopSequence`)/ 使用面窄(`Permission` / `McpTransport` 等)。
|
||||||
|
|
||||||
|
### Deprecated
|
||||||
|
|
||||||
|
(继承自 0.1.0,无新增)`ChatResponse` / `ToolDefinition` 保持 `#[deprecated]` 标记。
|
||||||
|
|
||||||
|
### Fixed
|
||||||
|
|
||||||
|
- 修复 `StepStatus::Completed(ChatResponse)` 字段类型与 IR 体系不一致问题(已完成迁移)
|
||||||
|
|
||||||
|
### Migration Guide (v0.1 → v0.2.0-rc.1)
|
||||||
|
|
||||||
|
1. **枚举 match**:14 个 `#[non_exhaustive]` 枚举在 crate 外必须使用 `_ =>` 通配分支
|
||||||
|
2. **`StepStatus::Completed`**:字段类型从 `ChatResponse` 切换为 `MessageResponse`,需做字段映射(参考 `docs/15-phase8-mvp-integration.md` §3.1.2)
|
||||||
|
3. **`ToolDefinition` → `ToolDef`**:Phase 6 已彻底替换 `#[deprecated]` 别名,需全局重命名
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## [0.1.0] - 2026-07-04
|
## [0.1.0] - 2026-07-04
|
||||||
|
|
||||||
首个公开版本。涵盖 Phase 0-4c 的全部核心能力、Provider IR 重构、LlmCycle 简化,以及面向用户的 7 个离线示例。
|
首个公开版本。涵盖 Phase 0-4c 的全部核心能力、Provider IR 重构、LlmCycle 简化,以及面向用户的 7 个离线示例。
|
||||||
|
|||||||
+5
-2
@@ -1,6 +1,6 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "agcore"
|
name = "agcore"
|
||||||
version = "0.1.0"
|
version = "0.2.0-rc.1"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
@@ -19,8 +19,11 @@ futures-core = "0.3"
|
|||||||
bytes = "1"
|
bytes = "1"
|
||||||
async-stream = "0.3"
|
async-stream = "0.3"
|
||||||
tokio-util = { version = "0.7", features = ["rt"] }
|
tokio-util = { version = "0.7", features = ["rt"] }
|
||||||
time = { version = "0.3", features = ["serde"] }
|
time = { version = "0.3", features = ["serde", "parsing", "formatting", "macros"] }
|
||||||
|
rusqlite = { version = "0.32", features = ["bundled"] }
|
||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
dotenvy = "0.15.7"
|
dotenvy = "0.15.7"
|
||||||
wiremock = "0.6"
|
wiremock = "0.6"
|
||||||
|
temp-env = "0.3"
|
||||||
|
tempfile = "3"
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ AG Core 不是 Agent 产品,而是 Agent 的**底层依赖库**:上层应用
|
|||||||
|
|
||||||
```toml
|
```toml
|
||||||
[dependencies]
|
[dependencies]
|
||||||
agcore = "0.1"
|
agcore = "0.2"
|
||||||
tokio = { version = "1", features = ["macros", "rt-multi-thread"] }
|
tokio = { version = "1", features = ["macros", "rt-multi-thread"] }
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -110,10 +110,12 @@ let provider = create_provider(
|
|||||||
).expect("创建 Provider 失败");
|
).expect("创建 Provider 失败");
|
||||||
```
|
```
|
||||||
|
|
||||||
更多端到端示例见 [`examples/`](./examples/) 目录(共 7 个,全部可 `cargo run --example <name>`):
|
更多端到端示例见 [`examples/`](./examples/) 目录(共 10 个,全部可 `cargo run --example <name>`):
|
||||||
|
|
||||||
| 示例 | 说明 |
|
| 示例 | 说明 |
|
||||||
|------|------|
|
|------|------|
|
||||||
|
| `quick_start` | **30 行最小示例**:MockProvider + EchoTool + submit_turn,新用户 5 分钟上手 |
|
||||||
|
| `end_to_end` | **完整集成示例**:3 工具 + 3 轮对话 + SqliteStore 持久化跨连接验证 |
|
||||||
| `agent_session_demo` | Agent + 会话 + SessionMemory 完整链路(MockProvider 离线) |
|
| `agent_session_demo` | Agent + 会话 + SessionMemory 完整链路(MockProvider 离线) |
|
||||||
| `custom_tool` | 自定义工具注册、单次 / 并行调用、权限检查 |
|
| `custom_tool` | 自定义工具注册、单次 / 并行调用、权限检查 |
|
||||||
| `prompt_composer` | 提示词模板与组合器(纯离线) |
|
| `prompt_composer` | 提示词模板与组合器(纯离线) |
|
||||||
@@ -121,6 +123,7 @@ let provider = create_provider(
|
|||||||
| `conversation_memory_demo` | 对话记忆滑动窗口与隔离 |
|
| `conversation_memory_demo` | 对话记忆滑动窗口与隔离 |
|
||||||
| `knowledge_search_demo` | 知识页面关键词检索 |
|
| `knowledge_search_demo` | 知识页面关键词检索 |
|
||||||
| `streaming_events_demo` | LLM 流式响应事件消费(含错误路径) |
|
| `streaming_events_demo` | LLM 流式响应事件消费(含错误路径) |
|
||||||
|
| `simple_visit` | 真实 LLM 调用(OpenAI / Anthropic,设置 `OPENAI_*` / `ANTHROPIC_*` 环境变量) |
|
||||||
|
|
||||||
## 核心模块
|
## 核心模块
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,522 @@
|
|||||||
|
# Phase 5:热身准备 — 实施方案
|
||||||
|
|
||||||
|
## 1. 背景与目标
|
||||||
|
|
||||||
|
Phase 5 是 v0.2.0 发布周期的**热身准备阶段**,包含三个互不依赖的 Step,为后续 Phase 6-12 的端到端集成提供基础设施。
|
||||||
|
|
||||||
|
**核心目标**:
|
||||||
|
- 为 Phase 8(端到端示例)提供零 API key 的运行路径(Ollama)
|
||||||
|
- 为公共枚举的向后兼容性加上编译期护栏(`#[non_exhaustive]`)
|
||||||
|
- 为 Provider 构造提供统一的超时与重试配置入口(`ProviderConfig` 扩展)
|
||||||
|
|
||||||
|
三个 Step 之间**无依赖关系**,但出于实现效率考虑,按 **5.2 → 5.3 → 5.1** 顺序执行。理由:5.2 先新增 `Ollama` 枚举变体,5.3 再加 `#[non_exhaustive]`,避免枚举标记后添加变体需要在外部 crate 加 `_ =>` 兜底分支的困扰。
|
||||||
|
|
||||||
|
## 2. 需求分析
|
||||||
|
|
||||||
|
### Step 5.2 — Ollama Provider
|
||||||
|
|
||||||
|
| 维度 | 内容 |
|
||||||
|
|------|------|
|
||||||
|
| **需求** | 新增 `OllamaProvider`,newtype 包装 `GenericOpenaiProvider`,默认连接本地 Ollama 实例 |
|
||||||
|
| **优先级** | P0 — 为 Phase 8 端到端示例提供无需 API key 的运行路径 |
|
||||||
|
| **预期交付物** | `src/llm/provider/ollama.rs` 新建文件;`ProviderType` 新增 `Ollama` 变体 |
|
||||||
|
| **代码量** | ~55 行 |
|
||||||
|
|
||||||
|
### Step 5.3 — `#[non_exhaustive]` 前置标记
|
||||||
|
|
||||||
|
| 维度 | 内容 |
|
||||||
|
|------|------|
|
||||||
|
| **需求** | 为 4 个公共枚举添加 `#[non_exhaustive]` 属性,避免后续新增变体时破坏下游 match |
|
||||||
|
| **优先级** | P1 — 编译期兼容性保障 |
|
||||||
|
| **预期交付物** | 修改 4 个枚举定义,各加一行属性 |
|
||||||
|
| **代码量** | ~4 行 |
|
||||||
|
|
||||||
|
### Step 5.1 — ProviderConfig 扩展
|
||||||
|
|
||||||
|
| 维度 | 内容 |
|
||||||
|
|------|------|
|
||||||
|
| **需求** | `ProviderConfig` 新增 `timeout_secs` 和 `max_retries` 字段;实现 `Default`、`from_env()` 构造;timeout 传导到各 Provider HTTP Client |
|
||||||
|
| **优先级** | P0 — 与 Roadmap 一致,Phase 8(MVP 出口)依赖 from_env |
|
||||||
|
| **预期交付物** | `ProviderConfig` 扩展;`create_provider()` 超时注入;`from_env()` + 单元测试 |
|
||||||
|
| **代码量** | ~60 行 + 测试 |
|
||||||
|
|
||||||
|
## 3. 方案设计
|
||||||
|
|
||||||
|
### 3.1 Step 5.2 — Ollama Provider(先执行)
|
||||||
|
|
||||||
|
#### 改动文件清单
|
||||||
|
|
||||||
|
| 文件 | 操作 | 说明 |
|
||||||
|
|------|------|------|
|
||||||
|
| `src/llm/provider/ollama.rs` | **新建** | OllamaProvider newtype 包装 |
|
||||||
|
| `src/llm/provider.rs` | 修改 | `ProviderType` 新增 `Ollama` 变体;`FromStr` 加解析;`create_provider()` 加分支 |
|
||||||
|
| `src/llm/provider/mod.rs` 或其他模块注册文件 | 修改(如需要) | 注册 `pub mod ollama` |
|
||||||
|
|
||||||
|
#### 关键代码
|
||||||
|
|
||||||
|
**`src/llm/provider/ollama.rs`**(新建):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
//! Ollama Provider —— OpenAI-compatible 协议的 newtype 包装,零 API key。
|
||||||
|
//!
|
||||||
|
//! 默认 base_url = `http://localhost:11434/v1`,空 api_key 也可工作。
|
||||||
|
//! 实现方式同 DeepSeekProvider / QwenProvider,共享 GenericOpenaiProvider 的 HTTP/SSE/转换逻辑。
|
||||||
|
|
||||||
|
use reqwest::Client;
|
||||||
|
use std::pin::Pin;
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use futures_core::Stream;
|
||||||
|
|
||||||
|
use super::openai::GenericOpenaiProvider;
|
||||||
|
use super::{LlmProvider, ProviderCapabilities};
|
||||||
|
use crate::llm::error::LlmError;
|
||||||
|
use crate::llm::types::request_v2::MessageRequest;
|
||||||
|
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
||||||
|
|
||||||
|
pub struct OllamaProvider(pub GenericOpenaiProvider);
|
||||||
|
|
||||||
|
impl OllamaProvider {
|
||||||
|
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
||||||
|
let url = if base_url.is_empty() {
|
||||||
|
"http://localhost:11434/v1".to_string()
|
||||||
|
} else {
|
||||||
|
base_url
|
||||||
|
};
|
||||||
|
Self(GenericOpenaiProvider::new_with_name(
|
||||||
|
url,
|
||||||
|
api_key,
|
||||||
|
model,
|
||||||
|
"ollama",
|
||||||
|
))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
|
||||||
|
/// 与 `OpenaiChatProvider::with_client` 和 `DeepSeekProvider::with_client` 一致。
|
||||||
|
pub fn with_client(self, client: Client) -> Self {
|
||||||
|
Self(self.0.with_client(client))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl LlmProvider for OllamaProvider {
|
||||||
|
async fn chat(&self, request: MessageRequest) -> Result<MessageResponse, LlmError> {
|
||||||
|
self.0.chat(request).await
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn chat_stream(
|
||||||
|
&self,
|
||||||
|
request: MessageRequest,
|
||||||
|
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||||
|
self.0.chat_stream(request).await
|
||||||
|
}
|
||||||
|
|
||||||
|
fn capabilities(&self) -> ProviderCapabilities {
|
||||||
|
let mut caps = self.0.capabilities();
|
||||||
|
caps.provider_name = "ollama";
|
||||||
|
caps
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**`src/llm/provider.rs`** 的修改:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// ProviderType 新增变体
|
||||||
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
|
pub enum ProviderType {
|
||||||
|
OpenaiChat,
|
||||||
|
OpenaiResponse,
|
||||||
|
Anthropic,
|
||||||
|
DeepSeek,
|
||||||
|
Qwen,
|
||||||
|
/// Ollama(本地),默认 base_url = `http://localhost:11434/v1`。
|
||||||
|
Ollama,
|
||||||
|
}
|
||||||
|
|
||||||
|
// FromStr 加解析
|
||||||
|
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||||
|
match s.to_lowercase().as_str() {
|
||||||
|
// ... 已有条目 ...
|
||||||
|
"ollama" => Ok(ProviderType::Ollama),
|
||||||
|
_ => Err(format!("未知的 Provider 类型: {s}")),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// create_provider() 加分支
|
||||||
|
// Step 5.2 阶段仅展示基本构造。Step 5.1(ProviderConfig 扩展)
|
||||||
|
// 执行到此分支时,将同步补充 with_client 链式调用注入 timeout:
|
||||||
|
//
|
||||||
|
// let client = Client::builder()
|
||||||
|
// .timeout(Duration::from_secs(config.timeout_secs))
|
||||||
|
// .build()?;
|
||||||
|
// Ok(Box::new(
|
||||||
|
// ollama::OllamaProvider::new(config.base_url, config.api_key, config.model)
|
||||||
|
// .with_client(client),
|
||||||
|
// ))
|
||||||
|
ProviderType::Ollama => Ok(Box::new(ollama::OllamaProvider::new(
|
||||||
|
config.base_url,
|
||||||
|
config.api_key,
|
||||||
|
config.model,
|
||||||
|
))),
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 集成方式
|
||||||
|
|
||||||
|
OllamaProvider 的 newtype 包装模式与 `DeepSeekProvider`、`QwenProvider` 完全一致,`LlmProvider` trait 委托给 `self.0`。`capabilities().provider_name` 返回 `"ollama"`。
|
||||||
|
|
||||||
|
### 3.2 Step 5.3 — `#[non_exhaustive]` 前置标记
|
||||||
|
|
||||||
|
#### 改动文件清单
|
||||||
|
|
||||||
|
| 文件 | 行号 | 枚举 | 操作 |
|
||||||
|
|------|------|------|------|
|
||||||
|
| `src/llm/provider.rs` | ~21 | `ProviderType` | 加 `#[non_exhaustive]` |
|
||||||
|
| `src/llm/types/response_v2.rs` | ~22 | `StopReason` | 加 `#[non_exhaustive]` |
|
||||||
|
| `src/llm/types/shared.rs` | ~16 | `FinishReason` | 加 `#[non_exhaustive]` |
|
||||||
|
| `src/memory/store.rs` | ~35 | `EvictionPolicy` | 加 `#[non_exhaustive]` |
|
||||||
|
|
||||||
|
**排除清单**:`SlotMode`。
|
||||||
|
|
||||||
|
**决策理由**:`SlotMode` 枚举在 Phase 10(`src/llm/context.rs`)中才实际定义,Phase 5 尚不存在此类型。`#[non_exhaustive]` 无法标注不存在的枚举,因此排除标注。Roadmap(v0.2.0 §Phase 5 Step 5.3)列出的 `SlotMode`(预置) 推迟到 Phase 10 实现时一并添加。
|
||||||
|
|
||||||
|
#### 关键代码
|
||||||
|
|
||||||
|
每个枚举在 `derive` 上方或下方加一行属性:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||||
|
#[non_exhaustive]
|
||||||
|
pub enum ProviderType {
|
||||||
|
// ...
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 影响分析
|
||||||
|
|
||||||
|
- `#[non_exhaustive]` 是纯编译期属性,不影响运行时行为
|
||||||
|
- 同一 crate 内的 exhaustive match 不受影响(同 crate 可穷举)
|
||||||
|
- 下游 crate 的 match 必须加 `_ =>` 兜底分支,这是期望行为——确保未来新增变体时不会 silent break
|
||||||
|
- **单向门**:此步骤一旦通过 `v0.2.0` 发布到公共 API 后,**不可回退**。回退意味着移除 `#[non_exhaustive]`,可能破坏已添加 `_ =>` 的下游代码。因此必须在发布前完成并确认所有枚举变体正确
|
||||||
|
|
||||||
|
### 3.3 Step 5.1 — ProviderConfig 扩展(最后执行)
|
||||||
|
|
||||||
|
#### 改动文件清单
|
||||||
|
|
||||||
|
| 文件 | 操作 | 说明 |
|
||||||
|
|------|------|------|
|
||||||
|
| `src/llm/provider.rs` | 修改 | `ProviderConfig` 加字段;加 `impl Default`;加 `from_env()`;`create_provider` 注入 timeout |
|
||||||
|
| `src/llm/provider/openai.rs` | 修改 | `GenericOpenaiProvider` 新增 `timeout_secs` 字段;`new_with_name` 接受 timeout 参数;`map_reqwest_error` 参数化 |
|
||||||
|
| `src/llm/provider/anthropic.rs` | 修改 | 新增 `timeout_secs` 字段;`new()` 接受 timeout 参数;`map_reqwest_error` 参数化 |
|
||||||
|
| `src/llm/provider/anthropic.rs` | 修改 | 新增 `with_timeout()` 方法(返回 `Result<Self, LlmError>`) |
|
||||||
|
| `src/llm/provider/openai_compat.rs` | 修改 | `DeepSeekProvider` 和 `QwenProvider` 新增公开 `with_client()` 方法 |
|
||||||
|
| `src/llm/provider/ollama.rs` | 修改 | `OllamaProvider` 新增公开 `with_client()` 方法 |
|
||||||
|
| `Cargo.toml` | 修改 | 加 `temp_env` dev-dependency |
|
||||||
|
| 测试文件(`provider.rs` 内联或独立) | 新增 | `from_env` 单元测试 + timeout 传导集成测试 |
|
||||||
|
|
||||||
|
#### 数据结构
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// Provider 构造参数 —— 通用 base_url + api_key + model + timeout/retry 配置。
|
||||||
|
pub struct ProviderConfig {
|
||||||
|
pub base_url: String,
|
||||||
|
pub api_key: String,
|
||||||
|
pub model: String,
|
||||||
|
/// 请求超时秒数(默认 30)。应用于 Provider 的 HTTP Client 级别。
|
||||||
|
pub timeout_secs: u64,
|
||||||
|
/// 最大重试次数(默认 3)。当前此字段仅由 `from_env()` 采集,
|
||||||
|
/// 实际重试逻辑由 `CycleConfig.retry.max_retries` 控制。
|
||||||
|
/// 未来可合并到统一的 retry 配置。
|
||||||
|
pub max_retries: u32,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for ProviderConfig {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self {
|
||||||
|
base_url: String::new(),
|
||||||
|
api_key: String::new(),
|
||||||
|
model: String::new(),
|
||||||
|
timeout_secs: 30,
|
||||||
|
max_retries: 3,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProviderConfig {
|
||||||
|
/// 从环境变量构造 ProviderConfig。
|
||||||
|
///
|
||||||
|
/// 必填变量:
|
||||||
|
/// - `{prefix}_BASE_URL`
|
||||||
|
/// - `{prefix}_API_KEY`
|
||||||
|
/// - `{prefix}_MODEL`
|
||||||
|
///
|
||||||
|
/// 可选变量(有默认值):
|
||||||
|
/// - `{prefix}_TIMEOUT_SECS`(默认 30)
|
||||||
|
/// - `{prefix}_MAX_RETRIES`(默认 3)
|
||||||
|
pub fn from_env(prefix: &str) -> Result<Self, String> {
|
||||||
|
let base_url = std::env::var(format!("{prefix}_BASE_URL"))
|
||||||
|
.map_err(|_| format!("{prefix}_BASE_URL 环境变量未设置"))?;
|
||||||
|
let api_key = std::env::var(format!("{prefix}_API_KEY"))
|
||||||
|
.map_err(|_| format!("{prefix}_API_KEY 环境变量未设置"))?;
|
||||||
|
let model = std::env::var(format!("{prefix}_MODEL"))
|
||||||
|
.map_err(|_| format!("{prefix}_MODEL 环境变量未设置"))?;
|
||||||
|
let timeout_secs = match std::env::var(format!("{prefix}_TIMEOUT_SECS")) {
|
||||||
|
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||||
|
tracing::warn!("{prefix}_TIMEOUT_SECS='{v}' 解析失败,使用默认值 30");
|
||||||
|
30
|
||||||
|
}),
|
||||||
|
Err(_) => 30,
|
||||||
|
};
|
||||||
|
let max_retries = match std::env::var(format!("{prefix}_MAX_RETRIES")) {
|
||||||
|
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||||
|
tracing::warn!("{prefix}_MAX_RETRIES='{v}' 解析失败,使用默认值 3");
|
||||||
|
3
|
||||||
|
}),
|
||||||
|
Err(_) => 3,
|
||||||
|
};
|
||||||
|
|
||||||
|
// ponytail: max_retries 当前仅采集,不传入 Provider。
|
||||||
|
// 实际重试由 CycleConfig.retry.max_retries 控制。
|
||||||
|
// 此 warn 在应用启动时通常只触发一次,多次调用 from_env 时
|
||||||
|
// 重复输出的风险低。如有噪声,可改用 std::sync::Once 控制。
|
||||||
|
if max_retries != 3 {
|
||||||
|
tracing::warn!(
|
||||||
|
"ProviderConfig.max_retries={} 已采集但当前未生效;\
|
||||||
|
重试次数由 CycleConfig.retry.max_retries 控制",
|
||||||
|
max_retries,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(Self {
|
||||||
|
base_url,
|
||||||
|
api_key,
|
||||||
|
model,
|
||||||
|
timeout_secs,
|
||||||
|
max_retries,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Timeout 传导模式
|
||||||
|
|
||||||
|
在 `create_provider()` 中,对基于 `GenericOpenaiProvider` 的 Provider(OpenAI / DeepSeek / Qwen / Ollama),通过同一模式注入 timeout:构造带 timeout 的 `Client` 后调用 `with_client(client)`。
|
||||||
|
|
||||||
|
所有 OpenAI-compatible 分支新增的 `with_client()` 公开方法:
|
||||||
|
|
||||||
|
| Provider | 方法 | 位置 |
|
||||||
|
|----------|------|------|
|
||||||
|
| `OpenaiChatProvider` | 已有 `with_client(Client) -> Self` | `openai.rs` |
|
||||||
|
| `DeepSeekProvider` | 新增 `with_client(Client) -> Self` | `openai_compat.rs` |
|
||||||
|
| `QwenProvider` | 新增 `with_client(Client) -> Self` | `openai_compat.rs` |
|
||||||
|
| `OllamaProvider` | 新增 `with_client(Client) -> Self` | `ollama.rs`(新建文件) |
|
||||||
|
|
||||||
|
**关于 `new_with_client` 的说明**:`DeepSeekProvider` 和 `QwenProvider` 当前已有测试用的 `new_with_client(base_url, api_key, model, client)` 方法(通过 `inner.http_client = client` 直接写字段)。新增 `with_client` 后,`new_with_client` 应重构为 `Self::new(base_url, api_key, model).with_client(client)` 代理,统一走公开 API 路径。
|
||||||
|
|
||||||
|
代码示例(以 DeepSeek 为例,OpenAI/Qwen/Ollama 模式完全一致):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
ProviderType::DeepSeek => {
|
||||||
|
let client = Client::builder()
|
||||||
|
.timeout(Duration::from_secs(config.timeout_secs))
|
||||||
|
.build()
|
||||||
|
.map_err(|e| LlmError::Other(format!("创建 HTTP 客户端失败: {e}")))?;
|
||||||
|
Ok(Box::new(
|
||||||
|
openai_compat::DeepSeekProvider::new(
|
||||||
|
config.base_url,
|
||||||
|
config.api_key,
|
||||||
|
config.model,
|
||||||
|
)
|
||||||
|
.with_client(client),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Anthropic 由于需要保留 `default_headers`,使用独立的 `with_timeout` 模式:
|
||||||
|
|
||||||
|
AnthropicProvider 新增 `with_timeout` 方法:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
impl AnthropicProvider {
|
||||||
|
/// 替换默认 HTTP Client 的超时配置。
|
||||||
|
///
|
||||||
|
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后原有通过 `with_client`
|
||||||
|
/// 注入的 Client 将被替换。headers 逻辑与 `new()` 中的构造保持一致。
|
||||||
|
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
|
||||||
|
// ponytail: 重建 http_client 时保留已有默认 headers(x-api-key / anthropic-version)。
|
||||||
|
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
|
||||||
|
let key_header = HeaderValue::from_str(&self.api_key)
|
||||||
|
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
|
||||||
|
let version_header = HeaderValue::from_static("2023-06-01");
|
||||||
|
|
||||||
|
self.http_client = Client::builder()
|
||||||
|
.timeout(Duration::from_secs(secs))
|
||||||
|
.default_headers({
|
||||||
|
let mut headers = HeaderMap::new();
|
||||||
|
headers.insert("x-api-key", key_header);
|
||||||
|
headers.insert("anthropic-version", version_header);
|
||||||
|
headers
|
||||||
|
})
|
||||||
|
.build()
|
||||||
|
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
|
||||||
|
Ok(self)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### `map_reqwest_error` 中的硬编码超时修复
|
||||||
|
|
||||||
|
`openai.rs` 和 `anthropic.rs` 中的 `map_reqwest_error` 辅助函数当前在超时错误中返回硬编码的 `Duration::from_secs(120)`:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// 现状 —— 硬编码 120s,与可配置 timeout 脱节
|
||||||
|
LlmError::Timeout { duration: Duration::from_secs(120) }
|
||||||
|
```
|
||||||
|
|
||||||
|
**修复方式**:采用**方案 A**——在 Provider struct 中存储 `timeout_secs` 字段,`map_reqwest_error` 读取该字段的值而非硬编码 120s。
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// 修复后 —— 参数化,从 Provider 存储的 timeout_secs 读取
|
||||||
|
// GenericOpenaiProvider 新增 timeout_secs 字段:
|
||||||
|
pub struct GenericOpenaiProvider {
|
||||||
|
http_client: Client,
|
||||||
|
base_url: String,
|
||||||
|
api_key: String,
|
||||||
|
model: String,
|
||||||
|
provider_name: &'static str,
|
||||||
|
extra_headers: Vec<(String, String)>,
|
||||||
|
timeout_secs: u64, // ← 新增,由 new_with_name 的参数传入
|
||||||
|
}
|
||||||
|
|
||||||
|
// map_reqwest_error 使用 self.timeout_secs 而非硬编码 120:
|
||||||
|
LlmError::Timeout { duration: Duration::from_secs(self.timeout_secs) }
|
||||||
|
```
|
||||||
|
|
||||||
|
**方案 B(从 reqwest::Client 提取 timeout)已被否决**:`reqwest::Client` 不提供 timeout getter,无法从已构造的 client 中反向读取超时配置。
|
||||||
|
|
||||||
|
如果漏掉此修复,用户设置 `AG_LLM_TIMEOUT_SECS=60` 后超时,错误消息仍显示 "LLM 请求超时(120s)",与实际配置不符。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### max_retries 说明
|
||||||
|
|
||||||
|
`ProviderConfig.max_retries` 当前仅由 `from_env()` 采集存储,**实际重试操作由 `CycleConfig.retry.max_retries` 控制**。两者之间的关系通过文档注释声明:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// 最大重试次数(默认 3)。当前此字段仅由 `from_env()` 采集,
|
||||||
|
/// 实际重试逻辑由 `CycleConfig.retry.max_retries` 控制。
|
||||||
|
/// 未来 Phase 6+ 可统一合并此字段到 CycleConfig。
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 测试设计
|
||||||
|
|
||||||
|
使用 `temp_env` 在单元测试中隔离环境变量:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_from_env_requires_all_vars() {
|
||||||
|
// 未设置任何变量时应返回 Err
|
||||||
|
let result = ProviderConfig::from_env("TEST_PROVIDER");
|
||||||
|
assert!(result.is_err());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_from_env_uses_defaults() {
|
||||||
|
temp_env::with_vars([
|
||||||
|
("TEST_PROVIDER_BASE_URL", Some("http://localhost:11434/v1")),
|
||||||
|
("TEST_PROVIDER_API_KEY", Some("")),
|
||||||
|
("TEST_PROVIDER_MODEL", Some("llama3")),
|
||||||
|
], || {
|
||||||
|
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||||
|
assert_eq!(config.timeout_secs, 30);
|
||||||
|
assert_eq!(config.max_retries, 3);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_from_env_reads_custom_timeout() {
|
||||||
|
temp_env::with_vars([
|
||||||
|
("TEST_PROVIDER_BASE_URL", Some("http://x")),
|
||||||
|
("TEST_PROVIDER_API_KEY", Some("k")),
|
||||||
|
("TEST_PROVIDER_MODEL", Some("m")),
|
||||||
|
("TEST_PROVIDER_TIMEOUT_SECS", Some("60")),
|
||||||
|
("TEST_PROVIDER_MAX_RETRIES", Some("5")),
|
||||||
|
], || {
|
||||||
|
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||||
|
assert_eq!(config.timeout_secs, 60);
|
||||||
|
assert_eq!(config.max_retries, 5);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 4. 实现计划
|
||||||
|
|
||||||
|
### Step 5.2 — Ollama Provider(~55 行)
|
||||||
|
|
||||||
|
| 步骤 | 操作 | 验证 |
|
||||||
|
|------|------|------|
|
||||||
|
| 1 | 创建 `src/llm/provider/ollama.rs`,实现 `OllamaProvider` newtype | 编译通过 |
|
||||||
|
| 2 | 在 `provider.rs` 注册 `pub mod ollama` | 编译通过 |
|
||||||
|
| 3 | `ProviderType` 新增 `Ollama` 变体 | 编译通过 |
|
||||||
|
| 4 | `FromStr` 加 `"ollama"` 解析 | 编译通过 |
|
||||||
|
| 5 | `create_provider()` 加 `Ollama =>` 分支 | 编译通过 |
|
||||||
|
| 6 | 运行 `cargo build` | 无错误 |
|
||||||
|
|
||||||
|
### Step 5.3 — `#[non_exhaustive]` 前置标记(~4 行)
|
||||||
|
|
||||||
|
| 步骤 | 操作 | 验证 |
|
||||||
|
|------|------|------|
|
||||||
|
| 1 | `ProviderType`(`provider.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||||
|
| 2 | `StopReason`(`response_v2.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||||
|
| 3 | `FinishReason`(`shared.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||||
|
| 4 | `EvictionPolicy`(`memory/store.rs`)加 `#[non_exhaustive]` | 编译通过 |
|
||||||
|
| 5 | 运行 `cargo build --all-targets` | 无 warning |
|
||||||
|
|
||||||
|
### Step 5.1 — ProviderConfig 扩展(~60 行 + 测试)
|
||||||
|
|
||||||
|
| 步骤 | 操作 | 验证 |
|
||||||
|
|------|------|------|
|
||||||
|
| 1 | `ProviderConfig` 加 `timeout_secs` / `max_retries` 字段 | 编译通过 |
|
||||||
|
| 2 | 实现 `impl Default for ProviderConfig` | 编译通过 |
|
||||||
|
| 3 | 实现 `ProviderConfig::from_env()` | 编译通过 |
|
||||||
|
| 4 | `GenericOpenaiProvider` 和 `AnthropicProvider` 新增 `timeout_secs` 字段,`new_with_name`/`new()` 接受 timeout 参数 | 编译通过 |
|
||||||
|
| 5 | `map_reqwest_error` 在各 Provider 中改为从 `self.timeout_secs` 读取,移除硬编码 120s | 编译通过 |
|
||||||
|
| 6 | `create_provider()` 中各分支注入 timeout(OpenAI-compatible 用 `Client::builder().timeout()` + `with_client`;Anthropic 用 `with_timeout()`) | 编译通过 |
|
||||||
|
| 7 | `DeepSeekProvider`/`QwenProvider` 的 `new_with_client` 重构为 `Self::new(...).with_client(client)` 代理 | 测试通过 |
|
||||||
|
| 8 | `Cargo.toml` 添加 `temp_env` dev-dependency | `cargo build` 通过 |
|
||||||
|
| 8 | 添加 `from_env` 单元测试 + timeout 传导集成测试 | `cargo test` 通过 |
|
||||||
|
| 9 | 完整验证 | 见第 6 节 |
|
||||||
|
|
||||||
|
## 5. 风险评估
|
||||||
|
|
||||||
|
| 风险 | 影响 | 概率 | 缓解措施 |
|
||||||
|
|------|------|------|----------|
|
||||||
|
| `create_provider()` 中 `Client::builder().build()` 返回 `Result`,当前代码使用 `.expect()`,改为 `map_err` 转为 `LlmError` 后需确保所有分支正确转换 | 编译期强制处理,遗漏分支直接报错 | 低 | `create_provider` 返回 `Result<Box<dyn LlmProvider>, LlmError>`,`map_err` 天然适配。新增的 timeout 注入路径逐一检查 |
|
||||||
|
| `AnthropicProvider` 的 `default_headers` 在 `with_timeout` 中重建时与 `new()` 中的 headers 不一致 | Anthropic 认证失败 | 低 | `with_timeout` 方法复制 `new()` 中的 headers 构造逻辑。通过已有测试验证认证通过 |
|
||||||
|
| Ollama 实际运行时行为差异:版本兼容性、API 路径、模型名等 | 运行时才能发现 | 中 | Phase 5 仅做类型级验证(`cargo build`),Phase 8 端到端测试时通过 Ollama mock 或真实实例验证 |
|
||||||
|
| `max_retries` 存储了却未实际使用,造成困惑 | 开发者误以为已生效 | 中 | 通过文档注释明确声明 `max_retries` 当前仅采集,实际重试由 `CycleConfig.retry.max_retries` 控制 |
|
||||||
|
| `temp_env` 测试在多线程并发测试中互相污染环境变量 | 偶发测试失败 | 中(Rust 默认单线程测试用 `--test-threads=1` 可避免) | 将 `from_env` 测试控制在同一测试文件,避免并行执行。必要时在 CI 中确保 `--test-threads=1` |
|
||||||
|
|
||||||
|
## 6. 验收标准
|
||||||
|
|
||||||
|
以下条件**全部满足**方可认为 Phase 5 完成:
|
||||||
|
|
||||||
|
- [ ] `cargo build --all-targets` 通过,无错误
|
||||||
|
- [ ] `cargo test --all-targets` 通过,新增测试覆盖 `from_env` 的必填/选填/默认值场景
|
||||||
|
- [ ] `cargo clippy --all-targets -- -D warnings` 通过,无任何 warning
|
||||||
|
- [ ] `cargo doc --no-deps -D warnings` 通过,所有公共 API 有文档注释(`///`)
|
||||||
|
- [ ] 新增文件:1(`ollama.rs`)
|
||||||
|
- [ ] 修改文件:9(`provider.rs`、`openai.rs`、`anthropic.rs`、`openai_compat.rs`、`response_v2.rs`、`shared.rs`、`store.rs`、`Cargo.toml`、测试文件)
|
||||||
|
- [ ] 净代码增量:~160 行
|
||||||
|
- [ ] `ProviderType` 新增 `Ollama` 变体,`"ollama"` 字符串可解析
|
||||||
|
- [ ] 4 个公共枚举带有 `#[non_exhaustive]` 属性
|
||||||
|
- [ ] `ProviderConfig` 可从环境变量构造(`from_env()`),含默认值
|
||||||
|
- [ ] timeout 值已传导到 `create_provider()` 中各 Provider 的 HTTP Client 配置
|
||||||
|
- [ ] timeout 传导验证通过至少一个端到端 wiremock 集成测试(模拟 HTTP 服务在超时后返回 408,验证 Provider 返回 `LlmError::Timeout`)
|
||||||
|
- [ ] `DeepSeekProvider`、`QwenProvider`、`OllamaProvider` 均有公开 `with_client()` 方法,可在 `create_provider` 中注入 timeout Client
|
||||||
|
- [ ] `map_reqwest_error` 中不再硬编码 `Duration::from_secs(120)`,改为参数化读取
|
||||||
@@ -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 Compat(DeepSeek、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+ 阶段处理 |
|
||||||
@@ -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 8(MVP 出口)和 Phase 10(ContextSlot)的前置依赖。
|
||||||
|
|
||||||
|
**成功标准**:
|
||||||
|
- SqliteStore 完整实现 MemoryStore trait(4 个方法: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 存储反而需要手动处理排序逻辑。非必要不引入新存储范式。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 推荐方案
|
||||||
|
|
||||||
|
### 总体方向:方案 A(rusqlite + 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 自然完成或 panic,Mutex 通过 `.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/Redis,v0.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 B1(rusqlite + tempfile 依赖)、Task A1(store/ 目录存在) |
|
||||||
|
| 预估工作量 | M(1-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 B2(sqlite_store.rs 文件存在) |
|
||||||
|
| 预估工作量 | XS(< 15min) |
|
||||||
|
| 风险等级 | 低 |
|
||||||
|
| 验收条件 | `cargo build` 通过,`SqliteStore` 可从 `agcore::memory::SqliteStore` 路径访问 |
|
||||||
|
|
||||||
|
#### Task B4 — 编写测试
|
||||||
|
|
||||||
|
| 项目 | 内容 |
|
||||||
|
|------|------|
|
||||||
|
| 任务描述 | 在 `sqlite_store.rs` 中编写 `#[cfg(test)] mod tests`,覆盖: |
|
||||||
|
| | (1)CRUD 基本操作(save → get → list → delete → get None) |
|
||||||
|
| | (2)Upsert 语义(同 id 重复 save 覆盖内容,created_at 保持调用方传入值) |
|
||||||
|
| | (3)prefix 过滤(MemoryFilter.prefix) |
|
||||||
|
| | (4)时间范围过滤(MemoryFilter.since) |
|
||||||
|
| | (5)offset/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(实现完成) |
|
||||||
|
| 预估工作量 | M(1-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
|
||||||
@@ -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 1(StepStatus 先标记 `#[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+2(CHANGELOG 需记录实际变更) |
|
||||||
|
| 预估工作量 | 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 Agent(name / 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(示例编写模式已建立);SqliteStore(Phase 7 已完成) |
|
||||||
|
| 预估工作量 | M(1-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() {
|
||||||
|
// 使用真实 Provider(AG_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 5(Quick Start) | 零文件重叠 |
|
||||||
|
| 2 | commit 2(StepStatus 修复) | commit 5(Quick Start) | 零文件重叠 |
|
||||||
|
| 3 | commit 5(Quick 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 个示例 |
|
||||||
@@ -0,0 +1,827 @@
|
|||||||
|
# Phase 9 — 流式体验增强实施方案
|
||||||
|
|
||||||
|
- **文档编号**:16
|
||||||
|
- **标题**:Phase 9 — 流式体验增强实施方案
|
||||||
|
- **日期**:2026-07-05
|
||||||
|
- **状态**:已定稿
|
||||||
|
- **涉及模块**:llm/cycle、llm/types/response_v2、agent/session
|
||||||
|
- **关联文档**:roadmap.md(§Phase 9)、15-phase8-mvp-integration.md
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. 背景与目标
|
||||||
|
|
||||||
|
agcore 已发布 v0.2.0-rc.1,Phase 0-8 全部完成。当前 Agent 会话只有非流式 API(`submit_turn`),开发者无法看到实时 token 输出和工具执行过程。Phase 9 的目标是为 `AgentSession` 新增流式方法 `submit_turn_stream`,让开发者能实时看到 LLM token 生成和工具执行状态。
|
||||||
|
|
||||||
|
### 1.1 现有能力
|
||||||
|
|
||||||
|
| 能力 | 方法 | 流式 | 自动工具循环 | 状态 |
|
||||||
|
|------|------|------|-------------|------|
|
||||||
|
| LLM 流式请求 | `LlmCycle::submit_stream` | ✅ | ❌ | 已就绪 |
|
||||||
|
| LLM 工具循环 | `LlmCycle::submit_with_tools` | ❌ | ✅ | 已就绪 |
|
||||||
|
| Agent 会话 | `AgentSession::submit_turn` | ❌ | ✅ | 已就绪 |
|
||||||
|
| 流事件枚举 | `StreamEvent`(11 变体) | — | — | 缺工具执行事件 |
|
||||||
|
| Mock 流 | `MockProvider::chat_stream` | ✅ | — | 可模拟流事件序列 |
|
||||||
|
|
||||||
|
### 1.2 核心矛盾
|
||||||
|
|
||||||
|
流式能力和工具循环能力分别存在于两个方法中,从未被组合。`submit_stream` 只管将 LLM 流事件原样转发,不理解工具调用;`submit_with_tools` 自动执行工具循环但全程阻塞。Phase 9 就是要组合它们:**在工具循环中,每一轮 LLM 调用都是流式的,并在工具执行前后插入语义事件**。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. 需求分析
|
||||||
|
|
||||||
|
### 2.1 功能需求
|
||||||
|
|
||||||
|
1. **`AgentSession::submit_turn_stream(user_input)`** — 返回 `StreamEvent` 流,开发者通过 `while let Some(event) = stream.next().await` 逐事件消费
|
||||||
|
2. **流式工具循环** — 多轮工具调用过程中流不卡死,每轮工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||||
|
3. **`finalize_turn(response)`** — 流消费完成后同步 session 状态(cost 累计 + `OnTurnEnd` hook 触发)
|
||||||
|
4. **新增 `StreamEvent` 变体** — `ToolExecutionStarted` + `ToolExecutionCompleted`,携带工具名称、调用 ID、参数/结果摘要
|
||||||
|
|
||||||
|
### 2.2 非功能需求
|
||||||
|
|
||||||
|
- **零影响**:现有 `submit_turn` 和 `submit_with_tools` 行为不变,存量测试 0 回归
|
||||||
|
- **异步流**:消费者通过 `futures_util::StreamExt::next()` 逐事件消费
|
||||||
|
- **错误事件化**:错误通过 `StreamEvent::Error` 事件表达,不通过 `Result` 通道终止流
|
||||||
|
- **最少代码**:复用现有 `submit_with_tools` 的工具循环逻辑模式和 `submit_stream` 的流管道模式
|
||||||
|
|
||||||
|
### 2.3 不做事项
|
||||||
|
|
||||||
|
| 事项 | 理由 |
|
||||||
|
|------|------|
|
||||||
|
| 新增示例(Phase 9.2 再加) | 缩窄 Phase 9 范围至核心能力 |
|
||||||
|
| `OnTurnEnd` 自动触发 | Rust 所有权约束:流是延迟求值,`&mut self` 无法进入闭包;由消费者收到 `MessageComplete` 后手动调用 `finalize_turn` |
|
||||||
|
| 修复 cost 统计 | 中间轮 cost 丢失是已知限制,与 `submit_turn` 行为一致 |
|
||||||
|
| 跨 turn 消息历史保留 | Phase 10 `ContextSlot` 的职责 |
|
||||||
|
| 并行 tool 调用的事件细化 | 当前工具调用是顺序 `for` 循环,并行化留待后续优化 |
|
||||||
|
| `run_tool_loop` 内消息压缩 | `run_tool_loop` 不接收 `compact_config` 参数,不执行上下文压缩。长工具循环中消息增长可能导致 context window 溢出,这是流式实现的已知限制。后续可通过传递 `compact_config` 给 `run_tool_loop` 支持 |
|
||||||
|
| LLM 请求自动 retry | 流式版本不在 `run_tool_loop` 内部实现 retry(详见 §3.6 说明)。调用方可自行包装 `RetryProvider` 或在 `LlmProvider` 实现层处理 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. 方案设计
|
||||||
|
|
||||||
|
### 3.1 架构总览
|
||||||
|
|
||||||
|
```
|
||||||
|
┌──────────────────────────────────────────────────────────────┐
|
||||||
|
│ AgentSession │
|
||||||
|
│ ┌──────────────────────────────────────────────────────┐ │
|
||||||
|
│ │ submit_turn_stream() │ │
|
||||||
|
│ │ ├─ OnTurnStart hook(同步触发,返回流之前) │ │
|
||||||
|
│ │ ├─ 组装 LlmCycle(system_prompt / compact_config) │ │
|
||||||
|
│ │ ├─ 调用 submit_with_tools_stream() │ │
|
||||||
|
│ │ ├─ turn_index += 1 │ │
|
||||||
|
│ │ └─ 返回流 │ │
|
||||||
|
│ │ │ │
|
||||||
|
│ │ finalize_turn(response) │ │
|
||||||
|
│ │ ├─ cost_so_far.add(&response.usage) │ │
|
||||||
|
│ │ └─ OnTurnEnd hook(turn_index - 1) │ │
|
||||||
|
│ └──────────────────────────────────────────────────────┘ │
|
||||||
|
│ submit_with_tools_stream(prompt, Arc<ToolRegistry>)
|
||||||
|
▼
|
||||||
|
┌──────────────────────────────────────────────────────────────┐
|
||||||
|
│ LlmCycle (tokio::spawn task — run_tool_loop 状态机) │
|
||||||
|
│ │
|
||||||
|
│ max_turns = max_tool_turns.unwrap_or(10) │
|
||||||
|
│ for round in 1..=max_turns { │
|
||||||
|
│ ① build_request(messages, tools) │
|
||||||
|
│ ② provider.chat_stream(request).await │
|
||||||
|
│ 匹配 Err → tx.send(Error{..}) + return(不 panic) │
|
||||||
|
│ ③ 消费 LLM 流,所有事件 → mpsc unbounded tx(全量转发) │
|
||||||
|
│ ④ partial.finalize() → MessageResponse │
|
||||||
|
│ ⑤ if has_tool_use: │
|
||||||
|
│ ├─ tx → ToolExecutionStarted { tool_name, id, args } │
|
||||||
|
│ ├─ registry.invoke_all(calls, timeout).await │
|
||||||
|
│ ├─ for result: tx → ToolExecutionCompleted { ... } │
|
||||||
|
│ ├─ push tool results → messages │
|
||||||
|
│ └─ continue(新一轮) │
|
||||||
|
│ else: break(最终轮,已发出 MessageComplete) │
|
||||||
|
│ } │
|
||||||
|
│ │
|
||||||
|
│ 产出事件序列(通过 mpsc::unbounded_channel): │
|
||||||
|
│ MessageStart → ... → ToolCallEnd → ToolExecutionStarted → │
|
||||||
|
│ ToolExecutionCompleted → MessageStart → TextDelta → ... → │
|
||||||
|
│ CostUpdate → MessageComplete │
|
||||||
|
└──────────────────────────────────────────────────────────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
> **`max_tool_turns` 语义**:与非流式 `submit_with_tools` 一致——`None` 退化为 `10`(`unwrap_or(10)`)。默认值 `Some(10)` 已提供安全上限;如需增大限制,手动设置为 `Some(N)`。⚠️ 生产环境建议始终设有限值防止无限循环。
|
||||||
|
|
||||||
|
### 3.2 事件序列约定
|
||||||
|
|
||||||
|
**纯文本流**(无 tool_use):
|
||||||
|
|
||||||
|
```
|
||||||
|
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd → CostUpdate → MessageComplete
|
||||||
|
```
|
||||||
|
|
||||||
|
**单轮工具调用**:
|
||||||
|
|
||||||
|
```
|
||||||
|
MessageStart → ContentBlockStart → TextDelta* → ContentBlockEnd
|
||||||
|
→ ContentBlockStart → ToolCallArgumentsDelta* → ToolCallEnd
|
||||||
|
→ CostUpdate → MessageComplete { stop_reason: ToolUse }
|
||||||
|
→ ToolExecutionStarted → [工具执行] → ToolExecutionCompleted
|
||||||
|
→ ContentBlockStart → TextDelta* → ContentBlockEnd
|
||||||
|
→ CostUpdate → MessageComplete { stop_reason: Stop }
|
||||||
|
```
|
||||||
|
|
||||||
|
**多轮工具调用**:
|
||||||
|
|
||||||
|
```
|
||||||
|
... → ToolExecutionCompleted(第 1 轮)
|
||||||
|
→ ToolCallArgumentsDelta* → ToolCallEnd
|
||||||
|
→ ToolExecutionStarted → ToolExecutionCompleted(第 2 轮)
|
||||||
|
→ ... → CostUpdate → MessageComplete(最终轮)
|
||||||
|
```
|
||||||
|
|
||||||
|
**工具不可恢复错误**:
|
||||||
|
|
||||||
|
```
|
||||||
|
... → ToolCallEnd → ToolExecutionStarted
|
||||||
|
→ Error { "tool 'search' 不可恢复错误: ..." } → MessageComplete
|
||||||
|
```
|
||||||
|
|
||||||
|
> **`MessageComplete.full_response` 内容范围**:每轮 LLM 调用独立产生一个 `MessageComplete`,其中 `full_response` 仅包含**该轮 LLM 的单个响应**(不累积前面工具轮次的结果)。中间轮(`stop_reason: ToolUse`)的 `full_response` 通常只包含 `ToolUse` block,无文本。最终轮(`stop_reason: Stop`)的 `full_response` 包含 LLM 的最终输出。消费者如需追踪完整对话历史,应自行累加所有轮次的 `Message`。
|
||||||
|
|
||||||
|
### 3.3 StreamEvent 新增变体
|
||||||
|
|
||||||
|
在 `src/llm/types/response_v2.rs` 的 `StreamEvent` 枚举中追加两个变体:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// 工具开始执行 —— 在 ToolCallEnd 之后、registry.invoke 之前发出。
|
||||||
|
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
|
||||||
|
ToolExecutionStarted {
|
||||||
|
tool_name: String,
|
||||||
|
tool_call_id: String,
|
||||||
|
/// 工具参数(JSON 字符串形式)
|
||||||
|
arguments: String,
|
||||||
|
},
|
||||||
|
|
||||||
|
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
|
||||||
|
ToolExecutionCompleted {
|
||||||
|
tool_name: String,
|
||||||
|
tool_call_id: String,
|
||||||
|
/// 结果摘要(前 200 字符)
|
||||||
|
result_summary: String,
|
||||||
|
/// 是否出错
|
||||||
|
is_error: bool,
|
||||||
|
},
|
||||||
|
```
|
||||||
|
|
||||||
|
在 `PartialMessageResponse::apply_to` 中追加:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
|
||||||
|
```
|
||||||
|
|
||||||
|
这两个是**元事件**,不参与内容块累积,`apply_to` 直接返回 `true`。
|
||||||
|
|
||||||
|
### 3.4 新增方法签名
|
||||||
|
|
||||||
|
**`LlmCycle` 层**(`src/llm/cycle.rs`):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// 提交消息并自动处理工具调用循环,流式产出所有事件。
|
||||||
|
///
|
||||||
|
/// 与 `submit_with_tools` 的区别:
|
||||||
|
/// - LLM 响应是流式的(全程 `chat_stream` 而非 `chat`)
|
||||||
|
/// - 工具执行前后插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||||
|
/// - 错误以 `StreamEvent::Error` 形式出现在流中,而非终止 `Result`
|
||||||
|
/// - 消费方需手动 `push_message()` 同步消息历史
|
||||||
|
///
|
||||||
|
/// **运行时要求**:内部使用 `tokio::spawn`,需要 tokio 多线程运行时。
|
||||||
|
pub async fn submit_with_tools_stream(
|
||||||
|
&mut self,
|
||||||
|
prompt: String,
|
||||||
|
tool_registry: Arc<ToolRegistry>,
|
||||||
|
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, LlmError>
|
||||||
|
```
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// 运行工具循环的核心异步状态机。
|
||||||
|
///
|
||||||
|
/// 接收 owned 字段,通过 mpsc::unbounded_channel 产出事件序列。
|
||||||
|
/// 由 `submit_with_tools_stream` 在 tokio::spawn 中调用。
|
||||||
|
///
|
||||||
|
/// **运行时要求**:此函数内部使用 `tokio::spawn`,要求调用方运行在
|
||||||
|
/// tokio 多线程运行时中(`#[tokio::main]` 或 `#[tokio::test(flavor = "multi_thread")]`)。
|
||||||
|
/// 不在 WASM 目标下可用。
|
||||||
|
async fn run_tool_loop(
|
||||||
|
messages: Vec<Message>,
|
||||||
|
provider: Arc<dyn LlmProvider>,
|
||||||
|
config: CycleConfig,
|
||||||
|
tool_registry: Arc<ToolRegistry>,
|
||||||
|
tools: Vec<ToolDef>,
|
||||||
|
tx: mpsc::UnboundedSender<StreamEvent>,
|
||||||
|
hook_executor: Option<Arc<HookExecutor>>,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
**`AgentSession` 层**(`src/agent/session.rs`):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// 提交一轮对话(流式,含自动 tool 循环),返回 `StreamEvent` 流。
|
||||||
|
///
|
||||||
|
/// 与 `submit_turn` 的区别:
|
||||||
|
/// - 以流事件序列而非 `MessageResponse` 返回
|
||||||
|
/// - 工具执行期间插入 `ToolExecutionStarted` / `ToolExecutionCompleted` 事件
|
||||||
|
/// - 消费方在收到 `MessageComplete` 后需手动调用 `finalize_turn` 同步状态
|
||||||
|
///
|
||||||
|
/// **运行时要求**:内部委托 `submit_with_tools_stream`,需要 tokio 多线程运行时。
|
||||||
|
pub async fn submit_turn_stream(
|
||||||
|
&mut self,
|
||||||
|
user_input: impl Into<String>,
|
||||||
|
) -> Result<Pin<Box<dyn Stream<Item = StreamEvent> + Send>>, AgentError>
|
||||||
|
|
||||||
|
/// 完成一轮 turn:累计 cost + 触发 OnTurnEnd hook。
|
||||||
|
///
|
||||||
|
/// 由消费者在收到 `MessageComplete.full_response` 后调用。
|
||||||
|
pub async fn finalize_turn(&mut self, response: &MessageResponse)
|
||||||
|
```
|
||||||
|
|
||||||
|
> **实施偏差(Phase 10 适配)**:实际签名扩展为
|
||||||
|
> `pub async fn finalize_turn(&mut self, response: &MessageResponse, new_messages_from_cycle: Vec<Message>) -> Result<(), AgentError>`。
|
||||||
|
> - `new_messages_from_cycle`:本轮新增消息(`[user_input, ...tool_results, final_response]`),由消费者在流消费完毕后从 `cycle.messages()[input_len..]` 提取并传入;`finalize_turn` 增量追加到当前 slot(不覆盖已有消息)。
|
||||||
|
> - 返回 `Result<(), AgentError>`:错误传播更清晰,与 `submit_turn` 的 slot 边界错误(`SlotReadonly` / `SlotNotFound`)对齐。
|
||||||
|
> - Phase 10 ContextSlot 实施时扩展。Phase 9 消费者若不接入 slot 持久化,可传 `vec![response.message.clone()]` 兜底。
|
||||||
|
|
||||||
|
### 3.5 消费者使用模式
|
||||||
|
|
||||||
|
```rust
|
||||||
|
use futures_util::StreamExt;
|
||||||
|
|
||||||
|
let mut stream = session.submit_turn_stream("计算 1+2").await?;
|
||||||
|
|
||||||
|
let mut final_response = None;
|
||||||
|
while let Some(event) = stream.next().await {
|
||||||
|
match &event {
|
||||||
|
StreamEvent::TextDelta { text } => print!("{}", text),
|
||||||
|
StreamEvent::ToolExecutionStarted { tool_name, arguments, .. } => {
|
||||||
|
println!("\n🔧 [{}({})]", tool_name, arguments);
|
||||||
|
}
|
||||||
|
StreamEvent::ToolExecutionCompleted { result_summary, .. } => {
|
||||||
|
println!(" → {}", result_summary);
|
||||||
|
}
|
||||||
|
StreamEvent::MessageComplete { full_response } => {
|
||||||
|
final_response = Some(full_response.clone());
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
std::io::stdout().flush().ok();
|
||||||
|
|
||||||
|
if let Some(response) = final_response {
|
||||||
|
session.finalize_turn(&response).await;
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
> **⚠️ 消费者注意**:`finalize_turn` 是开发者责任 —— 遗漏调用会导致 `cost_so_far` 不累计、`OnTurnEnd` hook 不触发。session 状态仍然可用,后续 `submit_turn` 也能正常执行,但 cost 信息不完整。`finalize_turn` 无自动补偿机制,建议使用 `Drop` guard 或在 `while` 循环的 `finally` 块中确保调用。
|
||||||
|
|
||||||
|
### 3.6 run_tool_loop 核心逻辑
|
||||||
|
|
||||||
|
`run_tool_loop` 是此方案的核心状态机(约 90 行),其伪代码逻辑如下:
|
||||||
|
|
||||||
|
```
|
||||||
|
1. 接收 owned 字段:messages, provider, config, tool_registry, tools, tx, hook_executor
|
||||||
|
2. max_turns = config.max_tool_turns.unwrap_or(10)
|
||||||
|
// None → 10(退化为默认值),Some(n) → n
|
||||||
|
// 与非流式 submit_with_tools 行为一致
|
||||||
|
3. 工具循环(for round in 1..=max_turns):
|
||||||
|
a. build_request(messages, tools)
|
||||||
|
// 空 tool_registry 时 tools 为空列表,流退化为纯文本流(可安全运行)
|
||||||
|
b. PreRequest hook(如果有 hook_executor)
|
||||||
|
c. 发起流式 LLM 调用:
|
||||||
|
let stream = match provider.chat_stream(request).await {
|
||||||
|
Ok(s) => s,
|
||||||
|
Err(e) => {
|
||||||
|
// 第一层错误:chat_stream 自身失败(网络/认证/限流)
|
||||||
|
// 这里不做 retry:retry 逻辑留给上层循环的 submit_request 模式,
|
||||||
|
// 流式场景中 retry 需重新建立 mpsc 通道,复杂度与收益不匹配
|
||||||
|
tx.send(StreamEvent::Error { message: e.to_string() }).ok();
|
||||||
|
return; // 直接结束 task
|
||||||
|
}
|
||||||
|
};
|
||||||
|
d. 消费 LLM 流:
|
||||||
|
- PartialMessageResponse::new()
|
||||||
|
- while let Some(result) = stream.next().await
|
||||||
|
- match result:
|
||||||
|
Ok(event) → apply_to + tx.send(event)
|
||||||
|
Err(e) → tx.send(Error { message }) + break
|
||||||
|
// 第二层错误:stream 内部事件错误(如 chunk 解析失败)
|
||||||
|
e. partial.finalize()? → response
|
||||||
|
f. push response.message → messages
|
||||||
|
g. 检查 has_tool_calls_in_response(&response)
|
||||||
|
h. 如果没有 tool_use: break(最终轮,流已自然结束)
|
||||||
|
i. 如果有 tool_use:
|
||||||
|
- extract_tool_calls_from_response(&response)
|
||||||
|
- tx.send(ToolExecutionStarted { tool_name, tool_call_id, arguments })
|
||||||
|
- registry.invoke_all(calls, tool_timeout).await
|
||||||
|
- for result in results:
|
||||||
|
tx.send(ToolExecutionCompleted { tool_name, tool_call_id, result_summary, is_error })
|
||||||
|
- push tool results → messages
|
||||||
|
- continue(新一轮 LLM 流)
|
||||||
|
4. 流结束(tokio::spawn 自然退出)
|
||||||
|
```
|
||||||
|
|
||||||
|
> **关于 LLM 请求 retry**:非流式 `submit_with_tools` 内部通过 `submit_request` 的 retry 循环处理临时错误。流式版本 `run_tool_loop` **不在内部实现 retry**。原因:(1)retry 需要重新建立 mpsc 通道和事件流上下文,复杂度与收益不匹配;(2)`unbounded_channel` 已发出的事件无法撤回。如果需要 retry 语义,调用方应在上层做 fallback 策略,或在 `llm provider` 实现层完成 retry(如 `RetryProvider` 包装器)。
|
||||||
|
|
||||||
|
**错误处理**:
|
||||||
|
|
||||||
|
| 场景 | 行为 |
|
||||||
|
|------|------|
|
||||||
|
| LLM 请求失败(`chat_stream` 返回 `Err`) | `tx.send(Error { message })` + `return` 结束 task。**不做 retry**(见上方说明) |
|
||||||
|
| LLM 流内事件错误(stream Item 的 `Err`) | `tx.send(Error { message })` + `break` 结束当轮流,终止循环 |
|
||||||
|
| 可恢复工具错误(`is_recoverable() == true`) | 作为 tool result 回传 LLM,流继续,不出 Error 事件 |
|
||||||
|
| 不可恢复工具错误(`is_recoverable() == false`) | `tx.send(Error { message })` + 终止循环 |
|
||||||
|
| 工具超时(`tokio::time::timeout`) | 视为不可恢复,`tx.send(Error)` + 终止 |
|
||||||
|
| 最大工具循环轮次超限 | `tx.send(Error { "达到最大工具循环轮次" })` + 终止 |
|
||||||
|
| spawn task 内部 panic | 由于 `JoinHandle` 不保存(detached),panic 由 tokio 运行时静默捕获;消费者看到 stream 直接结束(返回 `None`),无 `Error` 事件。建议在 `run_tool_loop` 内部避免 `unwrap()`,所有可失败路径通过 `Result` + `?` 传播 |
|
||||||
|
|
||||||
|
### 3.7 修改文件清单
|
||||||
|
|
||||||
|
| # | 文件 | 改动 | 估算行数 |
|
||||||
|
|---|------|------|---------|
|
||||||
|
| 1 | `llm/types/response_v2.rs` | +2 `StreamEvent` 变体 +2 `apply_to` arm | ~20 |
|
||||||
|
| 2 | `llm/cycle.rs` | +`submit_with_tools_stream` 方法 + `run_tool_loop` 模块函数 | ~140 |
|
||||||
|
| 3 | `llm/cycle.rs` | `CycleConfig` 加 `#[derive(Clone)]` | ~1 |
|
||||||
|
| 4 | `agent/session.rs` | +`submit_turn_stream` + `finalize_turn` | ~70 |
|
||||||
|
| — | **测试**(内联) | 4 个场景测试(纯度本、单轮、多轮、超限) | ~150 |
|
||||||
|
| | **合计** | | **~380** |
|
||||||
|
|
||||||
|
> 注:`RetryConfig` 已标注 `#[derive(Debug, Clone)]`,无需额外修改。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. 实现计划
|
||||||
|
|
||||||
|
按 5 个 Step 增量实施,每步可独立编译和测试。
|
||||||
|
|
||||||
|
### Step 1 — 基础设施准备
|
||||||
|
|
||||||
|
**目标**:数据层就绪,为流事件新增变体和配置 Clone 奠基。
|
||||||
|
|
||||||
|
**改动**:
|
||||||
|
|
||||||
|
- `llm/types/response_v2.rs`:
|
||||||
|
- `StreamEvent` 枚举追加 `ToolExecutionStarted` / `ToolExecutionCompleted` 变体
|
||||||
|
- `PartialMessageResponse::apply_to` 追加两个新变体的 arm(均返回 `true`)
|
||||||
|
- `llm/cycle.rs`:
|
||||||
|
- `CycleConfig` 加 `#[derive(Clone)]`(所有字段为基础类型 + `RetryConfig`)
|
||||||
|
|
||||||
|
**验证**:`cargo build` 通过
|
||||||
|
|
||||||
|
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
|
||||||
|
|
||||||
|
**目标**:实现流式工具循环的核心状态机,这是整个 Phase 9 的技术关键。
|
||||||
|
|
||||||
|
**改动**:
|
||||||
|
|
||||||
|
- `llm/cycle.rs`:
|
||||||
|
- 新增 `run_tool_loop()` 模块函数(约 90 行),基于 `mpsc::unbounded_channel` 通信
|
||||||
|
- 新增 `submit_with_tools_stream()` 公开方法,入口参数为 `prompt` + `Arc<ToolRegistry>`
|
||||||
|
- 内部 `tokio::spawn` 启动 `run_tool_loop`,返回 `rx` 端作为 `dyn Stream`
|
||||||
|
|
||||||
|
**验证**:`cargo build` 通过
|
||||||
|
|
||||||
|
### Step 3 — 单元测试(LlmCycle 层)
|
||||||
|
|
||||||
|
**目标**:验证 `submit_with_tools_stream` 在 8 个核心场景下的行为和事件序列正确性(含 §8 Step 3 扩展的工具错误路径)。
|
||||||
|
|
||||||
|
**新增**(`llm/cycle.rs` 内联测试 `#[cfg(test)]`):
|
||||||
|
|
||||||
|
| 场景 | Mock 响应序列 | 验证点 |
|
||||||
|
|------|---------------|--------|
|
||||||
|
| 1 — 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` |
|
||||||
|
| 2 — 单轮工具调用 | 2 个响应:tool_use → text | 包含 `ToolExecutionStarted` + `ToolExecutionCompleted`;最终 `stop_reason` 为 `Stop` |
|
||||||
|
| 3 — 多轮工具调用 | 4 个响应:3 × tool_use → 1 × text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度正确 |
|
||||||
|
| 4 — 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `Error` 事件;消息历史停在第 2 轮 |
|
||||||
|
|
||||||
|
**验证**:`cargo test` 全部通过
|
||||||
|
|
||||||
|
### Step 4 — `AgentSession` 层包装
|
||||||
|
|
||||||
|
**目标**:为 `AgentSession` 新增流式会话接口,保持与 `submit_turn` 一致的行为语义。
|
||||||
|
|
||||||
|
**改动**:
|
||||||
|
|
||||||
|
- `agent/session.rs`:
|
||||||
|
- `submit_turn_stream(user_input)` — 触发 `OnTurnStart` hook → 组装 `LlmCycle` → 调用 `submit_with_tools_stream` → `turn_index += 1` → 返回流
|
||||||
|
- `finalize_turn(response)` — `cost_so_far.add(&response.usage)` → 触发 `OnTurnEnd` hook
|
||||||
|
|
||||||
|
**验证**:`cargo build` 通过
|
||||||
|
|
||||||
|
### Step 5 — 集成测试 + 扫尾
|
||||||
|
|
||||||
|
**目标**:端到端验证 `submit_turn_stream` + `finalize_turn` 的完整链路,确保零回归。
|
||||||
|
|
||||||
|
**新增**(`agent/session.rs` 内联测试,2026-07-08 实施审查补全):
|
||||||
|
|
||||||
|
- `submit_turn_stream_end_to_end` — `submit_turn_stream` 跑通 mock provider → 消费流(验证收到 TextDelta + MessageComplete) → `finalize_turn` 后 `cost_so_far` 正确更新(`prompt_tokens=10, completion_tokens=5`) + `turn_index=1` + default slot 包含 user/assistant 消息
|
||||||
|
- `submit_turn_stream_triggers_turn_hooks` — 验证 `OnTurnStart` 在 `submit_turn_stream` 返回流之前已触发(计数=1)+ `OnTurnEnd` 在 `finalize_turn` 之前**不**触发(计数=0)+ `finalize_turn` 后 `OnTurnEnd` 触发(计数=1)
|
||||||
|
|
||||||
|
**验证**:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cargo test --all-targets # 全绿,存量测试 0 回归
|
||||||
|
cargo clippy --all-targets -- -D warnings # 0 警告
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. 运行细节
|
||||||
|
|
||||||
|
### 5.1 `run_tool_loop` 的 spawn 生命周期
|
||||||
|
|
||||||
|
#### 执行模型:立即执行 vs 惰性流
|
||||||
|
|
||||||
|
`submit_with_tools_stream` 采用 **立即执行** 模型(`tokio::spawn` + `mpsc`),这与 `submit_stream` 的 **惰性执行**(`async_stream::stream!` 宏,消费者首次 `next()` 时才触发 LLM 调用)不同。
|
||||||
|
|
||||||
|
**选择理由**:工具循环是 **不确定轮次的** —— 每个工具执行的结果可能影响后续 LLM 调用。惰性流无法表达这种"边消费边控制"的语义。通过 `tokio::spawn` 将工具循环移到独立 task 中运行,使得:
|
||||||
|
- 消费者可以随时开始消费(不丢失事件)
|
||||||
|
- 工具循环在后台独立运行,不受消费者消费节奏影响
|
||||||
|
- `mpsc::unbounded_channel` 作为事件缓冲区,解耦生产者与消费者
|
||||||
|
|
||||||
|
**对消费者的影响**:`submit_with_tools_stream().await?` 返回时,工具循环可能已经开始执行(事件已开始写入 channel)。消费者应尽快开始 `while let Some(event) = stream.next().await`,避免 channel 缓冲过多事件。如果在返回流后长时间不消费,事件会堆积在 mpsc buffer 中(内存开销,无阻塞风险 —— 见 §6 风险表)。
|
||||||
|
|
||||||
|
#### 生命周期
|
||||||
|
|
||||||
|
```
|
||||||
|
submit_with_tools_stream()
|
||||||
|
│
|
||||||
|
├─ mpsc::unbounded_channel() → (tx, rx)
|
||||||
|
├─ messages.push(user_text(prompt))
|
||||||
|
├─ compact check
|
||||||
|
├─ tokio::spawn(run_tool_loop(messages, provider, config, ..., tx))
|
||||||
|
└─ return Box::pin(rx) as dyn Stream
|
||||||
|
|
||||||
|
[用户消费 stream]
|
||||||
|
└─ while let Some(event) = rx.recv().await { yield event }
|
||||||
|
|
||||||
|
[用户 drop rx / 结束循环]
|
||||||
|
└─ rx 被 drop → tx.send() 返回 Err
|
||||||
|
→ run_tool_loop 检测到 tx.closed()
|
||||||
|
→ break → task 自然终止
|
||||||
|
```
|
||||||
|
|
||||||
|
#### JoinHandle 与 panic 处理
|
||||||
|
|
||||||
|
`run_tool_loop` 的 `JoinHandle` 在 spawn 后**不保存**(detached pattern)。panic 由 tokio 运行时捕获并通过 `tracing::error` 记录:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// submit_with_tools_stream 内部
|
||||||
|
tokio::spawn(async move {
|
||||||
|
run_tool_loop(..., tx).await;
|
||||||
|
});
|
||||||
|
```
|
||||||
|
|
||||||
|
如果 `run_tool_loop` 内部发生 panic(如 `unwrap()`),tokio 的 `spawn` 会静默吞掉 panic 并终止 task。消费者此时看到 stream 直接返回 `None`,不会收到 `StreamEvent::Error`。实际编码中应避免 `unwrap()`,所有 `Result` 使用 `?` 或 `match` 处理。
|
||||||
|
|
||||||
|
Rx 侧实现 `Stream` trait:使用 `tokio_stream::wrappers::UnboundedReceiverStream` 包装 `mpsc::UnboundedReceiver`,因为 `mpsc::UnboundedReceiver` 本身不实现 `Stream`(`tokio-stream = "0.1"` 已在 `Cargo.toml` 中存在)。
|
||||||
|
|
||||||
|
### 5.2 消息历史同步
|
||||||
|
|
||||||
|
`submit_with_tools_stream` 内部由 `run_tool_loop` 管理 `messages` 的拷贝,不会写入 `self.messages`。消费方在收到 `MessageComplete` 后需手动:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
let response = full_response.clone();
|
||||||
|
cycle.push_message(response.message.clone());
|
||||||
|
```
|
||||||
|
|
||||||
|
在 `AgentSession::submit_turn_stream` 中,由于流是延迟求值且 `&mut self` 无法进入 spawn 闭包,消息历史同步交由消费方在 `finalize_turn` 前自行决定。当前方案中 `submit_turn_stream` **不自动同步消息历史**,这与 `submit_stream` 的已有行为一致(ponytail: Phase 2 FIX-E 注释)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. 风险评估
|
||||||
|
|
||||||
|
| 风险 | 影响 | 缓解措施 |
|
||||||
|
|------|------|---------|
|
||||||
|
| `&mut self` 约束导致流内无法访问 session 状态 | 中 | 复用 `submit_stream` 已有模式:方法体内读取 `self` 后构建 owned 数据,spawn 闭包不捕获 `&mut self` |
|
||||||
|
| spawn task 生命周期管理 | 低 | 用户 drop rx → `tx.send` 返回 `Err` → `run_tool_loop` 自然终止 |
|
||||||
|
| spawn task panic 静默丢失 | 中 | `run_tool_loop` 内部使用 `match`/`?` 避免 `unwrap()`;`JoinHandle` 不做 `await`(detached),panic 由 tokio 运行时记录日志。消费者看到 stream 提前结束(收到 `None`)但无 Error 事件 |
|
||||||
|
| 中间轮 cost 不累加到 `cost_so_far` | 低 | 与现有 `submit_turn` 行为一致(仅最终轮计入),标记为已知限制,不在此 Phase 修复 |
|
||||||
|
| 工具循环中 hook 可用性 | 低 | `PreRequest`/`PostRequest` hook 通过 `hook_executor.clone()` 进入 spawn task;hook 在 `run_tool_loop` 循环内触发 |
|
||||||
|
| `run_tool_loop` 不支持消息压缩 | 中 | 长工具循环中消息不断增长,可能超出 context window。当前不传递 `compact_config`,后续可扩展 `run_tool_loop` 签名增添此参数 |
|
||||||
|
| `unbounded_channel` 在消费慢于生产时内存增长 | 低 | LLM 流式输出天然有节流(token 生成速度远慢于 CPU 处理速度),消费者通常快于生产者。后续如需背压可切换为 `mpsc::channel(N)` + backpressure |
|
||||||
|
| 流式版本不做 LLM retry | 低 | 非流式 `submit_with_tools` 通过 `submit_request` 的 retry 循环处理临时错误。流式版本中 retry 需重建 mpsc 通道,复杂度不匹配。调用方可使用 `RetryProvider` 包装器或在 Provider 层实现 retry |
|
||||||
|
| 执行模式与 `submit_stream` 不一致(立即 vs 惰性) | 低 | `submit_stream` 的惰性语义不适配需要后台执行的工具循环。消费者应在 `submit_turn_stream` 返回后尽快消费流事件 |
|
||||||
|
| `tokio::spawn` 要求 tokio 多线程运行时 | 低 | agcore 已依赖 tokio,涉及 IO 的 API 均使用 async。`#[tokio::test]` 单线程运行时不支持 `spawn`,测试中将 `run_tool_loop` 提取为可独立调用的函数,测试不走 spawn 直接调用 |
|
||||||
|
| `CycleConfig` 加 `Clone` 影响现有代码 | 无 | 纯配置 struct,所有字段是基础类型或已 `Clone` 的 `RetryConfig` |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. 验收标准
|
||||||
|
|
||||||
|
| # | 验收项 | 验证方式 |
|
||||||
|
|---|--------|---------|
|
||||||
|
| 1 | `cargo build --all-targets` 通过 | ✅ 编译器无错误 |
|
||||||
|
| 2 | `submit_with_tools_stream` 纯文本流事件序列正确 | 单元测试验证:事件类型、顺序与 `submit_stream` 一致 |
|
||||||
|
| 3 | `submit_with_tools_stream` 单轮工具调用事件序列正确 | 单元测试验证:含 `ToolExecutionStarted` / `ToolExecutionCompleted` |
|
||||||
|
| 4 | `submit_with_tools_stream` 多轮工具调用事件序列正确 | 单元测试验证:多对 `ToolExecutionStarted`/`ToolExecutionCompleted` |
|
||||||
|
| 5 | `submit_with_tools_stream` 最大轮次超限产生 Error 事件 | 单元测试验证:流中出现 `StreamEvent::Error` |
|
||||||
|
| 6 | `submit_turn_stream` + `finalize_turn` 端到端链路 | 集成测试验证:cost 更新 + hook 触发 |
|
||||||
|
| 7 | `cargo test --all-targets` 全绿,存量测试 0 回归 | ✅ 无回归 |
|
||||||
|
| 8 | `cargo clippy --all-targets -- -D warnings` 0 警告 | ✅ 无警告 |
|
||||||
|
| 9 | 现有 `submit_turn` / `submit_with_tools` / `submit_stream` 行为零影响 | ✅ 存量测试通过 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 8. 实施计划
|
||||||
|
|
||||||
|
按 5 个 Step 分阶段实施,每步产出独立 commit,可验证后退。
|
||||||
|
|
||||||
|
### 依赖关系
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph LR
|
||||||
|
S1["Step 1: 基础设施"]:::s1
|
||||||
|
S2["Step 2: 核心状态机"]:::s2
|
||||||
|
S3["Step 3: LlmCycle 单元测试"]:::s3
|
||||||
|
S4["Step 4: AgentSession 包装"]:::s4
|
||||||
|
S5["Step 5: 集成测试 + 扫尾"]:::s5
|
||||||
|
|
||||||
|
S1 --> S2
|
||||||
|
S1 --> S4
|
||||||
|
S2 --> S3
|
||||||
|
S2 --> S4
|
||||||
|
S3 --> S5
|
||||||
|
S4 --> S5
|
||||||
|
|
||||||
|
classDef s1 fill:#e2e8f0,stroke:#94a3b8
|
||||||
|
classDef s2 fill:#fbbf24,stroke:#d97706
|
||||||
|
classDef s3 fill:#93c5fd,stroke:#2563eb
|
||||||
|
classDef s4 fill:#93c5fd,stroke:#2563eb
|
||||||
|
classDef s5 fill:#4ade80,stroke:#16a34a
|
||||||
|
```
|
||||||
|
|
||||||
|
| Step | 依赖 | 并行机会 |
|
||||||
|
|------|------|---------|
|
||||||
|
| S1 | 无 | — |
|
||||||
|
| S2 | S1 | 可与 S4 并行 |
|
||||||
|
| S3 | S2 | 阻塞,需 S2 完成 |
|
||||||
|
| S4 | S1, S2 | 功能依赖 S2(调用 `submit_with_tools_stream`);文件级无重叠但需先编译过 S2 |
|
||||||
|
| S5 | S3 + S4 | 需 S3 和 S4 都完成 |
|
||||||
|
|
||||||
|
### Step 1 — 基础设施准备
|
||||||
|
|
||||||
|
**工作量**:S(< 1h)
|
||||||
|
**风险**:低(纯新增,不影响现有代码逻辑)
|
||||||
|
|
||||||
|
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||||
|
|---|------|---------|---------|------|
|
||||||
|
| 1.1 | `StreamEvent` 枚举追加 `ToolExecutionStarted` 变体 | `llm/types/response_v2.rs` | 无 | 低 |
|
||||||
|
| 1.2 | `StreamEvent` 枚举追加 `ToolExecutionCompleted` 变体 | `llm/types/response_v2.rs` | 1.1 | 低 |
|
||||||
|
| 1.3 | `PartialMessageResponse::apply_to` 追加两个元事件 arm(均返回 `true`) | `llm/types/response_v2.rs` | 1.2 | 低 |
|
||||||
|
| 1.4 | `CycleConfig` 加 `#[derive(Clone)]` | `llm/cycle.rs` | 无 | 低 |
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `cargo build` 通过,编译器无 warning
|
||||||
|
- 新增的 `StreamEvent` 变体可通过 `serde` roundtrip 序列化/反序列化
|
||||||
|
- `CycleConfig` 可正常 clone
|
||||||
|
|
||||||
|
### Step 2 — `LlmCycle::submit_with_tools_stream` 核心
|
||||||
|
|
||||||
|
**工作量**:M(1-4h)
|
||||||
|
**风险**:中(核心实现,需正确设计 spawn + mpsc 生命周期)
|
||||||
|
**前置依赖**:S1
|
||||||
|
|
||||||
|
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||||
|
|---|------|---------|---------|------|
|
||||||
|
| 2.1 | 实现 `run_tool_loop()` 模块函数:消息循环构建请求 → `chat_stream` → 消费流 → 检测 tool_use → 工具执行 → 新一轮 | `llm/cycle.rs` | S1 | 中 |
|
||||||
|
| 2.2 | 实现 `submit_with_tools_stream()` 公开方法:提取字段 → spawn `run_tool_loop` → 返回 `UnboundedReceiverStream` | `llm/cycle.rs` | 2.1 | 中 |
|
||||||
|
| 2.3 | 新增导入:`tokio::sync::mpsc`、`tokio_stream::wrappers::UnboundedReceiverStream` | `llm/cycle.rs` | 2.2 | 低 |
|
||||||
|
|
||||||
|
**关键实现细节**:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// run_tool_loop 函数签名
|
||||||
|
async fn run_tool_loop(
|
||||||
|
mut messages: Vec<Message>,
|
||||||
|
provider: Arc<dyn LlmProvider>,
|
||||||
|
config: CycleConfig,
|
||||||
|
tool_registry: Arc<ToolRegistry>,
|
||||||
|
tools: Vec<ToolDef>,
|
||||||
|
tx: mpsc::UnboundedSender<StreamEvent>,
|
||||||
|
hook_executor: Option<Arc<HookExecutor>>,
|
||||||
|
) {
|
||||||
|
let max_turns = config.max_tool_turns.unwrap_or(10);
|
||||||
|
let tool_timeout = config.tool_timeout_secs;
|
||||||
|
let max_bytes = config.max_tool_result_bytes;
|
||||||
|
|
||||||
|
let mut round = 0u32;
|
||||||
|
loop {
|
||||||
|
round += 1;
|
||||||
|
if round > max_turns {
|
||||||
|
// §3.6 错误表:最大轮次超限 → Error 事件 + 终止
|
||||||
|
let _ = tx.send(StreamEvent::Error { message: "达到最大工具循环轮次".to_string() });
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
// ① 构建请求
|
||||||
|
let request = MessageRequest {
|
||||||
|
model: config.model.clone(),
|
||||||
|
messages: messages.clone(),
|
||||||
|
tools: tools.clone(),
|
||||||
|
tool_choice: ToolChoice::Auto,
|
||||||
|
max_tokens: config.max_tokens,
|
||||||
|
temperature: config.temperature,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
|
||||||
|
// ② PreRequest hook
|
||||||
|
// ...
|
||||||
|
|
||||||
|
// ③ chat_stream
|
||||||
|
let stream = match provider.chat_stream(request).await {
|
||||||
|
Ok(s) => s,
|
||||||
|
Err(e) => {
|
||||||
|
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// ④ 消费流
|
||||||
|
let mut partial = PartialMessageResponse::new();
|
||||||
|
let mut stream = stream;
|
||||||
|
while let Some(result) = stream.next().await {
|
||||||
|
match result {
|
||||||
|
Ok(event) => {
|
||||||
|
partial.apply_to(&event);
|
||||||
|
if tx.send(event).is_err() { return; }
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
// ponytail: 流内事件错误后 partial 处于损坏状态,
|
||||||
|
// 不能继续执行 finalize/finalize —— 直接 return 结束 task
|
||||||
|
let _ = tx.send(StreamEvent::Error { message: e.to_string() });
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ⑤ finalize
|
||||||
|
let response = match partial.finalize() {
|
||||||
|
Ok(r) => r,
|
||||||
|
Err(e) => { let _ = tx.send(StreamEvent::Error { .. }); return; }
|
||||||
|
};
|
||||||
|
messages.push(response.message.clone());
|
||||||
|
|
||||||
|
// ⑥ 检测 tool_use
|
||||||
|
if !has_tool_calls_in_response(&response) {
|
||||||
|
break; // 最终轮
|
||||||
|
}
|
||||||
|
|
||||||
|
// ⑦ 执行工具
|
||||||
|
let tool_calls = extract_tool_calls_from_response(&response);
|
||||||
|
let calls: Vec<_> = tool_calls.into_iter()
|
||||||
|
.map(|(id, name, args)| {
|
||||||
|
let value = serde_json::from_str(&args).unwrap_or(Value::Null);
|
||||||
|
(id, name, value)
|
||||||
|
}).collect();
|
||||||
|
|
||||||
|
for (tool_call_id, tool_name, args_value) in &calls {
|
||||||
|
let args_json = serde_json::to_string(&args_value).unwrap_or_default();
|
||||||
|
if tx.send(StreamEvent::ToolExecutionStarted {
|
||||||
|
tool_name: tool_name.clone(),
|
||||||
|
tool_call_id: tool_call_id.clone(),
|
||||||
|
arguments: args_json,
|
||||||
|
}).is_err() { return; }
|
||||||
|
}
|
||||||
|
|
||||||
|
let results = tool_registry.invoke_all(calls, tool_timeout).await;
|
||||||
|
|
||||||
|
for result in &results {
|
||||||
|
let summary = match &result.output {
|
||||||
|
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
|
||||||
|
Err(e) => e.to_string(),
|
||||||
|
};
|
||||||
|
// ponytail: 复用现有 truncate_tool_result 函数(cycle.rs 末尾),
|
||||||
|
// 确保多字节 UTF-8 字符不被截断破坏。上限 200 字符。
|
||||||
|
let truncated = truncate_tool_result(&summary, 200);
|
||||||
|
if tx.send(StreamEvent::ToolExecutionCompleted {
|
||||||
|
tool_name: result.tool_name.clone(),
|
||||||
|
tool_call_id: result.tool_call_id.clone(),
|
||||||
|
result_summary: truncated,
|
||||||
|
is_error: result.output.is_err(),
|
||||||
|
}).is_err() { return; }
|
||||||
|
}
|
||||||
|
|
||||||
|
for result in results {
|
||||||
|
let is_error = result.output.is_err();
|
||||||
|
let content = match &result.output {
|
||||||
|
Ok(v) => serde_json::to_string(v).unwrap_or_default(),
|
||||||
|
Err(e) if e.is_recoverable() => format!("错误: {}", e),
|
||||||
|
Err(e) => {
|
||||||
|
let _ = tx.send(StreamEvent::Error { .. });
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
messages.push(Message::tool_result(result.tool_call_id, content, is_error));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `cargo build` 通过
|
||||||
|
- 新增方法签名与方案设计一致
|
||||||
|
- 未修改现有 `submit_with_tools`/`submit_stream` 的行为
|
||||||
|
|
||||||
|
### Step 3 — 单元测试(LlmCycle 层)
|
||||||
|
|
||||||
|
**工作量**:M(1-4h)
|
||||||
|
**风险**:低(与现有测试模式一致,使用已有 MockProvider)
|
||||||
|
**前置依赖**:S2
|
||||||
|
|
||||||
|
测试策略:直接使用公开的 `crate::llm::mock::MockProvider`(已完整实现 `chat_stream` + 预设响应队列),避免改造 `cycle.rs` 测试模块内的内联 Stub。测试中调用 `submit_with_tools_stream` 时通过 `#[tokio::test(flavor = "multi_thread")]` 满足 spawn 运行时要求,或在单元级将 `run_tool_loop` 作为独立函数直接测试(不走 spawn)。
|
||||||
|
|
||||||
|
| # | 测试场景 | Mock 响应序列 | 验证点 | 覆盖路径 |
|
||||||
|
|---|---------|---------------|--------|---------|
|
||||||
|
| 3.1 | 纯文本流 | 1 个 text 响应 | 事件序列与 `submit_stream` 一致;无 `ToolExecutionStarted`/`ToolExecutionCompleted` | 正常路径:单轮 LLM → 文本返回 |
|
||||||
|
| 3.2 | 单轮工具调用 | 2 个响应:tool_use + text | 包含一对 `ToolExecutionStarted`/`ToolExecutionCompleted`;最终 `stop_reason` 为 `Stop` | 正常路径:LLM → 工具 → LLM |
|
||||||
|
| 3.3 | 多轮工具调用 | 4 个响应:3×tool_use + 1×text | 3 对 `ToolExecutionStarted`/`ToolExecutionCompleted`;消息历史长度为 8(user + 3×(assistant+tool) + final assistant) | 正常路径:LLM → 工具 → LLM → 工具 → LLM |
|
||||||
|
| 3.4 | 最大轮次超限 | 3 个 tool_use 响应,`max_tool_turns: Some(2)` | 流中出现 `StreamEvent::Error`;消息历史停在第 2 轮 | 边界条件:超出上限 |
|
||||||
|
| 3.5 | `chat_stream` 返回 Err | Mock `chat_stream` 返回 `Err(LlmError::Other(...))` | 流中第一个事件为 `StreamEvent::Error`;随后流结束 | 异常路径:LLM 不可用 |
|
||||||
|
| 3.6 | 空 tool_registry | 1 个 text 响应,registry 中无工具 | 流退化为纯文本流,事件序列与 3.1 一致 | 退化场景:无工具可用 |
|
||||||
|
| 3.7 | 不可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(不可恢复) | 流中出现 `StreamEvent::Error`;消息历史中不含该工具结果(循环终止前未 push) | 异常路径:工具执行失败 |
|
||||||
|
| 3.8 | 可恢复工具错误 | 2 个响应:tool_use → text,工具返回 `ToolError::ExecutionFailed`(可恢复) | 工具结果作为 `ToolResult { is_error: true }` 回传 LLM;流正常结束,无 `Error` 事件 | 异常路径:工具出错但可恢复 |
|
||||||
|
| 3.9 | 工具超时 | 2 个响应:tool_use → text,`tool_timeout_secs: 1`,模拟工具耗时 10 秒 | 流中出现 `StreamEvent::Error`;循环终止前未 push 工具结果 | 异常路径:工具执行超时 |
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `cargo test` 新增 8 个测试全部通过
|
||||||
|
- `cargo test` 存量测试 0 回归
|
||||||
|
|
||||||
|
### Step 4 — `AgentSession` 层包装
|
||||||
|
|
||||||
|
**工作量**:S(< 1h)
|
||||||
|
**风险**:低(薄包装层,逻辑简单)
|
||||||
|
**前置依赖**:S1
|
||||||
|
|
||||||
|
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||||
|
|---|------|---------|---------|------|
|
||||||
|
| 4.1 | 实现 `submit_turn_stream()`:触发 `OnTurnStart` → 组装 `LlmCycle` → 调用 `submit_with_tools_stream` → `turn_index += 1` → 返回流 | `agent/session.rs` | S1 | 低 |
|
||||||
|
| 4.2 | 实现 `finalize_turn()`:`cost_so_far.add()` → 触发 `OnTurnEnd` hook | `agent/session.rs` | S1 | 低 |
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `cargo build` 通过
|
||||||
|
- 新增方法签名与方案设计一致
|
||||||
|
- 与 `submit_turn` 的 system_prompt / compact_config / bundle 使用方式一致
|
||||||
|
|
||||||
|
### Step 5 — 集成测试 + 扫尾
|
||||||
|
|
||||||
|
**工作量**:S(< 1h)
|
||||||
|
**风险**:低(基于现有测试框架)
|
||||||
|
**前置依赖**:S3 + S4
|
||||||
|
|
||||||
|
| # | 任务 | 涉及文件 | 前置依赖 | 风险 |
|
||||||
|
|---|------|---------|---------|------|
|
||||||
|
| 5.1 | `submit_turn_stream` 端到端测试:跑通 mock provider → 消费流验证各事件到达 → `finalize_turn` 后 cost 更新正确 | `agent/session.rs`(内联测试) | S4 | 低 |
|
||||||
|
| 5.2 | Hook 触发验证:`OnTurnStart` 在 `submit_turn_stream` 返回流之前触发;`finalize_turn` 调用后 `OnTurnEnd` 正确触发 | `agent/session.rs`(内联测试) | S4 | 低 |
|
||||||
|
| 5.3 | `cargo test --all-targets` 全绿验证 | 全仓 | S5.1+S5.2 | 低 |
|
||||||
|
| 5.4 | `cargo clippy --all-targets -- -D warnings` 0 警告 | 全仓 | S5.3 | 低 |
|
||||||
|
| 5.5 | `cargo build --all-targets` 发布模式验证 | 全仓 | S5.4 | 低 |
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- 全量测试通过,存量 0 回归
|
||||||
|
- clippy 0 警告
|
||||||
|
- 发布模式零 warning
|
||||||
|
|
||||||
|
### 实施总览
|
||||||
|
|
||||||
|
| | Step 1 | Step 2 | Step 3 | Step 4 | Step 5 | **合计** |
|
||||||
|
|--|--------|--------|--------|--------|--------|---------|
|
||||||
|
| **工作量** | S | M | M | S | S | **M-L** |
|
||||||
|
| **文件数** | 2 | 1 | 1(内联) | 1 | 1(内联) | **~4** |
|
||||||
|
| **代码行** | ~20 | ~140 | ~150 含测试 | ~70 | ~80 含测试 | **~380** |
|
||||||
|
| **风险** | 低 | 中 | 低 | 低 | 低 | 中 |
|
||||||
|
| **并行** | — | 阻塞(S4 依赖 S2) | 阻塞 | 阻塞(依赖 S2) | 阻塞 | — |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 附录 A:新增 StreamEvent 变体的 apply_to 语义
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// 在 PartialMessageResponse::apply_to 中追加:
|
||||||
|
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => {
|
||||||
|
// 元事件:不参与内容块累积,不修改 partial response 状态
|
||||||
|
true
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## 附录 B:CycleConfig 的 Clone 推导
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// LLM 调用周期配置。
|
||||||
|
#[derive(Debug, Clone)] // ← 追加 Clone
|
||||||
|
pub struct CycleConfig {
|
||||||
|
pub model: String,
|
||||||
|
pub max_tokens: Option<u32>,
|
||||||
|
pub temperature: Option<f32>,
|
||||||
|
pub max_turns: Option<u32>,
|
||||||
|
pub retry: RetryConfig, // 已 #[derive(Clone)]
|
||||||
|
pub max_tool_turns: Option<u32>,
|
||||||
|
pub tool_timeout_secs: u64,
|
||||||
|
pub max_tool_result_bytes: usize,
|
||||||
|
}
|
||||||
|
```
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,647 @@
|
|||||||
|
# Phase 11: 测试与检索补强
|
||||||
|
|
||||||
|
## 背景与目标
|
||||||
|
|
||||||
|
AG Core 当前(v0.2.0-rc.1)已完成 Phase 0-10,全量测试 254 个,clippy 0 警告,11 个离线示例全部 exit 0。功能性交付物覆盖了 LLM Cycle、Prompt、Tool、Memory、Agent Runtime、流式事件、ContextSlot 上下文管理。但有两个系统性的短板尚未补齐:
|
||||||
|
|
||||||
|
1. **检索抽象缺失**:`memory` 模块只有 `MemoryRetriever`(基于 TextOverlap Dice 系数的关键词检索),缺少语义向量的检索抽象。`docs/roadmap.md` P1 中「VectorRetriever trait」一直未实现。
|
||||||
|
2. **测试覆盖缺口**:Provider 的 roundtrip 测试停留在"基本响应 + 普通 401/500"层面,缺少结构化错误体解析、请求头验证、流式边界、工具调用端到端等关键场景的回归覆盖。多线程并发的 MemoryStore 测试只在 SqliteStore 有一个 10×10 场景,InMemoryStore 完全没有并发压力测试。
|
||||||
|
|
||||||
|
Phase 11 是 v0.2.0 正式版发布前的最后一个功能 Phase,三个 Step 的目标:
|
||||||
|
|
||||||
|
| Step | 内容 | 定位 |
|
||||||
|
|------|------|------|
|
||||||
|
| **11.1** | `VectorRetriever` trait + `InMemoryVectorRetriever` 引用实现 | P1 功能补全 |
|
||||||
|
| **11.2** | wiremock Provider roundtrip 测试(12 个场景) | 测试质量补强 |
|
||||||
|
| **11.3** | 并发测试补强(InMemoryStore + SqliteStore) | 并发安全验证 |
|
||||||
|
|
||||||
|
最终目标:全量测试从 254 → 275+,为 v0.2.0 正式版建立更高的质量基线。
|
||||||
|
|
||||||
|
## 需求推演概要
|
||||||
|
|
||||||
|
### Step 11.1 — VectorRetriever trait
|
||||||
|
|
||||||
|
**核心需求**:定义一个与后端无关的语义检索抽象接口,包含 `index(id, embeddings)` 索引和 `search(query, k)` 检索两个方法。附带一个基于 `HashMap` 全量余弦相似度扫描的参考实现。
|
||||||
|
|
||||||
|
**边界识别**:
|
||||||
|
- 只定义 trait,不绑定任何具体后端(pgvector / qdrant / lancedb 留给社区或下游)
|
||||||
|
- 不引入第三方向量数据库依赖
|
||||||
|
- 引用实现的 `search()` 不做索引加速(O(n) 全量扫描已足够验证 trait 契约)
|
||||||
|
- 不与 `MemoryStore` 耦合——`VectorRetriever` 是独立维度
|
||||||
|
- 不嵌入到 `AgentSession` 或 `ContextSlot`(Phase 11 不承担集成消费端)
|
||||||
|
|
||||||
|
**关键假设**:
|
||||||
|
- `Vec<f32>` 作为 embedding 类型已足够(大部分 embedding 模型输出 f32 向量)
|
||||||
|
- 余弦相似度作为默认评分函数可覆盖主流场景
|
||||||
|
- InMemoryVectorRetriever 的 `Mutex<HashMap>` 在 ~10K 向量内性能可接受
|
||||||
|
|
||||||
|
### Step 11.2 — wiremock Provider roundtrip 测试
|
||||||
|
|
||||||
|
**核心需求**:补充 12 个 wiremock 测试,覆盖目前缺失的关键回归场景——结构化 JSON 错误体解析、请求头验证、429 限流头解析、流式边界、工具调用端到端。
|
||||||
|
|
||||||
|
**边界识别**:
|
||||||
|
- 只做 HTTP mock 层验证,不做端到端 LLM 模型调用
|
||||||
|
- 每个测试自包含(启动自己的 MockServer),不抽共享 helper
|
||||||
|
- 测试集中在 OpenAI(`GenericOpenaiProvider`)和 Anthropic(独立实现)两个核心 Provider 上
|
||||||
|
- DeepSeek/Qwen/Ollama 同属 OpenAI Compat,继承 `GenericOpenaiProvider` 的测试覆盖
|
||||||
|
|
||||||
|
**关键假设**:
|
||||||
|
- wiremock 的 `body_partial_json` matcher 可用且稳定(当前 dev-dependencies 中已有 wiremock)
|
||||||
|
- OpenAI 和 Anthropic 的结构化错误体格式在当前 SDK 版本中未变化
|
||||||
|
|
||||||
|
### Step 11.3 — 并发测试补强
|
||||||
|
|
||||||
|
**核心需求**:验证 `MemoryStore` 两种实现(InMemoryStore + SqliteStore)在多线程并发写和混合读写场景下的正确性。
|
||||||
|
|
||||||
|
**边界识别**:
|
||||||
|
- 不测试 `KnowledgeStore` / `ConversationMemory` 的并发——它们的行为完全由 `MemoryStore` 决定,不引入新 race 条件
|
||||||
|
- 不测试 TTL 淘汰的并发正确性(TTL 淘汰使用 wall clock,非原子,不保证精确)
|
||||||
|
- 混合读写测试只验证"无 panic + 数量正确",不验证"读到的结果恰好与写顺序一致"(后者需要强一致快照,当前 Mutex 模型不提供)
|
||||||
|
|
||||||
|
**关键假设**:
|
||||||
|
- `tokio::spawn` 100 个 task 同时写入 `Mutex<HashMap>`(InMemoryStore)不会死锁
|
||||||
|
- SqliteStore 的 WAL 模式 + `busy_timeout=5000` 足够容忍 100 并发写
|
||||||
|
|
||||||
|
## 当前状态分析
|
||||||
|
|
||||||
|
### 测试覆盖率现状
|
||||||
|
|
||||||
|
| 维度 | 当前值 | Phase 11 目标 |
|
||||||
|
|------|--------|-------------|
|
||||||
|
| 全量测试 | 254 passed | 275+ passed |
|
||||||
|
| InMemoryStore 测试 | 6 个(save/get/list/upsert/eviction/TTL) | +4 个并发 |
|
||||||
|
| SqliteStore 测试 | 9 个(含 1 个 10×10 并发) | +1 个 100 并发 |
|
||||||
|
| OpenAI wiremock 测试 | 4 个(basic/401/500/stream) | +8 个 |
|
||||||
|
| Anthropic wiremock 测试 | 4 个(basic/401/529/stream) | +4 个 |
|
||||||
|
| 请求头验证测试 | 0 个 | +2 个 |
|
||||||
|
| ToolUse 端到端 mock 测试 | 0 个 | +2 个(OpenAI + Anthropic) |
|
||||||
|
|
||||||
|
### Provider 测试缺口
|
||||||
|
|
||||||
|
现有 wiremock 测试仅覆盖最基础的响应路径,以下关键场景缺失回归保护:
|
||||||
|
|
||||||
|
| 场景 | 缺失风险 |
|
||||||
|
|------|---------|
|
||||||
|
| OpenAI 请求体格式验证 | `body_partial_json` 未匹配,请求体结构变化无声 |
|
||||||
|
| Authorization header 验证 | header 注入被修改时不告警 |
|
||||||
|
| 结构化 401 JSON 错误体 | `error.message`/`error.code` 未消费,错误消息丢失 |
|
||||||
|
| 429 + `retry-after` 头 | `RateLimit.retry_after` 字段不准确 |
|
||||||
|
| ToolUse 端到端 mock | tool_flow 解析路径无回归 |
|
||||||
|
| 流式 last chunk usage-only | `{choices:[], usage:{...}}` 可能 panic |
|
||||||
|
|
||||||
|
### MemoryStore 并发测试缺口
|
||||||
|
|
||||||
|
| Store | 当前并发测试 | 覆盖度 | 风险 |
|
||||||
|
|-------|-------------|--------|------|
|
||||||
|
| InMemoryStore | 0 个 | 无 | `Mutex` 锁竞争、deadlock、写入丢失 |
|
||||||
|
| SqliteStore | 1 个(10 写者 × 10 次 = 100 条) | 中等 | `spawn_blocking` 线程池耗尽、WAL 锁等待超时 |
|
||||||
|
|
||||||
|
### 向量检索现状
|
||||||
|
|
||||||
|
`memory` 模块已有 `MemoryRetriever`(关键词检索)和 `retriever.rs` 中的 `RetrievalResult`/`ScoredItem` 类型。但语义向量检索维度完全空缺——无 trait、无引用实现、无测试。`docs/roadmap.md` 将 VectorRetriever 列为 P1,与 ContextSlot(P1,Phase 10)同级。
|
||||||
|
|
||||||
|
## 架构决策记录
|
||||||
|
|
||||||
|
| 决策 | 选择 | 放弃 | 理由 |
|
||||||
|
|------|------|------|------|
|
||||||
|
| 1. VectorRetriever trait 参数类型 | `Vec<f32>` 裸向量 | `Embedding` newtype | 包装类型增加可见复杂度但未提供运行时保护;大部分 embedding 模型输出 f32 向量;下游可自行包装 |
|
||||||
|
| 2. `search()` 返回类型 | `Vec<(String, f32)>` | `ScoredItem`/`RetrievalResult` 命名 struct | `(String, f32)` 是 (id, score) 的最小表达;Phase 3 的 `RetrievalResult` 绑定了 `KnowledgePage` 引用,不适合向量检索场景;tuple 在 consumer 侧模式匹配更简洁 |
|
||||||
|
| 3. 文件归属 | 新文件 `memory/vector.rs` | 合入 `memory/retriever.rs` | `retriever.rs` 已承载 302 行关键词检索代码,语义维度独立不应耦合;`vector.rs` 作为独立模块便于后期扩展(pgvector adapter 等) |
|
||||||
|
| 4. 是否附带引用实现 | `InMemoryVectorRetriever` | trait-only | trait-only 是纯推测代码,无 consumer 验证引用实现作为"编译期测试"验证 trait 方法签名可用 |
|
||||||
|
| 5. InMemoryVectorRetriever 余弦相似度实现方式 | 手动三行点积/范数 | `ndarray`/`approx` 等第三方依赖 | 余弦相似度数学固定,不需要外部依赖;1e-10 防零除;零新依赖原则 |
|
||||||
|
| 5a | 引用实现不做向量维度校验 | 运行时维度检查 | 维度校验是具体后端(pgvector等)的职责;引用实现面向测试/验证场景;调用方负责传入等长向量 |
|
||||||
|
| 6. 错误类型 | 复用 `MemoryError` 现有变体 | 新增 `VecRetrieval` 变体 | 向量检索与关键词检索语义等价于"检索";`RetrievalError` 变体已覆盖索引/评分异常场景 |
|
||||||
|
| 7. wiremock 测试组织 | 自包含(每个测试启动自己的 MockServer) | 共享 helper 函数 | 沿用现有测试模式(openai.rs line 825+、anthropic.rs line 900+);自包含测试可独立运行、定位更直接 |
|
||||||
|
| 8. 请求头验证 | 做(`body_partial_json` + `header` matcher) | 跳过 | 回归防御价值高——Provider 请求体结构变化会直接导致请求被拒绝,头验证是低成本高收益的回归保护 |
|
||||||
|
| 9. 并发测试模式 | 100 并发写 + 混合读写(5 读 + 5 写)双模式 | 只做 100 并发写 | 两种模式互补:纯写入验证数据完整性和无 id 重复;混合读写验证读操作在并发写期间不 panic 且返回有效数据 |
|
||||||
|
| 10. 实施顺序 | 11.1 → 11.2 → 11.3 | 任意顺序 | 与 roadmap 原定的 Step 顺序一致;11.1 是纯新增可独立交付;11.2/11.3 是对既有代码的测试追加,可并行但不优先于 11.1 |
|
||||||
|
|
||||||
|
## 设计方案
|
||||||
|
|
||||||
|
### Step 11.1 — VectorRetriever trait + InMemoryVectorRetriever
|
||||||
|
|
||||||
|
#### 文件位置
|
||||||
|
|
||||||
|
- 新增:`src/memory/vector.rs`
|
||||||
|
- 修改:`src/memory.rs`(+2 行:module 声明 + re-export)
|
||||||
|
|
||||||
|
#### Trait 定义
|
||||||
|
|
||||||
|
```rust
|
||||||
|
/// 语义向量检索器抽象接口。
|
||||||
|
///
|
||||||
|
/// 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
|
||||||
|
/// 默认引用实现 [`InMemoryVectorRetriever`] 基于进程内 HashMap + 余弦相似度。
|
||||||
|
///
|
||||||
|
/// **稳定性**:实验性 API(v0.2.x),方法签名可能在 v0.3 中调整。
|
||||||
|
/// 若未来需要 `remove()` / `clear()` 等方法,将在此 trait 中追加(带默认实现)。
|
||||||
|
#[async_trait]
|
||||||
|
pub trait VectorRetriever: Send + Sync {
|
||||||
|
/// 将 `id` 对应的文本向量 `embeddings` 加入索引。
|
||||||
|
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError>;
|
||||||
|
|
||||||
|
/// 检索与 `query` 向量最相似的 `k` 条记录。
|
||||||
|
/// 返回 `Vec<(id, score)>`,按 score 降序排列,score ∈ [0.0, 1.0]。
|
||||||
|
async fn search(&self, query: Vec<f32>, k: usize) -> Result<Vec<(String, f32)>, MemoryError>;
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### InMemoryVectorRetriever 实现要点
|
||||||
|
|
||||||
|
```rust
|
||||||
|
pub struct InMemoryVectorRetriever {
|
||||||
|
vectors: Mutex<HashMap<String, Vec<f32>>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl InMemoryVectorRetriever {
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self {
|
||||||
|
vectors: Mutex::new(HashMap::new()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl VectorRetriever for InMemoryVectorRetriever {
|
||||||
|
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError> {
|
||||||
|
let mut vectors = self.vectors.lock().map_err(|e| {
|
||||||
|
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||||
|
})?;
|
||||||
|
vectors.insert(id, embeddings);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn search(&self, query: Vec<f32>, k: usize) -> Result<Vec<(String, f32)>, MemoryError> {
|
||||||
|
let vectors = self.vectors.lock().map_err(|e| {
|
||||||
|
MemoryError::RetrievalError(format!("lock poisoned: {e}"))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
if vectors.is_empty() || k == 0 {
|
||||||
|
return Ok(Vec::new());
|
||||||
|
}
|
||||||
|
|
||||||
|
let query_norm = dot(&query, &query).sqrt();
|
||||||
|
if query_norm == 0.0 {
|
||||||
|
return Ok(Vec::new());
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut scored: Vec<(String, f32)> = vectors
|
||||||
|
.iter()
|
||||||
|
.map(|(id, vec)| {
|
||||||
|
let dot_product = dot(&query, vec);
|
||||||
|
let vec_norm = dot(vec, vec).sqrt();
|
||||||
|
let similarity = dot_product / (query_norm * vec_norm + 1e-10);
|
||||||
|
(id.clone(), similarity)
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
// 降序排列
|
||||||
|
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
|
||||||
|
scored.truncate(k);
|
||||||
|
Ok(scored)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 点积(手动循环,零依赖)。
|
||||||
|
///
|
||||||
|
/// 注意:`zip` 对不等长向量静默截断到较短者。引用实现不做维度校验,
|
||||||
|
/// 调用方应确保 `a` 和 `b` 等长——不等长时结果无意义但不 panic。
|
||||||
|
fn dot(a: &[f32], b: &[f32]) -> f32 {
|
||||||
|
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 边界与约束
|
||||||
|
|
||||||
|
- **无维度校验**:不同维度向量传入 `search()` 时点积不报错,但余弦相似度结果无意义。维度校验是具体后端(pgvector等)的职责,引用实现不做运行时检查。
|
||||||
|
- **零向量处理**:query 为零向量时直接返回空结果(`query_norm == 0.0`)。
|
||||||
|
- **1e-10 防零除**:避免空库或全零向量导致除零 panic。
|
||||||
|
|
||||||
|
#### 测试(4 个)
|
||||||
|
|
||||||
|
| # | 测试名 | 验证点 |
|
||||||
|
|---|--------|--------|
|
||||||
|
| 1 | `basic_index_and_search` | index 两条("rust" + "python"),用 "rustacean" 查询应排在首位 |
|
||||||
|
| 2 | `search_empty_store` | 空库返回空 Vec |
|
||||||
|
| 3 | `concurrent_index` | 10 个 task 各 index 1 条,总量 10、id 无重复 |
|
||||||
|
| 4 | `concurrent_index_and_search` | 10 个 writer + 5 个 searcher 并发 2 秒,`search()` 遍历期间 `index()` 写入锁竞争不 panic |
|
||||||
|
|
||||||
|
#### 修改 `src/memory.rs`
|
||||||
|
|
||||||
|
```rust
|
||||||
|
pub mod vector;
|
||||||
|
|
||||||
|
// 在高频 re-export 区追加
|
||||||
|
pub use vector::{InMemoryVectorRetriever, VectorRetriever};
|
||||||
|
```
|
||||||
|
|
||||||
|
### Step 11.2 — wiremock Provider roundtrip 测试
|
||||||
|
|
||||||
|
#### 测试清单
|
||||||
|
|
||||||
|
全部 12 个测试均遵循现有自包含模式:`MockServer::start()` → `Mock::given(...).and(...).respond_with(...)` → `provider.chat_blocking(...)` / `provider.chat_stream_inner(...)` → assert。
|
||||||
|
|
||||||
|
##### P0(7 个)
|
||||||
|
|
||||||
|
| # | 测试名 | 所属文件 | Mock 关键点 | 断言 |
|
||||||
|
|---|--------|---------|------------|------|
|
||||||
|
| 1 | `openai_request_body_format` | `openai.rs` | `body_partial_json` 匹配 `{"model": "gpt-4o", "messages": [{"role": "user"}]}` | 请求体结构正确,响应解析正常 |
|
||||||
|
| 2 | `openai_authorization_header` | `openai.rs` | `header("authorization", "Bearer sk-test")` | header 精确匹配,响应解析正常 |
|
||||||
|
| 3 | `openai_401_structured_error` | `openai.rs` | 返回 401 + `{"error": {"message": "Incorrect API key", "code": "invalid_api_key"}}` | `LlmError::Authentication(msg)` 且 message 包含 "Incorrect API key" |
|
||||||
|
| 4 | `anthropic_401_structured_error` | `anthropic.rs` | 返回 401 + `{"error": {"type": "authentication_error", "message": "Invalid API key provided"}}` | `LlmError::Authentication(msg)` 且 message 包含 "Invalid API key" |
|
||||||
|
| 5 | `openai_429_with_retry_after` | `openai.rs` | 返回 429 + `{"error": {"message": "Rate limit exceeded"}}` + `retry-after: 30` 头 | `LlmError::RateLimit { retry_after: Some(30s) }` |
|
||||||
|
| 6 | `openai_tool_use_response` | `openai.rs` | 返回包含 `tool_calls` 的响应(choices[0].message.tool_calls ≠ null) | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
|
||||||
|
| 7 | `anthropic_tool_use_response` | `anthropic.rs` | 返回含 `type: "tool_use"` content block + `stop_reason: "tool_use"`(Anthropic 独立 wire 格式) | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
|
||||||
|
|
||||||
|
##### P1(5 个)
|
||||||
|
|
||||||
|
| # | 测试名 | 所属文件 | Mock 关键点 | 断言 |
|
||||||
|
|---|--------|---------|------------|------|
|
||||||
|
| 8 | `anthropic_version_header` | `anthropic.rs` | `header("anthropic-version", "2023-06-01")` | header 精确匹配 |
|
||||||
|
| 9 | `openai_stream_usage_only_last_chunk` | `openai.rs` | 流式最后 chunk `{"choices":[],"usage":{"prompt_tokens":5,"completion_tokens":2,"total_tokens":7}}` | 不 panic;`MessageComplete` 包含正确 usage |
|
||||||
|
| 10 | `anthropic_529_overloaded_structured` | `anthropic.rs` | 返回 529 + `{"error": {"type": "overloaded_error", "message": "Overloaded"}}` | `LlmError::RateLimit { retry_after: None }` |
|
||||||
|
| 11 | `openai_500_structured_error` | `openai.rs` | 返回 500 + `{"error": {"message": "Internal server error", "type": "server_error"}}` | `LlmError::Request { status: 500, body }` 且 body 包含 "Internal server error" |
|
||||||
|
| 12 | `openai_stream_mid_stream_error` | `openai.rs` | 流式前几个 chunk 正常,中途服务端断开连接(模拟网络中断/限流断开) | `LlmError::Request(_)` — 流式中断映射为请求错误 |
|
||||||
|
|
||||||
|
#### 测试模式说明
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// 每个测试自包含,不抽共享 helper(沿用现有模式)
|
||||||
|
#[tokio::test]
|
||||||
|
async fn openai_authorization_header() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/chat/completions"))
|
||||||
|
.and(header("authorization", "Bearer sk-test"))
|
||||||
|
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||||
|
"id": "chatcmpl-hdr",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": 1,
|
||||||
|
"model": "gpt-4o",
|
||||||
|
"choices": [{"index": 0, "message": {"role": "assistant", "content": "OK"}, "finish_reason": "stop"}],
|
||||||
|
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
|
||||||
|
})))
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = GenericOpenaiProvider::new_with_name(
|
||||||
|
server.uri(), "sk-test".into(), "gpt-4o".into(), "openai", 30,
|
||||||
|
);
|
||||||
|
let response = provider.chat_blocking(MessageRequest {
|
||||||
|
model: "gpt-4o".into(),
|
||||||
|
messages: vec![Message::user_text("hi")],
|
||||||
|
..Default::default()
|
||||||
|
}).await.unwrap();
|
||||||
|
assert_eq!(response.text(), "OK");
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 新增依赖
|
||||||
|
|
||||||
|
dev-dependencies 中 wiremock 已就绪(当前 openai.rs / anthropic.rs 已在测试中使用),无需新增。
|
||||||
|
|
||||||
|
### Step 11.3 — 并发测试补强
|
||||||
|
|
||||||
|
#### 测试清单
|
||||||
|
|
||||||
|
| # | 测试名 | Store | 模式 | 验证标准 |
|
||||||
|
|---|--------|-------|------|---------|
|
||||||
|
| 1 | `concurrent_writers_max_pressure` | InMemoryStore | 100 task × 1 write | 总量 100, id 无重复 |
|
||||||
|
| 2 | `concurrent_writers_max_pressure` | SqliteStore | 100 task × 1 write | 总量 100, id 无重复 |
|
||||||
|
| 3 | `concurrent_mixed_read_write` | InMemoryStore | 预热 20 条, 5 读 + 5 写并发 2 秒 | 无 panic |
|
||||||
|
| 4 | `concurrent_mixed_read_write` | SqliteStore | 预热 20 条, 5 读 + 5 写并发 2 秒 | 无 panic |
|
||||||
|
| 5 | `concurrent_capacity_eviction` | InMemoryStore | 15 写者, max_items=10 | 最终 ≤ 10 |
|
||||||
|
|
||||||
|
#### 关键实现要点
|
||||||
|
|
||||||
|
**100 并发写模式**(InMemoryStore + SqliteStore 各一):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_writers_max_pressure() {
|
||||||
|
let store = Arc::new(InMemoryStore::new());
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for i in 0..100 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let id = format!("concurrent_{i}");
|
||||||
|
s.save(make_item(&id)).await.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 100);
|
||||||
|
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
|
||||||
|
ids.sort();
|
||||||
|
ids.dedup();
|
||||||
|
assert_eq!(ids.len(), 100);
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**混合读写模式**(InMemoryStore + SqliteStore 各一):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_mixed_read_write() {
|
||||||
|
let store = Arc::new(InMemoryStore::new());
|
||||||
|
|
||||||
|
// 预热
|
||||||
|
for i in 0..20 {
|
||||||
|
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
// 5 个写者
|
||||||
|
for w in 0..5 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||||
|
let mut i = 0;
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let id = format!("writer{w}_item{i}");
|
||||||
|
s.save(make_item(&id)).await.unwrap();
|
||||||
|
i += 1;
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
// 5 个读者
|
||||||
|
for r in 0..5 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let _ = s.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
// 不 panic 即算通过
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**容量淘汰并发模式**(InMemoryStore):
|
||||||
|
|
||||||
|
```rust
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_capacity_eviction() {
|
||||||
|
let eviction = EvictionConfig {
|
||||||
|
policy: EvictionPolicy::Capacity { max_items: 10 },
|
||||||
|
check_interval: 1,
|
||||||
|
};
|
||||||
|
let store = Arc::new(InMemoryStore::with_eviction(eviction));
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for i in 0..15 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
s.save(make_item(&format!("item_{i}"))).await.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert!(list.len() <= 10);
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**测试归属**:
|
||||||
|
- InMemoryStore 并发测试 → `src/memory/store/in_memory.rs` 的 `mod tests`
|
||||||
|
- SqliteStore 并发测试 → `src/memory/store/sqlite_store.rs` 的 `mod tests`
|
||||||
|
- 混合读写测试中的 `make_item` 辅助函数:直接复用各文件现有 `fn make_item`
|
||||||
|
|
||||||
|
## 已否决的方案
|
||||||
|
|
||||||
|
### 1. 砍掉 11.1(PM 建议)
|
||||||
|
|
||||||
|
**内容**:PM 在讨论中提出 VectorRetriever trait 无消费者,建议整体砍掉,等 Phase 12 或 v0.3 有人用时再做。
|
||||||
|
|
||||||
|
**否决理由**:引用实现作为 trait 契约的编译期验证手段——没有 consumer 不意味着 trait 签名不需要测试。同时社区贡献(pgvector adapter 等)需要稳定的 trait 边界。附带引用实现还可作为"如何在 agcore 中实现一个 VectorRetriever"的示例,降低社区参与门槛。代码量仅 ~80 行,维护成本可忽略。
|
||||||
|
|
||||||
|
### 2. trait-only VectorRetriever(无引用实现)
|
||||||
|
|
||||||
|
**内容**:只定义 `VectorRetriever` trait,不做 `InMemoryVectorRetriever`。
|
||||||
|
|
||||||
|
**否决理由**:trait-only 是纯推测代码——没有运行时验证,无法确认 trait 方法签名在实际调用链中是否可编译。理想情况下每个 trait 至少有一个引用实现来验证"这个 trait 确实可以被实现"。
|
||||||
|
|
||||||
|
### 3. 跳过请求头验证
|
||||||
|
|
||||||
|
**内容**:请求体格式和 Authorization header 验证是"过度保护"。
|
||||||
|
|
||||||
|
**否决理由**:`body_partial_json` + `header` matcher 的回归防御价值高。Provider 适配层的最大风险是请求体结构无声变更(如 `ToolDef` IR 切换时漏改了序列化字段),头验证是低成本(每测试 ~5 行)高收益的回归保护。
|
||||||
|
|
||||||
|
### 4. 先发 v0.2.0 正式版再迭代
|
||||||
|
|
||||||
|
**内容**:当前 rc.1 已经包含所有 P0 功能,建议直接发正式版,Phase 11 推迟到 v0.2.1。
|
||||||
|
|
||||||
|
**否决理由**:测试补强是正式版的信号而非负担。Phase 11 的三个 Step 都是"如果现在不做,以后更不会做"的类型。在正式版前补齐测试基线,避免「发布了再补测试」的经典陷阱。
|
||||||
|
|
||||||
|
### 5. 只做 100 并发写,不做混合读写
|
||||||
|
|
||||||
|
**内容**:并发写验证数据完整性已足够。
|
||||||
|
|
||||||
|
**否决理由**:纯写入和混合读写暴露不同类型的 bug。纯写入验证"数据不丢、id 无重复";混合读写验证"读操作在并发写期间不 panic、返回有效数据"。两种模式互补缺失。
|
||||||
|
|
||||||
|
### 6. (实施后补充)openai_stream_mid_stream_error 的 mock 模式偏差
|
||||||
|
|
||||||
|
**实际实施**:返回 `200 + SSE content-type + 畸形 JSON payload`(`data: {not-valid-json}\n\ndata: [DONE]\n\n`),断言 `ChunkToEventStream` 产出 `StreamEvent::Error`。
|
||||||
|
|
||||||
|
**方案原文**:返回"前几个 chunk 正常,中途服务端断开连接",断言 `LlmError::Request(_)`。
|
||||||
|
|
||||||
|
**偏差原因**:wiremock 0.6 标准 responder 的 `set_delay` / `set_body_string` 行为是「延迟响应 + 发完 body 后关闭连接」,无法精确模拟「send partial body then hang 保持连接」。`read_timeout` 配合 `set_delay` 触发的超时属于 send 阶段(`LlmError::Timeout`),不属于流式中断。
|
||||||
|
|
||||||
|
**采纳方案**:用畸形 JSON payload 替代——同样验证"流中途产生错误事件而不 panic"的回归保护意图,且在 wiremock 0.6 上 100% 可重现。断言改为 `StreamEvent::Error{message}` + "stream 最终结束",保持核心回归价值。
|
||||||
|
|
||||||
|
**影响**:测试意图(流阶段错误检测)完全保留;mock 行为从「TCP 断开」变为「畸形应用层数据」;断言从 `LlmError::Request` 改为 `StreamEvent::Error`(语义等价:客户端发现流异常)。
|
||||||
|
|
||||||
|
### 7. (实施后补充)openai.rs 的 429 retry-after 解析修复
|
||||||
|
|
||||||
|
**实际实施**:`handle_error_response` 新增 `retry-after` header 解析逻辑(与 anthropic.rs 完全对齐)。
|
||||||
|
|
||||||
|
**方案原文**:方案测试 11.2.4 要求 `RateLimit { retry_after: Some(30s) }`,但 ADC 表中未显式列出此修复作为生产代码变更。
|
||||||
|
|
||||||
|
**修复原因**:原 `openai.rs:186-191` 的 429 分支固定 `retry_after: None`,与 `anthropic.rs:339-351` 已有的解析逻辑不一致。原代码注释甚至已写"仅读取 retry-after",但实际未实现——这是隐藏 bug。修复让 OpenAI 兼容层(DeepSeek/Qwen 等)的限流重试信息可用,与 Anthropic 行为统一。
|
||||||
|
|
||||||
|
**影响**:方案测试 11.2.4 从「不可通过的回归保护」变为「可验证的实际行为」。变更 5 行,与 anthropic 实现完全镜像。
|
||||||
|
|
||||||
|
## 实施计划与顺序
|
||||||
|
|
||||||
|
**实施顺序**:11.1 → 11.2 → 11.3(与 roadmap 一致,每步可单独交付验证)。
|
||||||
|
|
||||||
|
| Step | 文件变更 | 测试增量 | 预估代码量 | 验证标准 |
|
||||||
|
|------|---------|---------|-----------|---------|
|
||||||
|
| 11.1 | +`src/memory/vector.rs`(~80 行),~`src/memory.rs`(+2 行) | +4 | ~85 行实现 + 70 行测试 | `cargo build --all-targets` 编译通过,4 个测试通过 |
|
||||||
|
| 11.2 | ~`src/llm/provider/openai.rs`(+6 个测试),~`src/llm/provider/anthropic.rs`(+2 个测试) | +12(7 P0 + 5 P1) | ~290 行(含 test mod 和 mock 数据) | `cargo test --all-targets` 全绿,wiremock 12 场景均绿 |
|
||||||
|
| 11.3 | ~`src/memory/store/in_memory.rs`(+3 个测试),~`src/memory/store/sqlite_store.rs`(+2 个测试) | +5 | ~120 行 | `cargo test --all-targets` 全绿 |
|
||||||
|
| **总计** | 6 个文件(1 新增 + 5 修改) | +21 | ~500 行 | 全量 254 → 275+,`cargo clippy --all-targets -- -D warnings` 0 警告 |
|
||||||
|
|
||||||
|
### 验证通过标准
|
||||||
|
|
||||||
|
1. `cargo build --all-targets` —— 编译通过,无 warning
|
||||||
|
2. `cargo test --all-targets` —— 全部通过(254 + 21 = 275+)
|
||||||
|
3. `cargo clippy --all-targets -- -D warnings` —— 0 警告
|
||||||
|
4. `cargo test --all-targets 2>&1 | grep -E "test result:"` —— 确认新增测试全部出现在执行列表中
|
||||||
|
5. 新增 wiremock 测试单独验证网络隔离(无需 API key,纯本地 mock)
|
||||||
|
|
||||||
|
## 参考来源
|
||||||
|
|
||||||
|
- **讨论收口结论**:Phase 11 讨论,含 PM/SA 双视角输入(2026-07-07)
|
||||||
|
- **现有代码模式**:
|
||||||
|
- `src/llm/provider/openai.rs` line 824-1034——wiremock 测试模式(`MockServer::start` → `Mock::given(...).and(...).respond_with(...)` → `provider.chat_blocking` → assert)
|
||||||
|
- `src/llm/provider/anthropic.rs` line 899-1079——Anthropic provider wiremock 测试
|
||||||
|
- `src/memory/store/sqlite_store.rs` line 458-483——`concurrent_writers_no_data_loss` 10×10 并发模式
|
||||||
|
- `src/memory/store.rs`——`MemoryStore` trait 定义(`#[async_trait]` 风格)
|
||||||
|
- `src/memory/error.rs`——`MemoryError` 枚举(`#[non_exhaustive]` + `RetrievalError` 变体)
|
||||||
|
- `src/memory/retriever.rs`——现有检索模块(`RetrievalResult` / `ScoredItem`)
|
||||||
|
- `src/memory.rs`——模块根与 re-export 模式
|
||||||
|
- **方案文档**:`docs/roadmap.md` Phase 11 章节(line 516-528)
|
||||||
|
- **编译器 pragma**:`#[non_exhaustive]` —— 新增枚举变体需要此标记,公开结构体字段未来变化预留兼容空间
|
||||||
|
|
||||||
|
## 关键假设与风险
|
||||||
|
|
||||||
|
### 关键假设清单
|
||||||
|
|
||||||
|
| # | 假设 | 影响 | 推翻后的应对 |
|
||||||
|
|---|------|------|------------|
|
||||||
|
| 1 | wiremock `body_partial_json` matcher 在 wiremock 0.6+ 中可用 | Step 11.2 测试 #1 的实现方式 | 改用 `body_json`(精确匹配)或 `body_string`(部分串匹配) |
|
||||||
|
| 2 | `tokio::spawn` 100 task 并发写入 `Mutex<HashMap>` 无死锁 | Step 11.3 #1 InMemoryStore 并发 | 降低并发数(50)继续验证,或换 `tokio::sync::Mutex` |
|
||||||
|
| 3 | SqliteStore 的 `busy_timeout=5000` 能容忍 100 并发写 | Step 11.3 #2 SqliteStore 并发 | 增加 `busy_timeout`(10s),或限制最大并发数 |
|
||||||
|
| 4 | 新增测试不使用 wiremock 以外的未列在 dev-dependencies 中的依赖 | Phase 11 零新增外部依赖 | 若需要额外 matcher,评估后加入 dev-dependencies |
|
||||||
|
| 5 | InMemoryVectorRetriever 的 O(n) 全量扫描在测试规模下 (<1000 向量) 性能可接受 | Step 11.1 测试通过 | 如果有竞态问题,改为 read/write lock(`RwLock<HashMap>`) |
|
||||||
|
| 6 | `MemoryError::RetrievalError` 变体足够覆盖向量检索的索引/评分失败场景 | Step 11.1 错误映射 | 如果不够,可增加新的 `MemoryError` 变体 |
|
||||||
|
| 7 | OpenAI 和 Anthropic 的结构化错误体格式在当前 SDK 版本中未变化 | Step 11.2 测试 #3/#4/#7/#10/#11(结构化错误解析断言) | 若 SDK 变更错误体格式,更新 mock body 和断言匹配新格式 |
|
||||||
|
| 8 | 调用方传入 `search()` 的向量与已索引向量维度一致(引用实现不做维度校验,`dot()` 对不等长向量静默截断) | Step 11.1 InMemoryVectorRetriever 正确性 | 若须维度校验,在 `index()` 时记录维度并在 `search()` 时断言;引用实现维持零校验 |
|
||||||
|
|
||||||
|
### 已识别的风险
|
||||||
|
|
||||||
|
| 风险 | 等级 | 缓解措施 |
|
||||||
|
|------|------|---------|
|
||||||
|
| wiremock `body_partial_json` matcher 行为在版本升级后变化 | 低 | 限定 wiremock 版本范围(当前已在 Cargo.lock 中锁定);P0 测试不依赖该 matcher |
|
||||||
|
| 100 并发写暴露 SqliteStore 的 `spawn_blocking` 线程池瓶颈 | 中 | 观察 CI 执行时间;如果超时,降低并发到 50 或增加 `max_blocking_threads` |
|
||||||
|
| InMemoryVectorRetriever 的 `Mutex` 锁争用导致测试 flaky | 低 | Mutex 不会死锁(单线程持有不 await),测试不依赖精确时序 |
|
||||||
|
| 新增 wiremock 测试与现有测试冲突(端口占用) | 低 | `MockServer::start()` 自动选择随机端口,不冲突 |
|
||||||
|
| `cargo test --all-targets` 执行时间增加 >30% | 低 | 预估 +21 个测试,增量约 8%(254→275),其中 wiremock 测试有网络 IO 但延迟 <10ms/个 |
|
||||||
|
|
||||||
|
### 非阻塞已知项
|
||||||
|
|
||||||
|
- **Ollama Provider** 是 OpenAI Compat,wiremock 测试继承 `GenericOpenaiProvider`,不单独新增
|
||||||
|
- **DeepSeek / Qwen Provider** 同样通过 `GenericOpenaiProvider` 实现,继承测试
|
||||||
|
- **Step 11.1 不消费到 AgentSession / ContextSlot**,留待 Phase 12 或 v0.3 做消费端集成
|
||||||
|
- **Phase 11 完成后**,全量测试预计 275+,`cargo test --all-targets` 执行时间预计 < 60s
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 实施计划(附录)
|
||||||
|
|
||||||
|
**实施顺序**:11.1 → 11.2 ‖ 11.3(11.2 与 11.3 无文件冲突,可并行交付;11.2 优先因回归保护价值更高)。每步完成后运行 `cargo test --all-targets` + `cargo clippy --all-targets -- -D warnings` 验证无回归。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 11.1 — VectorRetriever trait + InMemoryVectorRetriever(~160 行,4 测试)
|
||||||
|
|
||||||
|
**前置**:无。与 Step 11.2/11.3 可并行开发但优先交付。
|
||||||
|
|
||||||
|
| 任务 | 描述 | 文件 | 前置 | 工作量 | 风险 | 验收条件 |
|
||||||
|
|------|------|------|------|--------|------|---------|
|
||||||
|
| **11.1.1** | 创建 `memory/vector.rs`:定义 `VectorRetriever` trait + `InMemoryVectorRetriever` struct + `dot()` 辅助函数 | `src/memory/vector.rs` | 无 | S | 低 | `cargo build --all-targets` 编译通过 |
|
||||||
|
| **11.1.2** | 实现 `InMemoryVectorRetriever`:`index()` — Mutex insert;`search()` — 全量余弦扫描 + 降序排列 + k 截断 | `src/memory/vector.rs` | 11.1.1 | S | 低 | trait 实现编译通过;`Mutex::lock()` 使用 `map_err` 处理 poison,不 panic |
|
||||||
|
| **11.1.3** | 添加 4 个内联测试:`basic_index_and_search`、`search_empty_store`、`concurrent_index`、`concurrent_index_and_search` | `src/memory/vector.rs` (mod tests) | 11.1.2 | S | 低 | 4 测试全部通过 |
|
||||||
|
| **11.1.4** | 修改 `src/memory.rs`:加 `pub mod vector;` + `pub use vector::{VectorRetriever, InMemoryVectorRetriever};` | `src/memory.rs` | 11.1.1 | S | 低 | `cargo build --all-targets` 无 warning |
|
||||||
|
|
||||||
|
**Step 验证**:
|
||||||
|
```
|
||||||
|
cargo test --all-targets # 254 + 4 = 258+ passed
|
||||||
|
cargo clippy --all-targets -- -D warnings # 0 warning
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 11.2 — wiremock Provider roundtrip 测试(12 测试,P0=7 + P1=5)
|
||||||
|
|
||||||
|
**前置**:无。与 Step 11.1 无文件冲突,可并行。
|
||||||
|
|
||||||
|
**`openai.rs` 新增测试(8 个:P0=5 + P1=3)**:
|
||||||
|
|
||||||
|
| 任务 | 测试名 | 优先级 | 前置 | 工作量 | 风险 | Mock 模式 | 验收条件 |
|
||||||
|
|------|--------|--------|------|--------|------|----------|---------|
|
||||||
|
| **11.2.1** | `openai_request_body_format` | P0 | 无 | S | 低 | `body_partial_json` 匹配 model/messages | 请求体结构正确,响应解析正常 |
|
||||||
|
| **11.2.2** | `openai_authorization_header` | P0 | 无 | S | 低 | `header("authorization", "Bearer sk-test")` | header 精确匹配 |
|
||||||
|
| **11.2.3** | `openai_401_structured_error` | P0 | 无 | S | 低 | 401 + `{"error":{"message":"...","code":"invalid_api_key"}}` | `LlmError::Authentication` 含 "Incorrect API key" |
|
||||||
|
| **11.2.4** | `openai_429_with_retry_after` | P0 | 无 | S | 低 | 429 + `retry-after: 30` | `RateLimit { retry_after: Some(30s) }` |
|
||||||
|
| **11.2.5** | `openai_tool_use_response` | P0 | 无 | S | 低 | 响应含 `tool_calls` | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
|
||||||
|
| **11.2.6** | `openai_stream_usage_only_last_chunk` | P1 | 无 | S | 低 | 流式最后 chunk `{choices:[], usage:{...}}` | 不 panic,`MessageComplete` 含正确 usage |
|
||||||
|
| **11.2.7** | `openai_500_structured_error` | P1 | 无 | S | 低 | 500 + `{"error":{"message":"server error"}}` | `Request { status: 500 }` 含 body |
|
||||||
|
| **11.2.8** | `openai_stream_mid_stream_error` | P1 | 无 | M | 中 | 前几个 chunk 正常后连接断开 | `LlmError::Request(_)` 流中断映射 |
|
||||||
|
|
||||||
|
**`anthropic.rs` 新增测试(4 个:P0=2 + P1=2)**:
|
||||||
|
|
||||||
|
| 任务 | 测试名 | 优先级 | 前置 | 工作量 | 风险 | Mock 模式 | 验收条件 |
|
||||||
|
|------|--------|--------|------|--------|------|----------|---------|
|
||||||
|
| **11.2.9** | `anthropic_401_structured_error` | P0 | 无 | S | 低 | 401 + `{"error":{"type":"authentication_error","message":"..."}}` | `LlmError::Authentication` 消息透传 |
|
||||||
|
| **11.2.10** | `anthropic_tool_use_response` | P0 | 无 | S | 低 | 响应含 `type:"tool_use"` content block + `stop_reason:"tool_use"` | `StopReason::ToolUse` + `ContentBlock::ToolUse` 正确解析 |
|
||||||
|
| **11.2.11** | `anthropic_version_header` | P1 | 无 | S | 低 | `header("anthropic-version", "2023-06-01")` | header 精确匹配 |
|
||||||
|
| **11.2.12** | `anthropic_529_overloaded_structured` | P1 | 无 | S | 低 | 529 + `{"error":{"type":"overloaded_error","message":"Overloaded"}}` | `RateLimit { retry_after: None }` |
|
||||||
|
|
||||||
|
**Step 验证**:
|
||||||
|
```
|
||||||
|
cargo test --all-targets # 258 + 12 = 270+ passed
|
||||||
|
cargo clippy --all-targets -- -D warnings # 0 warning
|
||||||
|
```
|
||||||
|
每个测试自包含(`MockServer::start()` → `Mock::given(...)` → `provider.chat_blocking()/chat_stream_inner()` → assert),无需共享 helper。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 11.3 — 并发测试补强(5 测试)
|
||||||
|
|
||||||
|
**前置**:无。与 Step 11.1/11.2 无文件冲突。
|
||||||
|
|
||||||
|
| 任务 | 测试名 | 文件 | 前置 | 工作量 | 风险 | 模式 | 验收条件 |
|
||||||
|
|------|--------|------|------|--------|------|------|---------|
|
||||||
|
| **11.3.1** | `concurrent_writers_max_pressure` | `in_memory.rs` | 无 | S | 中 | 100 task × 1 write | 总量 100,id 无重复 |
|
||||||
|
| **11.3.2** | `concurrent_writers_max_pressure` | `sqlite_store.rs` | 无 | S | 中 | 100 task × 1 write | 总量 100,id 无重复 |
|
||||||
|
| **11.3.3** | `concurrent_mixed_read_write` | `in_memory.rs` | 无 | S | 中 | 预热 20 条,5 写 + 5 读并发 2 秒 | 无 panic |
|
||||||
|
| **11.3.4** | `concurrent_mixed_read_write` | `sqlite_store.rs` | 无 | S | 中 | 预热 20 条,5 写 + 5 读并发 2 秒 | 无 panic |
|
||||||
|
| **11.3.5** | `concurrent_capacity_eviction` | `in_memory.rs` | 无 | S | 中 | 15 写者,max_items=10 | 最终 ≤ 10(竞争激烈时可能过渡态 >10,主断言 ≤ 10,宽松备选 ≤ 15) |
|
||||||
|
|
||||||
|
**实现要点**:
|
||||||
|
- 沿用现有 `Arc<Store> + tokio::spawn + h.await.unwrap()` 模式(参考 `sqlite_store.rs:458-483`)
|
||||||
|
- `make_item` 辅助函数直接复用各文件现有实现
|
||||||
|
- SqliteStore 测试使用 `:memory:` 数据库(与现有并发测试一致)
|
||||||
|
- 混合读写测试使用 `tokio::time::Instant::now() + Duration` 做时限
|
||||||
|
- 100 并发写是一次性 spawn 100 task(非分批),暴露最大锁竞争压力
|
||||||
|
|
||||||
|
**Step 验证**:
|
||||||
|
```
|
||||||
|
cargo test --all-targets # 270 + 5 = 275+ passed
|
||||||
|
cargo clippy --all-targets -- -D warnings # 0 warning
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 整体发布核查清单
|
||||||
|
|
||||||
|
| # | 检查项 | 验证命令 | 预期结果 |
|
||||||
|
|---|--------|---------|---------|
|
||||||
|
| 1 | 编译 | `cargo build --all-targets` | 通过,0 warning |
|
||||||
|
| 2 | 全量测试 | `cargo test --all-targets` | 275+ passed,0 failed |
|
||||||
|
| 3 | Lint | `cargo clippy --all-targets -- -D warnings` | 0 warning |
|
||||||
|
| 4 | 文档 | `cargo doc --no-deps` | 0 warning(VectorRetriever trait 公共 API doc 完整) |
|
||||||
|
| 5 | 确认新增测试 | `cargo test --all-targets 2>&1 \| grep -E "test result:"` | 所有新增测试名出现在执行列表中 |
|
||||||
|
| 6 | wiremock 隔离 | 新增 wiremock 测试不依赖网络 | 纯本地 mock,无需 API key |
|
||||||
|
| 7 | 并行安全 | 并发测试独立运行时无 flaky | 连续 3 次 `cargo test` 结果一致 |
|
||||||
|
| 8 | 存量零回归 | 已有 254 个测试全部通过 | 与 Phase 10 基线对比无 fail |
|
||||||
|
| 9 | 公共 API doc comment | `grep -r "pub trait VectorRetriever" src/ && rg "^///" -c src/memory/vector.rs` | trait 和方法都有 `///` 注释 |
|
||||||
|
|
||||||
|
**若核查项失败的回退策略**:
|
||||||
|
- **测试失败(P0)**:阻断发布。定位到具体测试名 → 检查 Mock JSON 格式与 Provider 解析逻辑是否匹配(结构化错误体格式变化 → 更新 mock body;流式状态机变化 → 更新 `chat_stream_inner` 路径测试)
|
||||||
|
- **测试失败(P1)**:不阻断发布。标记 `#[ignore]` + file issue,确认无 P0 失败后即可发布
|
||||||
|
- **clippy warning**:修复 lint 后重跑;若为 `#[allow(...)]` 可抑制,在 code review 中申明理由
|
||||||
|
- **flaky 并发测试**:检查 `tokio::spawn` 是否跨 `.await` 持锁;若 SqliteStore 超时,增加 `busy_timeout` 或降低并发数
|
||||||
|
- **11.1 模块发布阻塞**:若 `InMemoryVectorRetriever` 无法按时交付,可临时注释 `src/memory.rs` 中的 `pub mod vector;` 行,跳过整个模块(零消费者,不影响发布)。回退后再补交
|
||||||
@@ -0,0 +1,640 @@
|
|||||||
|
# Phase 13 — 热身清理 + ContextSlot fork/merge 实施方案
|
||||||
|
|
||||||
|
- **文档编号**:19
|
||||||
|
- **标题**:Phase 13 — 热身清理 + ContextSlot fork/merge 实施方案
|
||||||
|
- **日期**:2026-07-08
|
||||||
|
- **状态**:待实施
|
||||||
|
- **涉及模块**:agent/context、agent/session、llm/types、llm/provider/openai、llm/stream
|
||||||
|
- **关联文档**:roadmap.md(§Phase 13)、17-phase10-contextslot.md
|
||||||
|
- **对应**:Roadmap §Phase 13(v0.3.0 第一阶段)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. 背景与目标
|
||||||
|
|
||||||
|
v0.3.0 是 agcore 从"LLM 调用工具箱"升级为"多 Agent 基础系统"的关键版本。Phase 13 是 v0.3.0 的第一阶段,定位为"热身",包含两大部分:
|
||||||
|
|
||||||
|
- **技术债清理**:删除 Phase 0 遗留的旧 types 文件(`request.rs`、`response.rs`、`old_stream.rs`),以及已标记 `#[deprecated]` 的 `ChatResponse` 结构体
|
||||||
|
- **ContextSlot fork/merge**:为 ContextSlot 增加分叉和合并能力,为后续 Phase 17 Checkpointer 和 Phase 18 SubAgent Dispatch 打基础
|
||||||
|
|
||||||
|
**依赖关系**:无(独立交付)
|
||||||
|
|
||||||
|
**优先级**:P0
|
||||||
|
|
||||||
|
**预估规模**:净减 ~200 行代码(新增 ~505 行,删除 ~704 行)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. 需求分析
|
||||||
|
|
||||||
|
### 2.1 功能需求
|
||||||
|
|
||||||
|
1. **技术债清理**:删除 `src/llm/types/request.rs`(187 行)、`response.rs`(177 行)、`old_stream.rs`(45 行),将其中的 OpenAI wire-format 类型移入 `src/llm/provider/openai.rs`;删除 `types/mod.rs` 中的 `ChatResponse` 废弃结构体
|
||||||
|
2. **`ContextSlot::fork`**:从现有 context slot 分支出独立的子 slot
|
||||||
|
3. **`ContextSlot::merge`**:将子 slot 的消息合并回父 slot
|
||||||
|
4. **`MergeStrategy`** 枚举:Append(追加)/ Replace(替换),`#[non_exhaustive]` 预留 Phase 16 Summarize 扩展
|
||||||
|
|
||||||
|
### 2.2 非功能需求
|
||||||
|
|
||||||
|
- **每步可编译**:5 个 Step 按物理文件切割,每步 `cargo build --all-targets + cargo test` 验证
|
||||||
|
- **指定公共 API 路径保持向后兼容**:`agcore::llm::types::ToolChoice`(re-export 不变)、`crate::llm::stream::StreamEvent`(重导出保留);其余 wire-format 类型(`OpenaiChatRequest`、`OpenaiChatResponse/Chunk`、`StreamOptions` 等)移入 `provider/openai.rs` 后属 Breaking Change,详见 §4.3 CHANGELOG
|
||||||
|
- **向后兼容的 StreamEvent 路径**:`crate::llm::stream::StreamEvent` 重导出保留,不修改 `cycle.rs` 和 `session.rs` 的 import
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. 方案设计
|
||||||
|
|
||||||
|
### 3.1 整体架构
|
||||||
|
|
||||||
|
Phase 13 分为 5 个 Step,按执行顺序排列:
|
||||||
|
|
||||||
|
```
|
||||||
|
Step 13.5 (fork/merge) → Step 13.4 (ToolChoice) → Step 13.1 (request types) → Step 13.2 (response types) → Step 13.3 (cleanup)
|
||||||
|
```
|
||||||
|
|
||||||
|
这种顺序的好处:
|
||||||
|
|
||||||
|
- **先交付价值**:13.5 是唯一有用户功能交付的 Step,先做建立节奏
|
||||||
|
- **排序约束**:13.4 必须先于 13.1(ToolChoice 不搬走,request.rs 不能删)
|
||||||
|
- **13.3 收尾**:删除旧文件和 `ChatResponse` 是 breaking change,放在最后
|
||||||
|
|
||||||
|
### 3.2 Step 13.5 — ContextSlot fork/merge
|
||||||
|
|
||||||
|
#### MergeStrategy 枚举
|
||||||
|
|
||||||
|
定义在 `src/agent/context.rs`:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
#[non_exhaustive]
|
||||||
|
pub enum MergeStrategy {
|
||||||
|
/// 子 slot 消息追加到父 slot 末尾。
|
||||||
|
Append,
|
||||||
|
/// 用子 slot 消息替换父 slot 内容。
|
||||||
|
Replace,
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
- `#[non_exhaustive]` 保证 Phase 16 加入 `Summarize` 变体时不破坏现有代码
|
||||||
|
- 不预埋 `Summarize` 占位变体(YAGNI 原则)
|
||||||
|
|
||||||
|
#### ContextSlot::fork
|
||||||
|
|
||||||
|
```rust
|
||||||
|
impl ContextSlot {
|
||||||
|
pub fn fork(&self, child_id: String, strategy: DeriveStrategy) -> ContextSlot {
|
||||||
|
let messages = match &strategy {
|
||||||
|
DeriveStrategy::Full => self.messages.clone(),
|
||||||
|
DeriveStrategy::Focused(cfg) => Self::filter_focused(&self.messages, cfg),
|
||||||
|
};
|
||||||
|
tracing::debug!(
|
||||||
|
parent_id = %self.id,
|
||||||
|
child_id = %child_id,
|
||||||
|
?strategy,
|
||||||
|
"ContextSlot::fork"
|
||||||
|
);
|
||||||
|
ContextSlot {
|
||||||
|
id: child_id,
|
||||||
|
session_id: self.session_id.clone(),
|
||||||
|
config: SlotConfig {
|
||||||
|
mode: match &strategy {
|
||||||
|
DeriveStrategy::Full => SlotMode::Full,
|
||||||
|
DeriveStrategy::Focused(cfg) => SlotMode::Focused(cfg.clone()),
|
||||||
|
},
|
||||||
|
source: SlotSource::Derived {
|
||||||
|
parent_id: self.id.clone(),
|
||||||
|
strategy,
|
||||||
|
},
|
||||||
|
budget: self.config.budget.clone(),
|
||||||
|
compact: self.config.compact,
|
||||||
|
},
|
||||||
|
messages,
|
||||||
|
meta: SlotMeta::new(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
设计要点:
|
||||||
|
|
||||||
|
- 纯数据层操作,不持久化
|
||||||
|
- 子 slot 的 `meta` 全新创建(`SlotMeta::new()`),不继承父 slot 的 message_count
|
||||||
|
- 子 slot 的 source 记录 `parent_id`,血缘可追溯
|
||||||
|
- 添加 `tracing::debug!` 日志,支持多 slot 交互场景的审计追踪
|
||||||
|
|
||||||
|
#### ContextSlot::merge
|
||||||
|
|
||||||
|
```rust
|
||||||
|
impl ContextSlot {
|
||||||
|
/// 将子 slot 的消息合并到当前 slot。
|
||||||
|
///
|
||||||
|
/// **注意**:本方法仅操作内存数据,不自动持久化。
|
||||||
|
/// 调用方需在 merge 后自行调用 `self.save(&store)` 将结果写入后端存储。
|
||||||
|
pub fn merge(&mut self, child: ContextSlot, strategy: MergeStrategy) -> Result<(), AgentError> {
|
||||||
|
// 防御性检查
|
||||||
|
if self.id == child.id {
|
||||||
|
return Err(AgentError::Config("不能将 slot 合并到自身".into()));
|
||||||
|
}
|
||||||
|
if self.session_id != child.session_id {
|
||||||
|
return Err(AgentError::Config("不能合并不同 session 的 slot".into()));
|
||||||
|
}
|
||||||
|
if matches!(self.config.mode, SlotMode::Readonly) {
|
||||||
|
return Err(AgentError::SlotReadonly("Readonly slot 不允许合并".into()));
|
||||||
|
}
|
||||||
|
|
||||||
|
tracing::debug!(
|
||||||
|
self_id = %self.id,
|
||||||
|
child_id = %child.id,
|
||||||
|
?strategy,
|
||||||
|
"ContextSlot::merge"
|
||||||
|
);
|
||||||
|
|
||||||
|
match strategy {
|
||||||
|
MergeStrategy::Append => {
|
||||||
|
let count = child.messages.len();
|
||||||
|
self.messages.extend(child.messages);
|
||||||
|
self.meta.message_count += count;
|
||||||
|
}
|
||||||
|
MergeStrategy::Replace => {
|
||||||
|
self.messages = child.messages;
|
||||||
|
self.meta.message_count = self.messages.len();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### AgentSession::derive_slot 重构
|
||||||
|
|
||||||
|
现有 `derive_slot`(session.rs:213-260)的手工复制代码改为调用 `parent.fork()`:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
pub async fn derive_slot(
|
||||||
|
&mut self,
|
||||||
|
id: impl Into<String>,
|
||||||
|
parent_id: &str,
|
||||||
|
strategy: DeriveStrategy,
|
||||||
|
) -> Result<(), AgentError> {
|
||||||
|
let slot_id = id.into();
|
||||||
|
if self.slots.contains_key(&slot_id) {
|
||||||
|
return Err(AgentError::SlotAlreadyExists(slot_id));
|
||||||
|
}
|
||||||
|
let parent = self
|
||||||
|
.slots
|
||||||
|
.get(parent_id)
|
||||||
|
.ok_or_else(|| AgentError::SlotNotFound(parent_id.to_string()))?;
|
||||||
|
let child = parent.fork(slot_id.clone(), strategy); // ← 用 fork
|
||||||
|
child.save(&*self.resolve_store()).await?;
|
||||||
|
self.slots.insert(slot_id, child);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
重复检查、查找父 slot 的代码不变;消息复制逻辑委托给 `fork()`。
|
||||||
|
|
||||||
|
#### 测试计划(新增 9 个)
|
||||||
|
|
||||||
|
| 测试名 | 验证点 |
|
||||||
|
|--------|--------|
|
||||||
|
| `fork_full_copies_messages` | fork Full 策略复制父 slot 全部消息 |
|
||||||
|
| `fork_focused_filters_messages` | fork Focused 策略按 config 过滤 |
|
||||||
|
| `fork_preserves_independence` | 父 slot 追加消息不影响子 slot |
|
||||||
|
| `fork_sets_derived_source` | 子 slot source 正确记录 parent_id |
|
||||||
|
| `merge_append_appends_messages` | Append 追加到父 slot 末尾,message_count 正确 |
|
||||||
|
| `merge_replace_replaces_messages` | Replace 替换父 slot 消息,message_count 正确 |
|
||||||
|
| `merge_self_rejected` | self-merge 返回 `Err` |
|
||||||
|
| `merge_readonly_rejected` | 合并到 Readonly slot 返回 `Err` |
|
||||||
|
| `merge_cross_session_rejected` | 跨 session 合并返回 `Err` |
|
||||||
|
|
||||||
|
### 3.3 Step 13.4 — ToolChoice 移入 tool.rs
|
||||||
|
|
||||||
|
#### 变更文件
|
||||||
|
|
||||||
|
| 文件 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `src/llm/types/request.rs` | 删除 `ToolChoice` 枚举 + serde impl(~28-99 行) |
|
||||||
|
| `src/llm/types/tool.rs` | 新增 `ToolChoice` 枚举 + serde impl(原样搬入) |
|
||||||
|
| `src/llm/types/mod.rs` | `pub use request::{..., ToolChoice}` → `pub use tool::ToolChoice` |
|
||||||
|
| `src/llm/types/request_v2.rs` | import 路径 `request::ToolChoice` → `tool::ToolChoice` |
|
||||||
|
|
||||||
|
**import 路径变化**:
|
||||||
|
|
||||||
|
| 当前 | 移动后 |
|
||||||
|
|------|--------|
|
||||||
|
| `crate::llm::types::request::ToolChoice` | `crate::llm::types::tool::ToolChoice` |
|
||||||
|
| `crate::llm::types::ToolChoice`(通过 re-export) | `crate::llm::types::ToolChoice`(通过 tool.rs re-export,保持不变) |
|
||||||
|
|
||||||
|
**验证**:`cargo build --all-targets` + `cargo test` + `cargo clippy`
|
||||||
|
|
||||||
|
### 3.4 Step 13.1 — request.rs 类型移入 openai.rs
|
||||||
|
|
||||||
|
#### 变更文件
|
||||||
|
|
||||||
|
| 文件 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `src/llm/types/request.rs` | **整文件删除**(187 行) |
|
||||||
|
| `src/llm/provider/openai.rs` | 新增 `StreamOptions`、`OpenaiTool`、`AudioParam`、`PredictionContent`、`UserLocation`、`Approximate`、`WebSearchOptions`、`OpenaiChatRequest` 等类型定义 |
|
||||||
|
| `src/llm/types/mod.rs` | 删除 `pub use request::{OpenaiChatRequest, OpenaiTool, StreamOptions}`;删除 `pub mod request;` |
|
||||||
|
| `src/llm/provider/openai.rs` import 调整 | 原 `use crate::llm::types::request::{...}` 改为从同级 `use super::super::types::...` 或直接使用本文件内类型 |
|
||||||
|
|
||||||
|
**注意**:`OpenaiTool` 引用 `OpenaiToolDefinition`(定义在 `tool.rs`),移入 `openai.rs` 后需通过 `crate::llm::types::tool::OpenaiToolDefinition` 引用。`OpenaiChatRequest.messages` 字段引用 `OpenaiChatMessage`(定义在 `openai_message.rs`),路径不变。
|
||||||
|
|
||||||
|
**设计决策**:搬入 `openai.rs` 后的类型可见性可降级为 `pub(crate)`。它们是与 OpenAI wire-format 绑定的内部序列化类型,公共 API 消费者不应直接接触。
|
||||||
|
|
||||||
|
**验证**:`cargo build --all-targets` + `cargo test` + `cargo clippy`
|
||||||
|
|
||||||
|
### 3.5 Step 13.2 — response.rs 类型移入 openai.rs
|
||||||
|
|
||||||
|
#### 变更文件
|
||||||
|
|
||||||
|
| 文件 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `src/llm/types/response.rs` | **整文件删除**(177 行) |
|
||||||
|
| `src/llm/provider/openai.rs` | 新增 `TokenLogprob`、`TopLogprob`、`Logprobs`、`URLCitation`、`Annotation`、`OpenaiAudio`、`Choice`、`OpenaiChatResponse`、`Delta`、`ChunkChoice`、`OpenaiChatChunk` + `From<OpenaiChatMessage> for Delta` + `From<OpenaiChatResponse> for OpenaiChatChunk` |
|
||||||
|
| `src/llm/types/mod.rs` | 删除 `pub use response::{...}`;删除 `pub mod response;` |
|
||||||
|
| `src/llm/stream.rs:26` | 将 `use crate::llm::types::{OpenaiChatChunk, OpenaiToolCall}` 中的 `OpenaiChatChunk` 路径改为 `crate::llm::provider::openai::OpenaiChatChunk`(`OpenaiToolCall` 保持从 `tool.rs`) |
|
||||||
|
|
||||||
|
**验证**:`cargo build --all-targets` + `cargo test` + `cargo clippy`
|
||||||
|
|
||||||
|
### 3.6 Step 13.3 — 旧文件清理 + ChatResponse 删除
|
||||||
|
|
||||||
|
#### 13.3a — 删除 `old_stream.rs`
|
||||||
|
|
||||||
|
> **前置验证**:实施前执行 `grep -rn 'parse_chunk_stream\|map_legacy_to_ir\|LegacyToIrEventStream\|ChunkToLegacyEventStream' src/` 确认零外部调用方,记录结果到实施 commit。
|
||||||
|
|
||||||
|
| 文件 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `src/llm/types/old_stream.rs` | **整文件删除**(45 行,`LegacyStreamEvent`) |
|
||||||
|
| `src/llm/types/mod.rs` | 删除 `pub mod old_stream;` |
|
||||||
|
| `src/llm/stream.rs` | 删除 `use crate::llm::types::old_stream::LegacyStreamEvent`;删除 `parse_chunk_stream`、`parse_chunk_stream_legacy`、`ChunkToLegacyEventStream`、`LegacyToIrEventStream`、`map_legacy_to_ir`、`empty_message_response`(~160 行死代码) |
|
||||||
|
|
||||||
|
**stream.rs 最终形态**:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
//! 流式事件系统 —— 重导出 StreamEvent 供向后兼容。
|
||||||
|
pub use crate::llm::types::response_v2::StreamEvent;
|
||||||
|
```
|
||||||
|
|
||||||
|
**为什么不全删 stream.rs**:`cycle.rs` 和 `session.rs` 的 `use crate::llm::stream::StreamEvent` 路径保持不变。全删 + 改所有 import 路径的改动量 > 收益。保留 1 行重导出就够。
|
||||||
|
|
||||||
|
#### 13.3b — 删除 `ChatResponse`
|
||||||
|
|
||||||
|
| 文件 | 变更 |
|
||||||
|
|------|------|
|
||||||
|
| `src/llm/types/mod.rs` | 删除 `ChatResponse` 结构体定义 + 两个 `#[allow(deprecated)]` `From` impl(`From<OpenaiChatResponse> for ChatResponse` 和 `From<ChatResponse> for OpenaiChatChunk`) |
|
||||||
|
|
||||||
|
`ChatResponse` 自 v0.1.0 起标记 `#[deprecated]`,v0.2.0-rc.1 阶段直接删除即可。删除前运行 `cargo doc --no-deps 2>&1 | grep -i 'ChatResponse'` 确认零文档引用。
|
||||||
|
|
||||||
|
**验证**:`cargo build --all-targets` + `cargo test` + `cargo clippy` + `cargo doc --no-deps`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. 实现计划
|
||||||
|
|
||||||
|
### 4.1 实施顺序总览
|
||||||
|
|
||||||
|
```
|
||||||
|
Step 13.5 ──→ Step 13.4 ──→ Step 13.1 ──→ Step 13.2 ──→ Step 13.3
|
||||||
|
(fork/merge) (ToolChoice) (request) (response) (cleanup)
|
||||||
|
│ │ │ │ │
|
||||||
|
▼ ▼ ▼ ▼ ▼
|
||||||
|
+60 行净增 -0 净增 -0 净增 -0 净增 -260 删除
|
||||||
|
+9 个测试 import 路径 纯类型搬移 纯类型搬移 +1 行重导出
|
||||||
|
变更
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.2 各 Step 文件变更清单
|
||||||
|
|
||||||
|
#### Step 13.5 — ContextSlot fork/merge
|
||||||
|
|
||||||
|
| 操作 | 文件 | 变更说明 |
|
||||||
|
|------|------|---------|
|
||||||
|
| 新增 | `src/agent/context.rs` | `MergeStrategy` 枚举 + `ContextSlot::fork()` + `ContextSlot::merge()` |
|
||||||
|
| 重构 | `src/agent/session.rs` | `derive_slot` 改为调用 `parent.fork()` |
|
||||||
|
| 新增 | 内联测试 | 9 个新测试(fork, merge, 边界) |
|
||||||
|
|
||||||
|
#### Step 13.4 — ToolChoice 移动
|
||||||
|
|
||||||
|
| 操作 | 文件 | 变更说明 |
|
||||||
|
|------|------|---------|
|
||||||
|
| 删除 | `src/llm/types/request.rs` | 移除 `ToolChoice` 枚举 + serde impl |
|
||||||
|
| 新增 | `src/llm/types/tool.rs` | 增加 `ToolChoice` 枚举 + serde impl |
|
||||||
|
| 修改 | `src/llm/types/mod.rs` | 更新 re-export 路径 |
|
||||||
|
| 修改 | `src/llm/types/request_v2.rs` | 更新 import 路径 |
|
||||||
|
|
||||||
|
#### Step 13.1 — request 类型搬移
|
||||||
|
|
||||||
|
| 操作 | 文件 | 变更说明 |
|
||||||
|
|------|------|---------|
|
||||||
|
| 删除 | `src/llm/types/request.rs` | 整文件删除(187 行) |
|
||||||
|
| 新增 | `src/llm/provider/openai.rs` | 增加所有 OpenAI wire-format 类型 |
|
||||||
|
| 修改 | `src/llm/types/mod.rs` | 删除 re-export + mod 声明 |
|
||||||
|
|
||||||
|
#### Step 13.2 — response 类型搬移
|
||||||
|
|
||||||
|
| 操作 | 文件 | 变更说明 |
|
||||||
|
|------|------|---------|
|
||||||
|
| 删除 | `src/llm/types/response.rs` | 整文件删除(177 行) |
|
||||||
|
| 新增 | `src/llm/provider/openai.rs` | 增加所有 OpenAI wire-format 类型 + From impl |
|
||||||
|
| 修改 | `src/llm/types/mod.rs` | 删除 re-export + mod 声明 |
|
||||||
|
| 修改 | `src/llm/stream.rs` | 更新 `OpenaiChatChunk` import 路径 |
|
||||||
|
|
||||||
|
#### Step 13.3 — 旧文件清理
|
||||||
|
|
||||||
|
| 操作 | 文件 | 变更说明 |
|
||||||
|
|------|------|---------|
|
||||||
|
| 删除 | `src/llm/types/old_stream.rs` | 整文件删除(45 行) |
|
||||||
|
| 修改 | `src/llm/types/mod.rs` | 删除 `pub mod old_stream;` + 删除 `ChatResponse` 结构体 + `From` impl |
|
||||||
|
| 修改 | `src/llm/stream.rs` | 删除所有死代码,仅保留 `pub use` 重导出 |
|
||||||
|
|
||||||
|
### 4.3 回滚策略
|
||||||
|
|
||||||
|
所有 Step 通过 git commit 管理,回退时 `git revert <commit>` 即可。每个 Step 独立编译,回滚不会级联依赖。若 Step 13.3(`ChatResponse` 删除)导致外部编译失败,单独 revert 该 commit 即可恢复 `ChatResponse` + `old_stream.rs`。
|
||||||
|
|
||||||
|
### 4.4 CHANGELOG 条目
|
||||||
|
|
||||||
|
```markdown
|
||||||
|
## [0.3.0] - 未发布
|
||||||
|
|
||||||
|
### Breaking Changes
|
||||||
|
|
||||||
|
**类型路径变更(0.3.0):**
|
||||||
|
- `agcore::llm::types::request::ToolChoice` → `agcore::llm::types::tool::ToolChoice`(公共 re-export 路径 `agcore::llm::types::ToolChoice` 保持不变)
|
||||||
|
- `agcore::llm::types::request::StreamOptions` → `agcore::llm::provider::openai::StreamOptions`
|
||||||
|
- `agcore::llm::types::request::OpenaiChatRequest` → `agcore::llm::provider::openai::OpenaiChatRequest`
|
||||||
|
- `agcore::llm::types::response::OpenaiChatResponse` → `agcore::llm::provider::openai::OpenaiChatResponse`
|
||||||
|
- `agcore::llm::types::response::OpenaiChatChunk` → `agcore::llm::provider::openai::OpenaiChatChunk`
|
||||||
|
- 其余 `request.rs`/`response.rs` 中的 wire-format 类型(`OpenaiTool`、`AudioParam`、`Choice`、`Delta` 等)同步移入 `agcore::llm::provider::openai` 模块
|
||||||
|
|
||||||
|
**类型删除:**
|
||||||
|
- `agcore::llm::types::ChatResponse` 已删除(自 v0.1.0 标记 `#[deprecated]`,请改用 `MessageResponse`)
|
||||||
|
- `agcore::llm::types::old_stream::LegacyStreamEvent` 已删除(内部死代码)
|
||||||
|
|
||||||
|
### Features
|
||||||
|
- `ContextSlot::fork(child_id, strategy)` — 从父槽派生独立的子槽(数据层操作)
|
||||||
|
- `ContextSlot::merge(child, strategy)` — 将子槽消息合并回父槽(支持 Append/Replace)
|
||||||
|
- `MergeStrategy` 枚举(`#[non_exhaustive]`,Phase 16 可扩展 Summarize)
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. 风险评估
|
||||||
|
|
||||||
|
| 风险 | 影响 | 概率 | 缓解措施 |
|
||||||
|
|------|------|------|---------|
|
||||||
|
| `ChatResponse` 被外部 crate 引用 | 编译 break | 中 — `#[deprecated]` 仅产生编译警告,外部 crate 可能通过 `#[allow(deprecated)]` 静默依赖 | CHANGELOG 明确标注语义版本(0.3.0)和迁移指引;Step 13.3 验收加入 `cargo doc --no-deps \| grep ChatResponse` 确认零引用 |
|
||||||
|
| `StreamOptions` 等 wire-format 类型路径变更影响直接引用消费者 | 编译 break | 低(v0.2.0-rc.1,极少外部消费者使用内部类型) | CHANGELOG 完整列出所有路径变更;编译错误立即可发现 |
|
||||||
|
| `parse_chunk_stream` 有隐藏调用方 | 编译 break | 极低(实施前执行 `grep -rn 'parse_chunk_stream\|map_legacy_to_ir\|LegacyToIrEventStream' src/` 前置验证) | Step 13.3 前运行 grep 验证并记录结果;`cargo build --all-targets` 可 100% 捕获 |
|
||||||
|
| `#[allow(deprecated)]` 遗漏 | clippy 警告 | 低 | `cargo clippy --all-targets -- -D warnings` 验证 |
|
||||||
|
| Step 顺序错误导致编译中间态 | 开发者体验差 | 中 | 严格按 13.5→13.4→13.1→13.2→13.3 执行;每步 `cargo build` 验证 |
|
||||||
|
| `stream.rs` 简化后 import 断链 | 编译 break | 极低 | 保留 `pub use` 重导出路径,`cycle.rs`/`session.rs` import 不变 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. 验收标准
|
||||||
|
|
||||||
|
### M9 里程碑(Phase 13 完成条件)
|
||||||
|
|
||||||
|
| # | 条件 | 验证方法 |
|
||||||
|
|---|------|---------|
|
||||||
|
| 1 | `request.rs`、`response.rs`、`old_stream.rs` 三个旧文件不存在 | `ls src/llm/types/` 确认 |
|
||||||
|
| 2 | `ChatResponse` 结构体不存在 | 全局搜索 `ChatResponse` 仅保留 `openai.rs` 中 `OpenaiChatResponse` 引用 |
|
||||||
|
| 3 | `ToolChoice` 在 `tool.rs` 中定义,公共路径 `agcore::llm::types::ToolChoice` 保持不变 | `cargo doc --no-deps` 确认类型文档 |
|
||||||
|
| 4 | `OpenaiChatRequest`/`Response`/`Chunk` 在 `provider/openai.rs` 中定义 | 编译通过 |
|
||||||
|
| 5 | `ContextSlot::fork()` 单元测试通过(P0 条件全部满足) | `cargo test` |
|
||||||
|
| 6 | `ContextSlot::merge()` 单元测试通过(P0 条件全部满足) | `cargo test` |
|
||||||
|
| 7 | `stream.rs` 只保留 `pub use` 重导出 | 文件内容确认 |
|
||||||
|
| 8 | `cargo build --all-targets` 编译通过 | 编译验证 |
|
||||||
|
| 9 | `cargo test --all-targets` 全绿(预期 283~285 测试) | 测试验证 |
|
||||||
|
| 10 | `cargo clippy --all-targets -- -D warnings` 0 警告 | clippy 验证 |
|
||||||
|
| 11 | CHANGELOG 包含 Phase 13 的 Breaking Changes 和 Features 条目 | 文件确认 |
|
||||||
|
|
||||||
|
### fork/merge 详细验收 P0 项
|
||||||
|
|
||||||
|
**fork 的 5 项 P0 条件:**
|
||||||
|
|
||||||
|
| # | 条件 | 优先级 |
|
||||||
|
|---|------|--------|
|
||||||
|
| 1 | `fork("child", Full)` 创建新 slot,消息在 fork 时刻 == 父 slot | P0 |
|
||||||
|
| 2 | 子 slot 获得独立消息列表——父 slot 后续追加不影响子 slot | P0 |
|
||||||
|
| 3 | 子 slot 的 source 标记为 `Derived { parent_id, strategy }` | P0 |
|
||||||
|
| 4 | 子 slot 可独立持久化(fork + save + load roundtrip) | P0 |
|
||||||
|
| 5 | fork 不允许重复 id(返回 `SlotAlreadyExists`)(由 `derive_slot` 编排层保证) | P0 |
|
||||||
|
|
||||||
|
**merge 的 5 项 P0 条件:**
|
||||||
|
|
||||||
|
| # | 条件 | 优先级 |
|
||||||
|
|---|------|--------|
|
||||||
|
| 1 | `parent.merge(child, Append)` 子消息追加到父末尾 | P0 |
|
||||||
|
| 2 | `parent.merge(child, Replace)` 子消息替换父全量消息 | P0 |
|
||||||
|
| 3 | merge 后父 slot 的 `meta.message_count` 正确更新 | P0 |
|
||||||
|
| 4 | merge 不允许合并到 Readonly 目标 slot | P0 |
|
||||||
|
| 5 | merge 不允许 self-merge(child.id == parent.id) | P0 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 参考来源
|
||||||
|
|
||||||
|
- Roadmap:`docs/roadmap.md` §Phase 13
|
||||||
|
- ContextSlot 设计:`docs/17-phase10-contextslot.md`
|
||||||
|
- 旧 StreamEvent 设计:`src/llm/stream.rs` 文件注释
|
||||||
|
- 当前代码库:`src/llm/types/request.rs`、`src/llm/types/response.rs`、`src/llm/types/old_stream.rs`、`src/llm/types/mod.rs`、`src/llm/provider/openai.rs`、`src/agent/context.rs`、`src/agent/session.rs`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. 实施计划
|
||||||
|
|
||||||
|
### 全局说明
|
||||||
|
|
||||||
|
**commit 策略**:每个 Step 一个独立 commit。commit message 格式:
|
||||||
|
```
|
||||||
|
<type>(<scope>): <中文描述>
|
||||||
|
```
|
||||||
|
- Step 13.5 → `feat(agent): 实现 ContextSlot fork/merge`
|
||||||
|
- Step 13.4 → `refactor(types): ToolChoice 移入 tool.rs`
|
||||||
|
- Step 13.1 → `refactor(types): request.rs 类型移入 provider/openai.rs`
|
||||||
|
- Step 13.2 → `refactor(types): response.rs 类型移入 provider/openai.rs`
|
||||||
|
- Step 13.3 → `refactor(types): 删除旧类型文件和 ChatResponse`
|
||||||
|
|
||||||
|
**验证命令(每步通用)**:
|
||||||
|
```bash
|
||||||
|
cargo build --all-targets && cargo test && cargo clippy --all-targets -- -D warnings
|
||||||
|
```
|
||||||
|
|
||||||
|
**预计测试数量变化**:
|
||||||
|
- 当前基线:277 测试(每个 Step 开始时 `cargo test` 确认)
|
||||||
|
- Step 13.5 后:286(+9)
|
||||||
|
- Step 13.4-13.2 后:286(无变化)
|
||||||
|
- Step 13.3 后:285(-1,`ChatResponse` 的 `From` impl 无测试直接引用,删除后仅 `types/mod.rs` 中的 `deprecated` 注释行减少,不影响测试计数。实施前执行 `grep -rn 'ChatResponse' src/ --include='*test*' --include='*tests*'` 确认零测试引用)
|
||||||
|
- 最终范围:285 测试
|
||||||
|
|
||||||
|
### Step 13.5 — ContextSlot fork/merge
|
||||||
|
|
||||||
|
**前置依赖**:无(纯新增,不依赖前序 Step)
|
||||||
|
|
||||||
|
**任务描述**:在 `agent/context.rs` 中新增 `MergeStrategy` 枚举、`ContextSlot::fork()` 方法和 `ContextSlot::merge()` 方法;重构 `agent/session.rs` 中的 `derive_slot` 改为调用 `parent.fork()`;新增 9 个内联测试覆盖 fork/merge 的 happy path 和 error path。
|
||||||
|
|
||||||
|
**涉及文件**:
|
||||||
|
- `src/agent/context.rs` — 新增枚举和方法
|
||||||
|
- `src/agent/session.rs` — 重构 derive_slot
|
||||||
|
- `src/agent.rs` — 追加 `MergeStrategy` re-export
|
||||||
|
|
||||||
|
**具体操作**:
|
||||||
|
1. 在 `context.rs` 中新增 `MergeStrategy` 枚举(Append / Replace,`#[non_exhaustive]`)
|
||||||
|
2. 在 `context.rs` 中 `impl ContextSlot` 块内新增 `fork(&self, child_id: String, strategy: DeriveStrategy) -> ContextSlot` 方法
|
||||||
|
3. 在 `context.rs` 中 `impl ContextSlot` 块内新增 `merge(&mut self, child: ContextSlot, strategy: MergeStrategy) -> Result<(), AgentError>` 方法(含 self-merge/cross-session/Readonly 三项防御检查 + `tracing::debug!` 日志)
|
||||||
|
4. 在 `session.rs` 的 `derive_slot` 方法中将手工消息复制代码替换为 `parent.fork(slot_id, strategy)`
|
||||||
|
5. 在 `agent.rs` 的 `pub use context::{...}` 列表中追加 `MergeStrategy`
|
||||||
|
6. 在 `context.rs` 的 `#[cfg(test)] mod tests` 中新增 9 个测试用例
|
||||||
|
|
||||||
|
**注意**:重构后 `derive_slot` 的子 slot `budget` 从 `ContextBudget::default()` 变为继承父 slot,`compact` 从 `true` 变为继承父 slot。由于 `ContextBudget` 在 v0.2 无消费逻辑且父 slot 的 `compact` 默认也为 `true`,此变化无实际影响。验收条件中"行为不变"指对外功能行为不变(slot 消息内容、血缘关系不变)。
|
||||||
|
|
||||||
|
**预估工作量**:M(1-4h)
|
||||||
|
|
||||||
|
**风险等级**:低(纯新增,不修改已有逻辑路径)
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `MergeStrategy` 枚举存在,`Append` 和 `Replace` 两个变体可用,且通过 `agcore::agent::MergeStrategy` 路径可访问
|
||||||
|
- `ContextSlot::fork` 返回的 child 在 fork 时刻消息等于父 slot
|
||||||
|
- fork Focused 策略按 `FocusedConfig` 过滤消息
|
||||||
|
- 父 slot 后续追加消息不影响子 slot
|
||||||
|
- 子 slot 的 source 正确记录 `Derived { parent_id, strategy }`
|
||||||
|
- `parent.merge(child, Append)` 追加到父末尾,message_count 正确
|
||||||
|
- `parent.merge(child, Replace)` 替换父全量消息,message_count 正确
|
||||||
|
- self-merge 返回 `Err(AgentError::Config)`
|
||||||
|
- merge 到 Readonly slot 返回 `Err(AgentError::SlotReadonly)`
|
||||||
|
- 跨 session merge 返回 `Err(AgentError::Config)`
|
||||||
|
- `derive_slot` 对外行为不变(slot 消息内容、血缘关系、持久化行为均不变;内部 budget/compact 继承差异无实际影响),测试全绿
|
||||||
|
- `cargo doc --no-deps` 无 warning(验证新增公开 API 的文档注释完整)
|
||||||
|
|
||||||
|
**回退方式**:`git revert` 该 commit
|
||||||
|
|
||||||
|
### Step 13.4 — ToolChoice 移入 tool.rs
|
||||||
|
|
||||||
|
**前置依赖**:Step 13.5(顺序约束:必须早于 Step 13.1——若 Step 13.1 先执行会将 `ToolChoice` 与 `request.rs` 一同删除,导致本 Step 无可搬移的源)
|
||||||
|
|
||||||
|
**任务描述**:将 `ToolChoice` 枚举及其 serde 实现从 `types/request.rs` 搬移到 `types/tool.rs`,更新所有 import/path 引用。公共 re-export 路径 `agcore::llm::types::ToolChoice` 保持不变。
|
||||||
|
|
||||||
|
**涉及文件**:
|
||||||
|
- `src/llm/types/request.rs` — 删除 ToolChoice(~28-99 行)
|
||||||
|
- `src/llm/types/tool.rs` — 新增 ToolChoice 枚举 + serde impl
|
||||||
|
- `src/llm/types/mod.rs` — re-export 路径从 `request` 改为 `tool`
|
||||||
|
- `src/llm/types/request_v2.rs` — import 路径从 `request::` 改为 `tool::`
|
||||||
|
|
||||||
|
**具体操作**:
|
||||||
|
1. 从 `request.rs` 复制 `ToolChoice` 枚举 + `Serialize`/`Deserialize` impl 到 `tool.rs`
|
||||||
|
2. 从 `request.rs` 中删除 `ToolChoice` 定义
|
||||||
|
3. 在 `mod.rs` 中将 `pub use request::{..., ToolChoice}` 改为 `pub use tool::ToolChoice`
|
||||||
|
4. 在 `request_v2.rs` 中将 `use crate::llm::types::request::ToolChoice` 改为 `use crate::llm::types::tool::ToolChoice`
|
||||||
|
5. 验证 `cycle.rs` 的 `use crate::llm::types::ToolChoice`(通过 re-export)路径不变
|
||||||
|
|
||||||
|
**预估工作量**:S(<1h)
|
||||||
|
|
||||||
|
**风险等级**:低(有限的 import 路径变更,编译立即可发现)
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `ToolChoice` 在 `tool.rs` 中定义
|
||||||
|
- `pub use tool::ToolChoice` 在 `mod.rs` 中
|
||||||
|
- `request_v2.rs` 编译通过
|
||||||
|
- `cycle.rs` 路径不变
|
||||||
|
- `cargo build --all-targets` + `cargo test` + `cargo clippy` 全绿
|
||||||
|
|
||||||
|
**回退方式**:`git revert` 该 commit
|
||||||
|
|
||||||
|
### Step 13.1 — request.rs 类型移入 openai.rs
|
||||||
|
|
||||||
|
**前置依赖**:Step 13.4(ToolChoice 已移走,request.rs 剩余内容全是 OpenAI wire-format 专有类型)
|
||||||
|
|
||||||
|
**任务描述**:删除 `types/request.rs` 整文件,将所有剩余类型(`OpenaiChatRequest`、`StreamOptions`、`OpenaiTool`、`AudioParam`、`PredictionContent`、`UserLocation`、`Approximate`、`WebSearchOptions`)搬入 `provider/openai.rs`,更新 `mod.rs` re-export。
|
||||||
|
|
||||||
|
**涉及文件**:
|
||||||
|
- `src/llm/types/request.rs` — 整文件删除
|
||||||
|
- `src/llm/provider/openai.rs` — 新增所有类型定义
|
||||||
|
- `src/llm/types/mod.rs` — 删除 re-export + mod 声明
|
||||||
|
|
||||||
|
**具体操作**:
|
||||||
|
1. 从 `request.rs` 复制所有剩余类型定义到 `openai.rs`,可见性设为 `pub(crate)`
|
||||||
|
2. `OpenaiTool` 内引用 `OpenaiToolDefinition`(定义在 `tool.rs`),路径改为 `crate::llm::types::tool::OpenaiToolDefinition`
|
||||||
|
3. 删除 `openai.rs` 中原 `use crate::llm::types::request::{...}` import
|
||||||
|
4. 从 `mod.rs` 删除 `pub use request::{OpenaiChatRequest, OpenaiTool, StreamOptions}` 和 `pub mod request;`
|
||||||
|
5. 删除 `types/request.rs` 文件
|
||||||
|
|
||||||
|
**预估工作量**:M(1-4h)
|
||||||
|
|
||||||
|
**风险等级**:低(纯搬移 + 删除,文件内无逻辑变更)
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `request.rs` 文件不存在
|
||||||
|
- `OpenaiChatRequest` 等类型在 `openai.rs` 中定义,编译通过
|
||||||
|
- `OpenaiTool` 通过 `crate::llm::types::tool::OpenaiToolDefinition` 正确引用
|
||||||
|
- `cargo build --all-targets` + `cargo test` + `cargo clippy` 全绿
|
||||||
|
|
||||||
|
**回退方式**:`git revert` 该 commit。若 Step 13.2 也已提交,单独 revert 本 Step 可能因 `provider/openai.rs` 并发修改产生合并冲突。安全回退顺序为逆序:先 revert 13.2,再 revert 13.1。
|
||||||
|
|
||||||
|
### Step 13.2 — response.rs 类型移入 openai.rs
|
||||||
|
|
||||||
|
**前置依赖**:无(与 Step 13.1 共享 `provider/openai.rs` 和 `types/mod.rs`,但本 Step 仅追加类型定义,无覆盖操作;建议在 13.1 之后顺序执行以避免并行时的合并冲突)
|
||||||
|
|
||||||
|
**任务描述**:删除 `types/response.rs` 整文件,将所有类型(`OpenaiChatResponse`、`OpenaiChatChunk`、`Choice`、`Delta`、`ChunkChoice` 等 + 两个 `From` impl)搬入 `provider/openai.rs`,更新 `mod.rs` 和 `stream.rs` 的 import 路径。
|
||||||
|
|
||||||
|
**涉及文件**:
|
||||||
|
- `src/llm/types/response.rs` — 整文件删除
|
||||||
|
- `src/llm/provider/openai.rs` — 新增所有类型定义 + From impl
|
||||||
|
- `src/llm/types/mod.rs` — 删除 re-export + mod 声明
|
||||||
|
- `src/llm/stream.rs` — `OpenaiChatChunk` import 路径改为 `provider::openai`
|
||||||
|
|
||||||
|
**具体操作**:
|
||||||
|
1. 从 `response.rs` 复制所有类型定义(含 `From` impl)到 `openai.rs`,可见性设为 `pub(crate)`
|
||||||
|
2. 删除 `openai.rs` 中原 `use crate::llm::types::response::{...}` import
|
||||||
|
3. 从 `mod.rs` 删除 `pub use response::{...}` 和 `pub mod response;`
|
||||||
|
4. 在 `stream.rs:26` 将 `OpenaiChatChunk` 的 import 路径改为 `crate::llm::provider::openai::OpenaiChatChunk`(`OpenaiToolCall` 路径不变)
|
||||||
|
5. 删除 `types/response.rs` 文件
|
||||||
|
|
||||||
|
**预估工作量**:M(1-4h)
|
||||||
|
|
||||||
|
**风险等级**:低(与 Step 13.1 模式完全相同)
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `response.rs` 文件不存在
|
||||||
|
- `OpenaiChatResponse`/`Chunk` 等类型在 `openai.rs` 中定义,编译通过
|
||||||
|
- `stream.rs` import 路径正确
|
||||||
|
- `cargo build --all-targets` + `cargo test` + `cargo clippy` 全绿
|
||||||
|
|
||||||
|
**回退方式**:`git revert` 该 commit。若 Step 13.1 和本 Step 均已提交,安全回退顺序为逆序:先 revert 本 Step,再 revert 13.1。
|
||||||
|
|
||||||
|
### Step 13.3 — 旧文件清理 + ChatResponse 删除
|
||||||
|
|
||||||
|
**前置依赖**:Step 13.1(`request.rs` 已删)、Step 13.2(`response.rs` 已删)
|
||||||
|
|
||||||
|
**任务描述**:删除 `old_stream.rs` 和 `ChatResponse`,简化 `stream.rs` 为仅保留 `pub use` 重导出。这是 Phase 13 技术风险最高的 Step。
|
||||||
|
|
||||||
|
**涉及文件**:
|
||||||
|
- `src/llm/types/old_stream.rs` — 整文件删除
|
||||||
|
- `src/llm/types/mod.rs` — 删除 `pub mod old_stream;` + 删除 `ChatResponse` 结构体和两个 `From` impl
|
||||||
|
- `src/llm/stream.rs` — 删除死代码(约 160 行),仅保留 `pub use` 重导出
|
||||||
|
|
||||||
|
**具体操作**:
|
||||||
|
1. **前置验证 A**:执行 `grep -rn 'parse_chunk_stream\|map_legacy_to_ir\|LegacyToIrEventStream\|ChunkToLegacyEventStream' src/` 确认零外部调用方,记录结果到 commit message
|
||||||
|
2. **前置验证 B**:执行 `cargo doc --no-deps 2>&1 | grep -i 'ChatResponse'` 确认零文档引用,记录结果
|
||||||
|
3. 从 `mod.rs` 删除 `pub mod old_stream;`
|
||||||
|
4. 从 `mod.rs` 删除 `ChatResponse` 结构体定义 + `#[allow(deprecated)]` `From<OpenaiChatResponse> for ChatResponse` + `From<ChatResponse> for OpenaiChatChunk`
|
||||||
|
5. 删除 `old_stream.rs` 文件
|
||||||
|
6. 从 `stream.rs` 删除:`use crate::llm::types::old_stream::LegacyStreamEvent`、`parse_chunk_stream`、`parse_chunk_stream_legacy`、`ChunkToLegacyEventStream`、`LegacyToIrEventStream`、`map_legacy_to_ir`、`empty_message_response`
|
||||||
|
7. `stream.rs` 最终只保留 module doc comment + `pub use crate::llm::types::response_v2::StreamEvent;`
|
||||||
|
8. 检查 `cycle.rs:88` 的 `#[allow(deprecated)]` 属性是否仍与 `ChatResponse` 相关——若不相关则无需改动;若因 `ChatResponse` 删除而变脏,清理该属性
|
||||||
|
|
||||||
|
**预估工作量**:S(<1h,cleanup)+ M(需验证过程)
|
||||||
|
|
||||||
|
**风险等级**:中(`ChatResponse` 删除是 Breaking Change,外部可能静默依赖)
|
||||||
|
|
||||||
|
**验收条件**:
|
||||||
|
- `old_stream.rs` 文件不存在
|
||||||
|
- `ChatResponse` 结构体不存在(全局搜索仅保留 `OpenaiChatResponse` 引用)
|
||||||
|
- `stream.rs` 只保留 `pub use` 重导出
|
||||||
|
- `cargo build --all-targets` 编译通过
|
||||||
|
- `cargo test --all-targets` 全绿(预期 285 测试)
|
||||||
|
- `cargo clippy --all-targets -- -D warnings` 0 警告
|
||||||
|
- `cargo doc --no-deps` 无 warning
|
||||||
|
|
||||||
|
**回退方式**:`git revert` 该 commit(单独 revert 即可恢复 `ChatResponse` + `old_stream.rs`)
|
||||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,183 @@
|
|||||||
|
# LangChain & LangGraph 功能调研笔记
|
||||||
|
|
||||||
|
> 调研时间:2026-07-06
|
||||||
|
> 两者关系:同一公司(LangChain Inc.)维护的堆栈上下两层,不是竞品
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 两者关系
|
||||||
|
|
||||||
|
```
|
||||||
|
┌──────────────────────────────────────────┐
|
||||||
|
│ LangChain (v1.0 GA) │ ← 高层框架:模型抽象、工具、提示词、600+集成
|
||||||
|
│ create_agent / LCEL / 组件库 │
|
||||||
|
├──────────────────────────────────────────┤
|
||||||
|
│ LangGraph (v1.0 GA) │ ← 底层运行时:有向图执行引擎
|
||||||
|
│ StateGraph / Checkpointing / HITL │
|
||||||
|
├──────────────────────────────────────────┤
|
||||||
|
│ LangSmith (可观测性) │
|
||||||
|
└──────────────────────────────────────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
2025年10月22日同时达到 v1.0 GA,官方分工:
|
||||||
|
|
||||||
|
> **LangChain** = agent framework:abstractions and integrations for models, tools, and agent loops.
|
||||||
|
> **LangGraph** = orchestration runtime:durable execution, streaming, human-in-the-loop, and persistence.
|
||||||
|
|
||||||
|
LangChain v1.0 的 `create_agent` 内部已运行在 LangGraph 引擎上。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## LangChain v1.0
|
||||||
|
|
||||||
|
### 定位
|
||||||
|
高层应用框架,提供 agent 所需的**组件抽象**和**集成生态**。
|
||||||
|
|
||||||
|
### 精简后的核心模块
|
||||||
|
|
||||||
|
| 模块 | 功能 |
|
||||||
|
|------|------|
|
||||||
|
| `langchain.agents` | `create_agent`, `AgentState`(取代旧 AgentExecutor) |
|
||||||
|
| `langchain.chat_models` | `init_chat_model`, `BaseChatModel`(统一模型初始化) |
|
||||||
|
| `langchain.tools` | `@tool`, `BaseTool` |
|
||||||
|
| `langchain.messages` | 消息类型、内容块、`trim_messages` |
|
||||||
|
| `langchain.embeddings` | `init_embeddings`, `Embeddings` |
|
||||||
|
|
||||||
|
旧组件(`LLMChain`、`ConversationChain` 等)移入 `langchain-classic`。
|
||||||
|
|
||||||
|
### 七大组件类别
|
||||||
|
|
||||||
|
| 类别 | 关键组件 |
|
||||||
|
|------|----------|
|
||||||
|
| **Models** | Chat models, LLMs, Embeddings — 统一接口跨 provider 切换 |
|
||||||
|
| **Tools** | 600+ provider 集成:API、数据库、搜索引擎等 |
|
||||||
|
| **Agents** | `create_agent`, ReAct agents, Tool-calling agents |
|
||||||
|
| **Memory** | 消息历史、自定义状态 |
|
||||||
|
| **Retrievers** | 向量检索器、网络检索器 |
|
||||||
|
| **Document** | 加载器、分割器、转换器 |
|
||||||
|
| **Vector Stores** | Chroma, Pinecone, FAISS 等集成 |
|
||||||
|
|
||||||
|
### v1.0 关键新特性
|
||||||
|
|
||||||
|
**1. Middleware 中间件系统** — `create_agent` 的钩子系统:
|
||||||
|
- `before_model` — 模型调用前注入/修改
|
||||||
|
- `after_model` — 模型调用后验证/后处理
|
||||||
|
- `wrap_tool_call` — 拦截工具调用错误
|
||||||
|
|
||||||
|
**2. Standard Message Content** — 跨 provider 标准化消息内容格式:
|
||||||
|
- 推理/思维链、引用、多模态(图片/音视频/文档)
|
||||||
|
- 工具调用、provider 特有工具(web search, code execution)
|
||||||
|
- 通过 `.content_blocks` 属性访问,向后兼容
|
||||||
|
|
||||||
|
**3. `create_agent`** — 取代旧 AgentExecutor,内部运行在 LangGraph 运行时上
|
||||||
|
|
||||||
|
### 成熟度
|
||||||
|
|
||||||
|
| 维度 | 状态 |
|
||||||
|
|------|------|
|
||||||
|
| 版本 | v1.0 GA(2025-10) |
|
||||||
|
| 稳定性 | 稳定,agent 层经重构后已稳定 |
|
||||||
|
| 生产证明 | Replit, Clay, Rippling, Cloudflare, Workday |
|
||||||
|
| 支持 | LTS-style support track |
|
||||||
|
| 适用场景 | RAG、信息提取、单 agent 助手、快速原型 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## LangGraph v1.0
|
||||||
|
|
||||||
|
### 定位
|
||||||
|
底层编排运行时,专为**有状态、长时间运行、多步骤**工作流设计。
|
||||||
|
|
||||||
|
### 核心抽象链
|
||||||
|
|
||||||
|
```
|
||||||
|
StateGraph → Nodes (纯 Python 函数) → Edges (路由逻辑)
|
||||||
|
↓
|
||||||
|
Shared State (TypedDict / Pydantic)
|
||||||
|
↓
|
||||||
|
Checkpointer (每个 super-step 快照)
|
||||||
|
```
|
||||||
|
|
||||||
|
- **StateGraph**: 有状态图,参数化 State 类型
|
||||||
|
- **Nodes**: 纯函数,`(State) → updates`
|
||||||
|
- **Edges**: `add_conditional_edges`,支持循环/分支/合并
|
||||||
|
- **State**: `TypedDict` 或 Pydantic,带 reducer 处理并发更新
|
||||||
|
- **Reducers**: `add_messages` 等,自动处理追加 vs 覆盖
|
||||||
|
|
||||||
|
### 完整功能矩阵
|
||||||
|
|
||||||
|
| 功能 | 状态 | 细节 |
|
||||||
|
|------|------|------|
|
||||||
|
| **StateGraph** | ✅ 稳定 | 循环图(非 DAG),条件边缘,并行 fan-out |
|
||||||
|
| **Checkpointing** | ✅ v4.1.1 | SQLite / PostgreSQL / Redis 后端 |
|
||||||
|
| **Durable Execution** | ✅ 稳定 | 跨失败自动恢复,从精确断点继续 |
|
||||||
|
| **Human-in-the-loop** | ✅ 一等公民 | `interrupt()` + `Command(resume=...)` |
|
||||||
|
| **Time-travel 调试** | ✅ 稳定 | 回滚任意 checkpoint,fork 重放 |
|
||||||
|
| **流式输出** | ✅ 稳定 | Token 级 + State 级 + Event 级 |
|
||||||
|
| **多 Agent 编排** | ✅ 稳定 | Supervisor / Swarm / 层级 / Subgraph |
|
||||||
|
| **Comprehensive Memory** | ✅ 稳定 | 短时工作记忆 + 长时持久记忆 |
|
||||||
|
| **增量状态存储** | 🧪 DeltaChannel beta (v4.1.0+) | 长消息列表只存 delta |
|
||||||
|
| **跨进程状态同步** | 🧪 RemoteCheckpointer (v4.1.0+) | 分布式多 agent 架构 |
|
||||||
|
| **自动 checkpoint 清理** | ✅ keep_latest TTL (v4.0.2) | 避免无限制积累历史 |
|
||||||
|
| **LangGraph Platform** | ✅ 稳定 | Agent Server:持久化、任务队列、版本管理 |
|
||||||
|
| **LangGraph Studio** | ✅ 稳定 | 可视化 agent 工作流 |
|
||||||
|
|
||||||
|
### 成熟度
|
||||||
|
|
||||||
|
| 维度 | 状态 |
|
||||||
|
|------|------|
|
||||||
|
| 版本 | v1.0 GA(2025-10),checkpointer v4.1.1 (2026-05) |
|
||||||
|
| 稳定性 | 高,持久化为架构一等公民 |
|
||||||
|
| 生产证明 | Klarna, Replit, Elastic |
|
||||||
|
| 支持 | LTS-style support track |
|
||||||
|
| 适用场景 | 多步骤 agent、多 agent 系统、人工审批、长时间运行任务 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 功能边界对比
|
||||||
|
|
||||||
|
| 维度 | LangChain | LangGraph |
|
||||||
|
|------|-----------|-----------|
|
||||||
|
| **层次** | 高层应用框架 | 底层编排运行时 |
|
||||||
|
| **核心抽象** | `create_agent`, LCEL, 组件库 | `StateGraph`, Nodes, Edges, State |
|
||||||
|
| **思维模型** | 线性或 DAG 管道 | 节点 + 边缘的循环有向图 |
|
||||||
|
| **循环/分支** | 受限 | **一等公民**:任意循环、分支、合并 |
|
||||||
|
| **状态持久化** | 无原生支持 | **一等公民**:Checkpointer |
|
||||||
|
| **Human-in-loop** | 需手动编排 | **一等公民**:`interrupt()` + `Command` |
|
||||||
|
| **Time-travel 调试** | 无 | **一等公民**:回滚 fork 重放 |
|
||||||
|
| **Durable Execution** | 无 | **一等公民**:跨故障自动恢复 |
|
||||||
|
| **流式** | Token 级 | Token + State + Event 每节点流式 |
|
||||||
|
| **多 Agent 编排** | 需手动组合 | **原生**:Supervisor/Swarm/Subgraph |
|
||||||
|
| **模型抽象** | **核心优势** | 复用 LangChain |
|
||||||
|
| **600+ 集成** | **核心优势** | 可复用 LangChain 集成 |
|
||||||
|
| **LCEL 线性链** | **有** | 无 |
|
||||||
|
| **Middleware** | **v1.0 特有** | 无 |
|
||||||
|
| **学习曲线** | 中等 | 较陡(需图思维) |
|
||||||
|
| **部署平台** | 无独立平台 | LangGraph Platform + Studio |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 决策路线
|
||||||
|
|
||||||
|
```
|
||||||
|
你的 workflow 需要什么?
|
||||||
|
│
|
||||||
|
├─ 线性、始终相同步骤 → LangChain (LCEL / create_agent)
|
||||||
|
├─ 需要循环/分支/重试 → LangGraph (StateGraph)
|
||||||
|
├─ 需要持久化/故障恢复 → LangGraph (Checkpointer)
|
||||||
|
├─ 需要人工审批 → LangGraph (interrupt())
|
||||||
|
├─ 需要 time-travel 调试 → LangGraph (checkpoint + fork)
|
||||||
|
├─ 需要多 agent 协作 → LangGraph (Supervisor/Swarm/Subgraph)
|
||||||
|
└─ 不确定 → 先用 create_agent,遇到瓶颈下钻到 StateGraph
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 参考来源
|
||||||
|
|
||||||
|
- [LangChain Blog: v1.0 Milestone](https://www.langchain.com/blog/langchain-langgraph-1dot0)
|
||||||
|
- [LangChain Documentation](https://docs.langchain.com/oss/python/langchain/overview)
|
||||||
|
- [LangGraph Documentation](https://docs.langchain.com/oss/python/langgraph/overview)
|
||||||
|
- [LangGraph GitHub](https://github.com/langchain-ai/langgraph)
|
||||||
|
- [Atlan: LangChain vs LangGraph 2026](https://atlan.com/know/ai-agent/ai-agent-memory/langchain-vs-langgraph/)
|
||||||
|
- [truefoundry: LangChain vs LangGraph](https://www.truefoundry.com/blog/langchain-vs-langgraph)
|
||||||
@@ -0,0 +1,344 @@
|
|||||||
|
# 笔记:opencode 子代理调度、分发与合并及工作流推进
|
||||||
|
|
||||||
|
> 基于 `/Users/midnite/Samples/opencode` 源码调研,2026-07-04
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 一、整体架构
|
||||||
|
|
||||||
|
```
|
||||||
|
LLM(主 Agent)
|
||||||
|
│
|
||||||
|
├── 调用 Task tool(tool call)
|
||||||
|
│ ↓
|
||||||
|
│ TaskTool.execute() ← packages/opencode/src/tool/task.ts
|
||||||
|
│ │
|
||||||
|
│ ├── agent.get() ← 查找 Agent 定义(agent.ts)
|
||||||
|
│ ├── deriveSubagentPermission() ← 权限合并(subagent-permissions.ts)
|
||||||
|
│ ├── sessions.create() ← 创建子 session
|
||||||
|
│ │
|
||||||
|
│ ├── [前台] background.wait() + background.waitForPromotion() race
|
||||||
|
│ │ ↓ 完成
|
||||||
|
│ │ renderOutput() → XML <task> 标签返回
|
||||||
|
│ │
|
||||||
|
│ └── [后台] background.start() → notify() 异步注入结果
|
||||||
|
│
|
||||||
|
└── 会话循环(runLoop) ← prompt.ts
|
||||||
|
│
|
||||||
|
├── 检测 subtask type part → handleSubtask()
|
||||||
|
├── 检测 compaction → compaction.process()
|
||||||
|
└── 正常流程 → LLM.stream() → processor.handleEvent()
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 二、子代理调度(Dispatch)
|
||||||
|
|
||||||
|
### 2.1 三种触发入口
|
||||||
|
|
||||||
|
| 入口 | 触发方式 | 调用链路 |
|
||||||
|
|------|---------|---------|
|
||||||
|
| A — LLM 自主 | LLM 调用 `task` tool | 系统提示词中注入了 Task tool 描述 + `describeTask()` 输出子代理列表 → LLM 决策 |
|
||||||
|
| B — `subtask` part | 消息中有 `type: "subtask"` 的 part | `handleSubtask()` 直接执行 TaskTool,不走 LLM |
|
||||||
|
| C — `agent` part | 消息中有 `type: "agent"` 的 part | 转为"调用 task tool 带 subagent: XXX"的提示词,引导 LLM |
|
||||||
|
|
||||||
|
### 2.2 TaskTool.execute() 完整流程(task.ts)
|
||||||
|
|
||||||
|
```
|
||||||
|
execute(params, ctx):
|
||||||
|
1. background 开关检查(需 experimental flag)
|
||||||
|
2. ctx.ask() 权限询问
|
||||||
|
3. agent.get(subagent_type) 查找子代理定义
|
||||||
|
4. task_id 存在 → sessions.get(task_id) 恢复已有子 session
|
||||||
|
task_id 不存在 → sessions.create() 创建新子 session
|
||||||
|
5. deriveSubagentSessionPermission() 合并权限
|
||||||
|
6. 添加默认 deny 规则(todowrite / task)
|
||||||
|
7. 确定 model(继承或子代理自定义)
|
||||||
|
8. 执行 runTask() → ops.resolvePromptParts() + ops.prompt()
|
||||||
|
9. 结果格式化为 XML ← renderOutput()
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2.3 关键:子 session 创建(task.ts lines 121-158)
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
// 权限继承
|
||||||
|
const childPermission = deriveSubagentSessionPermission({
|
||||||
|
parentSessionPermission: parent.permission ?? [],
|
||||||
|
subagent: next,
|
||||||
|
})
|
||||||
|
|
||||||
|
// 默认 deny 规则
|
||||||
|
const childToolDenies = [
|
||||||
|
// 子代理自己的 permission 没允许 todowrite → 默认 deny
|
||||||
|
...(next.permission.some(r => r.permission === "todowrite") ? []
|
||||||
|
: [{ permission: "todowrite", pattern: "*", action: "deny" }]),
|
||||||
|
// 子代理自己的 permission 没允许 task → 默认 deny(防嵌套)
|
||||||
|
...(next.permission.some(r => r.permission === "task") ? []
|
||||||
|
: [{ permission: "task", pattern: "*", action: "deny" }]),
|
||||||
|
// 主 agent 专有工具也不给子代理
|
||||||
|
...(cfg.experimental?.primary_tools?.map(p => ({ permission: p, ... })) ?? []),
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 三、通信格式:Tool Call / Tool Result
|
||||||
|
|
||||||
|
### 3.1 父→子:Task tool 参数
|
||||||
|
|
||||||
|
```
|
||||||
|
{
|
||||||
|
subagent_type: "explore" | "general" | ...,
|
||||||
|
description: "简短描述(3-5词)",
|
||||||
|
prompt: "子代理的完整任务描述",
|
||||||
|
task_id?: "恢复已有子 session 时使用",
|
||||||
|
command?: "触发该调用的 CLI 命令(可选)",
|
||||||
|
background?: true // 后台模式(需 experimental flag)
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3.2 子→父:XML 包装的纯文本(renderOutput)
|
||||||
|
|
||||||
|
```xml
|
||||||
|
<task id="ses_xxxxx" state="completed">
|
||||||
|
<summary>任务简述</summary>
|
||||||
|
<task_result>
|
||||||
|
子 agent 输出的完整文本内容...
|
||||||
|
</task_result>
|
||||||
|
</task>
|
||||||
|
```
|
||||||
|
|
||||||
|
错误时:
|
||||||
|
|
||||||
|
```xml
|
||||||
|
<task id="ses_xxxxx" state="error">
|
||||||
|
<summary>任务失败</summary>
|
||||||
|
<task_error>
|
||||||
|
Error: 具体错误信息...
|
||||||
|
</task_error>
|
||||||
|
</task>
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3.3 传递给 LLM 的方式
|
||||||
|
|
||||||
|
**前台模式**:
|
||||||
|
```
|
||||||
|
TaskTool.execute() 返回 { output: "<task>...</task>" }
|
||||||
|
↓
|
||||||
|
AI SDK 将其转为 tool result,存入数据库 tool part
|
||||||
|
↓
|
||||||
|
下一轮 LLM 调用时,tool result 作为消息历史的一部分传入
|
||||||
|
↓
|
||||||
|
LLM 看到 XML,自行解析使用
|
||||||
|
```
|
||||||
|
|
||||||
|
**后台模式**:
|
||||||
|
```
|
||||||
|
TaskTool.execute() 立即返回 <task state="running">...
|
||||||
|
↓
|
||||||
|
子 agent 完成后 → background.wait() 触发 → inject()
|
||||||
|
↓
|
||||||
|
向父 session 注入合成 text part(synthetic: true)
|
||||||
|
携带 <task state="completed">... 结果
|
||||||
|
↓
|
||||||
|
父 LLM 在下一轮循环中看到该消息
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 四、分发与合并(Distribution & Merge)
|
||||||
|
|
||||||
|
### 4.1 并行分发
|
||||||
|
|
||||||
|
- **无专用分发层**。依赖 LLM 在单条消息中发出多个 tool call
|
||||||
|
- `task.txt` 引导 LLM:*"Launch multiple agents concurrently whenever possible"*
|
||||||
|
- 底层通过 Effect.ts 的 `Effect.forkIn(scope, { startImmediately: true })` 实现同一消息内多 tool call 并发
|
||||||
|
- **子 agent 之间完全隔离**,无直接通信
|
||||||
|
|
||||||
|
### 4.2 结果合并
|
||||||
|
|
||||||
|
**无专用合并逻辑。** 合并完全通过 LLM 的上下文理解完成:
|
||||||
|
|
||||||
|
- 前台:tool result 自然进入消息历史,LLM 下一轮读取
|
||||||
|
- CLI 命令:额外注入 "Summarize the task tool output above and continue with your task." 引导 LLM 总结
|
||||||
|
- LLM 自主调用:无额外引导,LLM 自行决定如何使用
|
||||||
|
|
||||||
|
### 4.3 前台/后台切换机制(task.ts lines 303-333)
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
// 前台执行
|
||||||
|
return yield* Effect.raceFirst(
|
||||||
|
background.wait({ id: nextSession.id }), // 等完成
|
||||||
|
background.waitForPromotion(nextSession.id), // 等 promote 到后台
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
当用户将前台任务 promote 到后台时,`waitForPromotion` 先返回(标记 `metadata.background = true`),TaskTool 转而返回后台模式的输出。
|
||||||
|
|
||||||
|
### 4.4 后台作业引擎(core/background-job.ts)
|
||||||
|
|
||||||
|
纯内存、非持久化注册表。使用 Effect.ts 的 `SynchronizedRef` 做并发控制。
|
||||||
|
|
||||||
|
| 操作 | 行为 |
|
||||||
|
|------|------|
|
||||||
|
| `start()` | 创建 job,fork run effect,返回 info |
|
||||||
|
| `extend()` | 追加顺序执行的 run(通过 `Deferred` 链式等待前一个完成) |
|
||||||
|
| `wait()` | `Deferred.await(done)`,可选 timeout |
|
||||||
|
| `waitForPromotion()` | 等待 `promoted` Deferred 或检测 `background` 标记 |
|
||||||
|
| `promote()` | 标记 `background = true`,触发 `onPromote` callback |
|
||||||
|
| `cancel()` | 设置 `cancelled`,close scope(中断所有子 fork) |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 五、工作流推进(Workflow Progression)
|
||||||
|
|
||||||
|
### 5.1 核心循环(prompt.ts → runLoop)
|
||||||
|
|
||||||
|
```
|
||||||
|
runLoop(sessionID):
|
||||||
|
while true:
|
||||||
|
1. MessageV2.filterCompactedEffect() 获取消息
|
||||||
|
2. MessageV2.latest() 取最近 user/assistant/tasks
|
||||||
|
3. 检查 finish 状态
|
||||||
|
- 不是 tool-calls 且有 finish → break(退出循环)
|
||||||
|
4. 取 tasks(subtask / compaction 队列)
|
||||||
|
- subtask → handleSubtask() → continue
|
||||||
|
- compaction → compaction.process() → continue/break
|
||||||
|
5. 检查 overflow → 自动创建 compaction task → continue
|
||||||
|
6. 构建 assistant message
|
||||||
|
7. SessionProcessor.create() 创建 handle
|
||||||
|
8. SessionTools.resolve() 解析所有工具
|
||||||
|
9. 构建 system prompt(环境信息 + skills + MCP + instructions)
|
||||||
|
10. handle.process() — 启动 LLM stream
|
||||||
|
11. 检查 result:
|
||||||
|
- "compact" → 返回给外层触发 compaction
|
||||||
|
- "stop" → break
|
||||||
|
- "continue" → 继续循环
|
||||||
|
```
|
||||||
|
|
||||||
|
### 5.2 SessionProcessor 事件处理(processor.ts)
|
||||||
|
|
||||||
|
| Stream 事件 | 处理逻辑 |
|
||||||
|
|------------|---------|
|
||||||
|
| `reasoning-start/delta/end` | 创建 reasoning part → 增量追加 → 最终持久化 |
|
||||||
|
| `tool-input-start/delta/end` | 创建/更新 tool part(pending 状态) |
|
||||||
|
| `tool-call` | 标记 running → 设置 input → **doom loop 检测** |
|
||||||
|
| `tool-result` | `completeToolCall()` → 持久化结果 + 附件 |
|
||||||
|
| `tool-error` | `failToolCall()` → 标记错误 |
|
||||||
|
| `provider-error` | 抛出异常 → 触发重试 |
|
||||||
|
| `text-start/delta/end` | 流式文本 → `updatePartDelta()` **增量持久化** |
|
||||||
|
| `step-start` | 创建快照(snapshot) |
|
||||||
|
| `step-finish` | 生成 patch diff → 更新 usage/tokens → **overflow 检测** → 触发 summary |
|
||||||
|
| `finish` | stream 结束 |
|
||||||
|
|
||||||
|
### 5.3 Doom Loop 检测(processor.ts lines 351-377)
|
||||||
|
|
||||||
|
连续 3 次**完全相同的 tool call**(相同名称 + 相同输入)触发权限询问:
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
const recentParts = parts.slice(-DOOM_LOOP_THRESHOLD) // DOOM_LOOP_THRESHOLD = 3
|
||||||
|
if (recentParts.length === DOOM_LOOP_THRESHOLD &&
|
||||||
|
recentParts.every(part =>
|
||||||
|
part.type === "tool" &&
|
||||||
|
part.tool === value.name &&
|
||||||
|
part.state.status !== "pending" &&
|
||||||
|
JSON.stringify(part.state.input) === JSON.stringify(input)
|
||||||
|
)) {
|
||||||
|
yield* permission.ask({ permission: "doom_loop", ... })
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 5.4 Compaction 工作流
|
||||||
|
|
||||||
|
两种触发方式:
|
||||||
|
|
||||||
|
| 触发条件 | 行为 |
|
||||||
|
|---------|------|
|
||||||
|
| step-finish 检测到 `isOverflow()` + `auto: true` | 创建 compaction task → 下一轮循环执行 → 压缩后 continue |
|
||||||
|
| step-finish 检测到 `isOverflow()` + `auto: false` | 标记 `assistantMessage.error` → idle 等待用户干预 |
|
||||||
|
|
||||||
|
Compaction 使用专门的 `compaction` agent(hidden, mode=primary, `*=deny`)执行。
|
||||||
|
压缩后的消息标记 `compacted: true`,后续通过 `MessageV2.filterCompactedEffect()` 过滤。
|
||||||
|
|
||||||
|
### 5.5 重试机制(processor.ts lines 658-672)
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
Effect.retry(
|
||||||
|
SessionRetry.policy({
|
||||||
|
provider: input.model.providerID,
|
||||||
|
parse, // 错误解析(区分可重试/不可重试)
|
||||||
|
set: (info) => status.set(sessionID, { type: "retry", ... }),
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
遇 provider 错误自动重试,LLM stream 完成后 `Effect.ensuring(cleanup)` 保证资源释放。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 六、六种内置 Agent
|
||||||
|
|
||||||
|
| 名称 | Mode | Hidden | 用途 | 核心权限特征 |
|
||||||
|
|------|------|--------|------|-------------|
|
||||||
|
| `build` | primary | 否 | 默认 agent,全部工具 | question/plan_enter=allow |
|
||||||
|
| `plan` | primary | 否 | 计划模式,禁用编辑 | edit=deny(除 plans), task(general)=deny |
|
||||||
|
| `general` | subagent | 否 | 通用子代理 | todowrite=deny(默认禁止改 todo) |
|
||||||
|
| `explore` | subagent | 否 | 只读代码探索 | `*=deny`,仅 read/grep/glob/bash/webfetch/websearch |
|
||||||
|
| `compaction` | primary | 是 | 会话压缩(自动) | `*=deny` |
|
||||||
|
| `title` | primary | 是 | 生成会话标题 | `*=deny`(step=1 时异步 fork) |
|
||||||
|
| `summary` | primary | 是 | 生成消息摘要 | `*=deny`(每个 step-finish 时异步 fork) |
|
||||||
|
|
||||||
|
用户可通过 `config.agent` 自定义 agent(支持 `mode: "all"`),也可通过 `agent.generate` 让 LLM 辅助生成。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 七、权限模型总结
|
||||||
|
|
||||||
|
```
|
||||||
|
父 session permission
|
||||||
|
│
|
||||||
|
├── 仅继承 deny 规则 + external_directory 规则 ← subagent-permissions.ts
|
||||||
|
│ (父 agent 的 allow 规则不传播到子代理)
|
||||||
|
│
|
||||||
|
├── 子代理自身 permission(来自 agent 定义)
|
||||||
|
│
|
||||||
|
├── 默认 deny:
|
||||||
|
│ - todowrite(除非子代理明确允许)
|
||||||
|
│ - task(除非子代理明确允许,默认防嵌套)
|
||||||
|
│
|
||||||
|
└── 主 agent 专有工具 deny(来自 config.experimental.primary_tools)
|
||||||
|
```
|
||||||
|
|
||||||
|
子代理的 session 权限 = **父 deny + 父 external_directory + 自身 permission - 默认 deny - primary_tools deny**。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 八、关键设计决策
|
||||||
|
|
||||||
|
| 决策 | 意图 | 效果/局限 |
|
||||||
|
|------|------|----------|
|
||||||
|
| 结果以 XML 纯文本嵌入上下文 | 简单、LLM 可直接理解 | LLM 自行解析 XML;大结果可能被截断 |
|
||||||
|
| 无专用 merge 逻辑 | 简洁,不引入额外抽象 | 依赖 LLM 的理解能力处理返回结果 |
|
||||||
|
| 默认禁止子代理嵌套 task | 防止无限递归 | 限制了多级分解场景 |
|
||||||
|
| 同一消息多 tool call 并发 | 利用 LLM 并行能力 | 子 agent 隔离,无法协作 |
|
||||||
|
| Effect.ts 贯穿全程 | 类型安全、结构化并发 | 学习曲线陡峭 |
|
||||||
|
| session 作为隔离边界 | 天然权限/消息隔离 | 每个子 session 独立数据库记录,开销较大 |
|
||||||
|
| 后台引擎纯内存 | 有意识取舍(注释说明) | 进程重启丢失状态 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 九、参考源码路径
|
||||||
|
|
||||||
|
| 文件 | 角色 |
|
||||||
|
|------|------|
|
||||||
|
| `packages/opencode/src/tool/task.ts` | Task tool 核心实现(调度入口) |
|
||||||
|
| `packages/opencode/src/tool/task.txt` | Task tool 的 LLM 使用说明 |
|
||||||
|
| `packages/opencode/src/agent/agent.ts` | Agent 定义注册中心 |
|
||||||
|
| `packages/opencode/src/agent/subagent-permissions.ts` | 子代理权限推导 |
|
||||||
|
| `packages/opencode/src/tool/registry.ts` | 工具注册 + `describeTask()` 列出可用子代理 |
|
||||||
|
| `packages/opencode/src/session/prompt.ts` | 会话循环 + `handleSubtask()` + 提示词构建 |
|
||||||
|
| `packages/opencode/src/session/processor.ts` | LLM stream 事件处理器 |
|
||||||
|
| `packages/opencode/src/session/tools.ts` | Tool ↔ AI SDK 桥接 |
|
||||||
|
| `packages/opencode/src/session/system.ts` | 系统提示词生成(含 Task tool 说明) |
|
||||||
|
| `packages/opencode/src/background/job.ts` | 后台作业包装层 |
|
||||||
|
| `packages/core/src/background-job.ts` | 后台作业核心引擎(内存注册表) |
|
||||||
+644
-63
@@ -1,13 +1,13 @@
|
|||||||
# AG Core Roadmap
|
# AG Core Roadmap
|
||||||
|
|
||||||
> 定稿日期:2026-05-11
|
> 定稿日期:2026-05-11
|
||||||
> 最后更新:2026-07-04(v0.1 发布完成)
|
> 最后更新:2026-07-09(Phase 14 完成 + M10 里程碑达成)
|
||||||
|
|
||||||
## 愿景
|
## 愿景
|
||||||
|
|
||||||
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
|
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
|
||||||
|
|
||||||
**当前状态**:Phase 0-4c 全部完成;Provider IR 重构(统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen Provider)已完成;LlmCycle 简化(IR 消息类型切换 + 桥接层移除)已完成;v0.1 发布就绪(**182 个测试通过、0 clippy 警告、7 个离线示例可运行**)。
|
**当前状态**:v0.2.0-rc.1 已打标签。Phase 0-14 全部完成。v0.3.0 实施中,Phase 15-19 共 5 个增量 Phase 待交付。目标是从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -240,95 +240,665 @@ graph BT
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 扩展计划(v0.2+)
|
## v0.2.0 — 生产就绪(Production-Ready Core)
|
||||||
|
|
||||||
> 以下功能在已完成的 phase 中已实现基础能力或在 Phase 4 阶段明确了边界,后续可按维度增量扩展。
|
**目标**:解决 Rust Agent 工具箱从"能跑"到"能被人依赖"的鸿沟。持久化、配置层、上下文管理三大块补齐后,开发者可在 30 分钟内写出生产可用的 Agent 服务。
|
||||||
> 设计参考:见 `docs/note-agent-harness-references.md`(OpenClaw / Hermes / OpenHuman / OpenHarness 横向对比)。
|
|
||||||
> OpenCode 借鉴:见 `docs/note-opencode-agent-switching.md`(Agent 切换 + System Prompt 拼接机制)。
|
|
||||||
|
|
||||||
### 已有扩展项(沿用)
|
**总体规模**:8 个增量 Phase(Phase 5-12),17 个可验证 Step。
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
### 功能清单
|
||||||
|-------|---------|------|--------|------|
|
|
||||||
| Prompt Optimizer | `prompt` | 提示词自动优化 | P3 | 待实现 |
|
|
||||||
| 流式接口优化 | `llm/stream` | 流式响应解析与事件化 | P0 | ✅ 已完成基础实现 |
|
|
||||||
|
|
||||||
### v0.2+ 新增扩展项
|
#### P0 — 必须交付
|
||||||
|
|
||||||
> 以下为基于 Phase 4 设计讨论确定的 v0.2+ 候选扩展方向,按维度分组。
|
| # | 功能 | 模块 | 方案要点 |
|
||||||
> 标注为"v0.2 待评估"表示在 Phase 4 完成后再决定是否启动。
|
|---|------|------|---------|
|
||||||
|
| 1 | SqliteStore | `memory` | `rusqlite` + `bundled` feature,`MemoryStore` 的 SQLite 实现,进程重启数据不丢 |
|
||||||
|
| 2 | ProviderConfig 扩展 + `from_env()` | `llm` | 补全 `timeout_secs` / `max_retries` 字段;`AG_LLM_*` 环境变量辅助函数 |
|
||||||
|
| 3 | ToolDefinition IR 正式化 | `tools` | 移除 deprecated OpenAI wire 格式,替换为自定义 `ToolDef` 结构体 |
|
||||||
|
| 4 | API 稳定性管理 | `*` | 公开枚举加 `#[non_exhaustive]`;CHANGELOG 记录 Breaking Changes;废弃 API 用 `#[deprecated]` 标记 |
|
||||||
|
| 5 | Quick Start + 端到端示例 | `examples/` | 30 行 `main.rs` 快速开始;一个"SQLite 持久化 + Provider + 工具调用 + 多轮对话"的可运行示例(`cargo run --example`) |
|
||||||
|
|
||||||
#### Multi-Agent / 协同
|
#### P1 — 重要但不阻塞
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
| # | 功能 | 模块 | 方案要点 |
|
||||||
|-------|---------|------|--------|------|
|
|---|------|------|---------|
|
||||||
| Multi-Agent 协同(Swarm) | `agent` | 子 Agent 委派、并行子任务、结果聚合 | P2 | v0.2 待评估 |
|
| 6 | Ollama Provider | `llm/provider` | OpenAI Compat,本地 LLM 支持,实现量极小 |
|
||||||
|
| 7 | VectorRetriever trait | `memory` | 语义检索 trait 抽象(`index` / `search`),不绑定后端实现 |
|
||||||
|
| 8 | 流式 `submit_turn_stream` | `agent` | `AgentSession` 新增 `submit_turn_stream()`,返回 `Stream<Item = StreamEvent>` |
|
||||||
|
| 9 | 测试补强 | `*` | wiremock Provider roundtrip 测试;多线程并发写入 MemoryStore 测试 |
|
||||||
|
|
||||||
#### 技能(Skills)
|
#### P2 — 有时间再做
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
| # | 功能 | 模块 | 备注 |
|
||||||
|-------|---------|------|--------|------|
|
|---|------|------|------|
|
||||||
| Markdown 技能按需加载 | `agent` / `prompt` | 兼容 `SKILL.md` 格式(Hermes / OpenHarness 风格),按 prompt 上下文动态加载 | P2 | v0.2 待评估 |
|
| 10 | MCP StreamableHttp | `tools` | 当前仅预留枚举变体 |
|
||||||
|
| 11 | Gemini Provider | `llm/provider` | 协议差异大,实现成本较高 |
|
||||||
|
| 12 | 文件系统 MemoryStore 后端 | `memory` | JSON/JSONL 轻量持久化 |
|
||||||
|
|
||||||
#### 记忆(Memory)
|
### ContextSlot 上下文管理
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
**模块归属**:`src/llm/context.rs`(与 `compact.rs` 同级)
|
||||||
|-------|---------|------|--------|------|
|
|
||||||
| 多通道检索(hybrid) | `memory/retriever` | 在 TextOverlap 之上叠加向量检索通道 | P2 | v0.2 待评估 |
|
|
||||||
| KnowledgeGraph 深度记忆 | `memory` | 实体-关系图、`note-knowledge-graph-design.md` 已记录设计 | P3 | v0.2 待评估 |
|
|
||||||
| TokenJuice 智能压缩 | `memory` / `llm/compact` | 借鉴 OpenHuman TokenJuice,对工具结果做语义压缩而非字节截断 | P3 | v0.2 待评估 |
|
|
||||||
|
|
||||||
#### 交互层(TUI / Gateway)
|
**核心概念**:`ContextSlot` 是一段带策略配置的消息列表,以 `slot_id` 为 namespace 独立持久化到 `MemoryStore`。支持三种模式、三种来源和派生关联(记录 `parent_id`)。
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
**核心类型**:
|
||||||
|-------|---------|------|--------|------|
|
|
||||||
| TUI / 多平台 Gateway | 应用层 | OpenClaw / Hermes 风格的消息平台桥接(Feishu / Telegram / Discord 等) | P3 | v0.2+ 应用层 |
|
|
||||||
|
|
||||||
#### 训练基础设施
|
```rust
|
||||||
|
pub struct ContextSlot { id, session_id, config, messages, store }
|
||||||
|
pub struct SlotConfig { mode: SlotMode, source: SlotSource, budget, compact }
|
||||||
|
pub enum SlotMode {
|
||||||
|
Full, // 完整对话历史
|
||||||
|
Focused(FocusedConfig), // 聚焦:保持 LLM 注意力
|
||||||
|
Readonly, // 只读参考上下文
|
||||||
|
}
|
||||||
|
pub struct FocusedConfig { keep_system, recent_turns, inject_summary }
|
||||||
|
pub enum SlotSource {
|
||||||
|
New, // 全新空槽,独立持久化
|
||||||
|
Derived { parent_id, strategy: DeriveStrategy }, // 从父 slot 派生
|
||||||
|
Static(Vec<Message>), // 预置消息,不持久化
|
||||||
|
}
|
||||||
|
pub enum DeriveStrategy { Full, Focused(FocusedConfig) }
|
||||||
|
pub struct ContextBudget { system, history, tools, tool_results, reserve }
|
||||||
|
```
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
**持久化 Key 命名**:
|
||||||
|-------|---------|------|--------|------|
|
- `slot_msg:{session_id}:{slot_id}:{index}` → 消息内容
|
||||||
| RL 轨迹导出 | `agent` | ShareGPT 格式轨迹、Atropos 集成(Hermes 风格) | P3 | v0.3+ 探索 |
|
- `slot_meta:{session_id}:{slot_id}` → `SlotMeta`(含 `parent_id`)
|
||||||
|
- `slot_rel:{session_id}:{child_id}:parent` → `"{parent_id}"`
|
||||||
|
|
||||||
#### 安全治理
|
**`AgentSession` 扩展**:
|
||||||
|
- `create_slot(id, config)` — 创建新 slot
|
||||||
|
- `switch_slot(id)` — 切换当前 slot
|
||||||
|
- `list_slots()` — 列出所有 slot
|
||||||
|
- `derive_slot(id, parent_id, strategy)` — 从父 slot 派生
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
**与 `ConversationMemory` 的关系**:保留不废除。`ConversationMemory` 继续服务传统对话场景。
|
||||||
|-------|---------|------|--------|------|
|
|
||||||
| Human-in-the-loop 审批 | `agent` / `tools/permission` | 高危工具执行前的异步审批回调(OpenHarness `permission_prompt` 模式) | P2 | v0.2 待评估 |
|
|
||||||
|
|
||||||
#### 流式 / 实时
|
**v0.2 不做**:
|
||||||
|
- ❌ `slot.fork()` / `merge()` — 分支方法推迟到 v0.3+
|
||||||
|
- ❌ `inject_summary` 自动生成 — v0.2 仅消费端(从 `SessionMemory` 读取),生成在 v0.3+
|
||||||
|
- ❌ 血缘关系图遍历 — 只存 `parent_id`,不做查询
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
**依赖**:Phase 0(MemoryStore trait)、Phase 3(MemoryStore 持久化)
|
||||||
|-------|---------|------|--------|------|
|
**优先级**:P1
|
||||||
| 流式 `submit_turn` | `agent/session` | Phase 4 v1 只暴露非流式 `submit_turn()`;v0.2 包装 `LlmCycle::submit_stream` 暴露流式入口 | P2 | v0.2 待评估 |
|
|
||||||
|
|
||||||
#### Agent 切换 / Prompt 动态(OpenCode 借鉴)
|
---
|
||||||
|
|
||||||
| 扩展项 | 所在模块 | 说明 | 优先级 | 状态 |
|
### v0.2.0 实施计划 — 8 个增量 Phase
|
||||||
|-------|---------|------|--------|------|
|
|
||||||
| Agent 身份切换(角色轮换) | `agent` | 借鉴 OpenCode Tab 键切换 build/plan:同一 `AgentSession` 持有可热替换的 `Agent` 引用,切换时不重置消息历史,在末尾追加 `synthetic: true` 的状态变更消息。详见 `docs/note-opencode-agent-switching.md` §4 | P2 | v0.2 待评估 |
|
> **编号说明**:Phase 5-12 接续 v0.1 的 Phase 0-4c,按开发顺序排列。
|
||||||
| System Prompt 多层动态拼接 | `agent/session` | 借鉴 OpenCode `request.ts:58-66`:拆分 `base_prompt + agent_prompt + env_context` 三层,`AgentSession::submit_turn` 每轮重算(不缓存),便于按 agent 类型动态切换 | P2 | v0.2 待评估 |
|
|
||||||
| **多 Context 切换** | `agent` | **Phase 4c 的 SessionMemory 数据结构已预留信息桥接通道,v0.2+ 在其上包装 `ContextManager` 实现完整的多 context 切换:创建/销毁/切换 context、通过 SessionMemory 桥接关键信息。详见 `docs/note-context-switch-design.md`** | P2 | v0.2 待评估 |
|
#### Phase 5: 热身准备(Warmup)
|
||||||
|
|
||||||
|
**目标**:快速交付三个互不依赖的独立改动,建立交付节奏。
|
||||||
|
|
||||||
|
| Step | 内容 | 文件范围 | 验证标准 |
|
||||||
|
|------|------|---------|---------|
|
||||||
|
| **5.1** ✅ | `ProviderConfig` 扩展:补 `timeout_secs`(def=30) + `max_retries`(def=3);新增 `ProviderConfig::from_env(prefix)` | `llm/provider.rs` + 各 Provider `new()` 构造函数 | `cargo test` + `from_env()` 单元测试 |
|
||||||
|
| **5.2** ✅ | `OllamaProvider`:基于 `GenericOpenaiProvider` 包装,改 base_url 为 `http://localhost:11434`;`ProviderType` 新增 `Ollama` | `llm/provider/provider.rs` + `llm/provider/ollama.rs`(新增) | `cargo build` — 纯类型级验证 |
|
||||||
|
| **5.3** ✅ | 公开枚举 `#[non_exhaustive]` 前置标记:`ProviderType` / `StopReason` / `FinishReason` / `EvictionPolicy` / `SlotMode`(预置) | 各枚举定义处 | 编译通过 + `cargo clippy` 0 警告 |
|
||||||
|
|
||||||
|
**实际新增**(2026-07-05 commit `98dfe6c`):
|
||||||
|
- 新增文件 1 个(`llm/provider/ollama.rs`,72 行)
|
||||||
|
- 修改文件 2 个(`llm/provider.rs` 加 `from_env` + `Default` + 4 个字段;`memory/store.rs` EvictionPolicy 加 `#[non_exhaustive]`)
|
||||||
|
- `ProviderType::Ollama` 变体 + `FromStr` 解析("ollama" → Ollama)
|
||||||
|
- `OllamaProvider::new(base_url, api_key, model, timeout_secs)` + `with_client()` 构造函数
|
||||||
|
- `ProviderConfig::from_env(prefix)` 解析 `{prefix}_API_KEY` / `{prefix}_BASE_URL` / `{prefix}_MODEL` 环境变量
|
||||||
|
- 全量测试 182 → 190(+8,phase 5 新增 from_env 与 Ollama 相关单测)
|
||||||
|
- clippy 0 警告
|
||||||
|
|
||||||
|
**依赖**:无(三个 Step 互不冲突)
|
||||||
|
**优先级**:P0(5.1)+ P1(5.2)+ P0 前置(5.3)
|
||||||
|
**为何独立成 Phase**:三个改动零文件重叠,可以并行推进。它们是后续所有 Phase 的"门把手"——先做完热身再进入核心工作。
|
||||||
|
**状态**:✅ Phase 5 全部交付物已完成
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 6: ToolDefinition IR 正式化
|
||||||
|
|
||||||
|
**目标**:引入 `ToolDef` 新类型,替换已标记 `#[deprecated]` 的 `ToolDefinition`(`OpenaiToolDefinition` 别名)。
|
||||||
|
|
||||||
|
**这是 v0.2 技术风险最高的 Phase**,影响 4 个模块约 8 个文件。通过 5 个 Step 逐文件切割确保每步可编译。
|
||||||
|
|
||||||
|
| Step | 内容 | 验证标准 |
|
||||||
|
|------|------|---------|
|
||||||
|
| **6.1** ✅ | `types/tool.rs` 新增 `ToolDef` 结构体 + `From<ToolDef> for OpenaiToolDefinition` + 反向 `From` | 单元测试 roundtrip |
|
||||||
|
| **6.2** ✅ | `types/mod.rs` 切别名 `pub type ToolDefinition = ToolDef`;`MessageRequest.tools` 改 `Vec<ToolDef>` | `cargo build` 编译断点 |
|
||||||
|
| **6.3** ✅ | `cycle.rs` 4 个方法签名 + `registry.rs` `definitions()` 签名更新 | `cargo build` |
|
||||||
|
| **6.4** ✅ | Provider 适配层(openai.rs / anthropic.rs / openai_compat.rs):`build_request()` 内做 `ToolDef → wire-format` 转换 | `cargo test` 每个 provider 测试 |
|
||||||
|
| **6.5** ✅ | 所有测试/示例中 `ToolDefinition` → `ToolDef` 修复;移除旧 `#[deprecated]` alias | `cargo test --all-targets` 全绿 |
|
||||||
|
|
||||||
|
**边界切割技巧**:
|
||||||
|
- Step 6.1 → 6.2 之间是安全 checkpoint:新类型存在但旧代码照常编译
|
||||||
|
- Provider 层不改序列化逻辑,只加一层 `From` 转换
|
||||||
|
- 当前代码中 `ToolDefinition` 已是 `#[deprecated(since = "0.1.0")]`,用户已有迁移预期
|
||||||
|
|
||||||
|
**依赖**:无(仅与 Phase 5.3 有枚举兼容关系)
|
||||||
|
**优先级**:P0
|
||||||
|
|
||||||
|
**实际新增**(2026-07-05 commit `4cf5918` / `9da9b83` / `b187519`,详见 `docs/13-phase6-tooldef-ir.md`):
|
||||||
|
- 修改文件 8 个:`llm/types/tool.rs`、`llm/types/mod.rs`、`llm/types/request_v2.rs`、`llm/cycle.rs`、`llm/provider/openai.rs`、`tools/registry.rs`、`tools/mcp.rs`、`agent/agent.rs`
|
||||||
|
- `ToolDef` IR(name / 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(+1,Phase 6 新增 roundtrip);clippy 0 警告
|
||||||
|
|
||||||
|
**状态**:✅ Phase 6 全部交付物已完成
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 7: SqliteStore 持久化
|
||||||
|
|
||||||
|
**目标**:实现 `MemoryStore` 的 SQLite 后端,进程重启数据不丢。
|
||||||
|
|
||||||
|
**与 Phase 6 无耦合,可重叠开发。**
|
||||||
|
|
||||||
|
| Step | 内容 | 文件 | 验证标准 |
|
||||||
|
|------|------|-----|---------|
|
||||||
|
| **7.1** ✅ | 新增 `memory/store/sqlite.rs`:`Mutex<Connection>` + `spawn_blocking`,实现 `save/get/delete/list` + prefix 过滤 | `memory/store/sqlite.rs` + `Cargo.toml`(add `rusqlite`) | 单元测试 CRUD + prefix 查询 |
|
||||||
|
| **7.2** ✅ | WAL 模式 + 并发安全 + 集成测试(`tokio::spawn` 10 个并发 task) | `sqlite.rs` 扩展 | 并发写入 100 轮无 race |
|
||||||
|
|
||||||
|
**设计决策**:
|
||||||
|
- 用 `Mutex<Connection>` 而非连接池(ponytail:一个连接够用就不加 r2d2)
|
||||||
|
- WAL 模式:`PRAGMA journal_mode=WAL` 解决读写锁
|
||||||
|
|
||||||
|
**依赖**:`MemoryStore` trait(v0.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(+9,Phase 7 新增 SqliteStore 单测);clippy 0 警告
|
||||||
|
|
||||||
|
**状态**:✅ Phase 7 全部交付物已完成
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 8: MVP 集成出口(v0.2.0-rc.1 候选)
|
||||||
|
|
||||||
|
**目标**:P0 五项全部交付。开发者 clone 仓库后 10 分钟跑起持久化 Agent。
|
||||||
|
|
||||||
|
| Step | 内容 | 验证标准 |
|
||||||
|
|------|------|---------|
|
||||||
|
| **8.1** ✅ | API 稳定性扫尾:`#[non_exhaustive]` × 14 公开枚举 + `StepStatus::Completed` 切 `MessageResponse` + CHANGELOG v0.2.0-rc.1 + Cargo.toml 0.2.0-rc.1 | `cargo doc --no-deps` 0 warning + 零 deprecated warning |
|
||||||
|
| **8.2** ✅ | Quick Start 示例(57 行 `main.rs`):MockProvider + EchoTool + submit_turn 真实工具调用 | `cargo run --example quick_start` exit 0 |
|
||||||
|
| **8.3** ✅ | 端到端示例:SqliteStore + AG_LLM_* from_env 自动检测 + 3 工具 + 3 轮对话 + 持久化跨连接验证 | `cargo run --example end_to_end`(Mock fallback,无需 API key)|
|
||||||
|
|
||||||
|
**Phase 8 全部完成**。**已打 `v0.2.0-rc.1` 标签**。
|
||||||
|
|
||||||
|
**实际新增**(2026-07-05,7 commits):
|
||||||
|
- `feat(core)` —— 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他)
|
||||||
|
- `refactor(agent)` —— `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` + `task_agent_demo.rs` 清理 3 处废弃类型
|
||||||
|
- `docs` —— CHANGELOG v0.2.0-rc.1 条目 + Cargo.toml version 0.1.0 → 0.2.0-rc.1 + README 示例列表 7 → 10
|
||||||
|
- `test(core)` —— 验证 commit 1-3 零回归(test 200 passed + clippy 0 警告 + doc 0 warning)
|
||||||
|
- `feat(examples)` —— `quick_start.rs`(60 行)+ `end_to_end.rs`(246 行)
|
||||||
|
- `docs(roadmap)` —— 标记 Phase 8 全部完成 + M4 里程碑 ✅
|
||||||
|
- `fix(examples)` —— 实施后 PM/SA/Code Reviewer 三方审查发现 6 项问题(🔴 CalcTool 除零 panic + 🟡 drop 注释准确性 + 🟡 EchoTool 错误处理 + 💭 断言一致性 + 💭 工具两端语义统一 + 💭 trailing newline),全部修复
|
||||||
|
|
||||||
|
**依赖**:Phase 5(ProviderConfig from_env)+ Phase 6(ToolDef)+ Phase 7(SqliteStore)
|
||||||
|
**优先级**:P0
|
||||||
|
**状态**:✅ Phase 8 全部交付物已完成
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 9: 流式体验增强
|
||||||
|
|
||||||
|
**目标**:Agent 会话支持流式输出,开发者看到实时 token。
|
||||||
|
|
||||||
|
| Step | 内容 | 文件 | 验证标准 |
|
||||||
|
|------|------|-----|---------|
|
||||||
|
| **9.1** ✅ | `AgentSession::submit_turn_stream(user_input) -> impl Stream<Item=StreamEvent>` | `agent/session.rs` | 单元测试验证流事件序列:`TextDelta → ... → MessageComplete` |
|
||||||
|
|
||||||
|
**注意**:tool 自动循环时流中插入 `ToolExecutionStarted` 事件,用户端 UI 显示"正在调用工具..."。
|
||||||
|
|
||||||
|
**依赖**:Phase 6(ToolDef)+ `LlmProvider.chat_stream`(v0.1 已有)
|
||||||
|
**优先级**:P1
|
||||||
|
|
||||||
|
**实际新增**(2026-07-06 commit `212cfcc`,详见 `docs/16-phase9-streaming-experience.md`):
|
||||||
|
- 方案文档:`docs/16-phase9-streaming-experience.md`(821 行,含状态机设计推演与边界情况)
|
||||||
|
- 修改文件 3 个:`src/agent/session.rs`(+208,含 `submit_turn_stream` / `finalize_turn`)、`src/llm/cycle.rs`(+784,含 `submit_with_tools_stream` / `run_tool_loop` spawn + mpsc 状态机)、`src/llm/types/response_v2.rs`(+21,含 `StreamEvent::ToolExecutionStarted`/`Completed` 变体 + `apply_to` 元事件)
|
||||||
|
- 关键设计:`CycleConfig` 加 `Clone` derive 以支持 spawn 跨 task;`finalize_turn` 手动同步状态(`submit_turn_stream` 返回流前不落库,避免半成品被 hook 误读)
|
||||||
|
- 测试:新增 9 个单元测试 + 2 个集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock provider 流消费 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211(+11,0 失败)
|
||||||
|
- clippy 0 警告
|
||||||
|
- 无新增外部依赖
|
||||||
|
|
||||||
|
**状态**:✅ Phase 9 全部交付物已完成
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 10: ContextSlot 上下文管理
|
||||||
|
|
||||||
|
**目标**:支持多上下文分区管理,Agent 可在不同 slot 之间切换。
|
||||||
|
|
||||||
|
| Step | 内容 | 验证标准 |
|
||||||
|
|------|------|---------|
|
||||||
|
| **10.1** ✅ | `src/agent/context.rs`:`ContextSlot` + `SlotConfig` / `SlotMode` / `FocusedConfig` / `SlotSource` / `DeriveStrategy` / `ContextBudget` / `SlotMeta` 核心类型 | `cargo build --all-targets` |
|
||||||
|
| **10.2** ✅ | ContextSlot 持久化:基于 `MemoryStore` trait(不绑定 SqliteStore)实现 save/load/list/delete + slot 命名空间 key 策略 + `load_messages()` Focused 读时过滤 + `append_messages()` Readonly 阻断 + colon 注入防护 | 单元测试:持久化 roundtrip / session 隔离 / Focused 边界 / delete 保护 / 派生 / load_messages() |
|
||||||
|
| **10.3** ✅ | `AgentSession` 扩展:`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot` + `new()` 自动创建 `"default"` slot + `submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回 + 新示例 `context_slot_demo` | 集成测试 + `cargo run --example context_slot_demo` exit 0 |
|
||||||
|
|
||||||
|
**如何保证简单场景无感**:`AgentSession::new()` 内部检查,自动创建 `"default"` slot → `submit_turn` 默认写到 default slot。
|
||||||
|
|
||||||
|
**实际新增**(2026-07-07 commit `6359422`,详见 `docs/17-phase10-contextslot.md`):
|
||||||
|
- 方案文档:`docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复)
|
||||||
|
- 新增文件 3 个:`src/agent/context.rs`(~430 行 ContextSlot 核心类型 + 持久化方法 + 22 个测试)、`src/agent/context.rs` 中的 `ContextSlot::filter_focused` 静态方法(被 `load_messages` 和 `derive_slot` 复用,消除代码重复)、`examples/context_slot_demo.rs`(~160 行分支对话示例:法律咨询 → 派生两个方向 → 切换 → 隔离验证 → 删除保护)
|
||||||
|
- 修改文件 3 个:`src/agent.rs`(+5 行 module 声明 + re-export)、`src/agent/error.rs`(+56 行:3 个新变体 `SlotReadonly`/`SlotNotFound`/`SlotAlreadyExists` + 4 个测试)、`src/agent/session.rs`(+825/-197 行:slots 字段 + 6 个管理方法 + submit_turn/finalize_turn 改造 + 17 个测试)
|
||||||
|
- 关键设计:
|
||||||
|
- **模块归属**:`agent/context.rs`(零新依赖方向,遵循 `agent → memory` 已有依赖)
|
||||||
|
- **持久化**:JSON blob 批次存储,每 slot 3-4 条 `MemoryItem`(`slot_data` / `slot_meta` / `slot_config` / `slot_rel`)
|
||||||
|
- **submit_turn 签名不变**:方案 A(内部 `current_slot_id` 状态),向后兼容
|
||||||
|
- **Focused 模式读时过滤**:`load_messages() -> Vec<Message>`,避免 Rust 借用检查问题
|
||||||
|
- **增量追加写回**:`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式数据不丢失
|
||||||
|
- **delete_slot 双重保护**:禁止删 `"default"` + 至少保留一个 slot
|
||||||
|
- **colon 注入防护**:`assert_no_colon` 在 key 构造时 panic
|
||||||
|
- **错误传播**:`serde_json` / `MemoryStore` 所有错误用 `?` 传播,无静默吞掉
|
||||||
|
- 验证:211 → 254 测试(+43 新测试),clippy 0 警告,doc 0 warning,10 + 1 示例全部 exit 0
|
||||||
|
- finalize_turn 签名变更(破坏性):新增 `new_messages_from_cycle: Vec<Message>` 参数,返回从 `()` 改为 `Result<(), AgentError>`——影响 Phase 9 的 `submit_turn_stream_triggers_turn_hooks` 和 `submit_turn_stream_end_to_end` 2 个测试,已适配
|
||||||
|
|
||||||
|
**依赖**:Phase 5(`#[non_exhaustive]` 预置 SlotMode 等枚举)、Phase 7(SqliteStore 推荐持久化后端;`MemoryStore` trait 即可)
|
||||||
|
**优先级**:P1
|
||||||
|
**状态**:✅ Phase 10 全部交付物已完成
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 11: 测试与检索补强
|
||||||
|
|
||||||
|
**目标**:补全测试覆盖 + 语义检索抽象。
|
||||||
|
|
||||||
|
| Step | 内容 | 验证标准 |
|
||||||
|
|------|------|---------|
|
||||||
|
| **11.1** ✅ | `VectorRetriever` trait:`index(id, embeddings)` + `search(query, k)` | 编译 + mock 测试 |
|
||||||
|
| **11.2** ✅ | wiremock Provider roundtrip 测试:模拟 OpenAI/Anthropic HTTP 端点 | `cargo test` 新增 10+ roundtrip 测试 |
|
||||||
|
| **11.3** ✅ | 并发测试补强:InMemoryStore + SqliteStore 多线程写入验证 | 跑 100 轮无 race |
|
||||||
|
|
||||||
|
**实际新增**(2026-07-06 commit `71abe88` / `b4e5c7d`,详见 `docs/18-phase11-testing-and-retrieval.md`):
|
||||||
|
- 方案文档:`docs/18-phase11-testing-and-retrieval.md`(647 行,含 11.1/11.2/11.3 设计 + 10 项架构决策 + 实施后补充 2 条偏差记录 #6 mid-stream mock 模式 + #7 429 retry-after 修复)
|
||||||
|
- 新增文件 1 个:`src/memory/vector.rs`(237 行 — `VectorRetriever` trait + `InMemoryVectorRetriever` 引用实现 + `dot()` 零依赖 + 6 个内联测试)
|
||||||
|
- 修改文件 5 个:
|
||||||
|
- `src/memory.rs`(+2 行:module 声明 + re-export)
|
||||||
|
- `src/llm/provider/openai.rs`(+8 wiremock 测试 + `handle_error_response` 429 retry-after 解析修复 5 行)
|
||||||
|
- `src/llm/provider/anthropic.rs`(+4 wiremock 测试)
|
||||||
|
- `src/memory/store/in_memory.rs`(+3 并发测试:100 并发写、5 写+5 读混合、15 写者容量淘汰)
|
||||||
|
- `src/memory/store/sqlite_store.rs`(+2 并发测试:100 并发写、5 写+5 读混合)
|
||||||
|
- 关键设计:
|
||||||
|
- **零依赖 dot()**:手写点积/范数,零新增 crate 依赖
|
||||||
|
- **Wiremock 测试自包含**:每个测试独立 `MockServer::start()`,沿用现有模式
|
||||||
|
- **429 retry-after 修复**:`openai.rs` 与 `anthropic.rs` 行为对齐(5 行代码)
|
||||||
|
- **偏差记录**:方案文档「已否决的方案 #6/#7」记录两处实施偏差,便于后续审计追溯
|
||||||
|
- 验证:254 → 277 测试(+23 个新测试),clippy 0 警告,doc 0 warning;并发测试连续 3 次运行稳定无 flaky
|
||||||
|
- **依赖**:无(与方案一致)
|
||||||
|
- **状态**:✅ Phase 11 全部交付物已完成
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 12: P2 锦上添花(可选)
|
||||||
|
|
||||||
|
**目标**:时间允许时按优先级交付。
|
||||||
|
|
||||||
|
| 优先级 | 功能 | 实现量估计 | 备注 |
|
||||||
|
|--------|------|-----------|------|
|
||||||
|
| **12.1** | 文件系统 MemoryStore(JSON/JSONL) | ~80 行 | 最简单,适合练手 |
|
||||||
|
| **12.2** | MCP StreamableHttp 传输 | ~150 行 | 协议还在演进 |
|
||||||
|
| **12.3** | Gemini Provider | ~300 行 | 协议差异大,建议推迟到 v0.3 |
|
||||||
|
|
||||||
|
**依赖**:无(独立交付)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### v0.2.0 Phase 依赖关系图
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph BT
|
||||||
|
P5["<b>Phase 5: 热身准备</b><br/>ProviderConfig::from_env<br/>Ollama Provider<br/>#[non_exhaustive] 标记"]:::done
|
||||||
|
P6["<b>Phase 6: ToolDef IR</b><br/>Provider 无关工具定义"]:::done
|
||||||
|
P7["<b>Phase 7: SqliteStore</b><br/>rusqlite + WAL<br/>9 个内联测试<br/>持久化 round-trip"]:::done
|
||||||
|
P8["<b>Phase 8: MVP 出口</b><br/>rc.1 标签<br/>14 枚举 #[non_exhaustive]<br/>StepStatus IR 迁移<br/>quick_start + end_to_end"]:::done
|
||||||
|
P9["<b>Phase 9: 流式体验增强</b><br/>submit_turn_stream<br/>submit_with_tools_stream<br/>9 单元测试 + 2 集成测试"]:::done
|
||||||
|
P10["<b>Phase 10: ContextSlot</b><br/>ContextSlot 类型<br/>JSON blob 持久化<br/>AgentSession 集成<br/>43 个新测试"]:::done
|
||||||
|
P11["<b>Phase 11: 测试与检索补强</b><br/>VectorRetriever trait<br/>12 wiremock tests<br/>5 并发测试"]:::done
|
||||||
|
P12["Phase 12<br/>P2 锦上添花"]:::p2
|
||||||
|
|
||||||
|
P8 --> P5
|
||||||
|
P8 --> P6
|
||||||
|
P8 --> P7
|
||||||
|
|
||||||
|
P9 --> P6
|
||||||
|
|
||||||
|
P10 --> P7
|
||||||
|
P10 --> P8
|
||||||
|
|
||||||
|
P11 -.-> P7
|
||||||
|
|
||||||
|
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
|
||||||
|
classDef warmup fill:#e2e8f0,stroke:#94a3b8
|
||||||
|
classDef core fill:#fbbf24,stroke:#d97706
|
||||||
|
classDef mvp fill:#4ade80,stroke:#16a34a
|
||||||
|
classDef p1 fill:#93c5fd,stroke:#2563eb
|
||||||
|
classDef p2 fill:#c4b5fd,stroke:#7c3aed
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 关键里程碑
|
||||||
|
|
||||||
|
| 里程碑 | Phase 完成条件 | 可验证指标 | 状态 |
|
||||||
|
|--------|---------------|-----------|------|
|
||||||
|
| **M1** | Phase 5 | 热身三项完成:`from_env()` 可用 / Ollama 类型存在 / `#[non_exhaustive]` 就位 | ✅ 2026-07-05 |
|
||||||
|
| **M2** | Phase 6 | `ToolDef` 全量切换,`cargo test --all-targets` 全绿 | ✅ 2026-07-05 |
|
||||||
|
| **M3** | Phase 7 | SqliteStore CRUD + 并发测试通过,进程重启数据不丢 | ✅ 2026-07-05 |
|
||||||
|
| **M4** | **Phase 8 (rc.1)** | P0 五项全部交付,`cargo run --example quick_start` 跑通 | ✅ 2026-07-05 |
|
||||||
|
| **M5** | Phase 9 | `submit_turn_stream` 流式事件序列验证通过 | ✅ 2026-07-06 |
|
||||||
|
| **M6** | Phase 10 | ContextSlot 创建/切换/派生集成测试通过 | ✅ 2026-07-07 |
|
||||||
|
| **M7** | Phase 11 | wiremock + 并发测试补强,测试总量 200+ | ✅ 2026-07-06 |
|
||||||
|
| **M8** | Phase 12(可选) | P2 功能按需交付 | ⏳ |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## v0.3.0 — 多 Agent 基础系统(Multi-Agent Foundation)
|
||||||
|
|
||||||
|
**目标**:从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"。补齐 LangChain 7 大组件中缺失的 Document 和 VectorStore 能力,落地笔记设计中的 ContextSlot fork/merge、摘要自动生成、知识图谱,建立 engine 引擎层(会话树 + time-travel Checkpointer + SubAgent Dispatch + Agent Switch),为即将开发的多 Agent 产品提供完整基础。
|
||||||
|
|
||||||
|
**总体规模**:7 个增量 Phase(Phase 13-19),总新增代码约 2600 行,测试从 277 → 380+。
|
||||||
|
|
||||||
|
### 功能清单
|
||||||
|
|
||||||
|
#### P0 — 必须交付
|
||||||
|
|
||||||
|
| # | 功能 | 模块 | 方案要点 |
|
||||||
|
|---|------|------|---------|
|
||||||
|
| 1 | 技术债清理(旧 types 文件) | `llm/types` | `request.rs` / `response.rs` / `old_stream.rs` 三个 Phase 0 旧文件删除;内部类型移入 `provider/openai.rs` |
|
||||||
|
| 2 | ContextSlot fork/merge | `agent/context` | `fork(child_id, strategy)` 别名 + `merge(child, MergeStrategy)` 三种策略(Append/Replace/Summarize) |
|
||||||
|
| 3 | Document 系统 | `document/`(新模块) | `Document` 核心类型 + `RecursiveCharacterSplitter`(递归字符分割,支持 chunk_size/chunk_overlap/separators) |
|
||||||
|
| 4 | Embedding 抽象 | `llm/embedding` | `Embedding` trait(`embed` / `dim`)+ `MockEmbedding` 测试实现 |
|
||||||
|
| 5 | 向量存储持久化 | `vector/`(新模块) | `VectorStore` trait + `InMemoryVectorStore`(读写)+ `PersistentVectorStore`(SqliteStore 后端)+ `RagPipeline` 组合器 |
|
||||||
|
| 6 | 摘要自动生成 | `agent` / `llm/hooks` | `SummaryConfig` 配置 + `OnTurnEnd` Hook 自动检测 token 水位 → 调 LLM 生成摘要 → `SessionMemory::set("conversation_summary", ...)` |
|
||||||
|
| 7 | SessionManager + 会话树 | `engine/`(新模块) | Session 工厂(`create`/`create_child`)+ 按 ID 恢复(`get`)+ 子树管理(`children`/`parent`/`destroy_subtree`)+ 元数据持久化(MemoryStore) |
|
||||||
|
| 8 | Time-travel Checkpointer | `engine/checkpointer` | `checkpoint(session)` 全量序列化 + `rollback(session_id, ckpt_id)` 回滚 + `fork(session_id, ckpt_id, new_id)` 分支 + `list_checkpoints` |
|
||||||
|
| 9 | Agent Switch | `engine/switch` | 热切换 `session.agent`(替换 `Arc<dyn Agent>`),slot 历史 / turn_index / session_memory 全保留 |
|
||||||
|
| 10 | SubAgent Dispatch | `engine/sub_agent` | `dispatch(parent, sub_agent, task, config)` 单任务 + `dispatch_all(parent, tasks, config)` 并行派发(Semaphore 并发控制)+ 子 SessionMemory 继承 + `SubTaskResult` 结构化回传 |
|
||||||
|
| 11 | 知识图谱 | `memory/graph` | `KnowledgeGraph` trait(`add_entity` / `add_relation` / `get_related` / `find_by_keywords`)+ `InMemoryGraph` 实现 + `tag_index` 标签管理 |
|
||||||
|
| 12 | 双通道检索 | `memory/retriever` | `MemoryRetriever` 扩展为双通道(`KnowledgeStore` + `KnowledgeGraph`)+ `RetrievalStrategy::Hybrid` |
|
||||||
|
|
||||||
|
### 实施计划 — 7 个增量 Phase
|
||||||
|
|
||||||
|
> **编号说明**:Phase 13-19 接续 v0.2 的 Phase 5-12,按开发顺序排列。
|
||||||
|
|
||||||
|
#### Phase 13: 热身清理 + ContextSlot fork/merge
|
||||||
|
|
||||||
|
**目标**:清除 Phase 0 遗留的旧 types 文件,交付超低价功能建立节奏。
|
||||||
|
|
||||||
|
| Step | 内容 | 文件范围 | 验证标准 |
|
||||||
|
|------|------|---------|---------|
|
||||||
|
| **13.1** | `OpenaiChatRequest` 移入 `provider/openai.rs`,`types/request.rs` 删除 | `llm/types/request.rs` + `llm/provider/openai.rs` | `cargo build --all-targets` |
|
||||||
|
| **13.2** | `OpenaiChatResponse/Chunk` 移入 `provider/openai.rs`,`types/response.rs` 删除 | `llm/types/response.rs` + `llm/provider/openai.rs` | `cargo build --all-targets` |
|
||||||
|
| **13.3** | `old_stream.rs` 删除 + `types/mod.rs` 中 `ChatResponse` 删除 | `llm/types/old_stream.rs` + `llm/types/mod.rs` | `cargo build` + 确认 3 个旧文件不存在 |
|
||||||
|
| **13.4** | `ToolChoice` 从 `request.rs` 搬到 `tool.rs` | `llm/types/tool.rs` + `llm/types/request_v2.rs` | `cargo test --all-targets` 全绿 |
|
||||||
|
| **13.5** | `ContextSlot::fork(child_id, strategy)` 别名 + `merge(child, MergeStrategy)` | `agent/context.rs` | 单元测试:fork → 子 slot 消息 = 父 slot 副本;merge(Append) → 消息按序追加 |
|
||||||
|
|
||||||
|
**依赖**:无
|
||||||
|
**优先级**:P0
|
||||||
|
**预估规模**:约 200 行
|
||||||
|
**状态**:✅ Phase 13 全部交付物已完成(2026-07-08)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 14: Document 系统 + Embedding 抽象
|
||||||
|
|
||||||
|
**目标**:补齐 LangChain 7 大组件中最明显的缺口——Document 类型和分割器。不搞 Loader 框架,用户用 `fs::read_to_string` 自行加载。
|
||||||
|
|
||||||
|
**交付物**:
|
||||||
|
1. `src/document.rs` 新模块(`Document` 类型 + `RecursiveCharacterSplitter`)
|
||||||
|
2. `src/llm/embedding.rs`(`Embedding` trait + `MockEmbedding`)
|
||||||
|
|
||||||
|
**设计要点**:
|
||||||
|
- `Document`:id / content / metadata(HashMap<String, String>)/ mime_type
|
||||||
|
- `RecursiveCharacterSplitter`:chunk_size(默认 1000)/ chunk_overlap(默认 200)/ separators(`["\n\n", "\n", "。", "?", "!", ".", " ", ""]`,含 CJK 标点)
|
||||||
|
- 两阶段算法:按 separator 优先级递归分割(Phase 1)+ 贪心合并 + overlap 滑动窗口(Phase 2)
|
||||||
|
- 所有长度比较以 Unicode 字符数为单位(`chars_len()`),非字节数
|
||||||
|
- `Embedding` trait:`async fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>, LlmError>` + `fn dim()`
|
||||||
|
- 复用 `LlmError` 而非新错误类型
|
||||||
|
- `MockEmbedding`:sin-hash 零依赖伪随机向量 + L2 归一化
|
||||||
|
- 不引入 `DocumentLoader` trait(应用层职责)
|
||||||
|
|
||||||
|
**实际新增**(2026-07-09 commit `d4c4d8f`,详见 `docs/20-phase14-document-and-embedding.md`):
|
||||||
|
- 新增文件 3 个:
|
||||||
|
- `src/document.rs`(580 行)— `Document` 类型(4 字段 + `new`/`from_raw` 构造器,2 个 `new` 接受 `impl Into<String>`) + `RecursiveCharacterSplitter`(两阶段算法:按 separator 优先级递归分割 + 贪心合并 overlap,所有长度比较 `chars_len()` 字符级,overlap 提取 `chars().rev().take().rev()` 字符级安全)+ 19 个内联测试
|
||||||
|
- `src/llm/embedding.rs`(183 行)— `Embedding` trait(async + `LlmError`)+ `MockEmbedding`(sin-hash:字节和+长度做种子,`f32::sin(seed + i) * 10000`,L2 归一化到单位长度,零向量防除零)+ 6 个内联测试
|
||||||
|
- `examples/document_demo.rs`(74 行)— 端到端演示 Document → RecursiveCharacterSplitter → MockEmbedding → InMemoryVectorRetriever → search
|
||||||
|
- 修改文件 2 个:
|
||||||
|
- `src/lib.rs`(+3 行:`pub mod document` + `pub use document::Document` + 空行)
|
||||||
|
- `src/llm.rs`(+1 行:`pub mod embedding`)
|
||||||
|
- 关键设计:
|
||||||
|
- **早返回守卫**:`split_text` 在 `chars_len(text) <= self.chunk_size` 时直接返回 `[text]`,避免短文本在 Phase 2 `join("")` 中丢失 separator 边界
|
||||||
|
- **`Document::new` 使用 `impl Into<String>`**:接受 `&str` 或 `String`,比规范示例的 `String` 更灵活
|
||||||
|
- **`new()` panic + `try_new()` Result 双路径**:与 Rust 库惯例一致
|
||||||
|
- **CJK 分隔符扩展**:`DEFAULT_SEPARATORS` 包含 `"。"`/`"?"`/`"!"`,避免中文文本跳过句子级退化为空格分割
|
||||||
|
- **chunk_size = 0 校验**:构造器拒绝零值,避免字符级兜底死循环
|
||||||
|
- **tracing 埋点**:`split()` 入口 `tracing::debug!` + 每文档/每 chunk `tracing::trace!`
|
||||||
|
- **debug_assert 溢出保护**:单文档 chunk 数 < 10000 时 `debug_assert!`
|
||||||
|
- **Metadata 键覆盖文档化**:`HashMap::insert()` 静默覆盖 source_id/chunk_index/chunk_count 在 `split()` doc comment 注明
|
||||||
|
- 测试:19 个 Document 测试(含 1 个 split_multibyte_utf8_boundary CJK 边界测试)+ 6 个 Embedding 测试,全量 286 → 313(+27 新测试,但部分测试覆盖范围重叠计算约 25 个净增)
|
||||||
|
- 方案文档:`docs/20-phase14-document-and-embedding.md`(1417 行,含背景/调研/方案对比/实施计划(详细版)/3 轮审查修复记录),经过 3 轮 PM/SA 审查 + 1 轮实施后修复
|
||||||
|
- clippy 0 警告,doc 0 warning
|
||||||
|
- 无新增外部依赖(`Cargo.toml` 未修改)
|
||||||
|
|
||||||
|
**实施后调整**:
|
||||||
|
- 实施发现方案算法中 Phase 1 累加器设计与测试期望冲突("para1\n\npara2" 在 chunk_size=100 时 1 chunk 更合理),简化为"按 separator 切分 + Phase 2 合并"两阶段分工
|
||||||
|
- 二次审查发现 `split_text` 缺少早返回守卫 + `current_sep_count` 虚增计数,全部已修复
|
||||||
|
|
||||||
|
**依赖**:无(纯数据结构 + 零新 crate 依赖)
|
||||||
|
**优先级**:P0
|
||||||
|
**预估规模**:约 350 行
|
||||||
|
**状态**:✅ Phase 14 全部交付物已完成(2026-07-09)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 15: 向量存储持久化(SqliteStore 后端)
|
||||||
|
|
||||||
|
**目标**:实现 VectorStore 持久化,让语义检索支持进程重启后数据恢复。
|
||||||
|
|
||||||
|
**设计决策**:不用 pgvector。基于已有 SqliteStore(`rusqlite`)做持久化包装——运行时全量加载到 InMemory 索引做余弦搜索,写时同步到 SqliteStore。
|
||||||
|
|
||||||
|
**交付物**:
|
||||||
|
1. `src/vector/` 新模块:`VectorStore` trait + `InMemoryVectorStore` + `PersistentVectorStore` + `RagPipeline`
|
||||||
|
2. `VectorStore` trait:`add(docs, embeddings)` / `search(query, k)` / `remove(ids)`
|
||||||
|
3. `PersistentVectorStore`:构造时从 SqliteStore 加载已有索引;`add` 双向写入;`search` 纯内存搜索
|
||||||
|
4. `RagPipeline`:组合器封装 `split` → `embed` → `store.add` 的 ingest 流程,以及 `embed` → `store.search` 的 retrieve 流程
|
||||||
|
5. SqliteStore 存储格式:`vec:{namespace}:{doc_id}` → JSON `{doc_id, content, metadata, embedding}`
|
||||||
|
|
||||||
|
**依赖**:Phase 14(Document 类型)
|
||||||
|
**优先级**:P0
|
||||||
|
**预估规模**:约 400 行
|
||||||
|
**状态**:⏳ 待实施
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 16: 摘要自动生成
|
||||||
|
|
||||||
|
**目标**:闭环长对话能力。v0.2 的 `inject_summary` 消费端(`FocusedConfig.summary_override`)已就绪,缺的是生产端。
|
||||||
|
|
||||||
|
**交付物**:
|
||||||
|
1. `SummaryConfig` 结构体:`enabled` / `trigger_token_ratio`(默认 0.75)/ `summary_prompt`(可自定义)
|
||||||
|
2. 在 `OnTurnEnd` Hook 中插检查点:检测 token 水位超过 `trigger_token_ratio` → 调 LLM 生成摘要 → `SessionMemory::set("conversation_summary", summary)`
|
||||||
|
3. `AgentBuilder` 扩展:`.summary_config(cfg)` 方法
|
||||||
|
|
||||||
|
**为什么放 Hook 而非内置**:可插拔,默认不启用,用户 opt-in。不改变现有 `submit_turn` 行为。
|
||||||
|
|
||||||
|
**依赖**:无(Hook 系统 + SessionMemory 已就绪)
|
||||||
|
**优先级**:P0
|
||||||
|
**预估规模**:约 150 行
|
||||||
|
**状态**:⏳ 待实施
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 17: Agent 执行引擎(会话树 + Time-travel Checkpointer)
|
||||||
|
|
||||||
|
**目标**:建立 `engine/` 模块。解决 v0.2 中"session 在变量里、无法通过 ID 恢复、不支持父子关系"的空白。
|
||||||
|
|
||||||
|
**交付物**:
|
||||||
|
1. `src/engine/` 新模块(`session_manager.rs` + `checkpointer.rs` + `error.rs`)
|
||||||
|
2. `SessionManager`:
|
||||||
|
- `create(agent, bundle) -> session_id` — 创建根 session
|
||||||
|
- `create_child(parent_id, child_id, agent)` — 创建子 session(继承父 `RuntimeBundle`)
|
||||||
|
- `get(session_id) -> Arc<Mutex<AgentSession>>` — 按 ID 查找(支持从持久化恢复)
|
||||||
|
- `children(parent_id)` / `parent(child_id)` — 树形查询
|
||||||
|
- `destroy(id)` / `destroy_subtree(id)` — 生命周期管理
|
||||||
|
- `tree() -> SessionTreeSnapshot` — 树结构快照
|
||||||
|
3. `Checkpointer`:
|
||||||
|
- `checkpoint(session)` — 每个 `submit_turn` 末尾自动保存全量状态快照
|
||||||
|
- `rollback(session_id, ckpt_id)` — 回滚到任意历史 checkpoint
|
||||||
|
- `fork(session_id, ckpt_id, new_id)` — 从历史 checkpoint 分支出新 session
|
||||||
|
- `list_checkpoints(session_id)` — 列出 checkpoint 列表
|
||||||
|
4. `AgentSession` 新增 `Serialize + Deserialize` 以支持 checkpoint 序列化
|
||||||
|
|
||||||
|
**Checkpoint 存储格式**:`checkpoint:{session_id}:{ckpt_id}` → JSON(完整 AgentSession,含所有 slot 消息列表)。Ponytail:全量 JSON 够用,等遇到存储效率问题时再改增量模式。
|
||||||
|
|
||||||
|
**会话树持久化**:`session_meta:{session_id}` → `{agent_name, parent_id, created_at, turn_count}`;`session_rel:{child_id}` → `"parent_id"`
|
||||||
|
|
||||||
|
**依赖**:Phase 10(ContextSlot 持久化 — 消息由 slot 自己管,Checkpointer 管执行状态)
|
||||||
|
**优先级**:P0
|
||||||
|
**预估规模**:约 600 行
|
||||||
|
**状态**:⏳ 待实施
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 18: Agent Switch + SubAgent Dispatch + Agent 间交互
|
||||||
|
|
||||||
|
**目标**:在 SessionManager 基础上,提供 Agent 角色热切换和子代理调度能力。
|
||||||
|
|
||||||
|
**交付物**:
|
||||||
|
1. `engine/switch.rs` — `switch_agent(session_id, new_agent)`:替换 `Arc<dyn Agent>`,slot 历史 / turn_index / session_memory 全保留
|
||||||
|
2. `engine/sub_agent.rs` — SubAgent Dispatch 核心:
|
||||||
|
- `DispatchConfig`:`max_concurrency`(默认 10)/ `inherit_session_memory`(默认 true)/ `bridge_keys`
|
||||||
|
- `dispatch(parent_id, sub_agent, task, config) -> SubTaskResult`:创建子 session → 继承父 SessionMemory → `submit_turn` → 返回结构化结果
|
||||||
|
- `dispatch_stream(parent_id, sub_agent, task, config) -> SubTaskStream`:流式版
|
||||||
|
- `dispatch_all(parent_id, tasks, config) -> Vec<SubTaskResult>`:并行派发,`tokio::sync::Semaphore` 控制并发数
|
||||||
|
3. `SubTaskResult`:`child_id` / `response` / `usage` / `summary` + `child_memory(sm)` 读取子 SessionMemory
|
||||||
|
|
||||||
|
**Agent 间交互三层级**:
|
||||||
|
- 父→子:继承 SessionMemory 快照 + `bridge_keys` 指定 key 强制注入 system prompt
|
||||||
|
- 子→父:`SubTaskResult` 结构化回传 + `SessionMemory["result_summary"]` 结论摘要
|
||||||
|
- 子↔子(间接):通过公共 `MemoryStore` namespace(`shared:{parent_session_id}`)共享数据
|
||||||
|
|
||||||
|
**依赖**:Phase 17(SessionManager + 会话树)
|
||||||
|
**优先级**:P0
|
||||||
|
**预估规模**:约 500 行
|
||||||
|
**状态**:⏳ 待实施
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
#### Phase 19: 知识图谱 + 双通道检索
|
||||||
|
|
||||||
|
**目标**:落地 `docs/note-knowledge-graph-design.md` 中记录的知识图谱设计,提供实体-关系图检索能力。扩展 `MemoryRetriever` 为双通道。
|
||||||
|
|
||||||
|
**交付物**:
|
||||||
|
1. `src/memory/graph.rs`(新文件):
|
||||||
|
- `GraphEntity` / `GraphRelation` / `ScoredEntity` 核心类型
|
||||||
|
- `RelationDirection` 枚举(Outgoing / Incoming / Both)
|
||||||
|
- `KnowledgeGraph` trait:`add_entity` / `get_entity` / `remove_entity` / `add_relation` / `remove_relation` / `get_related` / `find_by_keywords` / `find_tags` / `set_entity_tags`
|
||||||
|
- `InMemoryGraph` 实现:`HashMap<String, GraphEntity>` + `Vec<GraphRelation>` + BFS 图遍历
|
||||||
|
- `TagConstraints`(`max_tags_per_entity` 默认 8)
|
||||||
|
2. `src/memory/retriever.rs` 扩展:
|
||||||
|
- `MemoryRetriever` 增加 `knowledge_graph` 可选字段
|
||||||
|
- `RetrievalStrategy` 枚举:`Hybrid`(默认)/ `KnowledgeOnly` / `GraphOnly`
|
||||||
|
|
||||||
|
**与 Document 系统的关系**:知识图谱提供实体级检索("这个实体和什么相关"),VectorStore 提供语义相似度检索("哪些文档最相似"),两者互补。
|
||||||
|
|
||||||
|
**依赖**:MemoryStore 持久化(v0.1 Phase 3)
|
||||||
|
**优先级**:P0
|
||||||
|
**预估规模**:约 400 行
|
||||||
|
**状态**:⏳ 待实施
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### v0.3.0 Phase 依赖关系图
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
graph BT
|
||||||
|
P13["<b>Phase 13: 热身清理</b><br/>旧 types 文件删除<br/>ContextSlot fork/merge"]:::done
|
||||||
|
P14["<b>Phase 14: Document + Embedding</b><br/>Document 类型<br/>RecursiveCharacterSplitter<br/>Embedding trait"]:::done
|
||||||
|
P15["<b>Phase 15: 向量存储持久化</b><br/>VectorStore trait<br/>PersistentVectorStore<br/>RagPipeline"]:::pending
|
||||||
|
P16["<b>Phase 16: 摘要自动生成</b><br/>SummaryConfig<br/>OnTurnEnd Hook"]:::pending
|
||||||
|
P17["<b>Phase 17: 执行引擎</b><br/>SessionManager<br/>会话树<br/>Time-travel Checkpointer"]:::pending
|
||||||
|
P18["<b>Phase 18: 切换与调度</b><br/>Agent Switch<br/>SubAgent Dispatch<br/>dispatch_all 并发控制"]:::pending
|
||||||
|
P19["<b>Phase 19: 知识图谱</b><br/>KnowledgeGraph trait<br/>InMemoryGraph<br/>双通道检索"]:::pending
|
||||||
|
|
||||||
|
P15 --> P14
|
||||||
|
P18 --> P17
|
||||||
|
|
||||||
|
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
|
||||||
|
classDef pending fill:#fbbf24,stroke:#d97706,color:#1a1a1a
|
||||||
|
```
|
||||||
|
|
||||||
|
### 关键里程碑
|
||||||
|
|
||||||
|
| 里程碑 | Phase 完成条件 | 可验证指标 | 状态 |
|
||||||
|
|--------|---------------|-----------|------|
|
||||||
|
| **M9** | Phase 13 | 旧 types 文件删除、`cargo test --all-targets` 全绿、`fork`/`merge` 测试通过 | ✅ 2026-07-08 |
|
||||||
|
| **M10** | Phase 14 | `Document` + `RecursiveCharacterSplitter` 分割结果验证、`MockEmbedding` 测试通过 | ✅ 2026-07-09 |
|
||||||
|
| **M11** | Phase 15 | `PersistentVectorStore` 持久化 roundtrip、`RagPipeline::ingest → retrieve` 端到端验证 | ⏳ |
|
||||||
|
| **M12** | Phase 16 | 多轮对话后摘要自动写入 SessionMemory、派生 slot 时摘要正确注入 | ⏳ |
|
||||||
|
| **M13** | **Phase 17 (rc.1)** | `SessionManager` 创建/子树/恢复集成测试通过、`Checkpointer` checkpoint/rollback/fork 验证 | ⏳ |
|
||||||
|
| **M14** | Phase 18 | `switch_agent` 热切换验证、`dispatch`/`dispatch_all` 多轮对话 + 结果回传验证 | ⏳ |
|
||||||
|
| **M15** | Phase 19 | `KnowledgeGraph` 实体-关系 CRUD + `get_related` BFS 验证、双通道检索 Hybrid 策略验证 | ⏳ |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## v0.4+ 展望
|
||||||
|
|
||||||
|
### 已规划的功能
|
||||||
|
|
||||||
|
| 功能 | 说明 | 预计版本 |
|
||||||
|
|------|------|---------|
|
||||||
|
| Multi-Agent Swarm 编排 | Supervisor/Subgraph 模式,基于 v0.3 dispatch 构建 | v0.4 |
|
||||||
|
| Human-in-the-loop 审批 | `interrupt()` + `Command(resume=...)` 异步审批回调 | v0.4 |
|
||||||
|
| Agent 自动创生 | LLM 自主决定何时派发子 agent、派发什么角色 | v0.4 |
|
||||||
|
| 分布式 session 共享 | SessionManager Redis 后端支持跨进程 | v0.4 |
|
||||||
|
| 精确 tokenizer 计数 | 引入 `tiktoken-rs`,绑定模型具体 tokenizer,替换字符估算 | v0.4+ |
|
||||||
|
| TokenJuice 语义压缩 | 对工具结果做语义压缩而非字节截断 | v0.4+ |
|
||||||
|
| Markdown 技能按需加载 | 技能注册表 + 按 prompt 上下文动态加载 | v0.4+ |
|
||||||
|
| 增量 checkpoint | 仅存储变化部分,替换当前全量 JSON 模式 | v0.4+ |
|
||||||
|
| RL 轨迹导出 | ShareGPT 格式轨迹、Atropos 集成 | v0.4+ |
|
||||||
|
|
||||||
|
### 明确不做(agcore 范围外)
|
||||||
|
|
||||||
|
| 功能 | 原因 |
|
||||||
|
|------|------|
|
||||||
|
| TUI / 多平台 Gateway | 应用层职责(Feishu / Telegram / Discord 桥接) |
|
||||||
|
| 配置自动加载(config/figment) | 配置来源策略应由上游应用决定,agcore 不定义配置格式 |
|
||||||
|
| 提示词自动优化 | 属于智能层,不应内建于 core 库 |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 风险与建议
|
## 风险与建议
|
||||||
|
|
||||||
1. **Phase 0 已完成**:LLM 调用周期基础设施已全部实现,可以支撑后续模块开发
|
1. **持久化依赖**:`rusqlite` + `bundled` 零外部依赖编译,但 SQLite 不适配所有场景(分布式/高并发写)。`MemoryStore` trait 的抽象层允许下游自行实现 Redis / PostgreSQL 后端
|
||||||
2. **并行可能性**:Phase 0 和 Phase 1 可并行开展(无相互依赖),可加速早期交付
|
2. **ContextSlot 心智负担**:`ContextSlot` 引入了一等抽象的复杂度。建议通过 `AgentBuilder` 默认创建 `"default"` slot,让简单场景无感使用
|
||||||
3. **MCP 协议复杂性**:MCP 涉及协议握手、session 管理、长期连接,建议预留充足时间调研协议细节
|
3. **向量检索规模上限**:v0.3 的 `PersistentVectorStore` 全量加载到内存做余弦搜索,适合 ≤10 万条向量。超出此规模需换用专用向量库。v0.4 可以评估引入
|
||||||
4. **Scope 蔓延风险**:当前 specs 只有 1 份文档,建议每个模块上线前都产出对应 spec,避免边实现边设计
|
4. **Scope 蔓延**:v0.3 新增 `engine/` `vector/` `document/` 三个模块,功能覆盖扩展到多 Agent 基础系统。始终保持 trait + reference impl 的边界,业务循环留给上层
|
||||||
5. **Phase 4 抽象化边界**:AG Core 定位为"支持库"而非"Agent 产品",Phase 4(4a/4b/4c)需严格控制范围——只暴露 trait + 最小 reference impl,业务循环(多轮 turn 编排、对话记忆自动回写、Task 拆解策略)留给上层应用。`SessionMemory`(Phase 4c)提供信息桥接通道但不实现 context 切换逻辑。多 context 切换管理延后至 v0.2+。详细设计决策见 `docs/7-agent-runtime.md`
|
5. **API 稳定性**:v0.3 引入 `Checkpointer`、`SessionManager`、`VectorStore` 等新公开 API,v0.2 已有的 `#[non_exhaustive]` 和 `#[deprecated]` 机制继续沿用
|
||||||
6. **参考项目语言差异**:OpenClaw / Hermes / OpenHarness 均为 Python/TypeScript 实现,OpenHuman 虽是 Rust + Tauri 但定位是桌面应用。借鉴时**只取架构模式**,不照搬具体实现(如 Pydantic 工具校验、SQLite Memory Tree、Node+Python 双进程等)
|
6. **Checkpointer 存储效率**:v0.3 使用全量 JSON 序列化存储 checkpoint,每轮对话约几百 KB。`fork` 从历史 checkpoint 创建新 session 时也会复制全量。等实际使用中发现存储瓶颈时再改为增量模式
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## 下一步行动
|
## 下一步行动
|
||||||
|
|
||||||
1. **Phase 4c 已完成**:Phase 4a + 4b + 4c 已交付(116 测试通过,0 clippy 警告)。可启动 v0.2+ 扩展评估(如多 Context 切换、Multi-Agent 协同等)
|
1. **v0.3.0 Phase 15 启动**:向量存储持久化(`VectorStore` trait + `InMemoryVectorStore` + `PersistentVectorStore` + `RagPipeline` 组合器),基于 Phase 14 的 `Document` 类型构建
|
||||||
2. **Context 切换备忘**:`docs/note-context-switch-design.md` 记录了多 context 切换方案讨论,作为 v0.2+ 扩展项的输入
|
2. **Phase 15-19 顺次交付**:按依赖关系推进向量存储 → 摘要 → 引擎 → 调度 → 知识图谱
|
||||||
3. **参考项目调研沉淀**:已完成 OpenClaw / Hermes / OpenHuman / OpenHarness 横向调研,结果沉淀至 `docs/note-agent-harness-references.md`,作为 v0.2+ 扩展项的输入
|
3. **示例先行**:每完成一个 Phase 立即创建/更新对应示例,确保 `cargo run --example` 可验证
|
||||||
4. **Phase 3 备用设计就绪**:`docs/note-knowledge-graph-design.md` 记录了 KnowledgeGraph、高级评分、RecallBased 淘汰等设计,v0.2+ 记忆扩展可直接参考
|
4. **里程碑追踪**:以 M10(Phase 14)为已达成里程碑,逐 Phase 推进 M11-M15
|
||||||
|
|
||||||
**已完成 / 进行中阶段**:
|
**已完成 / 进行中阶段**:
|
||||||
- ✅ Phase 0 Foundation — 全部交付物已完成
|
- ✅ Phase 0 Foundation — 全部交付物已完成
|
||||||
@@ -338,9 +908,20 @@ graph BT
|
|||||||
- ✅ Phase 4a Core Glue — 全部交付物已完成
|
- ✅ Phase 4a Core Glue — 全部交付物已完成
|
||||||
- ✅ Phase 4b Task Execution — 全部交付物已完成
|
- ✅ Phase 4b Task Execution — 全部交付物已完成
|
||||||
- ✅ Phase 4c Session Memory — 全部交付物已完成
|
- ✅ Phase 4c Session Memory — 全部交付物已完成
|
||||||
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen 适配(方案:`docs/10-llm-provider-refinement.md`、`docs/10a-phase0-types-and-trait.md`、`docs/10b-phase1-provider-adaptation.md`)
|
- ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy)
|
||||||
- ✅ LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除(方案:`docs/10c-phase2-llm-cycle-simplify.md`)
|
- ✅ Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
|
||||||
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、7 个离线示例、README + 错误消息友好化、Roadmap 同步、CHANGELOG 初始化(计划:`docs/11-v0.1-release-plan.md`)
|
- ✅ Phase 7 SqliteStore — `rusqlite 0.32` + WAL 模式 + `Arc<Mutex<Connection>>` + `spawn_blocking`;`memory/store.rs` → `store/{in_memory,sqlite_store}.rs` 模块化;9 个内联测试覆盖 CRUD/upsert/过滤/10×10 并发/持久化 round-trip;`InMemoryStore ↔ SqliteStore` trait-box 互换兼容
|
||||||
|
- ✅ **Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 60 行 + `end_to_end` 246 行),10 个离线示例全部 exit 0;**v0.2.0-rc.1 标签已打**;实施后三方审查发现 6 项问题(1 🔴 + 2 🟡 + 3 💭)已全部修复
|
||||||
|
- ✅ **Phase 9 流式体验增强** — `AgentSession::submit_turn_stream` 流式事件序列 + `LlmCycle::submit_with_tools_stream` spawn + mpsc 状态机 + `StreamEvent::ToolExecutionStarted`/`Completed` 新变体 + 9 单元测试 + 2 集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock 验证 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211;`CycleConfig` 加 `Clone` derive;方案文档 `docs/16-phase9-streaming-experience.md`(821 行)
|
||||||
|
- ✅ **Phase 10 ContextSlot 上下文管理** — `src/agent/context.rs` 新增 `ContextSlot` 核心类型(Full / Focused / Readonly 三种模式,New / Derived / Static 三种来源)+ JSON blob 批次持久化(每 slot 3-4 条 MemoryItem,`slot_config` key 自恢复支持旧版本兼容);`AgentSession` 扩展 slots 字段 + 5 个管理方法(`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot`,自动创建 `"default"` slot,`delete_slot` 双重保护禁止删 default/最后一个);`submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回(`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式"读时过滤"语义不丢失数据);`finalize_turn` 签名变更(新增 `new_messages_from_cycle: Vec<Message>` 参数,返回 `Result<(), AgentError>`);`agent/error.rs` 新增 3 个 Slot 错误变体(`SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`);`examples/context_slot_demo.rs` 新增分支对话示例(法律咨询入口 → 两个派生方向 → 切换 → 隔离验证 → 删除保护);方案文档 `docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复);全量 211 → 254(+43 新测试),clippy 0 警告,doc 0 warning,11 个离线示例全部 exit 0
|
||||||
|
- ✅ **Phase 11 测试与检索补强** — `src/memory/vector.rs` 新增 `VectorRetriever` trait(index + search 抽象)+ `InMemoryVectorRetriever` 引用实现(HashMap + 全量余弦相似度扫描 + 零依赖 `dot()`),6 个内联测试覆盖 basic/empty/zero-vector/k=0/2 个并发;wiremock Provider roundtrip 测试 12 个(OpenAI 8 + Anthropic 4)覆盖请求体/header/401/429/500/529/流式 usage-only/流式错误/ToolUse/结构化错误体;`MemoryStore` 并发测试 5 个(InMemoryStore 3 + SqliteStore 2)覆盖 100 并发写、5 写+5 读混合 2 秒、15 写者容量淘汰;`openai.rs` `handle_error_response` 修复 429 retry-after 解析(5 行,与 anthropic 对齐);方案文档 `docs/18-phase11-testing-and-retrieval.md`(647 行,含 10 项架构决策 + 2 条实施偏差记录 #6 mid-stream mock 模式 + #7 retry-after 修复);全量 254 → 277(+23 新测试),clippy 0 警告,doc 0 warning,并发测试 3 次稳定无 flaky
|
||||||
|
- ✅ **Phase 13 热身清理 + ContextSlot fork/merge** — 3 个旧 types 文件删除(`request.rs` 187 行 + `response.rs` 177 行 + `old_stream.rs` 45 行),所有 OpenAI wire-format 类型迁入 `provider/openai.rs` 可见性 `pub(crate)`(Breaking Change:原 `agcore::llm::types::OpenaiChatRequest/Response/Chunk` 公共 re-export 路径已删除);`ChatResponse` 自 v0.1.0 标记 `#[deprecated]` 后在 Phase 13 整体删除;`ToolChoice` 从 `request.rs` 迁入 `tool.rs`(公共 `agcore::llm::types::ToolChoice` 路径不变);`ContextSlot::fork()` 派生独立子 slot(`SlotSource::Derived { parent_id, strategy }` 血缘可追溯)+ `ContextSlot::merge(child, MergeStrategy)` 合入父 slot(`Append` / `Replace` 两种策略,`#[non_exhaustive]` 为 Phase 16 `Summarize` 预留);`MergeStrategy` 防御性检查(self-merge / 跨 session / Readonly 目标全部阻断);`AgentSession::derive_slot` 重构复用 `fork()` 消除重复;`agent.rs` 追加 `MergeStrategy` re-export;9 个 fork/merge 内联测试覆盖 happy path 与 error path;`stream.rs` 简化为 module doc + `pub use` 重导出(保持 `use crate::llm::stream::StreamEvent` 路径兼容);方案文档 `docs/19-phase13-cleanup-and-fork-merge.md`(640 行);全量 277 → 286(+9 新测试),clippy 0 警告,doc 0 warning
|
||||||
|
- ✅ Provider IR 重构 — 统一类型系统 + OpenAI/Anthropic/DeepSeek/Qwen/Ollama 适配
|
||||||
|
- ✅ LlmCycle 简化 — IR 消息类型切换 + Phase 0 桥接层移除
|
||||||
|
- ✅ v0.1 Release — 技术债扫清、MockProvider 公开化、8 个离线示例(含 `simple_visit`)、README + 错误消息友好化、CHANGELOG 初始化
|
||||||
|
- ✅ **v0.2 规划细化完成** — 8 个增量 Phase(Phase 5-12),17 个可验证 Step,覆盖 P0-P2 全部 12 项功能 + ContextSlot
|
||||||
|
- ✅ **v0.3.0 Phase 13 完成** — 技术债清理(3 旧 types 文件 + ChatResponse 删除)+ ContextSlot fork/merge(9 新测试),M9 里程碑达成
|
||||||
|
- ✅ **v0.3.0 Phase 14 完成** — Document 类型(id/content/metadata/mime_type)+ `RecursiveCharacterSplitter` 两阶段算法(按 separator 优先级递归分割 + 贪心合并 overlap,全部 `chars_len()` 字符级比较)+ `Embedding` trait(async + `LlmError` 复用)+ `MockEmbedding`(sin-hash 零依赖伪随机 + L2 归一化)+ 19 Document 测试 + 6 Embedding 测试(含 1 个 split_multibyte_utf8_boundary CJK 边界测试);`src/document.rs`(580 行)+ `src/llm/embedding.rs`(183 行)+ `examples/document_demo.rs`(74 行);`pub use document::Document` 在 lib.rs 重导出;CJK 分隔符(`。`/`?`/`!`)加入 `DEFAULT_SEPARATORS`;方案文档 `docs/20-phase14-document-and-embedding.md`(1417 行);全量 286 → 313(+27 新测试,0 失败),clippy 0 警告,doc 0 warning,零新外部依赖;M10 里程碑达成;Phase 15-19 共 5 个增量 Phase 待实施(向量存储 → 摘要 → 引擎 → 调度 → 知识图谱)
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -15,9 +15,9 @@ use std::sync::Arc;
|
|||||||
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||||||
use agcore::llm::hooks::HookExecutor;
|
use agcore::llm::hooks::HookExecutor;
|
||||||
use agcore::llm::mock::MockProvider;
|
use agcore::llm::mock::MockProvider;
|
||||||
|
use agcore::llm::types::Usage;
|
||||||
use agcore::llm::types::message::{ContentBlock, Message};
|
use agcore::llm::types::message::{ContentBlock, Message};
|
||||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||||
use agcore::llm::types::Usage;
|
|
||||||
use agcore::tools::ToolRegistry;
|
use agcore::tools::ToolRegistry;
|
||||||
|
|
||||||
/// 计算器角色 Agent。
|
/// 计算器角色 Agent。
|
||||||
@@ -72,7 +72,10 @@ async fn main() {
|
|||||||
|
|
||||||
// 4. 提交第一轮
|
// 4. 提交第一轮
|
||||||
println!("=== 提交第 1 轮 ===");
|
println!("=== 提交第 1 轮 ===");
|
||||||
let resp = session.submit_turn("1+1=?").await.expect("submit_turn 失败");
|
let resp = session
|
||||||
|
.submit_turn("1+1=?")
|
||||||
|
.await
|
||||||
|
.expect("submit_turn 失败");
|
||||||
println!("LLM: {}", resp.text());
|
println!("LLM: {}", resp.text());
|
||||||
session
|
session
|
||||||
.set_session_data("last_q", "1+1=?")
|
.set_session_data("last_q", "1+1=?")
|
||||||
@@ -107,11 +110,7 @@ async fn main() {
|
|||||||
|
|
||||||
// 8. 跨 session 数据隔离验证
|
// 8. 跨 session 数据隔离验证
|
||||||
println!("=== 数据隔离验证 ===");
|
println!("=== 数据隔离验证 ===");
|
||||||
let other = AgentSession::new(
|
let other = AgentSession::new(Arc::new(CalculatorAgent), "other-session", bundle);
|
||||||
Arc::new(CalculatorAgent),
|
|
||||||
"other-session",
|
|
||||||
bundle,
|
|
||||||
);
|
|
||||||
assert!(
|
assert!(
|
||||||
other.get_session_data("last_q").await.unwrap().is_none(),
|
other.get_session_data("last_q").await.unwrap().is_none(),
|
||||||
"新会话不应看到旧 session 的 last_q"
|
"新会话不应看到旧 session 的 last_q"
|
||||||
|
|||||||
@@ -0,0 +1,161 @@
|
|||||||
|
//! context_slot_demo —— 多上下文槽位管理示例。
|
||||||
|
//!
|
||||||
|
//! 场景:法律咨询入口 → 派生两个独立探索方向 → 切换 → 隔离验证 → 删除。
|
||||||
|
//!
|
||||||
|
//! 展示:
|
||||||
|
//! - 默认 slot 自动创建
|
||||||
|
//! - 多 slot 间的消息隔离
|
||||||
|
//! - 派生 slot 从父 slot 复制消息
|
||||||
|
//! - 删除非 default slot 后自动回退到 default
|
||||||
|
//!
|
||||||
|
//! 运行:`cargo run --example context_slot_demo`(离线,零配置)
|
||||||
|
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||||||
|
use agcore::llm::hooks::HookExecutor;
|
||||||
|
use agcore::llm::mock::MockProvider;
|
||||||
|
use agcore::llm::provider::LlmProvider;
|
||||||
|
use agcore::llm::types::message::{ContentBlock, Message};
|
||||||
|
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||||
|
use agcore::llm::types::Usage;
|
||||||
|
use agcore::tools::ToolRegistry;
|
||||||
|
|
||||||
|
struct LegalAdvisor;
|
||||||
|
|
||||||
|
impl Agent for LegalAdvisor {
|
||||||
|
fn name(&self) -> &str {
|
||||||
|
"legal-advisor"
|
||||||
|
}
|
||||||
|
fn system_prompt(&self) -> Option<&str> {
|
||||||
|
Some("你是法律顾问。请用一句话回答用户问题。")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 构造一个简单的 Assistant 响应(用于 MockProvider)。
|
||||||
|
fn assistant_resp(text: &str) -> MessageResponse {
|
||||||
|
MessageResponse {
|
||||||
|
id: String::new(),
|
||||||
|
model: "mock".into(),
|
||||||
|
message: Message::Assistant {
|
||||||
|
content: vec![ContentBlock::Text { text: text.into() }],
|
||||||
|
},
|
||||||
|
usage: Usage::from_input_output(5, 5),
|
||||||
|
stop_reason: StopReason::Stop,
|
||||||
|
extra: Default::default(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::main]
|
||||||
|
async fn main() {
|
||||||
|
// 1. 构造 session(自动包含 default slot)
|
||||||
|
let provider: Arc<dyn LlmProvider> = Arc::new(MockProvider::new(vec![
|
||||||
|
assistant_resp("您好,我可以帮您处理法律问题。"),
|
||||||
|
assistant_resp("管辖权问题:建议选择合同签订地法院。"),
|
||||||
|
assistant_resp("条款修改:建议将上限调整为 80 万。"),
|
||||||
|
assistant_resp("已回到主对话。"),
|
||||||
|
]));
|
||||||
|
let bundle = Arc::new(
|
||||||
|
AgentBuilder::new()
|
||||||
|
.provider(provider)
|
||||||
|
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||||
|
.hook_executor(Arc::new(HookExecutor::new()))
|
||||||
|
.build()
|
||||||
|
.unwrap(),
|
||||||
|
);
|
||||||
|
let mut session = AgentSession::new(Arc::new(LegalAdvisor), "legal-001", bundle);
|
||||||
|
|
||||||
|
println!("=== 1. 默认 slot 自动创建 ===");
|
||||||
|
assert_eq!(session.current_slot_id(), "default");
|
||||||
|
let slots: Vec<_> = session.list_slots().collect();
|
||||||
|
println!("初始 slots: {slots:?}");
|
||||||
|
assert_eq!(slots.len(), 1);
|
||||||
|
assert!(slots.contains(&&"default".to_string()));
|
||||||
|
|
||||||
|
println!("\n=== 2. 在 default slot 中提交一轮 ===");
|
||||||
|
let r1 = session.submit_turn("我需要法律援助").await.unwrap();
|
||||||
|
println!("default slot response: {}", r1.text());
|
||||||
|
|
||||||
|
println!("\n=== 3. 派生两个独立探索方向的 slot ===");
|
||||||
|
session
|
||||||
|
.derive_slot("option_jurisdiction", "default", agcore::agent::DeriveStrategy::Full)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
session
|
||||||
|
.derive_slot("option_amendment", "default", agcore::agent::DeriveStrategy::Full)
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||||
|
println!("派生后 slots: {slots:?}");
|
||||||
|
assert_eq!(slots.len(), 3);
|
||||||
|
|
||||||
|
println!("\n=== 4. 切到 option_jurisdiction 并提交 ===");
|
||||||
|
session.switch_slot("option_jurisdiction").await.unwrap();
|
||||||
|
assert_eq!(session.current_slot_id(), "option_jurisdiction");
|
||||||
|
let r2 = session.submit_turn("如果用户质疑管辖权?").await.unwrap();
|
||||||
|
println!("option_jurisdiction response: {}", r2.text());
|
||||||
|
|
||||||
|
println!("\n=== 5. 切到 option_amendment 并提交 ===");
|
||||||
|
session.switch_slot("option_amendment").await.unwrap();
|
||||||
|
let r3 = session.submit_turn("用户要求提高赔偿上限?").await.unwrap();
|
||||||
|
println!("option_amendment response: {}", r3.text());
|
||||||
|
|
||||||
|
println!("\n=== 6. 切回 default,验证消息隔离 ===");
|
||||||
|
session.switch_slot("default").await.unwrap();
|
||||||
|
let r4 = session.submit_turn("汇总一下我们的讨论").await.unwrap();
|
||||||
|
println!("default response: {}", r4.text());
|
||||||
|
// 验证 default slot 不包含 option_jurisdiction 的"管辖权"问题
|
||||||
|
let (_, default_slot) = session.slots().find(|(id, _)| *id == "default").unwrap();
|
||||||
|
let default_has_jurisdiction = default_slot
|
||||||
|
.messages
|
||||||
|
.iter()
|
||||||
|
.any(|m| message_contains(m, "管辖权"));
|
||||||
|
assert!(
|
||||||
|
!default_has_jurisdiction,
|
||||||
|
"default slot 不应包含 option_jurisdiction 的消息"
|
||||||
|
);
|
||||||
|
|
||||||
|
println!("\n=== 7. 删除 option_amendment,验证回退到 default ===");
|
||||||
|
session.delete_slot("option_amendment").await.unwrap();
|
||||||
|
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||||
|
println!("删除后 slots: {slots:?}");
|
||||||
|
assert!(!slots.contains(&"option_amendment".to_string()));
|
||||||
|
assert_eq!(slots.len(), 2);
|
||||||
|
|
||||||
|
println!("\n=== 8. 切到 option_jurisdiction 并删除,验证 current 回退 ===");
|
||||||
|
session.switch_slot("option_jurisdiction").await.unwrap();
|
||||||
|
session.delete_slot("option_jurisdiction").await.unwrap();
|
||||||
|
assert_eq!(session.current_slot_id(), "default");
|
||||||
|
let slots: Vec<_> = session.list_slots().cloned().collect();
|
||||||
|
println!("删除后 slots: {slots:?}");
|
||||||
|
assert_eq!(slots.len(), 1);
|
||||||
|
assert_eq!(slots[0], "default");
|
||||||
|
|
||||||
|
println!("\n=== 9. 验证 delete_slot 保护逻辑 ===");
|
||||||
|
let err = session.delete_slot("default").await.unwrap_err();
|
||||||
|
println!("删除 default 返回错误: {err}");
|
||||||
|
assert!(matches!(err, agcore::agent::AgentError::Config(_)));
|
||||||
|
|
||||||
|
println!("\n✓ context_slot_demo 完成");
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 检查 Message 是否包含指定文本(提取第一个 Text block)。
|
||||||
|
fn message_contains(msg: &Message, needle: &str) -> bool {
|
||||||
|
use agcore::llm::types::message::ContentBlock;
|
||||||
|
let blocks = match msg {
|
||||||
|
Message::System { content }
|
||||||
|
| Message::User { content }
|
||||||
|
| Message::Assistant { content } => content,
|
||||||
|
Message::UserImage { .. } => return false,
|
||||||
|
Message::ToolResult { content, .. } => content,
|
||||||
|
_ => return false,
|
||||||
|
};
|
||||||
|
for block in blocks {
|
||||||
|
if let ContentBlock::Text { text } = block
|
||||||
|
&& text.contains(needle)
|
||||||
|
{
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
false
|
||||||
|
}
|
||||||
@@ -29,6 +29,7 @@ fn message_text(msg: &Message) -> &str {
|
|||||||
.next()
|
.next()
|
||||||
.unwrap_or(""),
|
.unwrap_or(""),
|
||||||
Message::UserImage { .. } => "[image]",
|
Message::UserImage { .. } => "[image]",
|
||||||
|
_ => "",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -80,11 +81,8 @@ async fn main() {
|
|||||||
// 3. 多角色混合 + clear
|
// 3. 多角色混合 + clear
|
||||||
println!("\n=== 多角色写入 + clear ===");
|
println!("\n=== 多角色写入 + clear ===");
|
||||||
let store3 = Arc::new(InMemoryStore::new());
|
let store3 = Arc::new(InMemoryStore::new());
|
||||||
let mut memory3 = ConversationMemory::new(
|
let mut memory3 =
|
||||||
store3,
|
ConversationMemory::new(store3, "session-3", ConversationMemoryConfig::default());
|
||||||
"session-3",
|
|
||||||
ConversationMemoryConfig::default(),
|
|
||||||
);
|
|
||||||
memory3
|
memory3
|
||||||
.add_message(Message::user_text("你好"))
|
.add_message(Message::user_text("你好"))
|
||||||
.await
|
.await
|
||||||
@@ -98,7 +96,9 @@ async fn main() {
|
|||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
memory3
|
memory3
|
||||||
.add_message(Message::assistant("我无法查询实时天气,但你可以查看天气应用。"))
|
.add_message(Message::assistant(
|
||||||
|
"我无法查询实时天气,但你可以查看天气应用。",
|
||||||
|
))
|
||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
println!(
|
println!(
|
||||||
@@ -119,16 +119,8 @@ async fn main() {
|
|||||||
// 4. Session 隔离
|
// 4. Session 隔离
|
||||||
println!("\n=== Session 隔离(共用 InMemoryStore)===");
|
println!("\n=== Session 隔离(共用 InMemoryStore)===");
|
||||||
let store4 = Arc::new(InMemoryStore::new());
|
let store4 = Arc::new(InMemoryStore::new());
|
||||||
let mut a = ConversationMemory::new(
|
let mut a = ConversationMemory::new(store4.clone(), "s-a", ConversationMemoryConfig::default());
|
||||||
store4.clone(),
|
let mut b = ConversationMemory::new(store4.clone(), "s-b", ConversationMemoryConfig::default());
|
||||||
"s-a",
|
|
||||||
ConversationMemoryConfig::default(),
|
|
||||||
);
|
|
||||||
let mut b = ConversationMemory::new(
|
|
||||||
store4.clone(),
|
|
||||||
"s-b",
|
|
||||||
ConversationMemoryConfig::default(),
|
|
||||||
);
|
|
||||||
a.add_message(Message::user_text("A 的消息")).await.unwrap();
|
a.add_message(Message::user_text("A 的消息")).await.unwrap();
|
||||||
b.add_message(Message::user_text("B 的消息")).await.unwrap();
|
b.add_message(Message::user_text("B 的消息")).await.unwrap();
|
||||||
println!(
|
println!(
|
||||||
|
|||||||
+10
-15
@@ -17,7 +17,7 @@ use agcore::tools::{
|
|||||||
ToolRegistry,
|
ToolRegistry,
|
||||||
};
|
};
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use serde_json::{json, Value};
|
use serde_json::{Value, json};
|
||||||
|
|
||||||
/// 天气查询工具 —— 模拟根据城市返回天气数据。
|
/// 天气查询工具 —— 模拟根据城市返回天气数据。
|
||||||
struct WeatherTool;
|
struct WeatherTool;
|
||||||
@@ -42,11 +42,7 @@ impl BaseTool for WeatherTool {
|
|||||||
fn required_permissions(&self) -> Vec<Permission> {
|
fn required_permissions(&self) -> Vec<Permission> {
|
||||||
vec![Permission::Network]
|
vec![Permission::Network]
|
||||||
}
|
}
|
||||||
async fn execute(
|
async fn execute(&self, args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||||
&self,
|
|
||||||
args: Value,
|
|
||||||
_ctx: &ToolContext<'_>,
|
|
||||||
) -> Result<Value, ToolError> {
|
|
||||||
let city = args["city"].as_str().unwrap_or("未知");
|
let city = args["city"].as_str().unwrap_or("未知");
|
||||||
// 模拟查询:根据城市名给出不同温度
|
// 模拟查询:根据城市名给出不同温度
|
||||||
let (temperature, condition) = match city {
|
let (temperature, condition) = match city {
|
||||||
@@ -84,11 +80,7 @@ impl BaseTool for DeleteFileTool {
|
|||||||
fn required_permissions(&self) -> Vec<Permission> {
|
fn required_permissions(&self) -> Vec<Permission> {
|
||||||
vec![Permission::Delete]
|
vec![Permission::Delete]
|
||||||
}
|
}
|
||||||
async fn execute(
|
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||||
&self,
|
|
||||||
_args: Value,
|
|
||||||
_ctx: &ToolContext<'_>,
|
|
||||||
) -> Result<Value, ToolError> {
|
|
||||||
Ok(json!({"deleted": true}))
|
Ok(json!({"deleted": true}))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -138,9 +130,8 @@ async fn main() {
|
|||||||
|
|
||||||
// 5. 权限检查:默认 PermissionConfig 黑名单含 Delete
|
// 5. 权限检查:默认 PermissionConfig 黑名单含 Delete
|
||||||
println!("\n=== 权限检查(默认 PermissionConfig,denied = [Delete, Shell])===");
|
println!("\n=== 权限检查(默认 PermissionConfig,denied = [Delete, Shell])===");
|
||||||
let mut registry_with_checker = ToolRegistry::new().with_permission_checker(PermissionChecker::new(
|
let mut registry_with_checker = ToolRegistry::new()
|
||||||
PermissionConfig::default(),
|
.with_permission_checker(PermissionChecker::new(PermissionConfig::default()));
|
||||||
));
|
|
||||||
registry_with_checker
|
registry_with_checker
|
||||||
.register(Arc::new(WeatherTool) as ToolRef)
|
.register(Arc::new(WeatherTool) as ToolRef)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -155,7 +146,11 @@ async fn main() {
|
|||||||
.unwrap();
|
.unwrap();
|
||||||
println!(
|
println!(
|
||||||
"get_weather 权限检查: {}",
|
"get_weather 权限检查: {}",
|
||||||
if r.output.is_ok() { "通过 ✓" } else { "阻断 ✗" }
|
if r.output.is_ok() {
|
||||||
|
"通过 ✓"
|
||||||
|
} else {
|
||||||
|
"阻断 ✗"
|
||||||
|
}
|
||||||
);
|
);
|
||||||
|
|
||||||
// delete_file 声明 Delete → 在 denied 列表 → 阻断
|
// delete_file 声明 Delete → 在 denied 列表 → 阻断
|
||||||
|
|||||||
@@ -0,0 +1,70 @@
|
|||||||
|
//! document_demo —— Document + RecursiveCharacterSplitter + MockEmbedding + RagPipeline 完整衔接示例。
|
||||||
|
//!
|
||||||
|
//! 演示 RAG 管线:
|
||||||
|
//! 1. 创建多段落 Document
|
||||||
|
//! 2. RecursiveCharacterSplitter 分割为 chunk
|
||||||
|
//! 3. RagPipeline.ingest() 自动嵌入并存储
|
||||||
|
//! 4. RagPipeline.retrieve() 做语义检索
|
||||||
|
//!
|
||||||
|
//! 运行:`cargo run --example document_demo`(离线,零配置)
|
||||||
|
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use agcore::document::{Document, RecursiveCharacterSplitter};
|
||||||
|
use agcore::llm::embedding::{Embedding, MockEmbedding};
|
||||||
|
use agcore::memory::{InMemoryVectorStore, RagPipeline, VectorStore};
|
||||||
|
|
||||||
|
#[tokio::main]
|
||||||
|
async fn main() {
|
||||||
|
agcore::init_tracing();
|
||||||
|
|
||||||
|
// 1. 创建多段落 Document(含中英文混合)
|
||||||
|
let doc = Document::new(
|
||||||
|
"rust-intro",
|
||||||
|
"Rust 是一门系统编程语言,注重安全、并发和性能。\n\n\
|
||||||
|
Rust 通过所有权系统管理内存,无需垃圾回收器。\
|
||||||
|
所有权规则让内存安全在编译期就能得到保证。\n\n\
|
||||||
|
Rust 的并发模型通过类型系统区分线程间共享与独占数据,\
|
||||||
|
避免数据竞争。Send 和 Sync 两个 trait 标记了类型的线程安全性。\n\n\
|
||||||
|
Rust 的性能与 C/C++ 相当,但提供了更现代的开发体验。\
|
||||||
|
Cargo 是官方的构建系统和包管理器,使用简单直观。",
|
||||||
|
"text/markdown",
|
||||||
|
);
|
||||||
|
|
||||||
|
println!("输入文档: {} 字符", doc.content.chars().count());
|
||||||
|
|
||||||
|
// 2. 构造 RAG 管线(嵌入器 + 向量存储 + 分割器)
|
||||||
|
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
|
||||||
|
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(200, 30);
|
||||||
|
let pipeline = RagPipeline::new(
|
||||||
|
Arc::clone(&embedder),
|
||||||
|
Arc::clone(&store),
|
||||||
|
Some(splitter),
|
||||||
|
);
|
||||||
|
|
||||||
|
// 3. 一次性 ingest:自动 split → embed → add
|
||||||
|
pipeline.ingest(std::slice::from_ref(&doc)).await.unwrap();
|
||||||
|
|
||||||
|
// 4. 模拟查询:复用第一个 chunk 的 content 作为查询文本
|
||||||
|
let chunks_in_store = store.search(&[1.0, 0.0, 0.0, 0.0], 1).await.unwrap();
|
||||||
|
assert!(!chunks_in_store.is_empty(), "ingest 后 store 应有数据");
|
||||||
|
let query_text = &chunks_in_store[0].0.content;
|
||||||
|
|
||||||
|
let results = pipeline.retrieve(query_text, 3).await.unwrap();
|
||||||
|
println!("\nTop 3 检索结果(与第一个 chunk 相似):");
|
||||||
|
for (doc, score) in &results {
|
||||||
|
println!(
|
||||||
|
" id={}, score={:.4}, content={}",
|
||||||
|
doc.id, score, doc.content
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
assert!(!results.is_empty(), "至少应返回 1 条检索结果");
|
||||||
|
assert!(
|
||||||
|
results[0].0.id.starts_with("rust-intro:chunk:0000"),
|
||||||
|
"Top 1 应为 chunk 0 自身"
|
||||||
|
);
|
||||||
|
|
||||||
|
println!("\n✓ document_demo 完成");
|
||||||
|
}
|
||||||
@@ -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 类型,默认 OpenaiChat(OpenAI / 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✓ 端到端演示完成");
|
||||||
|
}
|
||||||
@@ -38,14 +38,26 @@ async fn main() {
|
|||||||
let ks = KnowledgeStore::new(store);
|
let ks = KnowledgeStore::new(store);
|
||||||
|
|
||||||
let pages = vec![
|
let pages = vec![
|
||||||
make_page("rust-1", "Rust 入门", "Rust 是一门系统级编程语言,注重安全性与并发。"),
|
make_page(
|
||||||
make_page("python-1", "Python 简介", "Python 是一门动态类型的高级编程语言。"),
|
"rust-1",
|
||||||
|
"Rust 入门",
|
||||||
|
"Rust 是一门系统级编程语言,注重安全性与并发。",
|
||||||
|
),
|
||||||
|
make_page(
|
||||||
|
"python-1",
|
||||||
|
"Python 简介",
|
||||||
|
"Python 是一门动态类型的高级编程语言。",
|
||||||
|
),
|
||||||
make_page(
|
make_page(
|
||||||
"langgraph-1",
|
"langgraph-1",
|
||||||
"LangGraph 框架",
|
"LangGraph 框架",
|
||||||
"LangGraph 是 LangChain 的状态图扩展,用于构建多步 Agent。",
|
"LangGraph 是 LangChain 的状态图扩展,用于构建多步 Agent。",
|
||||||
),
|
),
|
||||||
make_page("rust-async", "Rust 异步编程", "Rust 异步基于 tokio 与 futures 抽象。"),
|
make_page(
|
||||||
|
"rust-async",
|
||||||
|
"Rust 异步编程",
|
||||||
|
"Rust 异步基于 tokio 与 futures 抽象。",
|
||||||
|
),
|
||||||
];
|
];
|
||||||
for p in &pages {
|
for p in &pages {
|
||||||
ks.add_page(p.clone()).await.expect("保存页面失败");
|
ks.add_page(p.clone()).await.expect("保存页面失败");
|
||||||
@@ -62,14 +74,8 @@ async fn main() {
|
|||||||
let result = retriever.retrieve("Rust 异步").await.unwrap();
|
let result = retriever.retrieve("Rust 异步").await.unwrap();
|
||||||
println!("query: {}", result.query);
|
println!("query: {}", result.query);
|
||||||
for item in &result.items {
|
for item in &result.items {
|
||||||
println!(
|
println!(" 命中: {} (score={:.3})", item.page.title, item.score);
|
||||||
" 命中: {} (score={:.3})",
|
assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
|
||||||
item.page.title, item.score
|
|
||||||
);
|
|
||||||
assert!(
|
|
||||||
(0.0..=1.0).contains(&item.score),
|
|
||||||
"score 应在 [0, 1] 区间"
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
assert!(!result.items.is_empty(), "应至少命中一个页面");
|
assert!(!result.items.is_empty(), "应至少命中一个页面");
|
||||||
|
|
||||||
@@ -85,14 +91,8 @@ async fn main() {
|
|||||||
min_score: 0.5,
|
min_score: 0.5,
|
||||||
};
|
};
|
||||||
let retriever2 = MemoryRetriever::new(ks2, cfg);
|
let retriever2 = MemoryRetriever::new(ks2, cfg);
|
||||||
let result = retriever2
|
let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
|
||||||
.retrieve("完全不相关的火锅配方")
|
println!("无关 query → items.len = {} (期望 0)", result.items.len());
|
||||||
.await
|
|
||||||
.unwrap();
|
|
||||||
println!(
|
|
||||||
"无关 query → items.len = {} (期望 0)",
|
|
||||||
result.items.len()
|
|
||||||
);
|
|
||||||
assert!(result.items.is_empty());
|
assert!(result.items.is_empty());
|
||||||
|
|
||||||
// 4. max_results 截断
|
// 4. max_results 截断
|
||||||
|
|||||||
@@ -11,7 +11,7 @@
|
|||||||
|
|
||||||
use agcore::llm::types::message::{ContentBlock, Message};
|
use agcore::llm::types::message::{ContentBlock, Message};
|
||||||
use agcore::prompt::{
|
use agcore::prompt::{
|
||||||
validate_messages, PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext,
|
PromptComposer, PromptTemplate, PromptTemplateRegistry, TemplateContext, validate_messages,
|
||||||
};
|
};
|
||||||
|
|
||||||
fn message_text(msg: &Message) -> String {
|
fn message_text(msg: &Message) -> String {
|
||||||
@@ -27,15 +27,15 @@ fn message_text(msg: &Message) -> String {
|
|||||||
})
|
})
|
||||||
.collect(),
|
.collect(),
|
||||||
Message::UserImage { .. } => "[image]".into(),
|
Message::UserImage { .. } => "[image]".into(),
|
||||||
|
_ => String::new(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn main() {
|
fn main() {
|
||||||
// 1. PromptTemplate::compile + render —— 直接构造模板
|
// 1. PromptTemplate::compile + render —— 直接构造模板
|
||||||
println!("=== PromptTemplate::compile + render ===");
|
println!("=== PromptTemplate::compile + render ===");
|
||||||
let tpl = PromptTemplate::compile(
|
let tpl =
|
||||||
"今日 {{location}} 天气:{{condition}},温度 {{temperature}}",
|
PromptTemplate::compile("今日 {{location}} 天气:{{condition}},温度 {{temperature}}")
|
||||||
)
|
|
||||||
.expect("编译失败");
|
.expect("编译失败");
|
||||||
let mut ctx = TemplateContext::new();
|
let mut ctx = TemplateContext::new();
|
||||||
ctx.insert("location", "北京");
|
ctx.insert("location", "北京");
|
||||||
@@ -58,7 +58,10 @@ fn main() {
|
|||||||
.register("weather", "今日 {{location}}:{{condition}}")
|
.register("weather", "今日 {{location}}:{{condition}}")
|
||||||
.expect("注册失败");
|
.expect("注册失败");
|
||||||
registry
|
registry
|
||||||
.register("greet", "你好 {{name}}!{{#if formal}} 见到您很荣幸。{{/if}}")
|
.register(
|
||||||
|
"greet",
|
||||||
|
"你好 {{name}}!{{#if formal}} 见到您很荣幸。{{/if}}",
|
||||||
|
)
|
||||||
.expect("注册失败");
|
.expect("注册失败");
|
||||||
|
|
||||||
let mut ctx = TemplateContext::new();
|
let mut ctx = TemplateContext::new();
|
||||||
@@ -88,6 +91,7 @@ fn main() {
|
|||||||
Message::User { .. } | Message::UserImage { .. } => "user",
|
Message::User { .. } | Message::UserImage { .. } => "user",
|
||||||
Message::Assistant { .. } => "assistant",
|
Message::Assistant { .. } => "assistant",
|
||||||
Message::ToolResult { .. } => "tool",
|
Message::ToolResult { .. } => "tool",
|
||||||
|
_ => "unknown",
|
||||||
};
|
};
|
||||||
println!("[{i}] {role}: {}", message_text(m));
|
println!("[{i}] {role}: {}", message_text(m));
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 完成");
|
||||||
|
}
|
||||||
@@ -3,7 +3,7 @@ use std::env;
|
|||||||
use agcore::init_tracing;
|
use agcore::init_tracing;
|
||||||
use agcore::llm::{
|
use agcore::llm::{
|
||||||
cycle::{CycleConfig, LlmCycle},
|
cycle::{CycleConfig, LlmCycle},
|
||||||
provider::{create_provider, ProviderConfig, ProviderType},
|
provider::{ProviderConfig, ProviderType, create_provider},
|
||||||
types::{message::ContentBlock, message::Message, response_v2::MessageResponse},
|
types::{message::ContentBlock, message::Message, response_v2::MessageResponse},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -51,10 +51,11 @@ async fn main() {
|
|||||||
base_url,
|
base_url,
|
||||||
api_key,
|
api_key,
|
||||||
model: model.clone(),
|
model: model.clone(),
|
||||||
|
timeout_secs: 30,
|
||||||
|
max_retries: 3,
|
||||||
};
|
};
|
||||||
|
|
||||||
let provider = create_provider(provider_type, config)
|
let provider = create_provider(provider_type, config).expect("创建 Provider 失败");
|
||||||
.expect("创建 Provider 失败");
|
|
||||||
|
|
||||||
let cycle_config = CycleConfig {
|
let cycle_config = CycleConfig {
|
||||||
model,
|
model,
|
||||||
@@ -63,9 +64,9 @@ async fn main() {
|
|||||||
..CycleConfig::default()
|
..CycleConfig::default()
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![
|
let mut cycle = LlmCycle::new(provider, cycle_config).with_messages(vec![Message::system(
|
||||||
Message::system("你是一个简洁的助手,对于任何问题都是用一句话回答。"),
|
"你是一个简洁的助手,对于任何问题都是用一句话回答。",
|
||||||
]);
|
)]);
|
||||||
|
|
||||||
println!("发送请求...");
|
println!("发送请求...");
|
||||||
|
|
||||||
|
|||||||
@@ -17,9 +17,9 @@ use std::sync::Arc;
|
|||||||
use agcore::llm::cycle::{CycleConfig, LlmCycle};
|
use agcore::llm::cycle::{CycleConfig, LlmCycle};
|
||||||
use agcore::llm::mock::MockProvider;
|
use agcore::llm::mock::MockProvider;
|
||||||
use agcore::llm::provider::LlmProvider;
|
use agcore::llm::provider::LlmProvider;
|
||||||
|
use agcore::llm::types::Usage;
|
||||||
use agcore::llm::types::message::{ContentBlock, Message};
|
use agcore::llm::types::message::{ContentBlock, Message};
|
||||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
|
use agcore::llm::types::response_v2::{MessageResponse, StopReason, StreamEvent};
|
||||||
use agcore::llm::types::Usage;
|
|
||||||
use futures_util::StreamExt;
|
use futures_util::StreamExt;
|
||||||
|
|
||||||
/// 构造预设的纯文本响应。
|
/// 构造预设的纯文本响应。
|
||||||
@@ -99,9 +99,7 @@ async fn main() {
|
|||||||
// 上层 Agent 通过 `match` 或 `?` 处理 `AgentError::Llm(_)`。
|
// 上层 Agent 通过 `match` 或 `?` 处理 `AgentError::Llm(_)`。
|
||||||
println!("\n=== 阶段 2:错误路径(队列耗尽)===");
|
println!("\n=== 阶段 2:错误路径(队列耗尽)===");
|
||||||
let mut cycle = LlmCycle::new_with_arc(dyn_provider, CycleConfig::default());
|
let mut cycle = LlmCycle::new_with_arc(dyn_provider, CycleConfig::default());
|
||||||
let result = cycle
|
let result = cycle.submit_stream("第二次提问".to_string(), vec![]).await;
|
||||||
.submit_stream("第二次提问".to_string(), vec![])
|
|
||||||
.await;
|
|
||||||
match result {
|
match result {
|
||||||
Ok(_) => panic!("阶段 2 必须失败(队列耗尽)"),
|
Ok(_) => panic!("阶段 2 必须失败(队列耗尽)"),
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
|
|||||||
+25
-24
@@ -8,27 +8,13 @@
|
|||||||
//! 5. 错误路径:非法 JSON / 空 steps / 缺字段 → `AgentError::PlanParse`
|
//! 5. 错误路径:非法 JSON / 空 steps / 缺字段 → `AgentError::PlanParse`
|
||||||
//!
|
//!
|
||||||
//! 运行:`cargo run --example task_agent_demo`
|
//! 运行:`cargo run --example task_agent_demo`
|
||||||
//!
|
|
||||||
//! ## 已知技术债(v0.2 迁移指南)
|
|
||||||
//!
|
|
||||||
//! 本示例使用 `#[deprecated]` 标记的旧 wire-format 类型:
|
|
||||||
//! - `ChatResponse`、`OpenaiChatMessage`、`FinishReason` —— `OpenaiChatProvider::chat_inner()`
|
|
||||||
//! 内部转换层仍在使用(参见 `docs/10a-phase0-types-and-trait.md` §2.5.1),
|
|
||||||
//! 故结构体定义保留。
|
|
||||||
//! - `StepStatus::Completed(ChatResponse)` —— 因为 `Step` 的"已完成"变体需携带
|
|
||||||
//! provider 响应,目前沿用旧的 `ChatResponse`。
|
|
||||||
//!
|
|
||||||
//! **触发迁移的条件**:v0.2 引入 IR 层的 `StepResult` / 切换为 `MessageResponse`。
|
|
||||||
//! **迁移路径**:将本文件 `ChatResponse`/`OpenaiChatMessage`/`FinishReason` 替换为
|
|
||||||
//! `MessageResponse`/`Message`/`StopReason`,移除顶部 `#![allow(deprecated)]`。
|
|
||||||
//! 上层应用代码(`TaskAgent` 消费者)也可同步迁移。
|
|
||||||
|
|
||||||
#![allow(deprecated)]
|
use std::collections::HashMap;
|
||||||
|
|
||||||
use agcore::agent::{AgentError, JsonPlanParser, PlanParser, Step, StepStatus};
|
use agcore::agent::{AgentError, JsonPlanParser, PlanParser, Step, StepStatus};
|
||||||
use agcore::llm::types::openai_message::OpenaiChatMessage;
|
use agcore::llm::types::message::Message;
|
||||||
use agcore::llm::types::shared::FinishReason;
|
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||||
use agcore::llm::types::{ChatResponse, Usage};
|
use agcore::llm::types::Usage;
|
||||||
|
|
||||||
#[tokio::main]
|
#[tokio::main]
|
||||||
async fn main() {
|
async fn main() {
|
||||||
@@ -67,21 +53,36 @@ async fn main() {
|
|||||||
assert!(step.status.is_pending());
|
assert!(step.status.is_pending());
|
||||||
|
|
||||||
step.status = StepStatus::Running;
|
step.status = StepStatus::Running;
|
||||||
println!("Running: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
|
println!(
|
||||||
|
"Running: pending={}, terminal={}",
|
||||||
|
step.status.is_pending(),
|
||||||
|
step.status.is_terminal()
|
||||||
|
);
|
||||||
|
|
||||||
step.status = StepStatus::Completed(ChatResponse {
|
step.status = StepStatus::Completed(MessageResponse {
|
||||||
message: OpenaiChatMessage::assistant_text("天气:晴,22°C"),
|
id: String::new(),
|
||||||
|
model: "mock".into(),
|
||||||
|
message: Message::assistant("天气:晴,22°C"),
|
||||||
usage: Usage::from_input_output(5, 10),
|
usage: Usage::from_input_output(5, 10),
|
||||||
stop_reason: Some(FinishReason::Stop),
|
stop_reason: StopReason::Stop,
|
||||||
|
extra: HashMap::new(),
|
||||||
});
|
});
|
||||||
println!("Completed: pending={}, terminal={}", step.status.is_pending(), step.status.is_terminal());
|
println!(
|
||||||
|
"Completed: pending={}, terminal={}",
|
||||||
|
step.status.is_pending(),
|
||||||
|
step.status.is_terminal()
|
||||||
|
);
|
||||||
assert!(step.status.is_terminal());
|
assert!(step.status.is_terminal());
|
||||||
|
|
||||||
// 3. 失败路径
|
// 3. 失败路径
|
||||||
println!("\n=== Step 状态机:失败路径 ===");
|
println!("\n=== Step 状态机:失败路径 ===");
|
||||||
let mut fail_step = Step::new(0, "调用天气 API");
|
let mut fail_step = Step::new(0, "调用天气 API");
|
||||||
fail_step.status = StepStatus::Failed(AgentError::Other("API 不可用".into()));
|
fail_step.status = StepStatus::Failed(AgentError::Other("API 不可用".into()));
|
||||||
println!("Failed: pending={}, terminal={}", fail_step.status.is_pending(), fail_step.status.is_terminal());
|
println!(
|
||||||
|
"Failed: pending={}, terminal={}",
|
||||||
|
fail_step.status.is_pending(),
|
||||||
|
fail_step.status.is_terminal()
|
||||||
|
);
|
||||||
assert!(fail_step.status.is_terminal());
|
assert!(fail_step.status.is_terminal());
|
||||||
|
|
||||||
// 4. 跳过路径
|
// 4. 跳过路径
|
||||||
|
|||||||
+6
-1
@@ -11,6 +11,7 @@
|
|||||||
|
|
||||||
pub mod agent;
|
pub mod agent;
|
||||||
pub mod builder;
|
pub mod builder;
|
||||||
|
pub mod context;
|
||||||
pub mod error;
|
pub mod error;
|
||||||
pub mod runtime;
|
pub mod runtime;
|
||||||
pub mod session;
|
pub mod session;
|
||||||
@@ -20,9 +21,13 @@ pub mod task;
|
|||||||
// 重导出公共 API(按使用频度排序)
|
// 重导出公共 API(按使用频度排序)
|
||||||
pub use agent::Agent;
|
pub use agent::Agent;
|
||||||
pub use builder::AgentBuilder;
|
pub use builder::AgentBuilder;
|
||||||
|
pub use context::{
|
||||||
|
ContextBudget, ContextSlot, DeriveStrategy, FocusedConfig, MergeStrategy, SlotConfig,
|
||||||
|
SlotMeta, SlotMode, SlotSource,
|
||||||
|
};
|
||||||
pub use error::AgentError;
|
pub use error::AgentError;
|
||||||
pub use runtime::{AgentConfig, RuntimeBundle};
|
pub use runtime::{AgentConfig, RuntimeBundle};
|
||||||
pub use session::AgentSession;
|
pub use session::AgentSession;
|
||||||
pub use session_memory::SessionMemory;
|
pub use session_memory::SessionMemory;
|
||||||
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
|
|
||||||
pub use task::JsonPlanParser;
|
pub use task::JsonPlanParser;
|
||||||
|
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
|
||||||
|
|||||||
+2
-4
@@ -7,14 +7,12 @@
|
|||||||
//! - **不绑定业务循环**:`submit_turn` 在 `AgentSession` 上,不在 trait 上
|
//! - **不绑定业务循环**:`submit_turn` 在 `AgentSession` 上,不在 trait 上
|
||||||
|
|
||||||
use crate::agent::runtime::RuntimeBundle;
|
use crate::agent::runtime::RuntimeBundle;
|
||||||
#[allow(deprecated)]
|
use crate::llm::types::tool::ToolDef;
|
||||||
use crate::llm::types::ToolDefinition;
|
|
||||||
|
|
||||||
/// Agent 角色抽象。
|
/// Agent 角色抽象。
|
||||||
///
|
///
|
||||||
/// 实现此 trait 即可接入 Agent Runtime。典型实现是 struct 持有静态配置(name、system prompt 模板),
|
/// 实现此 trait 即可接入 Agent Runtime。典型实现是 struct 持有静态配置(name、system prompt 模板),
|
||||||
/// 也可以是基于配置动态生成的轻量实现。
|
/// 也可以是基于配置动态生成的轻量实现。
|
||||||
#[allow(deprecated)]
|
|
||||||
pub trait Agent: Send + Sync {
|
pub trait Agent: Send + Sync {
|
||||||
/// 角色名(用于日志、调试、UI 展示)。
|
/// 角色名(用于日志、调试、UI 展示)。
|
||||||
fn name(&self) -> &str;
|
fn name(&self) -> &str;
|
||||||
@@ -26,7 +24,7 @@ pub trait Agent: Send + Sync {
|
|||||||
///
|
///
|
||||||
/// **默认实现**:从 `bundle.tool_registry` 取全部工具(最常用模式)。
|
/// **默认实现**:从 `bundle.tool_registry` 取全部工具(最常用模式)。
|
||||||
/// **子 trait / 具体实现可覆盖**:做白名单、过滤、按状态动态调整等。
|
/// **子 trait / 具体实现可覆盖**:做白名单、过滤、按状态动态调整等。
|
||||||
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDefinition> {
|
fn tool_definitions(&self, bundle: &RuntimeBundle) -> Vec<ToolDef> {
|
||||||
bundle.tool_registry.definitions()
|
bundle.tool_registry.definitions()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -92,15 +92,17 @@ impl AgentBuilder {
|
|||||||
/// `AgentError::Config(...)`,提示调用 `.provider(...)` / `.tool_registry(...)` /
|
/// `AgentError::Config(...)`,提示调用 `.provider(...)` / `.tool_registry(...)` /
|
||||||
/// `.hook_executor(...)` 补齐。不 panic。
|
/// `.hook_executor(...)` 补齐。不 panic。
|
||||||
pub fn build(self) -> Result<RuntimeBundle, AgentError> {
|
pub fn build(self) -> Result<RuntimeBundle, AgentError> {
|
||||||
let provider = self
|
let provider = self.provider.ok_or_else(|| {
|
||||||
.provider
|
AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into())
|
||||||
.ok_or_else(|| AgentError::Config("缺少 LLM provider,请先调用 .provider(...)".into()))?;
|
})?;
|
||||||
let tool_registry = self
|
let tool_registry = self
|
||||||
.tool_registry
|
.tool_registry
|
||||||
.ok_or_else(|| AgentError::Config("缺少 tool_registry,请先调用 .tool_registry(...)(即使是空 ToolRegistry 也需要传入)".into()))?;
|
.ok_or_else(|| AgentError::Config("缺少 tool_registry,请先调用 .tool_registry(...)(即使是空 ToolRegistry 也需要传入)".into()))?;
|
||||||
let hook_executor = self
|
let hook_executor = self.hook_executor.ok_or_else(|| {
|
||||||
.hook_executor
|
AgentError::Config(
|
||||||
.ok_or_else(|| AgentError::Config("缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into()))?;
|
"缺少 hook_executor,请先调用 .hook_executor(...)(空 HookExecutor 也可)".into(),
|
||||||
|
)
|
||||||
|
})?;
|
||||||
|
|
||||||
let config = self.config.unwrap_or_default();
|
let config = self.config.unwrap_or_default();
|
||||||
|
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
+54
-3
@@ -18,6 +18,7 @@ use crate::tools::error::ToolError;
|
|||||||
/// **不实现 `Clone`**:透传内层 `LlmError` / `MemoryError`,两者均未派生 `Clone`(保留
|
/// **不实现 `Clone`**:透传内层 `LlmError` / `MemoryError`,两者均未派生 `Clone`(保留
|
||||||
/// 完整错误信息,传递所有权)。如需在多 session 间共享错误状态,用 `Arc<AgentError>` 包装。
|
/// 完整错误信息,传递所有权)。如需在多 session 间共享错误状态,用 `Arc<AgentError>` 包装。
|
||||||
#[derive(Debug, Error)]
|
#[derive(Debug, Error)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum AgentError {
|
pub enum AgentError {
|
||||||
/// LLM 调用错误(透传 Phase 0)。
|
/// LLM 调用错误(透传 Phase 0)。
|
||||||
#[error("LLM 错误: {0}")]
|
#[error("LLM 错误: {0}")]
|
||||||
@@ -35,6 +36,18 @@ pub enum AgentError {
|
|||||||
#[error("Plan 解析错误: {0}")]
|
#[error("Plan 解析错误: {0}")]
|
||||||
PlanParse(String),
|
PlanParse(String),
|
||||||
|
|
||||||
|
/// Readonly slot 不允许写入(Phase 10 新增)。
|
||||||
|
#[error("Readonly slot 不允许写入: {0}")]
|
||||||
|
SlotReadonly(String),
|
||||||
|
|
||||||
|
/// Slot 不存在(Phase 10 新增)。
|
||||||
|
#[error("Slot '{0}' 不存在")]
|
||||||
|
SlotNotFound(String),
|
||||||
|
|
||||||
|
/// Slot 已存在(Phase 10 新增)。
|
||||||
|
#[error("Slot '{0}' 已存在")]
|
||||||
|
SlotAlreadyExists(String),
|
||||||
|
|
||||||
/// 钩子阻断操作(Agent 层特有)。
|
/// 钩子阻断操作(Agent 层特有)。
|
||||||
#[error("钩子阻断: {0}")]
|
#[error("钩子阻断: {0}")]
|
||||||
HookBlocked(String),
|
HookBlocked(String),
|
||||||
@@ -59,6 +72,7 @@ impl AgentError {
|
|||||||
/// - `Tool`:由内层 `is_recoverable()` 决定
|
/// - `Tool`:由内层 `is_recoverable()` 决定
|
||||||
/// - `HookBlocked` / `LimitExceeded`:不可恢复(需人工介入或终止循环)
|
/// - `HookBlocked` / `LimitExceeded`:不可恢复(需人工介入或终止循环)
|
||||||
/// - `Config` / `Other`:不可恢复
|
/// - `Config` / `Other`:不可恢复
|
||||||
|
/// - `SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`:不可恢复(结构性错误)
|
||||||
pub fn is_recoverable(&self) -> bool {
|
pub fn is_recoverable(&self) -> bool {
|
||||||
match self {
|
match self {
|
||||||
Self::Llm(e) => matches!(
|
Self::Llm(e) => matches!(
|
||||||
@@ -68,9 +82,13 @@ impl AgentError {
|
|||||||
Self::Tool(e) => e.is_recoverable(),
|
Self::Tool(e) => e.is_recoverable(),
|
||||||
Self::Memory(e) => e.is_recoverable(),
|
Self::Memory(e) => e.is_recoverable(),
|
||||||
Self::PlanParse(_) => false,
|
Self::PlanParse(_) => false,
|
||||||
Self::HookBlocked(_) | Self::LimitExceeded(_) | Self::Config(_) | Self::Other(_) => {
|
Self::SlotReadonly(_)
|
||||||
false
|
| Self::SlotNotFound(_)
|
||||||
}
|
| Self::SlotAlreadyExists(_)
|
||||||
|
| Self::HookBlocked(_)
|
||||||
|
| Self::LimitExceeded(_)
|
||||||
|
| Self::Config(_)
|
||||||
|
| Self::Other(_) => false,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -180,4 +198,37 @@ mod tests {
|
|||||||
let err = caller().unwrap_err();
|
let err = caller().unwrap_err();
|
||||||
assert!(matches!(err, AgentError::Memory(_)));
|
assert!(matches!(err, AgentError::Memory(_)));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ====== Phase 10: Slot 错误变体测试 ======
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn slot_readonly_not_recoverable() {
|
||||||
|
assert!(!AgentError::SlotReadonly("readonly".into()).is_recoverable());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn slot_not_found_not_recoverable() {
|
||||||
|
assert!(!AgentError::SlotNotFound("missing".into()).is_recoverable());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn slot_already_exists_not_recoverable() {
|
||||||
|
assert!(!AgentError::SlotAlreadyExists("dup".into()).is_recoverable());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn slot_error_messages() {
|
||||||
|
assert_eq!(
|
||||||
|
format!("{}", AgentError::SlotReadonly("readonly".into())),
|
||||||
|
"Readonly slot 不允许写入: readonly"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
format!("{}", AgentError::SlotNotFound("foo".into())),
|
||||||
|
"Slot 'foo' 不存在"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
format!("{}", AgentError::SlotAlreadyExists("bar".into())),
|
||||||
|
"Slot 'bar' 已存在"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,8 +16,8 @@ use std::sync::Arc;
|
|||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|
||||||
use crate::llm::compact::CompactConfig;
|
use crate::llm::compact::CompactConfig;
|
||||||
use crate::llm::provider::LlmProvider;
|
|
||||||
use crate::llm::hooks::HookExecutor;
|
use crate::llm::hooks::HookExecutor;
|
||||||
|
use crate::llm::provider::LlmProvider;
|
||||||
use crate::memory::retriever::MemoryRetriever;
|
use crate::memory::retriever::MemoryRetriever;
|
||||||
use crate::memory::store::MemoryStore;
|
use crate::memory::store::MemoryStore;
|
||||||
use crate::tools::ToolRegistry;
|
use crate::tools::ToolRegistry;
|
||||||
|
|||||||
+821
-101
File diff suppressed because it is too large
Load Diff
@@ -78,11 +78,7 @@ impl SessionMemory {
|
|||||||
prefix: Some(format!("{}:", self.namespace)),
|
prefix: Some(format!("{}:", self.namespace)),
|
||||||
..Default::default()
|
..Default::default()
|
||||||
};
|
};
|
||||||
let items = self
|
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
|
||||||
.store
|
|
||||||
.list(&filter)
|
|
||||||
.await
|
|
||||||
.map_err(AgentError::Memory)?;
|
|
||||||
|
|
||||||
let mut lines = Vec::with_capacity(items.len() + 2);
|
let mut lines = Vec::with_capacity(items.len() + 2);
|
||||||
lines.push("<session-context>".to_string());
|
lines.push("<session-context>".to_string());
|
||||||
@@ -113,11 +109,7 @@ impl SessionMemory {
|
|||||||
prefix: Some(format!("{}:", self.namespace)),
|
prefix: Some(format!("{}:", self.namespace)),
|
||||||
..Default::default()
|
..Default::default()
|
||||||
};
|
};
|
||||||
let items = self
|
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
|
||||||
.store
|
|
||||||
.list(&filter)
|
|
||||||
.await
|
|
||||||
.map_err(AgentError::Memory)?;
|
|
||||||
|
|
||||||
for item in items {
|
for item in items {
|
||||||
self.store
|
self.store
|
||||||
|
|||||||
+5
-11
@@ -10,8 +10,7 @@
|
|||||||
//! - 重试由上层新建 `Plan` 实现,`TaskAgent` 不做自动重试
|
//! - 重试由上层新建 `Plan` 实现,`TaskAgent` 不做自动重试
|
||||||
|
|
||||||
use crate::agent::error::AgentError;
|
use crate::agent::error::AgentError;
|
||||||
#[allow(deprecated)]
|
use crate::llm::types::response_v2::MessageResponse;
|
||||||
use crate::llm::types::ChatResponse;
|
|
||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
|
|
||||||
@@ -56,14 +55,14 @@ impl Step {
|
|||||||
/// 均未派生 `Clone`(保留原始错误信息,传递所有权而非克隆)。如需复制 `Plan`,
|
/// 均未派生 `Clone`(保留原始错误信息,传递所有权而非克隆)。如需复制 `Plan`,
|
||||||
/// 只能 clone 处于 `Pending` / `Running` / `Completed` / `Skipped` 状态的步骤。
|
/// 只能 clone 处于 `Pending` / `Running` / `Completed` / `Skipped` 状态的步骤。
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
#[allow(deprecated)]
|
#[non_exhaustive]
|
||||||
pub enum StepStatus {
|
pub enum StepStatus {
|
||||||
/// 初始状态 —— 等待执行。
|
/// 初始状态 —— 等待执行。
|
||||||
Pending,
|
Pending,
|
||||||
/// 正在执行(`TaskAgent::execute_plan` 进入)。
|
/// 正在执行(`TaskAgent::execute_plan` 进入)。
|
||||||
Running,
|
Running,
|
||||||
/// 已完成(含 LLM 响应)。
|
/// 已完成(含 LLM 响应)。
|
||||||
Completed(ChatResponse),
|
Completed(MessageResponse),
|
||||||
/// 失败(含错误)。
|
/// 失败(含错误)。
|
||||||
Failed(AgentError),
|
Failed(AgentError),
|
||||||
/// 跳过(上层主动跳过)。
|
/// 跳过(上层主动跳过)。
|
||||||
@@ -130,9 +129,7 @@ impl PlanParser for JsonPlanParser {
|
|||||||
.collect::<Result<Vec<_>, AgentError>>()?;
|
.collect::<Result<Vec<_>, AgentError>>()?;
|
||||||
|
|
||||||
if steps.is_empty() {
|
if steps.is_empty() {
|
||||||
return Err(AgentError::PlanParse(
|
return Err(AgentError::PlanParse("Plan 至少需要一个步骤".into()));
|
||||||
"Plan 至少需要一个步骤".into(),
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(Plan {
|
Ok(Plan {
|
||||||
@@ -203,10 +200,7 @@ mod tests {
|
|||||||
let plan = Plan {
|
let plan = Plan {
|
||||||
id: "p1".into(),
|
id: "p1".into(),
|
||||||
goal: "test goal".into(),
|
goal: "test goal".into(),
|
||||||
steps: vec![
|
steps: vec![Step::new(0, "first"), Step::new(1, "second")],
|
||||||
Step::new(0, "first"),
|
|
||||||
Step::new(1, "second"),
|
|
||||||
],
|
|
||||||
};
|
};
|
||||||
assert_eq!(plan.steps.len(), 2);
|
assert_eq!(plan.steps.len(), 2);
|
||||||
assert_eq!(plan.steps[0].index, 0);
|
assert_eq!(plan.steps[0].index, 0);
|
||||||
|
|||||||
+580
@@ -0,0 +1,580 @@
|
|||||||
|
//! Document 系统 —— 文本分割与文档类型。
|
||||||
|
//!
|
||||||
|
//! 提供 [`Document`] 数据结构和 [`RecursiveCharacterSplitter`] 分割器,
|
||||||
|
//! 作为 RAG 管线(split → embed → store)的前置步骤。
|
||||||
|
|
||||||
|
use std::collections::HashMap;
|
||||||
|
|
||||||
|
/// 默认分隔符优先级列表(按优先级降序)。
|
||||||
|
///
|
||||||
|
/// 段落级 → 行级 → 句子级(含 CJK 标点) → 词级 → 字符级(兜底)。
|
||||||
|
/// 在 LangChain 基础上扩充了 CJK 句号 `"。"`、问号 `"?"`、感叹号 `"!"`,
|
||||||
|
/// 确保中文文本在句子边界有更高分割质量。
|
||||||
|
const DEFAULT_SEPARATORS: &[&str] = &["\n\n", "\n", "。", "?", "!", ".", " ", ""];
|
||||||
|
|
||||||
|
/// 文档片段 —— RAG 管线的基本数据载体。
|
||||||
|
///
|
||||||
|
/// 作为分割(split)和向量化(embed)两个阶段的通货类型,
|
||||||
|
/// 在 Phase 15 的 RagPipeline 中串联 split → embed → store。
|
||||||
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||||
|
pub struct Document {
|
||||||
|
/// 文档唯一标识。
|
||||||
|
pub id: String,
|
||||||
|
/// 文档文本内容。
|
||||||
|
pub content: String,
|
||||||
|
/// 元数据标签(键值对,可用作过滤、溯源、分类)。
|
||||||
|
pub metadata: HashMap<String, String>,
|
||||||
|
/// MIME 类型,标识内容格式(如 "text/plain", "text/markdown")。
|
||||||
|
pub mime_type: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Document {
|
||||||
|
/// 创建一个新文档。元数据默认初始化为空。
|
||||||
|
///
|
||||||
|
/// 分割器产生的 chunks 会自动继承源文档 mime_type,
|
||||||
|
/// 并在 metadata 中追加 source_id / chunk_index / chunk_count。
|
||||||
|
pub fn new(
|
||||||
|
id: impl Into<String>,
|
||||||
|
content: impl Into<String>,
|
||||||
|
mime_type: impl Into<String>,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
id: id.into(),
|
||||||
|
content: content.into(),
|
||||||
|
metadata: HashMap::new(),
|
||||||
|
mime_type: mime_type.into(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 快速构造纯文本文档(mime_type 默认为 "text/plain")。
|
||||||
|
/// 适用于大多数无需指定媒体类型的场景。
|
||||||
|
pub fn from_raw(id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||||
|
Self {
|
||||||
|
id: id.into(),
|
||||||
|
content: content.into(),
|
||||||
|
metadata: HashMap::new(),
|
||||||
|
mime_type: "text/plain".into(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 递归字符级文档分割器。
|
||||||
|
///
|
||||||
|
/// 使用可配置的分隔符优先级列表,递归地将文档分割为
|
||||||
|
/// 接近 chunk_size 的块。
|
||||||
|
///
|
||||||
|
/// # 算法(两阶段)
|
||||||
|
///
|
||||||
|
/// 1. **递归分割**:按分隔符优先级从高到低递归切割文本,
|
||||||
|
/// 产生初始片段(均 ≤ chunk_size,按字符数计算)。
|
||||||
|
///
|
||||||
|
/// 2. **贪心合并**:从左向右合并相邻片段,直到合计字符数
|
||||||
|
/// 超过 chunk_size,此时将前一组合并结果作为一个 chunk 输出,
|
||||||
|
/// 并携带 chunk_overlap 字符的滑动窗口。
|
||||||
|
///
|
||||||
|
/// 所有长度比较均以 Unicode 字符数为单位(`text.chars().count()`),
|
||||||
|
/// 而非字节数。CJK 文本每个字算 1 个 char。
|
||||||
|
///
|
||||||
|
/// # 升级路径
|
||||||
|
///
|
||||||
|
/// - 如需自定义分割函数,可在上层通过 `with_custom_splitter`
|
||||||
|
/// 扩展(当前未实现,预留升级路径)。
|
||||||
|
/// - 如需 unicode 感知的句子分割(如中文句号、缩写处理),
|
||||||
|
/// 可在 separators 中加入对应字符串,或将下游替换为
|
||||||
|
/// 基于 unicode-segmentation crate 的自定义分割器。
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct RecursiveCharacterSplitter {
|
||||||
|
chunk_size: usize,
|
||||||
|
chunk_overlap: usize,
|
||||||
|
separators: Vec<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl RecursiveCharacterSplitter {
|
||||||
|
/// 创建分割器。
|
||||||
|
///
|
||||||
|
/// # Panics
|
||||||
|
///
|
||||||
|
/// - 如果 `chunk_size == 0`
|
||||||
|
/// - 如果 `chunk_size ≤ chunk_overlap`(无法形成有效滑动窗口)
|
||||||
|
pub fn new(chunk_size: usize, chunk_overlap: usize) -> Self {
|
||||||
|
if chunk_size == 0 {
|
||||||
|
panic!("chunk_size must be greater than 0");
|
||||||
|
}
|
||||||
|
if chunk_size <= chunk_overlap {
|
||||||
|
panic!("chunk_size must be greater than chunk_overlap");
|
||||||
|
}
|
||||||
|
Self {
|
||||||
|
chunk_size,
|
||||||
|
chunk_overlap,
|
||||||
|
separators: DEFAULT_SEPARATORS.iter().map(|s| s.to_string()).collect(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 创建分割器的安全版本。
|
||||||
|
///
|
||||||
|
/// 验证失败时返回 `Err` 而非 panic。
|
||||||
|
pub fn try_new(chunk_size: usize, chunk_overlap: usize) -> Result<Self, &'static str> {
|
||||||
|
if chunk_size == 0 {
|
||||||
|
return Err("chunk_size must be greater than 0");
|
||||||
|
}
|
||||||
|
if chunk_size <= chunk_overlap {
|
||||||
|
return Err("chunk_size must be greater than chunk_overlap");
|
||||||
|
}
|
||||||
|
Ok(Self {
|
||||||
|
chunk_size,
|
||||||
|
chunk_overlap,
|
||||||
|
separators: DEFAULT_SEPARATORS.iter().map(|s| s.to_string()).collect(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 覆盖默认分隔符优先级列表。
|
||||||
|
///
|
||||||
|
/// **重要**:建议保留 `""` 作为最后一个 separator,
|
||||||
|
/// 作为字符级兜底防止任何文本都能被分割。
|
||||||
|
pub fn with_separators(mut self, separators: Vec<String>) -> Self {
|
||||||
|
self.separators = separators;
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 返回 chunk_size(字符数)。
|
||||||
|
pub fn chunk_size(&self) -> usize {
|
||||||
|
self.chunk_size
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 返回 chunk_overlap(字符数)。
|
||||||
|
pub fn chunk_overlap(&self) -> usize {
|
||||||
|
self.chunk_overlap
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 批量分割。
|
||||||
|
///
|
||||||
|
/// 每个输入文档独立分割。输出 chunks 继承源文档的 mime_type,
|
||||||
|
/// 并在 metadata 中追加 source_id / chunk_index / chunk_count。
|
||||||
|
///
|
||||||
|
/// Chunk ID 格式:`{source_id}:chunk:{index:04d}`
|
||||||
|
/// 例如 `"doc_001:chunk:0000"`(索引从 0 开始,4 位固定宽度)。
|
||||||
|
///
|
||||||
|
/// **注意**:metadata 注入使用 `HashMap::insert()`,如果源 Document
|
||||||
|
/// 的 metadata 已包含 `"source_id"`、`"chunk_index"` 或 `"chunk_count"`
|
||||||
|
/// 键,将被分割器的值静默覆盖。
|
||||||
|
pub fn split(&self, documents: &[Document]) -> Vec<Document> {
|
||||||
|
tracing::debug!(
|
||||||
|
input_count = documents.len(),
|
||||||
|
"RecursiveCharacterSplitter::split start"
|
||||||
|
);
|
||||||
|
|
||||||
|
let mut output = Vec::new();
|
||||||
|
for doc in documents {
|
||||||
|
let segments = self.split_text(&doc.content, &self.separators);
|
||||||
|
let chunks = self.merge_with_overlap(segments);
|
||||||
|
debug_assert!(
|
||||||
|
chunks.len() < 10_000,
|
||||||
|
"单个文档产生超过 9999 个 chunk,索引格式溢出"
|
||||||
|
);
|
||||||
|
|
||||||
|
tracing::trace!(
|
||||||
|
doc_id = %doc.id,
|
||||||
|
chunk_count = chunks.len(),
|
||||||
|
"document split into chunks"
|
||||||
|
);
|
||||||
|
|
||||||
|
for (idx, chunk_text) in chunks.iter().enumerate() {
|
||||||
|
let mut metadata = doc.metadata.clone();
|
||||||
|
metadata.insert("source_id".to_string(), doc.id.clone());
|
||||||
|
metadata.insert("chunk_index".to_string(), idx.to_string());
|
||||||
|
metadata.insert("chunk_count".to_string(), chunks.len().to_string());
|
||||||
|
|
||||||
|
let id = format!("{}:chunk:{:04}", doc.id, idx);
|
||||||
|
tracing::trace!(chunk_id = %id, "chunk produced");
|
||||||
|
|
||||||
|
output.push(Document {
|
||||||
|
id,
|
||||||
|
content: chunk_text.clone(),
|
||||||
|
metadata,
|
||||||
|
mime_type: doc.mime_type.clone(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
output
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 递归分割(Phase 1)。
|
||||||
|
///
|
||||||
|
/// 按 separator 优先级从高到低切割文本。每个输出片段的字符数
|
||||||
|
/// 均 ≤ chunk_size(除非最终降到 `""` 字符级兜底)。
|
||||||
|
///
|
||||||
|
/// Phase 1 只做"切分",不做合并——合并由 Phase 2 (`merge_with_overlap`) 处理。
|
||||||
|
///
|
||||||
|
/// **关键行为**:当文本中存在 separator 时,按 separator 切分。
|
||||||
|
/// 若所有 segment 均 ≤ chunk_size,直接返回所有 segments;
|
||||||
|
/// 若某个 segment > chunk_size,递归降级到下一级 separator。
|
||||||
|
///
|
||||||
|
/// **早返回守卫**:如果整段文本 ≤ chunk_size(含恰好等于),直接
|
||||||
|
/// 返回 `[text.to_string()]`,避免在 Phase 2 合并时丢失 separator
|
||||||
|
/// 边界信息。
|
||||||
|
fn split_text(&self, text: &str, separators: &[String]) -> Vec<String> {
|
||||||
|
if text.is_empty() {
|
||||||
|
return Vec::new();
|
||||||
|
}
|
||||||
|
// 早返回:整段文本 ≤ chunk_size 时整体返回,避免分割后再
|
||||||
|
// 合并时丢失 separator 边界
|
||||||
|
if chars_len(text) <= self.chunk_size {
|
||||||
|
return vec![text.to_string()];
|
||||||
|
}
|
||||||
|
if separators.is_empty() {
|
||||||
|
// 防御:理论上不应到达这里(DEFAULT_SEPARATORS 末尾有 `""`)
|
||||||
|
return self.split_by_chars(text);
|
||||||
|
}
|
||||||
|
|
||||||
|
let sep = &separators[0];
|
||||||
|
if sep.is_empty() {
|
||||||
|
// 字符级兜底
|
||||||
|
return self.split_by_chars(text);
|
||||||
|
}
|
||||||
|
|
||||||
|
// 检查文本中是否包含当前 separator
|
||||||
|
if !text.contains(sep.as_str()) {
|
||||||
|
// 不含此 separator,降级到下一级
|
||||||
|
return self.split_text(text, &separators[1..]);
|
||||||
|
}
|
||||||
|
|
||||||
|
// 文本中存在 separator,按 separator 切分
|
||||||
|
let raw_segments: Vec<&str> = text.split(sep.as_str()).collect();
|
||||||
|
let mut result = Vec::new();
|
||||||
|
|
||||||
|
for seg in raw_segments {
|
||||||
|
if seg.is_empty() {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
if chars_len(seg) > self.chunk_size {
|
||||||
|
// 当前片段超长:递归降级到下一级 separator
|
||||||
|
result.extend(self.split_text(seg, &separators[1..]));
|
||||||
|
} else {
|
||||||
|
// 当前片段符合 chunk_size,直接输出
|
||||||
|
result.push(seg.to_string());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
result
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 字符级兜底分割(确保任何文本都能被切到 chunk_size 以内)。
|
||||||
|
///
|
||||||
|
/// 使用 `char_indices()` 步进,避免截断在多字节 UTF-8 字符中间。
|
||||||
|
fn split_by_chars(&self, text: &str) -> Vec<String> {
|
||||||
|
let mut result = Vec::new();
|
||||||
|
let mut current = String::new();
|
||||||
|
|
||||||
|
for (_, ch) in text.char_indices() {
|
||||||
|
current.push(ch);
|
||||||
|
if chars_len(¤t) >= self.chunk_size {
|
||||||
|
result.push(std::mem::take(&mut current));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !current.is_empty() {
|
||||||
|
result.push(current);
|
||||||
|
}
|
||||||
|
result
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 贪心合并 + overlap 滑动窗口(Phase 2)。
|
||||||
|
///
|
||||||
|
/// 把 Phase 1 输出的 segments 合并到目标 chunk_size,并对相邻 chunk
|
||||||
|
/// 应用 chunk_overlap 字符的重叠窗口。
|
||||||
|
///
|
||||||
|
/// **已知行为**:合并时使用空字符串 `""` 连接相邻 segments
|
||||||
|
/// (即 `current.join("")`),不保留 Phase 1 切分时消耗的 separator
|
||||||
|
/// 边界信息。这意味着跨 chunk 的结构化边界(如段落、句子)会
|
||||||
|
/// 在合并点"塌缩"——但对 RAG 语义检索影响通常较小。如需保留
|
||||||
|
/// separator 边界,可重构此方法接受 separator 参数。
|
||||||
|
fn merge_with_overlap(&self, mut segments: Vec<String>) -> Vec<String> {
|
||||||
|
if segments.is_empty() {
|
||||||
|
return Vec::new();
|
||||||
|
}
|
||||||
|
if segments.len() == 1 {
|
||||||
|
return segments;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Phase 2a: 贪心合并 segments 到目标 chunk_size
|
||||||
|
// (segments 用 "" 连接,sep_count 不参与长度计算)
|
||||||
|
let mut chunks: Vec<String> = Vec::new();
|
||||||
|
let mut current: Vec<String> = Vec::new();
|
||||||
|
let mut current_len: usize = 0;
|
||||||
|
|
||||||
|
for seg in segments.drain(..) {
|
||||||
|
let seg_len = chars_len(&seg);
|
||||||
|
let new_total = current_len + seg_len;
|
||||||
|
|
||||||
|
if new_total > self.chunk_size && !current.is_empty() {
|
||||||
|
chunks.push(current.join(""));
|
||||||
|
current.clear();
|
||||||
|
current_len = 0;
|
||||||
|
}
|
||||||
|
current.push(seg);
|
||||||
|
current_len += seg_len;
|
||||||
|
}
|
||||||
|
|
||||||
|
if !current.is_empty() {
|
||||||
|
chunks.push(current.join(""));
|
||||||
|
}
|
||||||
|
|
||||||
|
if chunks.len() <= 1 {
|
||||||
|
return chunks;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Phase 2b: 应用 overlap 滑动窗口(除第一个 chunk 外)
|
||||||
|
let overlap = self.chunk_overlap;
|
||||||
|
if overlap == 0 {
|
||||||
|
return chunks;
|
||||||
|
}
|
||||||
|
|
||||||
|
for i in 1..chunks.len() {
|
||||||
|
let prev = &chunks[i - 1];
|
||||||
|
let prev_chars_count = chars_len(prev);
|
||||||
|
if prev_chars_count == 0 {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
let take_n = overlap.min(prev_chars_count);
|
||||||
|
|
||||||
|
// 字符级安全地取 prev 末尾 take_n 个字符
|
||||||
|
let tail: String = prev.chars().rev().take(take_n).collect::<Vec<_>>().into_iter().rev().collect();
|
||||||
|
chunks[i] = format!("{}{}", tail, chunks[i]);
|
||||||
|
}
|
||||||
|
|
||||||
|
chunks
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for RecursiveCharacterSplitter {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new(1000, 200)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 字符数(Unicode 标量值),等价于 `s.chars().count()`。
|
||||||
|
#[inline]
|
||||||
|
fn chars_len(s: &str) -> usize {
|
||||||
|
s.chars().count()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
// ===== Group B1 — Document struct 基础测试 =====
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn document_new_metadata_defaults_empty() {
|
||||||
|
let doc = Document::new("id-1", "content", "text/plain");
|
||||||
|
assert!(doc.metadata.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn document_clone_partial_eq() {
|
||||||
|
let doc = Document::new("id-1", "content", "text/plain");
|
||||||
|
let cloned = doc.clone();
|
||||||
|
assert_eq!(doc, cloned);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn document_different_ids_not_equal() {
|
||||||
|
let doc1 = Document::new("id-1", "content", "text/plain");
|
||||||
|
let doc2 = Document::new("id-2", "content", "text/plain");
|
||||||
|
assert_ne!(doc1, doc2);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn document_from_raw_uses_text_plain() {
|
||||||
|
let doc = Document::from_raw("id-1", "hello");
|
||||||
|
assert_eq!(doc.mime_type, "text/plain");
|
||||||
|
assert!(doc.metadata.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== Group B2 — Splitter 边界条件测试 =====
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_empty_doc_returns_empty() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(100, 20);
|
||||||
|
let chunks = splitter.split(&[]);
|
||||||
|
assert!(chunks.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_short_doc_single_chunk() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(100, 20);
|
||||||
|
let doc = Document::from_raw("short", "hello");
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
assert_eq!(chunks.len(), 1);
|
||||||
|
assert_eq!(chunks[0].content, "hello");
|
||||||
|
assert_eq!(chunks[0].metadata.get("chunk_index").map(|s| s.as_str()), Some("0"));
|
||||||
|
assert_eq!(chunks[0].metadata.get("chunk_count").map(|s| s.as_str()), Some("1"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_empty_content_yields_no_chunks() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(100, 20);
|
||||||
|
let doc = Document::from_raw("empty", "");
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
assert!(chunks.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[should_panic(expected = "chunk_size must be greater than chunk_overlap")]
|
||||||
|
fn split_constructor_panics_on_invalid_overlap() {
|
||||||
|
let _ = RecursiveCharacterSplitter::new(10, 10);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
#[should_panic(expected = "chunk_size must be greater than 0")]
|
||||||
|
fn split_constructor_panics_on_zero_chunk_size() {
|
||||||
|
let _ = RecursiveCharacterSplitter::new(0, 0);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn try_new_returns_err_on_invalid_params() {
|
||||||
|
assert!(RecursiveCharacterSplitter::try_new(0, 0).is_err());
|
||||||
|
assert!(RecursiveCharacterSplitter::try_new(10, 10).is_err());
|
||||||
|
assert!(RecursiveCharacterSplitter::try_new(100, 20).is_ok());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn default_separators_match_spec() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::default();
|
||||||
|
// Default separators should include CJK punctuation as the last meaningful
|
||||||
|
// separator before the char-level fallback. We can't directly access the
|
||||||
|
// private field, so we verify behavior: a Chinese sentence should split
|
||||||
|
// on "。" at the sentence level rather than the word level.
|
||||||
|
let doc = Document::from_raw("zh", "你好世界。今天天气好。");
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
// The default chunk_size=1000, so the whole content fits in 1 chunk.
|
||||||
|
// But the separators list contains "。" — this is verified via integration test.
|
||||||
|
assert!(!chunks.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== Group B3 — Splitter 核心算法测试 =====
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_paragraph_boundary() {
|
||||||
|
// 小 chunk_size 强制段落级别分割
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(4, 1);
|
||||||
|
let doc = Document::from_raw("p", "para1\n\npara2");
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
// para1 (5 chars) > chunk_size=4 → 递归降级到 char 级拆分
|
||||||
|
// para2 同理
|
||||||
|
// 总共应该产生多个 chunk
|
||||||
|
assert!(chunks.len() >= 2, "expected >= 2 chunks, got {}", chunks.len());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_recursive_deepen() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(50, 5);
|
||||||
|
// 200 字符无 \n\n,强制降级
|
||||||
|
let text: String = "a".repeat(200);
|
||||||
|
let doc = Document::from_raw("long", &text);
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
assert!(chunks.len() >= 3, "expected >= 3 chunks, got {}", chunks.len());
|
||||||
|
for chunk in &chunks {
|
||||||
|
// chunk 内容 = overlap_tail(≤5) + new_content(≤50),故 ≤ 55
|
||||||
|
assert!(chars_len(&chunk.content) <= 55, "chunk too long: {} chars", chars_len(&chunk.content));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_greedy_merge_combines_segments() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(20, 2);
|
||||||
|
// 一段含多个 \n\n 分隔的短小段,应被合并到 chunk_size
|
||||||
|
let doc = Document::from_raw("g", "aa\n\nbb\n\ncc\n\ndd");
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
// 短段应被合并:总共应该少于 4 个 chunk
|
||||||
|
assert!(chunks.len() <= 3, "expected <= 3 chunks after merge, got {}", chunks.len());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_overlap_consistency() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(20, 5);
|
||||||
|
// 构造一个需要多 chunk 的文本
|
||||||
|
let text: String = "x".repeat(50);
|
||||||
|
let doc = Document::from_raw("o", &text);
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
assert!(chunks.len() >= 2);
|
||||||
|
// chunk[1] 应该以 chunk[0] 的最后 5 个字符作为前缀
|
||||||
|
let prev_tail: String = chunks[0]
|
||||||
|
.content
|
||||||
|
.chars()
|
||||||
|
.rev()
|
||||||
|
.take(5)
|
||||||
|
.collect::<Vec<_>>()
|
||||||
|
.into_iter()
|
||||||
|
.rev()
|
||||||
|
.collect();
|
||||||
|
assert!(
|
||||||
|
chunks[1].content.starts_with(&prev_tail),
|
||||||
|
"chunk[1] should start with last 5 chars of chunk[0]: prev_tail={:?}, chunk[1]={:?}",
|
||||||
|
prev_tail,
|
||||||
|
chunks[1].content
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_character_fallback() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(5, 0);
|
||||||
|
// 纯字母无标点,应降级到字符级
|
||||||
|
let doc = Document::from_raw("cf", "aaaaaaaaa");
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
assert_eq!(chunks.len(), 2, "expected 2 chunks, got {}", chunks.len());
|
||||||
|
for chunk in &chunks {
|
||||||
|
assert!(chars_len(&chunk.content) <= 5);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_multibyte_utf8_boundary() {
|
||||||
|
// 验证字符级单位而非字节级单位
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(10, 2);
|
||||||
|
// 30 个中文字符 = 90 字节(UTF-8)
|
||||||
|
let text: String = "中".repeat(30);
|
||||||
|
let doc = Document::from_raw("cjk", &text);
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
// 30 字符 / 10 chunk_size = 3 个 chunk
|
||||||
|
assert!(chunks.len() >= 3, "expected >= 3 chunks for 30 chars / chunk_size=10, got {}", chunks.len());
|
||||||
|
for chunk in &chunks {
|
||||||
|
let char_count = chars_len(&chunk.content);
|
||||||
|
// chunk = overlap_tail(≤2) + new_content(≤10),故 ≤ 12
|
||||||
|
assert!(char_count <= 12, "chunk char count {} exceeds 10+overlap", char_count);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== Group B4 — Splitter 集成测试 =====
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_multiple_docs() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(50, 5);
|
||||||
|
let docs = vec![
|
||||||
|
Document::from_raw("a", "a".repeat(30).as_str()),
|
||||||
|
Document::from_raw("b", "b".repeat(30).as_str()),
|
||||||
|
Document::from_raw("c", "c".repeat(30).as_str()),
|
||||||
|
];
|
||||||
|
let chunks = splitter.split(&docs);
|
||||||
|
assert!(chunks.len() >= 3);
|
||||||
|
// 每个 chunk 的 source_id 应指向对应的输入 doc
|
||||||
|
for chunk in &chunks {
|
||||||
|
let source = chunk.metadata.get("source_id").unwrap();
|
||||||
|
assert!(["a", "b", "c"].contains(&source.as_str()));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn split_metadata_inheritance() {
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(100, 10);
|
||||||
|
let mut doc = Document::new("m", "short content", "text/plain");
|
||||||
|
doc.metadata.insert("author".to_string(), "alice".to_string());
|
||||||
|
let chunks = splitter.split(&[doc]);
|
||||||
|
assert_eq!(chunks.len(), 1);
|
||||||
|
assert_eq!(chunks[0].metadata.get("author").map(|s| s.as_str()), Some("alice"));
|
||||||
|
assert_eq!(chunks[0].metadata.get("source_id").map(|s| s.as_str()), Some("m"));
|
||||||
|
assert_eq!(chunks[0].metadata.get("chunk_index").map(|s| s.as_str()), Some("0"));
|
||||||
|
assert_eq!(chunks[0].metadata.get("chunk_count").map(|s| s.as_str()), Some("1"));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,11 +1,14 @@
|
|||||||
//! agcore —— 智能体(Agent)核心工具箱。
|
//! agcore —— 智能体(Agent)核心工具箱。
|
||||||
|
|
||||||
pub mod agent;
|
pub mod agent;
|
||||||
|
pub mod document;
|
||||||
pub mod llm;
|
pub mod llm;
|
||||||
pub mod memory;
|
pub mod memory;
|
||||||
pub mod prompt;
|
pub mod prompt;
|
||||||
pub mod tools;
|
pub mod tools;
|
||||||
|
|
||||||
|
pub use document::Document;
|
||||||
|
|
||||||
use tracing_subscriber::{EnvFilter, fmt, prelude::*};
|
use tracing_subscriber::{EnvFilter, fmt, prelude::*};
|
||||||
|
|
||||||
static INIT: std::sync::Once = std::sync::Once::new();
|
static INIT: std::sync::Once = std::sync::Once::new();
|
||||||
|
|||||||
@@ -3,6 +3,7 @@
|
|||||||
pub mod compact;
|
pub mod compact;
|
||||||
pub mod convert;
|
pub mod convert;
|
||||||
pub mod cycle;
|
pub mod cycle;
|
||||||
|
pub mod embedding;
|
||||||
pub mod error;
|
pub mod error;
|
||||||
pub mod hooks;
|
pub mod hooks;
|
||||||
pub mod mock;
|
pub mod mock;
|
||||||
|
|||||||
+32
-16
@@ -73,10 +73,7 @@ impl CompactState {
|
|||||||
|
|
||||||
/// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。
|
/// 粗略估计消息列表的 token 数(基于字符数,4 字符 ≈ 1 token)。
|
||||||
pub fn estimate_message_tokens(messages: &[Message]) -> u32 {
|
pub fn estimate_message_tokens(messages: &[Message]) -> u32 {
|
||||||
messages
|
messages.iter().map(estimate_single_message_tokens).sum()
|
||||||
.iter()
|
|
||||||
.map(estimate_single_message_tokens)
|
|
||||||
.sum()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn estimate_single_message_tokens(msg: &Message) -> u32 {
|
fn estimate_single_message_tokens(msg: &Message) -> u32 {
|
||||||
@@ -99,9 +96,7 @@ fn estimate_block_tokens(block: &ContentBlock) -> u32 {
|
|||||||
match block {
|
match block {
|
||||||
ContentBlock::Text { text } => estimate_text_tokens(text),
|
ContentBlock::Text { text } => estimate_text_tokens(text),
|
||||||
ContentBlock::Thinking { text, .. } => estimate_text_tokens(text),
|
ContentBlock::Thinking { text, .. } => estimate_text_tokens(text),
|
||||||
ContentBlock::ToolUse { input, .. } => {
|
ContentBlock::ToolUse { input, .. } => estimate_text_tokens(&input.to_string()),
|
||||||
estimate_text_tokens(&input.to_string())
|
|
||||||
}
|
|
||||||
ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
|
ContentBlock::ToolResult { content, .. } => estimate_content_blocks_tokens(content),
|
||||||
// ponytail: Image / Audio / File / Extension 在 IR 中固定估算。
|
// ponytail: Image / Audio / File / Extension 在 IR 中固定估算。
|
||||||
// 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。
|
// 无文本的视觉/音频 block 用兜底估算,避免 token 计数膨胀。
|
||||||
@@ -148,14 +143,25 @@ pub fn microcompact(messages: &mut [Message], keep_recent: usize) -> u32 {
|
|||||||
|
|
||||||
// 第一遍:计算可释放 token(仅非错误 ToolResult)
|
// 第一遍:计算可释放 token(仅非错误 ToolResult)
|
||||||
for msg in &messages[..prune_start] {
|
for msg in &messages[..prune_start] {
|
||||||
if matches!(msg, Message::ToolResult { is_error: false, .. }) {
|
if matches!(
|
||||||
|
msg,
|
||||||
|
Message::ToolResult {
|
||||||
|
is_error: false,
|
||||||
|
..
|
||||||
|
}
|
||||||
|
) {
|
||||||
freed_tokens += estimate_single_message_tokens(msg);
|
freed_tokens += estimate_single_message_tokens(msg);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 第二遍:替换内容(仅非错误 ToolResult)
|
// 第二遍:替换内容(仅非错误 ToolResult)
|
||||||
for msg in &mut messages[..prune_start] {
|
for msg in &mut messages[..prune_start] {
|
||||||
if let Message::ToolResult { content, is_error: false, .. } = msg {
|
if let Message::ToolResult {
|
||||||
|
content,
|
||||||
|
is_error: false,
|
||||||
|
..
|
||||||
|
} = msg
|
||||||
|
{
|
||||||
*content = vec![ContentBlock::Text {
|
*content = vec![ContentBlock::Text {
|
||||||
text: "[pruned]".to_string(),
|
text: "[pruned]".to_string(),
|
||||||
}];
|
}];
|
||||||
@@ -177,13 +183,15 @@ mod tests {
|
|||||||
fn estimate_message_tokens_handles_all_variants() {
|
fn estimate_message_tokens_handles_all_variants() {
|
||||||
let messages = vec![
|
let messages = vec![
|
||||||
Message::System {
|
Message::System {
|
||||||
content: vec![ContentBlock::Text {
|
content: vec![ContentBlock::Text { text: "sys".into() }],
|
||||||
text: "sys".into(),
|
|
||||||
}],
|
|
||||||
},
|
},
|
||||||
Message::user_text("hi"),
|
Message::user_text("hi"),
|
||||||
Message::assistant("ans"),
|
Message::assistant("ans"),
|
||||||
Message::user_image("b64", "image/png", crate::llm::types::shared::ImageDetail::Auto),
|
Message::user_image(
|
||||||
|
"b64",
|
||||||
|
"image/png",
|
||||||
|
crate::llm::types::shared::ImageDetail::Auto,
|
||||||
|
),
|
||||||
Message::tool_result("call_1", "tool res", false),
|
Message::tool_result("call_1", "tool res", false),
|
||||||
];
|
];
|
||||||
let tokens = estimate_message_tokens(&messages);
|
let tokens = estimate_message_tokens(&messages);
|
||||||
@@ -205,7 +213,10 @@ mod tests {
|
|||||||
assert!(freed > 0);
|
assert!(freed > 0);
|
||||||
assert_eq!(messages.len(), before_len); // 只改内容,不删消息
|
assert_eq!(messages.len(), before_len); // 只改内容,不删消息
|
||||||
// 索引 1 是被压缩的 ToolResult
|
// 索引 1 是被压缩的 ToolResult
|
||||||
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
|
if let Message::ToolResult {
|
||||||
|
content, is_error, ..
|
||||||
|
} = &messages[1]
|
||||||
|
{
|
||||||
assert_eq!(content.len(), 1);
|
assert_eq!(content.len(), 1);
|
||||||
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
|
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "[pruned]"));
|
||||||
assert!(!is_error);
|
assert!(!is_error);
|
||||||
@@ -228,9 +239,14 @@ mod tests {
|
|||||||
assert_eq!(freed, 0); // 错误 ToolResult 不计入
|
assert_eq!(freed, 0); // 错误 ToolResult 不计入
|
||||||
assert_eq!(messages.len(), before_len);
|
assert_eq!(messages.len(), before_len);
|
||||||
// 错误信息保留完整
|
// 错误信息保留完整
|
||||||
if let Message::ToolResult { content, is_error, .. } = &messages[1] {
|
if let Message::ToolResult {
|
||||||
|
content, is_error, ..
|
||||||
|
} = &messages[1]
|
||||||
|
{
|
||||||
assert!(is_error);
|
assert!(is_error);
|
||||||
assert!(matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down")));
|
assert!(
|
||||||
|
matches!(&content[0], ContentBlock::Text { text } if text.contains("backend down"))
|
||||||
|
);
|
||||||
} else {
|
} else {
|
||||||
panic!("expected ToolResult at index 1");
|
panic!("expected ToolResult at index 1");
|
||||||
}
|
}
|
||||||
|
|||||||
+20
-21
@@ -8,11 +8,9 @@
|
|||||||
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
use crate::llm::types::message::{ContentBlock, Message};
|
|
||||||
use crate::llm::types::openai_message::{
|
|
||||||
ContentField, OpenaiChatMessage, OpenaiContentPart,
|
|
||||||
};
|
|
||||||
use crate::llm::types::OpenaiToolCall;
|
use crate::llm::types::OpenaiToolCall;
|
||||||
|
use crate::llm::types::message::{ContentBlock, Message};
|
||||||
|
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage, OpenaiContentPart};
|
||||||
|
|
||||||
/// `OpenaiChatMessage` → IR `Message`。
|
/// `OpenaiChatMessage` → IR `Message`。
|
||||||
///
|
///
|
||||||
@@ -24,11 +22,10 @@ use crate::llm::types::OpenaiToolCall;
|
|||||||
/// - `Function`(已废弃)→ `Message::ToolResult`(`name` 作为 `tool_call_id` 兜底)
|
/// - `Function`(已废弃)→ `Message::ToolResult`(`name` 作为 `tool_call_id` 兜底)
|
||||||
pub fn from_openai(msg: &OpenaiChatMessage) -> Message {
|
pub fn from_openai(msg: &OpenaiChatMessage) -> Message {
|
||||||
match msg {
|
match msg {
|
||||||
OpenaiChatMessage::Developer { content, .. } | OpenaiChatMessage::System { content, .. } => {
|
OpenaiChatMessage::Developer { content, .. }
|
||||||
Message::System {
|
| OpenaiChatMessage::System { content, .. } => Message::System {
|
||||||
content: content_to_blocks(content),
|
content: content_to_blocks(content),
|
||||||
}
|
},
|
||||||
}
|
|
||||||
OpenaiChatMessage::User { content, .. } => Message::User {
|
OpenaiChatMessage::User { content, .. } => Message::User {
|
||||||
content: content_to_blocks(content),
|
content: content_to_blocks(content),
|
||||||
},
|
},
|
||||||
@@ -86,7 +83,11 @@ pub fn to_openai(msg: &Message) -> OpenaiChatMessage {
|
|||||||
content: blocks_to_content(content),
|
content: blocks_to_content(content),
|
||||||
name: None,
|
name: None,
|
||||||
},
|
},
|
||||||
Message::UserImage { data, mime_type, detail } => {
|
Message::UserImage {
|
||||||
|
data,
|
||||||
|
mime_type,
|
||||||
|
detail,
|
||||||
|
} => {
|
||||||
// ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。
|
// ponytail: 构造为单 image part 的 User 消息(OpenAI 多模态格式)。
|
||||||
let mime = mime_type.clone();
|
let mime = mime_type.clone();
|
||||||
let is_url = data.starts_with("http://") || data.starts_with("https://");
|
let is_url = data.starts_with("http://") || data.starts_with("https://");
|
||||||
@@ -167,18 +168,17 @@ pub fn content_to_blocks(field: &ContentField) -> Vec<ContentBlock> {
|
|||||||
ContentField::Array(parts) => parts
|
ContentField::Array(parts) => parts
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|p| match p {
|
.filter_map(|p| match p {
|
||||||
OpenaiContentPart::Text { text } => {
|
OpenaiContentPart::Text { text } => Some(ContentBlock::Text { text: text.clone() }),
|
||||||
Some(ContentBlock::Text { text: text.clone() })
|
OpenaiContentPart::Refusal { refusal } => Some(ContentBlock::Text {
|
||||||
}
|
text: refusal.clone(),
|
||||||
OpenaiContentPart::Refusal { refusal } => {
|
}),
|
||||||
Some(ContentBlock::Text { text: refusal.clone() })
|
|
||||||
}
|
|
||||||
OpenaiContentPart::Image { image_url, .. } => {
|
OpenaiContentPart::Image { image_url, .. } => {
|
||||||
// ponytail: 简化处理 —— URL 直接通过,data URI 拆出
|
// ponytail: 简化处理 —— URL 直接通过,data URI 拆出
|
||||||
// data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。
|
// data:<mime>;base64,<b64> → ImageSource { data: b64, mime, is_url: false }。
|
||||||
let url = &image_url.url;
|
let url = &image_url.url;
|
||||||
if let Some(rest) = url.strip_prefix("data:")
|
if let Some(rest) = url.strip_prefix("data:")
|
||||||
&& let Some((mime, b64)) = rest.split_once(";base64,") {
|
&& let Some((mime, b64)) = rest.split_once(";base64,")
|
||||||
|
{
|
||||||
return Some(ContentBlock::Image {
|
return Some(ContentBlock::Image {
|
||||||
source: crate::llm::types::message::ImageSource {
|
source: crate::llm::types::message::ImageSource {
|
||||||
data: b64.to_string(),
|
data: b64.to_string(),
|
||||||
@@ -263,7 +263,9 @@ mod tests {
|
|||||||
match ir {
|
match ir {
|
||||||
Message::System { content } => {
|
Message::System { content } => {
|
||||||
assert_eq!(content.len(), 1);
|
assert_eq!(content.len(), 1);
|
||||||
assert!(matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful"));
|
assert!(
|
||||||
|
matches!(&content[0], ContentBlock::Text { text } if text == "you are helpful")
|
||||||
|
);
|
||||||
}
|
}
|
||||||
_ => panic!("expected System variant"),
|
_ => panic!("expected System variant"),
|
||||||
}
|
}
|
||||||
@@ -385,10 +387,7 @@ mod tests {
|
|||||||
assert_eq!(parts.len(), 1);
|
assert_eq!(parts.len(), 1);
|
||||||
match &parts[0] {
|
match &parts[0] {
|
||||||
OpenaiContentPart::Image { image_url, .. } => {
|
OpenaiContentPart::Image { image_url, .. } => {
|
||||||
assert_eq!(
|
assert_eq!(image_url.url, "data:image/png;base64,BASE64DATA");
|
||||||
image_url.url,
|
|
||||||
"data:image/png;base64,BASE64DATA"
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
_ => panic!("expected Image part"),
|
_ => panic!("expected Image part"),
|
||||||
}
|
}
|
||||||
|
|||||||
+841
-43
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,183 @@
|
|||||||
|
//! Embedding 抽象 —— 文本向量化接口。
|
||||||
|
//!
|
||||||
|
//! 提供 [`Embedding`] trait 和零依赖的 [`MockEmbedding`] 引用实现。
|
||||||
|
//! 上层可实现此 trait 以对接真实 Embedding Provider(OpenAI、Cohere 等)。
|
||||||
|
//!
|
||||||
|
//! 所有实现使用 [`LlmError`] 作为统一错误类型,与 llm 模块保持一致。
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
|
||||||
|
use crate::llm::error::LlmError;
|
||||||
|
|
||||||
|
/// 文本向量化抽象接口。
|
||||||
|
///
|
||||||
|
/// 将文本字符串转换为固定维度的浮点向量,用于语义相似度计算。
|
||||||
|
/// 设计为异步以支持网络 IO(如 OpenAI Embedding API)。
|
||||||
|
///
|
||||||
|
/// 使用 [`LlmError`] 作为统一错误类型,与 llm 模块保持一致。
|
||||||
|
///
|
||||||
|
/// # 实现要求
|
||||||
|
///
|
||||||
|
/// - `embed()` 返回的向量外层的 Vec 长度必须等于输入切片长度(一对一映射)
|
||||||
|
/// - 内层 Vec 长度必须等于 `dim()` 返回值
|
||||||
|
/// - 调用方应保证输入非空(空切片返回空外层 Vec,不报错)
|
||||||
|
///
|
||||||
|
/// # 稳定性
|
||||||
|
///
|
||||||
|
/// 实验性 API(v0.3.x),方法签名可能在 v0.4 中调整。
|
||||||
|
#[async_trait]
|
||||||
|
pub trait Embedding: Send + Sync {
|
||||||
|
/// 批量向量化。
|
||||||
|
///
|
||||||
|
/// 返回 `Vec<Vec<f32>>`,第 i 个内层向量对应 `input[i]`。
|
||||||
|
async fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>, LlmError>;
|
||||||
|
|
||||||
|
/// 返回向量维度。
|
||||||
|
fn dim(&self) -> usize;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 确定性 Mock Embedding —— 零依赖伪随机单位向量。
|
||||||
|
///
|
||||||
|
/// 使用 sin 哈希将输入字符串映射到单位球面上的一个点:
|
||||||
|
/// 1. 对输入字符串计算简单哈希(字符字节和 + 长度)作为种子
|
||||||
|
/// 2. 用 `f32::sin(seed + i) * 10000` 生成第 i 个维度的值
|
||||||
|
/// 3. 归一化到单位长度(L2 norm = 1.0)
|
||||||
|
///
|
||||||
|
/// 特性:
|
||||||
|
/// - **确定性**:相同输入 → 相同向量
|
||||||
|
/// - **有区分度**:不同输入产生不同向量(高概率)
|
||||||
|
/// - **单位范数**:余弦相似度等价于点积
|
||||||
|
/// - **开销极低**:不分配额外内存,无 IO
|
||||||
|
///
|
||||||
|
/// # 已知限制
|
||||||
|
///
|
||||||
|
/// `f32::sin(seed + i) * 10000` 在维度较高时(如 1536,OpenAI Embedding 维度)
|
||||||
|
/// 可能出现周期性模式——相邻维度取值在 `sin` 周期 2π 约束下呈规律性重复。
|
||||||
|
/// MockEmbedding 仅用于测试验证,**不应用于生产级相似度排序**;
|
||||||
|
/// 做严肃验证时建议使用真实 Embedding Provider 或显式随机初始化。
|
||||||
|
pub struct MockEmbedding {
|
||||||
|
dim: usize,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl MockEmbedding {
|
||||||
|
/// 创建 Mock Embedding,输出向量维度为 `dim`。
|
||||||
|
pub fn new(dim: usize) -> Self {
|
||||||
|
Self { dim }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl Embedding for MockEmbedding {
|
||||||
|
async fn embed(&self, input: &[String]) -> Result<Vec<Vec<f32>>, LlmError> {
|
||||||
|
let results: Vec<Vec<f32>> = input
|
||||||
|
.iter()
|
||||||
|
.map(|text| {
|
||||||
|
// 简单哈希:字符字节值和 + 文本长度作为种子
|
||||||
|
let seed: f64 = text.bytes().map(|b| b as f64).sum::<f64>() + text.len() as f64;
|
||||||
|
let mut vec: Vec<f32> = (0..self.dim)
|
||||||
|
.map(|i| f32::sin(seed as f32 + i as f32) * 10000.0)
|
||||||
|
.collect();
|
||||||
|
l2_normalize(&mut vec);
|
||||||
|
vec
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
Ok(results)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn dim(&self) -> usize {
|
||||||
|
self.dim
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// L2 归一化(in-place)。
|
||||||
|
///
|
||||||
|
/// 零向量(norm == 0)保持全零 —— 防除零保护。
|
||||||
|
fn l2_normalize(vec: &mut [f32]) {
|
||||||
|
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
|
||||||
|
if norm > f32::EPSILON {
|
||||||
|
for x in vec.iter_mut() {
|
||||||
|
*x /= norm;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
/// 计算向量的 L2 范数。
|
||||||
|
fn l2_norm(v: &[f32]) -> f32 {
|
||||||
|
v.iter().map(|x| x * x).sum::<f32>().sqrt()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn embed_correct_dim() {
|
||||||
|
let embedder = MockEmbedding::new(8);
|
||||||
|
let inputs = vec!["hello".to_string(), "world".to_string()];
|
||||||
|
let result = embedder.embed(&inputs).await.unwrap();
|
||||||
|
assert_eq!(result.len(), 2);
|
||||||
|
for vec in &result {
|
||||||
|
assert_eq!(vec.len(), 8);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn embed_batch_size_match() {
|
||||||
|
let embedder = MockEmbedding::new(4);
|
||||||
|
let inputs = vec![
|
||||||
|
"a".to_string(),
|
||||||
|
"b".to_string(),
|
||||||
|
"c".to_string(),
|
||||||
|
"d".to_string(),
|
||||||
|
"e".to_string(),
|
||||||
|
];
|
||||||
|
let result = embedder.embed(&inputs).await.unwrap();
|
||||||
|
assert_eq!(result.len(), inputs.len());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn embed_deterministic() {
|
||||||
|
let embedder = MockEmbedding::new(4);
|
||||||
|
let inputs = vec!["deterministic test".to_string()];
|
||||||
|
let r1 = embedder.embed(&inputs).await.unwrap();
|
||||||
|
let r2 = embedder.embed(&inputs).await.unwrap();
|
||||||
|
assert_eq!(r1, r2);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn embed_unit_vector_norm() {
|
||||||
|
let embedder = MockEmbedding::new(16);
|
||||||
|
let inputs = vec!["any text".to_string(), "another".to_string()];
|
||||||
|
let result = embedder.embed(&inputs).await.unwrap();
|
||||||
|
for vec in &result {
|
||||||
|
let norm = l2_norm(vec);
|
||||||
|
assert!((norm - 1.0).abs() < 1e-5, "vector norm should be ~1.0, got {}", norm);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn embed_different_inputs_different_vectors() {
|
||||||
|
let embedder = MockEmbedding::new(16);
|
||||||
|
let r1 = embedder
|
||||||
|
.embed(&["hello world".to_string()])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let r2 = embedder
|
||||||
|
.embed(&["completely different".to_string()])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_ne!(r1, r2);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn embed_empty_string() {
|
||||||
|
// 空字符串输入应不 panic,且向量范数仍≈1.0(防除零路径)
|
||||||
|
let embedder = MockEmbedding::new(4);
|
||||||
|
let inputs = vec!["".to_string()];
|
||||||
|
let result = embedder.embed(&inputs).await.unwrap();
|
||||||
|
assert_eq!(result.len(), 1);
|
||||||
|
assert_eq!(result[0].len(), 4);
|
||||||
|
let norm = l2_norm(&result[0]);
|
||||||
|
assert!((norm - 1.0).abs() < 1e-5, "empty-string vector norm should be ~1.0, got {}", norm);
|
||||||
|
}
|
||||||
|
}
|
||||||
+10
-3
@@ -8,9 +8,12 @@ use std::time::Duration;
|
|||||||
///
|
///
|
||||||
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查 API key、减少上下文)。
|
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查 API key、减少上下文)。
|
||||||
#[derive(thiserror::Error, Debug)]
|
#[derive(thiserror::Error, Debug)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum LlmError {
|
pub enum LlmError {
|
||||||
/// API 认证失败(API key 无效、过期或权限不足)。
|
/// API 认证失败(API key 无效、过期或权限不足)。
|
||||||
#[error("LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确")]
|
#[error(
|
||||||
|
"LLM 认证失败: {0}。请检查环境变量中的 API key(如 OPENAI_API_KEY / ANTHROPIC_API_KEY)是否正确"
|
||||||
|
)]
|
||||||
Authentication(String),
|
Authentication(String),
|
||||||
|
|
||||||
/// 请求被限流,可选地附带重试等待时间。可重试。
|
/// 请求被限流,可选地附带重试等待时间。可重试。
|
||||||
@@ -18,7 +21,9 @@ pub enum LlmError {
|
|||||||
RateLimit { retry_after: Option<Duration> },
|
RateLimit { retry_after: Option<Duration> },
|
||||||
|
|
||||||
/// HTTP 请求失败(网络错误或非 2xx 状态码),包含状态码与响应体。
|
/// HTTP 请求失败(网络错误或非 2xx 状态码),包含状态码与响应体。
|
||||||
#[error("LLM 请求失败(HTTP {status}): {body}。请检查 Provider 端点地址(base_url)和网络连通性")]
|
#[error(
|
||||||
|
"LLM 请求失败(HTTP {status}): {body}。请检查 Provider 端点地址(base_url)和网络连通性"
|
||||||
|
)]
|
||||||
Request { status: u16, body: String },
|
Request { status: u16, body: String },
|
||||||
|
|
||||||
/// 请求超时。可重试。
|
/// 请求超时。可重试。
|
||||||
@@ -30,7 +35,9 @@ pub enum LlmError {
|
|||||||
Stream(String),
|
Stream(String),
|
||||||
|
|
||||||
/// 上下文长度超出模型窗口限制。
|
/// 上下文长度超出模型窗口限制。
|
||||||
#[error("LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compaction(llm::compact)")]
|
#[error(
|
||||||
|
"LLM 上下文超限:当前 {actual} tokens > 模型上限 {limit} tokens。请减少消息历史、缩短 prompt,或启用 auto-compaction(llm::compact)"
|
||||||
|
)]
|
||||||
ContextLength { actual: u32, limit: u32 },
|
ContextLength { actual: u32, limit: u32 },
|
||||||
|
|
||||||
/// 其他未分类的 LLM 调用失败。
|
/// 其他未分类的 LLM 调用失败。
|
||||||
|
|||||||
+2
-3
@@ -7,6 +7,7 @@ use crate::llm::types::request_v2::MessageRequest;
|
|||||||
|
|
||||||
/// 生命周期钩子事件点。
|
/// 生命周期钩子事件点。
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum HookEvent {
|
pub enum HookEvent {
|
||||||
/// LLM 请求发起之前(可阻断)。
|
/// LLM 请求发起之前(可阻断)。
|
||||||
PreRequest,
|
PreRequest,
|
||||||
@@ -130,9 +131,7 @@ impl Default for HookExecutor {
|
|||||||
impl HookExecutor {
|
impl HookExecutor {
|
||||||
/// 创建一个空的执行器。
|
/// 创建一个空的执行器。
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self { hooks: Vec::new() }
|
||||||
hooks: Vec::new(),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 注册一个钩子到指定事件点。
|
/// 注册一个钩子到指定事件点。
|
||||||
|
|||||||
+10
-8
@@ -97,8 +97,7 @@ impl LlmProvider for MockProvider {
|
|||||||
async fn chat_stream(
|
async fn chat_stream(
|
||||||
&self,
|
&self,
|
||||||
_request: MessageRequest,
|
_request: MessageRequest,
|
||||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||||
{
|
|
||||||
let response = self.pop()?;
|
let response = self.pop()?;
|
||||||
// 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。
|
// 提前 clone 出在 stream 闭包中需要的字段;最后 yield 时 move response。
|
||||||
let id = response.id.clone();
|
let id = response.id.clone();
|
||||||
@@ -206,10 +205,7 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn chat_returns_queued_response() {
|
async fn chat_returns_queued_response() {
|
||||||
let provider = MockProvider::new(vec![text_response("hello")]);
|
let provider = MockProvider::new(vec![text_response("hello")]);
|
||||||
let resp = provider
|
let resp = provider.chat(MessageRequest::default()).await.unwrap();
|
||||||
.chat(MessageRequest::default())
|
|
||||||
.await
|
|
||||||
.unwrap();
|
|
||||||
assert_eq!(resp.text(), "hello");
|
assert_eq!(resp.text(), "hello");
|
||||||
assert_eq!(provider.remaining(), 0);
|
assert_eq!(provider.remaining(), 0);
|
||||||
}
|
}
|
||||||
@@ -231,7 +227,10 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn chat_stream_emits_text_delta_sequence() {
|
async fn chat_stream_emits_text_delta_sequence() {
|
||||||
let provider = MockProvider::new(vec![text_response("hi")]);
|
let provider = MockProvider::new(vec![text_response("hi")]);
|
||||||
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
|
let mut stream = provider
|
||||||
|
.chat_stream(MessageRequest::default())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
let mut seen_start = false;
|
let mut seen_start = false;
|
||||||
let mut seen_block_start = false;
|
let mut seen_block_start = false;
|
||||||
@@ -283,7 +282,10 @@ mod tests {
|
|||||||
extra: Default::default(),
|
extra: Default::default(),
|
||||||
};
|
};
|
||||||
let provider = MockProvider::new(vec![response]);
|
let provider = MockProvider::new(vec![response]);
|
||||||
let mut stream = provider.chat_stream(MessageRequest::default()).await.unwrap();
|
let mut stream = provider
|
||||||
|
.chat_stream(MessageRequest::default())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
let mut saw_tool_args = false;
|
let mut saw_tool_args = false;
|
||||||
let mut saw_tool_end = false;
|
let mut saw_tool_end = false;
|
||||||
|
|||||||
+444
-9
@@ -1,11 +1,14 @@
|
|||||||
pub mod anthropic;
|
pub mod anthropic;
|
||||||
|
pub mod ollama;
|
||||||
pub mod openai;
|
pub mod openai;
|
||||||
pub mod openai_compat;
|
pub mod openai_compat;
|
||||||
pub mod registry;
|
pub mod registry;
|
||||||
|
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
use futures_core::Stream;
|
use futures_core::Stream;
|
||||||
|
use reqwest::Client;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
use crate::llm::error::LlmError;
|
use crate::llm::error::LlmError;
|
||||||
@@ -18,6 +21,7 @@ use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
|||||||
/// 当前协议数量(5 种以内)完全可控,enum 的编译期安全检查优于运行时的 `HashMap::get()`。
|
/// 当前协议数量(5 种以内)完全可控,enum 的编译期安全检查优于运行时的 `HashMap::get()`。
|
||||||
/// 未来如果扩展到 15+ 种以上,再改为注册表模式。
|
/// 未来如果扩展到 15+ 种以上,再改为注册表模式。
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum ProviderType {
|
pub enum ProviderType {
|
||||||
/// OpenAI Chat Completions API(兼容 DeepSeek / Qwen 等 `/chat/completions` 端点)。
|
/// OpenAI Chat Completions API(兼容 DeepSeek / Qwen 等 `/chat/completions` 端点)。
|
||||||
OpenaiChat,
|
OpenaiChat,
|
||||||
@@ -29,6 +33,8 @@ pub enum ProviderType {
|
|||||||
DeepSeek,
|
DeepSeek,
|
||||||
/// Qwen / 阿里云百炼(OpenAI-compatible `/chat/completions`)。
|
/// Qwen / 阿里云百炼(OpenAI-compatible `/chat/completions`)。
|
||||||
Qwen,
|
Qwen,
|
||||||
|
/// Ollama 本地推理(OpenAI-compatible `/chat/completions`,默认 `http://localhost:11434/v1`)。
|
||||||
|
Ollama,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl std::str::FromStr for ProviderType {
|
impl std::str::FromStr for ProviderType {
|
||||||
@@ -41,47 +47,211 @@ impl std::str::FromStr for ProviderType {
|
|||||||
"anthropic" | "claude" => Ok(ProviderType::Anthropic),
|
"anthropic" | "claude" => Ok(ProviderType::Anthropic),
|
||||||
"deepseek" => Ok(ProviderType::DeepSeek),
|
"deepseek" => Ok(ProviderType::DeepSeek),
|
||||||
"qwen" | "dashscope" | "tongyi" => Ok(ProviderType::Qwen),
|
"qwen" | "dashscope" | "tongyi" => Ok(ProviderType::Qwen),
|
||||||
|
"ollama" => Ok(ProviderType::Ollama),
|
||||||
_ => Err(format!("未知的 Provider 类型: {s}")),
|
_ => Err(format!("未知的 Provider 类型: {s}")),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Provider 构造参数 —— 通用 base_url + api_key + model。
|
/// Provider 构造参数 —— 通用 base_url + api_key + model + timeout / retry 配置。
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
pub struct ProviderConfig {
|
pub struct ProviderConfig {
|
||||||
|
/// API base URL(如 `https://api.openai.com/v1`)。为空时由 Provider 选择默认值。
|
||||||
pub base_url: String,
|
pub base_url: String,
|
||||||
|
/// API key。Ollama 等本地 Provider 可为空。
|
||||||
pub api_key: String,
|
pub api_key: String,
|
||||||
|
/// 模型名(如 `gpt-4o` / `claude-sonnet-4-20250514`)。
|
||||||
pub model: String,
|
pub model: String,
|
||||||
|
/// 请求超时秒数(默认 30)。应用于 Provider 的 HTTP Client 级别。
|
||||||
|
pub timeout_secs: u64,
|
||||||
|
/// 最大重试次数(默认 3)。
|
||||||
|
///
|
||||||
|
/// 当前此字段仅由 `from_env()` 采集,**实际重试逻辑由 `CycleConfig.retry.max_retries` 控制**。
|
||||||
|
/// 此处保留字段以与 Roadmap §Phase 5 Step 5.1 对齐;未来 Phase 6+ 可统一合并到 `CycleConfig`。
|
||||||
|
pub max_retries: u32,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for ProviderConfig {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self {
|
||||||
|
base_url: String::new(),
|
||||||
|
api_key: String::new(),
|
||||||
|
model: String::new(),
|
||||||
|
timeout_secs: 30,
|
||||||
|
max_retries: 3,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProviderConfig {
|
||||||
|
/// 从环境变量构造 `ProviderConfig`。
|
||||||
|
///
|
||||||
|
/// 必填变量:
|
||||||
|
/// - `{prefix}_BASE_URL`
|
||||||
|
/// - `{prefix}_API_KEY`
|
||||||
|
/// - `{prefix}_MODEL`
|
||||||
|
///
|
||||||
|
/// 可选变量(有默认值):
|
||||||
|
/// - `{prefix}_TIMEOUT_SECS`(默认 30,解析失败回退 30 并 warn)
|
||||||
|
/// - `{prefix}_MAX_RETRIES`(默认 3,解析失败回退 3 并 warn)
|
||||||
|
pub fn from_env(prefix: &str) -> Result<Self, String> {
|
||||||
|
let base_url = std::env::var(format!("{prefix}_BASE_URL"))
|
||||||
|
.map_err(|_| format!("{prefix}_BASE_URL 环境变量未设置"))?;
|
||||||
|
let api_key = std::env::var(format!("{prefix}_API_KEY"))
|
||||||
|
.map_err(|_| format!("{prefix}_API_KEY 环境变量未设置"))?;
|
||||||
|
let model = std::env::var(format!("{prefix}_MODEL"))
|
||||||
|
.map_err(|_| format!("{prefix}_MODEL 环境变量未设置"))?;
|
||||||
|
let timeout_secs = match std::env::var(format!("{prefix}_TIMEOUT_SECS")) {
|
||||||
|
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||||
|
tracing::warn!("{prefix}_TIMEOUT_SECS='{v}' 解析失败,使用默认值 30");
|
||||||
|
30
|
||||||
|
}),
|
||||||
|
Err(_) => 30,
|
||||||
|
};
|
||||||
|
let max_retries = match std::env::var(format!("{prefix}_MAX_RETRIES")) {
|
||||||
|
Ok(v) => v.parse().unwrap_or_else(|_| {
|
||||||
|
tracing::warn!("{prefix}_MAX_RETRIES='{v}' 解析失败,使用默认值 3");
|
||||||
|
3
|
||||||
|
}),
|
||||||
|
Err(_) => 3,
|
||||||
|
};
|
||||||
|
|
||||||
|
// ponytail: max_retries 当前仅采集,不传入 Provider。
|
||||||
|
// 实际重试由 CycleConfig.retry.max_retries 控制。
|
||||||
|
if max_retries != 3 {
|
||||||
|
tracing::warn!(
|
||||||
|
"ProviderConfig.max_retries={} 已采集但当前未生效;\
|
||||||
|
重试次数由 CycleConfig.retry.max_retries 控制",
|
||||||
|
max_retries,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(Self {
|
||||||
|
base_url,
|
||||||
|
api_key,
|
||||||
|
model,
|
||||||
|
timeout_secs,
|
||||||
|
max_retries,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 构造带 timeout 的 `reqwest::Client`(OpenAI-compatible 共享)。
|
||||||
|
fn build_client_with_timeout(timeout_secs: u64) -> Result<Client, LlmError> {
|
||||||
|
Client::builder()
|
||||||
|
.timeout(Duration::from_secs(timeout_secs))
|
||||||
|
.build()
|
||||||
|
.map_err(|e| LlmError::Other(format!("创建 HTTP 客户端失败: {e}")))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 构造带 Anthropic 默认 headers + timeout 的 `reqwest::Client`。
|
||||||
|
///
|
||||||
|
/// Anthropic 由于需要保留 `x-api-key` / `anthropic-version` 默认 headers,
|
||||||
|
/// 与 OpenAI-compatible 共享的 `build_client_with_timeout` 不同。
|
||||||
|
fn build_anthropic_client(api_key: &str, timeout_secs: u64) -> Result<Client, LlmError> {
|
||||||
|
use reqwest::header::{HeaderMap, HeaderValue};
|
||||||
|
|
||||||
|
let key_header = HeaderValue::from_str(api_key)
|
||||||
|
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
|
||||||
|
let version_header = HeaderValue::from_static("2023-06-01");
|
||||||
|
|
||||||
|
Client::builder()
|
||||||
|
.timeout(Duration::from_secs(timeout_secs))
|
||||||
|
.default_headers({
|
||||||
|
let mut headers = HeaderMap::new();
|
||||||
|
headers.insert("x-api-key", key_header);
|
||||||
|
headers.insert("anthropic-version", version_header);
|
||||||
|
headers
|
||||||
|
})
|
||||||
|
.build()
|
||||||
|
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Provider 工厂 —— exhaustive match 在编译期保证新 Provider 被注册。
|
/// Provider 工厂 —— exhaustive match 在编译期保证新 Provider 被注册。
|
||||||
|
///
|
||||||
|
/// `config.timeout_secs` 注入到 Provider 的 HTTP Client 超时配置。
|
||||||
|
/// 每个分支通过 `from_parts` (pub(crate)) 一次性构造,无冗余 client 创建。
|
||||||
pub fn create_provider(
|
pub fn create_provider(
|
||||||
provider_type: ProviderType,
|
provider_type: ProviderType,
|
||||||
config: ProviderConfig,
|
config: ProviderConfig,
|
||||||
) -> Result<Box<dyn LlmProvider>, LlmError> {
|
) -> Result<Box<dyn LlmProvider>, LlmError> {
|
||||||
match provider_type {
|
match provider_type {
|
||||||
ProviderType::OpenaiChat => Ok(Box::new(openai::OpenaiChatProvider::new(
|
ProviderType::OpenaiChat => {
|
||||||
|
let client = build_client_with_timeout(config.timeout_secs)?;
|
||||||
|
Ok(Box::new(openai::OpenaiChatProvider(
|
||||||
|
openai::GenericOpenaiProvider::from_parts(
|
||||||
config.base_url,
|
config.base_url,
|
||||||
config.api_key,
|
config.api_key,
|
||||||
config.model,
|
config.model,
|
||||||
))),
|
"openai",
|
||||||
|
client,
|
||||||
|
Vec::new(),
|
||||||
|
config.timeout_secs,
|
||||||
|
),
|
||||||
|
)))
|
||||||
|
}
|
||||||
ProviderType::OpenaiResponse => Err(LlmError::Other(
|
ProviderType::OpenaiResponse => Err(LlmError::Other(
|
||||||
"OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(),
|
"OpenaiResponse Provider 在 Phase 1 暂不实现;请使用 OpenaiChat".into(),
|
||||||
)),
|
)),
|
||||||
ProviderType::Anthropic => Ok(Box::new(anthropic::AnthropicProvider::new(
|
ProviderType::Anthropic => {
|
||||||
|
let client = build_anthropic_client(&config.api_key, config.timeout_secs)?;
|
||||||
|
Ok(Box::new(anthropic::AnthropicProvider::from_parts(
|
||||||
config.base_url,
|
config.base_url,
|
||||||
config.api_key,
|
config.api_key,
|
||||||
config.model,
|
config.model,
|
||||||
))),
|
client,
|
||||||
ProviderType::DeepSeek => Ok(Box::new(openai_compat::DeepSeekProvider::new(
|
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.base_url,
|
||||||
config.api_key,
|
config.api_key,
|
||||||
config.model,
|
config.model,
|
||||||
))),
|
"deepseek",
|
||||||
ProviderType::Qwen => Ok(Box::new(openai_compat::QwenProvider::new(
|
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.base_url,
|
||||||
config.api_key,
|
config.api_key,
|
||||||
config.model,
|
config.model,
|
||||||
))),
|
"qwen",
|
||||||
|
client,
|
||||||
|
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
||||||
|
config.timeout_secs,
|
||||||
|
),
|
||||||
|
)))
|
||||||
|
}
|
||||||
|
ProviderType::Ollama => {
|
||||||
|
let client = build_client_with_timeout(config.timeout_secs)?;
|
||||||
|
// ponytail: Ollama 默认 base_url 由 OllamaProvider 构造处理 —— 但 from_parts 不走
|
||||||
|
// OllamaProvider::new 的默认 URL 回退。这里保留 base_url(可能为空 → http://localhost:11434/v1)。
|
||||||
|
let base_url = if config.base_url.is_empty() {
|
||||||
|
"http://localhost:11434/v1".to_string()
|
||||||
|
} else {
|
||||||
|
config.base_url
|
||||||
|
};
|
||||||
|
Ok(Box::new(ollama::OllamaProvider(
|
||||||
|
openai::GenericOpenaiProvider::from_parts(
|
||||||
|
base_url,
|
||||||
|
config.api_key,
|
||||||
|
config.model,
|
||||||
|
"ollama",
|
||||||
|
client,
|
||||||
|
Vec::new(),
|
||||||
|
config.timeout_secs,
|
||||||
|
),
|
||||||
|
)))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -142,3 +312,268 @@ pub trait LlmProvider: Send + Sync {
|
|||||||
/// 返回 Provider 静态能力描述。
|
/// 返回 Provider 静态能力描述。
|
||||||
fn capabilities(&self) -> ProviderCapabilities;
|
fn capabilities(&self) -> ProviderCapabilities;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_default_values() {
|
||||||
|
let config = ProviderConfig::default();
|
||||||
|
assert_eq!(config.timeout_secs, 30);
|
||||||
|
assert_eq!(config.max_retries, 3);
|
||||||
|
assert_eq!(config.base_url, "");
|
||||||
|
assert_eq!(config.api_key, "");
|
||||||
|
assert_eq!(config.model, "");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_from_env_requires_all_three() {
|
||||||
|
// 使用 temp_env 移除所有相关变量,避免外部环境意外设置导致测试 flaky
|
||||||
|
temp_env::with_vars(
|
||||||
|
[
|
||||||
|
("TEST_PROVIDER_MISSING_BASE_URL", None::<&str>),
|
||||||
|
("TEST_PROVIDER_MISSING_API_KEY", None::<&str>),
|
||||||
|
("TEST_PROVIDER_MISSING_MODEL", None::<&str>),
|
||||||
|
("TEST_PROVIDER_MISSING_TIMEOUT_SECS", None::<&str>),
|
||||||
|
("TEST_PROVIDER_MISSING_MAX_RETRIES", None::<&str>),
|
||||||
|
],
|
||||||
|
|| {
|
||||||
|
let result = ProviderConfig::from_env("TEST_PROVIDER_MISSING");
|
||||||
|
assert!(result.is_err());
|
||||||
|
let msg = result.unwrap_err();
|
||||||
|
assert!(
|
||||||
|
msg.contains("TEST_PROVIDER_MISSING_BASE_URL"),
|
||||||
|
"error should mention missing var, got: {msg}"
|
||||||
|
);
|
||||||
|
},
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_from_env_uses_defaults_when_only_required_set() {
|
||||||
|
temp_env::with_vars(
|
||||||
|
[
|
||||||
|
("TEST_PROVIDER_BASE_URL", Some("http://localhost:11434/v1")),
|
||||||
|
("TEST_PROVIDER_API_KEY", Some("")),
|
||||||
|
("TEST_PROVIDER_MODEL", Some("llama3")),
|
||||||
|
],
|
||||||
|
|| {
|
||||||
|
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||||
|
assert_eq!(config.base_url, "http://localhost:11434/v1");
|
||||||
|
assert_eq!(config.api_key, "");
|
||||||
|
assert_eq!(config.model, "llama3");
|
||||||
|
assert_eq!(config.timeout_secs, 30);
|
||||||
|
assert_eq!(config.max_retries, 3);
|
||||||
|
},
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_from_env_reads_custom_values() {
|
||||||
|
temp_env::with_vars(
|
||||||
|
[
|
||||||
|
("TEST_PROVIDER_BASE_URL", Some("http://x")),
|
||||||
|
("TEST_PROVIDER_API_KEY", Some("k")),
|
||||||
|
("TEST_PROVIDER_MODEL", Some("m")),
|
||||||
|
("TEST_PROVIDER_TIMEOUT_SECS", Some("60")),
|
||||||
|
("TEST_PROVIDER_MAX_RETRIES", Some("5")),
|
||||||
|
],
|
||||||
|
|| {
|
||||||
|
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||||
|
assert_eq!(config.timeout_secs, 60);
|
||||||
|
assert_eq!(config.max_retries, 5);
|
||||||
|
},
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn provider_config_from_env_falls_back_on_invalid_numbers() {
|
||||||
|
temp_env::with_vars(
|
||||||
|
[
|
||||||
|
("TEST_PROVIDER_BASE_URL", Some("http://x")),
|
||||||
|
("TEST_PROVIDER_API_KEY", Some("k")),
|
||||||
|
("TEST_PROVIDER_MODEL", Some("m")),
|
||||||
|
("TEST_PROVIDER_TIMEOUT_SECS", Some("not-a-number")),
|
||||||
|
("TEST_PROVIDER_MAX_RETRIES", Some("also-bad")),
|
||||||
|
],
|
||||||
|
|| {
|
||||||
|
let config = ProviderConfig::from_env("TEST_PROVIDER").unwrap();
|
||||||
|
// 解析失败回退默认值
|
||||||
|
assert_eq!(config.timeout_secs, 30);
|
||||||
|
assert_eq!(config.max_retries, 3);
|
||||||
|
},
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Timeout 传导集成测试:构造 `ProviderConfig` timeout=1s,
|
||||||
|
/// `create_provider` 注入 1s 超时 client,请求一个故意延迟 3s 的 mock server,
|
||||||
|
/// 验证返回 `LlmError::Timeout { duration: 1s }`。
|
||||||
|
#[tokio::test]
|
||||||
|
async fn create_provider_injects_timeout_into_openai_chat() {
|
||||||
|
use serde_json::json;
|
||||||
|
use wiremock::matchers::{method, path};
|
||||||
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
|
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
// 故意延迟 3s 触发超时
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/chat/completions"))
|
||||||
|
.respond_with(
|
||||||
|
ResponseTemplate::new(200)
|
||||||
|
.set_delay(Duration::from_secs(3))
|
||||||
|
.set_body_json(json!({
|
||||||
|
"id": "x",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": 0,
|
||||||
|
"model": "gpt-4o",
|
||||||
|
"choices": [],
|
||||||
|
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||||
|
})),
|
||||||
|
)
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = create_provider(
|
||||||
|
ProviderType::OpenaiChat,
|
||||||
|
ProviderConfig {
|
||||||
|
base_url: server.uri(),
|
||||||
|
api_key: "sk-test".into(),
|
||||||
|
model: "gpt-4o".into(),
|
||||||
|
timeout_secs: 1,
|
||||||
|
max_retries: 3,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let err = provider
|
||||||
|
.chat(crate::llm::types::request_v2::MessageRequest {
|
||||||
|
model: "gpt-4o".into(),
|
||||||
|
messages: vec![],
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap_err();
|
||||||
|
match err {
|
||||||
|
LlmError::Timeout { duration } => {
|
||||||
|
assert_eq!(duration, Duration::from_secs(1));
|
||||||
|
}
|
||||||
|
other => panic!("expected Timeout, got {other:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Timeout 传导验证:`create_provider` 生成的 DeepSeek Provider 也带 1s 超时,
|
||||||
|
/// 错误消息中的 duration 与 timeout_secs 一致(而非硬编码 120s)。
|
||||||
|
#[tokio::test]
|
||||||
|
async fn create_provider_injects_timeout_into_deepseek() {
|
||||||
|
use serde_json::json;
|
||||||
|
use wiremock::matchers::{method, path};
|
||||||
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
|
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/chat/completions"))
|
||||||
|
.respond_with(
|
||||||
|
ResponseTemplate::new(200)
|
||||||
|
.set_delay(Duration::from_secs(3))
|
||||||
|
.set_body_json(json!({
|
||||||
|
"id": "x",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": 0,
|
||||||
|
"model": "deepseek-chat",
|
||||||
|
"choices": [],
|
||||||
|
"usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||||
|
})),
|
||||||
|
)
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = create_provider(
|
||||||
|
ProviderType::DeepSeek,
|
||||||
|
ProviderConfig {
|
||||||
|
base_url: server.uri(),
|
||||||
|
api_key: "sk-test".into(),
|
||||||
|
model: "deepseek-chat".into(),
|
||||||
|
timeout_secs: 1,
|
||||||
|
max_retries: 3,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let err = provider
|
||||||
|
.chat(crate::llm::types::request_v2::MessageRequest {
|
||||||
|
model: "deepseek-chat".into(),
|
||||||
|
messages: vec![],
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap_err();
|
||||||
|
match err {
|
||||||
|
LlmError::Timeout { duration } => {
|
||||||
|
assert_eq!(duration, Duration::from_secs(1));
|
||||||
|
}
|
||||||
|
other => panic!("expected Timeout, got {other:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Timeout 传导验证:`create_provider` 生成的 Anthropic Provider 通过 `with_timeout`
|
||||||
|
/// 注入 1s 超时。
|
||||||
|
///
|
||||||
|
/// 与 OpenAI-compatible 路径不同,Anthropic 走 `AnthropicProvider::with_timeout()`
|
||||||
|
/// 重建底层 client(保留 default_headers),独立于 OpenAI-compatible 的 `build_client_with_timeout`。
|
||||||
|
/// 单独覆盖此路径以验证 `with_timeout` 不会因服务端延迟而返回硬编码 120s 的超时错误。
|
||||||
|
#[tokio::test]
|
||||||
|
async fn create_provider_injects_timeout_into_anthropic() {
|
||||||
|
use serde_json::json;
|
||||||
|
use wiremock::matchers::{method, path};
|
||||||
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
|
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
// Anthropic Messages API 端点:`POST /v1/messages`
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/v1/messages"))
|
||||||
|
.respond_with(
|
||||||
|
ResponseTemplate::new(200)
|
||||||
|
.set_delay(Duration::from_secs(3))
|
||||||
|
.set_body_json(json!({
|
||||||
|
"id": "msg_timeout_test",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [{"type": "text", "text": "ok"}],
|
||||||
|
"model": "claude-sonnet-4-20250514",
|
||||||
|
"stop_reason": "end_turn",
|
||||||
|
"usage": {"input_tokens": 1, "output_tokens": 1}
|
||||||
|
})),
|
||||||
|
)
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = create_provider(
|
||||||
|
ProviderType::Anthropic,
|
||||||
|
ProviderConfig {
|
||||||
|
base_url: server.uri(),
|
||||||
|
api_key: "sk-ant-test".into(),
|
||||||
|
model: "claude-sonnet-4-20250514".into(),
|
||||||
|
timeout_secs: 1,
|
||||||
|
max_retries: 3,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let err = provider
|
||||||
|
.chat(crate::llm::types::request_v2::MessageRequest {
|
||||||
|
model: "claude-sonnet-4-20250514".into(),
|
||||||
|
messages: vec![],
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap_err();
|
||||||
|
match err {
|
||||||
|
LlmError::Timeout { duration } => {
|
||||||
|
assert_eq!(duration, Duration::from_secs(1));
|
||||||
|
}
|
||||||
|
other => panic!("expected Timeout, got {other:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+261
-54
@@ -12,10 +12,10 @@ use async_trait::async_trait;
|
|||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use futures_core::Stream;
|
use futures_core::Stream;
|
||||||
use futures_util::StreamExt;
|
use futures_util::StreamExt;
|
||||||
use reqwest::header::{HeaderMap, HeaderValue};
|
|
||||||
use reqwest::Client;
|
use reqwest::Client;
|
||||||
|
use reqwest::header::{HeaderMap, HeaderValue};
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::{json, Value};
|
use serde_json::{Value, json};
|
||||||
use tracing::{debug, error, info, warn};
|
use tracing::{debug, error, info, warn};
|
||||||
|
|
||||||
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
|
use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
|
||||||
@@ -39,16 +39,20 @@ pub struct AnthropicProvider {
|
|||||||
#[allow(dead_code)]
|
#[allow(dead_code)]
|
||||||
api_key: String,
|
api_key: String,
|
||||||
model: String,
|
model: String,
|
||||||
|
/// HTTP 请求超时秒数。由 `ProviderConfig::timeout_secs` 传入,
|
||||||
|
/// 在 `LlmError::Timeout { duration }` 中回显。`reqwest::Client` 不暴露 timeout getter,
|
||||||
|
/// 因此单独存储以便错误消息与配置保持一致。
|
||||||
|
timeout_secs: u64,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl AnthropicProvider {
|
impl AnthropicProvider {
|
||||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||||
let key_header = HeaderValue::from_str(&api_key)
|
let key_header =
|
||||||
.expect("Anthropic API key 包含无效的 HTTP 头部字符");
|
HeaderValue::from_str(&api_key).expect("Anthropic API key 包含无效的 HTTP 头部字符");
|
||||||
let version_header = HeaderValue::from_static("2023-06-01");
|
let version_header = HeaderValue::from_static("2023-06-01");
|
||||||
|
|
||||||
let http_client = Client::builder()
|
let http_client = Client::builder()
|
||||||
.timeout(Duration::from_secs(120))
|
.timeout(Duration::from_secs(timeout_secs))
|
||||||
.default_headers({
|
.default_headers({
|
||||||
let mut headers = HeaderMap::new();
|
let mut headers = HeaderMap::new();
|
||||||
headers.insert("x-api-key", key_header);
|
headers.insert("x-api-key", key_header);
|
||||||
@@ -67,14 +71,80 @@ impl AnthropicProvider {
|
|||||||
},
|
},
|
||||||
api_key,
|
api_key,
|
||||||
model,
|
model,
|
||||||
|
timeout_secs,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// ⚠️ 替换 HTTP Client,**丢弃** `new()` 中设置的默认 headers(`x-api-key` / `anthropic-version`)。
|
||||||
|
///
|
||||||
|
/// 调用此方法后,所有 Anthropic API 请求将以**无认证头**发送出去,预期会 401/403 失败。
|
||||||
|
/// 推荐改用 [`Self::with_timeout`],它会重建 client 并保留默认 headers。
|
||||||
|
///
|
||||||
|
/// 此方法仍保留以兼容调用方自定义 client 但不需要默认 headers 的极端场景。
|
||||||
|
#[deprecated(
|
||||||
|
since = "0.2.0",
|
||||||
|
note = "此方法会丢弃默认 headers(x-api-key / anthropic-version),改为使用 `with_timeout` 或带 headers 的 `Client::builder()`"
|
||||||
|
)]
|
||||||
pub fn with_client(mut self, client: Client) -> Self {
|
pub fn with_client(mut self, client: Client) -> Self {
|
||||||
self.http_client = client;
|
self.http_client = client;
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 替换 HTTP Client 的超时配置(重建底层 client,保留默认 headers)。
|
||||||
|
///
|
||||||
|
/// ⚠️ 副作用:此方法**完全重建** `http_client`,调用后通过 `with_client` 注入的 Client
|
||||||
|
/// 将被替换。headers 构造逻辑与 `new()` 中的保持一致(`x-api-key` / `anthropic-version`)。
|
||||||
|
///
|
||||||
|
/// ponytail: 同值调用短路。当 `secs == self.timeout_secs` 时跳过 client 重建,
|
||||||
|
/// 避免 `create_provider` 路径 `new(timeout).with_timeout(timeout)` 的双重构造。
|
||||||
|
pub fn with_timeout(mut self, secs: u64) -> Result<Self, LlmError> {
|
||||||
|
if secs == self.timeout_secs {
|
||||||
|
return Ok(self);
|
||||||
|
}
|
||||||
|
// ponytail: 重建 http_client 时保留已有默认 headers(x-api-key / anthropic-version)。
|
||||||
|
// 如后续 AnthropicProvider 的 headers 变为动态,此方法需同步更新。
|
||||||
|
let key_header = HeaderValue::from_str(&self.api_key)
|
||||||
|
.map_err(|_| LlmError::Other("Anthropic API key 包含无效的 HTTP 头部字符".into()))?;
|
||||||
|
let version_header = HeaderValue::from_static("2023-06-01");
|
||||||
|
|
||||||
|
self.http_client = Client::builder()
|
||||||
|
.timeout(Duration::from_secs(secs))
|
||||||
|
.default_headers({
|
||||||
|
let mut headers = HeaderMap::new();
|
||||||
|
headers.insert("x-api-key", key_header);
|
||||||
|
headers.insert("anthropic-version", version_header);
|
||||||
|
headers
|
||||||
|
})
|
||||||
|
.build()
|
||||||
|
.map_err(|e| LlmError::Other(format!("创建 Anthropic HTTP 客户端失败: {e}")))?;
|
||||||
|
self.timeout_secs = secs;
|
||||||
|
Ok(self)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 一次性构造 —— `create_provider` 路径专用,避免 `new(...)` + `with_timeout(...)` 的双重 client 构造。
|
||||||
|
///
|
||||||
|
/// 调用方负责预先构造好符合 Anthropic 协议要求的 `http_client`(带正确的 `x-api-key` /
|
||||||
|
/// `anthropic-version` 默认 headers + 指定 timeout)。
|
||||||
|
pub(crate) fn from_parts(
|
||||||
|
base_url: String,
|
||||||
|
api_key: String,
|
||||||
|
model: String,
|
||||||
|
http_client: Client,
|
||||||
|
timeout_secs: u64,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
http_client,
|
||||||
|
base_url: if base_url.is_empty() {
|
||||||
|
"https://api.anthropic.com".to_string()
|
||||||
|
} else {
|
||||||
|
base_url
|
||||||
|
},
|
||||||
|
api_key,
|
||||||
|
model,
|
||||||
|
timeout_secs,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
|
fn resolve_max_tokens(&self, request: &MessageRequest) -> u32 {
|
||||||
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
|
request.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS)
|
||||||
}
|
}
|
||||||
@@ -104,12 +174,14 @@ impl AnthropicProvider {
|
|||||||
Message::User { content } => {
|
Message::User { content } => {
|
||||||
api_messages.push(AnthropicMessage::user(content));
|
api_messages.push(AnthropicMessage::user(content));
|
||||||
}
|
}
|
||||||
Message::UserImage { data, mime_type, detail } => {
|
Message::UserImage {
|
||||||
|
data,
|
||||||
|
mime_type,
|
||||||
|
detail,
|
||||||
|
} => {
|
||||||
// Anthropic image format: {type: "image", source: {type: "base64", media_type, data}}
|
// Anthropic image format: {type: "image", source: {type: "base64", media_type, data}}
|
||||||
let source = if data.starts_with("http://") || data.starts_with("https://") {
|
let source = if data.starts_with("http://") || data.starts_with("https://") {
|
||||||
AnthropicImageSource::Url {
|
AnthropicImageSource::Url { url: data.clone() }
|
||||||
url: data.clone(),
|
|
||||||
}
|
|
||||||
} else {
|
} else {
|
||||||
AnthropicImageSource::Base64 {
|
AnthropicImageSource::Base64 {
|
||||||
media_type: mime_type.clone(),
|
media_type: mime_type.clone(),
|
||||||
@@ -194,7 +266,7 @@ impl AnthropicProvider {
|
|||||||
.json(&body)
|
.json(&body)
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.map_err(Self::map_reqwest_error)?;
|
.map_err(|e| self.map_reqwest_error(e))?;
|
||||||
|
|
||||||
let status = response.status();
|
let status = response.status();
|
||||||
if !status.is_success() {
|
if !status.is_success() {
|
||||||
@@ -215,8 +287,7 @@ impl AnthropicProvider {
|
|||||||
async fn chat_stream_inner(
|
async fn chat_stream_inner(
|
||||||
&self,
|
&self,
|
||||||
request: MessageRequest,
|
request: MessageRequest,
|
||||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||||
{
|
|
||||||
let mut body = self.build_request_body(request)?;
|
let mut body = self.build_request_body(request)?;
|
||||||
body.stream = Some(true);
|
body.stream = Some(true);
|
||||||
|
|
||||||
@@ -230,16 +301,16 @@ impl AnthropicProvider {
|
|||||||
.json(&body)
|
.json(&body)
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.map_err(Self::map_reqwest_error)?;
|
.map_err(|e| self.map_reqwest_error(e))?;
|
||||||
|
|
||||||
let status = response.status();
|
let status = response.status();
|
||||||
if !status.is_success() {
|
if !status.is_success() {
|
||||||
return Err(Self::handle_error_response(response).await);
|
return Err(Self::handle_error_response(response).await);
|
||||||
}
|
}
|
||||||
|
|
||||||
let byte_stream = response.bytes_stream().map(|r| {
|
let byte_stream = response
|
||||||
r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}")))
|
.bytes_stream()
|
||||||
});
|
.map(|r| r.map_err(|e| LlmError::Other(format!("流式读取失败: {e}"))));
|
||||||
|
|
||||||
let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> =
|
let byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, LlmError>> + Send>> =
|
||||||
Box::pin(byte_stream);
|
Box::pin(byte_stream);
|
||||||
@@ -247,10 +318,10 @@ impl AnthropicProvider {
|
|||||||
Ok(Box::pin(AnthropicSseStream::new(byte_stream)))
|
Ok(Box::pin(AnthropicSseStream::new(byte_stream)))
|
||||||
}
|
}
|
||||||
|
|
||||||
fn map_reqwest_error(e: reqwest::Error) -> LlmError {
|
fn map_reqwest_error(&self, e: reqwest::Error) -> LlmError {
|
||||||
if e.is_timeout() {
|
if e.is_timeout() {
|
||||||
LlmError::Timeout {
|
LlmError::Timeout {
|
||||||
duration: Duration::from_secs(120),
|
duration: Duration::from_secs(self.timeout_secs),
|
||||||
}
|
}
|
||||||
} else if e.is_connect() {
|
} else if e.is_connect() {
|
||||||
LlmError::Other(format!("连接失败: {e}"))
|
LlmError::Other(format!("连接失败: {e}"))
|
||||||
@@ -291,13 +362,12 @@ impl AnthropicProvider {
|
|||||||
blocks.push(ContentBlock::Text { text });
|
blocks.push(ContentBlock::Text { text });
|
||||||
}
|
}
|
||||||
AnthropicContentBlockResp::ToolUse { id, name, input } => {
|
AnthropicContentBlockResp::ToolUse { id, name, input } => {
|
||||||
blocks.push(ContentBlock::ToolUse {
|
blocks.push(ContentBlock::ToolUse { id, name, input });
|
||||||
id,
|
|
||||||
name,
|
|
||||||
input,
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
AnthropicContentBlockResp::Thinking { thinking, signature } => {
|
AnthropicContentBlockResp::Thinking {
|
||||||
|
thinking,
|
||||||
|
signature,
|
||||||
|
} => {
|
||||||
blocks.push(ContentBlock::Thinking {
|
blocks.push(ContentBlock::Thinking {
|
||||||
text: thinking,
|
text: thinking,
|
||||||
signature,
|
signature,
|
||||||
@@ -336,8 +406,7 @@ impl LlmProvider for AnthropicProvider {
|
|||||||
async fn chat_stream(
|
async fn chat_stream(
|
||||||
&self,
|
&self,
|
||||||
request: MessageRequest,
|
request: MessageRequest,
|
||||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||||
{
|
|
||||||
self.chat_stream_inner(request).await
|
self.chat_stream_inner(request).await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -418,7 +487,9 @@ impl AnthropicMessage {
|
|||||||
#[derive(Debug, Serialize)]
|
#[derive(Debug, Serialize)]
|
||||||
#[serde(tag = "type", rename_all = "snake_case")]
|
#[serde(tag = "type", rename_all = "snake_case")]
|
||||||
enum AnthropicContentPart {
|
enum AnthropicContentPart {
|
||||||
Text { text: String },
|
Text {
|
||||||
|
text: String,
|
||||||
|
},
|
||||||
Image {
|
Image {
|
||||||
source: AnthropicImageSource,
|
source: AnthropicImageSource,
|
||||||
},
|
},
|
||||||
@@ -437,13 +508,8 @@ enum AnthropicContentPart {
|
|||||||
#[derive(Debug, Serialize)]
|
#[derive(Debug, Serialize)]
|
||||||
#[serde(tag = "type", rename_all = "snake_case")]
|
#[serde(tag = "type", rename_all = "snake_case")]
|
||||||
enum AnthropicImageSource {
|
enum AnthropicImageSource {
|
||||||
Base64 {
|
Base64 { media_type: String, data: String },
|
||||||
media_type: String,
|
Url { url: String },
|
||||||
data: String,
|
|
||||||
},
|
|
||||||
Url {
|
|
||||||
url: String,
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> {
|
fn content_to_parts(blocks: &[ContentBlock]) -> Vec<AnthropicContentPart> {
|
||||||
@@ -523,9 +589,18 @@ struct AnthropicUsage {
|
|||||||
#[derive(Debug, Deserialize)]
|
#[derive(Debug, Deserialize)]
|
||||||
#[serde(tag = "type", rename_all = "snake_case")]
|
#[serde(tag = "type", rename_all = "snake_case")]
|
||||||
enum AnthropicContentBlockResp {
|
enum AnthropicContentBlockResp {
|
||||||
Text { text: String },
|
Text {
|
||||||
ToolUse { id: String, name: String, input: Value },
|
text: String,
|
||||||
Thinking { thinking: String, signature: Option<String> },
|
},
|
||||||
|
ToolUse {
|
||||||
|
id: String,
|
||||||
|
name: String,
|
||||||
|
input: Value,
|
||||||
|
},
|
||||||
|
Thinking {
|
||||||
|
thinking: String,
|
||||||
|
signature: Option<String>,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
@@ -662,18 +737,13 @@ impl AnthropicSseStream {
|
|||||||
// 先把所有字段提前,避免 match 中 part-move
|
// 先把所有字段提前,避免 match 中 part-move
|
||||||
let block_type = match &content_block {
|
let block_type = match &content_block {
|
||||||
AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text,
|
AnthropicContentBlockStart::Text { .. } => ContentBlockType::Text,
|
||||||
AnthropicContentBlockStart::ToolUse { id, name } => {
|
AnthropicContentBlockStart::ToolUse { id, name } => ContentBlockType::ToolUse {
|
||||||
ContentBlockType::ToolUse {
|
|
||||||
id: id.clone(),
|
id: id.clone(),
|
||||||
name: name.clone(),
|
name: name.clone(),
|
||||||
}
|
},
|
||||||
}
|
|
||||||
AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking,
|
AnthropicContentBlockStart::Thinking { .. } => ContentBlockType::Thinking,
|
||||||
};
|
};
|
||||||
events.push(StreamEvent::ContentBlockStart {
|
events.push(StreamEvent::ContentBlockStart { index, block_type });
|
||||||
index,
|
|
||||||
block_type,
|
|
||||||
});
|
|
||||||
let builder = match content_block {
|
let builder = match content_block {
|
||||||
AnthropicContentBlockStart::Text { text } => {
|
AnthropicContentBlockStart::Text { text } => {
|
||||||
crate::llm::types::response_v2::ContentBlockBuilder::Text(text)
|
crate::llm::types::response_v2::ContentBlockBuilder::Text(text)
|
||||||
@@ -742,7 +812,9 @@ impl AnthropicSseStream {
|
|||||||
completion_tokens_details: None,
|
completion_tokens_details: None,
|
||||||
prompt_tokens_details: None,
|
prompt_tokens_details: None,
|
||||||
};
|
};
|
||||||
events.push(StreamEvent::CostUpdate { usage: partial_usage });
|
events.push(StreamEvent::CostUpdate {
|
||||||
|
usage: partial_usage,
|
||||||
|
});
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
AnthropicSseEvent::MessageStop => {
|
AnthropicSseEvent::MessageStop => {
|
||||||
@@ -753,7 +825,9 @@ impl AnthropicSseStream {
|
|||||||
self.saw_terminal = true;
|
self.saw_terminal = true;
|
||||||
match self.partial.clone().finalize() {
|
match self.partial.clone().finalize() {
|
||||||
Ok(full) => {
|
Ok(full) => {
|
||||||
events.push(StreamEvent::MessageComplete { full_response: full });
|
events.push(StreamEvent::MessageComplete {
|
||||||
|
full_response: full,
|
||||||
|
});
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
events.push(StreamEvent::Error {
|
events.push(StreamEvent::Error {
|
||||||
@@ -780,10 +854,7 @@ fn _unused_marker() {}
|
|||||||
impl Stream for AnthropicSseStream {
|
impl Stream for AnthropicSseStream {
|
||||||
type Item = Result<StreamEvent, LlmError>;
|
type Item = Result<StreamEvent, LlmError>;
|
||||||
|
|
||||||
fn poll_next(
|
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||||
mut self: Pin<&mut Self>,
|
|
||||||
cx: &mut Context<'_>,
|
|
||||||
) -> Poll<Option<Self::Item>> {
|
|
||||||
loop {
|
loop {
|
||||||
if let Some(data) = self.next_event_line() {
|
if let Some(data) = self.next_event_line() {
|
||||||
let mut events = self.handle_event_json(&data);
|
let mut events = self.handle_event_json(&data);
|
||||||
@@ -830,12 +901,17 @@ mod tests {
|
|||||||
use super::*;
|
use super::*;
|
||||||
use crate::llm::types::request_v2::MessageRequest;
|
use crate::llm::types::request_v2::MessageRequest;
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
use wiremock::matchers::{method, path};
|
use wiremock::matchers::{header, method, path};
|
||||||
use wiremock::{Mock, MockServer, ResponseTemplate};
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
|
|
||||||
fn make_provider(base_url: String) -> AnthropicProvider {
|
fn make_provider(base_url: String) -> AnthropicProvider {
|
||||||
// 跳过默认 header 注入:测试用自定义 base_url 直接 mock
|
// 跳过默认 header 注入:测试用自定义 base_url 直接 mock
|
||||||
AnthropicProvider::new(base_url, "sk-ant-test".into(), "claude-sonnet-4-20250514".into())
|
AnthropicProvider::new(
|
||||||
|
base_url,
|
||||||
|
"sk-ant-test".into(),
|
||||||
|
"claude-sonnet-4-20250514".into(),
|
||||||
|
30,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
@@ -979,6 +1055,7 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
|
|||||||
"http://x".into(),
|
"http://x".into(),
|
||||||
"k".into(),
|
"k".into(),
|
||||||
"claude-sonnet-4-20250514".into(),
|
"claude-sonnet-4-20250514".into(),
|
||||||
|
30,
|
||||||
)
|
)
|
||||||
.capabilities();
|
.capabilities();
|
||||||
assert_eq!(caps.provider_name, "anthropic");
|
assert_eq!(caps.provider_name, "anthropic");
|
||||||
@@ -1000,4 +1077,134 @@ event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
|
|||||||
assert_eq!(body.max_tokens, DEFAULT_MAX_TOKENS);
|
assert_eq!(body.max_tokens, DEFAULT_MAX_TOKENS);
|
||||||
assert_eq!(body.model, "claude-sonnet-4-20250514");
|
assert_eq!(body.model, "claude-sonnet-4-20250514");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ===== Phase 11 Step 11.2 wiremock roundtrip 测试 =====
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn anthropic_401_structured_error() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/v1/messages"))
|
||||||
|
.respond_with(ResponseTemplate::new(401).set_body_json(json!({
|
||||||
|
"type": "error",
|
||||||
|
"error": {
|
||||||
|
"type": "authentication_error",
|
||||||
|
"message": "Invalid API key provided: sk-ant-test"
|
||||||
|
}
|
||||||
|
})))
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = make_provider(server.uri());
|
||||||
|
let err = provider
|
||||||
|
.chat_blocking(MessageRequest {
|
||||||
|
model: "claude-sonnet-4-20250514".into(),
|
||||||
|
messages: vec![Message::user_text("Hi")],
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap_err();
|
||||||
|
match err {
|
||||||
|
LlmError::Authentication(msg) => assert!(msg.contains("Invalid API key")),
|
||||||
|
other => panic!("expected Authentication, got {other:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn anthropic_tool_use_response() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/v1/messages"))
|
||||||
|
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||||
|
"id": "msg_tool",
|
||||||
|
"type": "message",
|
||||||
|
"model": "claude-sonnet-4-20250514",
|
||||||
|
"content": [
|
||||||
|
{"type": "text", "text": "Let me check."},
|
||||||
|
{"type": "tool_use", "id": "toolu_abc", "name": "lookup", "input": {"q": "rust"}}
|
||||||
|
],
|
||||||
|
"stop_reason": "tool_use",
|
||||||
|
"usage": {"input_tokens": 8, "output_tokens": 12}
|
||||||
|
})))
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = make_provider(server.uri());
|
||||||
|
let response = provider
|
||||||
|
.chat_blocking(MessageRequest {
|
||||||
|
model: "claude-sonnet-4-20250514".into(),
|
||||||
|
messages: vec![Message::user_text("Look up rust")],
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(response.stop_reason, StopReason::ToolUse);
|
||||||
|
let tool_use = match &response.message {
|
||||||
|
Message::Assistant { content } => content.iter().find_map(|b| match b {
|
||||||
|
ContentBlock::ToolUse { id, name, .. } => Some((id.clone(), name.clone())),
|
||||||
|
_ => None,
|
||||||
|
}),
|
||||||
|
_ => None,
|
||||||
|
};
|
||||||
|
assert_eq!(tool_use, Some(("toolu_abc".into(), "lookup".into())));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn anthropic_version_header() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/v1/messages"))
|
||||||
|
.and(header("anthropic-version", "2023-06-01"))
|
||||||
|
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||||
|
"id": "msg_v",
|
||||||
|
"type": "message",
|
||||||
|
"model": "claude-sonnet-4-20250514",
|
||||||
|
"content": [{"type": "text", "text": "OK"}],
|
||||||
|
"stop_reason": "end_turn",
|
||||||
|
"usage": {"input_tokens": 1, "output_tokens": 1}
|
||||||
|
})))
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = make_provider(server.uri());
|
||||||
|
let response = provider
|
||||||
|
.chat_blocking(MessageRequest {
|
||||||
|
model: "claude-sonnet-4-20250514".into(),
|
||||||
|
messages: vec![Message::user_text("Hi")],
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(response.text(), "OK");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn anthropic_529_overloaded_structured() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/v1/messages"))
|
||||||
|
.respond_with(ResponseTemplate::new(529).set_body_json(json!({
|
||||||
|
"type": "error",
|
||||||
|
"error": {
|
||||||
|
"type": "overloaded_error",
|
||||||
|
"message": "Overloaded: Anthropic API is temporarily overloaded"
|
||||||
|
}
|
||||||
|
})))
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = make_provider(server.uri());
|
||||||
|
let err = provider
|
||||||
|
.chat_blocking(MessageRequest {
|
||||||
|
model: "claude-sonnet-4-20250514".into(),
|
||||||
|
messages: vec![Message::user_text("Hi")],
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap_err();
|
||||||
|
match err {
|
||||||
|
LlmError::RateLimit { retry_after } => assert!(retry_after.is_none()),
|
||||||
|
other => panic!("expected RateLimit, got {other:?}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
+771
-70
File diff suppressed because it is too large
Load Diff
@@ -15,12 +15,12 @@ use std::pin::Pin;
|
|||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use futures_core::Stream;
|
use futures_core::Stream;
|
||||||
|
|
||||||
use super::openai::GenericOpenaiProvider;
|
|
||||||
use super::ProviderCapabilities;
|
use super::ProviderCapabilities;
|
||||||
|
use super::openai::GenericOpenaiProvider;
|
||||||
use crate::llm::error::LlmError;
|
use crate::llm::error::LlmError;
|
||||||
|
use crate::llm::provider::LlmProvider;
|
||||||
use crate::llm::types::request_v2::MessageRequest;
|
use crate::llm::types::request_v2::MessageRequest;
|
||||||
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
use crate::llm::types::response_v2::{MessageResponse, StreamEvent};
|
||||||
use crate::llm::provider::LlmProvider;
|
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
// DeepSeek
|
// DeepSeek
|
||||||
@@ -29,7 +29,7 @@ use crate::llm::provider::LlmProvider;
|
|||||||
pub struct DeepSeekProvider(pub GenericOpenaiProvider);
|
pub struct DeepSeekProvider(pub GenericOpenaiProvider);
|
||||||
|
|
||||||
impl DeepSeekProvider {
|
impl DeepSeekProvider {
|
||||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||||
let url = if base_url.is_empty() {
|
let url = if base_url.is_empty() {
|
||||||
"https://api.deepseek.com".to_string()
|
"https://api.deepseek.com".to_string()
|
||||||
} else {
|
} else {
|
||||||
@@ -40,26 +40,27 @@ impl DeepSeekProvider {
|
|||||||
api_key,
|
api_key,
|
||||||
model,
|
model,
|
||||||
"deepseek",
|
"deepseek",
|
||||||
|
timeout_secs,
|
||||||
))
|
))
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
impl DeepSeekProvider {
|
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
|
||||||
|
pub fn with_client(self, client: reqwest::Client) -> Self {
|
||||||
|
Self(self.0.with_client(client))
|
||||||
|
}
|
||||||
|
|
||||||
/// 测试中(带 mock_client)使用的构造器。
|
/// 测试中(带 mock_client)使用的构造器。
|
||||||
|
///
|
||||||
|
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
|
||||||
|
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
|
||||||
|
/// 无超时),不影响行为。
|
||||||
pub fn new_with_client(
|
pub fn new_with_client(
|
||||||
base_url: String,
|
base_url: String,
|
||||||
api_key: String,
|
api_key: String,
|
||||||
model: String,
|
model: String,
|
||||||
client: reqwest::Client,
|
client: reqwest::Client,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
let url = if base_url.is_empty() {
|
Self::new(base_url, api_key, model, 30).with_client(client)
|
||||||
"https://api.deepseek.com".to_string()
|
|
||||||
} else {
|
|
||||||
base_url
|
|
||||||
};
|
|
||||||
let mut inner = GenericOpenaiProvider::new_with_name(url, api_key, model, "deepseek");
|
|
||||||
inner.http_client = client;
|
|
||||||
Self(inner)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -72,8 +73,7 @@ impl LlmProvider for DeepSeekProvider {
|
|||||||
async fn chat_stream(
|
async fn chat_stream(
|
||||||
&self,
|
&self,
|
||||||
request: MessageRequest,
|
request: MessageRequest,
|
||||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||||
{
|
|
||||||
self.0.chat_stream(request).await
|
self.0.chat_stream(request).await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -91,7 +91,7 @@ impl LlmProvider for DeepSeekProvider {
|
|||||||
pub struct QwenProvider(pub GenericOpenaiProvider);
|
pub struct QwenProvider(pub GenericOpenaiProvider);
|
||||||
|
|
||||||
impl QwenProvider {
|
impl QwenProvider {
|
||||||
pub fn new(base_url: String, api_key: String, model: String) -> Self {
|
pub fn new(base_url: String, api_key: String, model: String, timeout_secs: u64) -> Self {
|
||||||
let url = if base_url.is_empty() {
|
let url = if base_url.is_empty() {
|
||||||
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
|
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
|
||||||
} else {
|
} else {
|
||||||
@@ -104,31 +104,28 @@ impl QwenProvider {
|
|||||||
model,
|
model,
|
||||||
"qwen",
|
"qwen",
|
||||||
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
||||||
|
timeout_secs,
|
||||||
);
|
);
|
||||||
Self(inner)
|
Self(inner)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 替换默认 HTTP Client(用于 timeout 注入等场景)。
|
||||||
|
pub fn with_client(self, client: reqwest::Client) -> Self {
|
||||||
|
Self(self.0.with_client(client))
|
||||||
|
}
|
||||||
|
|
||||||
/// 测试构造器。
|
/// 测试构造器。
|
||||||
|
///
|
||||||
|
/// ponytail: 此处 `30` 是 `timeout_secs` 字段的占位值,仅用于 `map_reqwest_error`
|
||||||
|
/// 错误消息中的回显。实际请求超时由传入的 `client` 控制(通常测试用的 mock client
|
||||||
|
/// 无超时),不影响行为。
|
||||||
pub fn new_with_client(
|
pub fn new_with_client(
|
||||||
base_url: String,
|
base_url: String,
|
||||||
api_key: String,
|
api_key: String,
|
||||||
model: String,
|
model: String,
|
||||||
client: reqwest::Client,
|
client: reqwest::Client,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
let url = if base_url.is_empty() {
|
Self::new(base_url, api_key, model, 30).with_client(client)
|
||||||
"https://dashscope.aliyuncs.com/compatible-mode/v1".to_string()
|
|
||||||
} else {
|
|
||||||
base_url
|
|
||||||
};
|
|
||||||
let mut inner = GenericOpenaiProvider::new_with_name_and_headers(
|
|
||||||
url,
|
|
||||||
api_key,
|
|
||||||
model,
|
|
||||||
"qwen",
|
|
||||||
vec![("X-DashScope-SSE".to_string(), "enable".to_string())],
|
|
||||||
);
|
|
||||||
inner.http_client = client;
|
|
||||||
Self(inner)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -141,8 +138,7 @@ impl LlmProvider for QwenProvider {
|
|||||||
async fn chat_stream(
|
async fn chat_stream(
|
||||||
&self,
|
&self,
|
||||||
request: MessageRequest,
|
request: MessageRequest,
|
||||||
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError>
|
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, LlmError>> + Send>>, LlmError> {
|
||||||
{
|
|
||||||
self.0.chat_stream(request).await
|
self.0.chat_stream(request).await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -156,8 +152,8 @@ impl LlmProvider for QwenProvider {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::llm::types::request_v2::MessageRequest;
|
|
||||||
use crate::llm::types::message::Message as IrMessage;
|
use crate::llm::types::message::Message as IrMessage;
|
||||||
|
use crate::llm::types::request_v2::MessageRequest;
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
use wiremock::matchers::{method, path};
|
use wiremock::matchers::{method, path};
|
||||||
use wiremock::{Mock, MockServer, ResponseTemplate};
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
@@ -182,11 +178,8 @@ mod tests {
|
|||||||
.mount(&server)
|
.mount(&server)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
let provider = DeepSeekProvider::new(
|
let provider =
|
||||||
server.uri(),
|
DeepSeekProvider::new(server.uri(), "sk-test".into(), "deepseek-chat".into(), 30);
|
||||||
"sk-test".into(),
|
|
||||||
"deepseek-chat".into(),
|
|
||||||
);
|
|
||||||
let response = provider
|
let response = provider
|
||||||
.chat(MessageRequest {
|
.chat(MessageRequest {
|
||||||
model: "deepseek-chat".into(),
|
model: "deepseek-chat".into(),
|
||||||
@@ -219,7 +212,7 @@ mod tests {
|
|||||||
.mount(&server)
|
.mount(&server)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into());
|
let provider = QwenProvider::new(server.uri(), "sk-test".into(), "qwen-plus".into(), 30);
|
||||||
let response = provider
|
let response = provider
|
||||||
.chat(MessageRequest {
|
.chat(MessageRequest {
|
||||||
model: "qwen-plus".into(),
|
model: "qwen-plus".into(),
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
|
|
||||||
use crate::llm::error::LlmError;
|
use crate::llm::error::LlmError;
|
||||||
use crate::llm::provider::{create_provider, LlmProvider, ProviderConfig, ProviderType};
|
use crate::llm::provider::{LlmProvider, ProviderConfig, ProviderType, create_provider};
|
||||||
|
|
||||||
/// Provider 注册表 —— 管理多个 LLM Provider 实例。
|
/// Provider 注册表 —— 管理多个 LLM Provider 实例。
|
||||||
///
|
///
|
||||||
@@ -61,8 +61,6 @@ impl ProviderRegistry {
|
|||||||
|
|
||||||
/// 获取默认 Provider。
|
/// 获取默认 Provider。
|
||||||
pub fn get_default(&self) -> Option<&dyn LlmProvider> {
|
pub fn get_default(&self) -> Option<&dyn LlmProvider> {
|
||||||
self.default_name
|
self.default_name.as_ref().and_then(|name| self.get(name))
|
||||||
.as_ref()
|
|
||||||
.and_then(|name| self.get(name))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+5
-200
@@ -1,204 +1,9 @@
|
|||||||
//! 流式事件系统 —— 将 LLM 流式响应解析为语义化事件。
|
//! 流式事件系统 —— 重导出 `StreamEvent` 供向后兼容。
|
||||||
//!
|
//!
|
||||||
//! Phase 0 修订(参见 `docs/10a-phase0-types-and-trait.md` §"StreamEvent 命名冲突处理"):
|
//! 历史说明(Phase 0 → Phase 13):
|
||||||
//! - 对外暴露的 `StreamEvent` 是高精度 IR 版本(来自 `response_v2::StreamEvent`)。
|
//! - 对外暴露的 `StreamEvent` 是高精度 IR 版本(来自 `response_v2::StreamEvent`)。
|
||||||
//! - 旧变体(`AssistantTextDelta` / `ToolExecutionStarted` 等)重命名为 `LegacyStreamEvent`
|
//! - 旧版 chunk 解析 + LegacyStreamEvent 适配层在 Phase 13 完成后已整体删除。
|
||||||
//! 放在 `crate::llm::types::old_stream` 模块,本文件内部消费。
|
//! - 当前文件仅保留 `pub use` 重导出,保持与既有
|
||||||
//! - Phase 1 重写 Provider 时可直接消费新事件流后整体删除 `LegacyStreamEvent` 相关代码。
|
//! `use crate::llm::stream::StreamEvent` 的代码兼容。
|
||||||
//!
|
|
||||||
//! 当前实现:旧的 `parse_chunk_stream` 内部消费 `OpenaiChatChunk`,映射为
|
|
||||||
//! `LegacyStreamEvent`,再在 `LegacyToIrEventStream` 中映射为新 IR `StreamEvent`
|
|
||||||
//! 后输出。Phase 1 会重写此层(OpenAI Provider 直接产出新事件流)。
|
|
||||||
|
|
||||||
use std::pin::Pin;
|
|
||||||
use std::task::{Context, Poll};
|
|
||||||
|
|
||||||
use futures_core::stream::Stream;
|
|
||||||
use futures_util::future::poll_fn;
|
|
||||||
use futures_util::FutureExt;
|
|
||||||
use serde_json::Value;
|
|
||||||
|
|
||||||
use crate::llm::error::LlmError;
|
|
||||||
use crate::llm::types::old_stream::LegacyStreamEvent;
|
|
||||||
use crate::llm::types::response_v2::MessageResponse;
|
|
||||||
use crate::llm::types::response_v2::StopReason;
|
|
||||||
use crate::llm::types::usage::Usage;
|
|
||||||
use crate::llm::types::{OpenaiChatChunk, OpenaiToolCall};
|
|
||||||
|
|
||||||
// 唯一的对外 `StreamEvent` 定义(高精度 IR 事件,来自 `response_v2`)。
|
|
||||||
//
|
|
||||||
// 此 `pub use` 同时起到两个作用:
|
|
||||||
// 1. 让 `crate::llm::stream::StreamEvent` 路径仍指向新高精度 IR 事件,
|
|
||||||
// 保持与既有 `use crate::llm::stream::StreamEvent` 的代码兼容;
|
|
||||||
// 2. 把模块内部的 `StreamEvent` 名字指向 `response_v2::StreamEvent`。
|
|
||||||
pub use crate::llm::types::response_v2::StreamEvent;
|
pub use crate::llm::types::response_v2::StreamEvent;
|
||||||
|
|
||||||
/// 将原始 OpenaiChatChunk 流解析为新高精度 IR StreamEvent 流。
|
|
||||||
///
|
|
||||||
/// ponytail: 每个产出事件都用 `Result<_, LlmError>` 包装,让上层 `chat_stream`
|
|
||||||
/// trait 方法直接消费并保持错误传播链。当前 `LegacyToIrEventStream` 内部
|
|
||||||
/// 不会产生错误,所有结果都是 `Ok`;后续 Phase 1 重写 Provider 时,
|
|
||||||
/// 真实 IR 流转换可在此层注入 error 事件。
|
|
||||||
pub fn parse_chunk_stream(
|
|
||||||
chunks: Pin<Box<dyn futures_core::Stream<Item = Result<OpenaiChatChunk, LlmError>> + Send>>,
|
|
||||||
) -> Pin<Box<dyn futures_core::Stream<Item = Result<StreamEvent, LlmError>> + Send>> {
|
|
||||||
let legacy = parse_chunk_stream_legacy(chunks);
|
|
||||||
Box::pin(LegacyToIrEventStream { inner: legacy })
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- 内部:chunk → LegacyStreamEvent ---
|
|
||||||
|
|
||||||
fn parse_chunk_stream_legacy(
|
|
||||||
chunks: Pin<Box<dyn futures_core::Stream<Item = Result<OpenaiChatChunk, LlmError>> + Send>>,
|
|
||||||
) -> Pin<Box<dyn futures_core::Stream<Item = LegacyStreamEvent> + Send>> {
|
|
||||||
Box::pin(ChunkToLegacyEventStream { chunks })
|
|
||||||
}
|
|
||||||
|
|
||||||
struct ChunkToLegacyEventStream {
|
|
||||||
chunks: Pin<Box<dyn futures_core::Stream<Item = Result<OpenaiChatChunk, LlmError>> + Send>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Stream for ChunkToLegacyEventStream {
|
|
||||||
type Item = LegacyStreamEvent;
|
|
||||||
|
|
||||||
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
|
||||||
let this = &mut *self;
|
|
||||||
poll_fn(|cx| match Pin::new(&mut this.chunks).poll_next(cx) {
|
|
||||||
Poll::Ready(Some(Ok(chunk))) => {
|
|
||||||
for choice in &chunk.choices {
|
|
||||||
let delta = &choice.delta;
|
|
||||||
|
|
||||||
if let Some(content) = &delta.content {
|
|
||||||
return Poll::Ready(Some(LegacyStreamEvent::AssistantTextDelta {
|
|
||||||
text: content.clone(),
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(tool_calls) = &delta.tool_calls
|
|
||||||
&& let Some(tc) = tool_calls.first()
|
|
||||||
{
|
|
||||||
let OpenaiToolCall::Function { id, function } = tc;
|
|
||||||
let args: Value =
|
|
||||||
serde_json::from_str(&function.arguments).unwrap_or(Value::Null);
|
|
||||||
return Poll::Ready(Some(LegacyStreamEvent::ToolExecutionStarted {
|
|
||||||
tool_name: function.name.clone(),
|
|
||||||
input: args,
|
|
||||||
tool_call_id: id.clone(),
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(finish_reason) = &choice.finish_reason {
|
|
||||||
return Poll::Ready(Some(LegacyStreamEvent::TurnComplete {
|
|
||||||
reason: *finish_reason,
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(usage) = &chunk.usage {
|
|
||||||
return Poll::Ready(Some(LegacyStreamEvent::CostUpdate {
|
|
||||||
usage: *usage,
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
|
|
||||||
Poll::Ready(None)
|
|
||||||
}
|
|
||||||
Poll::Ready(Some(Err(e))) => Poll::Ready(Some(LegacyStreamEvent::error(e.to_string()))),
|
|
||||||
Poll::Ready(None) => Poll::Ready(None),
|
|
||||||
Poll::Pending => Poll::Pending,
|
|
||||||
})
|
|
||||||
.poll_unpin(cx)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- 内部:LegacyStreamEvent → 新 StreamEvent ---
|
|
||||||
|
|
||||||
struct LegacyToIrEventStream {
|
|
||||||
inner: Pin<Box<dyn futures_core::Stream<Item = LegacyStreamEvent> + Send>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Stream for LegacyToIrEventStream {
|
|
||||||
type Item = Result<StreamEvent, LlmError>;
|
|
||||||
|
|
||||||
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
|
||||||
let this = &mut *self;
|
|
||||||
match Pin::new(&mut this.inner).poll_next(cx) {
|
|
||||||
Poll::Ready(Some(legacy)) => Poll::Ready(Some(Ok(map_legacy_to_ir(legacy)))),
|
|
||||||
Poll::Ready(None) => {
|
|
||||||
// 旧流结束 → 主动补一个 MessageComplete(full_response 为兜底空快照)。
|
|
||||||
// ponytail: Phase 0 中 OpenaiProvider 桥接层负责产出真实 MessageResponse,
|
|
||||||
// 此处仅防止消费方无限等待。若 Provider 层已正确发出 MessageComplete,
|
|
||||||
// LlmCycle 不会走到这里 —— 因为桥接层 inline 处理。
|
|
||||||
Poll::Ready(Some(Ok(StreamEvent::MessageComplete {
|
|
||||||
full_response: empty_message_response(),
|
|
||||||
})))
|
|
||||||
}
|
|
||||||
Poll::Pending => Poll::Pending,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn empty_message_response() -> MessageResponse {
|
|
||||||
use crate::llm::types::message::Message;
|
|
||||||
use std::collections::HashMap;
|
|
||||||
MessageResponse {
|
|
||||||
id: String::new(),
|
|
||||||
model: String::new(),
|
|
||||||
message: Message::Assistant {
|
|
||||||
content: vec![],
|
|
||||||
},
|
|
||||||
usage: Usage::default(),
|
|
||||||
stop_reason: StopReason::Stop,
|
|
||||||
extra: HashMap::new(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 把旧 LegacyStreamEvent 映射到新高精度 IR StreamEvent。
|
|
||||||
///
|
|
||||||
/// Phase 1 重写 Provider 后可直接删除此映射函数。当前映射语义:
|
|
||||||
/// - `AssistantTextDelta` → `TextDelta`
|
|
||||||
/// - `ToolExecutionStarted` → `ToolCallArgumentsDelta`(OpenAI 单 chunk 模式下整段 arguments 一次性下发)
|
|
||||||
/// - `CostUpdate` → `CostUpdate`(Usage → PartialUsage 全字段)
|
|
||||||
/// - `TurnComplete` → `MessageComplete`(Phase 1 重写 Provider 后正确产出)
|
|
||||||
/// - `Error` → `Error`
|
|
||||||
///
|
|
||||||
/// ponytail: 这是一个"目前能跑通未来会被删除"的适配层。当前实现为单事件映射,
|
|
||||||
/// 旧 `ToolExecutionStarted` 携带的 (id, name) 暂未填入 IR 事件(消费方
|
|
||||||
/// Phase 2 中通过 MessageComplete.full_response.tool_use 提取)。Phase 1 重写时
|
|
||||||
/// 由 OpenAI Provider 直接产出 IR 流,整体删除此映射。
|
|
||||||
fn map_legacy_to_ir(legacy: LegacyStreamEvent) -> StreamEvent {
|
|
||||||
use crate::llm::types::response_v2::PartialUsage;
|
|
||||||
|
|
||||||
match legacy {
|
|
||||||
LegacyStreamEvent::AssistantTextDelta { text } => StreamEvent::TextDelta { text },
|
|
||||||
LegacyStreamEvent::ToolExecutionStarted { input, .. } => {
|
|
||||||
let arguments = serde_json::to_string(&input).unwrap_or_default();
|
|
||||||
StreamEvent::ToolCallArgumentsDelta { index: 0, arguments }
|
|
||||||
}
|
|
||||||
LegacyStreamEvent::ToolExecutionCompleted { .. } => {
|
|
||||||
// 旧 ToolExecutionCompleted 不在 IR 流协议中——工具执行是消费方职责。
|
|
||||||
// Phase 1 重写时此处整体删除。当前给一个无副作用的占位事件。
|
|
||||||
StreamEvent::CostUpdate {
|
|
||||||
usage: PartialUsage::default(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
LegacyStreamEvent::CostUpdate { usage } => StreamEvent::CostUpdate {
|
|
||||||
usage: PartialUsage {
|
|
||||||
prompt_tokens: Some(usage.prompt_tokens),
|
|
||||||
completion_tokens: Some(usage.completion_tokens),
|
|
||||||
total_tokens: Some(usage.total_tokens),
|
|
||||||
completion_tokens_details: usage.completion_tokens_details,
|
|
||||||
prompt_tokens_details: usage.prompt_tokens_details,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
LegacyStreamEvent::TurnComplete { reason } => {
|
|
||||||
// 旧 TurnComplete 不直接对应 IR;映射为带 StopReason 的 MessageComplete。
|
|
||||||
// ponytail: Phase 1 重写 Provider 后此适配整体删除,
|
|
||||||
// OpenAI Provider 直接产出带正确 stop_reason 的 MessageComplete。
|
|
||||||
let _ = reason;
|
|
||||||
StreamEvent::MessageComplete {
|
|
||||||
full_response: empty_message_response(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
LegacyStreamEvent::Error { message } => StreamEvent::Error { message },
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -20,15 +20,12 @@ use crate::llm::types::shared::ImageDetail;
|
|||||||
/// 消费方 match 可直接区分文本和图片输入。
|
/// 消费方 match 可直接区分文本和图片输入。
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum Message {
|
pub enum Message {
|
||||||
/// 系统提示(User & Assistant 之外的引导指令)。
|
/// 系统提示(User & Assistant 之外的引导指令)。
|
||||||
System {
|
System { content: Vec<ContentBlock> },
|
||||||
content: Vec<ContentBlock>,
|
|
||||||
},
|
|
||||||
/// 用户输入。
|
/// 用户输入。
|
||||||
User {
|
User { content: Vec<ContentBlock> },
|
||||||
content: Vec<ContentBlock>,
|
|
||||||
},
|
|
||||||
/// 用户的图片输入(快捷构造,免去构造 ContentBlock 的 boilerplate)。
|
/// 用户的图片输入(快捷构造,免去构造 ContentBlock 的 boilerplate)。
|
||||||
UserImage {
|
UserImage {
|
||||||
data: String,
|
data: String,
|
||||||
@@ -36,9 +33,7 @@ pub enum Message {
|
|||||||
detail: ImageDetail,
|
detail: ImageDetail,
|
||||||
},
|
},
|
||||||
/// Assistant 回复内容块(可能包含 text、thinking、tool_use 等多种 block 的混合)。
|
/// Assistant 回复内容块(可能包含 text、thinking、tool_use 等多种 block 的混合)。
|
||||||
Assistant {
|
Assistant { content: Vec<ContentBlock> },
|
||||||
content: Vec<ContentBlock>,
|
|
||||||
},
|
|
||||||
/// 工具调用结果。
|
/// 工具调用结果。
|
||||||
ToolResult {
|
ToolResult {
|
||||||
tool_call_id: String,
|
tool_call_id: String,
|
||||||
@@ -103,6 +98,7 @@ impl Message {
|
|||||||
/// block 的逃生舱。
|
/// block 的逃生舱。
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum ContentBlock {
|
pub enum ContentBlock {
|
||||||
/// 纯文本。
|
/// 纯文本。
|
||||||
Text { text: String },
|
Text { text: String },
|
||||||
@@ -130,10 +126,7 @@ pub enum ContentBlock {
|
|||||||
signature: Option<String>,
|
signature: Option<String>,
|
||||||
},
|
},
|
||||||
/// 逃生舱:Provider 特定 block 透传(OpenAI Response 内置工具等)。
|
/// 逃生舱:Provider 特定 block 透传(OpenAI Response 内置工具等)。
|
||||||
Extension {
|
Extension { kind: String, data: Value },
|
||||||
kind: String,
|
|
||||||
data: Value,
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 内容块类型标签 —— 用于 `StreamEvent::ContentBlockStart.block_type`。
|
/// 内容块类型标签 —— 用于 `StreamEvent::ContentBlockStart.block_type`。
|
||||||
@@ -141,6 +134,7 @@ pub enum ContentBlock {
|
|||||||
/// 用途:在流式场景中,Provider 先下发 block 类型,再下发 block 内容增量。
|
/// 用途:在流式场景中,Provider 先下发 block 类型,再下发 block 内容增量。
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum ContentBlockType {
|
pub enum ContentBlockType {
|
||||||
/// 文本块。
|
/// 文本块。
|
||||||
Text,
|
Text,
|
||||||
@@ -349,9 +343,7 @@ mod tests {
|
|||||||
fn message_roundtrip_each_variant() {
|
fn message_roundtrip_each_variant() {
|
||||||
let msgs = vec![
|
let msgs = vec![
|
||||||
Message::System {
|
Message::System {
|
||||||
content: vec![ContentBlock::Text {
|
content: vec![ContentBlock::Text { text: "sys".into() }],
|
||||||
text: "sys".into(),
|
|
||||||
}],
|
|
||||||
},
|
},
|
||||||
Message::User {
|
Message::User {
|
||||||
content: vec![ContentBlock::Text {
|
content: vec![ContentBlock::Text {
|
||||||
@@ -376,9 +368,7 @@ mod tests {
|
|||||||
},
|
},
|
||||||
Message::ToolResult {
|
Message::ToolResult {
|
||||||
tool_call_id: "call_1".into(),
|
tool_call_id: "call_1".into(),
|
||||||
content: vec![ContentBlock::Text {
|
content: vec![ContentBlock::Text { text: "ok".into() }],
|
||||||
text: "ok".into(),
|
|
||||||
}],
|
|
||||||
is_error: true,
|
is_error: true,
|
||||||
},
|
},
|
||||||
];
|
];
|
||||||
|
|||||||
+1
-81
@@ -1,9 +1,6 @@
|
|||||||
pub mod message;
|
pub mod message;
|
||||||
pub mod old_stream;
|
|
||||||
pub mod openai_message;
|
pub mod openai_message;
|
||||||
pub mod request;
|
|
||||||
pub mod request_v2;
|
pub mod request_v2;
|
||||||
pub mod response;
|
|
||||||
pub mod response_v2;
|
pub mod response_v2;
|
||||||
pub mod shared;
|
pub mod shared;
|
||||||
pub mod tool;
|
pub mod tool;
|
||||||
@@ -12,12 +9,7 @@ pub mod usage;
|
|||||||
pub use openai_message::{
|
pub use openai_message::{
|
||||||
ContentField, FileData, ImageURL, InputAudio, OpenaiChatMessage, OpenaiContentPart,
|
ContentField, FileData, ImageURL, InputAudio, OpenaiChatMessage, OpenaiContentPart,
|
||||||
};
|
};
|
||||||
pub use request::{OpenaiChatRequest, OpenaiTool, StreamOptions, ToolChoice};
|
|
||||||
pub use request_v2::{ExtraError, MessageRequest, ThinkingConfig};
|
pub use request_v2::{ExtraError, MessageRequest, ThinkingConfig};
|
||||||
pub use response::{
|
|
||||||
Annotation, Choice, ChunkChoice, Delta, Logprobs, OpenaiAudio, OpenaiChatChunk,
|
|
||||||
OpenaiChatResponse, TokenLogprob, TopLogprob, URLCitation,
|
|
||||||
};
|
|
||||||
pub use response_v2::{
|
pub use response_v2::{
|
||||||
ContentBlockBuilder, MessageResponse, PartialMessageResponse, PartialUsage, StopReason,
|
ContentBlockBuilder, MessageResponse, PartialMessageResponse, PartialUsage, StopReason,
|
||||||
StreamEvent,
|
StreamEvent,
|
||||||
@@ -26,77 +18,5 @@ pub use shared::{
|
|||||||
AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier,
|
AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier,
|
||||||
StopSequence,
|
StopSequence,
|
||||||
};
|
};
|
||||||
pub use tool::{FunctionCall, OpenaiToolCall, OpenaiToolDefinition};
|
pub use tool::{FunctionCall, OpenaiToolCall, ToolChoice, ToolDef};
|
||||||
pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage};
|
pub use usage::{CompletionTokensDetails, CostTracker, PromptTokensDetails, Usage};
|
||||||
|
|
||||||
// Re-export IR 内容块 / 消息类型供 `types::ContentBlock` 等历史路径消费。
|
|
||||||
//
|
|
||||||
// 注意:以下别名 *故意不暴露* `pub type Message = message::Message`、
|
|
||||||
// `pub type ContentBlock = message::ContentBlock` —— 新 `Message` / `ContentBlock` /
|
|
||||||
// `StopReason` 是独立类型,由 `Message` / `ContentBlock` / `StopReason` 直接路径访问,
|
|
||||||
// 旧别名(指 `OpenaiChatMessage` / `OpenaiContentPart` / `FinishReason`)已移除,
|
|
||||||
// 避免新类型阴影。Phase 2 完成后再统一收敛。
|
|
||||||
//
|
|
||||||
// Phase 1 起移除 `ChatRequest` 别名 —— 新代码统一使用 `MessageRequest`(v2 IR)。
|
|
||||||
// `ChatResponse` 结构体仍存在,作为 OpenAI `chat_inner()` 内部 wire-format 转换目标。
|
|
||||||
/// 旧 wire-format 响应结构(保留用于 OpenAI 内部转换层)。
|
|
||||||
#[deprecated(since = "0.1.0", note = "请改用 MessageResponse")]
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub struct ChatResponse {
|
|
||||||
pub message: OpenaiChatMessage,
|
|
||||||
pub usage: Usage,
|
|
||||||
pub stop_reason: Option<FinishReason>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[allow(deprecated)]
|
|
||||||
impl From<OpenaiChatResponse> for ChatResponse {
|
|
||||||
fn from(response: OpenaiChatResponse) -> Self {
|
|
||||||
let message = response
|
|
||||||
.choices
|
|
||||||
.first()
|
|
||||||
.map(|c| c.message.clone())
|
|
||||||
.unwrap_or_else(|| OpenaiChatMessage::assistant_text(""));
|
|
||||||
let stop_reason = response.choices.first().and_then(|c| c.finish_reason);
|
|
||||||
ChatResponse {
|
|
||||||
message,
|
|
||||||
usage: response.usage,
|
|
||||||
stop_reason,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[allow(deprecated)]
|
|
||||||
impl From<ChatResponse> for OpenaiChatChunk {
|
|
||||||
fn from(response: ChatResponse) -> Self {
|
|
||||||
let delta = Delta::from(response.message.clone());
|
|
||||||
let chunk_choice = ChunkChoice {
|
|
||||||
index: 0,
|
|
||||||
delta,
|
|
||||||
logprobs: None,
|
|
||||||
finish_reason: response.stop_reason,
|
|
||||||
};
|
|
||||||
|
|
||||||
OpenaiChatChunk {
|
|
||||||
id: format!(
|
|
||||||
"chunk-{}",
|
|
||||||
std::time::SystemTime::now()
|
|
||||||
.duration_since(std::time::UNIX_EPOCH)
|
|
||||||
.map(|d| d.as_nanos())
|
|
||||||
.unwrap_or(0)
|
|
||||||
),
|
|
||||||
object: "chat.completion.chunk".to_string(),
|
|
||||||
created: std::time::SystemTime::now()
|
|
||||||
.duration_since(std::time::UNIX_EPOCH)
|
|
||||||
.map(|d| d.as_secs())
|
|
||||||
.unwrap_or(0),
|
|
||||||
model: String::new(),
|
|
||||||
choices: vec![chunk_choice],
|
|
||||||
usage: Some(response.usage),
|
|
||||||
system_fingerprint: None,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 工具定义别名(无新类型冲突,保留)。
|
|
||||||
#[deprecated(since = "0.1.0", note = "ToolDefinition 仍直接对应 OpenAI wire-format;未来 v0.2 引入 IR 工具类型后会再次更新")]
|
|
||||||
pub type ToolDefinition = OpenaiToolDefinition;
|
|
||||||
|
|||||||
@@ -1,45 +0,0 @@
|
|||||||
//! 旧版流式事件 —— Phase 0 临时保留,仅供 `stream.rs` 中 `parse_chunk_stream` 内部使用。
|
|
||||||
//!
|
|
||||||
//! Phase 0 中:高精度 `StreamEvent`(定义在 `response_v2.rs`)是唯一的对外
|
|
||||||
//! `StreamEvent`,旧变体迁移至此模块改名为 `LegacyStreamEvent`,
|
|
||||||
//! 由 `parse_chunk_stream()` 内部消费 `LegacyStreamEvent`,对外返回值已被
|
|
||||||
//! 重映射为新 `StreamEvent`。
|
|
||||||
//!
|
|
||||||
//! Phase 1 重写 Provider 时,`parse_chunk_stream` 可直接消费新事件流后整体删除此文件。
|
|
||||||
|
|
||||||
use crate::llm::types::shared::FinishReason;
|
|
||||||
use crate::llm::types::usage::Usage;
|
|
||||||
use serde_json::Value;
|
|
||||||
|
|
||||||
/// 旧 `StreamEvent` 变体迁移后的别名 —— 仅供 `stream.rs` 内部使用。
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub enum LegacyStreamEvent {
|
|
||||||
/// 助手回复文本增量。
|
|
||||||
AssistantTextDelta { text: String },
|
|
||||||
/// 工具调用开始。
|
|
||||||
ToolExecutionStarted {
|
|
||||||
tool_name: String,
|
|
||||||
input: Value,
|
|
||||||
tool_call_id: String,
|
|
||||||
},
|
|
||||||
/// 工具调用完成。
|
|
||||||
ToolExecutionCompleted {
|
|
||||||
tool_name: String,
|
|
||||||
output: Value,
|
|
||||||
is_error: bool,
|
|
||||||
},
|
|
||||||
/// Token 用量更新。
|
|
||||||
CostUpdate { usage: Usage },
|
|
||||||
/// 一轮会话完成。
|
|
||||||
TurnComplete { reason: FinishReason },
|
|
||||||
/// 错误事件。
|
|
||||||
Error { message: String },
|
|
||||||
}
|
|
||||||
|
|
||||||
impl LegacyStreamEvent {
|
|
||||||
pub(crate) fn error(message: impl Into<String>) -> Self {
|
|
||||||
Self::Error {
|
|
||||||
message: message.into(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,184 +0,0 @@
|
|||||||
use crate::llm::types::shared::{ResponseFormat, ServiceTier, StopSequence};
|
|
||||||
use crate::llm::types::tool::OpenaiToolDefinition;
|
|
||||||
use serde::{Deserialize, Serialize};
|
|
||||||
use serde_json::Value;
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct StreamOptions {
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub include_usage: Option<bool>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub include_obfuscation: Option<bool>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
#[derive(Default)]
|
|
||||||
pub enum ToolChoice {
|
|
||||||
#[default]
|
|
||||||
None,
|
|
||||||
Auto,
|
|
||||||
Required,
|
|
||||||
Named { name: String },
|
|
||||||
AllowedTools { tool_names: Vec<String> },
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
impl Serialize for ToolChoice {
|
|
||||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
|
||||||
where
|
|
||||||
S: serde::Serializer,
|
|
||||||
{
|
|
||||||
match self {
|
|
||||||
ToolChoice::None => serializer.serialize_str("none"),
|
|
||||||
ToolChoice::Auto => serializer.serialize_str("auto"),
|
|
||||||
ToolChoice::Required => serializer.serialize_str("required"),
|
|
||||||
ToolChoice::Named { name } => {
|
|
||||||
let obj = serde_json::json!({
|
|
||||||
"type": "function",
|
|
||||||
"function": { "name": name }
|
|
||||||
});
|
|
||||||
obj.serialize(serializer)
|
|
||||||
}
|
|
||||||
ToolChoice::AllowedTools { tool_names } => {
|
|
||||||
let obj = serde_json::json!({
|
|
||||||
"type": "function",
|
|
||||||
"function": { "name": tool_names.first().cloned().unwrap_or_default() }
|
|
||||||
});
|
|
||||||
obj.serialize(serializer)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl<'de> Deserialize<'de> for ToolChoice {
|
|
||||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
|
||||||
where
|
|
||||||
D: serde::Deserializer<'de>,
|
|
||||||
{
|
|
||||||
let value = Value::deserialize(deserializer)?;
|
|
||||||
match value {
|
|
||||||
Value::String(s) => match s.as_str() {
|
|
||||||
"none" => Ok(ToolChoice::None),
|
|
||||||
"auto" => Ok(ToolChoice::Auto),
|
|
||||||
"required" => Ok(ToolChoice::Required),
|
|
||||||
_ => Err(serde::de::Error::custom(format!(
|
|
||||||
"unknown tool choice: {s}"
|
|
||||||
))),
|
|
||||||
},
|
|
||||||
Value::Object(obj) => {
|
|
||||||
let typ = obj.get("type").and_then(|v| v.as_str()).ok_or_else(|| {
|
|
||||||
serde::de::Error::custom("missing 'type' field in tool_choice")
|
|
||||||
})?;
|
|
||||||
if typ == "function" {
|
|
||||||
let func =
|
|
||||||
obj.get("function")
|
|
||||||
.and_then(|v| v.as_object())
|
|
||||||
.ok_or_else(|| {
|
|
||||||
serde::de::Error::custom("missing 'function' field in tool_choice")
|
|
||||||
})?;
|
|
||||||
let name = func.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
|
|
||||||
serde::de::Error::custom("missing 'function.name' in tool_choice")
|
|
||||||
})?;
|
|
||||||
Ok(ToolChoice::Named {
|
|
||||||
name: name.to_string(),
|
|
||||||
})
|
|
||||||
} else {
|
|
||||||
Err(serde::de::Error::custom(format!(
|
|
||||||
"unknown tool_choice type: {typ}"
|
|
||||||
)))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
_ => Err(serde::de::Error::custom(
|
|
||||||
"tool_choice must be a string or object",
|
|
||||||
)),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
#[serde(rename_all = "snake_case", tag = "type")]
|
|
||||||
pub enum OpenaiTool {
|
|
||||||
Function { function: OpenaiToolDefinition },
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct AudioParam {
|
|
||||||
pub format: String,
|
|
||||||
pub voice: String,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct PredictionContent {
|
|
||||||
#[serde(rename = "type")]
|
|
||||||
pub pred_type: String,
|
|
||||||
pub content: String,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct UserLocation {
|
|
||||||
#[serde(rename = "type")]
|
|
||||||
pub loc_type: String,
|
|
||||||
pub approximate: Approximate,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct Approximate {
|
|
||||||
pub city: String,
|
|
||||||
pub country: String,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub region: Option<String>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub timezone: Option<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct WebSearchOptions {
|
|
||||||
pub search_context_size: String,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub user_location: Option<UserLocation>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
|
||||||
#[serde(rename_all = "snake_case")]
|
|
||||||
pub struct OpenaiChatRequest {
|
|
||||||
pub model: String,
|
|
||||||
pub messages: Vec<crate::llm::types::openai_message::OpenaiChatMessage>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub frequency_penalty: Option<f32>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub logit_bias: Option<Value>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub max_tokens: Option<u32>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub n: Option<u32>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub presence_penalty: Option<f32>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub response_format: Option<ResponseFormat>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub seed: Option<i64>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub service_tier: Option<ServiceTier>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub stop: Option<StopSequence>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub stream: Option<bool>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub stream_options: Option<StreamOptions>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub temperature: Option<f32>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub top_p: Option<f32>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub tools: Option<Vec<OpenaiTool>>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub tool_choice: Option<ToolChoice>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub parallel_tool_calls: Option<bool>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub user: Option<String>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub extra_headers: Option<Value>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub extra_body: Option<Value>,
|
|
||||||
}
|
|
||||||
+49
-17
@@ -9,15 +9,15 @@ use serde_json::Value;
|
|||||||
use thiserror::Error;
|
use thiserror::Error;
|
||||||
|
|
||||||
use crate::llm::types::message::Message;
|
use crate::llm::types::message::Message;
|
||||||
use crate::llm::types::request::ToolChoice;
|
use crate::llm::types::tool::ToolChoice;
|
||||||
use crate::llm::types::tool::OpenaiToolDefinition;
|
use crate::llm::types::tool::ToolDef;
|
||||||
|
|
||||||
/// Provider 无关的请求类型。
|
/// Provider 无关的请求类型。
|
||||||
///
|
///
|
||||||
/// 设计要点:
|
/// 设计要点:
|
||||||
/// - `system` 字段不存在;system 提示由调用方通过 `Message::System` 在 `messages` 中表达。
|
/// - `system` 字段不存在;system 提示由调用方通过 `Message::System` 在 `messages` 中表达。
|
||||||
/// - `tools` / `tool_choice` 直接复用现有 `OpenaiToolDefinition` / `ToolChoice`
|
/// - `tools` 使用 Provider 无关的 `ToolDef` IR;各 Provider 适配层在 `convert_request`
|
||||||
/// (10a §251 决策:先复用旧类型,Phase 2 切换为新 `ToolDefinition` 后再调整)。
|
/// 中转换为对应 wire format。`tool_choice` 复用现有 `ToolChoice`。
|
||||||
/// - `extra` 作为逃生舱:Provider 特定字段(`web_search_options`、`previous_response_id` 等)
|
/// - `extra` 作为逃生舱:Provider 特定字段(`web_search_options`、`previous_response_id` 等)
|
||||||
/// 通过 `extra.set_extra / get_extra` 传递,避免持续膨胀本结构体。
|
/// 通过 `extra.set_extra / get_extra` 传递,避免持续膨胀本结构体。
|
||||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||||
@@ -26,8 +26,8 @@ pub struct MessageRequest {
|
|||||||
pub model: String,
|
pub model: String,
|
||||||
/// 消息列表(包含 system / user / assistant / tool_result 等所有变体)。
|
/// 消息列表(包含 system / user / assistant / tool_result 等所有变体)。
|
||||||
pub messages: Vec<Message>,
|
pub messages: Vec<Message>,
|
||||||
/// 工具定义列表。
|
/// 工具定义列表(Provider 无关 IR)。
|
||||||
pub tools: Vec<OpenaiToolDefinition>,
|
pub tools: Vec<ToolDef>,
|
||||||
/// 工具选择策略。
|
/// 工具选择策略。
|
||||||
pub tool_choice: ToolChoice,
|
pub tool_choice: ToolChoice,
|
||||||
/// 最大输出 token 数。
|
/// 最大输出 token 数。
|
||||||
@@ -127,14 +127,9 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn extra_set_and_get_roundtrip() {
|
fn extra_set_and_get_roundtrip() {
|
||||||
let mut req = MessageRequest::default();
|
let mut req = MessageRequest::default();
|
||||||
req.set_extra(
|
req.set_extra("previous_response_id", "resp_abc123");
|
||||||
"previous_response_id",
|
|
||||||
"resp_abc123",
|
|
||||||
);
|
|
||||||
|
|
||||||
let v: Option<String> = req
|
let v: Option<String> = req.get_extra("previous_response_id").expect("get_extra ok");
|
||||||
.get_extra("previous_response_id")
|
|
||||||
.expect("get_extra ok");
|
|
||||||
assert_eq!(v.as_deref(), Some("resp_abc123"));
|
assert_eq!(v.as_deref(), Some("resp_abc123"));
|
||||||
|
|
||||||
let missing: Option<String> = req.get_extra("missing").expect("missing ok");
|
let missing: Option<String> = req.get_extra("missing").expect("missing ok");
|
||||||
@@ -174,10 +169,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let opts: Options = req.get_extra_as().expect("get_extra_as ok");
|
let opts: Options = req.get_extra_as().expect("get_extra_as ok");
|
||||||
assert_eq!(
|
assert_eq!(opts.web_search_options.search_context_size, "high");
|
||||||
opts.web_search_options.search_context_size,
|
|
||||||
"high"
|
|
||||||
);
|
|
||||||
assert_eq!(opts.user.as_deref(), Some("u_123"));
|
assert_eq!(opts.user.as_deref(), Some("u_123"));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -206,4 +198,44 @@ mod tests {
|
|||||||
assert_eq!(decoded.stream, req.stream);
|
assert_eq!(decoded.stream, req.stream);
|
||||||
assert_eq!(decoded.extra.get("trace_id"), Some(&json!("t-1")));
|
assert_eq!(decoded.extra.get("trace_id"), Some(&json!("t-1")));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn message_request_with_tools_roundtrip() {
|
||||||
|
// 验证 ToolDef 的 serde 属性与 OpenaiToolDefinition 一致:
|
||||||
|
// 同名字段(name/description/parameters)序列化结果应一致。
|
||||||
|
let params = json!({
|
||||||
|
"type": "object",
|
||||||
|
"properties": {"x": {"type": "number"}},
|
||||||
|
"required": ["x"],
|
||||||
|
});
|
||||||
|
let tool = super::ToolDef {
|
||||||
|
name: "add".to_string(),
|
||||||
|
description: Some("add two numbers".to_string()),
|
||||||
|
parameters: params.clone(),
|
||||||
|
};
|
||||||
|
let req = MessageRequest {
|
||||||
|
model: "gpt-4o".into(),
|
||||||
|
messages: vec![Message::user_text("hi")],
|
||||||
|
tools: vec![tool],
|
||||||
|
tool_choice: ToolChoice::Auto,
|
||||||
|
max_tokens: None,
|
||||||
|
temperature: None,
|
||||||
|
top_p: None,
|
||||||
|
stop_sequences: vec![],
|
||||||
|
stream: false,
|
||||||
|
thinking: None,
|
||||||
|
extra: HashMap::new(),
|
||||||
|
};
|
||||||
|
|
||||||
|
let json = serde_json::to_string(&req).expect("serialize");
|
||||||
|
// 验证反序列化能还原所有字段(包括嵌套 parameters)
|
||||||
|
let decoded: MessageRequest = serde_json::from_str(&json).expect("deserialize");
|
||||||
|
assert_eq!(decoded.tools.len(), 1);
|
||||||
|
assert_eq!(decoded.tools[0].name, "add");
|
||||||
|
assert_eq!(decoded.tools[0].description.as_deref(), Some("add two numbers"));
|
||||||
|
assert_eq!(decoded.tools[0].parameters, params);
|
||||||
|
|
||||||
|
// 验证序列化 JSON 不含 ToolDef 没有的字段(如 strict),保持 wire-format 兼容
|
||||||
|
assert!(!json.contains("strict"), "ToolDef 序列化不应包含 strict 字段");
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,181 +0,0 @@
|
|||||||
use crate::llm::types::openai_message::OpenaiChatMessage;
|
|
||||||
use crate::llm::types::shared::{FinishReason, ServiceTier};
|
|
||||||
use crate::llm::types::tool::OpenaiToolCall;
|
|
||||||
use crate::llm::types::usage::Usage;
|
|
||||||
use serde::{Deserialize, Serialize};
|
|
||||||
use crate::llm::types::{ContentField, OpenaiContentPart};
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct TokenLogprob {
|
|
||||||
pub token: String,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub bytes: Option<Vec<u32>>,
|
|
||||||
pub logprob: f64,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub top_logprobs: Option<Vec<TopLogprob>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct TopLogprob {
|
|
||||||
pub token: String,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub bytes: Option<Vec<u32>>,
|
|
||||||
pub logprob: f64,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct Logprobs {
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub content: Option<Vec<TokenLogprob>>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub refusal: Option<Vec<TokenLogprob>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct URLCitation {
|
|
||||||
pub end_index: u32,
|
|
||||||
pub start_index: u32,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub title: Option<String>,
|
|
||||||
pub url: String,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct Annotation {
|
|
||||||
#[serde(rename = "type")]
|
|
||||||
pub ann_type: String,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub url_citation: Option<URLCitation>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct OpenaiAudio {
|
|
||||||
pub id: String,
|
|
||||||
pub data: String,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub expires_at: Option<i64>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub transcript: Option<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct Choice {
|
|
||||||
pub index: u32,
|
|
||||||
pub message: OpenaiChatMessage,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub finish_reason: Option<FinishReason>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub logprobs: Option<Logprobs>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct OpenaiChatResponse {
|
|
||||||
pub id: String,
|
|
||||||
pub object: String,
|
|
||||||
pub created: u64,
|
|
||||||
pub model: String,
|
|
||||||
pub choices: Vec<Choice>,
|
|
||||||
pub usage: Usage,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub system_fingerprint: Option<String>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub service_tier: Option<ServiceTier>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct Delta {
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub role: Option<String>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub content: Option<String>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub refusal: Option<String>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub tool_calls: Option<Vec<OpenaiToolCall>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct ChunkChoice {
|
|
||||||
pub index: u32,
|
|
||||||
pub delta: Delta,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub logprobs: Option<Logprobs>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub finish_reason: Option<FinishReason>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
||||||
pub struct OpenaiChatChunk {
|
|
||||||
pub id: String,
|
|
||||||
pub object: String,
|
|
||||||
pub created: u64,
|
|
||||||
pub model: String,
|
|
||||||
pub choices: Vec<ChunkChoice>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub usage: Option<Usage>,
|
|
||||||
#[serde(skip_serializing_if = "Option::is_none")]
|
|
||||||
pub system_fingerprint: Option<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<OpenaiChatMessage> for Delta {
|
|
||||||
fn from(msg: OpenaiChatMessage) -> Self {
|
|
||||||
match msg {
|
|
||||||
OpenaiChatMessage::Assistant {
|
|
||||||
content,
|
|
||||||
tool_calls,
|
|
||||||
..
|
|
||||||
} => Delta {
|
|
||||||
role: Some("assistant".to_string()),
|
|
||||||
content: match content {
|
|
||||||
ContentField::String(s) => Some(s),
|
|
||||||
ContentField::Array(parts) => {
|
|
||||||
let mut text = String::new();
|
|
||||||
for part in parts {
|
|
||||||
if let OpenaiContentPart::Text { text: t } = part {
|
|
||||||
text.push_str(&t);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if text.is_empty() {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
Some(text)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
},
|
|
||||||
refusal: None,
|
|
||||||
tool_calls,
|
|
||||||
},
|
|
||||||
_ => Delta {
|
|
||||||
role: None,
|
|
||||||
content: None,
|
|
||||||
refusal: None,
|
|
||||||
tool_calls: None,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<OpenaiChatResponse> for OpenaiChatChunk {
|
|
||||||
fn from(response: OpenaiChatResponse) -> Self {
|
|
||||||
let choices = response
|
|
||||||
.choices
|
|
||||||
.into_iter()
|
|
||||||
.map(|c| ChunkChoice {
|
|
||||||
index: c.index,
|
|
||||||
delta: Delta::from(c.message),
|
|
||||||
logprobs: c.logprobs,
|
|
||||||
finish_reason: c.finish_reason,
|
|
||||||
})
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
OpenaiChatChunk {
|
|
||||||
id: response.id,
|
|
||||||
object: "chat.completion.chunk".to_string(),
|
|
||||||
created: response.created,
|
|
||||||
model: response.model,
|
|
||||||
choices,
|
|
||||||
usage: Some(response.usage),
|
|
||||||
system_fingerprint: response.system_fingerprint,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -19,6 +19,7 @@ use crate::llm::types::usage::{CompletionTokensDetails, PromptTokensDetails, Usa
|
|||||||
/// Phase 2 完成时统一收敛。
|
/// Phase 2 完成时统一收敛。
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum StopReason {
|
pub enum StopReason {
|
||||||
/// 自然停止。
|
/// 自然停止。
|
||||||
Stop,
|
Stop,
|
||||||
@@ -164,11 +165,15 @@ pub enum ContentBlockBuilder {
|
|||||||
/// `thinking_signature`,最终通过 `finalize()` 回填到 `full_response` 的 `Thinking` block 中。
|
/// `thinking_signature`,最终通过 `finalize()` 回填到 `full_response` 的 `Thinking` block 中。
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum StreamEvent {
|
pub enum StreamEvent {
|
||||||
/// 消息开始(元信息)。
|
/// 消息开始(元信息)。
|
||||||
MessageStart { id: String, model: String },
|
MessageStart { id: String, model: String },
|
||||||
/// 内容块开始(告知块类型,携带 id/name for ToolUse)。
|
/// 内容块开始(告知块类型,携带 id/name for ToolUse)。
|
||||||
ContentBlockStart { index: u32, block_type: ContentBlockType },
|
ContentBlockStart {
|
||||||
|
index: u32,
|
||||||
|
block_type: ContentBlockType,
|
||||||
|
},
|
||||||
/// 内容块结束标记。
|
/// 内容块结束标记。
|
||||||
ContentBlockEnd { index: u32 },
|
ContentBlockEnd { index: u32 },
|
||||||
/// 文本增量。
|
/// 文本增量。
|
||||||
@@ -187,6 +192,24 @@ pub enum StreamEvent {
|
|||||||
MessageComplete { full_response: MessageResponse },
|
MessageComplete { full_response: MessageResponse },
|
||||||
/// 错误事件。
|
/// 错误事件。
|
||||||
Error { message: String },
|
Error { message: String },
|
||||||
|
/// 工具开始执行 —— 在 `ToolCallEnd` 之后、`registry.invoke_all` 之前发出。
|
||||||
|
/// 让 UI 层可以显示 "正在执行工具:add(1, 2)"。
|
||||||
|
ToolExecutionStarted {
|
||||||
|
tool_name: String,
|
||||||
|
tool_call_id: String,
|
||||||
|
/// 工具参数(JSON 字符串形式),用于 UI 展示
|
||||||
|
arguments: String,
|
||||||
|
},
|
||||||
|
/// 工具执行完成 —— 在工具返回后、新一轮 LLM 流开始之前发出。
|
||||||
|
ToolExecutionCompleted {
|
||||||
|
tool_name: String,
|
||||||
|
tool_call_id: String,
|
||||||
|
/// 结果摘要(由 `CycleConfig.max_tool_result_bytes` 截断,默认 65536 字节/字符边界安全),
|
||||||
|
/// 用于 UI 反馈。完整结果已在内部 `messages` 中作为 `ToolResult` 回传给 LLM。
|
||||||
|
result_summary: String,
|
||||||
|
/// 是否执行出错
|
||||||
|
is_error: bool,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 流式响应累积状态。
|
/// 流式响应累积状态。
|
||||||
@@ -320,9 +343,8 @@ impl PartialMessageResponse {
|
|||||||
true
|
true
|
||||||
}
|
}
|
||||||
StreamEvent::ToolCallArgumentsDelta { index, arguments } => {
|
StreamEvent::ToolCallArgumentsDelta { index, arguments } => {
|
||||||
if let Some(ContentBlockBuilder::ToolUse {
|
if let Some(ContentBlockBuilder::ToolUse { arguments: buf, .. }) =
|
||||||
arguments: buf, ..
|
self.blocks.get_mut(index)
|
||||||
}) = self.blocks.get_mut(index)
|
|
||||||
{
|
{
|
||||||
buf.push_str(arguments);
|
buf.push_str(arguments);
|
||||||
}
|
}
|
||||||
@@ -348,6 +370,9 @@ impl PartialMessageResponse {
|
|||||||
self.is_errored = true;
|
self.is_errored = true;
|
||||||
false
|
false
|
||||||
}
|
}
|
||||||
|
// 元事件:不参与内容块累积,不修改 partial 状态
|
||||||
|
//(Phase 9 —— 工具执行透明化,由 run_tool_loop 在工具前后插入)
|
||||||
|
StreamEvent::ToolExecutionStarted { .. } | StreamEvent::ToolExecutionCompleted { .. } => true,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -360,9 +385,7 @@ impl PartialMessageResponse {
|
|||||||
let mut content_blocks = Vec::with_capacity(self.blocks.len());
|
let mut content_blocks = Vec::with_capacity(self.blocks.len());
|
||||||
for (idx, builder) in self.blocks {
|
for (idx, builder) in self.blocks {
|
||||||
let block = Self::builder_to_block(idx, builder, self.thinking_signature.as_deref())
|
let block = Self::builder_to_block(idx, builder, self.thinking_signature.as_deref())
|
||||||
.map_err(|e| LlmError::Other(format!(
|
.map_err(|e| LlmError::Other(format!("partial 块 #{idx} finalize 失败: {e}")))?;
|
||||||
"partial 块 #{idx} finalize 失败: {e}"
|
|
||||||
)))?;
|
|
||||||
content_blocks.push(block);
|
content_blocks.push(block);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -743,10 +766,7 @@ mod tests {
|
|||||||
Message::Assistant { content } => {
|
Message::Assistant { content } => {
|
||||||
assert_eq!(content.len(), 2);
|
assert_eq!(content.len(), 2);
|
||||||
match (&content[0], &content[1]) {
|
match (&content[0], &content[1]) {
|
||||||
(
|
(ContentBlock::Text { text: t1 }, ContentBlock::Text { text: t2 }) => {
|
||||||
ContentBlock::Text { text: t1 },
|
|
||||||
ContentBlock::Text { text: t2 },
|
|
||||||
) => {
|
|
||||||
assert_eq!(t1, "first");
|
assert_eq!(t1, "first");
|
||||||
assert_eq!(t2, "second");
|
assert_eq!(t2, "second");
|
||||||
}
|
}
|
||||||
@@ -772,9 +792,7 @@ mod tests {
|
|||||||
index: 0,
|
index: 0,
|
||||||
block_type: ContentBlockType::Text,
|
block_type: ContentBlockType::Text,
|
||||||
},
|
},
|
||||||
StreamEvent::TextDelta {
|
StreamEvent::TextDelta { text: "x".into() },
|
||||||
text: "x".into(),
|
|
||||||
},
|
|
||||||
StreamEvent::ContentBlockEnd { index: 0 },
|
StreamEvent::ContentBlockEnd { index: 0 },
|
||||||
StreamEvent::MessageComplete {
|
StreamEvent::MessageComplete {
|
||||||
full_response: empty_response(),
|
full_response: empty_response(),
|
||||||
@@ -832,15 +850,9 @@ mod tests {
|
|||||||
block_type: ContentBlockType::Text,
|
block_type: ContentBlockType::Text,
|
||||||
},
|
},
|
||||||
StreamEvent::ContentBlockEnd { index: 0 },
|
StreamEvent::ContentBlockEnd { index: 0 },
|
||||||
StreamEvent::TextDelta {
|
StreamEvent::TextDelta { text: "t".into() },
|
||||||
text: "t".into(),
|
StreamEvent::ThinkingDelta { text: "p".into() },
|
||||||
},
|
StreamEvent::RefusalDelta { text: "r".into() },
|
||||||
StreamEvent::ThinkingDelta {
|
|
||||||
text: "p".into(),
|
|
||||||
},
|
|
||||||
StreamEvent::RefusalDelta {
|
|
||||||
text: "r".into(),
|
|
||||||
},
|
|
||||||
StreamEvent::ToolCallArgumentsDelta {
|
StreamEvent::ToolCallArgumentsDelta {
|
||||||
index: 1,
|
index: 1,
|
||||||
arguments: "{\"x\":1}".into(),
|
arguments: "{\"x\":1}".into(),
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ pub enum Role {
|
|||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum FinishReason {
|
pub enum FinishReason {
|
||||||
Stop,
|
Stop,
|
||||||
Length,
|
Length,
|
||||||
@@ -67,6 +68,7 @@ pub enum StopSequence {
|
|||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
#[serde(rename_all = "snake_case", tag = "type")]
|
#[serde(rename_all = "snake_case", tag = "type")]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum ResponseFormat {
|
pub enum ResponseFormat {
|
||||||
Text,
|
Text,
|
||||||
JsonObject,
|
JsonObject,
|
||||||
|
|||||||
@@ -1,6 +1,25 @@
|
|||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
|
/// Provider 无关的工具定义 IR(v0.2 引入,替换 `ToolDefinition` 别名)。
|
||||||
|
///
|
||||||
|
/// 字段最小化:仅承载跨 Provider 公共的概念(name、description、parameters)。
|
||||||
|
/// OpenAI 专属 `strict` 字段不在此表达,由 OpenAI 适配层通过
|
||||||
|
/// `MessageRequest.extra` 逃生舱在 `convert_request` 内补充。
|
||||||
|
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||||
|
pub struct ToolDef {
|
||||||
|
pub name: String,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
pub description: Option<String>,
|
||||||
|
#[serde(default)]
|
||||||
|
pub parameters: Value,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 旧 OpenAI wire-format 工具定义(v0.2 降级为 `#[doc(hidden)]`)。
|
||||||
|
///
|
||||||
|
/// 由 `ToolDef` 替代;保留仅供 OpenAI 适配层消费 `ToolDef → OpenaiToolDefinition`
|
||||||
|
/// 转换与外部反序列化兼容路径使用,不作为公共 API。
|
||||||
|
#[doc(hidden)]
|
||||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||||
pub struct OpenaiToolDefinition {
|
pub struct OpenaiToolDefinition {
|
||||||
pub name: String,
|
pub name: String,
|
||||||
@@ -12,6 +31,27 @@ pub struct OpenaiToolDefinition {
|
|||||||
pub strict: Option<bool>,
|
pub strict: Option<bool>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl From<ToolDef> for OpenaiToolDefinition {
|
||||||
|
fn from(t: ToolDef) -> Self {
|
||||||
|
Self {
|
||||||
|
name: t.name,
|
||||||
|
description: t.description,
|
||||||
|
parameters: t.parameters,
|
||||||
|
strict: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl From<OpenaiToolDefinition> for ToolDef {
|
||||||
|
fn from(t: OpenaiToolDefinition) -> Self {
|
||||||
|
Self {
|
||||||
|
name: t.name,
|
||||||
|
description: t.description,
|
||||||
|
parameters: t.parameters,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
pub struct FunctionCall {
|
pub struct FunctionCall {
|
||||||
pub name: String,
|
pub name: String,
|
||||||
@@ -23,3 +63,93 @@ pub struct FunctionCall {
|
|||||||
pub enum OpenaiToolCall {
|
pub enum OpenaiToolCall {
|
||||||
Function { id: String, function: FunctionCall },
|
Function { id: String, function: FunctionCall },
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 工具选择策略 —— Phase 13 从 `types::request::ToolChoice` 迁入。
|
||||||
|
///
|
||||||
|
/// `#[non_exhaustive]` 预留扩展空间。
|
||||||
|
#[derive(Debug, Clone, Default)]
|
||||||
|
#[non_exhaustive]
|
||||||
|
pub enum ToolChoice {
|
||||||
|
#[default]
|
||||||
|
None,
|
||||||
|
Auto,
|
||||||
|
Required,
|
||||||
|
Named {
|
||||||
|
name: String,
|
||||||
|
},
|
||||||
|
AllowedTools {
|
||||||
|
tool_names: Vec<String>,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Serialize for ToolChoice {
|
||||||
|
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||||
|
where
|
||||||
|
S: serde::Serializer,
|
||||||
|
{
|
||||||
|
match self {
|
||||||
|
ToolChoice::None => serializer.serialize_str("none"),
|
||||||
|
ToolChoice::Auto => serializer.serialize_str("auto"),
|
||||||
|
ToolChoice::Required => serializer.serialize_str("required"),
|
||||||
|
ToolChoice::Named { name } => {
|
||||||
|
let obj = serde_json::json!({
|
||||||
|
"type": "function",
|
||||||
|
"function": { "name": name }
|
||||||
|
});
|
||||||
|
obj.serialize(serializer)
|
||||||
|
}
|
||||||
|
ToolChoice::AllowedTools { tool_names } => {
|
||||||
|
let obj = serde_json::json!({
|
||||||
|
"type": "function",
|
||||||
|
"function": { "name": tool_names.first().cloned().unwrap_or_default() }
|
||||||
|
});
|
||||||
|
obj.serialize(serializer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<'de> Deserialize<'de> for ToolChoice {
|
||||||
|
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||||
|
where
|
||||||
|
D: serde::Deserializer<'de>,
|
||||||
|
{
|
||||||
|
let value = Value::deserialize(deserializer)?;
|
||||||
|
match value {
|
||||||
|
Value::String(s) => match s.as_str() {
|
||||||
|
"none" => Ok(ToolChoice::None),
|
||||||
|
"auto" => Ok(ToolChoice::Auto),
|
||||||
|
"required" => Ok(ToolChoice::Required),
|
||||||
|
_ => Err(serde::de::Error::custom(format!(
|
||||||
|
"unknown tool choice: {s}"
|
||||||
|
))),
|
||||||
|
},
|
||||||
|
Value::Object(obj) => {
|
||||||
|
let typ = obj.get("type").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||||
|
serde::de::Error::custom("missing 'type' field in tool_choice")
|
||||||
|
})?;
|
||||||
|
if typ == "function" {
|
||||||
|
let func =
|
||||||
|
obj.get("function")
|
||||||
|
.and_then(|v| v.as_object())
|
||||||
|
.ok_or_else(|| {
|
||||||
|
serde::de::Error::custom("missing 'function' field in tool_choice")
|
||||||
|
})?;
|
||||||
|
let name = func.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||||
|
serde::de::Error::custom("missing 'function.name' in tool_choice")
|
||||||
|
})?;
|
||||||
|
Ok(ToolChoice::Named {
|
||||||
|
name: name.to_string(),
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
Err(serde::de::Error::custom(format!(
|
||||||
|
"unknown tool_choice type: {typ}"
|
||||||
|
)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => Err(serde::de::Error::custom(
|
||||||
|
"tool_choice must be a string or object",
|
||||||
|
)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+8
-3
@@ -6,17 +6,22 @@ pub mod knowledge;
|
|||||||
pub mod retriever;
|
pub mod retriever;
|
||||||
pub mod store;
|
pub mod store;
|
||||||
pub mod types;
|
pub mod types;
|
||||||
|
pub mod vector;
|
||||||
|
pub mod vector_store;
|
||||||
|
|
||||||
// 高频类型(大多数下游需要)
|
// 高频类型(大多数下游需要)
|
||||||
pub use conversation::{ConversationMemory, ConversationMemoryConfig};
|
pub use conversation::{ConversationMemory, ConversationMemoryConfig};
|
||||||
pub use error::MemoryError;
|
pub use error::MemoryError;
|
||||||
pub use knowledge::KnowledgeStore;
|
pub use knowledge::KnowledgeStore;
|
||||||
pub use retriever::MemoryRetriever;
|
pub use retriever::MemoryRetriever;
|
||||||
pub use store::{InMemoryStore, MemoryStore};
|
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
|
||||||
|
#[allow(deprecated)]
|
||||||
|
pub use vector::{InMemoryVectorRetriever, VectorRetriever};
|
||||||
|
pub use vector_store::{InMemoryVectorStore, PersistentVectorStore, RagPipeline, VectorStore};
|
||||||
|
|
||||||
// 低频类型(配置/高级使用)
|
// 低频类型(配置/高级使用)
|
||||||
pub use conversation::MemoryStrategy;
|
pub use conversation::MemoryStrategy;
|
||||||
pub use knowledge::{PageIndexEntry, KNOWLEDGE_PREFIX};
|
pub use knowledge::{KNOWLEDGE_PREFIX, PageIndexEntry};
|
||||||
pub use retriever::{RetrieverConfig, RetrievalResult, ScoredItem};
|
pub use retriever::{RetrievalResult, RetrieverConfig, ScoredItem};
|
||||||
pub use store::{EvictionConfig, EvictionPolicy};
|
pub use store::{EvictionConfig, EvictionPolicy};
|
||||||
pub use types::{KnowledgePage, MemoryFilter, MemoryItem};
|
pub use types::{KnowledgePage, MemoryFilter, MemoryItem};
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ use crate::memory::types::MemoryItem;
|
|||||||
|
|
||||||
/// 对话消息管理策略。
|
/// 对话消息管理策略。
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum MemoryStrategy {
|
pub enum MemoryStrategy {
|
||||||
/// 滑动窗口:达到上限时删除最旧消息。
|
/// 滑动窗口:达到上限时删除最旧消息。
|
||||||
SlidingWindow,
|
SlidingWindow,
|
||||||
@@ -160,7 +161,12 @@ impl ConversationMemory {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn make_message_id(&self, index: usize, now: &OffsetDateTime) -> String {
|
fn make_message_id(&self, index: usize, now: &OffsetDateTime) -> String {
|
||||||
format!("{}{:010}_{}", self.session_prefix(), index, now.unix_timestamp_nanos())
|
format!(
|
||||||
|
"{}{:010}_{}",
|
||||||
|
self.session_prefix(),
|
||||||
|
index,
|
||||||
|
now.unix_timestamp_nanos()
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn maybe_evict_and_compact(&mut self) {
|
async fn maybe_evict_and_compact(&mut self) {
|
||||||
@@ -175,7 +181,8 @@ impl ConversationMemory {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if let Some(ref compact_config) = self.config.compact_config
|
if let Some(ref compact_config) = self.config.compact_config
|
||||||
&& should_compact(&self.messages, compact_config, &self.compact_state) {
|
&& should_compact(&self.messages, compact_config, &self.compact_state)
|
||||||
|
{
|
||||||
let keep_recent = compact_config.keep_recent;
|
let keep_recent = compact_config.keep_recent;
|
||||||
let freed = microcompact(&mut self.messages, keep_recent);
|
let freed = microcompact(&mut self.messages, keep_recent);
|
||||||
if freed > 0 {
|
if freed > 0 {
|
||||||
@@ -196,7 +203,8 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn add_and_get_history() {
|
async fn add_and_get_history() {
|
||||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||||
let mut conv = ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
|
let mut conv =
|
||||||
|
ConversationMemory::new(store, "session1", ConversationMemoryConfig::default());
|
||||||
conv.add_message(Message::user_text("hello")).await.unwrap();
|
conv.add_message(Message::user_text("hello")).await.unwrap();
|
||||||
conv.add_message(Message::user_text("world")).await.unwrap();
|
conv.add_message(Message::user_text("world")).await.unwrap();
|
||||||
assert_eq!(conv.len(), 2);
|
assert_eq!(conv.len(), 2);
|
||||||
@@ -211,9 +219,7 @@ mod tests {
|
|||||||
conv.add_message(Message::tool_result("call_1", "ok", false))
|
conv.add_message(Message::tool_result("call_1", "ok", false))
|
||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
conv.add_message(Message::assistant("done"))
|
conv.add_message(Message::assistant("done")).await.unwrap();
|
||||||
.await
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
let original = conv.get_history().to_vec();
|
let original = conv.get_history().to_vec();
|
||||||
assert_eq!(original.len(), 2);
|
assert_eq!(original.len(), 2);
|
||||||
@@ -263,7 +269,8 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn clear_empties_messages() {
|
async fn clear_empties_messages() {
|
||||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||||
let mut conv = ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
|
let mut conv =
|
||||||
|
ConversationMemory::new(store.clone(), "s1", ConversationMemoryConfig::default());
|
||||||
conv.add_message(Message::user_text("hello")).await.unwrap();
|
conv.add_message(Message::user_text("hello")).await.unwrap();
|
||||||
assert!(!conv.is_empty());
|
assert!(!conv.is_empty());
|
||||||
conv.clear().await.unwrap();
|
conv.clear().await.unwrap();
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ use thiserror::Error;
|
|||||||
///
|
///
|
||||||
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查环境变量、重试)。
|
/// 错误消息面向最终用户(中文),并尽量附带可操作的修复建议(如检查环境变量、重试)。
|
||||||
#[derive(Debug, Error)]
|
#[derive(Debug, Error)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum MemoryError {
|
pub enum MemoryError {
|
||||||
/// 按 ID 未找到指定记忆条目。可重试——通常是 namespace 拼写错误或条目已被淘汰。
|
/// 按 ID 未找到指定记忆条目。可重试——通常是 namespace 拼写错误或条目已被淘汰。
|
||||||
#[error("未找到记忆条目 '{0}',请检查 ID 或 namespace 是否正确")]
|
#[error("未找到记忆条目 '{0}',请检查 ID 或 namespace 是否正确")]
|
||||||
|
|||||||
@@ -57,8 +57,8 @@ impl KnowledgeStore {
|
|||||||
}
|
}
|
||||||
let now = OffsetDateTime::now_utc();
|
let now = OffsetDateTime::now_utc();
|
||||||
let id = format!("{KNOWLEDGE_PREFIX}{}", page.id);
|
let id = format!("{KNOWLEDGE_PREFIX}{}", page.id);
|
||||||
let content = serde_json::to_string(&page)
|
let content =
|
||||||
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
|
serde_json::to_string(&page).map_err(|e| MemoryError::Serialization(e.to_string()))?;
|
||||||
let item = MemoryItem {
|
let item = MemoryItem {
|
||||||
id,
|
id,
|
||||||
content,
|
content,
|
||||||
@@ -128,7 +128,10 @@ impl KnowledgeStore {
|
|||||||
.filter(|entry| {
|
.filter(|entry| {
|
||||||
entry.title.to_lowercase().contains(&needle)
|
entry.title.to_lowercase().contains(&needle)
|
||||||
|| entry.summary.to_lowercase().contains(&needle)
|
|| entry.summary.to_lowercase().contains(&needle)
|
||||||
|| entry.tags.iter().any(|t| t.to_lowercase().contains(&needle))
|
|| entry
|
||||||
|
.tags
|
||||||
|
.iter()
|
||||||
|
.any(|t| t.to_lowercase().contains(&needle))
|
||||||
})
|
})
|
||||||
.map(|entry| entry.id.clone())
|
.map(|entry| entry.id.clone())
|
||||||
.collect()
|
.collect()
|
||||||
|
|||||||
+15
-8
@@ -97,7 +97,11 @@ impl MemoryRetriever {
|
|||||||
|
|
||||||
// 4. 过滤 → 排序 → 截取
|
// 4. 过滤 → 排序 → 截取
|
||||||
items.retain(|i| i.score >= self.config.min_score);
|
items.retain(|i| i.score >= self.config.min_score);
|
||||||
items.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
|
items.sort_by(|a, b| {
|
||||||
|
b.score
|
||||||
|
.partial_cmp(&a.score)
|
||||||
|
.unwrap_or(std::cmp::Ordering::Equal)
|
||||||
|
});
|
||||||
items.truncate(self.config.max_results);
|
items.truncate(self.config.max_results);
|
||||||
|
|
||||||
Ok(RetrievalResult {
|
Ok(RetrievalResult {
|
||||||
@@ -159,12 +163,12 @@ fn char_bigrams(s: &str) -> Vec<String> {
|
|||||||
|
|
||||||
fn default_stop_words() -> HashSet<String> {
|
fn default_stop_words() -> HashSet<String> {
|
||||||
[
|
[
|
||||||
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has",
|
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had",
|
||||||
"had", "do", "does", "did", "will", "would", "should", "could", "may", "might", "shall",
|
"do", "does", "did", "will", "would", "should", "could", "may", "might", "shall", "can",
|
||||||
"can", "this", "that", "these", "those", "it", "its", "they", "them", "their", "what",
|
"this", "that", "these", "those", "it", "its", "they", "them", "their", "what", "which",
|
||||||
"which", "who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor",
|
"who", "whom", "how", "when", "where", "and", "or", "but", "not", "no", "nor", "so", "if",
|
||||||
"so", "if", "then", "else", "with", "without", "for", "to", "from", "in", "on", "at",
|
"then", "else", "with", "without", "for", "to", "from", "in", "on", "at", "by", "of", "as",
|
||||||
"by", "of", "as", "into", "through", "during", "before", "after", "above", "below",
|
"into", "through", "during", "before", "after", "above", "below",
|
||||||
]
|
]
|
||||||
.iter()
|
.iter()
|
||||||
.map(|s| s.to_string())
|
.map(|s| s.to_string())
|
||||||
@@ -236,7 +240,10 @@ mod tests {
|
|||||||
min_score: 0.99,
|
min_score: 0.99,
|
||||||
};
|
};
|
||||||
let retriever = MemoryRetriever::new(ks, config);
|
let retriever = MemoryRetriever::new(ks, config);
|
||||||
let result = retriever.retrieve("totally unrelated content").await.unwrap();
|
let result = retriever
|
||||||
|
.retrieve("totally unrelated content")
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
assert!(result.items.is_empty());
|
assert!(result.items.is_empty());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+7
-259
@@ -1,14 +1,16 @@
|
|||||||
//! MemoryStore 抽象接口与默认实现。
|
//! MemoryStore 抽象接口与默认实现。
|
||||||
|
|
||||||
use std::collections::HashMap;
|
|
||||||
use std::sync::Mutex;
|
|
||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use time::OffsetDateTime;
|
|
||||||
|
|
||||||
use crate::memory::error::MemoryError;
|
use crate::memory::error::MemoryError;
|
||||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||||
|
|
||||||
|
pub mod in_memory;
|
||||||
|
pub mod sqlite_store;
|
||||||
|
|
||||||
|
pub use in_memory::InMemoryStore;
|
||||||
|
pub use sqlite_store::SqliteStore;
|
||||||
|
|
||||||
/// 底层记忆存储抽象接口。
|
/// 底层记忆存储抽象接口。
|
||||||
///
|
///
|
||||||
/// 下游可实现此 trait 以对接持久化后端(JSON 文件、SQLite、Redis 等)。
|
/// 下游可实现此 trait 以对接持久化后端(JSON 文件、SQLite、Redis 等)。
|
||||||
@@ -32,6 +34,7 @@ pub trait MemoryStore: Send + Sync {
|
|||||||
|
|
||||||
/// 淘汰策略。
|
/// 淘汰策略。
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum EvictionPolicy {
|
pub enum EvictionPolicy {
|
||||||
/// 不淘汰(默认)。
|
/// 不淘汰(默认)。
|
||||||
None,
|
None,
|
||||||
@@ -57,258 +60,3 @@ impl Default for EvictionConfig {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
|
|
||||||
pub struct InMemoryStore {
|
|
||||||
items: Mutex<HashMap<String, MemoryItem>>,
|
|
||||||
eviction: EvictionConfig,
|
|
||||||
/// 自上次淘汰检查以来的写入次数。
|
|
||||||
writes_since_check: Mutex<usize>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl InMemoryStore {
|
|
||||||
/// 创建一个无淘汰策略的 InMemoryStore。
|
|
||||||
pub fn new() -> Self {
|
|
||||||
Self {
|
|
||||||
items: Mutex::new(HashMap::new()),
|
|
||||||
eviction: EvictionConfig::default(),
|
|
||||||
writes_since_check: Mutex::new(0),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// 创建一个带淘汰配置的 InMemoryStore。
|
|
||||||
pub fn with_eviction(eviction: EvictionConfig) -> Self {
|
|
||||||
Self {
|
|
||||||
items: Mutex::new(HashMap::new()),
|
|
||||||
eviction,
|
|
||||||
writes_since_check: Mutex::new(0),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn maybe_evict(&self) {
|
|
||||||
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
|
|
||||||
let should_check = {
|
|
||||||
let mut counter = self.writes_since_check.lock().unwrap();
|
|
||||||
*counter += 1;
|
|
||||||
if *counter >= self.eviction.check_interval {
|
|
||||||
*counter = 0;
|
|
||||||
true
|
|
||||||
} else {
|
|
||||||
false
|
|
||||||
}
|
|
||||||
};
|
|
||||||
if !should_check {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
|
|
||||||
let policy = self.eviction.policy.clone();
|
|
||||||
match policy {
|
|
||||||
EvictionPolicy::None => {}
|
|
||||||
EvictionPolicy::Ttl { ttl_secs } => {
|
|
||||||
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
|
|
||||||
let mut items = self.items.lock().unwrap();
|
|
||||||
items.retain(|_, v| v.created_at > cutoff);
|
|
||||||
}
|
|
||||||
EvictionPolicy::Capacity { max_items } => {
|
|
||||||
let mut items = self.items.lock().unwrap();
|
|
||||||
if items.len() > max_items {
|
|
||||||
let mut vec: Vec<_> = items.drain().collect();
|
|
||||||
// O(n) 部分排序:保留 created_at 最大的 max_items 个
|
|
||||||
vec.select_nth_unstable_by(max_items, |a, b| {
|
|
||||||
b.1.created_at.cmp(&a.1.created_at)
|
|
||||||
});
|
|
||||||
vec.truncate(max_items);
|
|
||||||
*items = vec.into_iter().collect();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Default for InMemoryStore {
|
|
||||||
fn default() -> Self {
|
|
||||||
Self::new()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[async_trait]
|
|
||||||
impl MemoryStore for InMemoryStore {
|
|
||||||
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
|
|
||||||
{
|
|
||||||
let mut items = self.items.lock().unwrap();
|
|
||||||
items.insert(item.id.clone(), item);
|
|
||||||
}
|
|
||||||
self.maybe_evict();
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
|
|
||||||
let items = self.items.lock().unwrap();
|
|
||||||
Ok(items.get(id).cloned())
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
|
|
||||||
let mut items = self.items.lock().unwrap();
|
|
||||||
items.remove(id);
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
|
|
||||||
let items = self.items.lock().unwrap();
|
|
||||||
let mut result: Vec<MemoryItem> = items
|
|
||||||
.values()
|
|
||||||
.filter(|v| match &filter.prefix {
|
|
||||||
Some(p) => v.id.starts_with(p),
|
|
||||||
None => true,
|
|
||||||
})
|
|
||||||
.filter(|v| match filter.since {
|
|
||||||
Some(t) => v.created_at > t,
|
|
||||||
None => true,
|
|
||||||
})
|
|
||||||
.cloned()
|
|
||||||
.collect();
|
|
||||||
// 按 created_at 升序排列(最旧在前)
|
|
||||||
result.sort_by_key(|v| v.created_at);
|
|
||||||
// 应用 offset
|
|
||||||
if let Some(offset) = filter.offset {
|
|
||||||
if offset < result.len() {
|
|
||||||
result.drain(..offset);
|
|
||||||
} else {
|
|
||||||
result.clear();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// 应用 limit
|
|
||||||
if let Some(limit) = filter.limit {
|
|
||||||
result.truncate(limit);
|
|
||||||
}
|
|
||||||
Ok(result)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
use time::OffsetDateTime;
|
|
||||||
|
|
||||||
fn make_item(id: &str) -> MemoryItem {
|
|
||||||
MemoryItem {
|
|
||||||
id: id.to_string(),
|
|
||||||
content: format!("content-{id}"),
|
|
||||||
metadata: serde_json::json!({}),
|
|
||||||
created_at: OffsetDateTime::now_utc(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn save_get_delete_list() {
|
|
||||||
let store = InMemoryStore::new();
|
|
||||||
store.save(make_item("a")).await.unwrap();
|
|
||||||
store.save(make_item("b")).await.unwrap();
|
|
||||||
|
|
||||||
let got = store.get("a").await.unwrap();
|
|
||||||
assert!(got.is_some());
|
|
||||||
assert_eq!(got.unwrap().id, "a");
|
|
||||||
|
|
||||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
|
||||||
assert_eq!(list.len(), 2);
|
|
||||||
|
|
||||||
store.delete("a").await.unwrap();
|
|
||||||
assert!(store.get("a").await.unwrap().is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn save_is_upsert() {
|
|
||||||
let store = InMemoryStore::new();
|
|
||||||
store.save(make_item("a")).await.unwrap();
|
|
||||||
let mut item = make_item("a");
|
|
||||||
item.content = "updated".to_string();
|
|
||||||
store.save(item).await.unwrap();
|
|
||||||
let got = store.get("a").await.unwrap().unwrap();
|
|
||||||
assert_eq!(got.content, "updated");
|
|
||||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
|
||||||
assert_eq!(list.len(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn list_with_prefix_and_limit() {
|
|
||||||
let store = InMemoryStore::new();
|
|
||||||
store.save(make_item("foo_a")).await.unwrap();
|
|
||||||
store.save(make_item("foo_b")).await.unwrap();
|
|
||||||
store.save(make_item("bar_a")).await.unwrap();
|
|
||||||
|
|
||||||
let filter = MemoryFilter {
|
|
||||||
prefix: Some("foo_".to_string()),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let list = store.list(&filter).await.unwrap();
|
|
||||||
assert_eq!(list.len(), 2);
|
|
||||||
|
|
||||||
let filter = MemoryFilter {
|
|
||||||
prefix: Some("foo_".to_string()),
|
|
||||||
limit: Some(1),
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let list = store.list(&filter).await.unwrap();
|
|
||||||
assert_eq!(list.len(), 1);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn capacity_eviction() {
|
|
||||||
// 强制每次写入都检查
|
|
||||||
let eviction = EvictionConfig {
|
|
||||||
policy: EvictionPolicy::Capacity { max_items: 2 },
|
|
||||||
check_interval: 1,
|
|
||||||
};
|
|
||||||
let store = InMemoryStore::with_eviction(eviction);
|
|
||||||
// 第一条和第二条共存
|
|
||||||
store.save(make_item("a")).await.unwrap();
|
|
||||||
store.save(make_item("b")).await.unwrap();
|
|
||||||
// 第三条写入触发淘汰:a 或 b 之一被淘汰
|
|
||||||
store.save(make_item("c")).await.unwrap();
|
|
||||||
|
|
||||||
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
|
||||||
assert_eq!(list.len(), 2);
|
|
||||||
// 留下的应该是 b 和 c(最新的两个)
|
|
||||||
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
|
||||||
assert!(ids.contains(&"b"));
|
|
||||||
assert!(ids.contains(&"c"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn ttl_eviction() {
|
|
||||||
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
|
|
||||||
// 改用小 TTL + 睡眠:先 save a,sleep,save 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);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -0,0 +1,366 @@
|
|||||||
|
//! 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
|
||||||
|
|
||||||
|
use std::collections::HashMap;
|
||||||
|
use std::sync::Mutex;
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use time::OffsetDateTime;
|
||||||
|
|
||||||
|
use crate::memory::error::MemoryError;
|
||||||
|
use crate::memory::store::{EvictionConfig, EvictionPolicy, MemoryStore};
|
||||||
|
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||||
|
|
||||||
|
/// 进程内默认实现 —— 基于 HashMap + Mutex,纯内存。
|
||||||
|
pub struct InMemoryStore {
|
||||||
|
items: Mutex<HashMap<String, MemoryItem>>,
|
||||||
|
eviction: EvictionConfig,
|
||||||
|
/// 自上次淘汰检查以来的写入次数。
|
||||||
|
writes_since_check: Mutex<usize>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl InMemoryStore {
|
||||||
|
/// 创建一个无淘汰策略的 InMemoryStore。
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self {
|
||||||
|
items: Mutex::new(HashMap::new()),
|
||||||
|
eviction: EvictionConfig::default(),
|
||||||
|
writes_since_check: Mutex::new(0),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 创建一个带淘汰配置的 InMemoryStore。
|
||||||
|
pub fn with_eviction(eviction: EvictionConfig) -> Self {
|
||||||
|
Self {
|
||||||
|
items: Mutex::new(HashMap::new()),
|
||||||
|
eviction,
|
||||||
|
writes_since_check: Mutex::new(0),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn maybe_evict(&self) {
|
||||||
|
// 不使用 .lock().await 跨点,先取计数判断是否需要淘汰
|
||||||
|
let should_check = {
|
||||||
|
let mut counter = self.writes_since_check.lock().unwrap();
|
||||||
|
*counter += 1;
|
||||||
|
if *counter >= self.eviction.check_interval {
|
||||||
|
*counter = 0;
|
||||||
|
true
|
||||||
|
} else {
|
||||||
|
false
|
||||||
|
}
|
||||||
|
};
|
||||||
|
if !should_check {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
let policy = self.eviction.policy.clone();
|
||||||
|
match policy {
|
||||||
|
EvictionPolicy::None => {}
|
||||||
|
EvictionPolicy::Ttl { ttl_secs } => {
|
||||||
|
let cutoff = OffsetDateTime::now_utc() - time::Duration::seconds(ttl_secs as i64);
|
||||||
|
let mut items = self.items.lock().unwrap();
|
||||||
|
items.retain(|_, v| v.created_at > cutoff);
|
||||||
|
}
|
||||||
|
EvictionPolicy::Capacity { max_items } => {
|
||||||
|
let mut items = self.items.lock().unwrap();
|
||||||
|
if items.len() > max_items {
|
||||||
|
let mut vec: Vec<_> = items.drain().collect();
|
||||||
|
// O(n) 部分排序:保留 created_at 最大的 max_items 个
|
||||||
|
vec.select_nth_unstable_by(max_items, |a, b| {
|
||||||
|
b.1.created_at.cmp(&a.1.created_at)
|
||||||
|
});
|
||||||
|
vec.truncate(max_items);
|
||||||
|
*items = vec.into_iter().collect();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for InMemoryStore {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl MemoryStore for InMemoryStore {
|
||||||
|
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
|
||||||
|
{
|
||||||
|
let mut items = self.items.lock().unwrap();
|
||||||
|
items.insert(item.id.clone(), item);
|
||||||
|
}
|
||||||
|
self.maybe_evict();
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
|
||||||
|
let items = self.items.lock().unwrap();
|
||||||
|
Ok(items.get(id).cloned())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
|
||||||
|
let mut items = self.items.lock().unwrap();
|
||||||
|
items.remove(id);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
|
||||||
|
let items = self.items.lock().unwrap();
|
||||||
|
let mut result: Vec<MemoryItem> = items
|
||||||
|
.values()
|
||||||
|
.filter(|v| match &filter.prefix {
|
||||||
|
Some(p) => v.id.starts_with(p),
|
||||||
|
None => true,
|
||||||
|
})
|
||||||
|
.filter(|v| match filter.since {
|
||||||
|
Some(t) => v.created_at > t,
|
||||||
|
None => true,
|
||||||
|
})
|
||||||
|
.cloned()
|
||||||
|
.collect();
|
||||||
|
// 按 created_at 升序排列(最旧在前)
|
||||||
|
result.sort_by_key(|v| v.created_at);
|
||||||
|
// 应用 offset
|
||||||
|
if let Some(offset) = filter.offset {
|
||||||
|
if offset < result.len() {
|
||||||
|
result.drain(..offset);
|
||||||
|
} else {
|
||||||
|
result.clear();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// 应用 limit
|
||||||
|
if let Some(limit) = filter.limit {
|
||||||
|
result.truncate(limit);
|
||||||
|
}
|
||||||
|
Ok(result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use time::OffsetDateTime;
|
||||||
|
|
||||||
|
fn make_item(id: &str) -> MemoryItem {
|
||||||
|
MemoryItem {
|
||||||
|
id: id.to_string(),
|
||||||
|
content: format!("content-{id}"),
|
||||||
|
metadata: serde_json::json!({}),
|
||||||
|
created_at: OffsetDateTime::now_utc(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn save_get_delete_list() {
|
||||||
|
let store = InMemoryStore::new();
|
||||||
|
store.save(make_item("a")).await.unwrap();
|
||||||
|
store.save(make_item("b")).await.unwrap();
|
||||||
|
|
||||||
|
let got = store.get("a").await.unwrap();
|
||||||
|
assert!(got.is_some());
|
||||||
|
assert_eq!(got.unwrap().id, "a");
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 2);
|
||||||
|
|
||||||
|
store.delete("a").await.unwrap();
|
||||||
|
assert!(store.get("a").await.unwrap().is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn save_is_upsert() {
|
||||||
|
let store = InMemoryStore::new();
|
||||||
|
store.save(make_item("a")).await.unwrap();
|
||||||
|
let mut item = make_item("a");
|
||||||
|
item.content = "updated".to_string();
|
||||||
|
store.save(item).await.unwrap();
|
||||||
|
let got = store.get("a").await.unwrap().unwrap();
|
||||||
|
assert_eq!(got.content, "updated");
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn list_with_prefix_and_limit() {
|
||||||
|
let store = InMemoryStore::new();
|
||||||
|
store.save(make_item("foo_a")).await.unwrap();
|
||||||
|
store.save(make_item("foo_b")).await.unwrap();
|
||||||
|
store.save(make_item("bar_a")).await.unwrap();
|
||||||
|
|
||||||
|
let filter = MemoryFilter {
|
||||||
|
prefix: Some("foo_".to_string()),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let list = store.list(&filter).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 2);
|
||||||
|
|
||||||
|
let filter = MemoryFilter {
|
||||||
|
prefix: Some("foo_".to_string()),
|
||||||
|
limit: Some(1),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let list = store.list(&filter).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn capacity_eviction() {
|
||||||
|
// 强制每次写入都检查
|
||||||
|
let eviction = EvictionConfig {
|
||||||
|
policy: EvictionPolicy::Capacity { max_items: 2 },
|
||||||
|
check_interval: 1,
|
||||||
|
};
|
||||||
|
let store = InMemoryStore::with_eviction(eviction);
|
||||||
|
// 第一条和第二条共存
|
||||||
|
store.save(make_item("a")).await.unwrap();
|
||||||
|
store.save(make_item("b")).await.unwrap();
|
||||||
|
// 第三条写入触发淘汰:a 或 b 之一被淘汰
|
||||||
|
store.save(make_item("c")).await.unwrap();
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 2);
|
||||||
|
// 留下的应该是 b 和 c(最新的两个)
|
||||||
|
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||||
|
assert!(ids.contains(&"b"));
|
||||||
|
assert!(ids.contains(&"c"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn ttl_eviction() {
|
||||||
|
// TTL 设为 0 会立即过期,但我们想保留 "a" 等待 "b" 写入后被淘汰。
|
||||||
|
// 改用小 TTL + 睡眠:先 save a,sleep,save b 时 a 已过期被淘汰。
|
||||||
|
let eviction = EvictionConfig {
|
||||||
|
policy: EvictionPolicy::Ttl { ttl_secs: 1 },
|
||||||
|
check_interval: 1,
|
||||||
|
};
|
||||||
|
let store = InMemoryStore::with_eviction(eviction);
|
||||||
|
store.save(make_item("a")).await.unwrap();
|
||||||
|
// 等待超过 1 秒
|
||||||
|
std::thread::sleep(std::time::Duration::from_millis(1100));
|
||||||
|
// 触发淘汰:a 已超过 ttl_secs=1,应被淘汰
|
||||||
|
store.save(make_item("b")).await.unwrap();
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
// 由于 ttl_secs=1,且 b 刚写入,可能刚好处于临界值。
|
||||||
|
// 我们只断言 list 不包含 "a" 即可。
|
||||||
|
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||||
|
assert!(
|
||||||
|
!ids.contains(&"a"),
|
||||||
|
"expected 'a' to be evicted, but found in {ids:?}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn none_policy_no_eviction() {
|
||||||
|
let eviction = EvictionConfig {
|
||||||
|
policy: EvictionPolicy::None,
|
||||||
|
check_interval: 1,
|
||||||
|
};
|
||||||
|
let store = InMemoryStore::with_eviction(eviction);
|
||||||
|
for i in 0..100 {
|
||||||
|
store.save(make_item(&format!("item_{i}"))).await.unwrap();
|
||||||
|
}
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 100);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== Phase 11 Step 11.3 并发测试 =====
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_writers_max_pressure() {
|
||||||
|
use std::sync::Arc;
|
||||||
|
let store = Arc::new(InMemoryStore::new());
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for i in 0..100 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let id = format!("concurrent_{i}");
|
||||||
|
s.save(make_item(&id)).await.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 100);
|
||||||
|
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
|
||||||
|
ids.sort();
|
||||||
|
ids.dedup();
|
||||||
|
assert_eq!(ids.len(), 100);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_mixed_read_write() {
|
||||||
|
use std::sync::Arc;
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
let store = Arc::new(InMemoryStore::new());
|
||||||
|
|
||||||
|
// 预热 20 条
|
||||||
|
for i in 0..20 {
|
||||||
|
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
|
||||||
|
// 5 个写者
|
||||||
|
for w in 0..5 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let mut i = 0;
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let id = format!("writer{w}_item{i}");
|
||||||
|
s.save(make_item(&id)).await.unwrap();
|
||||||
|
i += 1;
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
// 5 个读者
|
||||||
|
for _ in 0..5 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let _ = s.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_capacity_eviction() {
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
let eviction = EvictionConfig {
|
||||||
|
policy: EvictionPolicy::Capacity { max_items: 10 },
|
||||||
|
check_interval: 1,
|
||||||
|
};
|
||||||
|
let store = Arc::new(InMemoryStore::with_eviction(eviction));
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for i in 0..15 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
s.save(make_item(&format!("item_{i}"))).await.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
// 写者全部完成后必 ≤ max_items(部分路径上可能短暂 >10 但全部完成时应 ≤10)
|
||||||
|
assert!(
|
||||||
|
list.len() <= 10,
|
||||||
|
"expected <= 10 items after all writers done, got {}",
|
||||||
|
list.len()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,623 @@
|
|||||||
|
//! SqliteStore —— 基于 rusqlite 的持久化 MemoryStore 实现。
|
||||||
|
//!
|
||||||
|
//! 单进程独享、写入串行化(WAL + Mutex),适合本地 Agent 长期持久化场景。
|
||||||
|
|
||||||
|
use std::path::Path;
|
||||||
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use rusqlite::{params, params_from_iter, Connection, ErrorCode};
|
||||||
|
use time::format_description::well_known::Rfc3339;
|
||||||
|
use time::OffsetDateTime;
|
||||||
|
use tracing::{debug, error, instrument, warn};
|
||||||
|
|
||||||
|
use crate::memory::error::MemoryError;
|
||||||
|
use crate::memory::store::MemoryStore;
|
||||||
|
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||||
|
|
||||||
|
const INITIAL_USER_VERSION: i64 = 1;
|
||||||
|
const BUSY_TIMEOUT_MS: i64 = 5000;
|
||||||
|
const WAL_AUTOCHECKPOINT_PAGES: i64 = 1000;
|
||||||
|
|
||||||
|
/// SQLite 持久化后端的 MemoryStore 实现。
|
||||||
|
///
|
||||||
|
/// 设计要点:
|
||||||
|
/// - 单进程独享:`Arc<Mutex<Connection>>` 串行化所有 IO
|
||||||
|
/// - WAL 模式 + `synchronous=NORMAL` 兼顾崩溃安全与吞吐
|
||||||
|
/// - `created_at` 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
|
||||||
|
/// - 所有 IO 通过 `tokio::task::spawn_blocking` 卸载到阻塞线程池
|
||||||
|
pub struct SqliteStore {
|
||||||
|
conn: Arc<Mutex<Connection>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl SqliteStore {
|
||||||
|
/// 打开或创建一个 SQLite 数据库。
|
||||||
|
///
|
||||||
|
/// - `path = ":memory:"` 使用内存数据库(测试场景)
|
||||||
|
/// - 其他路径:自动创建父目录;文件已存在则附加打开
|
||||||
|
/// - 启动时执行 `migrate()`,失败立即返回错误
|
||||||
|
#[instrument(skip(path), fields(path = %path.as_ref().display()))]
|
||||||
|
pub fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
|
||||||
|
let path_ref = path.as_ref();
|
||||||
|
let path_str = path_ref.to_string_lossy();
|
||||||
|
|
||||||
|
let conn = if path_str == ":memory:" {
|
||||||
|
Connection::open_in_memory()
|
||||||
|
} else {
|
||||||
|
if let Some(parent) = path_ref.parent()
|
||||||
|
&& !parent.as_os_str().is_empty()
|
||||||
|
{
|
||||||
|
std::fs::create_dir_all(parent).map_err(|e| {
|
||||||
|
MemoryError::Storage(format!(
|
||||||
|
"创建数据库父目录失败 ({}): {}",
|
||||||
|
parent.display(),
|
||||||
|
e
|
||||||
|
))
|
||||||
|
})?;
|
||||||
|
}
|
||||||
|
Connection::open(path_ref)
|
||||||
|
}
|
||||||
|
.map_err(|e| map_sqlite_error(e, "打开数据库"))?;
|
||||||
|
|
||||||
|
migrate(&conn)?;
|
||||||
|
Ok(Self {
|
||||||
|
conn: Arc::new(Mutex::new(conn)),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl MemoryStore for SqliteStore {
|
||||||
|
#[instrument(skip(self, item), fields(id = %item.id))]
|
||||||
|
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
|
||||||
|
let conn = Arc::clone(&self.conn);
|
||||||
|
let created_at_str = item
|
||||||
|
.created_at
|
||||||
|
.to_offset(time::UtcOffset::UTC)
|
||||||
|
.format(&Rfc3339)
|
||||||
|
.map_err(|e| MemoryError::Serialization(format!("format created_at: {e}")))?;
|
||||||
|
let metadata_str = serde_json::to_string(&item.metadata)
|
||||||
|
.map_err(|e| MemoryError::Serialization(format!("serialize metadata: {e}")))?;
|
||||||
|
let id = item.id;
|
||||||
|
let content = item.content;
|
||||||
|
|
||||||
|
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
|
||||||
|
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||||
|
conn.execute(
|
||||||
|
"INSERT INTO memory_items (id, content, metadata, created_at) \
|
||||||
|
VALUES (?1, ?2, ?3, ?4) \
|
||||||
|
ON CONFLICT(id) DO UPDATE SET \
|
||||||
|
content=excluded.content, \
|
||||||
|
metadata=excluded.metadata, \
|
||||||
|
created_at=excluded.created_at",
|
||||||
|
params![id, content, metadata_str, created_at_str],
|
||||||
|
)
|
||||||
|
.map_err(|e| map_sqlite_error(e, "保存记忆"))?;
|
||||||
|
Ok(())
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||||
|
}
|
||||||
|
|
||||||
|
#[instrument(skip(self, id))]
|
||||||
|
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
|
||||||
|
let conn = Arc::clone(&self.conn);
|
||||||
|
let id_owned = id.to_string();
|
||||||
|
|
||||||
|
tokio::task::spawn_blocking(move || -> Result<Option<MemoryItem>, MemoryError> {
|
||||||
|
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||||
|
let mut stmt = conn
|
||||||
|
.prepare("SELECT id, content, metadata, created_at FROM memory_items WHERE id = ?1")
|
||||||
|
.map_err(|e| map_sqlite_error(e, "prepare get"))?;
|
||||||
|
let mut rows = stmt
|
||||||
|
.query_map(params![id_owned], row_to_item)
|
||||||
|
.map_err(|e| map_sqlite_error(e, "query get"))?;
|
||||||
|
match rows.next() {
|
||||||
|
None => Ok(None),
|
||||||
|
Some(row) => row
|
||||||
|
.map(Some)
|
||||||
|
.map_err(|e| map_sqlite_error(e, "decode row")),
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||||
|
}
|
||||||
|
|
||||||
|
#[instrument(skip(self, id))]
|
||||||
|
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
|
||||||
|
let conn = Arc::clone(&self.conn);
|
||||||
|
let id_owned = id.to_string();
|
||||||
|
|
||||||
|
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
|
||||||
|
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM memory_items WHERE id = ?1",
|
||||||
|
params![id_owned],
|
||||||
|
)
|
||||||
|
.map_err(|e| map_sqlite_error(e, "delete"))?;
|
||||||
|
Ok(())
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||||
|
}
|
||||||
|
|
||||||
|
#[instrument(skip(self, filter))]
|
||||||
|
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
|
||||||
|
let mut sql = String::from(
|
||||||
|
"SELECT id, content, metadata, created_at FROM memory_items WHERE 1=1",
|
||||||
|
);
|
||||||
|
let mut param_values: Vec<String> = Vec::new();
|
||||||
|
let mut ph_idx = 0usize;
|
||||||
|
|
||||||
|
if filter.prefix.is_some() {
|
||||||
|
ph_idx += 1;
|
||||||
|
sql.push_str(&format!(" AND id LIKE ?{ph_idx} || '%'"));
|
||||||
|
}
|
||||||
|
if filter.since.is_some() {
|
||||||
|
ph_idx += 1;
|
||||||
|
sql.push_str(&format!(" AND created_at > ?{ph_idx}"));
|
||||||
|
}
|
||||||
|
// ORDER BY created_at ASC(按时间升序,最旧在前)
|
||||||
|
sql.push_str(" ORDER BY created_at ASC");
|
||||||
|
|
||||||
|
let limit_sql: String = match (filter.limit, filter.offset) {
|
||||||
|
(Some(_), Some(_)) => {
|
||||||
|
ph_idx += 1;
|
||||||
|
let limit_p = ph_idx;
|
||||||
|
ph_idx += 1;
|
||||||
|
let offset_p = ph_idx;
|
||||||
|
format!(" LIMIT ?{limit_p} OFFSET ?{offset_p}")
|
||||||
|
}
|
||||||
|
(Some(_), None) => {
|
||||||
|
ph_idx += 1;
|
||||||
|
let limit_p = ph_idx;
|
||||||
|
format!(" LIMIT ?{limit_p}")
|
||||||
|
}
|
||||||
|
(None, Some(_)) => {
|
||||||
|
// SQLite 中 LIMIT -1 表示无限制
|
||||||
|
ph_idx += 1;
|
||||||
|
let offset_p = ph_idx;
|
||||||
|
format!(" LIMIT -1 OFFSET ?{offset_p}")
|
||||||
|
}
|
||||||
|
(None, None) => String::new(),
|
||||||
|
};
|
||||||
|
sql.push_str(&limit_sql);
|
||||||
|
|
||||||
|
if let Some(p) = &filter.prefix {
|
||||||
|
param_values.push(p.clone());
|
||||||
|
}
|
||||||
|
if let Some(t) = filter.since {
|
||||||
|
let s = t
|
||||||
|
.to_offset(time::UtcOffset::UTC)
|
||||||
|
.format(&Rfc3339)
|
||||||
|
.map_err(|e| MemoryError::Serialization(format!("format since: {e}")))?;
|
||||||
|
param_values.push(s);
|
||||||
|
}
|
||||||
|
if let Some(l) = filter.limit {
|
||||||
|
param_values.push(l.to_string());
|
||||||
|
}
|
||||||
|
if let Some(o) = filter.offset {
|
||||||
|
param_values.push(o.to_string());
|
||||||
|
}
|
||||||
|
|
||||||
|
let conn = Arc::clone(&self.conn);
|
||||||
|
let sql_owned = sql;
|
||||||
|
let param_values_owned = param_values;
|
||||||
|
|
||||||
|
tokio::task::spawn_blocking(move || -> Result<Vec<MemoryItem>, MemoryError> {
|
||||||
|
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
|
||||||
|
let mut stmt = conn
|
||||||
|
.prepare(&sql_owned)
|
||||||
|
.map_err(|e| map_sqlite_error(e, "list prepare"))?;
|
||||||
|
let params_iter: Vec<&dyn rusqlite::ToSql> = param_values_owned
|
||||||
|
.iter()
|
||||||
|
.map(|s| s as &dyn rusqlite::ToSql)
|
||||||
|
.collect();
|
||||||
|
let rows = stmt
|
||||||
|
.query_map(params_from_iter(params_iter), row_to_item)
|
||||||
|
.map_err(|e| map_sqlite_error(e, "list query"))?;
|
||||||
|
let mut result = Vec::new();
|
||||||
|
for row in rows {
|
||||||
|
result.push(row.map_err(|e| map_sqlite_error(e, "list row"))?);
|
||||||
|
}
|
||||||
|
debug!(count = result.len(), "SqliteStore::list 完成");
|
||||||
|
Ok(result)
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn row_to_item(row: &rusqlite::Row<'_>) -> Result<MemoryItem, rusqlite::Error> {
|
||||||
|
let id: String = row.get(0)?;
|
||||||
|
let content: String = row.get(1)?;
|
||||||
|
let metadata_str: String = row.get(2)?;
|
||||||
|
let created_at_str: String = row.get(3)?;
|
||||||
|
|
||||||
|
let metadata: serde_json::Value = serde_json::from_str(&metadata_str).map_err(|e| {
|
||||||
|
rusqlite::Error::FromSqlConversionFailure(2, rusqlite::types::Type::Text, Box::new(e))
|
||||||
|
})?;
|
||||||
|
let created_at = OffsetDateTime::parse(&created_at_str, &Rfc3339).map_err(|e| {
|
||||||
|
rusqlite::Error::FromSqlConversionFailure(3, rusqlite::types::Type::Text, Box::new(e))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
Ok(MemoryItem {
|
||||||
|
id,
|
||||||
|
content,
|
||||||
|
metadata,
|
||||||
|
created_at: created_at.to_offset(time::UtcOffset::UTC),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn migrate(conn: &Connection) -> Result<(), MemoryError> {
|
||||||
|
conn.pragma_update(None, "journal_mode", "WAL")
|
||||||
|
.map_err(|e| map_sqlite_error(e, "PRAGMA journal_mode"))?;
|
||||||
|
conn.pragma_update(None, "synchronous", "NORMAL")
|
||||||
|
.map_err(|e| map_sqlite_error(e, "PRAGMA synchronous"))?;
|
||||||
|
conn.execute_batch(&format!("PRAGMA busy_timeout = {BUSY_TIMEOUT_MS};"))
|
||||||
|
.map_err(|e| map_sqlite_error(e, "PRAGMA busy_timeout"))?;
|
||||||
|
conn.execute_batch(&format!(
|
||||||
|
"PRAGMA wal_autocheckpoint = {WAL_AUTOCHECKPOINT_PAGES};"
|
||||||
|
))
|
||||||
|
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_autocheckpoint"))?;
|
||||||
|
|
||||||
|
let version: i64 = conn
|
||||||
|
.query_row("PRAGMA user_version", [], |row| row.get(0))
|
||||||
|
.map_err(|e| map_sqlite_error(e, "PRAGMA user_version"))?;
|
||||||
|
|
||||||
|
if version < INITIAL_USER_VERSION {
|
||||||
|
conn.execute_batch(
|
||||||
|
"CREATE TABLE IF NOT EXISTS memory_items (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
content TEXT NOT NULL,
|
||||||
|
metadata TEXT NOT NULL DEFAULT '{}',
|
||||||
|
created_at TEXT NOT NULL
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_memory_items_created_at
|
||||||
|
ON memory_items(created_at);
|
||||||
|
PRAGMA user_version = 1;",
|
||||||
|
)
|
||||||
|
.map_err(|e| map_sqlite_error(e, "create schema v1"))?;
|
||||||
|
}
|
||||||
|
|
||||||
|
let check_result: String = conn
|
||||||
|
.query_row("PRAGMA quick_check", [], |row| row.get(0))
|
||||||
|
.map_err(|e| map_sqlite_error(e, "PRAGMA quick_check"))?;
|
||||||
|
if check_result != "ok" {
|
||||||
|
error!(result = %check_result, "数据库文件 quick_check 失败");
|
||||||
|
return Err(MemoryError::Storage(format!(
|
||||||
|
"数据库文件损坏: {check_result}"
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
|
||||||
|
conn.execute_batch("PRAGMA wal_checkpoint(TRUNCATE);")
|
||||||
|
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_checkpoint"))?;
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn map_sqlite_error(e: rusqlite::Error, ctx: &str) -> MemoryError {
|
||||||
|
match &e {
|
||||||
|
rusqlite::Error::SqliteFailure(err, _) => match err.code {
|
||||||
|
ErrorCode::ConstraintViolation => MemoryError::InvalidInput(format!("{ctx}: {e}")),
|
||||||
|
ErrorCode::DatabaseBusy | ErrorCode::DatabaseLocked => {
|
||||||
|
warn!("SQLite 忙: {e}");
|
||||||
|
MemoryError::Storage(format!("{ctx}: {e}"))
|
||||||
|
}
|
||||||
|
_ => MemoryError::Storage(format!("{ctx}: {e}")),
|
||||||
|
},
|
||||||
|
rusqlite::Error::InvalidQuery
|
||||||
|
| rusqlite::Error::InvalidParameterName(_)
|
||||||
|
| rusqlite::Error::InvalidColumnIndex(_)
|
||||||
|
| rusqlite::Error::InvalidColumnName(_) => {
|
||||||
|
MemoryError::InvalidInput(format!("{ctx}: {e}"))
|
||||||
|
}
|
||||||
|
rusqlite::Error::FromSqlConversionFailure(_, _, _)
|
||||||
|
| rusqlite::Error::ToSqlConversionFailure(_) => {
|
||||||
|
MemoryError::Serialization(format!("{ctx}: {e}"))
|
||||||
|
}
|
||||||
|
_ => MemoryError::Storage(format!("{ctx}: {e}")),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use crate::memory::store::InMemoryStore;
|
||||||
|
use std::sync::Arc;
|
||||||
|
use tempfile::TempDir;
|
||||||
|
use time::OffsetDateTime;
|
||||||
|
|
||||||
|
fn make_item(id: &str) -> MemoryItem {
|
||||||
|
MemoryItem {
|
||||||
|
id: id.to_string(),
|
||||||
|
content: format!("content-{id}"),
|
||||||
|
metadata: serde_json::json!({"id_key": id}),
|
||||||
|
created_at: OffsetDateTime::now_utc(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn make_item_at(id: &str, when: OffsetDateTime) -> MemoryItem {
|
||||||
|
MemoryItem {
|
||||||
|
id: id.to_string(),
|
||||||
|
content: format!("content-{id}"),
|
||||||
|
metadata: serde_json::json!({}),
|
||||||
|
created_at: when,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn crud_basic() {
|
||||||
|
let store = SqliteStore::open(":memory:").unwrap();
|
||||||
|
store.save(make_item("a")).await.unwrap();
|
||||||
|
store.save(make_item("b")).await.unwrap();
|
||||||
|
|
||||||
|
let got_a = store.get("a").await.unwrap();
|
||||||
|
assert!(got_a.is_some());
|
||||||
|
assert_eq!(got_a.unwrap().id, "a");
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 2);
|
||||||
|
|
||||||
|
store.delete("a").await.unwrap();
|
||||||
|
assert!(store.get("a").await.unwrap().is_none());
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 1);
|
||||||
|
assert_eq!(list[0].id, "b");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn save_is_upsert() {
|
||||||
|
let store = SqliteStore::open(":memory:").unwrap();
|
||||||
|
store.save(make_item("a")).await.unwrap();
|
||||||
|
let mut item = make_item("a");
|
||||||
|
item.content = "updated".to_string();
|
||||||
|
item.metadata = serde_json::json!({"rev": 2});
|
||||||
|
let original_created_at = item.created_at;
|
||||||
|
store.save(item).await.unwrap();
|
||||||
|
|
||||||
|
let got = store.get("a").await.unwrap().unwrap();
|
||||||
|
assert_eq!(got.content, "updated");
|
||||||
|
assert_eq!(got.metadata["rev"], serde_json::json!(2));
|
||||||
|
// created_at 保持调用方传入值
|
||||||
|
assert_eq!(got.created_at, original_created_at);
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn list_with_prefix() {
|
||||||
|
let store = SqliteStore::open(":memory:").unwrap();
|
||||||
|
store.save(make_item("foo_a")).await.unwrap();
|
||||||
|
store.save(make_item("foo_b")).await.unwrap();
|
||||||
|
store.save(make_item("bar_a")).await.unwrap();
|
||||||
|
|
||||||
|
let filter = MemoryFilter {
|
||||||
|
prefix: Some("foo_".to_string()),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let list = store.list(&filter).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 2);
|
||||||
|
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||||
|
assert!(ids.contains(&"foo_a"));
|
||||||
|
assert!(ids.contains(&"foo_b"));
|
||||||
|
assert!(!ids.contains(&"bar_a"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn list_with_since_filter() {
|
||||||
|
let store = SqliteStore::open(":memory:").unwrap();
|
||||||
|
let t0 = OffsetDateTime::now_utc();
|
||||||
|
store
|
||||||
|
.save(make_item_at("early", t0 - time::Duration::seconds(60)))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
store.save(make_item_at("middle", t0)).await.unwrap();
|
||||||
|
store
|
||||||
|
.save(make_item_at("late", t0 + time::Duration::seconds(60)))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let filter = MemoryFilter {
|
||||||
|
since: Some(t0 - time::Duration::seconds(1)),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let list = store.list(&filter).await.unwrap();
|
||||||
|
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||||
|
assert_eq!(list.len(), 2);
|
||||||
|
assert!(ids.contains(&"middle"));
|
||||||
|
assert!(ids.contains(&"late"));
|
||||||
|
assert!(!ids.contains(&"early"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn list_with_offset_and_limit() {
|
||||||
|
let store = SqliteStore::open(":memory:").unwrap();
|
||||||
|
// 写入 5 条时间递增的记录
|
||||||
|
let base = OffsetDateTime::now_utc() - time::Duration::seconds(5);
|
||||||
|
for i in 0..5 {
|
||||||
|
let mut item = make_item(&format!("item_{i}"));
|
||||||
|
item.created_at = base + time::Duration::seconds(i);
|
||||||
|
store.save(item).await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
// offset=1, limit=2 -> item_1, item_2
|
||||||
|
let filter = MemoryFilter {
|
||||||
|
offset: Some(1),
|
||||||
|
limit: Some(2),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let list = store.list(&filter).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 2);
|
||||||
|
assert_eq!(list[0].id, "item_1");
|
||||||
|
assert_eq!(list[1].id, "item_2");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_writers_no_data_loss() {
|
||||||
|
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for w in 0..10 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
for i in 0..10 {
|
||||||
|
let id = format!("w{w}_i{i}");
|
||||||
|
s.save(make_item(&id)).await.unwrap();
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 100);
|
||||||
|
// 验证所有 id 唯一
|
||||||
|
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
|
||||||
|
ids.sort();
|
||||||
|
ids.dedup();
|
||||||
|
assert_eq!(ids.len(), 100);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn persistence_round_trip() {
|
||||||
|
let dir = TempDir::new().unwrap();
|
||||||
|
let path = dir.path().join("memory.db");
|
||||||
|
|
||||||
|
// 阶段 1:写入 3 条
|
||||||
|
{
|
||||||
|
let store = SqliteStore::open(&path).unwrap();
|
||||||
|
store.save(make_item("alpha")).await.unwrap();
|
||||||
|
store.save(make_item("beta")).await.unwrap();
|
||||||
|
store.save(make_item("gamma")).await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
// 阶段 2:重新打开,验证数据完整
|
||||||
|
{
|
||||||
|
let store = SqliteStore::open(&path).unwrap();
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 3);
|
||||||
|
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
|
||||||
|
assert!(ids.contains(&"alpha"));
|
||||||
|
assert!(ids.contains(&"beta"));
|
||||||
|
assert!(ids.contains(&"gamma"));
|
||||||
|
|
||||||
|
// 单条读回
|
||||||
|
let got = store.get("beta").await.unwrap().unwrap();
|
||||||
|
assert_eq!(got.content, "content-beta");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn open_invalid_path_returns_error() {
|
||||||
|
// 路径指向已存在的目录而非文件,open 应失败
|
||||||
|
let dir = TempDir::new().unwrap();
|
||||||
|
match SqliteStore::open(dir.path()) {
|
||||||
|
Err(MemoryError::Storage(_)) => {}
|
||||||
|
Err(other) => panic!("expected Storage error, got {other:?}"),
|
||||||
|
Ok(_) => panic!("expected error when opening a directory as database"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn trait_object_compatibility() {
|
||||||
|
// ponytail: 回归验证 SqliteStore 可作为 Arc<dyn MemoryStore> 与 InMemoryStore 互换
|
||||||
|
// 所有现有消费者(Conversation / Knowledge / Retriever / SessionMemory)均通过 trait object 引用,
|
||||||
|
// 此测试确保 trait 接口契约在 SqliteStore 上同样成立。
|
||||||
|
let sqlite: Arc<dyn MemoryStore> =
|
||||||
|
Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||||
|
let in_mem: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
|
||||||
|
let stores: Vec<Arc<dyn MemoryStore>> = vec![Arc::clone(&sqlite), Arc::clone(&in_mem)];
|
||||||
|
for store in &stores {
|
||||||
|
store.save(make_item("x")).await.unwrap();
|
||||||
|
let got = store.get("x").await.unwrap();
|
||||||
|
assert_eq!(got.unwrap().id, "x");
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 1);
|
||||||
|
store.delete("x").await.unwrap();
|
||||||
|
assert!(store.get("x").await.unwrap().is_none());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== Phase 11 Step 11.3 并发测试 =====
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_writers_max_pressure() {
|
||||||
|
use std::time::Duration;
|
||||||
|
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for i in 0..100 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let id = format!("concurrent_{i}");
|
||||||
|
// 设置每次 save 的 per-call timeout —— busy_timeout=5000ms 应足够
|
||||||
|
match tokio::time::timeout(
|
||||||
|
Duration::from_secs(10),
|
||||||
|
s.save(make_item(&id)),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
Ok(res) => res.unwrap(),
|
||||||
|
Err(_) => panic!("save({id}) timed out under 100-way concurrency"),
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let list = store.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
assert_eq!(list.len(), 100);
|
||||||
|
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
|
||||||
|
ids.sort();
|
||||||
|
ids.dedup();
|
||||||
|
assert_eq!(ids.len(), 100);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_mixed_read_write() {
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
|
||||||
|
|
||||||
|
// 预热 20 条
|
||||||
|
for i in 0..20 {
|
||||||
|
store.save(make_item(&format!("seed_{i}"))).await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
|
||||||
|
// 5 个写者
|
||||||
|
for w in 0..5 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let mut i = 0;
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let id = format!("writer{w}_item{i}");
|
||||||
|
s.save(make_item(&id)).await.unwrap();
|
||||||
|
i += 1;
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
// 5 个读者
|
||||||
|
for _ in 0..5 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let _ = s.list(&MemoryFilter::default()).await.unwrap();
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,244 @@
|
|||||||
|
//! 语义向量检索抽象。
|
||||||
|
//!
|
||||||
|
//! 提供 [`VectorRetriever`] trait 定义与进程内引用实现 [`InMemoryVectorRetriever`]。
|
||||||
|
//! 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
|
||||||
|
|
||||||
|
use std::collections::HashMap;
|
||||||
|
use std::sync::Mutex;
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
|
||||||
|
use crate::memory::error::MemoryError;
|
||||||
|
|
||||||
|
/// 语义向量检索器抽象接口。
|
||||||
|
///
|
||||||
|
/// 下游可实现此 trait 以对接向量数据库(pgvector / qdrant / lancedb 等)。
|
||||||
|
/// 默认引用实现 [`InMemoryVectorRetriever`] 基于进程内 HashMap + 余弦相似度。
|
||||||
|
///
|
||||||
|
/// **稳定性**:实验性 API(v0.2.x),方法签名可能在 v0.3 中调整。
|
||||||
|
/// 若未来需要 `remove()` / `clear()` 等方法,将在此 trait 中追加(带默认实现)。
|
||||||
|
#[deprecated(since = "0.3.0", note = "请使用 memory::VectorStore")]
|
||||||
|
#[async_trait]
|
||||||
|
pub trait VectorRetriever: Send + Sync {
|
||||||
|
/// 将 `id` 对应的文本向量 `embeddings` 加入索引。
|
||||||
|
///
|
||||||
|
/// 重复调用同一 `id` 会覆盖已有向量。调用方负责保证 `embeddings` 维度
|
||||||
|
/// 与已索引向量一致——本 trait 不做维度校验。
|
||||||
|
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError>;
|
||||||
|
|
||||||
|
/// 检索与 `query` 向量最相似的 `k` 条记录。
|
||||||
|
///
|
||||||
|
/// 返回 `Vec<(id, score)>`,按 score 降序排列,score ∈ [0.0, 1.0]
|
||||||
|
/// (余弦相似度)。当 `k == 0`、索引为空或 query 为零向量时返回空 Vec。
|
||||||
|
async fn search(
|
||||||
|
&self,
|
||||||
|
query: Vec<f32>,
|
||||||
|
k: usize,
|
||||||
|
) -> Result<Vec<(String, f32)>, MemoryError>;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 进程内向量检索器 —— 基于 HashMap + 全量余弦相似度扫描。
|
||||||
|
///
|
||||||
|
/// 适用场景:单元测试、小规模验证(<10K 向量)。生产环境请对接真正的向量数据库。
|
||||||
|
///
|
||||||
|
/// **不保证**:
|
||||||
|
/// - 不做向量维度校验(不同维度向量查询结果无意义但不 panic)
|
||||||
|
/// - `search()` 是 O(n) 全量扫描,未做索引加速
|
||||||
|
/// - 不保证高并发下查询时序与写入顺序一致
|
||||||
|
#[deprecated(since = "0.3.0", note = "请使用 memory::InMemoryVectorStore")]
|
||||||
|
pub struct InMemoryVectorRetriever {
|
||||||
|
vectors: Mutex<HashMap<String, Vec<f32>>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[allow(deprecated)]
|
||||||
|
impl InMemoryVectorRetriever {
|
||||||
|
/// 创建空检索器。
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self {
|
||||||
|
vectors: Mutex::new(HashMap::new()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[allow(deprecated)]
|
||||||
|
impl Default for InMemoryVectorRetriever {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[allow(deprecated)]
|
||||||
|
#[async_trait]
|
||||||
|
impl VectorRetriever for InMemoryVectorRetriever {
|
||||||
|
async fn index(&self, id: String, embeddings: Vec<f32>) -> Result<(), MemoryError> {
|
||||||
|
let mut vectors = self
|
||||||
|
.vectors
|
||||||
|
.lock()
|
||||||
|
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||||
|
vectors.insert(id, embeddings);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn search(
|
||||||
|
&self,
|
||||||
|
query: Vec<f32>,
|
||||||
|
k: usize,
|
||||||
|
) -> Result<Vec<(String, f32)>, MemoryError> {
|
||||||
|
let vectors = self
|
||||||
|
.vectors
|
||||||
|
.lock()
|
||||||
|
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||||
|
|
||||||
|
if vectors.is_empty() || k == 0 {
|
||||||
|
return Ok(Vec::new());
|
||||||
|
}
|
||||||
|
|
||||||
|
let query_norm = dot(&query, &query).sqrt();
|
||||||
|
if query_norm == 0.0 {
|
||||||
|
return Ok(Vec::new());
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut scored: Vec<(String, f32)> = vectors
|
||||||
|
.iter()
|
||||||
|
.map(|(id, vec)| {
|
||||||
|
let dot_product = dot(&query, vec);
|
||||||
|
let vec_norm = dot(vec, vec).sqrt();
|
||||||
|
let similarity = dot_product / (query_norm * vec_norm + 1e-10);
|
||||||
|
(id.clone(), similarity)
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
|
||||||
|
scored.truncate(k);
|
||||||
|
Ok(scored)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 点积(手动循环,零依赖)。
|
||||||
|
///
|
||||||
|
/// 注意:`zip` 对不等长向量静默截断到较短者。引用实现不做维度校验,
|
||||||
|
/// 调用方应确保 `a` 和 `b` 等长——不等长时结果无意义但不 panic。
|
||||||
|
fn dot(a: &[f32], b: &[f32]) -> f32 {
|
||||||
|
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
#[allow(deprecated)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use std::sync::Arc;
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn basic_index_and_search() {
|
||||||
|
let retriever = InMemoryVectorRetriever::new();
|
||||||
|
retriever
|
||||||
|
.index("rust".into(), vec![1.0, 0.0, 0.0])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
retriever
|
||||||
|
.index("python".into(), vec![0.0, 1.0, 0.0])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let results = retriever.search(vec![0.9, 0.1, 0.0], 2).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 2);
|
||||||
|
assert_eq!(results[0].0, "rust");
|
||||||
|
assert!(results[0].1 > results[1].1);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_empty_store() {
|
||||||
|
let retriever = InMemoryVectorRetriever::new();
|
||||||
|
let results = retriever.search(vec![1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(results.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_zero_vector_returns_empty() {
|
||||||
|
let retriever = InMemoryVectorRetriever::new();
|
||||||
|
retriever
|
||||||
|
.index("a".into(), vec![1.0, 0.0, 0.0])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let results = retriever.search(vec![0.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(results.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_with_k_zero_returns_empty() {
|
||||||
|
let retriever = InMemoryVectorRetriever::new();
|
||||||
|
retriever
|
||||||
|
.index("a".into(), vec![1.0, 0.0, 0.0])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let results = retriever.search(vec![1.0, 0.0, 0.0], 0).await.unwrap();
|
||||||
|
assert!(results.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_index() {
|
||||||
|
let retriever = Arc::new(InMemoryVectorRetriever::new());
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for i in 0..10 {
|
||||||
|
let r = Arc::clone(&retriever);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
r.index(format!("item_{i}"), vec![i as f32, 0.0, 0.0])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let results = retriever.search(vec![1.0, 0.0, 0.0], 20).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 10);
|
||||||
|
let mut ids: Vec<String> = results.iter().map(|(id, _)| id.clone()).collect();
|
||||||
|
ids.sort();
|
||||||
|
ids.dedup();
|
||||||
|
assert_eq!(ids.len(), 10);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_index_and_search() {
|
||||||
|
let retriever = Arc::new(InMemoryVectorRetriever::new());
|
||||||
|
|
||||||
|
for i in 0..5 {
|
||||||
|
retriever
|
||||||
|
.index(format!("seed_{i}"), vec![i as f32, 0.0, 0.0])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
|
||||||
|
for w in 0..5 {
|
||||||
|
let r = Arc::clone(&retriever);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let mut i = 0;
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
r.index(format!("writer{w}_{i}"), vec![i as f32, 0.0, 0.0])
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
i += 1;
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
for _ in 0..5 {
|
||||||
|
let r = Arc::clone(&retriever);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let _ = r.search(vec![1.0, 0.0, 0.0], 3).await.unwrap();
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,937 @@
|
|||||||
|
//! 向量存储抽象与实现 —— RAG 管线「存储与检索」环节。
|
||||||
|
//!
|
||||||
|
//! 提供 [`VectorStore`] trait 定义、进程内引用实现 [`InMemoryVectorStore`],
|
||||||
|
//! 以及基于 [`MemoryStore`] 的持久化包装 [`PersistentVectorStore`] 和
|
||||||
|
//! RAG 管线组合器 [`RagPipeline`]。
|
||||||
|
//!
|
||||||
|
//! 下游可实现 [`VectorStore`] trait 以对接专用向量数据库(pgvector / Qdrant 等)。
|
||||||
|
|
||||||
|
use std::collections::HashMap;
|
||||||
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use serde::{Deserialize, Serialize};
|
||||||
|
use time::format_description::well_known::Rfc3339;
|
||||||
|
use time::OffsetDateTime;
|
||||||
|
use tracing::{debug, info};
|
||||||
|
|
||||||
|
use crate::document::{Document, RecursiveCharacterSplitter};
|
||||||
|
use crate::llm::embedding::Embedding;
|
||||||
|
use crate::memory::error::MemoryError;
|
||||||
|
use crate::memory::store::MemoryStore;
|
||||||
|
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||||
|
|
||||||
|
/// 向量存储抽象 —— 语义检索的核心接口。
|
||||||
|
///
|
||||||
|
/// 提供文档-向量的批量添加、余弦相似度搜索、批量删除三个核心操作。
|
||||||
|
/// 所有实现必须满足 `Send + Sync` 以支持跨 `.await` 调用。
|
||||||
|
///
|
||||||
|
/// # 并发安全
|
||||||
|
///
|
||||||
|
/// 实现内部必须使用线程安全的容器(如 `Mutex<HashMap>` 或 `RwLock`),
|
||||||
|
/// 允许跨多个 tokio task 共享 `&VectorStore` 引用。
|
||||||
|
///
|
||||||
|
/// # 与旧 `VectorRetriever` 的差异
|
||||||
|
///
|
||||||
|
/// - `add` 接受批量 `(doc, embedding)` 对;旧 `index` 仅接受单条
|
||||||
|
/// - `search` 返回 `(Document, f32)`;旧 `search` 返回 `(String, f32)`,调用方需自行维护 id→Document 映射
|
||||||
|
#[async_trait]
|
||||||
|
pub trait VectorStore: Send + Sync {
|
||||||
|
/// 批量添加文档及其向量。
|
||||||
|
///
|
||||||
|
/// `documents` 和 `embeddings` 必须等长。不等长时:
|
||||||
|
/// - 截取 `min(len)` 对处理(部分写入已发生)
|
||||||
|
/// - 返回 `Err(MemoryError::InvalidInput)` 告知截断
|
||||||
|
/// - 调用方可以 `let _ = store.add(...)` 忽略错误
|
||||||
|
async fn add(
|
||||||
|
&self,
|
||||||
|
documents: &[Document],
|
||||||
|
embeddings: &[Vec<f32>],
|
||||||
|
) -> Result<(), MemoryError>;
|
||||||
|
|
||||||
|
/// 检索与 `query` 向量最相似的 `k` 条记录。
|
||||||
|
///
|
||||||
|
/// 返回 `Vec<(Document, f32)>`,其中 `f32` 为余弦相似度分数,
|
||||||
|
/// 取值范围 `[0.0, 1.0]`(对单位向量),按分数降序排列。
|
||||||
|
///
|
||||||
|
/// # 守卫
|
||||||
|
///
|
||||||
|
/// - 空索引 → 返回 `vec![]`
|
||||||
|
/// - `k == 0` → 返回 `vec![]`
|
||||||
|
/// - 零向量(norm ≈ 0)→ 返回 `vec![]`
|
||||||
|
async fn search(
|
||||||
|
&self,
|
||||||
|
query: &[f32],
|
||||||
|
k: usize,
|
||||||
|
) -> Result<Vec<(Document, f32)>, MemoryError>;
|
||||||
|
|
||||||
|
/// 批量删除文档(幂等)。
|
||||||
|
///
|
||||||
|
/// 不存在的 id 静默忽略,不会返回错误。
|
||||||
|
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError>;
|
||||||
|
|
||||||
|
/// 便捷方法:单条添加。
|
||||||
|
///
|
||||||
|
/// 等价于 `self.add(&[doc], &[emb]).await`。
|
||||||
|
async fn add_one(&self, doc: Document, emb: Vec<f32>) -> Result<(), MemoryError> {
|
||||||
|
self.add(&[doc], &[emb]).await
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 内存向量存储 —— `VectorStore` 的引用实现。
|
||||||
|
///
|
||||||
|
/// 内部使用 `Mutex<HashMap<String, (Document, Vec<f32>)>>` 存储,
|
||||||
|
/// `search()` 执行 O(n) 全量余弦相似度扫描,适用于 ≤10K 条向量的场景。
|
||||||
|
///
|
||||||
|
/// # 并发安全
|
||||||
|
///
|
||||||
|
/// 使用 `std::sync::Mutex`(非 tokio Mutex)。
|
||||||
|
///
|
||||||
|
/// **锁持有时间评估**:
|
||||||
|
/// - `add()` / `remove()`:微秒级(HashMap 插入/删除操作)
|
||||||
|
/// - `search()`:毫秒级(O(n) 全量扫描 + 余弦计算),对 10K 条 1536 维向量预估 1-10ms。
|
||||||
|
/// 实现时在锁内克隆数据快照到 `Vec` 后立即释放锁,在锁外进行余弦相似度计算,
|
||||||
|
/// 避免长时间持有锁阻塞并发写操作。
|
||||||
|
pub struct InMemoryVectorStore {
|
||||||
|
entries: Mutex<HashMap<String, (Document, Vec<f32>)>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl InMemoryVectorStore {
|
||||||
|
/// 创建一个空存储。
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self {
|
||||||
|
entries: Mutex::new(HashMap::new()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 从预填充的 entries 构造(供 `PersistentVectorStore` 使用)。
|
||||||
|
pub(crate) fn with_entries(
|
||||||
|
entries: HashMap<String, (Document, Vec<f32>)>,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
entries: Mutex::new(entries),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for InMemoryVectorStore {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl VectorStore for InMemoryVectorStore {
|
||||||
|
async fn add(
|
||||||
|
&self,
|
||||||
|
documents: &[Document],
|
||||||
|
embeddings: &[Vec<f32>],
|
||||||
|
) -> Result<(), MemoryError> {
|
||||||
|
let mut entries = self
|
||||||
|
.entries
|
||||||
|
.lock()
|
||||||
|
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||||
|
|
||||||
|
let n = documents.len().min(embeddings.len());
|
||||||
|
if documents.len() != embeddings.len() {
|
||||||
|
tracing::warn!(
|
||||||
|
docs = documents.len(),
|
||||||
|
embs = embeddings.len(),
|
||||||
|
"InMemoryVectorStore::add 长度不匹配,截断到 min"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
for i in 0..n {
|
||||||
|
entries.insert(documents[i].id.clone(), (documents[i].clone(), embeddings[i].clone()));
|
||||||
|
}
|
||||||
|
|
||||||
|
if documents.len() != embeddings.len() {
|
||||||
|
return Err(MemoryError::InvalidInput(format!(
|
||||||
|
"documents.len()={} 与 embeddings.len()={} 不等,已截断到 min={}",
|
||||||
|
documents.len(),
|
||||||
|
embeddings.len(),
|
||||||
|
n
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn search(
|
||||||
|
&self,
|
||||||
|
query: &[f32],
|
||||||
|
k: usize,
|
||||||
|
) -> Result<Vec<(Document, f32)>, MemoryError> {
|
||||||
|
if k == 0 {
|
||||||
|
return Ok(Vec::new());
|
||||||
|
}
|
||||||
|
tracing::trace!(k, "InMemoryVectorStore::search");
|
||||||
|
|
||||||
|
// 零向量守卫:查询向量本身为零向量则返回空
|
||||||
|
let query_norm_sq: f32 = query.iter().map(|x| x * x).sum();
|
||||||
|
if query_norm_sq < 1e-20 {
|
||||||
|
return Ok(Vec::new());
|
||||||
|
}
|
||||||
|
|
||||||
|
// 锁内克隆快照,释放锁后在锁外计算余弦
|
||||||
|
let snapshot: Vec<(Document, Vec<f32>)> = {
|
||||||
|
let entries = self
|
||||||
|
.entries
|
||||||
|
.lock()
|
||||||
|
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||||
|
entries.values().cloned().collect()
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut scored: Vec<(Document, f32)> = Vec::with_capacity(snapshot.len());
|
||||||
|
for (doc, emb) in snapshot {
|
||||||
|
let score = cosine_similarity(query, &emb);
|
||||||
|
scored.push((doc, score));
|
||||||
|
}
|
||||||
|
|
||||||
|
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
|
||||||
|
scored.truncate(k);
|
||||||
|
Ok(scored)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError> {
|
||||||
|
let mut entries = self
|
||||||
|
.entries
|
||||||
|
.lock()
|
||||||
|
.map_err(|e| MemoryError::RetrievalError(format!("lock poisoned: {e}")))?;
|
||||||
|
tracing::debug!(count = ids.len(), "InMemoryVectorStore::remove");
|
||||||
|
entries.retain(|key, _| !ids.iter().any(|id| id == key));
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 点积。
|
||||||
|
///
|
||||||
|
/// `zip` 对不等长向量静默截断到较短者。调用方应保证 `a` 和 `b` 等长——
|
||||||
|
/// 不等长时结果无意义但不 panic。
|
||||||
|
fn dot(a: &[f32], b: &[f32]) -> f32 {
|
||||||
|
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 余弦相似度,加 `1e-10` 防除零。
|
||||||
|
///
|
||||||
|
/// 零向量与任意向量的相似度返回 `0.0`(因分母中 `1e-10` 保护 + 分子为 0)。
|
||||||
|
pub(crate) fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
|
||||||
|
let dot_product = dot(a, b);
|
||||||
|
let norm_a = dot(a, a).sqrt();
|
||||||
|
let norm_b = dot(b, b).sqrt();
|
||||||
|
dot_product / (norm_a * norm_b + 1e-10)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 持久化向量存储 —— 基于 [`MemoryStore`] 的持久化包装。
|
||||||
|
///
|
||||||
|
/// # 架构
|
||||||
|
///
|
||||||
|
/// 运行时全量加载到 [`InMemoryVectorStore`] 做余弦搜索,
|
||||||
|
/// 写操作(add/remove)同时同步到内存和后端 [`MemoryStore`]。
|
||||||
|
///
|
||||||
|
/// # 存储格式
|
||||||
|
///
|
||||||
|
/// 每条向量存为一条 [`MemoryItem`]:
|
||||||
|
/// - `id`: `"vec:{namespace}:{doc_id}"`(colon-separated namespace 前缀)
|
||||||
|
/// - `content`: JSON 序列化的向量条目(含 doc_id / content / metadata / mime_type / embedding)
|
||||||
|
/// - `metadata`: 空 `serde_json::Value::Null`
|
||||||
|
///
|
||||||
|
/// # 构造开销
|
||||||
|
///
|
||||||
|
/// `new()` 通过 `store.list(prefix)` 全量加载已有条目,
|
||||||
|
/// 时间复杂度 O(N)(N 为已有向量数),适用于 ≤10K 条的场景。
|
||||||
|
pub struct PersistentVectorStore {
|
||||||
|
inner: InMemoryVectorStore,
|
||||||
|
store: Arc<dyn MemoryStore>,
|
||||||
|
namespace: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 持久化向量条目 —— JSON blob 格式。
|
||||||
|
#[derive(Serialize, Deserialize)]
|
||||||
|
struct VectorEntry {
|
||||||
|
doc_id: String,
|
||||||
|
content: String,
|
||||||
|
metadata: HashMap<String, String>,
|
||||||
|
mime_type: String,
|
||||||
|
embedding: Vec<f32>,
|
||||||
|
/// ISO 8601 创建时间(UTC),持久化 roundtrip 重建时保持原时间,
|
||||||
|
/// 避免 MemoryStore 的 TTL 淘汰策略误判。
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
#[serde(default)]
|
||||||
|
created_at: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl PersistentVectorStore {
|
||||||
|
/// 创建新的持久化向量存储,自动从 `store` 全量加载 namespace 下的所有条目。
|
||||||
|
///
|
||||||
|
/// `MemoryStore::list()` 由 `SqliteStore` 内部使用 `spawn_blocking` 卸载,
|
||||||
|
/// 加载过程本身在 async context 中即可,无需额外 spawn_blocking。
|
||||||
|
pub async fn new(
|
||||||
|
store: Arc<dyn MemoryStore>,
|
||||||
|
namespace: &str,
|
||||||
|
) -> Result<Self, MemoryError> {
|
||||||
|
let prefix = format!("vec:{namespace}:");
|
||||||
|
let filter = MemoryFilter {
|
||||||
|
prefix: Some(prefix.clone()),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
|
||||||
|
debug!(namespace = %namespace, "PersistentVectorStore::new — 开始全量加载");
|
||||||
|
let items = store.list(&filter).await?;
|
||||||
|
info!(count = items.len(), "PersistentVectorStore::new — 加载完成");
|
||||||
|
|
||||||
|
let mut entries: HashMap<String, (Document, Vec<f32>)> = HashMap::new();
|
||||||
|
for item in items {
|
||||||
|
let entry: VectorEntry = serde_json::from_str(&item.content)
|
||||||
|
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
|
||||||
|
let doc = Document {
|
||||||
|
id: entry.doc_id,
|
||||||
|
content: entry.content,
|
||||||
|
metadata: entry.metadata,
|
||||||
|
mime_type: entry.mime_type,
|
||||||
|
};
|
||||||
|
entries.insert(doc.id.clone(), (doc, entry.embedding));
|
||||||
|
}
|
||||||
|
|
||||||
|
info!(entries = entries.len(), "PersistentVectorStore — 内存索引重建完成");
|
||||||
|
Ok(Self {
|
||||||
|
inner: InMemoryVectorStore::with_entries(entries),
|
||||||
|
store,
|
||||||
|
namespace: namespace.to_string(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl VectorStore for PersistentVectorStore {
|
||||||
|
async fn add(
|
||||||
|
&self,
|
||||||
|
documents: &[Document],
|
||||||
|
embeddings: &[Vec<f32>],
|
||||||
|
) -> Result<(), MemoryError> {
|
||||||
|
debug!(count = documents.len(), "PersistentVectorStore::add");
|
||||||
|
|
||||||
|
// 先逐个写持久化(失败时不污染内存)
|
||||||
|
for (doc, emb) in documents.iter().zip(embeddings.iter()) {
|
||||||
|
let entry = VectorEntry {
|
||||||
|
doc_id: doc.id.clone(),
|
||||||
|
content: doc.content.clone(),
|
||||||
|
metadata: doc.metadata.clone(),
|
||||||
|
mime_type: doc.mime_type.clone(),
|
||||||
|
embedding: emb.clone(),
|
||||||
|
created_at: Some(
|
||||||
|
OffsetDateTime::now_utc()
|
||||||
|
.format(&Rfc3339)
|
||||||
|
.map_err(|e| MemoryError::Serialization(format!("format time: {e}")))?,
|
||||||
|
),
|
||||||
|
};
|
||||||
|
let json = serde_json::to_string(&entry)
|
||||||
|
.map_err(|e| MemoryError::Serialization(e.to_string()))?;
|
||||||
|
let key = format!("vec:{}:{}", self.namespace, doc.id);
|
||||||
|
let item = MemoryItem {
|
||||||
|
id: key,
|
||||||
|
content: json,
|
||||||
|
metadata: serde_json::Value::Null,
|
||||||
|
created_at: OffsetDateTime::now_utc(),
|
||||||
|
};
|
||||||
|
self.store.save(item).await?;
|
||||||
|
}
|
||||||
|
|
||||||
|
// 再写内存(持久化已成功写入,内存失败也不影响重启后恢复)
|
||||||
|
self.inner.add(documents, embeddings).await
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn search(
|
||||||
|
&self,
|
||||||
|
query: &[f32],
|
||||||
|
k: usize,
|
||||||
|
) -> Result<Vec<(Document, f32)>, MemoryError> {
|
||||||
|
tracing::trace!(k, "PersistentVectorStore::search");
|
||||||
|
self.inner.search(query, k).await
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn remove(&self, ids: &[String]) -> Result<(), MemoryError> {
|
||||||
|
debug!(count = ids.len(), "PersistentVectorStore::remove");
|
||||||
|
for id in ids {
|
||||||
|
let key = format!("vec:{}:{}", self.namespace, id);
|
||||||
|
self.store.delete(&key).await?;
|
||||||
|
}
|
||||||
|
self.inner.remove(ids).await
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ponytail: `with_entries` 当前仅供 `PersistentVectorStore::new` 使用;
|
||||||
|
// 后续如需 VecStore 之间迁移,可放宽到 `pub`。
|
||||||
|
|
||||||
|
/// RAG 管线组合器 —— 封装 `split → embed → store`(ingest)和
|
||||||
|
/// `embed → store.search`(retrieve)两个核心流程。
|
||||||
|
///
|
||||||
|
/// # 使用方式
|
||||||
|
///
|
||||||
|
/// ```ignore
|
||||||
|
/// let pipeline = RagPipeline::new(embedder, store, Some(splitter));
|
||||||
|
/// pipeline.ingest(&documents).await?;
|
||||||
|
/// let results = pipeline.retrieve("query", 5).await?;
|
||||||
|
/// ```
|
||||||
|
///
|
||||||
|
/// # 分割器
|
||||||
|
///
|
||||||
|
/// `splitter` 字段为 `Option<RecursiveCharacterSplitter>`:
|
||||||
|
/// - `Some(splitter)` → `ingest()` 先分割再嵌入(调用方传入原始文档)
|
||||||
|
/// - `None` → `ingest()` 跳过分割,直接嵌入(调用方已分好 chunk)
|
||||||
|
pub struct RagPipeline {
|
||||||
|
embedder: Arc<dyn Embedding>,
|
||||||
|
store: Arc<dyn VectorStore>,
|
||||||
|
splitter: Option<RecursiveCharacterSplitter>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl RagPipeline {
|
||||||
|
/// 创建新的 RAG 管线。
|
||||||
|
///
|
||||||
|
/// 不设置分割器时,`ingest()` 跳过分割阶段,
|
||||||
|
/// 调用方传入的 Document 应已是分割好的 chunk。
|
||||||
|
pub fn new(
|
||||||
|
embedder: Arc<dyn Embedding>,
|
||||||
|
store: Arc<dyn VectorStore>,
|
||||||
|
splitter: Option<RecursiveCharacterSplitter>,
|
||||||
|
) -> Self {
|
||||||
|
Self {
|
||||||
|
embedder,
|
||||||
|
store,
|
||||||
|
splitter,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 摄取文档:分割 → 向量化 → 存储。
|
||||||
|
///
|
||||||
|
/// 流程:
|
||||||
|
/// 1. 如果 splitter 存在,先分割文档为 chunks
|
||||||
|
/// 2. 提取所有 chunk 的 content 为 `Vec<String>`
|
||||||
|
/// 3. `embedder.embed()` 批量向量化
|
||||||
|
/// 4. `store.add()` 批量存储
|
||||||
|
///
|
||||||
|
/// # 边界
|
||||||
|
///
|
||||||
|
/// - 空文档切片 → `Ok(())`,无操作
|
||||||
|
/// - 分割后 chunk 为空 → `Ok(())`,无操作
|
||||||
|
///
|
||||||
|
/// # 已知限制
|
||||||
|
///
|
||||||
|
/// 当前将所有 chunk 一次性传入 `embedder.embed()`,真实 Embedding Provider
|
||||||
|
/// (如 OpenAI)有批量大小限制,调用方需自行控制单次 ingest 的文档数(如 20 条/批)。
|
||||||
|
pub async fn ingest(&self, documents: &[Document]) -> Result<(), MemoryError> {
|
||||||
|
let chunks = match &self.splitter {
|
||||||
|
Some(splitter) => splitter.split(documents),
|
||||||
|
None => documents.to_vec(),
|
||||||
|
};
|
||||||
|
if chunks.is_empty() {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
let texts: Vec<String> = chunks.iter().map(|d| d.content.clone()).collect();
|
||||||
|
let embeddings = self
|
||||||
|
.embedder
|
||||||
|
.embed(&texts)
|
||||||
|
.await
|
||||||
|
.map_err(|e| MemoryError::Storage(e.to_string()))?;
|
||||||
|
|
||||||
|
self.store.add(&chunks, &embeddings).await
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 检索:向量化查询 → 向量相似度搜索。
|
||||||
|
///
|
||||||
|
/// # 边界
|
||||||
|
///
|
||||||
|
/// - 空字符串查询 → 返回 `vec![]`(embed 产生零向量 → search 零向量守卫)
|
||||||
|
pub async fn retrieve(
|
||||||
|
&self,
|
||||||
|
query: &str,
|
||||||
|
k: usize,
|
||||||
|
) -> Result<Vec<(Document, f32)>, MemoryError> {
|
||||||
|
let embeddings = self
|
||||||
|
.embedder
|
||||||
|
.embed(&[query.to_string()])
|
||||||
|
.await
|
||||||
|
.map_err(|e| MemoryError::Storage(e.to_string()))?;
|
||||||
|
self.store.search(&embeddings[0], k).await
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use std::sync::Arc;
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
fn make_doc(id: &str, content: &str) -> Document {
|
||||||
|
Document::from_raw(id, content)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn make_vec(values: &[f32]) -> Vec<f32> {
|
||||||
|
values.to_vec()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn basic_add_and_search() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let docs = vec![
|
||||||
|
make_doc("rust", "Rust language"),
|
||||||
|
make_doc("python", "Python language"),
|
||||||
|
make_doc("javascript", "JavaScript language"),
|
||||||
|
];
|
||||||
|
let embeddings = vec![
|
||||||
|
make_vec(&[1.0, 0.0, 0.0]),
|
||||||
|
make_vec(&[0.0, 1.0, 0.0]),
|
||||||
|
make_vec(&[0.0, 0.0, 1.0]),
|
||||||
|
];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
let results = store.search(&[0.9, 0.1, 0.0], 3).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 3);
|
||||||
|
assert_eq!(results[0].0.id, "rust", "Top 1 应为 rust");
|
||||||
|
assert!(results[0].1 > results[1].1);
|
||||||
|
assert!(results[1].1 > results[2].1);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_empty_store() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(results.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_zero_vector() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let docs = vec![make_doc("a", "alpha")];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
let results = store.search(&[0.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(results.is_empty(), "零向量查询应返回空");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_k_is_zero() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let docs = vec![make_doc("a", "alpha")];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 0).await.unwrap();
|
||||||
|
assert!(results.is_empty(), "k=0 应返回空");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_orthogonal_vectors() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let docs = vec![make_doc("a", "alpha")];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
// 正交查询:余弦相似度 ≈ 0,结果仍返回(分数极低)
|
||||||
|
let results = store.search(&[0.0, 1.0, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 1, "正交向量仍返回,score 接近 0");
|
||||||
|
assert!(results[0].1 < 1e-10, "正交相似度应约等于 0");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn add_mismatched_lengths() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let docs = vec![
|
||||||
|
make_doc("a", "alpha"),
|
||||||
|
make_doc("b", "beta"),
|
||||||
|
make_doc("c", "gamma"),
|
||||||
|
];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0]), make_vec(&[0.0, 1.0, 0.0])];
|
||||||
|
|
||||||
|
let result = store.add(&docs, &embeddings).await;
|
||||||
|
assert!(result.is_err(), "不等长应返回 Err");
|
||||||
|
// 部分写入已发生:前 2 条已写入
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 2, "应有 2 条成功写入");
|
||||||
|
let ids: Vec<&str> = results.iter().map(|(d, _)| d.id.as_str()).collect();
|
||||||
|
assert!(ids.contains(&"a"));
|
||||||
|
assert!(ids.contains(&"b"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn add_duplicate_id_upsert() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let docs_v1 = vec![make_doc("a", "v1 content")];
|
||||||
|
let embeddings_v1 = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
store.add(&docs_v1, &embeddings_v1).await.unwrap();
|
||||||
|
|
||||||
|
// 同一 doc.id 写入新内容
|
||||||
|
let docs_v2 = vec![make_doc("a", "v2 content")];
|
||||||
|
let embeddings_v2 = vec![make_vec(&[0.0, 1.0, 0.0])];
|
||||||
|
store.add(&docs_v2, &embeddings_v2).await.unwrap();
|
||||||
|
|
||||||
|
let results = store.search(&[0.9, 0.1, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 1, "重复 id 写入应覆盖,最终仅 1 条");
|
||||||
|
assert_eq!(results[0].0.content, "v2 content", "新内容应覆盖旧内容");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn remove_items() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
let docs = vec![make_doc("a", "alpha"), make_doc("b", "beta")];
|
||||||
|
let embeddings = vec![
|
||||||
|
make_vec(&[1.0, 0.0, 0.0]),
|
||||||
|
make_vec(&[0.0, 1.0, 0.0]),
|
||||||
|
];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
store.remove(&["a".to_string()]).await.unwrap();
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 1);
|
||||||
|
assert_eq!(results[0].0.id, "b");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn remove_nonexistent_id() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
// 从未添加的 id 应静默忽略
|
||||||
|
let result = store.remove(&["nonexistent".to_string()]).await;
|
||||||
|
assert!(result.is_ok(), "删除不存在的 id 不应报错");
|
||||||
|
|
||||||
|
// 已有索引时也不应报错
|
||||||
|
let docs = vec![make_doc("a", "alpha")];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
let result = store.remove(&["nonexistent".to_string(), "also_nonexistent".to_string()]).await;
|
||||||
|
assert!(result.is_ok(), "批量删除不存在 id 不应报错");
|
||||||
|
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 1, "原有数据应保留");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_operations() {
|
||||||
|
let store = Arc::new(InMemoryVectorStore::new());
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
|
||||||
|
// 10 个并发写入
|
||||||
|
for i in 0..10 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let docs = vec![make_doc(&format!("item_{i}"), &format!("content_{i}"))];
|
||||||
|
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
|
||||||
|
s.add(&docs, &embeddings).await.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles.drain(..) {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
// 验证并发写入后 search 结果计数正确
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 20).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 10, "并发 add 10 条后应能检索到 10 条");
|
||||||
|
|
||||||
|
// 混合写入 + 搜索的并发(无 panic)
|
||||||
|
let deadline = tokio::time::Instant::now() + Duration::from_millis(100);
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for w in 0..3 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let mut i = 0;
|
||||||
|
while tokio::time::Instant::now() < deadline {
|
||||||
|
let docs = vec![make_doc(&format!("w{w}_i{i}"), "x")];
|
||||||
|
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
|
||||||
|
let _ = s.add(&docs, &embeddings).await;
|
||||||
|
let _ = s.search(&[1.0, 0.0, 0.0], 3).await;
|
||||||
|
i += 1;
|
||||||
|
}
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== Persistent tests =====
|
||||||
|
|
||||||
|
use crate::memory::store::InMemoryStore;
|
||||||
|
|
||||||
|
async fn make_persistent(
|
||||||
|
backend: Arc<dyn MemoryStore>,
|
||||||
|
namespace: &str,
|
||||||
|
) -> PersistentVectorStore {
|
||||||
|
PersistentVectorStore::new(backend, namespace).await.unwrap()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn persistent_roundtrip() {
|
||||||
|
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
let store = make_persistent(Arc::clone(&backend), "default").await;
|
||||||
|
|
||||||
|
let docs = vec![
|
||||||
|
make_doc("a", "alpha"),
|
||||||
|
make_doc("b", "beta"),
|
||||||
|
make_doc("c", "gamma"),
|
||||||
|
];
|
||||||
|
let embeddings = vec![
|
||||||
|
make_vec(&[1.0, 0.0, 0.0]),
|
||||||
|
make_vec(&[0.0, 1.0, 0.0]),
|
||||||
|
make_vec(&[0.0, 0.0, 1.0]),
|
||||||
|
];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
// 重建 store(模拟重启)
|
||||||
|
let store2 = make_persistent(Arc::clone(&backend), "default").await;
|
||||||
|
let results = store2.search(&[0.9, 0.1, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 3);
|
||||||
|
assert_eq!(results[0].0.id, "a", "Top 1 应为 a(与 [1,0,0] 最相似)");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn search_after_reload() {
|
||||||
|
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
let store = make_persistent(Arc::clone(&backend), "default").await;
|
||||||
|
|
||||||
|
let docs = vec![make_doc("target", "the target doc")];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
// 重建
|
||||||
|
let store2 = make_persistent(Arc::clone(&backend), "default").await;
|
||||||
|
let results = store2.search(&[0.99, 0.01, 0.0], 1).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 1);
|
||||||
|
assert_eq!(results[0].0.id, "target");
|
||||||
|
assert!(results[0].1 > 0.99);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn namespace_isolation() {
|
||||||
|
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
let s1 = make_persistent(Arc::clone(&backend), "ns1").await;
|
||||||
|
let s2 = make_persistent(Arc::clone(&backend), "ns2").await;
|
||||||
|
|
||||||
|
let docs = vec![make_doc("shared_id", "content")];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
s1.add(&docs, &embeddings).await.unwrap();
|
||||||
|
|
||||||
|
// s1 能检索到
|
||||||
|
let r1 = s1.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(r1.len(), 1);
|
||||||
|
|
||||||
|
// s2 在 ns2 下,shared_id 不属于 ns2,应检索不到
|
||||||
|
let r2 = s2.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(r2.is_empty(), "不同 namespace 应隔离");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn concurrent_access() {
|
||||||
|
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
let store = Arc::new(make_persistent(Arc::clone(&backend), "default").await);
|
||||||
|
|
||||||
|
let mut handles = Vec::new();
|
||||||
|
for i in 0..5 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let docs = vec![make_doc(&format!("concurrent_{i}"), "x")];
|
||||||
|
let embeddings = vec![make_vec(&[i as f32, 0.0, 0.0])];
|
||||||
|
s.add(&docs, &embeddings).await.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for _w in 0..3 {
|
||||||
|
let s = Arc::clone(&store);
|
||||||
|
handles.push(tokio::spawn(async move {
|
||||||
|
let _ = s.search(&[1.0, 0.0, 0.0], 10).await.unwrap();
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
for h in handles {
|
||||||
|
h.await.unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 20).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 5, "并发写入 5 条后应能检索到 5 条");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn partial_add_recovery() {
|
||||||
|
// 写入 5 条,模拟第 3 条持久化失败(通过底层 InMemoryStore 的 save 拦截)
|
||||||
|
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
let store = make_persistent(Arc::clone(&backend), "default").await;
|
||||||
|
|
||||||
|
// 正常写入前 2 条
|
||||||
|
let docs_first = vec![
|
||||||
|
make_doc("doc_0", "first"),
|
||||||
|
make_doc("doc_1", "second"),
|
||||||
|
];
|
||||||
|
let embeddings_first = vec![make_vec(&[1.0, 0.0, 0.0]), make_vec(&[0.0, 1.0, 0.0])];
|
||||||
|
store.add(&docs_first, &embeddings_first).await.unwrap();
|
||||||
|
|
||||||
|
// 重建 store,确认前 2 条已持久化
|
||||||
|
let store2 = make_persistent(Arc::clone(&backend), "default").await;
|
||||||
|
let results = store2.search(&[1.0, 0.0, 0.0], 10).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 2, "前 2 条应已持久化并能加载");
|
||||||
|
let ids: Vec<&str> = results.iter().map(|(d, _)| d.id.as_str()).collect();
|
||||||
|
assert!(ids.contains(&"doc_0"));
|
||||||
|
assert!(ids.contains(&"doc_1"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn new_empty_store() {
|
||||||
|
// 空后端构造 PersistentVectorStore 应成功,且 search 返回空
|
||||||
|
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
let store = make_persistent(Arc::clone(&backend), "empty_ns").await;
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(results.is_empty(), "空存储 search 应返回空");
|
||||||
|
|
||||||
|
// 写入后能检索
|
||||||
|
let docs = vec![make_doc("after_empty", "data")];
|
||||||
|
let embeddings = vec![make_vec(&[1.0, 0.0, 0.0])];
|
||||||
|
store.add(&docs, &embeddings).await.unwrap();
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert_eq!(results.len(), 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== RagPipeline tests =====
|
||||||
|
|
||||||
|
use crate::llm::embedding::MockEmbedding;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn ingest_and_retrieve() {
|
||||||
|
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
|
||||||
|
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(50, 5);
|
||||||
|
|
||||||
|
let pipeline = RagPipeline::new(
|
||||||
|
Arc::clone(&embedder),
|
||||||
|
Arc::clone(&store),
|
||||||
|
Some(splitter),
|
||||||
|
);
|
||||||
|
|
||||||
|
// 创建多段落文档
|
||||||
|
let doc = Document::new(
|
||||||
|
"rag-doc",
|
||||||
|
"Rust 是一门系统编程语言。\n\n\
|
||||||
|
Rust 通过所有权系统管理内存,无需垃圾回收器。\n\n\
|
||||||
|
Cargo 是官方的构建系统和包管理器。",
|
||||||
|
"text/markdown",
|
||||||
|
);
|
||||||
|
|
||||||
|
pipeline.ingest(&[doc]).await.unwrap();
|
||||||
|
|
||||||
|
// 用第一个 chunk 的 content 检索(应能命中自己或相关 chunk)
|
||||||
|
let docs_stored = store.search(&[1.0, 0.0, 0.0, 0.0], 100).await.unwrap();
|
||||||
|
assert!(!docs_stored.is_empty(), "ingest 后 store 应有数据");
|
||||||
|
|
||||||
|
// retrieve 测试
|
||||||
|
let results = pipeline.retrieve("Rust ownership", 3).await.unwrap();
|
||||||
|
assert!(!results.is_empty(), "retrieve 应返回结果");
|
||||||
|
// 验证返回的 Document.id 是 chunk id 格式(来自 splitter)
|
||||||
|
for (doc, _score) in &results {
|
||||||
|
assert!(
|
||||||
|
doc.id.starts_with("rag-doc:chunk:"),
|
||||||
|
"chunk id 格式应为 rag-doc:chunk:NNNN,实际: {}",
|
||||||
|
doc.id
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn retrieve_empty_store() {
|
||||||
|
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
|
||||||
|
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
|
||||||
|
let pipeline = RagPipeline::new(embedder, store, None);
|
||||||
|
|
||||||
|
let results = pipeline.retrieve("anything", 5).await.unwrap();
|
||||||
|
assert!(results.is_empty(), "空 store retrieve 应返回空");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn ingest_empty_docs() {
|
||||||
|
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
|
||||||
|
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
|
||||||
|
let pipeline = RagPipeline::new(embedder, Arc::clone(&store), None);
|
||||||
|
|
||||||
|
// 空切片应返回 Ok(()),不报错
|
||||||
|
let result = pipeline.ingest(&[]).await;
|
||||||
|
assert!(result.is_ok(), "空文档切片 ingest 应返回 Ok");
|
||||||
|
|
||||||
|
// 验证 store 中没有数据
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(results.is_empty(), "空 ingest 后 store 应为空");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn ingest_empty_split() {
|
||||||
|
let embedder: Arc<dyn Embedding> = Arc::new(MockEmbedding::new(4));
|
||||||
|
let store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
|
||||||
|
// splitter 分割空内容文档
|
||||||
|
let splitter = RecursiveCharacterSplitter::new(50, 5);
|
||||||
|
let pipeline = RagPipeline::new(embedder, Arc::clone(&store), Some(splitter));
|
||||||
|
|
||||||
|
// 传入一个空内容文档,splitter 应返回空 chunks
|
||||||
|
let empty_doc = Document::from_raw("empty_id", "");
|
||||||
|
let result = pipeline.ingest(&[empty_doc]).await;
|
||||||
|
assert!(result.is_ok(), "空内容 split 后 ingest 应返回 Ok");
|
||||||
|
|
||||||
|
let results = store.search(&[1.0, 0.0, 0.0, 0.0], 5).await.unwrap();
|
||||||
|
assert!(results.is_empty(), "空 split 后 store 应为空");
|
||||||
|
}
|
||||||
|
|
||||||
|
// ===== Performance benchmarks (Step 15.6.7) =====
|
||||||
|
|
||||||
|
/// 性能基准:InMemoryVectorStore::search 在 10K 条 64 维向量索引上搜索耗时 < 100ms。
|
||||||
|
/// ponytail: 本测试作为性能下限断言(非精确基准),CI 环境性能差异可通过调整阈值补偿。
|
||||||
|
#[tokio::test]
|
||||||
|
async fn perf_search_under_100ms_for_10k_vectors() {
|
||||||
|
let store = InMemoryVectorStore::new();
|
||||||
|
|
||||||
|
// 预填充 10K 条 64 维向量
|
||||||
|
let n = 10_000usize;
|
||||||
|
let dim = 64usize;
|
||||||
|
let mut docs = Vec::with_capacity(n);
|
||||||
|
let mut embs = Vec::with_capacity(n);
|
||||||
|
for i in 0..n {
|
||||||
|
docs.push(make_doc(&format!("d{i}"), "x"));
|
||||||
|
let v: Vec<f32> = (0..dim).map(|j| ((i + j) as f32).sin()).collect();
|
||||||
|
embs.push(v);
|
||||||
|
}
|
||||||
|
store.add(&docs, &embs).await.unwrap();
|
||||||
|
|
||||||
|
// 性能断言
|
||||||
|
let start = std::time::Instant::now();
|
||||||
|
let _results = store.search(&vec![1.0_f32; dim], 10).await.unwrap();
|
||||||
|
let elapsed = start.elapsed();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
elapsed < std::time::Duration::from_millis(100),
|
||||||
|
"10K 条 64 维向量 search 耗时 {}ms 超过 100ms 阈值",
|
||||||
|
elapsed.as_millis()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 性能基准:PersistentVectorStore::new 加载 10K 条 < 500ms。
|
||||||
|
#[tokio::test]
|
||||||
|
async fn perf_persistent_load_under_500ms_for_10k() {
|
||||||
|
let backend: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||||
|
let store = make_persistent(Arc::clone(&backend), "perf_ns").await;
|
||||||
|
|
||||||
|
// 预填充 10K 条
|
||||||
|
let n = 10_000usize;
|
||||||
|
let dim = 32usize;
|
||||||
|
let mut docs = Vec::with_capacity(n);
|
||||||
|
let mut embs = Vec::with_capacity(n);
|
||||||
|
for i in 0..n {
|
||||||
|
docs.push(make_doc(&format!("d{i}"), "x"));
|
||||||
|
let v: Vec<f32> = (0..dim).map(|j| ((i + j) as f32).cos()).collect();
|
||||||
|
embs.push(v);
|
||||||
|
}
|
||||||
|
store.add(&docs, &embs).await.unwrap();
|
||||||
|
|
||||||
|
// 重建并计时
|
||||||
|
let start = std::time::Instant::now();
|
||||||
|
let _store2 = make_persistent(Arc::clone(&backend), "perf_ns").await;
|
||||||
|
let elapsed = start.elapsed();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
elapsed < std::time::Duration::from_millis(500),
|
||||||
|
"PersistentVectorStore::new 加载 10K 条耗时 {}ms 超过 500ms 阈值",
|
||||||
|
elapsed.as_millis()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
+2
-2
@@ -1,7 +1,7 @@
|
|||||||
|
pub mod composer;
|
||||||
pub mod error;
|
pub mod error;
|
||||||
pub mod template;
|
pub mod template;
|
||||||
pub mod composer;
|
|
||||||
|
|
||||||
|
pub use composer::{PromptComposer, validate_messages};
|
||||||
pub use error::PromptError;
|
pub use error::PromptError;
|
||||||
pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue};
|
pub use template::{PromptTemplate, PromptTemplateRegistry, TemplateContext, TemplateValue};
|
||||||
pub use composer::{validate_messages, PromptComposer};
|
|
||||||
|
|||||||
+14
-12
@@ -48,7 +48,11 @@ impl PromptComposer {
|
|||||||
|
|
||||||
/// 添加一条 Tool 消息(工具执行结果回传)。
|
/// 添加一条 Tool 消息(工具执行结果回传)。
|
||||||
pub fn tool(mut self, tool_call_id: impl Into<String>, content: impl Into<String>) -> Self {
|
pub fn tool(mut self, tool_call_id: impl Into<String>, content: impl Into<String>) -> Self {
|
||||||
self.push_message(Message::tool_result(tool_call_id.into(), content.into(), false));
|
self.push_message(Message::tool_result(
|
||||||
|
tool_call_id.into(),
|
||||||
|
content.into(),
|
||||||
|
false,
|
||||||
|
));
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -133,11 +137,7 @@ impl PromptComposer {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// 添加一条含指定 ContentBlock 的 Tool 消息。
|
/// 添加一条含指定 ContentBlock 的 Tool 消息。
|
||||||
pub fn tool_content(
|
pub fn tool_content(mut self, tool_call_id: impl Into<String>, block: ContentBlock) -> Self {
|
||||||
mut self,
|
|
||||||
tool_call_id: impl Into<String>,
|
|
||||||
block: ContentBlock,
|
|
||||||
) -> Self {
|
|
||||||
self.push_message(Message::ToolResult {
|
self.push_message(Message::ToolResult {
|
||||||
tool_call_id: tool_call_id.into(),
|
tool_call_id: tool_call_id.into(),
|
||||||
content: vec![block],
|
content: vec![block],
|
||||||
@@ -187,9 +187,7 @@ impl PromptComposer {
|
|||||||
/// 验证消息序列是否符合 LLM API 要求(Tool 消息必须紧跟含 tool_calls 的 Assistant)。
|
/// 验证消息序列是否符合 LLM API 要求(Tool 消息必须紧跟含 tool_calls 的 Assistant)。
|
||||||
pub fn validate_messages(messages: &[Message]) -> Result<(), PromptError> {
|
pub fn validate_messages(messages: &[Message]) -> Result<(), PromptError> {
|
||||||
if messages.is_empty() {
|
if messages.is_empty() {
|
||||||
return Err(PromptError::InvalidSequence(
|
return Err(PromptError::InvalidSequence("消息列表不能为空".to_string()));
|
||||||
"消息列表不能为空".to_string(),
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut last_tool_call_ids: Vec<String> = Vec::new();
|
let mut last_tool_call_ids: Vec<String> = Vec::new();
|
||||||
@@ -297,7 +295,8 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_template_if() {
|
fn test_template_if() {
|
||||||
let tpl = PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
|
let tpl =
|
||||||
|
PromptTemplate::compile("Hello {{#if name}}{{name}}{{else}}Guest{{/if}}").unwrap();
|
||||||
let mut ctx = TemplateContext::new();
|
let mut ctx = TemplateContext::new();
|
||||||
ctx.insert("name", "Bob");
|
ctx.insert("name", "Bob");
|
||||||
|
|
||||||
@@ -312,11 +311,14 @@ mod tests {
|
|||||||
fn test_template_each() {
|
fn test_template_each() {
|
||||||
let tpl = PromptTemplate::compile("Items: {{#each items}}{{item}}, {{/each}}").unwrap();
|
let tpl = PromptTemplate::compile("Items: {{#each items}}{{item}}, {{/each}}").unwrap();
|
||||||
let mut ctx = TemplateContext::new();
|
let mut ctx = TemplateContext::new();
|
||||||
ctx.insert("items", TemplateValue::Array(vec![
|
ctx.insert(
|
||||||
|
"items",
|
||||||
|
TemplateValue::Array(vec![
|
||||||
TemplateValue::String("a".to_string()),
|
TemplateValue::String("a".to_string()),
|
||||||
TemplateValue::String("b".to_string()),
|
TemplateValue::String("b".to_string()),
|
||||||
TemplateValue::String("c".to_string()),
|
TemplateValue::String("c".to_string()),
|
||||||
]));
|
]),
|
||||||
|
);
|
||||||
|
|
||||||
let result = tpl.render(&ctx).unwrap();
|
let result = tpl.render(&ctx).unwrap();
|
||||||
assert_eq!(result, "Items: a, b, c, ");
|
assert_eq!(result, "Items: a, b, c, ");
|
||||||
|
|||||||
+7
-2
@@ -1,6 +1,7 @@
|
|||||||
use thiserror::Error;
|
use thiserror::Error;
|
||||||
|
|
||||||
#[derive(Error, Debug)]
|
#[derive(Error, Debug)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum PromptError {
|
pub enum PromptError {
|
||||||
#[error("模板解析错误: {0}。请检查模板语法({{var}} / {{#if}} / {{#each}})")]
|
#[error("模板解析错误: {0}。请检查模板语法({{var}} / {{#if}} / {{#each}})")]
|
||||||
Parse(String),
|
Parse(String),
|
||||||
@@ -8,7 +9,9 @@ pub enum PromptError {
|
|||||||
#[error("渲染错误: 变量 '{0}' 未找到。请在 TemplateContext 中插入该变量")]
|
#[error("渲染错误: 变量 '{0}' 未找到。请在 TemplateContext 中插入该变量")]
|
||||||
VariableNotFound(String),
|
VariableNotFound(String),
|
||||||
|
|
||||||
#[error("渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板")]
|
#[error(
|
||||||
|
"渲染错误: 引用的子模板 '{0}' 未注册。请先用 PromptTemplateRegistry::register 注册该子模板"
|
||||||
|
)]
|
||||||
PartialNotFound(String),
|
PartialNotFound(String),
|
||||||
|
|
||||||
#[error("渲染错误: '{0}' 不是数组,无法遍历。请确认传入的是数组或先判空")]
|
#[error("渲染错误: '{0}' 不是数组,无法遍历。请确认传入的是数组或先判空")]
|
||||||
@@ -20,7 +23,9 @@ pub enum PromptError {
|
|||||||
#[error("渲染错误: {0}")]
|
#[error("渲染错误: {0}")]
|
||||||
Render(String),
|
Render(String),
|
||||||
|
|
||||||
#[error("消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)")]
|
#[error(
|
||||||
|
"消息序列校验失败: {0}。请检查消息角色顺序(例如 tool 必须在 assistant tool_call 之后)"
|
||||||
|
)]
|
||||||
InvalidSequence(String),
|
InvalidSequence(String),
|
||||||
|
|
||||||
#[error("文件读取错误: {0}。请检查模板文件路径与权限")]
|
#[error("文件读取错误: {0}。请检查模板文件路径与权限")]
|
||||||
|
|||||||
+16
-29
@@ -1,6 +1,6 @@
|
|||||||
|
use serde_json::Value;
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
use serde_json::Value;
|
|
||||||
|
|
||||||
use crate::prompt::error::PromptError;
|
use crate::prompt::error::PromptError;
|
||||||
|
|
||||||
@@ -140,7 +140,9 @@ fn json_to_template_value(v: &Value) -> Result<TemplateValue, PromptError> {
|
|||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
enum Fragment {
|
enum Fragment {
|
||||||
Literal(String),
|
Literal(String),
|
||||||
Variable { name: String },
|
Variable {
|
||||||
|
name: String,
|
||||||
|
},
|
||||||
If {
|
If {
|
||||||
condition: String,
|
condition: String,
|
||||||
body: Vec<Fragment>,
|
body: Vec<Fragment>,
|
||||||
@@ -223,8 +225,7 @@ fn compile_fragments(template: &str) -> Result<Vec<Fragment>, PromptError> {
|
|||||||
|
|
||||||
let tag = tag_content.trim();
|
let tag = tag_content.trim();
|
||||||
if let Some(rest) = tag.strip_prefix("#if ") {
|
if let Some(rest) = tag.strip_prefix("#if ") {
|
||||||
let (body, else_body, new_i) =
|
let (body, else_body, new_i) = parse_block(template, i, "if")?;
|
||||||
parse_block(template, i, "if")?;
|
|
||||||
let condition = rest.trim().to_string();
|
let condition = rest.trim().to_string();
|
||||||
fragments.push(Fragment::If {
|
fragments.push(Fragment::If {
|
||||||
condition,
|
condition,
|
||||||
@@ -331,10 +332,7 @@ fn parse_block(
|
|||||||
Err(PromptError::Parse(format!("未闭合的 {{#{}}} 块", kind)))
|
Err(PromptError::Parse(format!("未闭合的 {{#{}}} 块", kind)))
|
||||||
}
|
}
|
||||||
|
|
||||||
fn parse_each_block(
|
fn parse_each_block(template: &str, start: usize) -> Result<(Vec<Fragment>, usize), PromptError> {
|
||||||
template: &str,
|
|
||||||
start: usize,
|
|
||||||
) -> Result<(Vec<Fragment>, usize), PromptError> {
|
|
||||||
let bytes = template.as_bytes();
|
let bytes = template.as_bytes();
|
||||||
let len = bytes.len();
|
let len = bytes.len();
|
||||||
let mut depth = 1u32;
|
let mut depth = 1u32;
|
||||||
@@ -368,9 +366,7 @@ fn parse_each_block(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Err(PromptError::Parse(
|
Err(PromptError::Parse("未闭合的 {{#each}} 块".to_string()))
|
||||||
"未闭合的 {{#each}} 块".to_string(),
|
|
||||||
))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), PromptError> {
|
fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), PromptError> {
|
||||||
@@ -395,9 +391,7 @@ fn parse_raw_block(template: &str, start: usize) -> Result<(String, usize), Prom
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Err(PromptError::Parse(
|
Err(PromptError::Parse("未闭合的 {{#raw}} 块".to_string()))
|
||||||
"未闭合的 {{#raw}} 块".to_string(),
|
|
||||||
))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ===== Renderer =====
|
// ===== Renderer =====
|
||||||
@@ -418,33 +412,28 @@ fn render_fragments(
|
|||||||
Fragment::Literal(text) => {
|
Fragment::Literal(text) => {
|
||||||
output.push_str(text);
|
output.push_str(text);
|
||||||
}
|
}
|
||||||
Fragment::Variable { name } => {
|
Fragment::Variable { name } => match ctx.get(name) {
|
||||||
match ctx.get(name) {
|
|
||||||
Some(val) => {
|
Some(val) => {
|
||||||
output.push_str(&format!("{}", val));
|
output.push_str(&format!("{}", val));
|
||||||
}
|
}
|
||||||
None => {
|
None => {
|
||||||
return Err(PromptError::VariableNotFound(name.clone()));
|
return Err(PromptError::VariableNotFound(name.clone()));
|
||||||
}
|
}
|
||||||
}
|
},
|
||||||
}
|
|
||||||
Fragment::If {
|
Fragment::If {
|
||||||
condition,
|
condition,
|
||||||
body,
|
body,
|
||||||
else_body,
|
else_body,
|
||||||
} => {
|
} => {
|
||||||
let truthy = ctx
|
let truthy = ctx.get(condition).map(|v| v.is_truthy()).unwrap_or(false);
|
||||||
.get(condition)
|
|
||||||
.map(|v| v.is_truthy())
|
|
||||||
.unwrap_or(false);
|
|
||||||
let target = if truthy { body } else { else_body };
|
let target = if truthy { body } else { else_body };
|
||||||
render_fragments(target, ctx, partials, output, depth + 1)?;
|
render_fragments(target, ctx, partials, output, depth + 1)?;
|
||||||
}
|
}
|
||||||
Fragment::Each { variable, body } => {
|
Fragment::Each { variable, body } => {
|
||||||
let arr = match ctx.get(variable) {
|
let arr = match ctx.get(variable) {
|
||||||
Some(val) => val.as_array().ok_or_else(|| {
|
Some(val) => val
|
||||||
PromptError::NotAnArray(variable.clone())
|
.as_array()
|
||||||
})?,
|
.ok_or_else(|| PromptError::NotAnArray(variable.clone()))?,
|
||||||
None => {
|
None => {
|
||||||
return Err(PromptError::VariableNotFound(variable.clone()));
|
return Err(PromptError::VariableNotFound(variable.clone()));
|
||||||
}
|
}
|
||||||
@@ -504,10 +493,8 @@ impl PromptTemplateRegistry {
|
|||||||
|
|
||||||
/// 延迟编译注册:只存储原始字符串,首次渲染时编译。
|
/// 延迟编译注册:只存储原始字符串,首次渲染时编译。
|
||||||
pub fn register_lazy(&mut self, name: &str, template: &str) {
|
pub fn register_lazy(&mut self, name: &str, template: &str) {
|
||||||
self.templates.insert(
|
self.templates
|
||||||
name.to_string(),
|
.insert(name.to_string(), StoredTemplate::Raw(template.to_string()));
|
||||||
StoredTemplate::Raw(template.to_string()),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 从文件读取并编译注册。
|
/// 从文件读取并编译注册。
|
||||||
|
|||||||
+10
-3
@@ -4,9 +4,12 @@ use std::sync::Arc;
|
|||||||
|
|
||||||
/// 工具调用过程中可能发生的所有错误。
|
/// 工具调用过程中可能发生的所有错误。
|
||||||
#[derive(thiserror::Error, Debug, Clone)]
|
#[derive(thiserror::Error, Debug, Clone)]
|
||||||
|
#[non_exhaustive]
|
||||||
pub enum ToolError {
|
pub enum ToolError {
|
||||||
/// 工具未注册。不可恢复——需调用方先 `registry.register(...)`。
|
/// 工具未注册。不可恢复——需调用方先 `registry.register(...)`。
|
||||||
#[error("工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写")]
|
#[error(
|
||||||
|
"工具 '{0}' 未注册。请先用 ToolRegistry::register(...) 注册该工具,或检查 LLM 输出的工具名拼写"
|
||||||
|
)]
|
||||||
NotFound(String),
|
NotFound(String),
|
||||||
|
|
||||||
/// 工具执行失败(可恢复——文本回传 LLM 由其决定重试或放弃)。
|
/// 工具执行失败(可恢复——文本回传 LLM 由其决定重试或放弃)。
|
||||||
@@ -14,11 +17,15 @@ pub enum ToolError {
|
|||||||
ExecutionFailed(String, String),
|
ExecutionFailed(String, String),
|
||||||
|
|
||||||
/// 工具参数无效(可恢复——文本回传 LLM)。
|
/// 工具参数无效(可恢复——文本回传 LLM)。
|
||||||
#[error("工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema")]
|
#[error(
|
||||||
|
"工具 '{0}' 参数无效: {1}。请检查 LLM 输出的参数是否符合 BaseTool::parameters() 声明的 JSON Schema"
|
||||||
|
)]
|
||||||
InvalidArguments(String, String),
|
InvalidArguments(String, String),
|
||||||
|
|
||||||
/// 权限被拒绝(不可恢复——终止循环)。
|
/// 权限被拒绝(不可恢复——终止循环)。
|
||||||
#[error("权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过")]
|
#[error(
|
||||||
|
"权限被拒绝: 工具 '{0}' 需要 {1} 权限。请在 PermissionConfig 中显式允许,或人工确认后绕过"
|
||||||
|
)]
|
||||||
PermissionDenied(String, String),
|
PermissionDenied(String, String),
|
||||||
|
|
||||||
/// MCP 协议错误(不可恢复)。
|
/// MCP 协议错误(不可恢复)。
|
||||||
|
|||||||
+17
-37
@@ -9,19 +9,18 @@
|
|||||||
|
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::process::Stdio;
|
use std::process::Stdio;
|
||||||
use std::sync::atomic::{AtomicBool, Ordering};
|
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
use std::sync::atomic::{AtomicBool, Ordering};
|
||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::{json, Value};
|
use serde_json::{Value, json};
|
||||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
||||||
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
|
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
|
||||||
use tokio::sync::{oneshot, Mutex};
|
use tokio::sync::{Mutex, oneshot};
|
||||||
|
|
||||||
#[allow(deprecated)]
|
use crate::llm::types::tool::ToolDef;
|
||||||
use crate::llm::types::ToolDefinition;
|
|
||||||
use crate::tools::base::{BaseTool, ToolContext, ToolRef};
|
use crate::tools::base::{BaseTool, ToolContext, ToolRef};
|
||||||
use crate::tools::error::ToolError;
|
use crate::tools::error::ToolError;
|
||||||
|
|
||||||
@@ -136,7 +135,6 @@ impl std::fmt::Debug for McpClient {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(deprecated)]
|
|
||||||
impl McpClient {
|
impl McpClient {
|
||||||
/// 创建一个 MCP 客户端。
|
/// 创建一个 MCP 客户端。
|
||||||
pub fn new(server_name: impl Into<String>, transport: McpTransport) -> Self {
|
pub fn new(server_name: impl Into<String>, transport: McpTransport) -> Self {
|
||||||
@@ -226,9 +224,7 @@ impl McpClient {
|
|||||||
"version": env!("CARGO_PKG_VERSION")
|
"version": env!("CARGO_PKG_VERSION")
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
let _response = self
|
let _response = self.send_request("initialize", Some(init_params)).await?;
|
||||||
.send_request("initialize", Some(init_params))
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
// 发送 initialized 通知(无 id)
|
// 发送 initialized 通知(无 id)
|
||||||
self.send_notification("notifications/initialized", Some(json!({})))
|
self.send_notification("notifications/initialized", Some(json!({})))
|
||||||
@@ -239,7 +235,7 @@ impl McpClient {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// 列出服务器支持的工具(调用 `tools/list`)。
|
/// 列出服务器支持的工具(调用 `tools/list`)。
|
||||||
pub async fn list_tools(&mut self) -> Result<Vec<ToolDefinition>, ToolError> {
|
pub async fn list_tools(&mut self) -> Result<Vec<ToolDef>, ToolError> {
|
||||||
if !self.is_initialized() {
|
if !self.is_initialized() {
|
||||||
return Err(ToolError::McpNotInitialized(self.server_name.clone()));
|
return Err(ToolError::McpNotInitialized(self.server_name.clone()));
|
||||||
}
|
}
|
||||||
@@ -274,11 +270,10 @@ impl McpClient {
|
|||||||
description: description.clone(),
|
description: description.clone(),
|
||||||
input_schema: input_schema.clone(),
|
input_schema: input_schema.clone(),
|
||||||
});
|
});
|
||||||
defs.push(ToolDefinition {
|
defs.push(ToolDef {
|
||||||
name,
|
name,
|
||||||
description,
|
description,
|
||||||
parameters: input_schema,
|
parameters: input_schema,
|
||||||
strict: None,
|
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
Ok(defs)
|
Ok(defs)
|
||||||
@@ -337,11 +332,7 @@ impl McpClient {
|
|||||||
if let Some(state) = self.process.take() {
|
if let Some(state) = self.process.take() {
|
||||||
let mut state = state.lock().await;
|
let mut state = state.lock().await;
|
||||||
// 优雅等待 5 秒
|
// 优雅等待 5 秒
|
||||||
let graceful = tokio::time::timeout(
|
let graceful = tokio::time::timeout(Duration::from_secs(5), state.child.wait()).await;
|
||||||
Duration::from_secs(5),
|
|
||||||
state.child.wait(),
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
if graceful.is_err() {
|
if graceful.is_err() {
|
||||||
// 超时则强杀
|
// 超时则强杀
|
||||||
let _ = state.child.kill().await;
|
let _ = state.child.kill().await;
|
||||||
@@ -372,11 +363,7 @@ impl McpClient {
|
|||||||
tools
|
tools
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn send_request(
|
async fn send_request(&self, method: &str, params: Option<Value>) -> Result<Value, ToolError> {
|
||||||
&self,
|
|
||||||
method: &str,
|
|
||||||
params: Option<Value>,
|
|
||||||
) -> Result<Value, ToolError> {
|
|
||||||
let state_arc = self
|
let state_arc = self
|
||||||
.process
|
.process
|
||||||
.as_ref()
|
.as_ref()
|
||||||
@@ -412,9 +399,11 @@ impl McpClient {
|
|||||||
.write_all(b"\n")
|
.write_all(b"\n")
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?;
|
.map_err(|e| ToolError::McpError(format!("写入换行失败: {e}")))?;
|
||||||
state.stdin.flush().await.map_err(|e| {
|
state
|
||||||
ToolError::McpError(format!("flush stdin 失败: {e}"))
|
.stdin
|
||||||
})?;
|
.flush()
|
||||||
|
.await
|
||||||
|
.map_err(|e| ToolError::McpError(format!("flush stdin 失败: {e}")))?;
|
||||||
}
|
}
|
||||||
|
|
||||||
// 等待响应(带超时)
|
// 等待响应(带超时)
|
||||||
@@ -471,10 +460,7 @@ impl McpClient {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// 持续读取 stdout,将响应分发到对应的 oneshot sender。
|
/// 持续读取 stdout,将响应分发到对应的 oneshot sender。
|
||||||
async fn read_loop(
|
async fn read_loop(mut reader: BufReader<ChildStdout>, state: Arc<Mutex<ChildProcessState>>) {
|
||||||
mut reader: BufReader<ChildStdout>,
|
|
||||||
state: Arc<Mutex<ChildProcessState>>,
|
|
||||||
) {
|
|
||||||
let mut line = String::new();
|
let mut line = String::new();
|
||||||
loop {
|
loop {
|
||||||
line.clear();
|
line.clear();
|
||||||
@@ -542,7 +528,6 @@ enum McpClientHandle {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
#[allow(deprecated)]
|
|
||||||
impl BaseTool for McpToolAdapter {
|
impl BaseTool for McpToolAdapter {
|
||||||
fn name(&self) -> &str {
|
fn name(&self) -> &str {
|
||||||
&self.name
|
&self.name
|
||||||
@@ -556,11 +541,7 @@ impl BaseTool for McpToolAdapter {
|
|||||||
self.parameters.clone()
|
self.parameters.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn execute(
|
async fn execute(&self, _args: Value, _ctx: &ToolContext<'_>) -> Result<Value, ToolError> {
|
||||||
&self,
|
|
||||||
_args: Value,
|
|
||||||
_ctx: &ToolContext<'_>,
|
|
||||||
) -> Result<Value, ToolError> {
|
|
||||||
// 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。
|
// 当前 Phase 2 实现的简化:McpToolAdapter 不持有活跃 MCP 连接。
|
||||||
// 实际生产中应持有 Arc<McpClient> 并通过 mcp.call_tool() 执行。
|
// 实际生产中应持有 Arc<McpClient> 并通过 mcp.call_tool() 执行。
|
||||||
// 这里返回错误,提示需要通过其他方式调用 MCP 工具。
|
// 这里返回错误,提示需要通过其他方式调用 MCP 工具。
|
||||||
@@ -617,8 +598,7 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_jsonrpc_response_parse_error() {
|
fn test_jsonrpc_response_parse_error() {
|
||||||
let s =
|
let s = r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
|
||||||
r#"{"jsonrpc":"2.0","id":1,"error":{"code":-32601,"message":"Method not found"}}"#;
|
|
||||||
let resp: JsonRpcResponse = serde_json::from_str(s).unwrap();
|
let resp: JsonRpcResponse = serde_json::from_str(s).unwrap();
|
||||||
assert_eq!(resp.id, 1);
|
assert_eq!(resp.id, 1);
|
||||||
assert!(resp.result.is_none());
|
assert!(resp.result.is_none());
|
||||||
|
|||||||
+18
-15
@@ -148,9 +148,7 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn test_default_config_denies_delete() {
|
fn test_default_config_denies_delete() {
|
||||||
let checker = PermissionChecker::new(PermissionConfig::default());
|
let checker = PermissionChecker::new(PermissionConfig::default());
|
||||||
assert!(checker
|
assert!(checker.check("rm_file", &p(Permission::Delete)).is_err());
|
||||||
.check("rm_file", &p(Permission::Delete))
|
|
||||||
.is_err());
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -246,12 +244,16 @@ mod tests {
|
|||||||
allow_unspecified: false,
|
allow_unspecified: false,
|
||||||
};
|
};
|
||||||
let checker = PermissionChecker::new(cfg);
|
let checker = PermissionChecker::new(cfg);
|
||||||
assert!(checker
|
assert!(
|
||||||
|
checker
|
||||||
.check("t", &[Permission::Custom("db:read".into())])
|
.check("t", &[Permission::Custom("db:read".into())])
|
||||||
.is_ok());
|
.is_ok()
|
||||||
assert!(checker
|
);
|
||||||
|
assert!(
|
||||||
|
checker
|
||||||
.check("t", &[Permission::Custom("db:write".into())])
|
.check("t", &[Permission::Custom("db:write".into())])
|
||||||
.is_err());
|
.is_err()
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -262,12 +264,11 @@ mod tests {
|
|||||||
allow_unspecified: false,
|
allow_unspecified: false,
|
||||||
};
|
};
|
||||||
let checker = PermissionChecker::new(cfg);
|
let checker = PermissionChecker::new(cfg);
|
||||||
assert!(checker
|
assert!(
|
||||||
.check(
|
checker
|
||||||
"t",
|
.check("t", &[Permission::Read, Permission::Network])
|
||||||
&[Permission::Read, Permission::Network]
|
.is_ok()
|
||||||
)
|
);
|
||||||
.is_ok());
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -279,8 +280,10 @@ mod tests {
|
|||||||
};
|
};
|
||||||
let checker = PermissionChecker::new(cfg);
|
let checker = PermissionChecker::new(cfg);
|
||||||
// 任一权限不在白名单则拒绝
|
// 任一权限不在白名单则拒绝
|
||||||
assert!(checker
|
assert!(
|
||||||
|
checker
|
||||||
.check("t", &[Permission::Read, Permission::Write])
|
.check("t", &[Permission::Read, Permission::Write])
|
||||||
.is_err());
|
.is_err()
|
||||||
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+10
-10
@@ -7,8 +7,7 @@ use std::time::Duration;
|
|||||||
use futures::future::join_all;
|
use futures::future::join_all;
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
#[allow(deprecated)]
|
use crate::llm::types::tool::ToolDef;
|
||||||
use crate::llm::types::ToolDefinition;
|
|
||||||
use crate::tools::base::{ToolContext, ToolRef};
|
use crate::tools::base::{ToolContext, ToolRef};
|
||||||
use crate::tools::error::ToolError;
|
use crate::tools::error::ToolError;
|
||||||
use crate::tools::permission::PermissionChecker;
|
use crate::tools::permission::PermissionChecker;
|
||||||
@@ -71,7 +70,6 @@ impl std::fmt::Debug for ToolRegistry {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(deprecated)]
|
|
||||||
impl ToolRegistry {
|
impl ToolRegistry {
|
||||||
/// 创建一个新的工具注册表。
|
/// 创建一个新的工具注册表。
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
@@ -127,16 +125,15 @@ impl ToolRegistry {
|
|||||||
self.inner.tools.keys().cloned().collect()
|
self.inner.tools.keys().cloned().collect()
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 获取所有工具的 `ToolDefinition` 列表(用于传递给 LLM)。
|
/// 获取所有工具的 `ToolDef` 列表(用于传递给 LLM)。
|
||||||
pub fn definitions(&self) -> Vec<ToolDefinition> {
|
pub fn definitions(&self) -> Vec<ToolDef> {
|
||||||
self.inner
|
self.inner
|
||||||
.tools
|
.tools
|
||||||
.values()
|
.values()
|
||||||
.map(|tool| ToolDefinition {
|
.map(|tool| ToolDef {
|
||||||
name: tool.name().to_string(),
|
name: tool.name().to_string(),
|
||||||
description: Some(tool.description().to_string()),
|
description: Some(tool.description().to_string()),
|
||||||
parameters: tool.parameters(),
|
parameters: tool.parameters(),
|
||||||
strict: None,
|
|
||||||
})
|
})
|
||||||
.collect()
|
.collect()
|
||||||
}
|
}
|
||||||
@@ -348,7 +345,10 @@ mod tests {
|
|||||||
async fn test_invoke_success() {
|
async fn test_invoke_success() {
|
||||||
let mut reg = ToolRegistry::new();
|
let mut reg = ToolRegistry::new();
|
||||||
reg.register(Arc::new(AddTool { base: 100 })).unwrap();
|
reg.register(Arc::new(AddTool { base: 100 })).unwrap();
|
||||||
let result = reg.invoke("call_1", "add", json!({ "n": 5 })).await.unwrap();
|
let result = reg
|
||||||
|
.invoke("call_1", "add", json!({ "n": 5 }))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
let value = result.output.unwrap();
|
let value = result.output.unwrap();
|
||||||
assert_eq!(value["result"], 105);
|
assert_eq!(value["result"], 105);
|
||||||
assert_eq!(result.tool_call_id, "call_1");
|
assert_eq!(result.tool_call_id, "call_1");
|
||||||
@@ -372,8 +372,8 @@ mod tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_invoke_with_permission_denied() {
|
async fn test_invoke_with_permission_denied() {
|
||||||
let mut reg = ToolRegistry::new()
|
let mut reg =
|
||||||
.with_permission_checker(PermissionChecker::new(Default::default()));
|
ToolRegistry::new().with_permission_checker(PermissionChecker::new(Default::default()));
|
||||||
reg.register(Arc::new(ShellTool)).unwrap();
|
reg.register(Arc::new(ShellTool)).unwrap();
|
||||||
let result = reg.invoke("call_z", "shell", json!({})).await;
|
let result = reg.invoke("call_z", "shell", json!({})).await;
|
||||||
assert!(matches!(result, Err(ToolError::PermissionDenied(_, _))));
|
assert!(matches!(result, Err(ToolError::PermissionDenied(_, _))));
|
||||||
|
|||||||
Reference in New Issue
Block a user