Compare commits
23
Commits
4348e4bf3e
...
v0.3.0
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
28d6a1c166 | ||
|
|
6676322666 | ||
|
|
9d73f525d0 | ||
|
|
7e72e102a2 | ||
|
|
5bb349d177 | ||
|
|
46de111965 | ||
|
|
cb922b03de | ||
|
|
34eec9f546 | ||
|
|
1d51dcdfe0 | ||
|
|
fbbf8bf6e5 | ||
|
|
cc1c68b69d | ||
|
|
209932e3b5 | ||
|
|
32d886f870 | ||
|
|
b04427e83f | ||
|
|
d4c4d8fa3c | ||
|
|
4686063ca8 | ||
|
|
c36668071e | ||
|
|
d4f27b5865 | ||
|
|
1c0e1e0ed1 | ||
|
|
760de46623 | ||
|
|
f8df6a9421 | ||
|
|
802518b5fe | ||
|
|
993118f661 |
@@ -2,6 +2,69 @@
|
||||
|
||||
本项目所有重要变更均记录于此文件。格式参考 [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 个面向新用户的集成示例。
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "agcore"
|
||||
version = "0.2.0-rc.1"
|
||||
version = "0.3.0"
|
||||
edition = "2024"
|
||||
|
||||
[dependencies]
|
||||
|
||||
@@ -241,6 +241,12 @@ pub async fn submit_turn_stream(
|
||||
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
|
||||
@@ -410,10 +416,10 @@ if let Some(response) = final_response {
|
||||
|
||||
**目标**:端到端验证 `submit_turn_stream` + `finalize_turn` 的完整链路,确保零回归。
|
||||
|
||||
**新增**(`agent/session.rs` 内联测试):
|
||||
**新增**(`agent/session.rs` 内联测试,2026-07-08 实施审查补全):
|
||||
|
||||
- **场景**:`submit_turn_stream` 跑通 mock provider → 消费流(验证各事件到达) → `finalize_turn` 后 cost 更新正确
|
||||
- **场景**:verify `OnTurnStart` hook 在 `submit_turn_stream` 返回流之前已触发
|
||||
- `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)
|
||||
|
||||
**验证**:
|
||||
|
||||
|
||||
@@ -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,471 @@
|
||||
# Phase 16 — 摘要自动生成
|
||||
|
||||
## 背景与目标
|
||||
|
||||
### 问题
|
||||
|
||||
长对话场景中,用户与 Agent 交互 30+ 轮后,消息历史长度远超模型上下文窗口,导致:
|
||||
|
||||
- LLM 被迫丢弃早期上下文,对话丧失连贯性
|
||||
- 开发者需要手动管理摘要逻辑(调 LLM → 写 SessionMemory → 注入 FocusedConfig)
|
||||
- v0.2 的 `FocusedConfig.summary_override` 消费端已就绪,但生产端是空的——用户只能手动设字符串
|
||||
|
||||
### 目标
|
||||
|
||||
闭环长对话的"上下文压缩"链路:
|
||||
|
||||
```
|
||||
[消费端 v0.2 已就绪] FocusedConfig.summary_override → filter_focused() 注入摘要
|
||||
[生产端 v0.3 补齐] token 水位检测 → LLM 摘要生成 → 自动写入 summary_override
|
||||
```
|
||||
|
||||
### 成功标准
|
||||
|
||||
1. 开发者只需在 `AgentBuilder` 中链式调用 `.summary_config(cfg)` 即可启用
|
||||
2. 长对话(如 30+ 轮或 token 水位超过 `max_context_tokens * trigger_token_ratio`)自动触发摘要,下轮 `load_messages()` 返回值包含 `[上下文摘要] {summary}`
|
||||
3. 摘要生成不改变 `submit_turn` 行为(opt-in、静默失败、不阻断主流程)
|
||||
4. 零新外部依赖
|
||||
|
||||
---
|
||||
|
||||
## 需求分析
|
||||
|
||||
### 功能需求
|
||||
|
||||
| # | 需求 | 优先级 | 说明 |
|
||||
|---|------|--------|------|
|
||||
| F1 | `SummaryConfig` 配置结构体 | P0 | `trigger_token_ratio` / `max_context_tokens` / `summary_prompt` / `debounce_turns` / `summary_model` / `max_tool_result_chars` |
|
||||
| F2 | Token 水位自动检测 | P0 | 每轮 OnTurnEnd 之后检查 `cost_so_far` 是否超过 `max * ratio` |
|
||||
| F3 | LLM 摘要生成 | P0 | 复用 `self.bundle.provider`,单次无工具 LLM 调用 |
|
||||
| F4 | 摘要写入 FocusedConfig | P0 | 更新 `summary_override` + `slot.save()` 持久化 |
|
||||
| F5 | 摘要全局快照 | P0 | 同步写入 `SessionMemory::set("conversation_summary", summary)` |
|
||||
| F6 | 防抖机制 | P0 | 两次摘要之间至少间隔 `debounce_turns` 轮(默认 3) |
|
||||
| F7 | 流式路径对称支持 | P0 | `finalize_turn` 中插入相同检查点 |
|
||||
| F8 | 公开 API:`get_conversation_summary()` | P1 | 读取 SessionMemory 中最新的摘要 |
|
||||
|
||||
### 非功能需求
|
||||
|
||||
| # | 需求 | 指标 |
|
||||
|---|------|------|
|
||||
| N1 | 零外部依赖 | 不修改 `Cargo.toml` |
|
||||
| N2 | 向后兼容 | 未设置 `SummaryConfig` 时行为零变化 |
|
||||
| N3 | 静默失败 | 摘要 LLM 调用失败不阻断 `submit_turn` |
|
||||
| N4 | 摘要延迟 | 首次摘要 LLM 调用 ≤ 3s(依赖 provider 响应速度) |
|
||||
|
||||
---
|
||||
|
||||
## 当前状态分析
|
||||
|
||||
### 消费端已就绪
|
||||
|
||||
`FocusedConfig.summary_override`(`src/agent/context.rs`)已在 Phase 10 实现,当前消费逻辑:
|
||||
|
||||
```
|
||||
filter_focused() → 若 cfg.summary_override = Some(text) → 在消息列表末尾插入
|
||||
Message::system("[上下文摘要] {text}")
|
||||
```
|
||||
|
||||
文档注释明确标注:`// v0.3 将支持 Hook 驱动的自动摘要生成`
|
||||
|
||||
### 代码上下文
|
||||
|
||||
| 模块 | 文件 | 状态 | 与 Phase 16 的关系 |
|
||||
|------|------|------|-------------------|
|
||||
| FocusedConfig | `agent/context.rs` | ✅ 消费端 | 摘要写入 `summary_override` 即生效 |
|
||||
| OnTurnEnd | `agent/session.rs:345` | ✅ 触发点 | 摘要检查点插在此之后 |
|
||||
| CostTracker | `llm/cycle/usage.rs` | ✅ 累计 token | 水位检测的数据源 |
|
||||
| SessionMemory | `agent/session_memory.rs` | ✅ set/get | 摘要全局快照存储 |
|
||||
| AgentBuilder | `agent/builder.rs` | ✅ 链式构造 | 新增 `.summary_config()` |
|
||||
| AgentConfig | `agent/runtime.rs` | ✅ 配置结构 | 新增 `summary_config` 字段 |
|
||||
| LlmProvider | `llm/provider.rs` | ✅ Trait | 摘要 LLM 调用复用 provider |
|
||||
| ContextSlot.save | `agent/context.rs:251` | ✅ 持久化 | 更新 config 后写回 |
|
||||
|
||||
---
|
||||
|
||||
## 可选方案推演
|
||||
|
||||
### 方案 A(推荐):内联检查点
|
||||
|
||||
**做法**:在 `submit_turn` 和 `finalize_turn` 中,OnTurnEnd 触发之后、`turn_index` 递增之前,插入以下逻辑:
|
||||
|
||||
```rust
|
||||
if let Some(ref sc) = self.bundle.config.summary_config
|
||||
&& self.should_summarize(sc)
|
||||
{
|
||||
// clone 所需数据(释放 &self 借用)
|
||||
let provider = Arc::clone(&self.bundle.provider);
|
||||
let messages = self.slots.get(&self.current_slot_id)
|
||||
.map(|s| s.messages.clone()).unwrap_or_default();
|
||||
let prompt = sc.summary_prompt.clone();
|
||||
let model = sc.summary_model.clone();
|
||||
|
||||
// 调关联函数(不持有 &self)
|
||||
match Self::generate_summary(&provider, &messages, &prompt, model.as_deref(), sc.max_tool_result_chars).await {
|
||||
Ok(text) => {
|
||||
// 更新 FocusedConfig + 持久化
|
||||
if let Some(slot) = self.slots.get_mut(&self.current_slot_id) {
|
||||
if let SlotMode::Focused(ref mut cfg) = slot.config.mode {
|
||||
cfg.summary_override = Some(text.clone());
|
||||
}
|
||||
let _ = slot.save(&*self.resolve_store()).await;
|
||||
}
|
||||
// 全局快照
|
||||
let _ = self.session_memory.set("conversation_summary", &text).await;
|
||||
self.last_summary_turn = self.turn_index;
|
||||
}
|
||||
Err(e) => tracing::error!("摘要自动生成失败 (turn={}): {}", self.turn_index, e),
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**优点**:
|
||||
- 代码路径最短最清晰(~50 行核心逻辑)
|
||||
- 直接访问所有需要的数据(`cost_so_far`、`slots`、`provider`、`session_memory`)
|
||||
- 流式和同步版本统一处理
|
||||
- `Option<SummaryConfig>` 本身已提供 opt-in/opt-out
|
||||
- 不改变 Hook 系统签名
|
||||
|
||||
**缺点**:
|
||||
- 摘要 LLM 调用延长了 `submit_turn` 的延迟(约 1-3s)
|
||||
- 违反"Hook 哲学"(但 `Option` 配置已足够提供可插拔性)
|
||||
|
||||
### 方案 B(否决):扩展 HookContext
|
||||
|
||||
**做法**:在 `HookContext` 中增加 `messages: &[Message]`、`usage: &Usage`、`provider: Arc<dyn LlmProvider>` 字段,让 OnTurnEnd Hook 实现者自行做摘要。
|
||||
|
||||
**否决原因**:
|
||||
1. **生命周期冲突**:`&[Message]` 要求 Hook 调用点消息已就绪但未被 `&mut self` 借用——在 `submit_turn` 第 7 步(slot.save)后消息已就绪,但 to pass `&[Message]` 到 HookContext 需要与 `slot.messages` 的不可变引用共存,而 `submit_turn` 流程中后续步骤需要 `&mut self`
|
||||
2. **流式路径不可行**:`finalize_turn` 触发 OnTurnEnd 时 cycle 已销毁,消息只能从 slot 获取,但 slot 在 `append_messages` 后已被 `&mut` 借用
|
||||
3. **`Arc<dyn LlmProvider>` 的 `'static` 需求**与 `HookContext<'a>` 的设计冲突
|
||||
|
||||
### 方案 C(否决):后台 spawn 异步摘要
|
||||
|
||||
**做法**:token 检测通过后,`tokio::spawn` 后台任务做摘要生成和写入。
|
||||
|
||||
**否决原因**:
|
||||
1. **写入冲突**:后台任务无法获取 `&mut AgentSession` 来更新 slot config
|
||||
2. **绕过方式增加复杂度**:后台任务需要直接操作 `Arc<dyn MemoryStore>` 的原始 key(`slot_config:{session_id}:{slot_id}`),绕过了 `ContextSlot::save()` 的封装
|
||||
3. **并发风险**:如果前一轮摘要尚未完成而下一轮 `finalize_turn` 又触发,可能导致覆盖写
|
||||
|
||||
---
|
||||
|
||||
## 推荐方案(内联检查点)
|
||||
|
||||
### 架构图
|
||||
|
||||
```
|
||||
submit_turn(user_input)
|
||||
│
|
||||
├─ 1. Readonly 检查
|
||||
├─ 2. OnTurnStart hook
|
||||
├─ 3. slot.load_messages() ← 历史摘要已注入(如有)
|
||||
├─ 4. LlmCycle.submit_with_tools
|
||||
├─ 5. cost_so_far.add(usage)
|
||||
├─ 6. slot.append_messages + save
|
||||
├─ 7. OnTurnEnd hook ← 纯通知,不做摘要
|
||||
│
|
||||
├─ [8.5] 摘要检查点 ──────────────────────────────┐
|
||||
│ ├─ should_summarize(cfg) │
|
||||
│ │ ├─ cost_so_far >= max * ratio? │
|
||||
│ │ └─ turn - last_summary >= debounce? │
|
||||
│ │ │
|
||||
│ ├─ generate_summary() ← 新 LlmCycle │
|
||||
│ │ ├─ format_messages_as_text() │
|
||||
│ │ ├─ replace {messages} │
|
||||
│ │ └─ submit_messages(无 tools) │
|
||||
│ │ │
|
||||
│ └─ 成功 → 更新 summary_override + save │
|
||||
│ → SessionMemory.set() │
|
||||
│ → last_summary_turn = turn_index │
|
||||
│ (流式路径用 saturating_sub(1) 修正) │
|
||||
│ 失败 → tracing::error! 静默 │
|
||||
│ │
|
||||
├─ 9. turn_index++
|
||||
└─ 10. return Ok(response)
|
||||
```
|
||||
|
||||
### 模块划分
|
||||
|
||||
**新增文件**:`src/agent/summary.rs`
|
||||
|
||||
```
|
||||
src/agent/summary.rs
|
||||
├── SummaryConfig // 摘要自动生成配置
|
||||
├── format_messages_as_text() // 消息 → 纯文本(简洁版)
|
||||
└── DEFAULT_SUMMARY_PROMPT // 默认 prompt 模板
|
||||
```
|
||||
|
||||
**修改文件**:
|
||||
|
||||
| 文件 | 改动 |
|
||||
|------|------|
|
||||
| `agent/runtime.rs` | `AgentConfig` 新增 `summary_config: Option<SummaryConfig>` |
|
||||
| `agent/builder.rs` | 新增 `summary_config(cfg)` 方法 |
|
||||
| `agent/session.rs` | 新增 `last_summary_turn` 字段;`submit_turn` / `finalize_turn` 插入检查点;关联函数 `generate_summary`;`get_conversation_summary()` |
|
||||
| `agent.rs` | `pub mod summary` + re-export |
|
||||
|
||||
**不变的文件**(无需改动):
|
||||
|
||||
| 文件 | 原因 |
|
||||
|------|------|
|
||||
| `llm/hooks.rs` | 内联方案不扩展 HookContext |
|
||||
| `llm/cycle.rs` | 摘要调用通过 `submit_messages` 独立使用 |
|
||||
| `agent/context.rs` | `FocusedConfig` 消费端已在 Phase 10 就绪 |
|
||||
| `Cargo.toml` | 零新外部依赖 |
|
||||
|
||||
### 核心接口定义
|
||||
|
||||
**`SummaryConfig`**(`agent/summary.rs`):
|
||||
|
||||
```rust
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SummaryConfig {
|
||||
/// Token 水位触发比例(0.0 ~ 1.0)。默认 0.75。
|
||||
pub trigger_token_ratio: f64,
|
||||
/// 模型上下文窗口大小(token)。默认 32_000,覆盖大部分开源模型。
|
||||
/// 修改为匹配实际使用模型的上下文窗口。
|
||||
/// ⚠️ 设置为超过模型窗口的值会导致摘要永远不触发。
|
||||
pub max_context_tokens: u32,
|
||||
/// 摘要 prompt 模板。`{messages}` 将被替换为对话历史文本。
|
||||
pub summary_prompt: String,
|
||||
/// 防抖轮次。默认 3。
|
||||
pub debounce_turns: u32,
|
||||
/// 摘要生成模型(None = 沿用主 provider 默认模型)。
|
||||
/// 默认 None。推荐设为便宜模型(如 "gpt-4o-mini")以节省成本。
|
||||
pub summary_model: Option<String>,
|
||||
/// 单个 ToolResult 在格式化时保留的最大字符数。默认 500。
|
||||
/// 超过此值从尾部截断。字符级安全(`chars().take()`)。
|
||||
pub max_tool_result_chars: usize,
|
||||
}
|
||||
```
|
||||
|
||||
**`generate_summary`**(`AgentSession` 关联函数):
|
||||
|
||||
```rust
|
||||
impl AgentSession {
|
||||
async fn generate_summary(
|
||||
provider: &Arc<dyn LlmProvider>,
|
||||
messages: &[Message],
|
||||
prompt_template: &str,
|
||||
summary_model: Option<&str>,
|
||||
max_tool_result_chars: usize,
|
||||
) -> Result<String, LlmError> { ... }
|
||||
}
|
||||
```
|
||||
|
||||
**`should_summarize`**(`AgentSession` 方法):
|
||||
|
||||
```rust
|
||||
fn should_summarize(&self, cfg: &SummaryConfig) -> bool {
|
||||
self.turn_index - self.last_summary_turn >= cfg.debounce_turns
|
||||
&& self.cost_so_far.total().total_tokens as f64
|
||||
>= cfg.max_context_tokens as f64 * cfg.trigger_token_ratio
|
||||
}
|
||||
```
|
||||
|
||||
### 消息格式化(简洁版)
|
||||
|
||||
`format_messages_as_text` 输出格式:
|
||||
|
||||
```
|
||||
System: 你是一个翻译助手
|
||||
User: 把这段英文翻译成中文
|
||||
Assistant: 请提供英文文本 [Tool: translate]
|
||||
Tool Result: 这是中文翻译
|
||||
User: 谢谢
|
||||
Assistant: 不客气
|
||||
```
|
||||
|
||||
处理规则:
|
||||
- `ContentBlock::Text { text }` → 直接拼接
|
||||
- `ContentBlock::ToolUse { name, .. }` → `[Tool: {name}]`(不显示参数 JSON)
|
||||
- `Message::ToolResult { content, is_error, tool_call_id }` → `Tool Result [{tool_call_id}]:` / `Tool Error [{tool_call_id}]:`,便于多工具场景下关联调用的返回
|
||||
- ToolResult 文本截断到前 `max_tool_result_chars` 个 Unicode 字符(`chars().take(n)`,字符级安全,避免多字节截断)
|
||||
- 整段对话若超过 30K 字符,从前面截断(优先保留最新消息)
|
||||
- `Message::UserImage { .. }` → `User: [image]`
|
||||
- 非 Text block(Image / Audio / File 等)统一标记为 `[{kind}]`
|
||||
- 每条消息一行,空行分隔
|
||||
|
||||
---
|
||||
|
||||
## 实现计划
|
||||
|
||||
### Step 16.1 — `SummaryConfig` 结构体
|
||||
|
||||
**文件**:新增 `src/agent/summary.rs`
|
||||
|
||||
**内容**:
|
||||
- `SummaryConfig` 结构体定义(6 个字段 + doc comments)
|
||||
- `DEFAULT_SUMMARY_PROMPT` 常量(约 100 字中文 prompt,含 `{messages}` 占位符)
|
||||
- `impl Default for SummaryConfig`
|
||||
- `format_messages_as_text(messages: &[Message]) -> String` 辅助函数
|
||||
|
||||
**验证**:`cargo build`
|
||||
|
||||
### Step 16.2 — `AgentConfig` 扩展 + `AgentBuilder` 方法
|
||||
|
||||
**文件**:`src/agent/runtime.rs` + `src/agent/builder.rs`
|
||||
|
||||
**改动**:
|
||||
- `AgentConfig` 新增字段:`pub summary_config: Option<SummaryConfig>`
|
||||
- `AgentBuilder` 新增方法:
|
||||
```rust
|
||||
pub fn summary_config(mut self, cfg: SummaryConfig) -> Self {
|
||||
let mut config = self.config.take().unwrap_or_default();
|
||||
config.summary_config = Some(cfg);
|
||||
self.config = Some(config);
|
||||
self
|
||||
}
|
||||
```
|
||||
|
||||
**验证**:`AgentBuilder` 单元测试 + `cargo test`
|
||||
|
||||
### Step 16.3 — `AgentSession` 新字段 + 检查点
|
||||
|
||||
**文件**:`src/agent/session.rs`
|
||||
|
||||
**改动**:
|
||||
|
||||
1. `AgentSession` 新增字段:`last_summary_turn: u32`(初始化 0)
|
||||
2. `submit_turn` 中 OnTurnEnd 之后、turn_index 之前插入检查点
|
||||
3. `finalize_turn` 中 OnTurnEnd 之后插入对称检查点。注意:流式路径中 `turn_index` 已在 `submit_turn_stream` 中递增,检查点赋值使用 `self.turn_index.saturating_sub(1)`(与 `OnTurnEnd` hook 保持一致)。
|
||||
4. 关联函数 `generate_summary`:
|
||||
- 接收 `provider`、`messages`、`prompt_template`、`summary_model`、`max_tool_result_chars`
|
||||
- 入口守卫:`messages.is_empty()` 时直接返回 `Ok(String::new())`
|
||||
- 构造 `LlmCycle`(`max_tokens = Some(1024)`)
|
||||
- 调 `cycle.submit_messages(vec![Message::user_text(prompt)], vec![])`
|
||||
- 提取 text 返回
|
||||
5. 公开 API:`get_conversation_summary()` → `self.session_memory.get("conversation_summary")`
|
||||
|
||||
**验证**:`cargo build --all-targets`
|
||||
|
||||
### Step 16.4 — re-export
|
||||
|
||||
**文件**:`src/agent.rs`
|
||||
|
||||
**改动**:
|
||||
```rust
|
||||
pub mod summary;
|
||||
pub use summary::SummaryConfig;
|
||||
```
|
||||
|
||||
**验证**:`cargo test --all-targets`
|
||||
|
||||
### Step 16.5 — 测试
|
||||
|
||||
| 测试 | 验证点 | 方式 |
|
||||
|------|--------|------|
|
||||
| `summary_config_defaults` | 默认值正确 | 单元测试 |
|
||||
| `summary_not_generated_below_threshold` | token < 阈值时不触发 | `MockProvider` + `Usage::from_input_output(10, 5)` |
|
||||
| `summary_generated_above_threshold` | token ≥ 阈值时触发 | 设置 `max_context_tokens=20` + `trigger_token_ratio=0.5` |
|
||||
| `summary_debounce_works` | debounce 内不重复 | 强行触发摘要后验证 3 轮内不触发 |
|
||||
| `summary_injected_into_focused` | Focused 模式 `load_messages()` 含 `[上下文摘要]` | 检查 Message 内容 |
|
||||
| `summary_written_to_session_memory` | `get_session_data("conversation_summary")` 有值 | 集成测试 |
|
||||
| `summary_not_injected_in_full_mode` | Full 模式不改 slot config | 验证 `summary_override` 为 None |
|
||||
| `summary_failure_does_not_block` | LLM error 不阻断 `submit_turn` | MockProvider 返回错误 |
|
||||
| `summary_stream_path` | 流式路径 `finalize_turn` 正确触发 | `submit_turn_stream` 端到端 |
|
||||
| `summary_format_messages` | 格式化输出结构正确 | 单元测试验证格式 |
|
||||
| `summary_skipped_for_empty_messages` | 空消息不调用 LLM | `generate_summary` 直接返回 `""` |
|
||||
| `summary_not_generated_if_max_context_unreachable` | `max_context_tokens` 过大时不触发 | 验证条件不满足 |
|
||||
|
||||
**验证**:`cargo test --all-targets` 全绿
|
||||
|
||||
---
|
||||
|
||||
## 规模估算
|
||||
|
||||
| 组件 | 纯实现 | 测试 | 合计 |
|
||||
|------|--------|------|------|
|
||||
| `agent/summary.rs`(SummaryConfig + format_messages + 默认 prompt + 截断守卫) | 60 | 10 | 70 |
|
||||
| `agent/runtime.rs`(1 个字段) | 3 | — | 3 |
|
||||
| `agent/builder.rs`(1 个方法) | 8 | 3 | 11 |
|
||||
| `agent/session.rs`(检查点 + generate_summary + get_conversation_summary) | 40 | 100 | 140 |
|
||||
| `agent.rs`(module 声明 + re-export) | 3 | — | 3 |
|
||||
| **合计** | **109** | **113** | **~222** |
|
||||
|
||||
---
|
||||
|
||||
## 风险评估
|
||||
|
||||
### 已知风险
|
||||
|
||||
| 风险 | 概率 | 影响 | 缓解措施 |
|
||||
|------|------|------|---------|
|
||||
| **同步阻塞**:摘要 LLM 调用延长 submit_turn 延迟 | 高 | 长对话用户多等 1-3s | 对于已达 75% 水位的长对话,用户感知可接受;所有错误静默处理 |
|
||||
| **默认模型不兼容**:非 OpenAI 用户未设置 `summary_model` 但默认 `None` 沿用主模型 | 低 | 无影响 | `summary_model` 默认 `None`,沿用主 provider 默认模型,零兼容问题 |
|
||||
| **无限循环**:摘要不减少 cost_so_far,每轮都超阈值 | 中 | 频繁 LLM 调用浪费 token | `debounce_turns=3` 强制隔断;`last_summary_turn` 记录确保了间隔。注意:摘要 token 不计入 `cost_so_far`(独立 LlmCycle),阈值不会因摘要本身加速膨胀 |
|
||||
| **Focusd 模式摘要位置**:注入为 `system` 消息排在列表末尾 | 低 | LLM 近因效应,摘要可能过度受关注 | 这是 v0.2 消费端的设计选择,Phase 16 不改变 |
|
||||
| **SessionMemory key 冲突**:用户手动写入 `"conversation_summary"` 会被覆盖 | 低 | 数据被摘要覆盖 | 文档建议用户自定义 key;或未来使用 namespaced key |
|
||||
| **可观测性盲区**:`tracing::warn!` 依赖用户配置了 tracing subscriber | 中 | 失败静默不可见 | 提升到 `tracing::error!` 级别,或加 `eprintln!` fallback |
|
||||
|
||||
### 不做的事
|
||||
|
||||
- ❌ 不扩展 `HookContext`
|
||||
- ❌ 不引入 `tokio::spawn` 后台摘要
|
||||
- ❌ 不做增量摘要(`SummaryStrategy::Incremental` 留待 v0.4)
|
||||
- ❌ 不改 `filter_focused()` 的摘要注入位置
|
||||
- ❌ 不追踪摘要 token 消耗(`summary_cost_so_far`)
|
||||
- ❌ 不添加运行时 prompt 校验(不检查 `{messages}` 是否存在)
|
||||
- ❌ 不添加 `MergeStrategy::Summarize` 变体(`context.rs:108` 预占注释将在实施时同步移除或更新)
|
||||
|
||||
---
|
||||
|
||||
## 验收标准
|
||||
|
||||
### 编译与测试
|
||||
|
||||
| 检查项 | 指标 |
|
||||
|--------|------|
|
||||
| `cargo build --all-targets` | ✅ 通过 |
|
||||
| `cargo test --all-targets` | ✅ 全量通过(预计 335 → ~345,新增 ~10 测试) |
|
||||
| `cargo clippy --all-targets -- -D warnings` | ✅ 0 警告 |
|
||||
| 测试覆盖范围 | F1-F8、N1-N4 |
|
||||
|
||||
### 功能验收场景
|
||||
|
||||
**场景 1:启用摘要后的长对话**
|
||||
|
||||
```rust
|
||||
let session = AgentSession::new(agent, "session-1", Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(registry)
|
||||
.hook_executor(executor)
|
||||
.summary_config(SummaryConfig {
|
||||
max_context_tokens: 100,
|
||||
trigger_token_ratio: 0.5,
|
||||
debounce_turns: 2,
|
||||
..Default::default()
|
||||
})
|
||||
.build()?
|
||||
));
|
||||
session.submit_turn("msg 1").await?;
|
||||
// ... submit_turn 多次直到 token 超 50 ...
|
||||
// 第 N 轮:摘要自动生成
|
||||
let summary = session.get_session_data("conversation_summary").await?;
|
||||
assert!(summary.is_some());
|
||||
// Focused 模式下 load_messages 包含摘要
|
||||
```
|
||||
|
||||
**场景 2:不启用时零影响**
|
||||
|
||||
```rust
|
||||
let session = AgentSession::new(agent, "session-2", bundle); // 无 summary_config
|
||||
for i in 0..50 {
|
||||
session.submit_turn(&format!("msg {}", i)).await?;
|
||||
}
|
||||
// 没有摘要产生,没有额外的 LLM 调用
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 参考来源
|
||||
|
||||
- Phase 10 方案文档:`docs/17-phase10-contextslot.md`(§5 FocusedConfig 消费端设计)
|
||||
- Phase 14 方案文档:`docs/20-phase14-document-and-embedding.md`(Provider 复用模式)
|
||||
- 当前代码:`src/agent/session.rs`(submit_turn 流程,OnTurnEnd 位置)
|
||||
- 当前代码:`src/llm/cycle.rs`(submit_messages 签名)
|
||||
- 当前代码:`src/agent/context.rs`(FocusedConfig.summary_override + filter_focused 消费逻辑)
|
||||
- 当前代码:`src/agent/runtime.rs`(AgentConfig 结构)
|
||||
- 当前代码:`src/agent/builder.rs`(Builder 链式模式)
|
||||
- 当前代码:`src/agent/session_memory.rs`(set/get API)
|
||||
@@ -0,0 +1,775 @@
|
||||
# Phase 17 — Agent 执行引擎
|
||||
|
||||
- **文档编号**:23
|
||||
- **标题**:Phase 17 — Agent 执行引擎(Engine)
|
||||
- **日期**:2026-07-15
|
||||
- **状态**:**审查修复完成,待第二轮复审**
|
||||
- **涉及模块**:`engine/`(新建,含 `session_manager` / `checkpointer` / `snapshot` / `error`)、`agent/session`、`agent/context`、`llm/types/usage`
|
||||
- **关联文档**:`docs/17-phase10-contextslot.md`、`docs/22-phase16-summary-auto-generation.md`、`docs/roadmap.md`
|
||||
- **审查记录**:第 1 轮 PM Director + SA Director 审查 → 6 🔴 阻塞问题,全部修复。详见 §变更记录。
|
||||
|
||||
---
|
||||
|
||||
## 背景与目标
|
||||
|
||||
### 问题
|
||||
|
||||
agcore v0.3.0 开发中,已完成 Phase 13-16(Phase 0-12 全部完成)。当前测试 353 个,全部通过,clippy 0 警告。
|
||||
|
||||
当前 `AgentSession` 存在以下空白:
|
||||
|
||||
1. **Session 在变量中**:`AgentSession` 实例仅在内存中存在,无法通过 session ID 从存储恢复
|
||||
2. **无父子关系**:session 之间相互独立,无法表达"子会话继承父会话"的树形关系
|
||||
3. **无检查点**:无法在任意时刻给 session 拍快照,出错后无法回滚到历史状态
|
||||
4. **不可序列化**:`AgentSession` 持有 `Arc<dyn Agent>` 和 `Arc<RuntimeBundle>`,无法直接序列化持久化
|
||||
|
||||
### 目标
|
||||
|
||||
建立 `engine/` 模块,补齐 session 生命周期的管理能力。具体包括:
|
||||
|
||||
1. **Session 工厂 + 按 ID 恢复**:`SessionManager::create()` / `get()`,session 创建后可通过 ID 从存储重建
|
||||
2. **父子 session 树形关系**:`create_child()` / `children()` / `parent()`,支持树形会话拓扑
|
||||
3. **生命周期管理**:`destroy()` 清理 session 及其存储记录
|
||||
4. **Time-travel Checkpointer**:`checkpoint()` / `rollback()` / `list_checkpoints()`,支持任意时刻状态快照与回滚
|
||||
5. **序列化支持**:通过 `SessionSnapshot` 独立 struct 间接实现 `AgentSession` 的快照持久化
|
||||
|
||||
### 成功标准
|
||||
|
||||
1. Session 创建后可通 ID 从存储恢复(`get()` 返回完整状态的 `AgentSession`)
|
||||
2. 父子 session 关系可查询(`children()` / `parent()`),数据正确隔离
|
||||
3. Checkpoint 拍快照后可完全恢复到该时刻状态(turn_index、cost_so_far、slots 一致)
|
||||
4. 零新外部依赖,全量测试 353 → ~385-390
|
||||
5. `cargo test --all-targets` 全绿,`cargo clippy` 0 警告
|
||||
|
||||
---
|
||||
|
||||
## 当前状态分析
|
||||
|
||||
### 模块现状
|
||||
|
||||
| 模块 | 文件 | 状态 | 与 Phase 17 的关系 |
|
||||
|------|------|------|-------------------|
|
||||
| `AgentSession` | `agent/session.rs` | ✅ 已实现 | 需扩展 `to_snapshot()` / `from_snapshot()` |
|
||||
| `ContextSlot` | `agent/context.rs` | ✅ 已实现(持久化、fork/merge/save/load) | 需加 `Serialize` / `Deserialize` derive |
|
||||
| `CostTracker` | `llm/types/usage.rs` | ✅ 已实现 | 需加 `Clone` + `Serialize` / `Deserialize` derive |
|
||||
| `MergeStrategy` | `agent/context.rs` | ✅ 已实现 | 需加 `Serialize` / `Deserialize` derive |
|
||||
| `MemoryStore` trait | `memory/store.rs` | ✅ 已实现 | Checkpointer 的存储后端 |
|
||||
| `RuntimeBundle` | `agent/runtime.rs` | ✅ 已实现(依赖注入容器) | `from_snapshot()` 需注入 `agent` 和 `bundle` |
|
||||
| `InMemoryStore` | `memory/store.rs` | ✅ 已实现 | 测试用存储后端 |
|
||||
| `SqliteStore` | `memory/sqlite_store.rs` | ✅ 已实现(Phase 7) | 生产环境存储后端 |
|
||||
| `Message` | `llm/types/message.rs` | ✅ 已有 `Serialize` / `Deserialize` | 可直接序列化 |
|
||||
|
||||
### AgentSession 关键字段
|
||||
|
||||
```rust
|
||||
pub struct AgentSession {
|
||||
pub session_id: String,
|
||||
pub agent: Arc<dyn Agent>, // ❌ 不可序列化
|
||||
bundle: Arc<RuntimeBundle>, // ❌ 不可序列化
|
||||
turn_index: u32, // ✅ 可序列化
|
||||
cost_so_far: CostTracker, // ⚠️ 需加 derive
|
||||
pub session_memory: SessionMemory, // ⚠️ 间接序列化
|
||||
slots: HashMap<String, ContextSlot>, // ⚠️ 需加 derive
|
||||
current_slot_id: String, // ✅ 可序列化
|
||||
last_summary_turn: Option<u32>, // ✅ 可序列化
|
||||
}
|
||||
```
|
||||
|
||||
核心制约:`Arc<dyn Agent>` 和 `Arc<RuntimeBundle>` 无法 `Serialize` / `Deserialize`,必须通过独立 snapshot struct + 外部注入重建。
|
||||
|
||||
### 关键假设(设计分析 — 需实施后验证)
|
||||
|
||||
以下假设在方案设计中做出,标注验证方式。实施 Step 1-3 后应逐项确认。
|
||||
|
||||
| # | 假设 | 验证方式 |
|
||||
|---|------|---------|
|
||||
| 1 | `submit_turn_stream` 内部 `tokio::spawn` 不持有 `&mut self` → 可通过 `Arc<Mutex<AgentSession>>` 安全共享 | 代码审查覆盖 `submit_with_tools_stream` → `run_tool_loop` 的 spawn 捕获列表;确认所有捕获变量为 owned 数据 |
|
||||
| 2 | `CostTracker` 加 `Clone` 不破坏现有代码 | 编译验证(`cargo build --all-targets`);检查 `CostTracker` 的所有消费方(`session.rs` 中只读引用) |
|
||||
| 3 | `ContextSlot` 加 `Serialize` / `Deserialize` 不影响现有 `save` / `load` 路径 | 现有 `save()` 直接序列化 `self.messages` / `self.meta` / `self.config`,不走 `ContextSlot` 整体 serde → 两组路径可共存 |
|
||||
| 4 | `Message` 已有 `Serialize` / `Deserialize` → 可直接嵌套序列化 | 代码确认(`message.rs` L21 已有 derive) |
|
||||
| 5 | `EngineError` 不需要 `derive Serialize` → 纯运行时错误类型 | Checkpoint 只存 `SessionSnapshot`,不存错误枚举 |
|
||||
| 6 | `MemoryStore` 操作是可靠的——失败时返回 `EngineError::Memory` 透传错误 | 当前不内置 store 重试逻辑;调用方负责 retry 或 failover |
|
||||
| 7 | session_id 使用 UUID v4 自动生成,冲突概率可忽略 | 实施确定 ID 生成方案(`uuid::Uuid::new_v4()` 或 时间戳+计数器无依赖方案) |
|
||||
| 8 | session_memory 当前只支持字符串值;未来支持复杂类型时 `SessionMemoryEntry` 的 `value` 字段需改用 `serde_json::Value` | 已预留在注释中 |
|
||||
|
||||
---
|
||||
|
||||
## 调研发现
|
||||
|
||||
### 可选方案对比
|
||||
|
||||
#### 方案 A(推荐):SessionSnapshot + 组合式架构
|
||||
|
||||
**做法**:用一个独立 `SessionSnapshot` struct 存储可序列化状态,避开 `Arc<dyn Agent>` 的序列化限制。`Checkpointer` 作为独立 struct,`SessionManager` 组合持有 `Checkpointer`。
|
||||
|
||||
**优点**:
|
||||
- 不污染 `AgentSession` 主类型,序列化逻辑与运行逻辑分离
|
||||
- `Checkpointer` 独立可测,不依赖 `SessionManager`
|
||||
- 组合关系清晰:`SessionManager` 持有 `Checkpointer`
|
||||
- 所有字段使用 `#[serde(default)]` 宽松反序列化,前向兼容
|
||||
|
||||
**缺点**:
|
||||
- 需要额外同步逻辑:`to_snapshot()` / `from_snapshot()` 双向转换
|
||||
|
||||
#### 方案 B(已否决):直接给 AgentSession derive Serialize
|
||||
|
||||
**做法**:给 `AgentSession` 加 `#[derive(Serialize)]`,用 `#[serde(skip)]` 跳过 `agent` 和 `bundle`。
|
||||
|
||||
**否决原因**:
|
||||
1. `#[serde(skip)]` 跳过了 2 个核心字段,序列化后的结果名不副实
|
||||
2. 技术债重:主类型获得"跳过一半字段"的诡异 serde 行为,未来维护者可能误以为 `AgentSession` 可整体序列化/反序列化
|
||||
3. 反序列化时 `agent` 和 `bundle` 缺失,仍需外部注入 → 不如直接使用独立的 snapshot struct
|
||||
|
||||
#### 方案 C(已否决):Checkpointer 作为 SessionManager 内部方法
|
||||
|
||||
**做法**:将 `checkpoint` / `rollback` 直接作为 `SessionManager` 的方法。
|
||||
|
||||
**否决原因**:
|
||||
1. 违反单一职责原则(SRP):`SessionManager` 承担 session 生命周期 + 检查点管理双重责任
|
||||
2. 破坏独立可测试性:检查点逻辑与 `SessionManager` 耦合
|
||||
3. `rollback` 返回后自动注册到 `SessionManager`,但调用方可能不需要注册
|
||||
4. 应返回 `AgentSession` 让调用方决定如何处理
|
||||
|
||||
### 技术决策清单
|
||||
|
||||
| 编号 | 决策项 | 选择 | 理由 |
|
||||
|------|--------|------|------|
|
||||
| D1 | 序列化方式 | `SessionSnapshot` 独立 struct | 不污染 `AgentSession`,序列化逻辑与运行逻辑分离 |
|
||||
| D2 | 并发模型 | `tokio::sync::Mutex` | 安全跨 `.await`,与 `AgentSession` 现有模式一致 |
|
||||
| D3 | 模块拆分 | `Checkpointer` 独立 + `SessionManager` 组合 | 独立可测,SRP 合规 |
|
||||
| D4 | 存储格式 | 全量 JSON | 简洁可靠,ponytail:>500 轮再优化为增量 |
|
||||
| D5 | Key 命名 | `session:{id}:meta` / `ckpt:{id}:{ckpt_id}` | 与 `slot_data:` 风格一致,prefix 查询友好 |
|
||||
| D6 | Checkpoint 触发 | `SessionManager` 封装方法中自动;同步写入 + `tracing::error!` 记录失败 | `AgentSession` 保持纯净;不提供强持久化保证(显式调 `checkpointer.checkpoint()` 确认) |
|
||||
| D7 | 序列化兼容 | `#[serde(default)]` 宽松 | 防前向破坏,新增字段自动兼容旧快照 |
|
||||
| D8 | 流式 checkpoint 时序 | 仅在 `finalize_turn` 时创建 checkpoint | `submit_turn_stream` 返回流时不做 checkpoint;客户端断开后不留下半成品 checkpoint 污染 |
|
||||
| D9 | `SessionManager` trait | 不需要 | YAGNI,无多后端需求 |
|
||||
| D10 | `CostTracker` / `ContextSlot` / `MergeStrategy` derive | 加 `Clone` + `Serialize` / `Deserialize` | 共约 7 行改动,支持快照序列化 |
|
||||
|
||||
### MVP 范围
|
||||
|
||||
| 做(Phase 17 首批) | 推迟 |
|
||||
|---------------------|------|
|
||||
| ① `SessionManager`: `create` / `get` / `create_child` / `children` / `parent` / `destroy` / `replace` / `recover` | ① `destroy_subtree` — 首次只做单节点 `destroy`。父被销毁后子 session 的 `parent()` 返回 `None`(允许孤儿)。调用方如需级联删除应自行遍历。 |
|
||||
| ② `Checkpointer`: `checkpoint` / `rollback` / `list_checkpoints` / `delete_all` | ② `tree()` — `children()` + `parent()` 组合查询在 v0.3 够用;Phase 18 SubAgent Dispatch 需要全量树快照时再补。 |
|
||||
| ③ `SessionSnapshot` + `to_snapshot()` / `from_snapshot()`(位于 `engine/snapshot.rs`)+ `restore_memory()` | ③ `Checkpointer::fork` — 推迟理由:`fork` 底层可拆解为 `rollback` + `create_child`,当前 Checkpointer + SessionManager 已提供原始能力。`fork` 作为高层 API 等价于约 30 行组合代码,风险可控延后到 Phase 18。若产品认为 fork 是 time-travel MVP 的必要项,可重新划入 Phase 17。 |
|
||||
| ④ `EngineError`(含 `MemoryError` 透传) | |
|
||||
| ⑤ 涉及的 derive 改动(`CostTracker` + `ContextSlot` + `MergeStrategy`) | |
|
||||
|
||||
**变更记录**(审查修复):
|
||||
- `create()` / `create_child()` 返回类型改为 `Result<String, EngineError>`
|
||||
- `get()` 改为仅内存查询,新增 `recover()` 显式恢复方法
|
||||
- 新增 `replace()` 方法支持 rollback 后无缝切换
|
||||
- MVP 推迟列补充 `tree()`(含推迟理由)、完善 `destroy_subtree`(定义孤儿语义)、
|
||||
补充 `fork` 推迟理由(含技术拆解和产品权衡)
|
||||
|
||||
---
|
||||
|
||||
## 推荐方案
|
||||
|
||||
### 架构概览
|
||||
|
||||
```
|
||||
┌──────────────────────────────────────────────┐
|
||||
│ Engine │
|
||||
│ ┌────────────────┐ ┌──────────────────┐ │
|
||||
│ │ SessionManager │──│ Checkpointer │ │
|
||||
│ │ │ │ │ │
|
||||
│ │ create() │ │ checkpoint() │ │
|
||||
│ │ get() │ │ rollback() │ │
|
||||
│ │ create_child() │ │ list_checkpoints│ │
|
||||
│ │ children() │ │ │ │
|
||||
│ │ parent() │ └──────────────────┘ │
|
||||
│ │ destroy() │ │
|
||||
│ └────────┬───────┘ │
|
||||
│ │ 组合 │
|
||||
│ │ 持有 │
|
||||
│ ▼ │
|
||||
│ ┌────────────────┐ │
|
||||
│ │ MemoryStore │ ── 存储后端 │
|
||||
│ └────────────────┘ │
|
||||
└──────────────────────────────────────────────┘
|
||||
|
||||
▼
|
||||
┌──────────────────┐
|
||||
│ SessionSnapshot │ ── 可序列化的状态快照
|
||||
│ (to/from │
|
||||
│ AgentSession) │
|
||||
└──────────────────┘
|
||||
```
|
||||
|
||||
### 模块划分
|
||||
|
||||
**新增文件**(5 个):
|
||||
|
||||
```
|
||||
src/engine/
|
||||
├── mod.rs # 约 30 行:模块根 + pub use 重导出
|
||||
├── session_manager.rs # 约 300 行:SessionManager 实现(含 replace/recover)
|
||||
├── checkpointer.rs # 约 220 行:Checkpointer 实现
|
||||
├── snapshot.rs # 约 50 行:SessionSnapshot + SessionMemoryEntry 定义
|
||||
└── error.rs # 约 70 行:EngineError 枚举
|
||||
```
|
||||
|
||||
**修改文件**(5 个):
|
||||
|
||||
| 文件 | 改动量 | 内容 |
|
||||
|------|--------|------|
|
||||
| `src/agent/session.rs` | +~80 行 | `to_snapshot()` / `from_snapshot()` / `restore_memory()` |
|
||||
| `src/agent/context.rs` | +4 行 | `ContextSlot` + `MergeStrategy` 加 `Serialize` / `Deserialize` |
|
||||
| `src/llm/types/usage.rs` | +3 行 | `CostTracker` 加 `Clone` + `Serialize` / `Deserialize` |
|
||||
| `src/lib.rs` | +2 行 | `pub mod engine` 声明 |
|
||||
| `examples/engine_demo.rs` | +~100 行(新增) | 端到端示例(含 rollback + replace 流程) |
|
||||
|
||||
### SessionSnapshot(位于 `engine/snapshot.rs`)
|
||||
|
||||
设计决策:`SessionSnapshot` 是 engine 层为持久化引入的序列化 DTO,定义在 `engine/snapshot.rs` 而非 `agent/session.rs`,保持依赖方向为 `engine → agent`。
|
||||
|
||||
```rust
|
||||
/// SessionMemory 条目的可序列化形式(保留元数据与时间戳)。
|
||||
#[derive(Serialize, Deserialize, Clone)]
|
||||
struct SessionMemoryEntry {
|
||||
pub value: String,
|
||||
#[serde(default)]
|
||||
pub metadata: serde_json::Value,
|
||||
#[serde(default)]
|
||||
pub created_at: Option<i64>, // Unix 时间戳秒;Option 兼容旧快照
|
||||
}
|
||||
|
||||
/// AgentSession 的可序列化快照。
|
||||
///
|
||||
/// 不持有 `Arc<dyn Agent>` 和 `Arc<RuntimeBundle>` —— 这两个由调用方在
|
||||
/// `from_snapshot()` 时注入。所有字段使用 `#[serde(default)]` 确保前向兼容。
|
||||
///
|
||||
/// **变更记录**(审查修复):
|
||||
/// - 位置从 `agent/session.rs` 移至 `engine/snapshot.rs`
|
||||
/// - `session_memory_data` 从 `HashMap<String, String>` 改为 `HashMap<String, SessionMemoryEntry>`
|
||||
/// 保留 metadata 和 created_at,避免恢复后时间戳丢失
|
||||
#[derive(Serialize, Deserialize, Clone)]
|
||||
pub(crate) struct SessionSnapshot {
|
||||
pub session_id: String,
|
||||
pub agent_name: String,
|
||||
pub turn_index: u32,
|
||||
#[serde(default)]
|
||||
pub cost_so_far: CostTracker,
|
||||
#[serde(default)]
|
||||
pub slots: HashMap<String, ContextSlot>,
|
||||
pub current_slot_id: String,
|
||||
pub last_summary_turn: Option<u32>,
|
||||
#[serde(default)]
|
||||
pub session_memory_data: HashMap<String, SessionMemoryEntry>,
|
||||
}
|
||||
```
|
||||
|
||||
### AgentSession 扩展方法
|
||||
|
||||
```rust
|
||||
impl AgentSession {
|
||||
/// 将当前状态拍平为 SessionSnapshot。
|
||||
///
|
||||
/// **需要 async**:因为 session_memory 的数据存储在 `MemoryStore` 中,读取需要异步 I/O。
|
||||
/// 可通过 `SessionMemory::list_entries()` 获取完整条目(含 metadata/created_at):
|
||||
///
|
||||
/// ```ignore
|
||||
/// let entries = self.session_memory.list_entries().await?;
|
||||
/// for (key, value, metadata, created_at) in entries {
|
||||
/// map.insert(key, SessionMemoryEntry { value, metadata, created_at: Some(created_at) });
|
||||
/// }
|
||||
/// ```
|
||||
/// `from_snapshot` 保持同步(构造器不应做 I/O),`to_snapshot` 做 async(快照输出可 I/O)—
|
||||
/// 两个方向不矛盾,设计上各自成立。
|
||||
pub async fn to_snapshot(&self) -> SessionSnapshot {
|
||||
// 拍平 session_memory → HashMap<String, SessionMemoryEntry>(通过 list_entries)
|
||||
// 复制 slots / cost_so_far / turn_index 等可序列化字段
|
||||
}
|
||||
|
||||
/// 从 SessionSnapshot + agent + bundle 重建 AgentSession。
|
||||
///
|
||||
/// **纯同步重建**:只做内存数据结构恢复(slots/turn_index/cost_so_far 等),
|
||||
/// 不执行任何 I/O。session_memory 的持久层恢复由 `restore_memory()` 完成。
|
||||
///
|
||||
/// 调用方负责:
|
||||
/// - 提供与 `agent_name` 对应的 `Arc<dyn Agent>`
|
||||
/// - 提供合法的 `Arc<RuntimeBundle>`
|
||||
///
|
||||
/// 返回 `Result` 以传播序列化反序列化错误(如 JSON 格式不兼容)。
|
||||
pub fn from_snapshot(
|
||||
snapshot: SessionSnapshot,
|
||||
agent: Arc<dyn Agent>,
|
||||
bundle: Arc<RuntimeBundle>,
|
||||
) -> Result<Self, EngineError> {
|
||||
// session_memory_data 存入临时字段(不写 store)
|
||||
// 重建 slots HashMap
|
||||
// 恢复 turn_index / cost_so_far / last_summary_turn
|
||||
}
|
||||
|
||||
/// 将 snapshot 中的 session_memory_data 写回持久层。
|
||||
/// 从 `from_snapshot()` 中剥离的异步操作,调用方显式 await。
|
||||
/// 放置在 `restore_memory` 而非构造函数中,确保构造函数是纯同步的。
|
||||
///
|
||||
/// **错误处理**:逐条写入,某条失败时返回 Err 但不回滚已写入的条目。
|
||||
/// 调用方可选择重试或忽略(不影响 AgentSession 内存状态)。
|
||||
pub async fn restore_memory(&self) -> Result<(), EngineError>;
|
||||
}
|
||||
```
|
||||
|
||||
**标准使用流程**:
|
||||
```rust
|
||||
// rollback:四步走
|
||||
let snapshot = cp.rollback_load(session_id, ckpt_id).await?; // ① 从存储读
|
||||
let session = AgentSession::from_snapshot(snapshot, agent, bundle)?; // ② 同步重建
|
||||
session.restore_memory().await?; // ③ 恢复持久层
|
||||
sm.replace(session_id, session).await?; // ④ 注册到 Manager
|
||||
```
|
||||
|
||||
**checkpoint 流程**(自动或显式调用):
|
||||
```rust
|
||||
// checkpoint 内部:
|
||||
let snapshot = session.to_snapshot().await; // async:从 MemoryStore 读取 session_memory
|
||||
cp.save(snapshot).await?;
|
||||
```
|
||||
|
||||
### Checkpointer 公开 API
|
||||
|
||||
```rust
|
||||
/// 检查点元数据。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CkptMeta {
|
||||
pub ckpt_id: String,
|
||||
pub session_id: String,
|
||||
pub turn_index: u32,
|
||||
pub created_at: u64, // Unix 时间戳,秒
|
||||
}
|
||||
|
||||
/// Time-travel 检查点管理器。
|
||||
///
|
||||
/// **不依赖 SessionManager**,可独立使用。直接操作 MemoryStore。
|
||||
/// 存储 key 格式:`ckpt:{session_id}:{ckpt_id}` → SessionSnapshot JSON
|
||||
pub struct Checkpointer {
|
||||
store: Arc<dyn MemoryStore>,
|
||||
}
|
||||
|
||||
impl Checkpointer {
|
||||
/// 创建新检查点。返回 ckpt_id。
|
||||
pub async fn checkpoint(&self, session: &AgentSession) -> Result<String, EngineError>;
|
||||
|
||||
/// 回滚到指定检查点。返回恢复后的 AgentSession。
|
||||
///
|
||||
/// 调用方需提供 `agent` 和 `bundle`(与 SessionSnapshot 反序列化的要求一致)。
|
||||
/// rollback 不自动注册到任何 SessionManager——调用方决定如何处理返回的 session。
|
||||
pub async fn rollback(
|
||||
&self,
|
||||
session_id: &str,
|
||||
ckpt_id: &str,
|
||||
agent: Arc<dyn Agent>,
|
||||
bundle: Arc<RuntimeBundle>,
|
||||
) -> Result<AgentSession, EngineError>;
|
||||
|
||||
/// 列出某 session 的所有检查点(按创建时间降序)。
|
||||
pub async fn list_checkpoints(&self, session_id: &str)
|
||||
-> Result<Vec<CkptMeta>, EngineError>;
|
||||
|
||||
/// 删除某 session 的所有检查点(session 被 destroy 时调用)。
|
||||
pub async fn delete_all(&self, session_id: &str) -> Result<(), EngineError>;
|
||||
}
|
||||
```
|
||||
|
||||
**注意**:`Checkpointer::fork()` 推迟到 Phase 18(详见 MVP 范围表)。
|
||||
|
||||
**关于 Checkpointer 的独立可用性**:Snapshot 数据的读写(`checkpoint` / `list_checkpoints`)不依赖 SessionManager,可直接用 `Checkpointer` 操作 MemoryStore。但 `rollback()` 重建 AgentSession 需要调用方提供与 session_id 匹配的 `Arc<dyn Agent>` 和 `Arc<RuntimeBundle>`——调用方需自行管理 agent→session 的映射(或通过 `SessionMeta.agent_name` 查询注册表)。
|
||||
|
||||
### SessionManager 公开 API
|
||||
|
||||
```rust
|
||||
/// SessionManager 配置。
|
||||
pub struct SessionManagerConfig {
|
||||
/// 每次 submit_turn 后是否自动 checkpoint(默认 true)。
|
||||
pub auto_checkpoint: bool,
|
||||
/// 默认 RuntimeBundle,用于从存储重建 session 时的 bundle 注入。
|
||||
/// 如果为 None,`recover()` 需要调用方手动传入 bundle。
|
||||
pub default_bundle: Option<Arc<RuntimeBundle>>,
|
||||
}
|
||||
|
||||
impl Default for SessionManagerConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
auto_checkpoint: true,
|
||||
default_bundle: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Session 生命周期管理器。
|
||||
///
|
||||
/// 组合持有 Checkpointer,提供 session 的 CRUD、树形关系查询和自动检查点。
|
||||
/// 内部用 `HashMap<String, Arc<tokio::sync::Mutex<AgentSession>>>` 管理活跃 session。
|
||||
/// 存储 key 格式:`session:{session_id}:meta` → SessionMeta JSON
|
||||
///
|
||||
/// **锁契约**:
|
||||
/// - 所有写操作(create/destroy/replace)内部先完成 HashMap 操作,释放 RwLock 后再调用
|
||||
/// Checkpointer/MemoryStore 的异步 I/O。调用方不应假设某个操作持有跨 .await 点的锁。
|
||||
/// - `get()` 返回 `Arc<Mutex<AgentSession>>` 后立即释放 RwLock 读锁,调用方持有的是
|
||||
/// session 级别的 Mutex 锁而非管理器级别的锁。
|
||||
pub struct SessionManager {
|
||||
sessions: RwLock<HashMap<String, Arc<tokio::sync::Mutex<AgentSession>>>>,
|
||||
checkpointer: Checkpointer,
|
||||
store: Arc<dyn MemoryStore>,
|
||||
config: SessionManagerConfig,
|
||||
}
|
||||
|
||||
impl SessionManager {
|
||||
/// 创建新 session。session_id 由内部自动生成(UUID v4)。
|
||||
/// 持久化 SessionMeta 后注册到 sessions HashMap。
|
||||
pub async fn create(
|
||||
&self,
|
||||
agent: Arc<dyn Agent>,
|
||||
bundle: Arc<RuntimeBundle>,
|
||||
) -> Result<String, EngineError>;
|
||||
|
||||
/// 从父 session 创建子 session(继承父的 RuntimeBundle,Arc::clone 共享引用)。
|
||||
/// session_id 由内部自动生成(UUID v4)。
|
||||
/// 如果 `parent_id` 不存在,返回 `EngineError::SessionNotFound(parent_id)`。
|
||||
pub async fn create_child(
|
||||
&self,
|
||||
parent_id: &str,
|
||||
agent: Arc<dyn Agent>,
|
||||
) -> Result<String, EngineError>;
|
||||
|
||||
/// 按 ID 获取 session(仅查内存,不自动从存储恢复)。
|
||||
/// 冷启动时 `get()` 未命中返回 `EngineError::SessionNotFound`。
|
||||
/// 如需从存储恢复,使用 `recover()` 方法。
|
||||
pub async fn get(
|
||||
&self,
|
||||
session_id: &str,
|
||||
) -> Result<Arc<tokio::sync::Mutex<AgentSession>>, EngineError>;
|
||||
|
||||
/// 从存储恢复 session。需要调用方提供 agent 和 bundle(与 SessionSnapshot
|
||||
/// 反序列化的要求一致)。
|
||||
/// 恢复后自动注册到 sessions HashMap(与 create 的行为一致)。
|
||||
pub async fn recover(
|
||||
&self,
|
||||
session_id: &str,
|
||||
agent: Arc<dyn Agent>,
|
||||
bundle: Arc<RuntimeBundle>,
|
||||
) -> Result<Arc<tokio::sync::Mutex<AgentSession>>, EngineError>;
|
||||
|
||||
/// 替换 SessionManager 中指定 session_id 的 AgentSession 实例。
|
||||
/// 用于 Checkpointer::rollback() 后的无缝切换:
|
||||
/// ```ignore
|
||||
/// let rolled_back = cp.rollback(sid, ckpt_id, agent.clone(), bundle.clone()).await?;
|
||||
/// sm.replace(sid, rolled_back).await?;
|
||||
/// ```
|
||||
/// 内部执行:内存替换 + 写回 SessionMeta。
|
||||
pub async fn replace(
|
||||
&self,
|
||||
session_id: &str,
|
||||
session: AgentSession,
|
||||
) -> Result<(), EngineError>;
|
||||
|
||||
/// 查询某 parent 的所有直接子 session 的 ID 列表。
|
||||
pub async fn children(&self, parent_id: &str) -> Result<Vec<String>, EngineError>;
|
||||
|
||||
/// 查询某 child session 的 parent ID。
|
||||
/// 如果 parent 已被销毁,返回 `Ok(None)`(允许孤儿 session 存在)。
|
||||
pub async fn parent(&self, child_id: &str) -> Result<Option<String>, EngineError>;
|
||||
|
||||
/// 销毁 session:从内存移除 + 清理 SessionMeta + 清理检查点。
|
||||
///
|
||||
/// **父子关系处理**:允许孤儿 session 存在(子 session 的 parent_id 仍指向已删除的父,
|
||||
/// 但 `parent()` 返回 `None`)。不递归删除子 session——调用方如需级联删除应自行遍历。
|
||||
pub async fn destroy(&self, session_id: &str) -> Result<(), EngineError>;
|
||||
|
||||
/// 暴露 Checkpointer 引用(调用方可直接操作检查点)。
|
||||
pub fn checkpointer(&self) -> &Checkpointer;
|
||||
}
|
||||
```
|
||||
|
||||
**变更记录**(审查修复):
|
||||
- `create()` 返回类型从 `String` 改为 `Result<String, EngineError>`
|
||||
- `create()` / `create_child()` session_id 统一为内部自动生成(UUID v4)
|
||||
- `get()` 改为"仅查内存",新增 `recover()` 显式恢复方法
|
||||
- 新增 `replace()` 方法支持 rollback 后的无缝替换
|
||||
- `destroy()` 明确孤儿策略:允许孤儿存在,不递归删除
|
||||
- `create_child()` 不再接受 `child_id` 参数(统一自动生成)
|
||||
- 锁契约明确化为 struct doc comment
|
||||
- `SessionManagerConfig` 新增 `default_bundle` 字段为后续扩展预留
|
||||
|
||||
### SessionMeta
|
||||
|
||||
```rust
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct SessionMeta {
|
||||
pub session_id: String,
|
||||
pub agent_name: String,
|
||||
pub parent_id: Option<String>,
|
||||
pub created_at: u64, // Unix 时间戳,秒
|
||||
pub turn_count: u32,
|
||||
}
|
||||
```
|
||||
|
||||
### 存储 Key 命名
|
||||
|
||||
| Key 模式 | 内容 | 说明 |
|
||||
|----------|------|------|
|
||||
| `session:{session_id}:meta` | `SessionMeta` JSON | session 元数据,含 parent_id |
|
||||
| `ckpt:{session_id}:{ckpt_id}` | `SessionSnapshot` JSON | 全量检查点,含 slots |
|
||||
|
||||
风格与 `ContextSlot` 的 `slot_data:{session_id}:{slot_id}` 一致:`前缀:session_id:后缀`。
|
||||
|
||||
**关于两种持久化路径共存**:`ContextSlot::save()`(增量消息持久化)和 `Checkpointer::checkpoint()`(全量快照)是互补的"增量基线 vs 全量备份"关系:
|
||||
- `ContextSlot::save()` 每轮追加消息到 slot 存储(增量),是进程重启后消息不丢的基线
|
||||
- `Checkpointer::checkpoint()` 全量序列化 session 状态(含所有 slot 消息),是 time-travel 回滚的快照
|
||||
- rollback 时优先使用 checkpoint 的 snapshot 数据(一致性保证),不依赖 slot 持久化中的消息状态
|
||||
|
||||
### EngineError
|
||||
|
||||
```rust
|
||||
#[derive(Debug, Error)]
|
||||
#[non_exhaustive]
|
||||
pub enum EngineError {
|
||||
/// 指定 session_id 不存在。
|
||||
/// 适用场景:get() 内存未命中、create_child() parent 不存在、destroy() 操作不存在的 session。
|
||||
#[error("Session not found: {0}")]
|
||||
SessionNotFound(String),
|
||||
|
||||
/// 创建 session 时 ID 已存在(自动生成 ID 时通常不会触发)。
|
||||
#[error("Session already exists: {0}")]
|
||||
SessionAlreadyExists(String),
|
||||
|
||||
/// 指定 ckpt_id 不存在。
|
||||
#[error("Checkpoint not found: {0}")]
|
||||
CheckpointNotFound(String),
|
||||
|
||||
/// 存储错误(透传 MemoryError)。
|
||||
/// Checkpointer 和 SessionManager 的所有 MemoryStore 操作通过此变体传播错误。
|
||||
/// 与项目既有模式一致(对比 AgentError:直接 #[from] LlmError/ToolError/MemoryError)。
|
||||
#[from]
|
||||
#[error("存储错误: {0}")]
|
||||
Memory(#[from] MemoryError),
|
||||
|
||||
/// 序列化/反序列化失败(serde_json/snapshot 格式错误)。
|
||||
#[error("序列化错误: {0}")]
|
||||
Serialization(String),
|
||||
|
||||
/// Agent 错误(透传 AgentError)。
|
||||
#[from]
|
||||
#[error("Agent 错误: {0}")]
|
||||
Agent(#[from] AgentError),
|
||||
}
|
||||
```
|
||||
|
||||
### 并发模型
|
||||
|
||||
`SessionManager` 内部使用 `tokio::sync::RwLock` 保护 `sessions: HashMap`:
|
||||
|
||||
```rust
|
||||
pub struct SessionManager {
|
||||
sessions: RwLock<HashMap<String, Arc<tokio::sync::Mutex<AgentSession>>>>,
|
||||
// ... 其他字段
|
||||
}
|
||||
```
|
||||
|
||||
- `RwLock` 适合读多写少的场景(`get()` 高频 > `create()` / `destroy()`)
|
||||
- `get()` 返回 `Arc<Mutex<AgentSession>>` 后立即释放 RwLock 读锁,调用方持有的是 session 级别的 Mutex 锁而非管理器级别的锁。**不持有 RwLock 跨越 .await**
|
||||
- 所有写操作(`create`/`destroy`/`replace`)先完成 HashMap 操作(持有写锁),释放 RwLock 后再调用 Checkpointer/MemoryStore 的异步 I/O
|
||||
- 返回的 `AgentSession` 用 `Arc<tokio::sync::Mutex<AgentSession>>` 包裹,支持跨 `.await` 的安全可变访问
|
||||
- `Checkpointer` 无锁(纯函数式操作 MemoryStore,依赖其内部实现)
|
||||
|
||||
---
|
||||
|
||||
## 实施建议
|
||||
|
||||
### 阶段划分(共 7 步)
|
||||
|
||||
```
|
||||
Step 1: 前置 derive 改动 → step-1-branch
|
||||
Step 2: EngineError + 模块骨架 → step-2-branch
|
||||
Step 3: SessionSnapshot + 扩展 → step-3-branch
|
||||
Step 4: Checkpointer → step-4-branch
|
||||
Step 5: SessionManager → step-5-branch
|
||||
Step 6: 自动 checkpoint 集成 → step-6-branch
|
||||
Step 7: 示例 + 测试补强 → step-7-branch
|
||||
```
|
||||
|
||||
#### Step 1:前置 derive 改动
|
||||
|
||||
- **文件**:`src/llm/types/usage.rs`、`src/agent/context.rs`(×2)
|
||||
- **内容**:
|
||||
- `CostTracker`:`#[derive(Debug, Default)]` → `#[derive(Debug, Default, Clone, Serialize, Deserialize)]`
|
||||
- `ContextSlot`:`#[derive(Debug, Clone)]` → `#[derive(Debug, Clone, Serialize, Deserialize)]`
|
||||
- `MergeStrategy`:`#[derive(Debug, Clone)]` → `#[derive(Debug, Clone, Serialize, Deserialize)]`
|
||||
- **验证**:`cargo build --all-targets` 编译通过
|
||||
|
||||
#### Step 2:EngineError + 模块骨架
|
||||
|
||||
- **文件**:
|
||||
- `src/engine/error.rs`(新增):`EngineError` 枚举定义
|
||||
- `src/engine/mod.rs`(新增):模块根声明 + `pub use` 重导出 `EngineError` / `SessionManager` / `Checkpointer` / `CkptMeta`
|
||||
- `src/lib.rs`(修改):加 `pub mod engine;`
|
||||
- **验证**:`cargo build --all-targets && cargo clippy --all-targets -- -D warnings`
|
||||
|
||||
#### Step 3:SessionSnapshot + AgentSession 扩展
|
||||
|
||||
- **文件**:`src/engine/snapshot.rs`(新增,来自 SA 审查建议)、`src/agent/session.rs`
|
||||
- **内容**:
|
||||
- `src/engine/snapshot.rs`:`SessionMemoryEntry` 结构体(含 `value`/`metadata`/`created_at`)、`SessionSnapshot` 结构体定义(`pub(crate)`)
|
||||
- `src/agent/session.rs`:`pub async fn to_snapshot(&self) -> SessionSnapshot`(**异步**,通过 `SessionMemory::list_entries()` 读取完整 session_memory 条目,复制 slots/cost_so_far/各标量字段)
|
||||
- `pub fn from_snapshot(snapshot, agent, bundle) -> Result<Self, EngineError>`(**纯同步**,不写 store;session_memory_data 暂存于内存,不写入持久层)
|
||||
- `pub async fn restore_memory(&self) -> Result<(), EngineError>`(异步,将 from_snapshot 暂存的 session_memory_data 写回持久层;逐条写入,失败时记录 error 但不回滚已写入条目)
|
||||
- `SessionMemory` 新增 `list_entries()` 方法返回 `Vec<(String, String, serde_json::Value, i64)>`(含 value/metadata/created_at),供 `to_snapshot` 消费
|
||||
- **验证**:单元测试 roundtrip(`to_snapshot().await` → `from_snapshot()` → 关键字段一致);`restore_memory` 幂等性测试
|
||||
|
||||
#### Step 4:Checkpointer
|
||||
|
||||
- **文件**:`src/engine/checkpointer.rs`(新增)
|
||||
- **内容**:
|
||||
- `Checkpointer` 结构体(持有 `Arc<dyn MemoryStore>`)
|
||||
- `CkptMeta` 结构体
|
||||
- `checkpoint()`:生成 ckpt_id(时间戳+计数器方案优先,ponytail;`uuid` 备选,需加依赖),`session.to_snapshot()` → JSON → 存 `ckpt:{session_id}:{ckpt_id}`
|
||||
- `rollback_load()`(两阶段 rollback 的第一阶段):读取 JSON → 反序列化为 `SessionSnapshot` → 返回 `SessionSnapshot`
|
||||
- 调用方拿到 `SessionSnapshot` 后,自行调用 `AgentSession::from_snapshot()`(纯同步)+ `restore_memory()`(异步)+ `SessionManager::replace()`(注册)
|
||||
- `list_checkpoints()`:prefix 查询 `ckpt:{session_id}:` → 反序列化 `CkptMeta`(从 snapshot JSON 中提取 `turn_index` / `created_at`)→ 按时间降序
|
||||
- `delete_all()`:prefix 查询 + 逐个删除
|
||||
- **验证**:3-5 个单元测试(checkpoint roundtrip / rollback_load 反序列化正确 / list 排序 / delete_all 幂等性)
|
||||
|
||||
#### Step 5:SessionManager
|
||||
|
||||
- **文件**:`src/engine/session_manager.rs`(新增)
|
||||
- **内容**:
|
||||
- `SessionManagerConfig` 结构体(含 `auto_checkpoint: bool` + `default_bundle: Option<Arc<RuntimeBundle>>`)
|
||||
- `SessionMeta` 结构体(`pub(crate)`)
|
||||
- `SessionManager` 结构体(`RwLock<HashMap<...>>` + `Checkpointer` + `store` + `config`)
|
||||
- `create()`:内部自动生成 session_id(UUID v4),`AgentSession::new()` → 存 `SessionMeta` → 注册到 `sessions` HashMap → `Ok(session_id)`
|
||||
- `create_child()`:验证 parent 存在 → 自动生成 child session_id → 设置 `parent_id` → `create()` 流程
|
||||
- `get()`:**仅查内存**,未命中返回 `SessionNotFound`(不自动从存储恢复)
|
||||
- `recover(session_id, agent, bundle)`:从存储读取 `SessionMeta` + 调 `Checkpointer` 最近 checkpoint → 重建 `AgentSession` → 注册到 HashMap
|
||||
- `replace(session_id, session)`:内存替换(覆盖 Mutex 中的 AgentSession)+ 写回 SessionMeta
|
||||
- `children(parent_id)`:prefix 查询 `session:{parent_id}:` → 过滤 `parent_id` 匹配 → 返回 child_id 列表
|
||||
- `parent(child_id)`:读 `SessionMeta.parent_id`,父已被销毁时返回 `Ok(None)`
|
||||
- `destroy(session_id)`:移除内存记录 → 删除 `SessionMeta` → 调 `Checkpointer::delete_all()`。**允许孤儿 session 存在**(不递归删除子 session)
|
||||
- **验证**:8-10 个单元测试(CRUD / recover 恢复 / replace 替换 / 树形关系 / session 隔离 / destroy 后 get 失败 / 孤儿 parent 返回 None)
|
||||
|
||||
#### Step 6:自动 checkpoint 集成
|
||||
|
||||
- **文件**:`src/engine/session_manager.rs`(扩展)
|
||||
- **内容**:
|
||||
- 在 `SessionManager` 上添加封装方法 `submit_turn(session_id, user_input)`,内部:
|
||||
1. `get(session_id)` 获取 session
|
||||
2. `session.lock().await.submit_turn(user_input).await`
|
||||
3. 如果 `config.auto_checkpoint == true`,同步调用 `checkpointer.checkpoint(&session).await`
|
||||
- checkpoint 失败时通过 `tracing::error!` 记录,不阻断 `submit_turn` 的 `Ok` 返回
|
||||
- 调用方如需强持久化保证,应显式调用 `checkpointer.checkpoint()` 并处理其 `Result`
|
||||
- 流式路径:仅在 `finalize_turn` 时创建 checkpoint(`submit_turn_stream` 返回流时不做 checkpoint)
|
||||
- 客户端断开连接导致 `finalize_turn` 未被调用时,保持上一个 checkpoint 的状态,不留下半成品 checkpoint 污染
|
||||
- `auto_checkpoint` 配置控制开关
|
||||
- **验证**:集成测试(`submit_turn` → `list_checkpoints` 中可查到新 checkpoint);关闭 `auto_checkpoint` 时不产生 checkpoint
|
||||
|
||||
#### Step 7:示例 + 测试补强 + Tracing 埋点
|
||||
|
||||
- **文件**:`examples/engine_demo.rs`(新增,~100 行)
|
||||
- **示例流程**:
|
||||
1. `SessionManager::create` → submit_turn
|
||||
2. `Checkpointer::checkpoint` → list_checkpoints
|
||||
3. `Checkpointer::rollback` + `AgentSession::restore_memory` + `SessionManager::replace`
|
||||
4. 验证回滚后 turn_index 和 cost 恢复到 checkpoint 时刻
|
||||
- **Tracing 埋点**(每个关键操作添加 `tracing` 日志,与项目既有风格一致):
|
||||
- `Checkpointer::checkpoint()` 成功时:`tracing::info!(ckpt_id, turn_index, snapshot_size, "checkpoint created")`
|
||||
- `Checkpointer::rollback()` 成功时:`tracing::info!(ckpt_id, session_id, turn_index, "rolled back")`
|
||||
- `Checkpointer::list_checkpoints` → `tracing::debug!(session_id, count)`
|
||||
- `SessionManager::create` → `tracing::info!(session_id, agent_name, "session created")`
|
||||
- `SessionManager::destroy` → `tracing::info!(session_id, "session destroyed")`
|
||||
- `SessionManager::get` / `recover` / `replace` → `tracing::debug!(session_id, ...)`
|
||||
- 序列化错误 / 存储错误 → `tracing::error!(session_id, error, ...)`
|
||||
- **补充测试**(12-15 个):
|
||||
- 空 slot checkpoint → rollback 后消息为空
|
||||
- Destroy 后再 checkpoint → 返回 `SessionNotFound`
|
||||
- 跨 session 检查点隔离(session A checkpoint 不影响 session B)
|
||||
- 序列化版本兼容(`#[serde(default)]` 兜底:缺少新字段的旧 snapshot 可正常反序列化)
|
||||
- 10 并发 session 创建/销毁(RwLock 写锁争用验证)
|
||||
- 父子 session 消息隔离(子 session 写数据不污染父 session)
|
||||
- `restore_memory` 幂等性(重复调用不产生重复数据)
|
||||
- `from_snapshot` 纯同步验证(检查构造过程中无 async 调用路径)
|
||||
- **验证**:`cargo test --all-targets` 全绿 + `cargo clippy` 0 警告
|
||||
|
||||
### 高层建议
|
||||
|
||||
1. **Step 1 应先行独立提交**:derive 改动可能触发整个 crate 的重新编译,与其他步骤分开可减少冲突
|
||||
2. **`get()` 只查内存,`recover()` 用于存储恢复**:`get()` 不自动从存储重建(因无 `agent`/`bundle` 通道)。冷启动后先 `create()` 再 `get()`,或显式调用 `recover(session_id, agent, bundle)`
|
||||
3. **ckpt_id 生成**:使用 `uuid::Uuid::new_v4()`(需在 `Cargo.toml` `[dependencies]` 中添加 `uuid = { version = "1", features = ["v4"] }`),或走无新增依赖方案:`format!("{}_{}", session_id, timestamp_nanos)` 结合单调计数器。建议优先走无新增依赖方案(ponytail)
|
||||
4. **SessionManager 的 RwLock 粒度**:避免持写锁时调 `checkpointer`(涉及 I/O),锁范围应仅限于 HashMap 操作;`get()` 返回 `Arc` 后立即释放读锁
|
||||
5. **自动 checkpoint 的持久化语义**:自动 checkpoint 采用`同步写入 + tracing::error! 记录失败` 模式(与 Phase 16 `maybe_summarize` 的静默模式一致)。**不提供强持久化保证**——调用方如需确保 checkpoint 成功,应显式调用 `checkpointer.checkpoint()` 并处理其 `Result`
|
||||
6. **ContextSlot 持久化与 Checkpointer 快照的关系**:两者是"增量基线 vs 全量备份"的互补关系。`ContextSlot::save()` 负责每轮追加消息到 slot 存储(增量),`Checkpointer::checkpoint()` 负责全量序列化 session 状态(快照)。rollback 时优先使用 checkpoint 数据(一致性),不依赖 slot 持久化的消息状态
|
||||
7. **`from_snapshot` 后调用 `restore_memory`**:`AgentSession::from_snapshot()` 是纯同步的,不写 store;写回 session_memory 需要显式 `await session.restore_memory()`。三步全流程:`from_snapshot → restore_memory → replace`
|
||||
|
||||
### @Chart 提示
|
||||
|
||||
```
|
||||
flowchart TD
|
||||
subgraph "engine/"
|
||||
SM[SessionManager]
|
||||
CP[Checkpointer]
|
||||
EE[EngineError]
|
||||
end
|
||||
|
||||
subgraph "现有模块"
|
||||
AS[AgentSession]
|
||||
CS[ContextSlot]
|
||||
CT[CostTracker]
|
||||
MS[MemoryStore]
|
||||
end
|
||||
|
||||
SM -->|组合持有| CP
|
||||
SM -->|RwLock 保护| HM[(sessions HashMap)]
|
||||
CP -->|持久化| MS
|
||||
AS -->|to_snapshot| SS[SessionSnapshot]
|
||||
SS -->|from_snapshot| AS
|
||||
|
||||
SM -->|get / create / destroy| AS
|
||||
CP -->|checkpoint / rollback| AS
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 变更记录(审查修复)
|
||||
|
||||
| 日期 | 变更 | 触发 |
|
||||
|------|------|------|
|
||||
| 2026-07-15 | **🔴 `to_snapshot(&self)` 从同步改为 `pub async fn`** | SA 第 2 轮审查:同步方法无法 async 读 MemoryStore;需通过 `SessionMemory::list_entries()` 获取完整条目 |
|
||||
| 2026-07-15 | **🔴 `docs/roadmap.md` Phase 17 交付物列表同步更新** | PM 第 2 轮审查:Roadmap 仍使用旧版范围(`tree()`/`fork()`/`destroy_subtree()` 未推迟,`create()` 签名未更新,缺 `recover()`/`replace()`) |
|
||||
| 2026-07-15 | **`SessionMemory::list_entries()` 新增方法** | SA 第 2 轮审查:`to_snapshot` 需要读取完整 entry 数据,现有 API 只返回 `Option<String>` |
|
||||
| 2026-07-15 | **`to_snapshot` 注释清理:移除错误的 Cell/RefCell 方案** | SA 第 2 轮审查:同步方法中无法通过 Cell/RefCell 绕开 async |
|
||||
|
||||
| 日期 | 变更 | 触发 |
|
||||
|------|------|------|
|
||||
| 2026-07-15 | **🔴 `SessionManager::get()` 改为仅查内存,新增 `recover()` 显式恢复方法** | SA 审查:get() "从存储恢复"不可实现(无 agent/bundle 通道) |
|
||||
| 2026-07-15 | **🔴 `from_snapshot()` 改为纯同步构造 + 分离 `restore_memory()` 异步方法;返回 `Result`** | SA 审查:异步 I/O + 返回 Self 导致脏数据 |
|
||||
| 2026-07-15 | **🔴 `session_memory_data` 从 `HashMap<String, String>` 改为 `HashMap<String, SessionMemoryEntry>`** | SA 审查:拍平丢失 metadata/created_at |
|
||||
| 2026-07-15 | **🔴 `EngineError` 新增 `Memory(#[from] MemoryError)` 透传变体** | SA 审查:缺少 MemoryError 透传 |
|
||||
| 2026-07-15 | **🔴 `tree()` 在 MVP 推迟列补充(含推迟理由)** | PM 审查:Roadmap L781 需求完全未提及 |
|
||||
| 2026-07-15 | **🔴 `fork()` 推迟理由补充(技术拆解 + 产品权衡)** | PM 审查:推迟理由不充分 |
|
||||
| 2026-07-15 | **🔴 新增 `SessionManager::replace()` API 支持 rollback 后无缝切换** | PM 审查:rollback 后 session 无法替换到 Manager |
|
||||
| 2026-07-15 | **`SessionSnapshot` 移至 `engine/snapshot.rs`** | SA 审查:DTO 应放在 engine 层,保持依赖方向 engine→agent |
|
||||
| 2026-07-15 | **uuid 依赖修正:改为"时间戳+计数器优先,uuid 备选"** | SA 审查:文档声称"已有依赖"但 Cargo.toml 不含 |
|
||||
| 2026-07-15 | **160KB 具体数字删除(替换为保守上限描述)** | SA 审查:无测量依据 |
|
||||
| 2026-07-15 | **"关键假设(已验证)"改为"设计分析" + 验证方式** | PM 审查:"已验证"字面与实际不符 |
|
||||
| 2026-07-15 | **流式 checkpoint 时序明确定义:仅在 `finalize_turn` 时创建** | PM 审查:时序未定义 |
|
||||
| 2026-07-15 | **自动 checkpoint 语义:`tracing::error!` 模式,非强持久化** | SA 审查:fire-and-forget 不可靠 |
|
||||
| 2026-07-15 | **`create()` 返回 `Result<String, EngineError>` + 统一自动生成 ID** | PM 审查:返回 String 不能表达错误 |
|
||||
| 2026-07-15 | **`destroy()` 明确孤儿策略:允许孤儿,不递归删除,parent() 返回 None** | PM+SA 审查:孤儿语义未定义 |
|
||||
| 2026-07-15 | **`create_child()` 不再接受 `child_id`(统一自动生成)** | PM 审查:ID 策略不一致 |
|
||||
| 2026-07-15 | **`RuntimeBundle` 继承语义补充(`Arc::clone` 共享引用)** | PM 审查:继承语义未定义 |
|
||||
| 2026-07-15 | **并发模型补充 RwLock 锁范围注释** | SA 审查:跨 await 风险缺文档 |
|
||||
| 2026-07-15 | **Checkpointer 独立可用性约束标注** | SA 审查:rollback 重建需要 agent+bundle |
|
||||
| 2026-07-15 | **ContextSlot 与 Checkpointer 两种持久化路径关系补充说明** | SA 审查:共存缺说明 |
|
||||
| 2026-07-15 | **Step 7 扩充:示例流程 + 12-15 个边界测试 + Tracing 埋点规划** | SA 审查:缺 tracing 规划 |
|
||||
| 2026-07-15 | **`SessionManagerConfig` 新增 `default_bundle` 字段** | PM 审查:未来扩展预留 |
|
||||
| 2026-07-15 | **项目文件新增/修改数量同步更新(5 新增 + 5 修改,~725 行)** | 全部审查修复导致文件范围变化 |
|
||||
|
||||
## 参考来源
|
||||
|
||||
- Phase 10 方案文档:`docs/17-phase10-contextslot.md`(ContextSlot 持久化设计,Phase 17 的前置依赖)
|
||||
- Phase 16 方案文档:`docs/22-phase16-summary-auto-generation.md`(上一 Phase 的实施风格参考)
|
||||
- 当前代码:`src/agent/session.rs`(AgentSession 当前实现,`to_snapshot` / `from_snapshot` 扩展点)
|
||||
- 当前代码:`src/agent/context.rs`(ContextSlot 当前实现,derive 改动点)
|
||||
- 当前代码:`src/llm/types/usage.rs`(CostTracker 当前实现,derive 改动点)
|
||||
- 当前代码:`src/lib.rs`(模块注册点)
|
||||
- 当前代码:`src/agent.rs`(模块组织风格参考)
|
||||
@@ -0,0 +1,700 @@
|
||||
# Phase 18:Agent 角色热切换与子代理调度
|
||||
|
||||
## 背景与目标
|
||||
|
||||
### 问题空间
|
||||
|
||||
agcore 已完整交付 Phase 0-17,具备 SessionManager 会话生命周期管理、会话树(父子层次)、Checkpointer 检查点、流式输出、ContextSlot 上下文分区、MemoryStore 持久化、摘要自动生成等能力。当前 Session 在 `create()` 时绑定一个 `Arc<dyn Agent>`,此后无法变更角色;会话间调度仅通过 `create_child()` + `submit_turn()` 手动编排,缺乏内建的子代理派发机制。
|
||||
|
||||
Phase 18 要解决两个正交但关联的问题:
|
||||
|
||||
1. **Agent 角色热切换**:运行时替换 session 绑定的 Agent,保留上下文(slot 历史、turn_index、session_memory、cost_so_far)
|
||||
2. **子代理调度**:在 SessionManager 上提供声明式的 `dispatch` / `dispatch_stream` / `dispatch_all` API,支持父子 session 间的 Memory 继承、bridge_keys 注入、并发控制、结构化回传
|
||||
|
||||
### 目标
|
||||
|
||||
- 提供 `SessionManager::switch_agent(session_id, new_agent)`,替换 `Arc<dyn Agent>`,全量保留 slot / turn_index / session_memory
|
||||
- 提供 `DispatchConfig` / `SubTaskResult` / `SubTaskStreamEvent` 类型以及 `dispatch` / `dispatch_stream` / `dispatch_all` 三个核心方法
|
||||
- 实现父转子三层级交互:父->子(Memory 快照继承 + bridge_keys)、子->父(SubTaskResult 结构化回传 + result_summary)、子<->子(shared namespace)
|
||||
- 产出 4 个端到端示例:`agent_switch_demo` / `sub_agent_dispatch_demo` / `bridge_keys_demo` / `dispatch_stream_demo`
|
||||
|
||||
### 依赖与优先级
|
||||
|
||||
- **依赖**:Phase 17(SessionManager + 会话树 + Checkpointer)[高]
|
||||
- **优先级**:P0
|
||||
- **预估规模**:约 720 行核心 + 210 行测试 + 450 行示例
|
||||
- **审查修复**:第 1 轮审查修复(Finalize 方案重写 + 4 个 🔴 阻塞 + 8 个 🟡 改进)
|
||||
|
||||
---
|
||||
|
||||
## 当前状态分析
|
||||
|
||||
### 现有架构中的关键接入点
|
||||
|
||||
| 接入点 | 位置 | 可用性 | 分析 |
|
||||
|--------|------|--------|------|
|
||||
| `AgentSession.agent` | `src/agent/session.rs:53` | `pub` 字段 | 可直接替换,无需新增 setter [高] |
|
||||
| `AgentSession.session_id` | `src/agent/session.rs:51` | `pub` 字段 | 子 session 创建后可读取 [高] |
|
||||
| `AgentSession.session_memory` | `src/agent/session.rs:58` | `pub` 字段 | dispatch 后父可读子 memory [高] |
|
||||
| `SessionManager::create_child(parent_id, agent)` | `src/engine/session_manager.rs:202-239` | `pub async` | dispatch 可直接复用,bundle 继承避免重复构造 [高] |
|
||||
| `SessionMemory::list_entries()` | `src/agent/session_memory.rs:81-110` | `pub async` | 可获取全量条目用于父子继承 / bridge_keys 过滤 [高] |
|
||||
| `SessionMemory::set_with_meta()` | `src/agent/session_memory.rs:61-79` | `pub async` | 子 session 写入继承数据时保留原始 metadata [高] |
|
||||
| `CostTracker` | `src/llm/cycle.rs` | derive `Clone` | SubTaskResult 可直接 clone usage [高] |
|
||||
| `SessionMeta` 持久化 | `src/engine/session_manager.rs:32-67` | `pub(crate)` | 以 `session:{id}:meta` key 存到 MemoryStore;switch 后需更新 agent_name [高] |
|
||||
| `SessionManager::save_session_meta()` | `src/engine/session_manager.rs:131-142` | `async fn` (非 pub) | switch_agent 需要类似的 meta 更新能力;考虑提取为 `pub(crate)` [高] |
|
||||
| `SessionManager::destroy()` | `src/engine/session_manager.rs:495-512` | `pub async` | dispatch 失败时清理子 session 可直接复用 [高] |
|
||||
| `SessionManager.sessions` | `src/engine/session_manager.rs:92` | `pub(crate)` RwLock<HashMap> | switch_agent 和 dispatch 的 get / replace 操作均依赖此字段 [高] |
|
||||
| `EngineError` 枚举 | `src/engine/error.rs:16-46` | `#[non_exhaustive]` | 已有 6 个变体,需追加 DispatchFailed / SwitchFailed / SubAgentStreamError [中] |
|
||||
| `futures-core` / `futures-util` | `Cargo.toml:17-18` | 已引入 | dispatch_stream 返回 `Pin<Box<dyn Stream>>` 所需依赖已就绪,无需新增 [高] |
|
||||
|
||||
### Agent trait 与 SessionManager 之间的关系
|
||||
|
||||
```
|
||||
AgentSession {
|
||||
agent: Arc<dyn Agent>, // 可替换
|
||||
session_memory: SessionMemory, // 可继承(clone backend)
|
||||
slots: HashMap<String, ContextSlot>,
|
||||
turn_index: u32,
|
||||
cost_so_far: CostTracker,
|
||||
// ... 其余内部字段
|
||||
}
|
||||
|
||||
SessionManager {
|
||||
sessions: RwLock<HashMap<String, Arc<Mutex<AgentSession>>>>,
|
||||
checkpointer: Checkpointer,
|
||||
store: Arc<dyn MemoryStore>,
|
||||
config: SessionManagerConfig,
|
||||
}
|
||||
```
|
||||
|
||||
`AgentSession.agent` 是 `pub` 字段,这意味着 `switch_agent` 只需 `get()` → `lock()` → 替换 `agent` → 写回 meta。Route 明确、无架构阻力 [高]。
|
||||
|
||||
### 锁契约(需格外注意)
|
||||
|
||||
`src/engine/session_manager.rs:4-12` 记录了锁契约:**不持有 RwLock 跨越 `.await`**。所有 `.await` 点必须在 RwLock guard drop 之后。这意味着:
|
||||
|
||||
- `switch_agent`:读锁 `get()` 返回 `Arc<Mutex<AgentSession>>` 后释放,然后 lock session 级别的 Mutex → 替换 agent → 释放 Mutex → save_session_meta(I/O)[高]
|
||||
- `dispatch`:读锁 `get()` 父 session → 释放 → `create_child`(内部写锁)→ lock 子 session → inherit memory → submit_turn → 释放 [高]
|
||||
- 不会引入新的死锁风险 [高]
|
||||
|
||||
### 现有测试覆盖
|
||||
|
||||
SessionManager 已有 906 行(含 12 个测试),覆盖 create/get/destroy/create_child/replace/recover/children/parent/并发创建等场景:`src/engine/session_manager.rs:527-906`。Phase 18 新增测试不修改这些已有测试。
|
||||
|
||||
---
|
||||
|
||||
## 调研发现
|
||||
|
||||
### 1. bridge_keys 注入位置
|
||||
|
||||
**问题**:bridge_keys 本质是父 session 想注入到子 agent prompt 中的上下文数据。它应该放在哪里?
|
||||
|
||||
**调研来源**:
|
||||
- `docs/note-opencode-subagent-dispatch.md` — 明确反对修改 Agent 的 `system_prompt()` [高]
|
||||
- `src/agent/agent.rs:21` — `system_prompt()` 返回 `&str`,无状态变更能力 [高]
|
||||
- `src/agent/session_memory.rs` — `SessionMemory::set()` 提供 key-value 写入,子 agent 可读 [高]
|
||||
|
||||
**结论**:bridge_keys 通过 SessionMemory 副本继承 + 过滤注入,不碰 `system_prompt()`。子 agent 通过 `get_session_data(key)` 读取桥接数据 [高]。
|
||||
|
||||
### 2. SessionMemory 继承策略
|
||||
|
||||
**问题**:子 agent 启动时,父 session 的 SessionMemory 如何传递?
|
||||
|
||||
**方案 A —— 引用共享**:父子共享同一 `SessionMemory` 实例(Arc clone 后端)。优点是零拷贝,缺点是父子隔离被破坏 [中]。
|
||||
|
||||
**方案 B —— 快照副本**:父调用 `list_entries()` 获取全量条目,子通过 `set_with_meta()` 写入自己的 namespace。优点是隔离性强,缺点是 O(n) 拷贝开销 [高]。
|
||||
|
||||
**来源**:
|
||||
- `src/agent/session.rs:503-524` — `to_snapshot()` 已实现类似的 list_entries → HashMap 拍平 [高]
|
||||
- `src/agent/session_memory.rs:61-79` — `set_with_meta()` 可保留原始 metadata [高]
|
||||
- `docs/note-opencode-subagent-dispatch.md` — SA 建议副本策略 [中]
|
||||
|
||||
**结论**:采用方案 B(快照副本),隔离性优先。`inherit_session_memory` 内部使用 `list_entries()` → 按 `bridge_keys` 过滤 → `set_with_meta()` 写入子 namespace [高]。
|
||||
|
||||
### 3. dispatch_all 部分成功语义
|
||||
|
||||
**问题**:当一批子代理中部分失败时,dispatch_all 应该整体失败还是返回部分成功的 `Vec`?
|
||||
|
||||
**来源**:`docs/note-opencode-subagent-dispatch.md` — PM 和 SA 一致认为应返回部分成功语义 [高]。
|
||||
|
||||
**结论**:返回 `Vec<Result<SubTaskResult, EngineError>>`。调用方可迭代检查每个结果,失败条目保留 checkpoint 以便审计 [高]。
|
||||
|
||||
### 4. dispatch_stream 生命周期
|
||||
|
||||
**问题**:`dispatch_stream` 需要返回一个流,流内部要做 `submit_turn_stream` + `finalize_active_stream`。session 所有权和生命周期如何管理?
|
||||
|
||||
**来源**:
|
||||
- `src/engine/session_manager.rs:390-412` — `submit_turn_stream` 的锁模式:短持锁获取流后立即释放 [高]
|
||||
- `src/agent/session.rs:384-445` — `submit_turn_stream` 自身不持有跨 await 的锁 [高]
|
||||
- `docs/note-opencode-subagent-dispatch.md` — SA 建议在 AgentSession 新增 `finalize_active_stream()` 内部方法 [中]
|
||||
|
||||
**结论**:`dispatch_stream` 使用 `&Arc<Self>` 签名 + `tokio::spawn`。内部流管道:create_child → inherit_memory → submit_turn_stream → mpsc channel 转发事件 → 流消费完毕后调用 `finalize_active_stream()`。session 通过 `Arc<Mutex<AgentSession>>` 在 spawned task 中持有 [高]。
|
||||
|
||||
### 5. Cargo.toml 依赖分析
|
||||
|
||||
**来源**:
|
||||
- `Cargo.toml:17` — `futures-util = "0.3"` 已在依赖中 [高]
|
||||
- `Cargo.toml:15` — `tokio-stream = "0.1"` 已在依赖中 [高]
|
||||
- `Cargo.toml:16` — `futures = "0.3"` 已在依赖中 [高]
|
||||
|
||||
**结论**:dispatch_stream 所需的 `StreamExt` / `ReceiverStream` 所需的基础设施已全部就绪,无需新增任何依赖 [高]。
|
||||
|
||||
---
|
||||
|
||||
## 可选方案
|
||||
|
||||
### 方案 A:switch_agent 作为 AgentSession 方法 vs SessionManager 方法
|
||||
|
||||
| 维度 | A1: AgentSession 方法 | A2: SessionManager 方法 |
|
||||
|------|-----------------------|------------------------|
|
||||
| 实现位置 | `agent/session.rs` | `engine/switch.rs` |
|
||||
| 职责归属 | session 实例级 | 管理器级 |
|
||||
| 能否更新 SessionMeta | 不能(无 store 引用) | 能(有 store + checkpointer) |
|
||||
| 能否做自动 checkpoint | 不能(无 checkpointer) | 能 |
|
||||
| 与 create_child / replace 对齐 | 不对齐(create 在 SM) | 对齐(都在 SM) |
|
||||
|
||||
**来源**:
|
||||
- `src/agent/session.rs:49-71` — AgentSession 不持有 store / checkpointer 引用 [高]
|
||||
- `src/engine/session_manager.rs:325-347` — `replace()` 是 SM 方法,涉及 meta 持久化 [高]
|
||||
- roadmap lines 838 — `switch_agent(session_id, new_agent)` 签名暗示 SM 方法 [中]
|
||||
|
||||
**结论**:采用 A2(SessionManager 方法)。AgentSession 没有 store 引用,无法更新 SessionMeta。独立文件 `engine/switch.rs` 作为 SessionManager 的 impl 块。
|
||||
|
||||
### 方案 B:bridge_keys 注入方式
|
||||
|
||||
| 维度 | B1: SessionMemory 副本继承 + 过滤 | B2: 修改 Agent trait |
|
||||
|------|------------------------------------|----------------------|
|
||||
| 系统 prompt 侵入性 | 无 | 需新增 `set_bridge_data()` 方法 |
|
||||
| switch_agent 兼容性 | 天然兼容(与 agent 解耦) | switch 后需重新注入 |
|
||||
| 实现复杂度 | 一个私有辅助函数 | 需改 Agent trait + 所有实现 |
|
||||
| 测试增量 | 小(只测 `inherit_session_memory`) | 大(需测所有 Agent impl) |
|
||||
|
||||
**来源**:
|
||||
- `src/agent/agent.rs:16-30` — Agent trait 当前仅 3 个方法,简洁 [高]
|
||||
- `docs/note-opencode-subagent-dispatch.md` — "bridge_keys 通过 slot 注入,不修改 Agent system_prompt" [高]
|
||||
|
||||
**结论**:采用 B1。隔离关注点:Agent 负责"角色",SessionMemory 负责"桥接数据"。
|
||||
|
||||
### 方案 C:dispatch_stream 返回类型
|
||||
|
||||
| 维度 | C1: `Pin<Box<dyn Stream<Item=SubTaskStreamEvent>+Send>>` | C2: 自定义 struct 包装 |
|
||||
|------|----------------------------------------------------------|------------------------|
|
||||
| 与现有 API 一致性 | 与 `submit_turn_stream` 一致 [高] | 不一致 |
|
||||
| 调用方灵活性 | 直接 `.next()` + StreamExt | 需解包装 |
|
||||
| 实现复杂度 | 直接返回 stream | 需额外 struct + 方法 |
|
||||
| 可组合性 | 高(可直接 map/filter/collect) | 低 |
|
||||
|
||||
**来源**:
|
||||
- `src/engine/session_manager.rs:390-412` — `submit_turn_stream` 返回 `Pin<Box<dyn Stream<Item=StreamEvent>+Send>>` [高]
|
||||
|
||||
**结论**:采用 C1。保持一致的模式,调用方可以 `StreamExt::collect` / `map` 等。
|
||||
|
||||
### 否决方案
|
||||
|
||||
| 方案 | 否决原因 |
|
||||
|------|----------|
|
||||
| switch_agent 做自动 checkpoint | 与 auto_checkpoint 语义不一致(submit_turn 才触发),用户可手动 checkpoint。来源:`src/engine/session_manager.rs:357-384` auto_checkpoint 仅在 submit_turn/finalize_turn 触发 |
|
||||
| AgentSession 中的 `Agent` 用 `Box<dyn Agent>` | 与现有 `Arc<dyn Agent>` 不一致,且 SessionSnapshot 不序列化 agent(`src/agent/session.rs:502`)。来源:`src/agent/session.rs:53` |
|
||||
| dispatch_all 返回所有成功再返回 | 需要调用方等待全部完成才能拿到第一个结果。Rust 已有 `JoinSet` / `FuturesUnordered` 可选,但 v0.3 先保持简单 |
|
||||
| child_memory 加额外权限控制 | 子 session 是父创建的,父天然有 destroy / read 权限。来源:`docs/note-opencode-subagent-dispatch.md` PM 明确"不做额外权限控制" |
|
||||
| 子 session 失败时保留 checkpoint | `destroy()` 调 `checkpointer.delete_all`(`session_manager.rs:508`),不保留持久化残留。失败路径的调试信息通过 `tracing::error!` 日志记录 |
|
||||
|
||||
---
|
||||
|
||||
## 推荐方案
|
||||
|
||||
### 整体架构
|
||||
|
||||
```
|
||||
SessionManager (existing)
|
||||
├── switch_agent(id, new_agent) → engine/switch.rs
|
||||
├── dispatch(parent, agent, task, cfg) → engine/sub_agent.rs
|
||||
├── dispatch_stream(parent, agent, task, cfg) → engine/sub_agent.rs
|
||||
└── dispatch_all(parent, tasks, cfg) → engine/sub_agent.rs
|
||||
```
|
||||
|
||||
`switch.rs` 和 `sub_agent.rs` 均为 SessionManager 的 `impl` 块文件,通过 `pub mod` 在 `engine/mod.rs` 中注册 [高]。
|
||||
|
||||
### 决策清单
|
||||
|
||||
| # | 决策 | 结论 | 理由 |
|
||||
|---|------|------|------|
|
||||
| D1 | switch_agent 位置 | SessionManager 方法,在 `engine/switch.rs` | 需 store 更新 SessionMeta,AgentSession 无 store 引用 |
|
||||
| D2 | bridge_keys 注入方式 | SessionMemory 副本继承 + 过滤,不碰 system_prompt | 概念正交,switch_agent 友好 |
|
||||
| D3 | Memory 继承策略 | 快照副本(list_entries → set_with_meta) | 父子隔离优先,O(n) 拷贝可接受 |
|
||||
| D4 | dispatch_all 返回类型 | `Vec<Result<SubTaskResult, EngineError>>` | 部分成功语义,Rust-idiomatic |
|
||||
| D5 | dispatch_stream 返回类型 | `Pin<Box<dyn Stream<Item=SubTaskStreamEvent>+Send>>` | 与 `submit_turn_stream` 一致 |
|
||||
| D6 | switch_agent 的 lock 策略 | get() 读锁立即释放 → Mutex lock → 替换 → 释放 Mutex → I/O | 严格遵循已有锁契约 |
|
||||
| D7 | switch checkpoint | 不自动 checkpoint | 与 `auto_checkpoint` 语义一致(仅 submit_turn/finalize_turn 触发) |
|
||||
| D8 | dispatch 失败清理 | destroy 子 session | 不留僵尸 session |
|
||||
| D9 | dispatch 失败时 checkpoint | destroy 清理全部(含 checkpoint) | `destroy()` 内部调 `checkpointer.delete_all`,不留持久化垃圾 |
|
||||
| D10 | dispatch_all 并发控制 | `tokio::sync::Semaphore` | 轻量、内建、语义清晰 |
|
||||
| D11 | dispatch_all / dispatch_stream 签名 | `self: &Arc<Self>` | 满足 `tokio::spawn` `'static` 约束 |
|
||||
| D12 | 子 session 创建时 bundle | 从父 session 的 `RuntimeBundle` clone | 复用 `create_child` 已有逻辑 |
|
||||
|
||||
### 设计理由详述
|
||||
|
||||
**D1 为什么 switch_agent 必须放在 SessionManager 下**:因为 switch 后需要更新持久化的 SessionMeta(`agent_name` 变化),而 `save_session_meta` 需要 `&self.store`。AgentSession 不持有 store 引用(纯内存对象)。如果放在 AgentSession 上,要么给它加 store 引用(开历史倒车),要么让调用方手动调 `save_session_meta`(容易遗漏)[高]。
|
||||
|
||||
**D3 为什么选副本而非引用**:隔离性优先原则。父 session 可能在子运行期间 `set` 新数据,引用共享会导致子看到父的运行时中间状态;副本确保子看到的是 dispatch 时刻的稳定快照。性能方面,SessionMemory 条目数通常 < 100,O(n) 拷贝可忽略 [高]。
|
||||
|
||||
**D5 dispatch_stream 方案**:核心挑战是 session 生命周期管理和消息 finalize。方案使用 spawn task + `tokio::sync::mpsc::unbounded_channel`。spawn task 通过事件追踪重建消息列表:记录 `user_input` 作为首条 `UserMessage`,从 `StreamEvent::ToolExecutionCompleted` 事件提取工具结果,从 `StreamEvent::MessageComplete` 提取完整响应。流结束后直接调用 `AgentSession::finalize_turn(response, new_messages).await`。当 receiver 端 drop 时 sender 侧的 `send()` 错误会被捕获,task 内清理。`dispatch_stream` 返回的 stream 发出 `SubTaskStreamEvent::ChildCreated`(先导)+ `Stream(StreamEvent)`(中间,透传)+ `Completed(SubTaskResult)`(最终),消费者无需额外调 finalize [高]。
|
||||
|
||||
## 实施建议
|
||||
|
||||
### 阶段划分
|
||||
|
||||
共 9 个步骤,建议依次实施,不可并行。总预估时间由实现者在实施时评估。
|
||||
|
||||
#### Step 1 — `error.rs` 扩展
|
||||
|
||||
**文件**:`src/engine/error.rs`
|
||||
**内容**:EngineError 追加 3 个变体
|
||||
- `DispatchFailed(String)` — 子代理调度通用失败
|
||||
- `SwitchFailed(String)` — 角色切换失败
|
||||
- `SubAgentStreamError { child_id: String, detail: String }` — 流式调度中的子代理错误
|
||||
|
||||
**验证**:`cargo build` 成功。
|
||||
**注意**:已存在的 `#[non_exhaustive]` 属性确保这不是 breaking change [高]。
|
||||
|
||||
#### Step 2 — `session.rs` 无变更
|
||||
|
||||
**文件**:`src/agent/session.rs` — 不修改现有 API。
|
||||
|
||||
**审查发现**:第一轮审查确认 `finalize_active_stream()` 假设不成立。`submit_turn_stream`(`session.rs:384-445`)返回 stream 后 `LlmCycle` 即被 drop(`cycle.rs:647` 通过 `std::mem::take` 移出消息),不存在"active stream 内部状态"可读取。
|
||||
|
||||
**结论**:改为在 `dispatch_stream` 的 spawn task 中**从 StreamEvent 序列重建消息列表**,直接调用已有的 `finalize_turn(response, new_messages).await`。详见 Step 7 第 3 项。
|
||||
|
||||
**验证**:不修改 `session.rs`,Step 7 实施前 `cargo build` 可通过。
|
||||
|
||||
#### Step 3 — `switch.rs`
|
||||
|
||||
**文件**:`src/engine/switch.rs`
|
||||
**内容**:
|
||||
|
||||
```rust
|
||||
impl SessionManager {
|
||||
/// 热切换指定 session 的 Agent 角色。
|
||||
///
|
||||
/// - 保留 slot 历史 / turn_index / session_memory / cost_so_far
|
||||
/// - 自动更新 SessionMeta 中的 agent_name(保持原始 created_at / parent_id)
|
||||
/// - 不自动 checkpoint(与 `auto_checkpoint` 语义一致:仅 submit_turn 触发)
|
||||
/// - **注意**: 切换后新的 system_prompt 将与已有对话历史共存。
|
||||
/// 建议在切换后发送一条明确的上下文过渡提示
|
||||
/// (如"你现在以新角色 X 的身份继续对话")作为切换后的首条输入。
|
||||
/// - **安全提示**: `AgentSession.agent` 是 `pub` 字段可直接访问,
|
||||
/// 绕过 `switch_agent` 直接修改会导致 SessionMeta 中的 agent_name
|
||||
/// 与内存状态不一致,请始终使用此方法。
|
||||
pub async fn switch_agent(
|
||||
&self,
|
||||
session_id: &str,
|
||||
new_agent: Arc<dyn Agent>,
|
||||
) -> Result<(), EngineError> {
|
||||
// 1. get session(RwLock 读锁,返回后释放)
|
||||
let session = self.get(session_id).await?;
|
||||
|
||||
// 2. lock Mutex,替换 agent,读 name + turn_index
|
||||
let (agent_name, turn_index) = {
|
||||
let mut guard = session.lock().await;
|
||||
guard.agent = new_agent;
|
||||
(guard.agent.name().to_string(), guard.turn_index())
|
||||
}; // 释放 Mutex
|
||||
|
||||
// 3. 读取原始 SessionMeta(用于保留 created_at / parent_id)
|
||||
let existing_meta = self
|
||||
.load_session_meta(session_id)
|
||||
.await?
|
||||
.ok_or_else(|| EngineError::SessionNotFound(session_id.to_string()))?;
|
||||
|
||||
// 4. 构造新 meta 并持久化(I/O,无锁)
|
||||
let meta = SessionMeta {
|
||||
session_id: session_id.to_string(),
|
||||
agent_name,
|
||||
parent_id: existing_meta.parent_id,
|
||||
created_at: existing_meta.created_at,
|
||||
turn_count: turn_index,
|
||||
};
|
||||
self.save_session_meta(&meta).await?;
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
agent_name = %meta.agent_name,
|
||||
previous_agent = %existing_meta.agent_name,
|
||||
"agent switched"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**测试**(预计 4 个):
|
||||
1. 基本切换:switch 后 `agent.name()` 返回新 name
|
||||
2. 上下文保留:turn_index / session_memory / slot 历史均不变
|
||||
3. SessionMeta 持久化:`load_session_meta` 验证 agent_name 已更新
|
||||
4. 不存在的 session:返回 `SessionNotFound`
|
||||
|
||||
**验证**:`cargo test --all-targets` + clippy
|
||||
|
||||
#### Step 4 — `sub_agent.rs` 类型
|
||||
|
||||
**文件**:`src/engine/sub_agent.rs`
|
||||
**内容**:3 个类型定义
|
||||
|
||||
```rust
|
||||
/// 子代理调度配置。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct DispatchConfig {
|
||||
/// 最大并发数(dispatch_all 用)。默认 10。
|
||||
pub max_concurrency: usize,
|
||||
/// 是否继承父 SessionMemory。默认 true。
|
||||
pub inherit_session_memory: bool,
|
||||
/// 桥接 key 列表:
|
||||
/// - `None` = 不继承任何父 SessionMemory
|
||||
/// - `Some(vec![])` = 继承全部父 SessionMemory
|
||||
/// - `Some(keys)` = 仅继承指定的 keys
|
||||
/// 默认 `None`(零继承),显式选择加入。
|
||||
pub bridge_keys: Option<Vec<String>>,
|
||||
/// 子↔子共享 namespace。如果为 `Some(prefix)`,
|
||||
/// 子 agent 可通过 `session.get_session_data(key)` 访问
|
||||
/// `shared:{prefix}:{key}` 命名空间的数据。
|
||||
/// 默认 `Some(parent_session_id)`。
|
||||
pub shared_namespace: Option<String>,
|
||||
}
|
||||
|
||||
impl Default for DispatchConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_concurrency: 10,
|
||||
inherit_session_memory: true,
|
||||
bridge_keys: None,
|
||||
shared_namespace: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 子代理执行结果。
|
||||
///
|
||||
/// dispatch 成功后子 session **保留在 SessionManager 中**,调用方可
|
||||
/// 通过 `sm.get(&result.child_id)` 获取子 session 引用,进而通过
|
||||
/// `session_memory()` 读取子 SessionMemory(如 "result_summary")。
|
||||
#[derive(Debug)]
|
||||
pub struct SubTaskResult {
|
||||
/// 子 session ID。可通过此 ID 在 SessionManager 中读取子 session。
|
||||
pub child_id: String,
|
||||
/// LLM 最终响应。
|
||||
pub response: MessageResponse,
|
||||
/// 本次调用的 token 用量。
|
||||
pub usage: CostTracker,
|
||||
/// 可选摘要(读取子 session_memory 中的 "result_summary")。
|
||||
pub summary: Option<String>,
|
||||
}
|
||||
```
|
||||
|
||||
```rust
|
||||
/// 流式子代理调度事件。
|
||||
#[derive(Debug)]
|
||||
pub enum SubTaskStreamEvent {
|
||||
/// 子 session 已创建(携带 child_id)。
|
||||
ChildCreated { child_id: String },
|
||||
/// LLM 流事件(透传)。
|
||||
Stream(StreamEvent),
|
||||
/// 执行完成(携带完整结果)。
|
||||
Completed(SubTaskResult),
|
||||
}
|
||||
```
|
||||
|
||||
`SubTaskStreamEvent` 需实现 `Display` 和 `std::error::Error`(`Completed` 和 `ChildCreated` 不触发错误路径,`Display` 仅用于调试日志)[中]。
|
||||
|
||||
**验证**:`cargo build`
|
||||
|
||||
#### Step 5 — `sub_agent.rs dispatch` 核心
|
||||
|
||||
**文件**:`src/engine/sub_agent.rs`
|
||||
**内容**:
|
||||
|
||||
```rust
|
||||
impl SessionManager {
|
||||
/// 私有辅助:从父 session memory 继承条目到子 session。
|
||||
///
|
||||
/// **一致性模型**:捕获的是调用时刻的父 session_memory 快照。
|
||||
/// 即使在 `list_entries()` 返回后、`set_with_meta()` 写入前
|
||||
/// 父 session 被并发写入新数据,子 session 也**不会**看到这些
|
||||
/// 新数据(快照副本的内生特征)。[审查确认]
|
||||
async fn inherit_session_memory(
|
||||
&self,
|
||||
parent_id: &str,
|
||||
child_id: &str,
|
||||
config: &DispatchConfig,
|
||||
) -> Result<(), EngineError> { /* ... */ }
|
||||
|
||||
/// 派发一个子任务,返回结构化结果。
|
||||
pub async fn dispatch(
|
||||
&self,
|
||||
parent_id: &str,
|
||||
sub_agent: Arc<dyn Agent>,
|
||||
task: impl Into<String>,
|
||||
config: DispatchConfig,
|
||||
) -> Result<SubTaskResult, EngineError> { /* ... */ }
|
||||
}
|
||||
```
|
||||
|
||||
**dispatch 流程**:
|
||||
1. `create_child(parent_id, sub_agent)` → 获取 child_id
|
||||
2. 若 `config.inherit_session_memory == true` → `inherit_session_memory(parent_id, child_id, config)`
|
||||
3. `submit_turn(child_id, task)` → 获取 response
|
||||
4. 读取 `"result_summary"`(可选)
|
||||
5. 返回 `SubTaskResult`(子 session 保留在 SessionManager 中,可通过 `sm.get(&child_id)` 读取 child_memory)
|
||||
6. 失败路径:`let _ = self.destroy(&child_id).await; tracing::error!(...)`(静默吞掉清理错误,**原始 EngineError 优先**;`destroy` 会清理子 session 的 SessionMeta + checkpoint 条目,不留僵尸)
|
||||
|
||||
**测试**(预计 5 个):
|
||||
1. 基本调度:子 agent 返回预期响应
|
||||
2. bridge_keys 过滤:仅指定的 key 被继承
|
||||
3. memory 继承:父 set 的值子可读到
|
||||
4. submit_turn 失败:错误传播 + 子 session 被销毁
|
||||
5. 无效 parent_id:返回 `SessionNotFound`
|
||||
|
||||
**验证**:`cargo test --all-targets`
|
||||
|
||||
#### Step 6 — `sub_agent.rs dispatch_all`
|
||||
|
||||
**文件**:`src/engine/sub_agent.rs`
|
||||
**内容**:
|
||||
|
||||
```rust
|
||||
impl SessionManager {
|
||||
/// 并行派发一批子任务。
|
||||
pub async fn dispatch_all(
|
||||
self: &Arc<Self>,
|
||||
parent_id: &str,
|
||||
tasks: Vec<(Arc<dyn Agent>, String)>,
|
||||
config: DispatchConfig,
|
||||
) -> Vec<Result<SubTaskResult, EngineError>> { /* ... */ }
|
||||
}
|
||||
```
|
||||
|
||||
**设计要点**:
|
||||
- 使用 `tokio::sync::Semaphore` 限制并发数(默认 `config.max_concurrency`)
|
||||
- **Semaphore acquire 在 spawn 内**:`let permit = semaphore.clone().acquire_owned().await;` — permit 所有权转移到 spawned task。避免 spawn N 个 task 时全量分配 Future 内存 [审查修复]
|
||||
- 每个 task `tokio::spawn` + `Arc<Self>` clone
|
||||
- 内部调用 `dispatch` 的同类逻辑(create_child → inherit → submit_turn)
|
||||
- **indexed 收集**:预分配 `Vec<Option<Result<...>>>` 按 `tasks` 索引填入,维持输入顺序。不使用排序(排序需等所有 child_id 生成后)[审查修复]
|
||||
- 每个结果独立:`Ok(SubTaskResult)` 或 `Err(EngineError)`
|
||||
|
||||
**测试**(预计 4 个):
|
||||
1. 并行 3 个全部成功
|
||||
2. 部分失败(MockProvider 对特定 task 返回错误)
|
||||
3. Semaphore 上限验证(max_concurrency=1 时串行执行)
|
||||
4. 空 tasks 列表
|
||||
|
||||
**验证**:`cargo test --all-targets`
|
||||
|
||||
#### Step 7 — `sub_agent.rs dispatch_stream`
|
||||
|
||||
**文件**:`src/engine/sub_agent.rs`
|
||||
**内容**:
|
||||
|
||||
```rust
|
||||
impl SessionManager {
|
||||
pub async fn dispatch_stream(
|
||||
self: &Arc<Self>,
|
||||
parent_id: &str,
|
||||
sub_agent: Arc<dyn Agent>,
|
||||
task: impl Into<String>,
|
||||
config: DispatchConfig,
|
||||
) -> Result<
|
||||
Pin<Box<dyn Stream<Item = SubTaskStreamEvent> + Send>>,
|
||||
EngineError,
|
||||
> { /* ... */ }
|
||||
}
|
||||
```
|
||||
|
||||
**设计要点**:
|
||||
- 同步部分(lock 外):create_child + inherit_memory
|
||||
- 获取 Stream 后通过 `tokio::sync::mpsc::unbounded_channel` 转发事件(与 LLM stream 内部背压策略一致,避免有界 channel 的 sender 阻塞风险)[审查修复]
|
||||
- spawn task 持有 `Arc<Mutex<AgentSession>>` 消费 LLM stream
|
||||
- **消息重建机制**(替代已移除的 `finalize_active_stream()`):[审查修复]
|
||||
```
|
||||
// 在 spawn task 中:
|
||||
let mut new_messages: Vec<Message> = vec![Message::user_text(&task)];
|
||||
let mut final_response: Option<MessageResponse> = None;
|
||||
|
||||
while let Some(event) = llm_stream.next().await {
|
||||
// 转发事件到输出 channel
|
||||
tx.send(SubTaskStreamEvent::Stream(event.clone()))?;
|
||||
// 从 ToolExecutionCompleted 构造 ToolResult 消息
|
||||
if let StreamEvent::ToolExecutionCompleted { tool_name, tool_call_id, input, output } = &event {
|
||||
new_messages.push(Message::tool_result(tool_call_id, tool_name, output));
|
||||
}
|
||||
// 捕获最终响应
|
||||
if let StreamEvent::MessageComplete(ref resp) = event {
|
||||
final_response = Some(resp.clone());
|
||||
}
|
||||
}
|
||||
// 流结束后,追加 assistant 消息并 finalize
|
||||
if let Some(response) = &final_response {
|
||||
new_messages.push(response.message.clone());
|
||||
child_session.lock().await
|
||||
.finalize_turn(response, new_messages).await?;
|
||||
}
|
||||
```
|
||||
- 事件序列:`ChildCreated` → `Stream(StreamEvent)` × N → `Completed(SubTaskResult)`
|
||||
- 消费者 drop receiver → unbounded channel sender 错误 → task 自动退出
|
||||
|
||||
**测试**(预计 4 个):
|
||||
1. 事件序列验证:收到 ChildCreated → 至少一个 Stream → Completed
|
||||
2. 错误传播:LLM 内部错误 → 正确映射到 error 事件
|
||||
3. receiver dropped:drop receiver 后 task 正确退出,不 panic
|
||||
4. finalize 正确性:Completed 中的 usage / summary 正确
|
||||
|
||||
**验证**:`cargo test --all-targets`
|
||||
|
||||
#### Step 8 — `mod.rs` + 集成验证
|
||||
|
||||
**文件**:`src/engine/mod.rs`
|
||||
**内容**:追加 `pub mod switch;` 和 `pub mod sub_agent;` + `pub use`
|
||||
|
||||
```rust
|
||||
pub mod checkpointer;
|
||||
pub mod error;
|
||||
pub mod session_manager;
|
||||
pub mod snapshot;
|
||||
pub mod switch; // <-- 新增
|
||||
pub mod sub_agent; // <-- 新增
|
||||
```
|
||||
|
||||
**验证**:
|
||||
1. `cargo test --all-targets` — 374+ 测试全部通过
|
||||
2. `cargo clippy --all-targets -- -D warnings` — 0 警告
|
||||
3. `cargo doc --no-deps` — 0 warning
|
||||
|
||||
#### Step 9 — 示例
|
||||
|
||||
**文件 1**:`examples/agent_switch_demo.rs`(约 80 行)
|
||||
- 创建 session → submit_turn(角色 A)→ switch_agent(角色 B)→ submit_turn(角色 B)→ 验证上下文保留
|
||||
- 演示目的:证明 switch_agent 保留 slot 历史 / turn_index / session_memory
|
||||
|
||||
**文件 2**:`examples/sub_agent_dispatch_demo.rs`(约 150 行)
|
||||
- 父 session → dispatch_all 3 个子 agent(研究、写作、审校)→ 收集结果 → 父汇总
|
||||
- 树形验证:`children(parent_id)` 返回 3 个子 ID
|
||||
- 演示目的:多 agent 协作完整链路
|
||||
|
||||
**文件 3**:`examples/bridge_keys_demo.rs`(约 140 行)
|
||||
- 父设置 SessionMemory(key: "project_goal", "constraints")→ dispatch + bridge_keys → 子 agent 通过 `get_session_data` 读取
|
||||
- 子↔子交互:父通过 `DispatchConfig.shared_namespace` 设定共享命名空间,子 A 写入 `shared:{parent_id}:fact_x`,子 B 通过约定 key 读取
|
||||
- 演示目的:bridge_keys 过滤机制 + 父子数据桥接 + 子↔子共享 namespace
|
||||
|
||||
**文件 4 — 新增**:`examples/dispatch_stream_demo.rs`(约 100 行)
|
||||
- 父 session → dispatch_stream 单个子 agent → 消费 `SubTaskStreamEvent` 序列
|
||||
- 验证收到 `ChildCreated` + 至少一个 `Stream` + `Completed` 事件
|
||||
- 输出 `SubTaskResult.child_id` / `usage` / `summary`,验证消息重建和 finalize 正确性
|
||||
- 演示 `receiver dropped` 场景:中途 drop receiver 后 task 正确退出不 panic
|
||||
- 演示目的:dispatch_stream 的事件序列 + finalize 完整性验证
|
||||
|
||||
**验证**:4 个示例全部 `cargo run --example` exit 0
|
||||
|
||||
### 高层实施建议
|
||||
|
||||
1. **Step 1 优先于所有步骤**:Error 扩展是所有后续步骤的基础,无依赖可并行 [高]
|
||||
2. **Step 3 独立性强**:switch_agent 不依赖 dispatch 的任何类型,可单独实施和测试 [高]
|
||||
3. **Step 4 是 Step 5-7 的前置**:类型定义不依赖其他逻辑,建议在 Step 3 完成后立即实施 [高]
|
||||
4. **Step 5-7 按复杂度递增**:dispatch → dispatch_all → dispatch_stream。dispatch_all 复用 dispatch 的核心逻辑;dispatch_stream 是最复杂的,建议最后实施 [高]
|
||||
5. **Step 8 集成验证不可跳过**:clippy + doc 全量验证确保无回归 [高]
|
||||
6. **Step 9 在所有核心完成后实施**:示例是验收标准的一部分,PM 确认 3 个递进示例 [中]
|
||||
7. **全量测试密码**:实施过程中持续 `cargo test --all-targets`,不在最后统一修复 [高]
|
||||
|
||||
### 风险矩阵
|
||||
|
||||
| # | 风险 | 等级 | 可能性 | 对策 |
|
||||
|---|------|------|--------|------|
|
||||
| R1 | dispatch_stream 的消息重建:从 StreamEvent 序列重建 `new_messages` 的完整性 | 🟡 中 | 低 | 已移除 `finalize_active_stream()` 方案。spawn task 通过追踪 `ToolExecutionCompleted` / `MessageComplete` 事件重建消息列表。完整性由 `MessageComplete` 事件保证 |
|
||||
| R2 | `tokio::spawn` `'static` + SessionManager 引用 | 🟡 中 | 低 | dispatch_all 和 dispatch_stream 签名已明确用 `&Arc<Self>`;调用方包装 `Arc<SessionManager>` |
|
||||
| R3 | 子 session 创建成功但后续 submit_turn 失败 | 🟡 中 | 中 | dispatch 内 `destroy(child_id)` 放在 `?` 前确保清理;通过 `let child_id = ...;` 先绑定,再 `let r = submit_turn(...).await`,失败时 `destroy(&child_id).await?` 清理 |
|
||||
| R4 | 部分失败时孤儿 checkpoint 数据 | 🟢 低 | 必然 | **正向利用**:保留用于调试审计,存储开销可忽略 |
|
||||
| R5 | 并发 dispatch_all 中任务 panic | 🟡 中 | 低 | `tokio::spawn` 的 `JoinHandle` 通过 `.await` 捕获 panic;panic 传播到 `dispatch_all` 内作为 `Err` 返回 |
|
||||
| R6 | bridge_keys 中不存在的 key | 🟢 低 | 中 | 静默跳过(与 `List_entries` 返回全量后再过滤,不存在的 key 自然不会出现在结果中) |
|
||||
|
||||
### 架构图(文本示意)
|
||||
|
||||
```
|
||||
SessionManager
|
||||
/ | \
|
||||
/ | \
|
||||
switch_agent dispatch dispatch_stream
|
||||
| | |
|
||||
v v v
|
||||
AgentSession create_child create_child
|
||||
.agent = new inherit_mem inherit_mem
|
||||
SessionMeta submit_turn submit_turn_stream
|
||||
更新 返回结果 mpsc 转发事件
|
||||
finalize_on_complete
|
||||
|
||||
交互层次:
|
||||
父 -> 子: SessionMemory snapshot + bridge_keys 过滤
|
||||
子 -> 父: SubTaskResult { child_id, response, usage, summary }
|
||||
子 <-> 子: shared:{parent_session_id} namespace
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 参考来源
|
||||
|
||||
### 代码路径
|
||||
|
||||
| 文件 | 用途 |
|
||||
|------|------|
|
||||
| `src/agent/session.rs` | AgentSession 定义、`agent` pub 字段(L53)、`session_memory` pub 字段(L58)、`submit_turn_stream`(L384)、`finalize_turn`(L456) |
|
||||
| `src/agent/agent.rs` | Agent trait 定义(3 个方法) |
|
||||
| `src/agent/session_memory.rs` | SessionMemory: `set`(L50)、`set_with_meta`(L61)、`list_entries`(L81) |
|
||||
| `src/engine/session_manager.rs` | SessionManager: `create_child`(L202)、`get`(L244)、`destroy`(L495)、`save_session_meta`(L131)、`load_session_meta`(L144)、SessionMeta(L32)、锁契约(L4-12) |
|
||||
| `src/engine/error.rs` | EngineError 枚举(当前 6 变体,`#[non_exhaustive]`) |
|
||||
| `src/engine/mod.rs` | 模块注册 |
|
||||
| `src/engine/snapshot.rs` | SessionSnapshot(from_snapshot / to_snapshot 所需) |
|
||||
| `src/llm/stream.rs` | StreamEvent 枚举 |
|
||||
| `Cargo.toml` | 依赖声明(`futures-util` L17、`tokio-stream` L15、`futures-core` L18) |
|
||||
|
||||
### 文档路径
|
||||
|
||||
| 文档 | 用途 |
|
||||
|------|------|
|
||||
| `docs/roadmap-v0.3.0.md` §Phase 18 | Phase 18 原始需求(交付物、交互层级、优先级) |
|
||||
| `docs/note-opencode-agent-switching.md` | Agent 热切换调研笔记(桥接方案分析、生命周期讨论) |
|
||||
| `docs/note-opencode-subagent-dispatch.md` | SubAgent Dispatch 调研笔记(PM/SA 建议、设计推演) |
|
||||
| `docs/23-phase17-agent-execution-engine.md` | Phase 17 方案文档(SessionManager 设计背景) |
|
||||
| `docs/7-agent-runtime.md` | Agent 运行时设计文档(Session 与 Agent 的关系) |
|
||||
| `docs/17-phase10-contextslot.md` | ContextSlot 上下文管理(Phase 10) |
|
||||
| `docs/24-phase18-agent-switch-and-dispatch.md` | 本文档 — 第 1 轮审查修复记录 |
|
||||
|
||||
### 决策轨迹
|
||||
|
||||
| 决策 | 参考来源 | 置信度 |
|
||||
|------|----------|--------|
|
||||
| switch_agent 在 SessionManager 而非 AgentSession | `src/agent/session.rs` AgentSession 无 store 引用 | 高 |
|
||||
| bridge_keys 通过 SessionMemory 副本,不碰 system_prompt | `docs/note-opencode-subagent-dispatch.md` PM/SA 建议 | 高 |
|
||||
| dispatch_all 返回 `Vec<Result<..>>` 部分成功 | `docs/note-opencode-subagent-dispatch.md` PM/SA 一致 | 高 |
|
||||
| dispatch_stream 返回 `Pin<Box<dyn Stream>>` | `src/engine/session_manager.rs` `submit_turn_stream` 签名一致 | 高 |
|
||||
| 失败时 destroy 子 session 清理全部(含 checkpoint) | `session_manager.rs:508` destroy 调 `checkpointer.delete_all` | 高 |
|
||||
| dispatch_all 用 `&Arc<Self>` 签名 | `tokio::spawn` `'static` 约束 | 高 |
|
||||
| child_memory 不做额外权限控制 | `docs/note-opencode-subagent-dispatch.md` PM 明确 | 高 |
|
||||
| `#[non_exhaustive]` 已存在 -> 新增 EngineError 变体不是 breaking change | `src/engine/error.rs:15` | 高 |
|
||||
| `futures-util` 已存在,无需新增依赖 | `Cargo.toml:17` | 高 |
|
||||
|
||||
### 审查修复轨迹(第 1 轮)
|
||||
|
||||
| # | 问题 | 🔴/🟡 | 修复内容 |
|
||||
|---|------|--------|---------|
|
||||
| F1 | `finalize_active_stream()` 假设不成立:`submit_turn_stream` 返回后 cycle 被 drop | 🔴 | 移除 Step 2 的 `finalize_active_stream()`,改为 spawn task 内从 StreamEvent 重建消息列表直接调 `finalize_turn()` |
|
||||
| F2 | `SubTaskResult` 缺 child_memory 访问路径 | 🔴 | `SubTaskResult.child_id` 可经由 `sm.get()` 读取子 session。构型 doc comment 增加说明 |
|
||||
| F3 | dispatch 失败路径 `destroy` 自身 I/O 可能失败,覆盖原始错误 | 🔴 | 改用 `let _ = destroy` + `tracing::error!`,原始 `EngineError` 优先 |
|
||||
| F4 | switch_agent SessionMeta 构造中 `created_at`/`parent_id` 用占位符 | 🔴 | 改用 `load_session_meta` 读取原始值,`turn_count` 从 `guard.turn_index()` 读取 |
|
||||
| F5 | D9 与 `destroy()` 实现矛盾(方案说保留,代码说删除) | 🟡 | D9 修正为"destroy 清理全部",否决条目同步更新 |
|
||||
| F6 | `bridge_keys` 默认值安全反直觉(空=全量继承) | 🟡 | 类型改为 `Option<Vec<String>>`,`None` = 不继承(默认),`Some(vec![])` = 全量 |
|
||||
| F7 | Semaphore acquire 位置未指定 | 🟡 | 指定 `acquire_owned()` 在 spawn 内 + indexed 收集维持输入顺序 |
|
||||
| F8 | dispatch_stream 缺少独立示例 | 🟡 | Step 9 追加 `dispatch_stream_demo` |
|
||||
| F9 | mpsc channel 背压策略未指定 | 🟡 | 改用 `unbounded_channel`,与 LLM stream 内部模式一致 |
|
||||
| F10 | 子↔子交互层缺少实现细节 | 🟡 | `DispatchConfig.shared_namespace` 字段 + 示例 3 演示 |
|
||||
| F11 | inherit_session_memory 竞态窗口未文档化 | 🟡 | doc comment 声明快照一致性模型 |
|
||||
| F12 | switch_agent system_prompt 断裂风险未说明 | 🟡 | doc comment 增加使用建议 + 安全提示 |
|
||||
|
||||
---
|
||||
|
||||
*本文档对应的实施步骤记录在 `docs/roadmap-v0.3.0.md` §Phase 18,实施完成后同步更新 roadmap 状态。*
|
||||
@@ -0,0 +1,652 @@
|
||||
# Phase 19:知识图谱 + 双通道检索
|
||||
|
||||
## 背景与目标
|
||||
|
||||
### 问题空间
|
||||
|
||||
agcore v0.3.0 已交付 Phase 0-18,记忆系统具备 `KnowledgeStore`(页面级内容检索)和 `VectorStore`(向量语义检索),但缺少实体-关系维度的关联检索能力。用户搜索"X 与什么相关"时,现有系统无法返回实体间的拓扑关系。
|
||||
|
||||
`docs/note-knowledge-graph-design.md` 已记录完整的知识图谱设计,Phase 19 将其落地为可编译、可测试的模块。
|
||||
|
||||
### 目标
|
||||
|
||||
- 新增 `memory/graph.rs`,实现 `KnowledgeGraph` trait + `InMemoryGraph` 内存实现
|
||||
- 扩展 `MemoryRetriever` 为双通道:KnowledgeStore(内容)+ KnowledgeGraph(实体关系)
|
||||
- 通过 `RetrievalStrategy` 枚举控制通道选择(Hybrid / KnowledgeOnly / GraphOnly)
|
||||
- 统一 `RetrievalResult.items` 为 `Vec<RetrievalItem>`,enum 变体区分类别
|
||||
- 标签管理 API 预留(无自动提取流程,Agent 层显式写入)
|
||||
- Phase 19 仅提供底层 CRUD 接口,实体/关系的写入由 Agent 层(如 LLM 提取)在后续 Phase 中接入。当前无自动填充流程,需 Agent 显式调用 `add_entity`/`add_relation`。
|
||||
|
||||
### 与现有模块的定位关系
|
||||
|
||||
```
|
||||
KnowledgeStore: 页面级内容("什么是 X") ← Phase 6 已有
|
||||
VectorStore: 向量语义(相似度检索) ← Phase 15 已有
|
||||
KnowledgeGraph: 实体级关系("X 与什么相关") ← Phase 19 新增
|
||||
MemoryRetriever: 统一检索入口 ← Phase 19 扩展为双通道
|
||||
```
|
||||
|
||||
### 依赖与优先级
|
||||
|
||||
- **依赖**:Phase 6(KnowledgeStore)[高]、Phase 15(VectorStore 模式参考)[低]
|
||||
- **优先级**:P0(v0.3.0 最后一个 Phase)
|
||||
- **预估规模**:约 600 行核心 + 200 行测试
|
||||
|
||||
---
|
||||
|
||||
## 需求分析
|
||||
|
||||
### 功能需求
|
||||
|
||||
| ID | 需求 | 优先级 |
|
||||
|----|------|--------|
|
||||
| F1 | `GraphEntity` / `GraphRelation` / `RelationDirection` 类型定义 | P0 |
|
||||
| F2 | `KnowledgeGraph` trait(10 个 async 方法) | P0 |
|
||||
| F3 | `InMemoryGraph` 实现(HashMap + Vec + tag_index) | P0 |
|
||||
| F4 | BFS 图遍历(防环、权重衰减、方向过滤) | P0 |
|
||||
| F5 | `RetrievalItem` / `RetrievalStrategy` / `RetrievalResult` 扩展 | P0 |
|
||||
| F6 | `MemoryRetriever` 双通道(`tokio::join!` 并行) | P0 |
|
||||
| F7 | 标签管理(set_entity_tags / find_tags / entity_count_by_tag) | P1(预留) |
|
||||
|
||||
### 非功能需求
|
||||
|
||||
| ID | 需求 | 说明 |
|
||||
|----|------|------|
|
||||
| NF1 | 零新依赖 | 纯 std + tokio + 已有 crate |
|
||||
| NF2 | 异步安全 | `InMemoryGraph` 内部 Mutex 保护,trait 方法 async |
|
||||
| NF3 | 类型安全 | 不新增 `MemoryError` 变体,复用现有 5 个 |
|
||||
| NF4 | 向后兼容 | `MemoryRetriever::new()` 签名不变,可选链式注入 graph |
|
||||
| NF5 | Breaking change 受控 | `RetrievalResult.items` 类型变化,需在 CHANGELOG 标注 |
|
||||
|
||||
---
|
||||
|
||||
## 方案设计
|
||||
|
||||
### 3.1 数据模型
|
||||
|
||||
#### GraphEntity
|
||||
|
||||
```rust
|
||||
/// 图谱实体 —— 表示一个可被关联检索的节点。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct GraphEntity {
|
||||
/// 唯一标识(如 "person:rust-dev-01")。
|
||||
pub id: String,
|
||||
/// 实体名称(用于展示和关键词匹配)。
|
||||
pub name: String,
|
||||
/// 实体类型("person" | "concept" | "project" | ...)。
|
||||
pub entity_type: String,
|
||||
/// 一句话描述。
|
||||
pub description: String,
|
||||
/// 检索标签(全小写,原子词,由 Agent 层显式写入)。
|
||||
pub tags: Vec<String>,
|
||||
/// 任意附加属性(与 PersistentVectorStore.metadata 保持一致)。
|
||||
pub properties: HashMap<String, String>,
|
||||
}
|
||||
```
|
||||
|
||||
#### GraphRelation
|
||||
|
||||
```rust
|
||||
/// 图谱关系 —— 连接两个实体的有向边。
|
||||
///
|
||||
/// 无 `id` 字段,用 `(source_id, target_id, relation_type)` 三元组唯一标识。
|
||||
/// 提供 `composite_key()` 作为派生 id,满足未来独立 id 需求。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct GraphRelation {
|
||||
/// 源实体 ID。
|
||||
pub source_id: String,
|
||||
/// 目标实体 ID。
|
||||
pub target_id: String,
|
||||
/// 关系类型("works_on" | "part_of" | "related_to" | ...)。
|
||||
pub relation_type: String,
|
||||
/// 关系强度 [0.0, 1.0],用于 BFS 评分衰减。
|
||||
pub weight: f32,
|
||||
}
|
||||
|
||||
impl GraphRelation {
|
||||
/// 复合键:`source_id:target_id:relation_type`,用于去重和查找。
|
||||
pub fn composite_key(&self) -> String {
|
||||
format!("{}:{}:{}", self.source_id, self.target_id, self.relation_type)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
#### RelationDirection
|
||||
|
||||
```rust
|
||||
/// 关系遍历方向。
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum RelationDirection {
|
||||
/// 仅出边:source_id → target_id(默认)。
|
||||
Outgoing,
|
||||
/// 仅入边:target_id → source_id。
|
||||
Incoming,
|
||||
/// 双向遍历。
|
||||
Both,
|
||||
}
|
||||
|
||||
impl Default for RelationDirection {
|
||||
fn default() -> Self {
|
||||
Self::Outgoing
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
#### ScoredEntity
|
||||
|
||||
```rust
|
||||
/// 带评分的实体 + 路径信息。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ScoredEntity {
|
||||
pub entity: GraphEntity,
|
||||
/// 基于图距离的评分 [0.0, 1.0],沿路径权重乘积衰减。
|
||||
pub score: f32,
|
||||
/// 从查询实体到当前实体的 ID 路径(用于可解释性)。
|
||||
pub path: Vec<String>,
|
||||
}
|
||||
```
|
||||
|
||||
#### TagConstraints
|
||||
|
||||
```rust
|
||||
/// 标签约束配置。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TagConstraints {
|
||||
/// 每个实体最多标签数(默认 8)。
|
||||
pub max_tags_per_entity: usize,
|
||||
}
|
||||
|
||||
impl Default for TagConstraints {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_tags_per_entity: 8,
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 3.2 KnowledgeGraph trait
|
||||
|
||||
```rust
|
||||
/// 知识图谱抽象 —— 实体-关系存储与图遍历检索。
|
||||
///
|
||||
/// 所有方法 `async + Send + Sync`,支持跨 `.await` 调用。
|
||||
/// 复用 `MemoryError`,不新增变体。
|
||||
#[async_trait]
|
||||
pub trait KnowledgeGraph: Send + Sync {
|
||||
// ── 实体管理 ──
|
||||
|
||||
/// 添加或更新实体(upsert 语义)。
|
||||
async fn add_entity(&self, entity: GraphEntity) -> Result<(), MemoryError>;
|
||||
|
||||
/// 按 ID 获取实体,不存在返回 `Ok(None)`。
|
||||
async fn get_entity(&self, id: &str) -> Result<Option<GraphEntity>, MemoryError>;
|
||||
|
||||
/// 删除实体及其所有关联关系。
|
||||
async fn remove_entity(&self, id: &str) -> Result<(), MemoryError>;
|
||||
|
||||
// ── 关系管理 ──
|
||||
|
||||
/// 添加关系(若复合键已存在则覆盖 weight)。
|
||||
async fn add_relation(&self, relation: GraphRelation) -> Result<(), MemoryError>;
|
||||
|
||||
/// 按复合键删除关系。
|
||||
async fn remove_relation(
|
||||
&self,
|
||||
source_id: &str,
|
||||
target_id: &str,
|
||||
relation_type: &str,
|
||||
) -> Result<(), MemoryError>;
|
||||
|
||||
/// 从指定实体出发,BFS 遍历 depth 层,返回关联实体(带评分)。
|
||||
///
|
||||
/// - `direction`:遍历方向(Outgoing / Incoming / Both)
|
||||
/// - `relation_types`:可选过滤,仅遍历指定关系类型
|
||||
async fn get_related(
|
||||
&self,
|
||||
entity_id: &str,
|
||||
depth: usize,
|
||||
direction: RelationDirection,
|
||||
relation_types: Option<&[&str]>,
|
||||
) -> Result<Vec<ScoredEntity>, MemoryError>;
|
||||
|
||||
// ── 检索 ──
|
||||
|
||||
/// 按关键词子串匹配实体(不区分大小写),与 KnowledgeStore.search 一致。
|
||||
async fn find_by_keywords(&self, keywords: &[String]) -> Result<Vec<GraphEntity>, MemoryError>;
|
||||
|
||||
// ── 标签管理(预留接口,Agent 层显式写入) ──
|
||||
|
||||
/// 按前缀查找已有标签(用于标签复用)。
|
||||
async fn find_tags(&self, prefix: &str) -> Result<Vec<String>, MemoryError>;
|
||||
|
||||
/// 设置实体标签(替换式,保留前 max_tags_per_entity 个)。
|
||||
/// 返回实际设置的标签数。
|
||||
async fn set_entity_tags(
|
||||
&self,
|
||||
entity_id: &str,
|
||||
tags: Vec<String>,
|
||||
) -> Result<usize, MemoryError>;
|
||||
|
||||
/// 按标签统计实体数量。
|
||||
async fn entity_count_by_tag(&self, tag: &str) -> Result<usize, MemoryError>;
|
||||
|
||||
/// 获取标签约束配置。
|
||||
fn tag_constraints(&self) -> TagConstraints;
|
||||
}
|
||||
```
|
||||
|
||||
### 3.3 InMemoryGraph 实现
|
||||
|
||||
#### 内部结构
|
||||
|
||||
```rust
|
||||
/// 内存知识图谱实现 —— 纯内存,无持久化。
|
||||
///
|
||||
/// 生命周期跟随实例;持久化路径参考 InMemoryVectorStore → PersistentVectorStore 演进模式。
|
||||
pub struct InMemoryGraph {
|
||||
/// 内部状态(单一锁结构,避免嵌套锁死锁)
|
||||
inner: Mutex<GraphInner>,
|
||||
/// 标签约束
|
||||
constraints: TagConstraints,
|
||||
}
|
||||
|
||||
struct GraphInner {
|
||||
/// id → entity
|
||||
entities: HashMap<String, GraphEntity>,
|
||||
/// 所有关系(线性扫描,实测 5000 条 ≈ 1-50µs,无需邻接表索引)
|
||||
relations: Vec<GraphRelation>,
|
||||
/// tag → entity_ids(反向索引,用于 find_tags / entity_count_by_tag)
|
||||
tag_index: HashMap<String, HashSet<String>>,
|
||||
}
|
||||
```
|
||||
|
||||
#### BFS 遍历算法
|
||||
|
||||
```rust
|
||||
async fn get_related(
|
||||
&self,
|
||||
entity_id: &str,
|
||||
depth: usize,
|
||||
direction: RelationDirection,
|
||||
relation_types: Option<&[&str]>,
|
||||
) -> Result<Vec<ScoredEntity>, MemoryError> {
|
||||
// 1. 验证起点存在
|
||||
let inner = self.inner.lock().unwrap();
|
||||
if !inner.entities.contains_key(entity_id) {
|
||||
return Err(MemoryError::NotFound(entity_id.to_string()));
|
||||
}
|
||||
|
||||
// 2. BFS 初始化
|
||||
let mut visited: HashSet<String> = HashSet::new();
|
||||
let mut result: Vec<ScoredEntity> = Vec::new();
|
||||
// 队列:(entity_id, score, path)
|
||||
let mut queue: VecDeque<(String, f32, Vec<String>)> = VecDeque::new();
|
||||
|
||||
queue.push_back((entity_id.to_string(), 1.0, vec![entity_id.to_string()]));
|
||||
visited.insert(entity_id.to_string());
|
||||
|
||||
// 3. BFS 逐层遍历
|
||||
for _ in 0..depth {
|
||||
let mut next_queue: VecDeque<(String, f32, Vec<String>)> = VecDeque::new();
|
||||
|
||||
while let Some((current_id, score, path)) = queue.pop_front() {
|
||||
// 筛选与 current_id 相关的关系
|
||||
for rel in inner.relations.iter() {
|
||||
// 方向过滤
|
||||
let (match_source, match_target) = match direction {
|
||||
RelationDirection::Outgoing => (&rel.source_id, &rel.target_id),
|
||||
RelationDirection::Incoming => (&rel.target_id, &rel.source_id),
|
||||
RelationDirection::Both => {
|
||||
if rel.source_id == current_id {
|
||||
(&rel.source_id, &rel.target_id)
|
||||
} else if rel.target_id == current_id {
|
||||
(&rel.target_id, &rel.source_id)
|
||||
} else {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
if *match_source != current_id {
|
||||
continue;
|
||||
}
|
||||
|
||||
// 关系类型过滤
|
||||
if let Some(types) = relation_types {
|
||||
if !types.contains(&rel.relation_type.as_str()) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
let neighbor_id = match_target.clone();
|
||||
if visited.contains(&neighbor_id) {
|
||||
continue;
|
||||
}
|
||||
visited.insert(neighbor_id.clone());
|
||||
|
||||
// 权重乘积衰减
|
||||
let new_score = score * rel.weight;
|
||||
let mut new_path = path.clone();
|
||||
new_path.push(neighbor_id.clone());
|
||||
|
||||
result.push(ScoredEntity {
|
||||
entity: inner.entities.get(&neighbor_id).cloned()
|
||||
.ok_or_else(|| MemoryError::NotFound(neighbor_id.clone()))?,
|
||||
score: new_score,
|
||||
path: new_path.clone(),
|
||||
});
|
||||
|
||||
next_queue.push_back((neighbor_id, new_score, new_path));
|
||||
}
|
||||
}
|
||||
|
||||
queue = next_queue;
|
||||
}
|
||||
|
||||
// 4. 按分数降序排列
|
||||
result.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
|
||||
Ok(result)
|
||||
}
|
||||
```
|
||||
|
||||
`depth=0` 时仅验证起点实体存在,返回空关联列表(不遍历任何边)。
|
||||
|
||||
**BFS 关键设计点**:
|
||||
|
||||
| 特性 | 处理方式 |
|
||||
|------|----------|
|
||||
| 环路 | `visited: HashSet<String>` 已访问集合防环 |
|
||||
| 评分衰减 | 沿路径 `score *= rel.weight`,权重乘积 |
|
||||
| 多路径 | BFS 天然先到先得,同一实体只保留首次到达路径 |
|
||||
| 关系类型过滤 | `relation_types: Option<&[&str]>`,`None` 表示不过滤 |
|
||||
| 方向过滤 | `RelationDirection` 枚举,`Both` 时双向检查 |
|
||||
|
||||
> 以上性能数据为基于算法复杂度的估算值(O(R) 线性扫描,R=关系数),实际性能需通过基准测试验证。建议在实现后添加 `#[bench]` 或 criterion 基准测试。
|
||||
|
||||
### 3.4 检索扩展
|
||||
|
||||
#### RetrievalStrategy
|
||||
|
||||
```rust
|
||||
/// 检索策略 —— 控制双通道分流。
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub enum RetrievalStrategy {
|
||||
/// 并行 KnowledgeStore + KnowledgeGraph,合并排序(默认)。
|
||||
#[default]
|
||||
Hybrid,
|
||||
/// 仅 KnowledgeStore。
|
||||
KnowledgeOnly,
|
||||
/// 仅 KnowledgeGraph。
|
||||
GraphOnly,
|
||||
}
|
||||
```
|
||||
|
||||
#### RetrievalItem
|
||||
|
||||
```rust
|
||||
/// 统一检索条目 —— enum 变体区分类别,两通道分数均在 [0,1] 区间。
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum RetrievalItem {
|
||||
/// 知识页面(来自 KnowledgeStore)。
|
||||
KnowledgePage {
|
||||
page: KnowledgePage,
|
||||
/// TextOverlap 评分 [0.0, 1.0]。
|
||||
score: f32,
|
||||
},
|
||||
/// 图谱实体(来自 KnowledgeGraph)。
|
||||
GraphEntity {
|
||||
entity: crate::memory::graph::GraphEntity,
|
||||
/// 图距离评分 [0.0, 1.0]。
|
||||
score: f32,
|
||||
/// 从查询实体到当前实体的 ID 路径。
|
||||
path: Vec<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl RetrievalItem {
|
||||
/// 统一分数(用于合并排序)。
|
||||
pub fn score(&self) -> f32 {
|
||||
match self {
|
||||
Self::KnowledgePage { score, .. } => *score,
|
||||
Self::GraphEntity { score, .. } => *score,
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
> **注意**:两个通道的分数维度不同(TextOverlap vs 图距离),合并排序仅用于统一返回,不代表跨通道可比性。
|
||||
|
||||
#### 向后兼容导出(ScoredItem)
|
||||
|
||||
```rust
|
||||
// ── 向后兼容导出 ──
|
||||
|
||||
/// 旧版带评分的知识页面检索结果(已废弃)。
|
||||
///
|
||||
/// 请迁移到 `RetrievalItem::KnowledgePage { page, score }`。
|
||||
#[deprecated(since = "0.3.0", note = "使用 RetrievalItem::KnowledgePage 代替")]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ScoredItem {
|
||||
pub page: KnowledgePage,
|
||||
pub score: f32,
|
||||
}
|
||||
|
||||
// 在 memory.rs 模块根的重导出中保留:
|
||||
// #[allow(deprecated)]
|
||||
// pub use retriever::ScoredItem;
|
||||
```
|
||||
|
||||
#### RetrievalResult 更新
|
||||
|
||||
```rust
|
||||
/// 检索结果。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RetrievalResult {
|
||||
/// 统一条目列表,按分数降序排列。
|
||||
pub items: Vec<RetrievalItem>,
|
||||
pub query: String,
|
||||
/// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。
|
||||
pub strategy: RetrievalStrategy,
|
||||
}
|
||||
```
|
||||
|
||||
#### MemoryRetriever 扩展
|
||||
|
||||
```rust
|
||||
pub struct MemoryRetriever {
|
||||
knowledge_store: KnowledgeStore,
|
||||
/// 可选知识图谱(None 时退化为单通道)。
|
||||
knowledge_graph: Option<Arc<dyn KnowledgeGraph>>,
|
||||
/// 检索策略(默认 Hybrid)。
|
||||
strategy: RetrievalStrategy,
|
||||
config: RetrieverConfig,
|
||||
stop_words: HashSet<String>,
|
||||
}
|
||||
|
||||
impl MemoryRetriever {
|
||||
/// 创建新的 MemoryRetriever(保持向后兼容)。
|
||||
pub fn new(knowledge_store: KnowledgeStore, config: RetrieverConfig) -> Self {
|
||||
Self {
|
||||
knowledge_store,
|
||||
knowledge_graph: None,
|
||||
strategy: RetrievalStrategy::default(),
|
||||
config,
|
||||
stop_words: default_stop_words(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 注入知识图谱,启用双通道检索。
|
||||
pub fn with_knowledge_graph(mut self, graph: Arc<dyn KnowledgeGraph>) -> Self {
|
||||
self.knowledge_graph = Some(graph);
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置检索策略。
|
||||
pub fn with_strategy(mut self, strategy: RetrievalStrategy) -> Self {
|
||||
self.strategy = strategy;
|
||||
self
|
||||
}
|
||||
|
||||
/// 检索相关记忆(双通道)。
|
||||
pub async fn retrieve(&self, query: &str) -> Result<RetrievalResult, MemoryError> {
|
||||
if query.is_empty() {
|
||||
return Ok(RetrievalResult {
|
||||
items: Vec::new(),
|
||||
query: query.to_string(),
|
||||
strategy: self.strategy.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
let keywords = extract_keywords(query, &self.stop_words);
|
||||
let has_graph = self.knowledge_graph.is_some();
|
||||
|
||||
// 按策略分流
|
||||
match (&self.strategy, has_graph) {
|
||||
// 仅知识页面
|
||||
(RetrievalStrategy::KnowledgeOnly, _) | (_, false) => {
|
||||
let items = self.search_knowledge_store(query, &keywords).await?;
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
strategy: RetrievalStrategy::KnowledgeOnly,
|
||||
})
|
||||
}
|
||||
// 仅图谱
|
||||
(RetrievalStrategy::GraphOnly, true) => {
|
||||
let graph = self.knowledge_graph.as_ref().unwrap();
|
||||
let items = self.search_graph(query, &keywords, graph).await?;
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
strategy: self.strategy.clone(),
|
||||
})
|
||||
}
|
||||
// 混合:并行执行,合并排序
|
||||
(RetrievalStrategy::Hybrid, true) => {
|
||||
let graph = self.knowledge_graph.as_ref().unwrap();
|
||||
let (kp_items, g_items) = tokio::join!(
|
||||
self.search_knowledge_store(query, &keywords),
|
||||
self.search_graph(query, &keywords, graph),
|
||||
);
|
||||
|
||||
let mut items = kp_items?;
|
||||
items.extend(g_items?);
|
||||
items.sort_by(|a, b| {
|
||||
b.score().partial_cmp(&a.score())
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
items.truncate(self.config.max_results);
|
||||
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
strategy: self.strategy.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### 3.5 标签管理
|
||||
|
||||
#### 标签索引维护
|
||||
|
||||
`tag_index: HashMap<String, HashSet<String>>` 维护 tag → entity_ids 反向映射:
|
||||
|
||||
- **`set_entity_tags`**:先清除旧标签的反向引用,再写入新标签。超出 `max_tags_per_entity` 时截断。
|
||||
- **`find_tags`**:遍历 `tag_index.keys()`,按前缀过滤。
|
||||
- **`entity_count_by_tag`**:直接返回 `tag_index.get(tag).map_or(0, |s| s.len())`。
|
||||
|
||||
#### 实现要点
|
||||
|
||||
```rust
|
||||
async fn set_entity_tags(
|
||||
&self,
|
||||
entity_id: &str,
|
||||
tags: Vec<String>,
|
||||
) -> Result<usize, MemoryError> {
|
||||
let mut inner = self.inner.lock().unwrap();
|
||||
let entity = inner.entities.get_mut(entity_id)
|
||||
.ok_or_else(|| MemoryError::NotFound(entity_id.to_string()))?;
|
||||
|
||||
// 清除旧标签的反向引用
|
||||
for old_tag in &entity.tags {
|
||||
if let Some(ids) = inner.tag_index.get_mut(old_tag) {
|
||||
ids.remove(entity_id);
|
||||
if ids.is_empty() {
|
||||
inner.tag_index.remove(old_tag);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 截断到 max_tags_per_entity
|
||||
let max = self.constraints.max_tags_per_entity;
|
||||
let new_tags: Vec<String> = tags.into_iter().take(max).collect();
|
||||
|
||||
// 写入新标签的反向引用
|
||||
for tag in &new_tags {
|
||||
inner.tag_index.entry(tag.clone())
|
||||
.or_default()
|
||||
.insert(entity_id.to_string());
|
||||
}
|
||||
|
||||
entity.tags = new_tags.clone();
|
||||
Ok(new_tags.len())
|
||||
}
|
||||
```
|
||||
|
||||
#### 标签复用流程(文档说明)
|
||||
|
||||
```
|
||||
LLM 提取候选标签 → 对每个候选:
|
||||
graph.find_tags(candidate.lowercase())
|
||||
├─ 命中已有标签 → 复用
|
||||
└─ 无匹配 → 注册新标签
|
||||
```
|
||||
|
||||
> **标注**:当前无自动提取流程,需 Agent 层显式调用 `set_entity_tags`。
|
||||
|
||||
---
|
||||
|
||||
## 实现计划
|
||||
|
||||
| Step | 内容 | 文件范围 | 验证标准 | 预估行数 |
|
||||
|------|------|----------|----------|----------|
|
||||
| 1 | `graph.rs` 核心类型:GraphEntity / GraphRelation / RelationDirection / ScoredEntity / TagConstraints | `src/memory/graph.rs` | `cargo check` 编译通过 | ~80 |
|
||||
| 2 | `KnowledgeGraph` trait 定义(10 个 async 方法) | `src/memory/graph.rs` | trait 编译通过,无未实现方法 | ~90 |
|
||||
| 3 | `InMemoryGraph` 实现 + BFS 遍历 | `src/memory/graph.rs` | 单元测试:添加实体/关系、BFS 遍历、方向过滤、类型过滤 | ~250 |
|
||||
| 4 | 标签管理实现(set_entity_tags / find_tags / entity_count_by_tag) | `src/memory/graph.rs` | 单元测试:标签增删查、截断、反向索引维护 | ~80 |
|
||||
| 5 | `retriever.rs` 扩展:RetrievalItem / RetrievalStrategy / MemoryRetriever 改造 | `src/memory/retriever.rs` | `cargo check` + 双通道检索测试 | ~180 |
|
||||
| 6 | `memory.rs` 模块根更新 + 重导出 | `src/memory.rs` | `cargo check`,pub use 无编译错误 | ~10 |
|
||||
| 7 | 内联测试 | `src/memory/graph.rs` + `src/memory/retriever.rs` | `cargo test --all-targets` 全绿 | ~150 |
|
||||
|
||||
**总预估**:约 840 行(核心 640 + 测试 200)
|
||||
|
||||
---
|
||||
|
||||
## 风险评估
|
||||
|
||||
| 风险 | 影响 | 缓解措施 |
|
||||
|------|------|----------|
|
||||
| 标签 API 无消费者 | 低 — 预留接口,不影响核心功能 | 文档标注"Agent 层显式写入",后续 Phase 接入 |
|
||||
| 评分不可比 | 中 — TextOverlap vs 图距离维度不同 | `RetrievalItem` enum 变体分离,合并排序仅统一返回,文档注明维度差异 |
|
||||
| BFS 性能 | 低 — 5000 关系遍历 ≈ 1-50µs | 不引入邻接表索引,等实测超过 1ms 再优化 |
|
||||
| Breaking change | 中 — `RetrievalResult.items` 类型变化 | CHANGELOG 标注,`ScoredItem` 保留为 `pub` 兼容导出(deprecate) |
|
||||
| Mutex 竞争 | 低 — InMemoryGraph 单实例场景 | 读多写少, Mutex 性能足够;后续可升级 RwLock |
|
||||
|
||||
---
|
||||
|
||||
## 验收标准
|
||||
|
||||
| 检查项 | 标准 | 验证命令 |
|
||||
|--------|------|----------|
|
||||
| 编译 | 0 error | `cargo check --all-targets` |
|
||||
| 测试 | 全绿,测试数从 ~391 增至 ~410+ | `cargo test --all-targets` |
|
||||
| Clippy | 0 warning | `cargo clippy --all-targets -- -D warnings` |
|
||||
| 文档 | 0 warning | `cargo doc --no-deps` |
|
||||
| BFS 覆盖 | 所有边界条件:空图、单实体、环路、深度 0、方向过滤、类型过滤 | 内联测试 |
|
||||
| 双通道 | Hybrid / KnowledgeOnly / GraphOnly 三种策略功能正确 | 内联测试 |
|
||||
| 向后兼容 | `MemoryRetriever::new()` 签名不变,现有调用无需修改 | `cargo check` 无 breaking error |
|
||||
@@ -0,0 +1,109 @@
|
||||
# AG Core Roadmap — Unsorted
|
||||
|
||||
> 本文件存放**尚未归到任何具体版本**的 roadmap 内容:跨版本的全局视图、面向未来的展望、风险与建议、阶段总回顾。
|
||||
>
|
||||
> **已分版本的内容**:请查阅
|
||||
> - [`roadmap-v0.1.0.md`](./roadmap-v0.1.0.md) — Phase 0–4c + v0.1.0 Release
|
||||
> - [`roadmap-v0.2.0.md`](./roadmap-v0.2.0.md) — Phase 5–12 + v0.2.0-rc.1
|
||||
> - [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) — Phase 13–19(13-18 已完成,19 待实施)
|
||||
>
|
||||
> 返回总入口:[`roadmap.md`](./roadmap.md)
|
||||
|
||||
---
|
||||
|
||||
## 全局愿景
|
||||
|
||||
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
|
||||
|
||||
**当前状态**:v0.2.0-rc.1 已打标签。Phase 0-18 全部完成。v0.3.0 实施中,Phase 19 共 1 个增量 Phase 待交付。目标是从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"。
|
||||
|
||||
---
|
||||
|
||||
## 模块完整性评估
|
||||
|
||||
| 功能领域 | 方案状态 | 文档位置 | 实现优先级 |
|
||||
|---------|---------|---------|-----------|
|
||||
| LLM 调用周期 | ✅ 完整 | `specs/llm-call-lifecycle.md` | P0 |
|
||||
| 提示词工程 | ✅ 完整 | `docs/4-prompt-engineering.md` | P1 |
|
||||
| 工具系统 + 权限 | ✅ 完整 | `docs/5-tool-system.md` | P1 |
|
||||
| 记忆检索 | ✅ 完整 | `docs/6-memory-system.md` | P2 |
|
||||
| Agent 运行时(4a 胶水层) | ✅ 已实现 | `docs/7-agent-runtime.md` | P2 |
|
||||
| 生命周期钩子 | ✅ 完整 | `docs/3-phase0-remaining.md` | P0(LLM Cycle 扩展) |
|
||||
| Provider 注册发现 | ✅ 完整 | `docs/3-phase0-remaining.md` | P0(Provider 接口扩展) |
|
||||
| 流式事件系统 | ✅ 完整 | `docs/3-phase0-remaining.md` | P0(流式接口前置) |
|
||||
|
||||
|
||||
---
|
||||
|
||||
## 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. **持久化依赖**:`rusqlite` + `bundled` 零外部依赖编译,但 SQLite 不适配所有场景(分布式/高并发写)。`MemoryStore` trait 的抽象层允许下游自行实现 Redis / PostgreSQL 后端
|
||||
2. **ContextSlot 心智负担**:`ContextSlot` 引入了一等抽象的复杂度。建议通过 `AgentBuilder` 默认创建 `"default"` slot,让简单场景无感使用
|
||||
3. **向量检索规模上限**:v0.3 的 `PersistentVectorStore` 全量加载到内存做余弦搜索,适合 ≤10 万条向量。超出此规模需换用专用向量库。v0.4 可以评估引入
|
||||
4. **Scope 蔓延**:v0.3 新增 `agent/summary` `document/` `engine/` `memory/vector_store` 模块,功能覆盖扩展到多 Agent 基础系统。始终保持 trait + reference impl 的边界,业务循环留给上层(Phase 16 已交付 `agent/summary` 摘要生产端 + `format_messages_as_text` 简洁版格式化 + 30K 字符整体截断保留最新;Phase 18 已交付 `engine/switch_agent` 热切换 + `engine/sub_agent` 调度全栈(dispatch / dispatch_all / dispatch_stream);实施后两轮审查 PASS,0 🔴 阻塞)
|
||||
5. **API 稳定性**:v0.3 引入 `Checkpointer`、`SessionManager`、`VectorStore` 等新公开 API,v0.2 已有的 `#[non_exhaustive]` 和 `#[deprecated]` 机制继续沿用
|
||||
6. **Checkpointer 存储效率**:v0.3 使用全量 JSON 序列化存储 checkpoint,每轮对话约几百 KB。`fork` 从历史 checkpoint 创建新 session 时也会复制全量。等实际使用中发现存储瓶颈时再改为增量模式
|
||||
|
||||
---
|
||||
|
||||
## 下一步行动
|
||||
|
||||
1. **v0.3.0 Phase 19 启动**:KnowledgeGraph + 双通道检索,落地 `docs/note-knowledge-graph-design.md` 中记录的知识图谱设计
|
||||
2. **Phase 19 收尾**:完成 v0.3.0 最后一个 Phase 后准备 rc.1 标签 + CHANGELOG
|
||||
3. **示例先行**:完成 Phase 19 后立即创建对应的 knowledge_graph_demo 示例,确保 `cargo run --example` 可验证
|
||||
4. **里程碑追踪**:以 M13(Phase 17)+ M14(Phase 18)为已达成里程碑,逐 Phase 推进 M15
|
||||
|
||||
---
|
||||
|
||||
**已完成 / 进行中阶段**:
|
||||
- ✅ Phase 0 Foundation — 全部交付物已完成
|
||||
- ✅ Phase 1 Prompt Engineering — 全部交付物已完成
|
||||
- ✅ Phase 2 Tool System — 全部交付物已完成
|
||||
- ✅ Phase 3 Memory System — 全部交付物已完成
|
||||
- ✅ Phase 4a Core Glue — 全部交付物已完成
|
||||
- ✅ Phase 4b Task Execution — 全部交付物已完成
|
||||
- ✅ Phase 4c Session Memory — 全部交付物已完成
|
||||
- ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy)
|
||||
- ✅ Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
|
||||
- ✅ Phase 7 SqliteStore — `rusqlite 0.32` + WAL 模式 + `Arc<Mutex<Connection>>` + `spawn_blocking`;`memory/store.rs` → `store/{in_memory,sqlite_store}.rs` 模块化;9 个内联测试覆盖 CRUD/upsert/过滤/10×10 并发/持久化 round-trip;`InMemoryStore ↔ SqliteStore` trait-box 互换兼容
|
||||
- ✅ **Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 60 行 + `end_to_end` 246 行),10 个离线示例全部 exit 0;**v0.2.0-rc.1 标签已打**;实施后三方审查发现 6 项问题(1 🔴 + 2 🟡 + 3 💭)已全部修复
|
||||
- ✅ **Phase 9 流式体验增强** — `AgentSession::submit_turn_stream` 流式事件序列 + `LlmCycle::submit_with_tools_stream` spawn + mpsc 状态机 + `StreamEvent::ToolExecutionStarted`/`Completed` 新变体 + 9 单元测试 + 2 集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock 验证 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211;`CycleConfig` 加 `Clone` derive;方案文档 `docs/16-phase9-streaming-experience.md`(821 行)
|
||||
- ✅ **Phase 10 ContextSlot 上下文管理** — `src/agent/context.rs` 新增 `ContextSlot` 核心类型(Full / Focused / Readonly 三种模式,New / Derived / Static 三种来源)+ JSON blob 批次持久化(每 slot 3-4 条 MemoryItem,`slot_config` key 自恢复支持旧版本兼容);`AgentSession` 扩展 slots 字段 + 5 个管理方法(`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot`,自动创建 `"default"` slot,`delete_slot` 双重保护禁止删 default/最后一个);`submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回(`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式"读时过滤"语义不丢失数据);`finalize_turn` 签名变更(新增 `new_messages_from_cycle: Vec<Message>` 参数,返回 `Result<(), AgentError>`);`agent/error.rs` 新增 3 个 Slot 错误变体(`SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`);`examples/context_slot_demo.rs` 新增分支对话示例(法律咨询入口 → 两个派生方向 → 切换 → 隔离验证 → 删除保护);方案文档 `docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复);全量 211 → 254(+43 新测试),clippy 0 警告,doc 0 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]` 预留扩展);`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 里程碑达成
|
||||
- ✅ **v0.3.0 Phase 15 完成** — `VectorStore` trait(`add`/`search`/`remove`/`add_one`,返回 `(Document, f32)` 消除调用方 id→Document 维护开销)+ `InMemoryVectorStore`(`Mutex<HashMap>` + 余弦全量扫描 + 预计算 L2 norm 缓存)+ `PersistentVectorStore`(构造时全量加载,先写持久化后写内存,持久化失败时内存不污染重启自动恢复,`remove` 幽灵数据窗口已知)+ `RagPipeline` 组合器(ingest: split→embed→store.add / retrieve: embed→store.search,`splitter: Option<RecursiveCharacterSplitter>` 灵活切换);`src/memory/vector_store.rs`(937 行,19 个内联测试覆盖 14 场景含 2 个性能基准)+ 零新外部依赖(纯 Rust `dot()` 余弦);旧 `VectorRetriever`/`InMemoryVectorRetriever` 标注 `#[deprecated(since = "0.3.0")]` 迁移路径清晰;`search_orthogonal_vectors` 返回 1 条 score≈0(文档已同步修正不过滤低分向量);方案文档 `docs/21-phase15-vector-store-persistence.md`(1570 行,经 3 轮审查 + 文档-代码一致化修复);全量 313 → 335(+22 新测试),clippy 0 警告,doc 0 warning;M11 里程碑达成
|
||||
- ✅ **v0.3.0 Phase 16 完成** — `SummaryConfig` 配置结构体(6 个字段:`trigger_token_ratio=0.75` / `max_context_tokens=32_000` / `summary_prompt` / `debounce_turns=3` / `summary_model=None` / `max_tool_result_chars=500`,默认 `None` 沿用主模型避断裂非 OpenAI 用户)+ `AgentBuilder::summary_config(cfg)` 链式方法 + `AgentConfig.summary_config: Option<SummaryConfig>` 字段;`AgentSession` 新增 `last_summary_turn: Option<u32>` 字段(首次不受防抖约束,`should_summarize` 用 `Option` 哨兵实现)+ `maybe_summarize(current_turn)` 内联检查点(OnTurnEnd 之后 / `turn_index` 之前,对称 `submit_turn` / `finalize_turn` 两个入口,流式路径 `saturating_sub(1)` 修正)+ 关联函数 `generate_summary`(构造独立 `LlmCycle` 调 `submit_messages` 传 `vec![Message::user_text(prompt)]`,`max_tokens=1024`,空消息守卫直接返回空串)+ 公开 API `get_conversation_summary()`;`src/agent/summary.rs`(~240 行,含 8 个 SummaryConfig/`format_messages_as_text` 内联测试——默认值/空输入/系统用户助理/ToolResult(含 `tool_call_id`)/工具调用/Unicode 安全截断/整体 30K 截断保留最新;有效字符数截断多字节安全,droptest 验证保留尾部消息)+ `src/agent/session.rs` 注入 10 个摘要集成测试(默认值不触发 / 超阈值触发 / 防抖阻止重复 / SessionMemory 写入 / Full 模式不注入 / 失败不阻断主流程 / 流式路径触发 / 默认配置零影响 / **Focused `summary_override` 写入正向验证** / **空消息不调用 LLM** / **巨型 `max_context_tokens` 永不触发**);`format_messages_as_text` 简洁版消息格式化(`[Tool: name]` + `Tool Result [id]:` + ToolResult 字符级 `chars().take(max_tool_result_chars)` 截断 + 整段 30K 总长度截断从头部保留最新);所有错误静默(失败用 `tracing::error!`,成功用 `tracing::info!(turn, summary_len)`);`MergeStrategy` 注释中过时 "Summarize 指向"与 `context.rs:78` "v0.3 将支持 Hook 驱动" 过时注释在实施时同步移除/更新;方案文档 `docs/22-phase16-summary-auto-generation.md`(471 行),实施后**两轮审查 PASS**:第一轮 PM/SA 审查 11 项问题修复 + 第二轮实施审查 9 项问题修复(🔴 `generate_summary` 空消息 bug + 🟡 W4 流式路径防抖 + 🟡 W2 模型硬编码 + 🟡 W5 Full 模式无谓 save + 🟡 W3 30K 截断 + 🟡 W6 成功无日志 + 🟡 W1/W7 测试补全 + 💭 注释同步);零新外部依赖;全量 335 → **353**(+18 新测试,含二次审查增补 4 个),clippy 0 警告,doc 0 warning,`quick_start` 示例正常 exit 0;**M12 里程碑达成** + 第二轮审查门禁 PASS
|
||||
- ✅ **v0.3.0 Phase 17 完成** — 新建 `src/engine/` 模块(5 文件:`mod.rs`/`error.rs`/`snapshot.rs`/`checkpointer.rs`/`session_manager.rs`),实现 **SessionManager**(10 个公开方法:`create`/`create_child`/`get`/`recover`/`replace`/`children`/`parent`/`destroy`/`submit_turn`/`submit_turn_stream`/`finalize_turn_stream`,内部 `RwLock<HashMap>` + `Arc<tokio::sync::Mutex<AgentSession>>` + `Checkpointer` 组合)和 **Checkpointer**(5 个公开方法:`checkpoint`/`rollback_load`/`list_checkpoints`/`delete_all`/`latest_snapshot`);`SessionSnapshot` 独立 struct 避开 `Arc<dyn Agent>` 不可序列化,配套 `SessionMemoryEntry` 保留 metadata/created_at;`AgentSession` 扩展三段式快照(`to_snapshot` async 读 MemoryStore + `from_snapshot` 纯同步构造 + `restore_memory` &mut self async 写回持久层);`SessionMemory` 新增 `list_entries()` 和 `set_with_meta()` 方法(恢复时保留完整 entry 数据);存储 key 风格统一为 `session:{id}:meta` / `ckpt:{id}:{ckpt_id}`(与 `slot_data:` 风格一致);`EngineError` 6 个变体(含 `Memory(#[from] MemoryError)` 透传 + `Agent(#[from] AgentError)`);`CkptMeta` 加 `created_at_nanos` 字段确保同秒内精确降序排序;ckpt_id 用纳秒+单调计数器生成(零外部依赖,ponytail);session_id 用纳秒+计数器自动生成(统一策略,UUID v4 备选);自动 checkpoint 失败 `tracing::error!` 不阻断主流程(不提供强持久化保证);流式 checkpoint 仅在 `finalize_turn_stream` 创建(不留半成品污染);孤儿策略:`destroy()` 不递归删除子 session,父被销毁后 `parent()` 返回 `Ok(None)`;3 处 derive 改动(`CostTracker` + `ContextSlot` + `MergeStrategy` 加 serde,`CostTracker` 额外加 `Clone`);`SessionManager::recover` + `replace` 内部自动 `restore_memory` 写回持久层;零新外部依赖;方案文档 `docs/23-phase17-agent-execution-engine.md`(775 行,经两轮 PM+SA 审查 + 实施后第三轮 PM+SA+Code Reviewer 三方联合审查),实施后**两轮审查门禁 PASS**:第一轮修复 6 🔴 + 第二轮修复 2 🔴(to_snapshot 同步→async + Roadmap 同步)+ 实施后修复 8 个 🟡(restore_memory metadata/created_at 完整恢复 + &mut self 签名 + 死代码清理 + 3 个边界测试 + tracing 补全 + 文档语义统一 + 示例 rollback 一致性 assert);15 个 `SessionManager` 内联测试(CRUD/recover/replace/树形/孤儿/auto_checkpoint on-off)+ 6 个 `Checkpointer` 内联测试(roundtrip/不存在的 ckpt/同秒降序/delete_all 幂等/latest/隔离)+ 1 个 `snapshot_deserialize_with_minimal_fields` 序列化兼容测试;全量 353 → **374**(+21 新测试),clippy 0 警告,doc 0 warning,`engine_demo` 示例端到端演示 create→submit_turn→checkpoint→rollback→replace→destroy 全链路并验证 rollback 一致性;**M13 里程碑达成** + 两轮审查门禁 PASS
|
||||
- ✅ **v0.3.0 Phase 18 完成** — 新增 `src/engine/switch.rs`(222 行)实现 `SessionManager::switch_agent()` 热切换(替换 `Arc<dyn Agent>`,slot 历史 / `turn_index` / `session_memory` / `cost_so_far` 全部保留,同步更新 `SessionMeta.agent_name` 到持久层,`created_at` / `parent_id` 保持原始不可变)+ 新增 `src/engine/sub_agent.rs`(1071 行)实现 4 个公开方法(`dispatch` / `dispatch_all` / `dispatch_stream` 与前述 `switch_agent` 共 4 个 Phase 18 核心 API)+ 3 个公开类型(`DispatchConfig` / `SubTaskResult` / `SubTaskStreamEvent`);`DispatchConfig` 4 字段(`max_concurrency=10` / `inherit_session_memory=true` / `bridge_keys=None` / `shared_namespace=None`)+ 三态 `bridge_keys` 语义(`None` = 不继承 / `Some(vec![])` = 全部 / `Some(keys)` = 指定 keys)+ 约定式 `shared_namespace` 子↔子共享(`shared:{prefix}:{key}`)不触发自动注入;`dispatch` 流程:`create_child` → `inherit_session_memory`(快照语义)→ `submit_turn` → 返回 `SubTaskResult`;`dispatch_all` `tokio::sync::Semaphore` 并发控制 + `Vec<Result<...>>` 部分成功语义按输入顺序 indexed 收集;`dispatch_stream` `unbounded_channel` + spawn task 消息重建 + `finalize_turn` 后台落库(明确不参与 `auto_checkpoint` 防重复);`SubTaskStreamEvent` 事件序列:`ChildCreated` → `Stream(StreamEvent) × N` → `Completed(SubTaskResult)` 或 `Error { child_id, error }`;`EngineError` 新增 `DispatchFailed(#[source] String)` 变体 + `CostTracker` 加 `From<Usage>` 转换;`save_session_meta` / `load_session_meta` 改 `pub(crate)` 供 `switch.rs` 调用;4 个端到端示例:`agent_switch_demo`(115 行)+ `sub_agent_dispatch_demo`(141 行)+ `bridge_keys_demo`(197 行)+ `dispatch_stream_demo`(121 行)全部 exit 0;17 个内联测试(4 switch + 5 dispatch + 4 dispatch_all + 4 dispatch_stream);零新外部依赖;方案文档 `docs/24-phase18-agent-switch-and-dispatch.md`(700 行);全量 374 → **391**(+17 新测试,0 失败),clippy 0 警告,doc 0 warning;**M14 里程碑达成**
|
||||
@@ -0,0 +1,242 @@
|
||||
# AG Core Roadmap — v0.1.0
|
||||
|
||||
> 本文件聚焦 **v0.1.0 版本** 的规划与交付(Phase 0–4c),已于 2026-07-04 完成发布。
|
||||
> 返回总入口:[`roadmap.md`](./roadmap.md)
|
||||
|
||||
## v0.1.0 愿景
|
||||
|
||||
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
|
||||
|
||||
## v0.1.0 总体范围
|
||||
|
||||
**总体规模**:5 个主体 Phase(Phase 0–4c)+ Provider IR 重构 + LlmCycle 简化 + v0.1 Release 收尾,182 个测试全绿,clippy 0 警告,7 个离线示例全 exit 0。
|
||||
|
||||
---
|
||||
|
||||
### Phase 0 — Foundation(基础设施)
|
||||
|
||||
**目标**:实现 LLM 调用周期的核心功能,作为所有上层模块的基础。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `llm/types.rs` — 核心数据类型(Message, ContentBlock, ChatRequest/Response, ToolDefinition, StopReason)
|
||||
2. ✅ `llm/error.rs` — 错误体系(LlmError 枚举,可重试/不可重试判断)
|
||||
3. ✅ `llm/provider.rs` + `llm/provider/openai.rs` — Provider 接口 + OpenAI 兼容实现
|
||||
4. ✅ `llm/provider/registry.rs` — ProviderRegistry(多 Provider 注册发现)
|
||||
5. ✅ `llm/cycle.rs` + `llm/cycle/{retry,usage}.rs` — 生命周期引擎(重试策略 + 用量追踪)
|
||||
6. ✅ `llm/hooks.rs` — HookExecutor 接口(生命周期钩子)
|
||||
7. ✅ `llm/stream.rs` — StreamEvents 流式事件系统(AssistantTextDelta, ToolExecutionStarted 等)
|
||||
8. ✅ `llm/compact.rs` — Auto-compaction(上下文自动压缩)
|
||||
9. ✅ `Cargo.toml` — 添加依赖(tokio, reqwest, serde, thiserror, async-trait, tracing)
|
||||
|
||||
**依赖**:无
|
||||
|
||||
**优先级**:Must Have
|
||||
|
||||
**预估规模**:约 1000 行核心代码
|
||||
|
||||
**状态**:✅ Phase 0 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 1 — Prompt Engineering(提示词工程)
|
||||
|
||||
**目标**:提供提示词的组合、模板化与优化能力。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `prompt.rs` + `prompt/` 模块
|
||||
2. ✅ `PromptTemplate` — 模板引擎(支持变量插值、条件渲染)
|
||||
3. ✅ `PromptComposer` — 提示词组合器(拼接 system/user/assistant 消息)
|
||||
4. ✅ `docs/4-prompt-engineering.md` — 方案文档
|
||||
|
||||
**依赖**:无(可与 Phase 0 并行)
|
||||
|
||||
**优先级**:Should Have
|
||||
|
||||
**预估规模**:约 400 行代码
|
||||
|
||||
**状态**:✅ Phase 1 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 2 — Tool System(工具系统)
|
||||
|
||||
**目标**:实现 MCP 协议集成与自定义工具注册、调用、权限控制。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `tools.rs` + `tools/` 模块(base/registry/permission/mcp/error)
|
||||
2. ✅ `ToolRegistry` — 工具注册表(注册、发现、调用、并行执行、超时控制)
|
||||
3. ✅ `BaseTool` trait — 工具抽象接口(含 ToolContext 执行上下文)
|
||||
4. ✅ `McpClient` — MCP 协议客户端(stdio transport,StreamableHttp 预留)
|
||||
5. ✅ `PermissionChecker` — 工具执行权限检查(白名单/黑名单/自定义权限)
|
||||
6. ✅ `docs/5-tool-system.md` — 方案设计文档
|
||||
7. ✅ 扩展 `llm/cycle.rs` 支持自动 tool 循环(`submit_with_tools()` + `submit_request()` + `maybe_compact()`)
|
||||
8. ✅ `ToolError` — 结构化错误体系(含 `is_recoverable()` 分类)
|
||||
|
||||
**依赖**:Phase 0(LlmProvider 接口传递 tool definitions)、Phase 1(提示词可能需要注入工具描述)
|
||||
|
||||
**优先级**:Should Have
|
||||
|
||||
**预估规模**:约 900 行代码(实际约 1500 行)
|
||||
|
||||
**状态**:✅ Phase 2 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 3 — Memory System(记忆系统)
|
||||
|
||||
**目标**:提供对话记忆的存储、检索与管理能力。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `memory.rs` + `memory/` 模块(store / conversation / knowledge / retriever / error / types)
|
||||
2. ✅ `MemoryStore` trait + `InMemoryStore` — 记忆存储抽象(可插拔后端)+ 默认实现
|
||||
3. ✅ `ConversationMemory` — 对话记忆管理(sliding window / 全量),复用 `llm::compact`
|
||||
4. ✅ `KnowledgeStore` — 知识页面存储(具体 struct,非 trait,基于 MemoryStore)
|
||||
5. ✅ `MemoryRetriever` — 记忆检索器(TextOverlap Dice 系数评分,单通道)
|
||||
6. ✅ `docs/6-memory-system.md` — 方案设计文档
|
||||
7. ✅ `docs/note-knowledge-graph-design.md` — KnowledgeGraph 等 Phase 4 备用设计
|
||||
8. ✅ `EvictionPolicy` — 支持 None / Ttl / Capacity 三种淘汰策略
|
||||
|
||||
**依赖**:Phase 0(llm::compact 复用)、Cargo.toml 新增 `time` 依赖
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 700 行代码(实际约 1242 行,含测试)
|
||||
|
||||
**状态**:✅ Phase 3 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 4a — Agent Core Glue(核心胶水层)
|
||||
|
||||
**目标**:提供最小可用的 Agent Runtime——把 Phase 0-3 的能力"装配"成 `AgentSession::submit_turn`。上层可基于 4a 构建多轮对话应用。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `agent.rs` + `agent/` 模块(7 个文件:agent/error/runtime/builder/session/task + 模块根)
|
||||
2. ✅ `Agent` trait — 智能体角色定义(name / system_prompt / tool_definitions)
|
||||
3. ✅ `AgentSession` — 会话实例(绑定 `Arc<dyn Agent>` + `RuntimeBundle` + 内联 HashMap session_data)
|
||||
4. ✅ `RuntimeBundle` — 显式依赖注入容器(不含 session_memory_backend)
|
||||
5. ✅ `AgentBuilder` — 链式构造入口(不含 session_memory_backend)
|
||||
6. ✅ `AgentError` — 统一错误类型(7 个变体:Llm / Tool / Memory / HookBlocked / LimitExceeded / Config / Other;不含 PlanParse)
|
||||
7. ✅ `Plan` / `Step` / `StepStatus` — 纯数据结构(不含任何解析逻辑)
|
||||
8. ✅ Hook 事件扩展:OnTurnStart / OnTurnEnd + turn_index 字段
|
||||
9. ✅ `docs/7-agent-runtime.md` — 方案设计文档(含 4a/4b/4c 分阶段计划)
|
||||
|
||||
**实际新增**:
|
||||
- 新增文件 7 个(agent.rs + agent/{agent, error, runtime, builder, session, task}.rs)
|
||||
- 修改文件 3 个(lib.rs +1 行;llm/hooks.rs +13 行追加变体/字段;llm/cycle.rs 内部字段 Box→Arc + 新增 `new_with_arc` 公共方法)
|
||||
- 实际代码量约 800 行(含测试;纯实现约 470 行——略高于方案预估 440 行,因 AgentSession 的 tests 模块内联 MockProvider/StubAgent 等辅助结构)
|
||||
- 新增内联测试 22 个;全量测试 84 → 109(0 失败)
|
||||
- clippy 0 警告(agent 模块)
|
||||
- 无新增外部依赖
|
||||
|
||||
**依赖**:Phase 0, 1, 2, 3
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 440 行代码
|
||||
|
||||
**状态**:✅ Phase 4a 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 4b — Task Execution(任务执行)
|
||||
|
||||
**目标**:在 Phase 4a 基础上,赋予智能体"拆解目标 → 逐步执行"的能力。
|
||||
|
||||
**前置条件**:Phase 4a 已完成。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `TaskAgent` trait — `run(goal)` 自主式 + `execute_plan(plan)` 外部驱动式
|
||||
2. ✅ `PlanParser` trait + `JsonPlanParser` 参考实现
|
||||
3. ✅ `AgentError` 追加 PlanParse 变体(共 7 个变体)
|
||||
4. ✅ Hook 事件扩展:OnPlanStepComplete + plan_step_index 字段
|
||||
|
||||
**依赖**:Phase 4a
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 200 行代码(增量)
|
||||
|
||||
**实际新增**:
|
||||
- 修改文件 2 个(llm/hooks.rs +5 行;agent/error.rs +10 行)
|
||||
- 新增代码约 150 行(含测试;纯实现约 90 行)
|
||||
- 新增内联测试 4 个;全量测试 109 → 113(0 失败)
|
||||
- clippy 0 警告
|
||||
- 无新增外部依赖
|
||||
|
||||
**状态**:✅ Phase 4b 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 4c — Session Memory(会话级记忆)
|
||||
|
||||
**目标**:提供会话级 key-value 记忆,作为 session 内各 context 之间的信息桥接通道。
|
||||
|
||||
**前置条件**:Phase 4a 已完成(可与 Phase 4b 并行)。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `SessionMemory` struct — 基于 `MemoryStore`,按 session_id namespace 隔离
|
||||
2. ✅ `RuntimeBundle` + `AgentBuilder` 扩展 `session_memory_backend` 字段
|
||||
3. ✅ `AgentSession` 替换内联 HashMap 为完整 `SessionMemory`
|
||||
|
||||
**依赖**:Phase 4a(Phase 3 MemoryStore)
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 115 行代码(增量)
|
||||
|
||||
**实际新增**:
|
||||
- 新增文件 1 个(agent/session_memory.rs)
|
||||
- 修改文件 4 个(agent/runtime.rs +5 行;agent/builder.rs +10 行;agent/session.rs +30 行;agent.rs +2 行)
|
||||
- 新增代码约 180 行(含测试;纯实现约 100 行)
|
||||
- 新增内联测试 3 个;全量测试 113 → 116(0 失败)
|
||||
- clippy 0 警告
|
||||
- 无新增外部依赖
|
||||
|
||||
**状态**:✅ Phase 4c 全部交付物已完成
|
||||
|
||||
---
|
||||
```mermaid
|
||||
graph BT
|
||||
P0["<b>Phase 0: Foundation</b><br/>LLM Cycle<br/>ProviderRegistry<br/>HookExecutor<br/>StreamEvents<br/>Auto-compaction"]:::done
|
||||
P1["<b>Phase 1: Prompt Engineering</b><br/>PromptTemplate<br/>PromptComposer"]:::done
|
||||
P2["<b>Phase 2: Tool System</b><br/>Tool Registry<br/>PermissionChecker<br/>MCP Client"]:::done
|
||||
P3["<b>Phase 3: Memory System</b><br/>MemoryStore<br/>ConversationMemory<br/>KnowledgeStore"]:::done
|
||||
P4a["<b>Phase 4a: Core Glue</b><br/>AgentSession<br/>RuntimeBundle<br/>Plan/Step 纯数据"]:::done
|
||||
P4b["<b>Phase 4b: Task Execution</b><br/>TaskAgent<br/>PlanParser<br/>JsonPlanParser"]:::done
|
||||
P4c["<b>Phase 4c: Session Memory</b><br/>SessionMemory"]:::done
|
||||
|
||||
P1 --> P0
|
||||
P2 --> P0
|
||||
P3 --> P0
|
||||
P2 --> P1
|
||||
P4a --> P1
|
||||
P4a --> P2
|
||||
P4a --> P3
|
||||
P4b --> P4a
|
||||
P4c --> P4a
|
||||
|
||||
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
|
||||
classDef pending fill:#fbbf24,stroke:#d97706,color:#1a1a1a
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## v0.1 发布里程碑(2026-07-04)
|
||||
|
||||
**质量基线**:
|
||||
|
||||
| 指标 | 数值 |
|
||||
|------|------|
|
||||
| `cargo build --all-targets` | ✅ 通过 |
|
||||
| `cargo test --all-targets` | ✅ **182 passed / 0 failed** |
|
||||
| `cargo clippy --all-targets -- -D warnings` | ✅ 0 警告 |
|
||||
| 离线示例(`cargo run --example`) | ✅ 7 个全部 exit 0 |
|
||||
|
||||
**关键交付**:
|
||||
1. **Provider IR 重构** — 统一 `Message` / `ContentBlock` / `MessageRequest` / `MessageResponse` 类型层;4 个 Provider 适配(OpenAI Chat / Anthropic Messages / DeepSeek / Qwen);`LlmProvider` trait 签名同步切换
|
||||
2. **LlmCycle 简化** — `LlmCycle` 内部消息类型切到 IR 层;移除 Phase 0 的 `OpenaiChatMessage ↔ Message` 桥接;测试从 116 → 182(含 provider 测试)
|
||||
3. **`MockProvider` 公开化** — `agcore::llm::mock::MockProvider` 支持 `chat` + `chat_stream`,无需 API key 即可运行示例
|
||||
4. **7 个离线示例** — `prompt_composer` / `custom_tool` / `agent_session_demo` / `task_agent_demo` / `conversation_memory_demo` / `knowledge_search_demo` / `streaming_events_demo`
|
||||
5. **错误消息友好化** — `AgentError` / `LlmError` / `ToolError` / `MemoryError` / `PromptError` 全部面向最终用户改写(给出可操作的建议)
|
||||
6. **文档完整** — README 完整版(快速上手 + 架构图 + 环境变量)、Apache-2.0 LICENSE
|
||||
@@ -0,0 +1,378 @@
|
||||
# AG Core Roadmap — v0.2.0
|
||||
|
||||
> 本文件聚焦 **v0.2.0 版本** 的规划与交付(Phase 5–12)。已打 `v0.2.0-rc.1` 标签。
|
||||
> 返回总入口:[`roadmap.md`](./roadmap.md)
|
||||
|
||||
## v0.2.0 愿景
|
||||
|
||||
从"LLM 调用工具箱"升级为"生产可用的 Agent 服务"。解决 Rust Agent 工具箱从"能跑"到"能被人依赖"的鸿沟——持久化、配置层、上下文管理三大块补齐后,开发者可在 30 分钟内写出生产可用的 Agent 服务。
|
||||
|
||||
## v0.2.0 总体范围
|
||||
|
||||
**总体规模**:8 个增量 Phase(Phase 5–12),17 个可验证 Step,约 2000+ 行新增代码,测试 182 → 277+。
|
||||
|
||||
---
|
||||
|
||||
## v0.2.0 — 生产就绪(Production-Ready Core)
|
||||
|
||||
**目标**:解决 Rust Agent 工具箱从"能跑"到"能被人依赖"的鸿沟。持久化、配置层、上下文管理三大块补齐后,开发者可在 30 分钟内写出生产可用的 Agent 服务。
|
||||
|
||||
**总体规模**:8 个增量 Phase(Phase 5-12),17 个可验证 Step。
|
||||
|
||||
### 功能清单
|
||||
|
||||
#### P0 — 必须交付
|
||||
|
||||
| # | 功能 | 模块 | 方案要点 |
|
||||
|---|------|------|---------|
|
||||
| 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`) |
|
||||
|
||||
#### P1 — 重要但不阻塞
|
||||
|
||||
| # | 功能 | 模块 | 方案要点 |
|
||||
|---|------|------|---------|
|
||||
| 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 测试 |
|
||||
|
||||
#### P2 — 有时间再做
|
||||
|
||||
| # | 功能 | 模块 | 备注 |
|
||||
|---|------|------|------|
|
||||
| 10 | MCP StreamableHttp | `tools` | 当前仅预留枚举变体 |
|
||||
| 11 | Gemini Provider | `llm/provider` | 协议差异大,实现成本较高 |
|
||||
| 12 | 文件系统 MemoryStore 后端 | `memory` | JSON/JSONL 轻量持久化 |
|
||||
|
||||
### ContextSlot 上下文管理
|
||||
|
||||
**模块归属**:`src/llm/context.rs`(与 `compact.rs` 同级)
|
||||
|
||||
**核心概念**:`ContextSlot` 是一段带策略配置的消息列表,以 `slot_id` 为 namespace 独立持久化到 `MemoryStore`。支持三种模式、三种来源和派生关联(记录 `parent_id`)。
|
||||
|
||||
**核心类型**:
|
||||
|
||||
```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}` → 消息内容
|
||||
- `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` 继续服务传统对话场景。
|
||||
|
||||
**v0.2 不做**:
|
||||
- ❌ `slot.fork()` / `merge()` — 分支方法推迟到 v0.3+
|
||||
- ❌ `inject_summary` 自动生成 — v0.2 仅消费端(从 `SessionMemory` 读取),生成在 v0.3+
|
||||
- ❌ 血缘关系图遍历 — 只存 `parent_id`,不做查询
|
||||
|
||||
**依赖**:Phase 0(MemoryStore trait)、Phase 3(MemoryStore 持久化)
|
||||
**优先级**:P1
|
||||
|
||||
---
|
||||
|
||||
### v0.2.0 实施计划 — 8 个增量 Phase
|
||||
|
||||
> **编号说明**:Phase 5-12 接续 v0.1 的 Phase 0-4c,按开发顺序排列。
|
||||
|
||||
#### Phase 5: 热身准备(Warmup)
|
||||
|
||||
**目标**:快速交付三个互不依赖的独立改动,建立交付节奏。
|
||||
|
||||
| Step | 内容 | 文件范围 | 验证标准 |
|
||||
|------|------|---------|---------|
|
||||
| **5.1** ✅ | `ProviderConfig` 扩展:补 `timeout_secs`(def=30) + `max_retries`(def=3);新增 `ProviderConfig::from_env(prefix)` | `llm/provider.rs` + 各 Provider `new()` 构造函数 | `cargo test` + `from_env()` 单元测试 |
|
||||
| **5.2** ✅ | `OllamaProvider`:基于 `GenericOpenaiProvider` 包装,改 base_url 为 `http://localhost:11434`;`ProviderType` 新增 `Ollama` | `llm/provider/provider.rs` + `llm/provider/ollama.rs`(新增) | `cargo build` — 纯类型级验证 |
|
||||
| **5.3** ✅ | 公开枚举 `#[non_exhaustive]` 前置标记:`ProviderType` / `StopReason` / `FinishReason` / `EvictionPolicy` / `SlotMode`(预置) | 各枚举定义处 | 编译通过 + `cargo clippy` 0 警告 |
|
||||
|
||||
**实际新增**(2026-07-05 commit `98dfe6c`):
|
||||
- 新增文件 1 个(`llm/provider/ollama.rs`,72 行)
|
||||
- 修改文件 2 个(`llm/provider.rs` 加 `from_env` + `Default` + 4 个字段;`memory/store.rs` EvictionPolicy 加 `#[non_exhaustive]`)
|
||||
- `ProviderType::Ollama` 变体 + `FromStr` 解析("ollama" → Ollama)
|
||||
- `OllamaProvider::new(base_url, api_key, model, timeout_secs)` + `with_client()` 构造函数
|
||||
- `ProviderConfig::from_env(prefix)` 解析 `{prefix}_API_KEY` / `{prefix}_BASE_URL` / `{prefix}_MODEL` 环境变量
|
||||
- 全量测试 182 → 190(+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 功能按需交付 | ⏳ |
|
||||
@@ -0,0 +1,367 @@
|
||||
# AG Core Roadmap — v0.3.0
|
||||
|
||||
> 本文件聚焦 **v0.3.0 版本** 的规划与交付(Phase 13–19)。Phase 13-19 全部完成,v0.3.0 交付完毕。
|
||||
> 返回总入口:[`roadmap.md`](./roadmap.md)
|
||||
|
||||
## v0.3.0 愿景
|
||||
|
||||
从"LLM 调用工具箱"升级为"能构建多 Agent 协作、RAG、长记忆 Agent 产品的基础系统"。补齐 LangChain 7 大组件中缺失的 Document 和 VectorStore 能力,落地笔记设计中的 ContextSlot fork/merge、摘要自动生成、知识图谱,建立 engine 引擎层(会话树 + time-travel Checkpointer + SubAgent Dispatch + Agent Switch),为即将开发的多 Agent 产品提供完整基础。
|
||||
|
||||
## v0.3.0 总体范围
|
||||
|
||||
**总体规模**:7 个增量 Phase(Phase 13–19),总新增代码约 2600 行,测试从 277 → 427。7 个 Phase 全部完成(M9-M15 已达成),v0.3.0 交付完毕。
|
||||
|
||||
---
|
||||
|
||||
## 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 → 427。
|
||||
|
||||
### 功能清单
|
||||
|
||||
#### 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/memory/vector_store.rs`(937 行)—— `VectorStore` trait + `InMemoryVectorStore` + `PersistentVectorStore` + `RagPipeline`
|
||||
2. `VectorStore` trait:`add(&[Document], &[Vec<f32>])` 批量 / `search(query, k)` 返回 `(Document, f32)` / `remove(ids)` 幂等
|
||||
3. `PersistentVectorStore`:构造时从 `MemoryStore` 全量加载已有索引;`add` 先写持久化后写内存(持久化失败时内存不污染,重启自动恢复);`search` 纯内存余弦搜索(快照 clone + 锁外计算)
|
||||
4. `RagPipeline`:组合器封装 `split → embed → store.add`(ingest)和 `embed → store.search`(retrieve)两条管线
|
||||
5. 存储格式:`vec:{namespace}:{doc_id}` → JSON `{doc_id, content, metadata, embedding, created_at}`,通过 `MemoryStore` 通用接口读写
|
||||
6. `src/memory/vector.rs` 旧 `VectorRetriever` trait + `InMemoryVectorRetriever` 标注 `#[deprecated(since = "0.3.0")]`,迁移路径指向 `VectorStore` / `InMemoryVectorStore`
|
||||
|
||||
**实际新增**(2026-07-09 commit `32d886f`):
|
||||
- 新增文件 1 个:`src/memory/vector_store.rs`(937 行,含 19 个内联测试)
|
||||
- 修改文件 3 个:`src/memory/vector.rs`(+4 行 deprecated 标注);`src/memory.rs`(+pub mod vector_store + 4 个 pub use re-export);`examples/document_demo.rs`(迁移到 RagPipeline ingest+retrieve)
|
||||
- 零新外部依赖(`Cargo.toml` 未修改)
|
||||
- 全量测试 313 → 335(+22,Phase 15 新增 19 测试 + 部分重叠计数 22 净增);clippy 0 警告,doc 0 warning
|
||||
- 设计文档:`docs/21-phase15-vector-store-persistence.md`(1570 行,经 3 轮审查 + 文档-代码不一致修复:`search_orthogonal_vectors` 返回 1 条 score≈0 而非空列表)
|
||||
|
||||
**依赖**:Phase 14(Document 类型 + Embedding trait)
|
||||
**优先级**:P0
|
||||
**预估规模**:约 400 行(实际约 937 行纯实现 + 测试)
|
||||
**状态**:✅ Phase 15 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### Phase 16: 摘要自动生成
|
||||
|
||||
**目标**:闭环长对话能力。v0.2 的 `inject_summary` 消费端(`FocusedConfig.summary_override`)已就绪,缺的是生产端。
|
||||
|
||||
**交付物**:
|
||||
1. `SummaryConfig` 结构体:`trigger_token_ratio`(默认 0.75) / `max_context_tokens`(默认 32_000)/ `summary_prompt`(默认中文 `DEFAULT_SUMMARY_PROMPT` 含 `{messages}`) / `debounce_turns`(默认 3) / `summary_model`(默认 `None` 沿用主模型) / `max_tool_result_chars`(默认 500)
|
||||
2. 在 `submit_turn` / `finalize_turn` 中 OnTurnEnd 之后插入**内联检查点**(非 Hook 扩展):`should_summarize`(水位 + 防抖,首次不受防抖约束)→ `generate_summary` 关联函数(新 `LlmCycle` + `submit_messages` 无工具调用)→ 更新 `FocusedConfig.summary_override` + `slot.save()` 持久化 + `SessionMemory::set("conversation_summary", summary)` 全局快照
|
||||
3. `AgentBuilder` 扩展:`.summary_config(cfg)` 方法(不覆盖整个 `AgentConfig`)
|
||||
4. 公开 API:`get_conversation_summary() -> Result<Option<String>, AgentError>`
|
||||
5. `format_messages_as_text()` 简洁版消息格式化(System/User/Assistant + `[Tool: name]` + `Tool Result [id]:` 截断到 `max_tool_result_chars` 字符)
|
||||
|
||||
**设计决策**:内联于 `submit_turn` 流程而非 Hook 扩展(因为 HookContext 无法携带 `&mut self` 引用更新 slot config,且流式路径的 `finalize_turn` 中 `cycle` 已销毁)。`Option<SummaryConfig>` 的 opt-in 机制已足够提供可插拔性,不改变 Hook 系统签名。流式路径中 `submit_turn_stream` 已将 `turn_index` 提前 ++1,检查点使用 `saturating_sub(1)` 修正。
|
||||
|
||||
**依赖**:无(`submit_turn` 流程 + `CostTracker` + `SessionMemory` + `LlmProvider` 均已就绪)
|
||||
**优先级**:P0
|
||||
**预估规模**:约 220 行(实际约 250 行,含 11 个内联测试)
|
||||
**方案文档**:`docs/22-phase16-summary-auto-generation.md`(471 行,经 PM/SA 审查 11 项修复 + 实施后第二轮审查 9 项修复全部完成)
|
||||
**状态**:✅ Phase 16 全部交付物已完成(含实施后 PM/SA/Code Reviewer 第二轮审查 PASS)
|
||||
|
||||
**实施后审查修复记录**(共 9 项):
|
||||
- 🔴 B1:`generate_summary` 调用 `submit_messages(Vec::new(), vec![])` 发送空消息列表 → 移除 `with_messages()`,直接 `submit_messages(vec![Message::user_text(prompt)], vec![])`
|
||||
- 🟡 W4:`should_summarize` 使用 `self.turn_index` 而非 `current_turn` 参数 → 改签名接收 `current_turn`,流式路径防抖准确
|
||||
- 🟡 W2:`summary_model` 硬编码 `unwrap_or("gpt-4o")` → 改为条件赋值,`None` 时沿用 `CycleConfig::default()`
|
||||
- 🟡 W5:Full 模式 `slot.save()` 无谓调用 → 移入 `SlotMode::Focused` 分支内
|
||||
- 🟡 W3:`format_messages_as_text` 缺 30K 整体截断 → 新增 `MAX_TOTAL_CHARS=30_000` + `truncate_total_chars`,优先保留最新
|
||||
- 🟡 W6:摘要成功无日志 → 添加 `tracing::info!(turn, summary_len, "摘要自动生成成功")`
|
||||
- 🟡 W1/W7:缺 3 个测试 → 新增 `format_total_charset_truncation_keeps_recent` / `summary_written_to_focused_slot_config` / `summary_skipped_for_empty_messages` / `summary_not_generated_if_max_context_unreachable`
|
||||
- 💭 `context.rs:78` 过时注释("v0.3 将支持 Hook 驱动")→ 更新为"v0.3 Phase 16 起 AgentBuilder 内联检查点自动生成摘要"
|
||||
|
||||
**第二轮审查门禁**:PASS(0 🔴 阻塞)。`cargo test --all-targets` **353 passed / 0 failed**,clippy 0 警告,doc 0 warning。
|
||||
|
||||
---
|
||||
|
||||
#### Phase 17: Agent 执行引擎(会话树 + Time-travel Checkpointer)
|
||||
|
||||
**目标**:建立 `engine/` 模块。解决 v0.2 中"session 在变量里、无法通过 ID 恢复、不支持父子关系"的空白。
|
||||
|
||||
**方案文档**:`docs/23-phase17-agent-execution-engine.md`
|
||||
|
||||
**交付物**:
|
||||
1. `src/engine/` 新模块(`session_manager.rs` + `checkpointer.rs` + `snapshot.rs` + `error.rs`)
|
||||
2. `SessionManager`:
|
||||
- `create(agent, bundle) -> Result<String, EngineError>` — 创建根 session(UUID v4 自动生成 ID)
|
||||
- `create_child(parent_id, agent) -> Result<String, EngineError>` — 创建子 session(继承父 `RuntimeBundle`,`Arc::clone` 共享引用)
|
||||
- `get(session_id) -> Result<Arc<Mutex<AgentSession>>, EngineError>` — 按 ID 查找(仅查内存,不自动从存储恢复)
|
||||
- `recover(session_id, agent, bundle) -> Result<Arc<Mutex<AgentSession>>, EngineError>` — 从存储恢复 session
|
||||
- `replace(session_id, session) -> Result<(), EngineError>` — 替换已有 session 实例(用于 rollback 后切换)
|
||||
- `children(parent_id)` / `parent(child_id)` — 树形查询
|
||||
- `destroy(id)` — 生命周期管理(允许孤儿 session 存在,不递归删除子 session)
|
||||
3. `Checkpointer`:
|
||||
- `checkpoint(session)` — 每个 `submit_turn` 末尾自动保存全量状态快照
|
||||
- `rollback_load(session_id, ckpt_id) -> SessionSnapshot` — 读取 checkpoint JSON 为 snapshot(不重建 AgentSession)
|
||||
- `list_checkpoints(session_id)` — 列出 checkpoint 列表
|
||||
- `delete_all(session_id)` — 清理某 session 所有 checkpoint
|
||||
- `fork()` 推迟(底层可拆解为 `rollback` + `create_child`,作为高层 API 等价于约 30 行组合代码,已具备原始能力)
|
||||
4. `SessionSnapshot` 独立 struct(位于 `engine/snapshot.rs`)—— 避开 `Arc<dyn Agent>` 不可序列化的限制,通过 `to_snapshot()` / `from_snapshot()` 双向转换实现 AgentSession 快照持久化
|
||||
- `to_snapshot()`(async,从 `SessionMemory` 读取完整数据)+ `from_snapshot()`(纯同步构造)+ `restore_memory()`(async 写回持久层)
|
||||
5. `EngineError` 枚举(含 `MemoryError` 透传变体 与项目既有 `AgentError` 风格一致)
|
||||
|
||||
**Checkpoint 存储格式**:`ckpt:{session_id}:{ckpt_id}` → `SessionSnapshot` JSON(全量 session 状态,含所有 slot 消息列表)。Ponytail:全量 JSON 够用,等遇到存储效率问题时再改增量模式。
|
||||
|
||||
**会话树持久化**:`session:{session_id}:meta` → `SessionMeta` JSON(`{agent_name, parent_id, created_at, turn_count}`)
|
||||
|
||||
**依赖**:Phase 10(ContextSlot 持久化 — 消息由 slot 自己管,Checkpointer 管执行状态)
|
||||
**优先级**:P0
|
||||
**预估规模**:约 700 行(5 新增文件 + 5 修改文件)
|
||||
**状态**:✅ Phase 17 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
**实际新增**(2026-07-15,3 commits + 实施审查修复一轮):
|
||||
|
||||
- **新增 5 文件(`src/engine/`)**:`mod.rs`(19 行)+ `error.rs`(43 行)+ `snapshot.rs`(35 行)+ `checkpointer.rs`(373 行)+ `session_manager.rs`(910 行含测试)
|
||||
- **修改 5 文件**:
|
||||
- `src/llm/types/usage.rs` — `CostTracker` 加 `Clone, Serialize, Deserialize`(3 行)
|
||||
- `src/agent/context.rs` — `ContextSlot` + `MergeStrategy` 加 `Serialize, Deserialize`(4 行)
|
||||
- `src/agent/session_memory.rs` — 新增 `list_entries()` + `set_with_meta()` 方法
|
||||
- `src/agent/session.rs` — 新增 `to_snapshot()` (async) / `from_snapshot()` (sync) / `restore_memory()` (&mut self, async) / `has_pending_memory_restore()` + 公开 `bundle()` accessor
|
||||
- `src/lib.rs` — `pub mod engine`
|
||||
- **新增 1 示例**:`examples/engine_demo.rs`(~210 行,端到端演示 create → submit_turn → checkpoint → list → rollback_load → from_snapshot → restore_memory → replace → destroy 全链路,含 rollback 一致性 assert)
|
||||
- **依赖**:零新外部依赖(ponytail:ckpt_id 用纳秒+计数器生成,session_id 同理)
|
||||
- **测试**:353 → **374**(+21 引擎内联测试:Checkpointer 6 个 + SessionManager 15 个)
|
||||
- **质量基线**:`cargo test --all-targets` 374 passed / 0 failed;`cargo clippy --all-targets -- -D warnings` 0 警告;`cargo doc --no-deps` 0 warning;`cargo run --example engine_demo` exit 0
|
||||
- **关键设计决策落地**:
|
||||
- `SessionSnapshot` 独立 struct(避开 `Arc<dyn Agent>` 不可序列化)
|
||||
- `to_snapshot` async + `from_snapshot` 纯同步 + `restore_memory` async 三段式分离
|
||||
- `session_memory_data` 改用 `HashMap<String, SessionMemoryEntry>`(保留 metadata/created_at)
|
||||
- `EngineError::Memory(#[from] MemoryError)` 透传变体
|
||||
- `EngineManager` 锁契约:所有写操作先 HashMap 再 I/O(或反之,destroy 反向)
|
||||
- 自动 checkpoint 失败 `tracing::error!` 不阻断主流程(不提供强持久化保证)
|
||||
- ckpt_id 时间戳+纳秒+计数器无外部依赖(`created_at_nanos` 字段确保同秒内精确排序)
|
||||
- 孤儿策略:`destroy()` 不递归删除子 session;父被销毁后 `parent()` 返回 `Ok(None)`
|
||||
- **实施审查通过**:经过 PM + SA + Code Reviewer 三方联合审查 → 1 轮修复 → 全部 🟡 警告关闭
|
||||
- **M13 里程碑达成** — Phase 17 rc.1 标签可打(v0.3.0 第二个 Phase)
|
||||
|
||||
---
|
||||
|
||||
#### 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 行
|
||||
|
||||
**实际新增**(2026-07-15 commit `46de111`,详见 `docs/24-phase18-agent-switch-and-dispatch.md`):
|
||||
- 方案文档:`docs/24-phase18-agent-switch-and-dispatch.md`(700 行,含 Agent Switch 与 SubAgent Dispatch 的设计推演)
|
||||
- 新增文件 2 个:
|
||||
- `src/engine/switch.rs`(222 行)— `SessionManager::switch_agent()` 热切换:替换 `Arc<dyn Agent>`,slot 历史 / turn_index / session_memory / cost_so_far 全部保留,同步更新 `SessionMeta.agent_name` 到持久层
|
||||
- `src/engine/sub_agent.rs`(1071 行)— SubAgent 调度完整实现:4 个公开方法 + 3 个公开类型
|
||||
- 修改文件 4 个:
|
||||
- `src/engine/mod.rs`(+5 行:`pub mod switch; pub mod sub_agent;` + `pub use sub_agent::{DispatchConfig, SubTaskResult, SubTaskStreamEvent};`)
|
||||
- `src/engine/error.rs`(+1 变体:`DispatchFailed(#[source] String)`)
|
||||
- `src/engine/session_manager.rs`(+2 处可见性:`save_session_meta` / `load_session_meta` 改 `pub(crate)` 供 `switch.rs` 使用)
|
||||
- `src/llm/types/usage.rs`(+`From<Usage>` 实现供 `SubTaskResult.usage` 字段构造)
|
||||
- 新增 4 个示例:
|
||||
- `examples/agent_switch_demo.rs`(115 行)— Agent 热切换演示
|
||||
- `examples/sub_agent_dispatch_demo.rs`(141 行)— dispatch / dispatch_all 并行派发演示
|
||||
- `examples/bridge_keys_demo.rs`(197 行)— bridge_keys 过滤的 SessionMemory 继承演示
|
||||
- `examples/dispatch_stream_demo.rs`(121 行)— dispatch_stream 流式派发演示
|
||||
- 关键设计:
|
||||
- **`switch_agent` 锁契约**:先 `get` session → 锁 `Mutex` 替换 agent 并读取 turn_index → 释放 Mutex → 读/写 `SessionMeta`(无锁 IO),最大限度减少锁竞争
|
||||
- **`SessionMeta` 保留原则**:切换 `agent_name` 字段,但 `created_at` / `parent_id` 保留原始(血缘不可变)
|
||||
- **`switch_agent` 不自动 checkpoint**:与 `auto_checkpoint` 语义一致(仅 `submit_turn` / `finalize_turn` 触发),避免每次角色切换产生冗余 checkpoint
|
||||
- **`inherit_session_memory` 快照语义**:捕获调用时刻的父 session_memory 快照,子 session 写回后即使父被并发写入也不传播(防止非确定性结果)
|
||||
- **`bridge_keys` 三态语义**:`None` = 不继承任何(安全默认)/ `Some(vec![])` = 继承全部 / `Some(keys)` = 仅继承指定 key
|
||||
- **`shared_namespace` 约定式共享**:纯约定字段,不触发自动注入逻辑,子 agent 显式 `session.set_session_data("shared:{prefix}:{key}", value)` 写入
|
||||
- **`dispatch_all` 部分成功语义**:`Vec<Result<SubTaskResult, EngineError>>` 按输入顺序 indexed 收集,task panic 通过 `DispatchFailed` 哨兵占位(不破坏顺序一致性)
|
||||
- **`dispatch_stream` 后台 finalize**:spawn task 内部调 `finalize_turn()` 落库,明确不参与 `auto_checkpoint`(避免与流式 checkpoint 重复)
|
||||
- **`SubTaskStreamEvent` 事件序列**:`ChildCreated { child_id }` → `Stream(StreamEvent) × N` → `Completed(SubTaskResult)` 或 `Error { child_id, error }`
|
||||
- **`SUBTASK_NAMESPACE` 防误注入**:子 session 注入到 SessionManager 时使用 `subtask:` prefix 避免与 SessionMeta 的 `session:{id}:meta` 冲突
|
||||
- 测试:+17 内联测试(4 switch + 5 dispatch + 4 dispatch_all + 4 dispatch_stream),全量 374 → **391 passed / 0 failed**(+17,0 失败)
|
||||
- 质量基线:`cargo test --all-targets` 391 passed / 0 failed;`cargo clippy --all-targets -- -D warnings` 0 警告;`cargo doc --no-deps` 0 warning;4 个示例全部 exit 0
|
||||
- 零新外部依赖(ponytail:与 Phase 17 一致)
|
||||
|
||||
**状态**:✅ Phase 18 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
#### 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 行(实际约 720 行核心 + 200 行测试)
|
||||
**方案文档**:`docs/25-phase19-knowledge-graph-and-retrieval.md`(652 行,经 PM/SA 双轮审查 PASS)
|
||||
**状态**:✅ Phase 19 全部交付物已完成(2026-07-17)
|
||||
|
||||
**实际新增**(2026-07-17):
|
||||
- 新增文件 2 个:
|
||||
- `src/memory/graph.rs`(~580 行)- `GraphEntity`(id/name/entity_type/description/tags/properties)+ `GraphRelation`(无 id 字段,`composite_key()` 派生)+ `RelationDirection`(`#[derive(Default)]` + `#[default]` Outgoing)+ `ScoredEntity`(含 path 路径)+ `TagConstraints`(max_tags_per_entity 默认 8)+ `KnowledgeGraph` trait(10 个 async 方法)+ `InMemoryGraph`(`Mutex<GraphInner>` 单一锁结构,避免嵌套锁死锁)+ BFS 图遍历(visited 防环 + 权重乘积衰减 + 多路径先到先得 + depth=0 返回空)+ 标签管理(tag_index 反向索引)+ 23 个内联测试
|
||||
- `examples/knowledge_graph_demo.rs`(~140 行)- 端到端演示:构建图谱 -> BFS 遍历 -> 标签管理 -> Hybrid/GraphOnly 双通道检索
|
||||
- 修改文件 3 个:
|
||||
- `src/memory/retriever.rs` - `RetrievalStrategy` 枚举(Hybrid 默认 / KnowledgeOnly / GraphOnly)+ `RetrievalItem` enum(统一列表,`score()` 方法)+ `RetrievalResult` 新增 `strategy` 字段(反映实际执行策略)+ `MemoryRetriever` 双通道(`with_knowledge_graph` / `with_strategy` 链式构造)+ `search_knowledge_store` / `search_graph` 私有方法 + `tokio::join!` 并行 + 旧 `ScoredItem` 标注 `#[deprecated]` + `RetrieverConfig` 新增 `graph_depth`(默认 2)+ 13 个内联测试
|
||||
- `src/memory.rs` - `pub mod graph` + 重导出 7 个图类型 + 更新 retriever 重导出
|
||||
- `examples/knowledge_search_demo.rs` - 适配新 API(`RetrievalItem` enum match + `RetrieverConfig.graph_depth`)
|
||||
- 关键设计:
|
||||
- **`Mutex<GraphInner>` 单一锁结构** - 避免 `set_entity_tags` 嵌套锁死锁风险(审查修复)
|
||||
- **`RetrievalResult.strategy` 反映实际执行策略** - graph 未注入时退化为 `KnowledgeOnly`(审查修复)
|
||||
- **`GraphRelation` 无 id 字段** + `composite_key()` 派生方法
|
||||
- **BFS**:`visited` 防环 + 权重乘积衰减 + 多路径先到先得 + `depth=0` 返回空
|
||||
- **零新外部依赖**(ponytail 风格)
|
||||
- 测试:391 -> **427 passed / 0 failed**(+36 新测试:23 graph + 13 retriever)
|
||||
- 质量基线:`cargo test --all-targets` 427 passed / 0 failed;`cargo clippy --all-targets -- -D warnings` 0 警告;`cargo doc --no-deps` 0 warning;`cargo run --example knowledge_graph_demo` exit 0
|
||||
|
||||
---
|
||||
|
||||
### 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<br/>19 新测试"]:::done
|
||||
P16["<b>Phase 16: 摘要自动生成</b><br/>SummaryConfig<br/>内联检查点<br/>首次防抖跳过<br/>18 新测试"]:::done
|
||||
P17["<b>Phase 17: 执行引擎</b><br/>SessionManager<br/>会话树<br/>Time-travel Checkpointer<br/>21 新测试"]:::done
|
||||
P18["<b>Phase 18: 切换与调度</b><br/>Agent Switch<br/>SubAgent Dispatch<br/>dispatch_all 并发控制<br/>17 新测试"]:::done
|
||||
P19["<b>Phase 19: 知识图谱</b><br/>KnowledgeGraph trait<br/>InMemoryGraph<br/>双通道检索"]:::done
|
||||
|
||||
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` 端到端验证 | ✅ 2026-07-09 |
|
||||
| **M12** | Phase 16 | 多轮对话后摘要自动写入 SessionMemory、派生 slot 时摘要正确注入 + 第二轮实施审查 PASS | ✅ 2026-07-10 |
|
||||
| **M13** | **Phase 17 (rc.1)** | `SessionManager` 创建/recover/replace/子树/销毁集成测试通过、`Checkpointer` checkpoint/rollback/list_checkpoints 验证(`fork` 推迟,按需时引入)| ✅ 2026-07-15 |
|
||||
| **M14** | Phase 18 | `switch_agent` 热切换验证(slot / turn_index / session_memory 保留)、`dispatch` / `dispatch_all` 并行派发 + Semaphore 顺序、`dispatch_stream` 流事件序列验证 | ✅ 2026-07-15 |
|
||||
| **M15** | Phase 19 | `KnowledgeGraph` 实体-关系 CRUD + `get_related` BFS 验证、双通道检索 Hybrid 策略验证 | ✅ 2026-07-17 |
|
||||
+15
-686
@@ -1,692 +1,21 @@
|
||||
# AG Core Roadmap
|
||||
|
||||
> 定稿日期:2026-05-11
|
||||
> 最后更新:2026-07-06
|
||||
> 拆分式 roadmap:按版本归档 + 未归类内容
|
||||
> 最后更新:2026-07-17(v0.3.0 Phase 19 完成 + M15 里程碑达成 + v0.3.0 全部交付完毕)
|
||||
|
||||
## 愿景
|
||||
## 文件索引
|
||||
|
||||
AG Core 定位为构建 AI 智能体的底层工具箱,通过模块化、可插拔的架构,提供大模型调用、提示词工程、工具系统、记忆检索四大核心能力,支持快速组合出符合业务需求的智能体应用。
|
||||
| 文件 | 范围 | 状态 |
|
||||
|------|------|------|
|
||||
| [`roadmap-v0.1.0.md`](./roadmap-v0.1.0.md) | v0.1.0 计划与交付 — Phase 0–4c + v0.1.0 Release | ✅ 已发布 2026-07-04 |
|
||||
| [`roadmap-v0.2.0.md`](./roadmap-v0.2.0.md) | v0.2.0 计划与交付 — Phase 5–12 + v0.2.0-rc.1 | 🟡 Phase 5-11 已完成;Phase 12 P2 锦上添花可选 |
|
||||
| [`roadmap-v0.3.0.md`](./roadmap-v0.3.0.md) | v0.3.0 计划与交付 - Phase 13–19 | ✅ Phase 13-19 全部完成,v0.3.0 交付完毕 |
|
||||
| [`roadmap-unsorted.md`](./roadmap-unsorted.md) | 未归到任何版本的内容 — 全局愿景、当前状态、模块完整性、v0.4+ 展望、风险与建议、下一步行动、阶段总回顾 | — |
|
||||
|
||||
**当前状态**:v0.1.0 已发布(2026-07-04)。Phase 0-11 全部完成,v0.2.0-rc.1 已打标签。Provider IR 重构 + LlmCycle 简化 + 11 个离线示例(含 `quick_start` 30 行最小示例、`end_to_end` 完整集成示例、`context_slot_demo` 分支对话示例)+ SqliteStore 持久化 + 14 个公开枚举 `#[non_exhaustive]` 护栏 + `StepStatus` IR 迁移 + `submit_turn_stream` 流式体验 + ContextSlot 多上下文分区管理 + `VectorRetriever` 语义检索 trait + 12 个 wiremock Provider roundtrip 测试 + 5 个并发测试已交付。下一步进入 v0.2.0 正式版打 tag 流程。
|
||||
## 阅读建议
|
||||
|
||||
---
|
||||
|
||||
## 模块完整性评估
|
||||
|
||||
| 功能领域 | 方案状态 | 文档位置 | 实现优先级 |
|
||||
|---------|---------|---------|-----------|
|
||||
| LLM 调用周期 | ✅ 完整 | `specs/llm-call-lifecycle.md` | P0 |
|
||||
| 提示词工程 | ✅ 完整 | `docs/4-prompt-engineering.md` | P1 |
|
||||
| 工具系统 + 权限 | ✅ 完整 | `docs/5-tool-system.md` | P1 |
|
||||
| 记忆检索 | ✅ 完整 | `docs/6-memory-system.md` | P2 |
|
||||
| Agent 运行时(4a 胶水层) | ✅ 已实现 | `docs/7-agent-runtime.md` | P2 |
|
||||
| 生命周期钩子 | ✅ 完整 | `docs/3-phase0-remaining.md` | P0(LLM Cycle 扩展) |
|
||||
| Provider 注册发现 | ✅ 完整 | `docs/3-phase0-remaining.md` | P0(Provider 接口扩展) |
|
||||
| 流式事件系统 | ✅ 完整 | `docs/3-phase0-remaining.md` | P0(流式接口前置) |
|
||||
|
||||
---
|
||||
|
||||
## 分阶段 Roadmap
|
||||
|
||||
### Phase 0 — Foundation(基础设施)
|
||||
|
||||
**目标**:实现 LLM 调用周期的核心功能,作为所有上层模块的基础。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `llm/types.rs` — 核心数据类型(Message, ContentBlock, ChatRequest/Response, ToolDefinition, StopReason)
|
||||
2. ✅ `llm/error.rs` — 错误体系(LlmError 枚举,可重试/不可重试判断)
|
||||
3. ✅ `llm/provider.rs` + `llm/provider/openai.rs` — Provider 接口 + OpenAI 兼容实现
|
||||
4. ✅ `llm/provider/registry.rs` — ProviderRegistry(多 Provider 注册发现)
|
||||
5. ✅ `llm/cycle.rs` + `llm/cycle/{retry,usage}.rs` — 生命周期引擎(重试策略 + 用量追踪)
|
||||
6. ✅ `llm/hooks.rs` — HookExecutor 接口(生命周期钩子)
|
||||
7. ✅ `llm/stream.rs` — StreamEvents 流式事件系统(AssistantTextDelta, ToolExecutionStarted 等)
|
||||
8. ✅ `llm/compact.rs` — Auto-compaction(上下文自动压缩)
|
||||
9. ✅ `Cargo.toml` — 添加依赖(tokio, reqwest, serde, thiserror, async-trait, tracing)
|
||||
|
||||
**依赖**:无
|
||||
|
||||
**优先级**:Must Have
|
||||
|
||||
**预估规模**:约 1000 行核心代码
|
||||
|
||||
**状态**:✅ Phase 0 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 1 — Prompt Engineering(提示词工程)
|
||||
|
||||
**目标**:提供提示词的组合、模板化与优化能力。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `prompt.rs` + `prompt/` 模块
|
||||
2. ✅ `PromptTemplate` — 模板引擎(支持变量插值、条件渲染)
|
||||
3. ✅ `PromptComposer` — 提示词组合器(拼接 system/user/assistant 消息)
|
||||
4. ✅ `docs/4-prompt-engineering.md` — 方案文档
|
||||
|
||||
**依赖**:无(可与 Phase 0 并行)
|
||||
|
||||
**优先级**:Should Have
|
||||
|
||||
**预估规模**:约 400 行代码
|
||||
|
||||
**状态**:✅ Phase 1 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 2 — Tool System(工具系统)
|
||||
|
||||
**目标**:实现 MCP 协议集成与自定义工具注册、调用、权限控制。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `tools.rs` + `tools/` 模块(base/registry/permission/mcp/error)
|
||||
2. ✅ `ToolRegistry` — 工具注册表(注册、发现、调用、并行执行、超时控制)
|
||||
3. ✅ `BaseTool` trait — 工具抽象接口(含 ToolContext 执行上下文)
|
||||
4. ✅ `McpClient` — MCP 协议客户端(stdio transport,StreamableHttp 预留)
|
||||
5. ✅ `PermissionChecker` — 工具执行权限检查(白名单/黑名单/自定义权限)
|
||||
6. ✅ `docs/5-tool-system.md` — 方案设计文档
|
||||
7. ✅ 扩展 `llm/cycle.rs` 支持自动 tool 循环(`submit_with_tools()` + `submit_request()` + `maybe_compact()`)
|
||||
8. ✅ `ToolError` — 结构化错误体系(含 `is_recoverable()` 分类)
|
||||
|
||||
**依赖**:Phase 0(LlmProvider 接口传递 tool definitions)、Phase 1(提示词可能需要注入工具描述)
|
||||
|
||||
**优先级**:Should Have
|
||||
|
||||
**预估规模**:约 900 行代码(实际约 1500 行)
|
||||
|
||||
**状态**:✅ Phase 2 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 3 — Memory System(记忆系统)
|
||||
|
||||
**目标**:提供对话记忆的存储、检索与管理能力。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `memory.rs` + `memory/` 模块(store / conversation / knowledge / retriever / error / types)
|
||||
2. ✅ `MemoryStore` trait + `InMemoryStore` — 记忆存储抽象(可插拔后端)+ 默认实现
|
||||
3. ✅ `ConversationMemory` — 对话记忆管理(sliding window / 全量),复用 `llm::compact`
|
||||
4. ✅ `KnowledgeStore` — 知识页面存储(具体 struct,非 trait,基于 MemoryStore)
|
||||
5. ✅ `MemoryRetriever` — 记忆检索器(TextOverlap Dice 系数评分,单通道)
|
||||
6. ✅ `docs/6-memory-system.md` — 方案设计文档
|
||||
7. ✅ `docs/note-knowledge-graph-design.md` — KnowledgeGraph 等 Phase 4 备用设计
|
||||
8. ✅ `EvictionPolicy` — 支持 None / Ttl / Capacity 三种淘汰策略
|
||||
|
||||
**依赖**:Phase 0(llm::compact 复用)、Cargo.toml 新增 `time` 依赖
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 700 行代码(实际约 1242 行,含测试)
|
||||
|
||||
**状态**:✅ Phase 3 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 4a — Agent Core Glue(核心胶水层)
|
||||
|
||||
**目标**:提供最小可用的 Agent Runtime——把 Phase 0-3 的能力"装配"成 `AgentSession::submit_turn`。上层可基于 4a 构建多轮对话应用。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `agent.rs` + `agent/` 模块(7 个文件:agent/error/runtime/builder/session/task + 模块根)
|
||||
2. ✅ `Agent` trait — 智能体角色定义(name / system_prompt / tool_definitions)
|
||||
3. ✅ `AgentSession` — 会话实例(绑定 `Arc<dyn Agent>` + `RuntimeBundle` + 内联 HashMap session_data)
|
||||
4. ✅ `RuntimeBundle` — 显式依赖注入容器(不含 session_memory_backend)
|
||||
5. ✅ `AgentBuilder` — 链式构造入口(不含 session_memory_backend)
|
||||
6. ✅ `AgentError` — 统一错误类型(7 个变体:Llm / Tool / Memory / HookBlocked / LimitExceeded / Config / Other;不含 PlanParse)
|
||||
7. ✅ `Plan` / `Step` / `StepStatus` — 纯数据结构(不含任何解析逻辑)
|
||||
8. ✅ Hook 事件扩展:OnTurnStart / OnTurnEnd + turn_index 字段
|
||||
9. ✅ `docs/7-agent-runtime.md` — 方案设计文档(含 4a/4b/4c 分阶段计划)
|
||||
|
||||
**实际新增**:
|
||||
- 新增文件 7 个(agent.rs + agent/{agent, error, runtime, builder, session, task}.rs)
|
||||
- 修改文件 3 个(lib.rs +1 行;llm/hooks.rs +13 行追加变体/字段;llm/cycle.rs 内部字段 Box→Arc + 新增 `new_with_arc` 公共方法)
|
||||
- 实际代码量约 800 行(含测试;纯实现约 470 行——略高于方案预估 440 行,因 AgentSession 的 tests 模块内联 MockProvider/StubAgent 等辅助结构)
|
||||
- 新增内联测试 22 个;全量测试 84 → 109(0 失败)
|
||||
- clippy 0 警告(agent 模块)
|
||||
- 无新增外部依赖
|
||||
|
||||
**依赖**:Phase 0, 1, 2, 3
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 440 行代码
|
||||
|
||||
**状态**:✅ Phase 4a 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 4b — Task Execution(任务执行)
|
||||
|
||||
**目标**:在 Phase 4a 基础上,赋予智能体"拆解目标 → 逐步执行"的能力。
|
||||
|
||||
**前置条件**:Phase 4a 已完成。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `TaskAgent` trait — `run(goal)` 自主式 + `execute_plan(plan)` 外部驱动式
|
||||
2. ✅ `PlanParser` trait + `JsonPlanParser` 参考实现
|
||||
3. ✅ `AgentError` 追加 PlanParse 变体(共 7 个变体)
|
||||
4. ✅ Hook 事件扩展:OnPlanStepComplete + plan_step_index 字段
|
||||
|
||||
**依赖**:Phase 4a
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 200 行代码(增量)
|
||||
|
||||
**实际新增**:
|
||||
- 修改文件 2 个(llm/hooks.rs +5 行;agent/error.rs +10 行)
|
||||
- 新增代码约 150 行(含测试;纯实现约 90 行)
|
||||
- 新增内联测试 4 个;全量测试 109 → 113(0 失败)
|
||||
- clippy 0 警告
|
||||
- 无新增外部依赖
|
||||
|
||||
**状态**:✅ Phase 4b 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
### Phase 4c — Session Memory(会话级记忆)
|
||||
|
||||
**目标**:提供会话级 key-value 记忆,作为 session 内各 context 之间的信息桥接通道。
|
||||
|
||||
**前置条件**:Phase 4a 已完成(可与 Phase 4b 并行)。
|
||||
|
||||
**交付物**:
|
||||
1. ✅ `SessionMemory` struct — 基于 `MemoryStore`,按 session_id namespace 隔离
|
||||
2. ✅ `RuntimeBundle` + `AgentBuilder` 扩展 `session_memory_backend` 字段
|
||||
3. ✅ `AgentSession` 替换内联 HashMap 为完整 `SessionMemory`
|
||||
|
||||
**依赖**:Phase 4a(Phase 3 MemoryStore)
|
||||
|
||||
**优先级**:Could Have
|
||||
|
||||
**预估规模**:约 115 行代码(增量)
|
||||
|
||||
**实际新增**:
|
||||
- 新增文件 1 个(agent/session_memory.rs)
|
||||
- 修改文件 4 个(agent/runtime.rs +5 行;agent/builder.rs +10 行;agent/session.rs +30 行;agent.rs +2 行)
|
||||
- 新增代码约 180 行(含测试;纯实现约 100 行)
|
||||
- 新增内联测试 3 个;全量测试 113 → 116(0 失败)
|
||||
- clippy 0 警告
|
||||
- 无新增外部依赖
|
||||
|
||||
**状态**:✅ Phase 4c 全部交付物已完成
|
||||
|
||||
---
|
||||
|
||||
## 依赖关系图
|
||||
|
||||
```mermaid
|
||||
graph BT
|
||||
P0["<b>Phase 0: Foundation</b><br/>LLM Cycle<br/>ProviderRegistry<br/>HookExecutor<br/>StreamEvents<br/>Auto-compaction"]:::done
|
||||
P1["<b>Phase 1: Prompt Engineering</b><br/>PromptTemplate<br/>PromptComposer"]:::done
|
||||
P2["<b>Phase 2: Tool System</b><br/>Tool Registry<br/>PermissionChecker<br/>MCP Client"]:::done
|
||||
P3["<b>Phase 3: Memory System</b><br/>MemoryStore<br/>ConversationMemory<br/>KnowledgeStore"]:::done
|
||||
P4a["<b>Phase 4a: Core Glue</b><br/>AgentSession<br/>RuntimeBundle<br/>Plan/Step 纯数据"]:::done
|
||||
P4b["<b>Phase 4b: Task Execution</b><br/>TaskAgent<br/>PlanParser<br/>JsonPlanParser"]:::done
|
||||
P4c["<b>Phase 4c: Session Memory</b><br/>SessionMemory"]:::done
|
||||
|
||||
P1 --> P0
|
||||
P2 --> P0
|
||||
P3 --> P0
|
||||
P2 --> P1
|
||||
P4a --> P1
|
||||
P4a --> P2
|
||||
P4a --> P3
|
||||
P4b --> P4a
|
||||
P4c --> P4a
|
||||
|
||||
classDef done fill:#4ade80,stroke:#16a34a,color:#1a1a1a
|
||||
classDef pending fill:#fbbf24,stroke:#d97706,color:#1a1a1a
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## v0.2.0 — 生产就绪(Production-Ready Core)
|
||||
|
||||
**目标**:解决 Rust Agent 工具箱从"能跑"到"能被人依赖"的鸿沟。持久化、配置层、上下文管理三大块补齐后,开发者可在 30 分钟内写出生产可用的 Agent 服务。
|
||||
|
||||
**总体规模**:8 个增量 Phase(Phase 5-12),17 个可验证 Step。
|
||||
|
||||
### 功能清单
|
||||
|
||||
#### P0 — 必须交付
|
||||
|
||||
| # | 功能 | 模块 | 方案要点 |
|
||||
|---|------|------|---------|
|
||||
| 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`) |
|
||||
|
||||
#### P1 — 重要但不阻塞
|
||||
|
||||
| # | 功能 | 模块 | 方案要点 |
|
||||
|---|------|------|---------|
|
||||
| 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 测试 |
|
||||
|
||||
#### P2 — 有时间再做
|
||||
|
||||
| # | 功能 | 模块 | 备注 |
|
||||
|---|------|------|------|
|
||||
| 10 | MCP StreamableHttp | `tools` | 当前仅预留枚举变体 |
|
||||
| 11 | Gemini Provider | `llm/provider` | 协议差异大,实现成本较高 |
|
||||
| 12 | 文件系统 MemoryStore 后端 | `memory` | JSON/JSONL 轻量持久化 |
|
||||
|
||||
### ContextSlot 上下文管理
|
||||
|
||||
**模块归属**:`src/llm/context.rs`(与 `compact.rs` 同级)
|
||||
|
||||
**核心概念**:`ContextSlot` 是一段带策略配置的消息列表,以 `slot_id` 为 namespace 独立持久化到 `MemoryStore`。支持三种模式、三种来源和派生关联(记录 `parent_id`)。
|
||||
|
||||
**核心类型**:
|
||||
|
||||
```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}` → 消息内容
|
||||
- `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` 继续服务传统对话场景。
|
||||
|
||||
**v0.2 不做**:
|
||||
- ❌ `slot.fork()` / `merge()` — 分支方法推迟到 v0.3+
|
||||
- ❌ `inject_summary` 自动生成 — v0.2 仅消费端(从 `SessionMemory` 读取),生成在 v0.3+
|
||||
- ❌ 血缘关系图遍历 — 只存 `parent_id`,不做查询
|
||||
|
||||
**依赖**:Phase 0(MemoryStore trait)、Phase 3(MemoryStore 持久化)
|
||||
**优先级**:P1
|
||||
|
||||
---
|
||||
|
||||
### v0.2.0 实施计划 — 8 个增量 Phase
|
||||
|
||||
> **编号说明**:Phase 5-12 接续 v0.1 的 Phase 0-4c,按开发顺序排列。
|
||||
|
||||
#### Phase 5: 热身准备(Warmup)
|
||||
|
||||
**目标**:快速交付三个互不依赖的独立改动,建立交付节奏。
|
||||
|
||||
| Step | 内容 | 文件范围 | 验证标准 |
|
||||
|------|------|---------|---------|
|
||||
| **5.1** ✅ | `ProviderConfig` 扩展:补 `timeout_secs`(def=30) + `max_retries`(def=3);新增 `ProviderConfig::from_env(prefix)` | `llm/provider.rs` + 各 Provider `new()` 构造函数 | `cargo test` + `from_env()` 单元测试 |
|
||||
| **5.2** ✅ | `OllamaProvider`:基于 `GenericOpenaiProvider` 包装,改 base_url 为 `http://localhost:11434`;`ProviderType` 新增 `Ollama` | `llm/provider/provider.rs` + `llm/provider/ollama.rs`(新增) | `cargo build` — 纯类型级验证 |
|
||||
| **5.3** ✅ | 公开枚举 `#[non_exhaustive]` 前置标记:`ProviderType` / `StopReason` / `FinishReason` / `EvictionPolicy` / `SlotMode`(预置) | 各枚举定义处 | 编译通过 + `cargo clippy` 0 警告 |
|
||||
|
||||
**实际新增**(2026-07-05 commit `98dfe6c`):
|
||||
- 新增文件 1 个(`llm/provider/ollama.rs`,72 行)
|
||||
- 修改文件 2 个(`llm/provider.rs` 加 `from_env` + `Default` + 4 个字段;`memory/store.rs` EvictionPolicy 加 `#[non_exhaustive]`)
|
||||
- `ProviderType::Ollama` 变体 + `FromStr` 解析("ollama" → Ollama)
|
||||
- `OllamaProvider::new(base_url, api_key, model, timeout_secs)` + `with_client()` 构造函数
|
||||
- `ProviderConfig::from_env(prefix)` 解析 `{prefix}_API_KEY` / `{prefix}_BASE_URL` / `{prefix}_MODEL` 环境变量
|
||||
- 全量测试 182 → 190(+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+ 展望
|
||||
|
||||
### 已规划的功能
|
||||
|
||||
| 功能 | 说明 | 预计版本 |
|
||||
|------|------|---------|
|
||||
| ContextSlot 分支(fork/merge) | 在决策点 fork 出子上下文,分支独立演进,可合并/丢弃 | v0.3 |
|
||||
| 摘要自动生成 | Hook 驱动,`OnTurnEnd` 自动将对话摘要写入 `SessionMemory`,`inject_summary` 消费端已在 v0.2 就绪 | v0.3 |
|
||||
| 知识图谱 | 实体-关系图,`docs/note-knowledge-graph-design.md` 已记录设计 | v0.3+ |
|
||||
| Multi-Agent 协同(Swarm) | 子 Agent 委派、并行子任务、结果聚合 | v0.4+ |
|
||||
| 精确 tokenizer 计数 | 绑定具体模型的 tokenizer 计数,替代当前的字符估算 | v0.3+ |
|
||||
| 血缘关系图遍历 | 以 `parent_id` 为基础,提供 slot 血缘链查询 | v0.3+ |
|
||||
| Markdown 技能按需加载 | 兼容 `SKILL.md` 格式,按 prompt 上下文动态加载 | v0.3+ |
|
||||
| TokenJuice 语义压缩 | 对工具结果做语义压缩而非字节截断 | v0.3+ |
|
||||
| Human-in-the-loop 审批 | 高危工具执行前的异步审批回调 | v0.3+ |
|
||||
| RL 轨迹导出 | ShareGPT 格式轨迹、Atropos 集成 | v0.4+ |
|
||||
|
||||
### 明确不做(agcore 范围外)
|
||||
|
||||
| 功能 | 原因 |
|
||||
|------|------|
|
||||
| TUI / 多平台 Gateway | 应用层职责(Feishu / Telegram / Discord 桥接) |
|
||||
| 配置自动加载(config/figment) | 配置来源策略应由上游应用决定,agcore 不定义配置格式 |
|
||||
| 提示词自动优化 | 属于智能层,不应内建于 core 库 |
|
||||
|
||||
---
|
||||
|
||||
## 风险与建议
|
||||
|
||||
1. **持久化依赖**:`rusqlite` + `bundled` 零外部依赖编译,但 SQLite 不适配所有场景(分布式/高并发写)。`MemoryStore` trait 的抽象层允许下游自行实现 Redis / PostgreSQL 后端
|
||||
2. **ContextSlot 心智负担**:`ContextSlot` 引入了一等抽象的复杂度。建议通过 `AgentBuilder` 默认创建 `"default"` slot,让简单场景无感使用
|
||||
3. **向量检索生态**:`VectorRetriever` trait-only 不绑定实现,需社区贡献或用户自行适配 pgvector / qdrant / lancedb
|
||||
4. **Scope 蔓延**:agcore 定位为"支持库"而非"Agent 产品",始终以 trait + reference impl 为边界,业务循环留给上层
|
||||
5. **API 稳定性**:v0.2 引入 `#[non_exhaustive]` 和 `#[deprecated]` 机制,但不承诺 SemVer 稳定——仍在快速迭代期
|
||||
|
||||
---
|
||||
|
||||
## 下一步行动
|
||||
|
||||
1. **v0.2.0 正式版打 tag**:Phase 8-11 全部完成,去掉 rc 后缀打 `v0.2.0` 正式版标签;CHANGELOG 整理 + Cargo.toml version 0.2.0-rc.1 → 0.2.0
|
||||
2. **Phase 12 评估**(可选):P2 锦上添花三项(文件系统 MemoryStore / MCP StreamableHttp / Gemini Provider)按需选做
|
||||
3. **示例先行**:v0.2 范围内每完成一个 Phase 立即更新对应示例,验证通过后再合入
|
||||
4. **里程碑追踪**:以 Phase 11(已完成,2026-07-06)为最新节点,逐 Phase 验收
|
||||
|
||||
**已完成 / 进行中阶段**:
|
||||
- ✅ Phase 0 Foundation — 全部交付物已完成
|
||||
- ✅ Phase 1 Prompt Engineering — 全部交付物已完成
|
||||
- ✅ Phase 2 Tool System — 全部交付物已完成
|
||||
- ✅ Phase 3 Memory System — 全部交付物已完成
|
||||
- ✅ Phase 4a Core Glue — 全部交付物已完成
|
||||
- ✅ Phase 4b Task Execution — 全部交付物已完成
|
||||
- ✅ Phase 4c Session Memory — 全部交付物已完成
|
||||
- ✅ Phase 5 Warmup — ProviderConfig::from_env + OllamaProvider + `#[non_exhaustive]` 前置标记(ProviderType / StopReason / FinishReason / EvictionPolicy)
|
||||
- ✅ Phase 6 ToolDefinition IR — `ToolDef` 新类型 + 双向 `From` 转换 + 别名彻底移除 + `#[allow(deprecated)]` 清理(cycle/registry/mcp/agent);Anthropic 零改动;roundtrip 测试覆盖
|
||||
- ✅ Phase 7 SqliteStore — `rusqlite 0.32` + WAL 模式 + `Arc<Mutex<Connection>>` + `spawn_blocking`;`memory/store.rs` → `store/{in_memory,sqlite_store}.rs` 模块化;9 个内联测试覆盖 CRUD/upsert/过滤/10×10 并发/持久化 round-trip;`InMemoryStore ↔ SqliteStore` trait-box 互换兼容
|
||||
- ✅ **Phase 8 MVP 集成出口** — 14 个公开枚举追加 `#[non_exhaustive]`(P0 核心 IR + P0 Error + P1 其他) + `StepStatus::Completed(ChatResponse)` → `Completed(MessageResponse)` 迁移 + CHANGELOG v0.2.0-rc.1 + 2 个新示例(`quick_start` 60 行 + `end_to_end` 246 行),10 个离线示例全部 exit 0;**v0.2.0-rc.1 标签已打**;实施后三方审查发现 6 项问题(1 🔴 + 2 🟡 + 3 💭)已全部修复
|
||||
- ✅ **Phase 9 流式体验增强** — `AgentSession::submit_turn_stream` 流式事件序列 + `LlmCycle::submit_with_tools_stream` spawn + mpsc 状态机 + `StreamEvent::ToolExecutionStarted`/`Completed` 新变体 + 9 单元测试 + 2 集成测试(含 `submit_turn_stream_end_to_end` 端到端 mock 验证 + `submit_turn_stream_triggers_turn_hooks` Hook 触发验证),全量 200 → 211;`CycleConfig` 加 `Clone` derive;方案文档 `docs/16-phase9-streaming-experience.md`(821 行)
|
||||
- ✅ **Phase 10 ContextSlot 上下文管理** — `src/agent/context.rs` 新增 `ContextSlot` 核心类型(Full / Focused / Readonly 三种模式,New / Derived / Static 三种来源)+ JSON blob 批次持久化(每 slot 3-4 条 MemoryItem,`slot_config` key 自恢复支持旧版本兼容);`AgentSession` 扩展 slots 字段 + 5 个管理方法(`create_slot` / `switch_slot` / `list_slots` / `derive_slot` / `delete_slot`,自动创建 `"default"` slot,`delete_slot` 双重保护禁止删 default/最后一个);`submit_turn`/`finalize_turn` 改造为基于当前 slot 的增量追加写回(`cycle.messages()[input_len..]` 提取本轮新增消息,确保 Focused 模式"读时过滤"语义不丢失数据);`finalize_turn` 签名变更(新增 `new_messages_from_cycle: Vec<Message>` 参数,返回 `Result<(), AgentError>`);`agent/error.rs` 新增 3 个 Slot 错误变体(`SlotReadonly` / `SlotNotFound` / `SlotAlreadyExists`);`examples/context_slot_demo.rs` 新增分支对话示例(法律咨询入口 → 两个派生方向 → 切换 → 隔离验证 → 删除保护);方案文档 `docs/17-phase10-contextslot.md`(1227 行,含 §5 推荐方案、§6 实施建议、§9 实施计划,经过 4 轮方案/计划/实施审查 + 1 轮非阻塞建议修复);全量 211 → 254(+43 新测试),clippy 0 警告,doc 0 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
|
||||
- ✅ 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.1 发布里程碑(2026-07-04)
|
||||
|
||||
**质量基线**:
|
||||
|
||||
| 指标 | 数值 |
|
||||
|------|------|
|
||||
| `cargo build --all-targets` | ✅ 通过 |
|
||||
| `cargo test --all-targets` | ✅ **182 passed / 0 failed** |
|
||||
| `cargo clippy --all-targets -- -D warnings` | ✅ 0 警告 |
|
||||
| 离线示例(`cargo run --example`) | ✅ 7 个全部 exit 0 |
|
||||
|
||||
**关键交付**:
|
||||
1. **Provider IR 重构** — 统一 `Message` / `ContentBlock` / `MessageRequest` / `MessageResponse` 类型层;4 个 Provider 适配(OpenAI Chat / Anthropic Messages / DeepSeek / Qwen);`LlmProvider` trait 签名同步切换
|
||||
2. **LlmCycle 简化** — `LlmCycle` 内部消息类型切到 IR 层;移除 Phase 0 的 `OpenaiChatMessage ↔ Message` 桥接;测试从 116 → 182(含 provider 测试)
|
||||
3. **`MockProvider` 公开化** — `agcore::llm::mock::MockProvider` 支持 `chat` + `chat_stream`,无需 API key 即可运行示例
|
||||
4. **7 个离线示例** — `prompt_composer` / `custom_tool` / `agent_session_demo` / `task_agent_demo` / `conversation_memory_demo` / `knowledge_search_demo` / `streaming_events_demo`
|
||||
5. **错误消息友好化** — `AgentError` / `LlmError` / `ToolError` / `MemoryError` / `PromptError` 全部面向最终用户改写(给出可操作的建议)
|
||||
6. **文档完整** — README 完整版(快速上手 + 架构图 + 环境变量)、Apache-2.0 LICENSE
|
||||
- **按版本顺序追溯历史**:v0.1.0 → v0.2.0 → v0.3.0
|
||||
- **了解产品演进全貌**:从 `roadmap-unsorted.md` 顶部开始读
|
||||
- **查找特定 Phase**:每个版本文件内按 Phase 编号顺序排列
|
||||
- **了解项目当前关注点**:从 `roadmap-unsorted.md` 的「下一步行动」开始
|
||||
- **未来规划视野**:从 `roadmap-unsorted.md` 的「v0.4+ 展望」开始
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
//! agent_switch_demo —— Agent 角色热切换示例。
|
||||
//!
|
||||
//! 演示:
|
||||
//! 1. 创建 session(绑定 Analyst agent)
|
||||
//! 2. 提交一轮对话(角色 A 输出"分析数据")
|
||||
//! 3. switch_agent 切换为 Reporter agent
|
||||
//! 4. 提交第二轮对话(角色 B 基于已有上下文输出"报告")
|
||||
//! 5. 验证:turn_index 连续、session_memory 保留、slot 历史保留
|
||||
//!
|
||||
//! 运行:`cargo run --example agent_switch_demo`
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::agent::{Agent, AgentBuilder};
|
||||
use agcore::engine::SessionManager;
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::types::Usage;
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use agcore::memory::store::InMemoryStore;
|
||||
use agcore::tools::ToolRegistry;
|
||||
|
||||
struct AnalystAgent;
|
||||
struct ReporterAgent;
|
||||
|
||||
impl Agent for AnalystAgent {
|
||||
fn name(&self) -> &str {
|
||||
"analyst"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are a data analyst. Analyze the input concisely.")
|
||||
}
|
||||
}
|
||||
|
||||
impl Agent for ReporterAgent {
|
||||
fn name(&self) -> &str {
|
||||
"reporter"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are a report writer. Write concise reports based on context.")
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_text(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: Usage::from_input_output(8, 4),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
println!("=== Agent Switch Demo ===\n");
|
||||
|
||||
// 1. 准备组件
|
||||
let store: Arc<dyn agcore::memory::store::MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
let provider = Arc::new(MockProvider::new(vec![
|
||||
assistant_text("Analyst: data analyzed (Q3 sales up 15%)"),
|
||||
assistant_text("Reporter: report drafted (3 paragraphs)"),
|
||||
]));
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.session_memory_backend(store.clone())
|
||||
.build()
|
||||
.expect("RuntimeBundle 装配失败"),
|
||||
);
|
||||
|
||||
let analyst: Arc<dyn Agent> = Arc::new(AnalystAgent);
|
||||
let reporter: Arc<dyn Agent> = Arc::new(ReporterAgent);
|
||||
|
||||
let sm = Arc::new(SessionManager::new(store));
|
||||
let session_id = sm.create(analyst, bundle.clone()).await.expect("create");
|
||||
println!("[1] session created: {session_id}");
|
||||
|
||||
// 2. Analyst 跑一轮
|
||||
let resp1 = sm
|
||||
.submit_turn(&session_id, "Analyze Q3 sales data")
|
||||
.await
|
||||
.expect("submit_turn 1");
|
||||
println!("[2] analyst turn 1: {:?}", resp1.text());
|
||||
|
||||
// 3. 切换到 Reporter
|
||||
sm.switch_agent(&session_id, reporter)
|
||||
.await
|
||||
.expect("switch_agent");
|
||||
println!("[3] agent switched to 'reporter'");
|
||||
|
||||
// 4. Reporter 跑一轮(基于已有上下文)
|
||||
let resp2 = sm
|
||||
.submit_turn(&session_id, "Write a report based on the analysis")
|
||||
.await
|
||||
.expect("submit_turn 2");
|
||||
println!("[4] reporter turn 2: {:?}", resp2.text());
|
||||
|
||||
// 5. 验证 turn_index 连续
|
||||
let (turn_index, agent_name_owned) = {
|
||||
let session = sm.get(&session_id).await.unwrap();
|
||||
let guard = session.lock().await;
|
||||
(guard.turn_index(), guard.agent.name().to_string())
|
||||
};
|
||||
println!("\n[verify] turn_index = {turn_index}, agent = {agent_name_owned}");
|
||||
assert_eq!(turn_index, 2, "turn_index should be 2 after 2 turns");
|
||||
assert_eq!(agent_name_owned, "reporter", "current agent should be reporter");
|
||||
println!("✓ context preserved across agent switch");
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
//! bridge_keys_demo —— bridge_keys 过滤 + 子↔子共享 namespace 示例。
|
||||
//!
|
||||
//! 演示:
|
||||
//! 1. 父 session 设置 SessionMemory(key: "project_goal", "constraints", "noise")
|
||||
//! 2. dispatch + bridge_keys = ["project_goal", "constraints"] → 只继承这两个
|
||||
//! 3. 验证子 session 读到的 session_memory 与过滤一致
|
||||
//! 4. 演示子↔子共享 namespace:dispatch 时设 `shared_namespace`,
|
||||
//! 子 A 写入 `shared:{parent_id}:fact_x`,子 B 通过约定 key 读取
|
||||
//!
|
||||
//! 运行:`cargo run --example bridge_keys_demo`
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::agent::{Agent, AgentBuilder};
|
||||
use agcore::engine::{DispatchConfig, SessionManager};
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::types::Usage;
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use agcore::memory::store::InMemoryStore;
|
||||
use agcore::tools::ToolRegistry;
|
||||
|
||||
struct WorkerAgent;
|
||||
|
||||
impl Agent for WorkerAgent {
|
||||
fn name(&self) -> &str {
|
||||
"worker"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are a worker.")
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_text(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: Usage::from_input_output(8, 4),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
println!("=== Bridge Keys Demo ===\n");
|
||||
|
||||
let store: Arc<dyn agcore::memory::store::MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
// 3 个 dispatch 调用需要 3 个 mock response
|
||||
let provider = Arc::new(MockProvider::new(vec![
|
||||
assistant_text("Worker 1: done"),
|
||||
assistant_text("Worker 2: done"),
|
||||
assistant_text("Worker 3: done"),
|
||||
]));
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.session_memory_backend(store.clone())
|
||||
.build()
|
||||
.expect("RuntimeBundle"),
|
||||
);
|
||||
|
||||
let worker: Arc<dyn Agent> = Arc::new(WorkerAgent);
|
||||
|
||||
let sm = Arc::new(SessionManager::new(store));
|
||||
let parent_id = sm
|
||||
.create(worker.clone(), bundle.clone())
|
||||
.await
|
||||
.expect("create");
|
||||
println!("[1] parent session created: {parent_id}");
|
||||
|
||||
// 父 session_memory 写入 3 个 key
|
||||
{
|
||||
let session = sm.get(&parent_id).await.unwrap();
|
||||
let mut guard = session.lock().await;
|
||||
guard
|
||||
.set_session_data("project_goal", "Build a fast compiler")
|
||||
.await
|
||||
.unwrap();
|
||||
guard
|
||||
.set_session_data("constraints", "Rust, no unsafe")
|
||||
.await
|
||||
.unwrap();
|
||||
guard
|
||||
.set_session_data("noise", "should NOT be inherited")
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
println!("[2] parent set 3 keys: project_goal, constraints, noise");
|
||||
|
||||
// dispatch + bridge_keys 过滤
|
||||
let config = DispatchConfig {
|
||||
bridge_keys: Some(vec!["project_goal".to_string(), "constraints".to_string()]),
|
||||
..Default::default()
|
||||
};
|
||||
let result = sm
|
||||
.dispatch(&parent_id, worker.clone(), "do work", config)
|
||||
.await
|
||||
.expect("dispatch");
|
||||
println!("[3] dispatched sub-agent (child_id={})\n", &result.child_id[..20]);
|
||||
|
||||
// 验证过滤效果
|
||||
let child_session = sm.get(&result.child_id).await.unwrap();
|
||||
let child_guard = child_session.lock().await;
|
||||
let inherited_goal = child_guard.session_memory().get("project_goal").await.unwrap();
|
||||
let inherited_constraint = child_guard.session_memory().get("constraints").await.unwrap();
|
||||
let filtered_noise = child_guard.session_memory().get("noise").await.unwrap();
|
||||
drop(child_guard);
|
||||
|
||||
println!("[verify] inherited keys in child session:");
|
||||
println!(" - project_goal: {:?}", inherited_goal);
|
||||
println!(" - constraints: {:?}", inherited_constraint);
|
||||
println!(" - noise: {:?} (should be None)", filtered_noise);
|
||||
|
||||
assert_eq!(inherited_goal, Some("Build a fast compiler".to_string()));
|
||||
assert_eq!(inherited_constraint, Some("Rust, no unsafe".to_string()));
|
||||
assert_eq!(filtered_noise, None, "noise should be filtered out");
|
||||
|
||||
println!("\n✓ bridge_keys filtering works correctly");
|
||||
|
||||
// =============== 第二部分:子↔子共享 namespace(convention)===============
|
||||
println!("\n=== Part 2: Child↔Child Shared Namespace (convention) ===\n");
|
||||
|
||||
// 关键点:`SessionMemory::get`/`set` 通过 session 自身 namespace 隔离
|
||||
// (每个 session 一个独立 namespace),所以"子↔子共享"不能直接通过 SessionMemory。
|
||||
// 真正的子↔子共享需要直接操作底层 MemoryStore,或由上层应用维护一个
|
||||
// 跨 session 的"共享通道"(例如独立的 namespace + 所有子 session 知道 key 前缀)。
|
||||
//
|
||||
// 本 demo 演示通过 DispatchConfig.shared_namespace(convention-based):
|
||||
// - `shared_namespace: Some(prefix)` 作为约定标记,告知子 agent
|
||||
// "你的数据共享 namespace 是 shared:{prefix}:*"
|
||||
// - 子 agent 自行通过 `sm.store()` 直接操作 MemoryStore(绕过 SessionMemory 的 namespace 隔离)
|
||||
//
|
||||
// 演示 2 个子 agent 通过约定 namespace prefix 共享数据。
|
||||
|
||||
let shared_ns_config = DispatchConfig {
|
||||
bridge_keys: Some(vec![]),
|
||||
shared_namespace: Some("parent-123".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
// dispatch 第一个子 agent
|
||||
let _researcher_result = sm
|
||||
.dispatch(
|
||||
&parent_id,
|
||||
worker.clone(),
|
||||
"research task",
|
||||
shared_ns_config.clone(),
|
||||
)
|
||||
.await
|
||||
.expect("dispatch researcher");
|
||||
|
||||
// 子 A 通过 `sm.store()` 直接写入共享 namespace key
|
||||
// (约定 prefix: "shared:parent-123:")
|
||||
let shared_key = "shared:parent-123:fact_architecture";
|
||||
sm.store()
|
||||
.save(agcore::memory::types::MemoryItem {
|
||||
id: shared_key.to_string(),
|
||||
content: "Microservices with event sourcing".to_string(),
|
||||
metadata: serde_json::json!({}),
|
||||
created_at: time::OffsetDateTime::now_utc(),
|
||||
})
|
||||
.await
|
||||
.expect("save shared fact");
|
||||
println!("[4] researcher wrote {shared_key}");
|
||||
|
||||
// dispatch 第二个子 agent
|
||||
let _writer_result = sm
|
||||
.dispatch(
|
||||
&parent_id,
|
||||
worker.clone(),
|
||||
"writing task",
|
||||
shared_ns_config,
|
||||
)
|
||||
.await
|
||||
.expect("dispatch writer");
|
||||
|
||||
// 子 B 通过 `sm.store()` 直接读取共享 namespace key
|
||||
let read_item = sm.store().get(shared_key).await.expect("get");
|
||||
let read_fact = read_item.map(|i| i.content);
|
||||
println!("[5] writer reads {shared_key} = {read_fact:?}");
|
||||
|
||||
assert_eq!(
|
||||
read_fact,
|
||||
Some("Microservices with event sourcing".to_string()),
|
||||
"writer should read researcher's shared fact"
|
||||
);
|
||||
|
||||
println!("\n✓ child↔child shared namespace works correctly (via MemoryStore convention)");
|
||||
println!("\n=== All Bridge Keys Demo checks passed ===");
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
//! dispatch_stream_demo —— 流式子代理调度示例。
|
||||
//!
|
||||
//! 演示:
|
||||
//! 1. 创建父 session
|
||||
//! 2. dispatch_stream 单个子 agent
|
||||
//! 3. 消费 SubTaskStreamEvent 序列
|
||||
//! 4. 验证事件序列:ChildCreated → Stream(...) × N → Completed
|
||||
//! 5. 验证:完成时 turn_index 已递增(finalize 副作用)
|
||||
//!
|
||||
//! 运行:`cargo run --example dispatch_stream_demo`
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::agent::{Agent, AgentBuilder};
|
||||
use agcore::engine::{SessionManager, SubTaskStreamEvent};
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::types::Usage;
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use agcore::memory::store::InMemoryStore;
|
||||
use agcore::tools::ToolRegistry;
|
||||
use futures_util::StreamExt;
|
||||
|
||||
struct StreamWorkerAgent;
|
||||
|
||||
impl Agent for StreamWorkerAgent {
|
||||
fn name(&self) -> &str {
|
||||
"stream_worker"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are a streaming worker.")
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_text(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: Usage::from_input_output(8, 4),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
println!("=== Dispatch Stream Demo ===\n");
|
||||
|
||||
let store: Arc<dyn agcore::memory::store::MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
let provider = Arc::new(MockProvider::new(vec![assistant_text("streamed response")]));
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.session_memory_backend(store.clone())
|
||||
.build()
|
||||
.expect("RuntimeBundle"),
|
||||
);
|
||||
|
||||
let worker: Arc<dyn Agent> = Arc::new(StreamWorkerAgent);
|
||||
|
||||
let sm = Arc::new(SessionManager::new(store));
|
||||
let parent_id = sm
|
||||
.create(worker.clone(), bundle.clone())
|
||||
.await
|
||||
.expect("create");
|
||||
println!("[1] parent session: {parent_id}");
|
||||
|
||||
// dispatch_stream
|
||||
let mut stream = sm
|
||||
.dispatch_stream(&parent_id, worker, "do streaming work", Default::default())
|
||||
.await
|
||||
.expect("dispatch_stream");
|
||||
|
||||
println!("[2] consuming SubTaskStreamEvent sequence...\n");
|
||||
let mut saw_child_created = false;
|
||||
let mut saw_stream_count = 0;
|
||||
let mut completed = None;
|
||||
|
||||
while let Some(event) = stream.next().await {
|
||||
match event {
|
||||
SubTaskStreamEvent::ChildCreated { child_id } => {
|
||||
println!(" → ChildCreated({})", &child_id[..20]);
|
||||
saw_child_created = true;
|
||||
}
|
||||
SubTaskStreamEvent::Stream(_) => {
|
||||
saw_stream_count += 1;
|
||||
}
|
||||
SubTaskStreamEvent::Completed(r) => {
|
||||
println!(" → Completed(child_id={}, {} tokens)", &r.child_id[..20], r.usage.total().total_tokens);
|
||||
completed = Some(r);
|
||||
break;
|
||||
}
|
||||
SubTaskStreamEvent::Error { child_id, error } => {
|
||||
panic!("unexpected error: child_id={child_id}, error={error}");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let result = completed.expect("Completed should arrive");
|
||||
|
||||
// 验证事件序列
|
||||
assert!(saw_child_created, "ChildCreated should be received");
|
||||
assert!(saw_stream_count > 0, "at least one Stream event");
|
||||
println!("\n[3] received {} stream events", saw_stream_count);
|
||||
|
||||
// 验证 finalize 已发生(turn_index 递增)
|
||||
let child_session = sm.get(&result.child_id).await.unwrap();
|
||||
let child_guard = child_session.lock().await;
|
||||
let child_turn_index = child_guard.turn_index();
|
||||
drop(child_guard);
|
||||
assert_eq!(child_turn_index, 1, "turn_index should increment after finalize");
|
||||
println!("[4] child session turn_index = {child_turn_index} (finalize works)");
|
||||
|
||||
println!("\n✓ dispatch_stream completed successfully");
|
||||
}
|
||||
@@ -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,252 @@
|
||||
//! engine_demo —— SessionManager + Checkpointer 端到端示例。
|
||||
//!
|
||||
//! 演示:
|
||||
//! 1. SessionManager::create 创建 session
|
||||
//! 2. SessionManager::submit_turn(auto_checkpoint=true 自动写 checkpoint)
|
||||
//! 3. SessionManager::create_child 创建子 session
|
||||
//! 4. children() / parent() 树形查询
|
||||
//! 5. Checkpointer::list_checkpoints 列出所有 checkpoint
|
||||
//! 6. SessionManager::recover 从 checkpoint 恢复(模拟进程重启)
|
||||
//! 7. AgentSession::to_snapshot + SessionManager::replace 演示 rollback 流程
|
||||
//! 8. SessionManager::destroy 清理
|
||||
//!
|
||||
//! 运行:`cargo run --example engine_demo`
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::agent::{Agent, AgentBuilder, AgentSession};
|
||||
use agcore::engine::SessionManager;
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::types::Usage;
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use agcore::memory::store::InMemoryStore;
|
||||
use agcore::tools::ToolRegistry;
|
||||
|
||||
struct DemoAgent;
|
||||
|
||||
impl Agent for DemoAgent {
|
||||
fn name(&self) -> &str {
|
||||
"demo"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("你是 demo agent,每轮回复一句话。")
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_text(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: Usage::from_input_output(8, 4),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
// 1. 准备底层组件
|
||||
let store: Arc<dyn agcore::memory::store::MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
let provider = Arc::new(MockProvider::new(vec![
|
||||
assistant_text("turn 1 response"),
|
||||
assistant_text("turn 2 response"),
|
||||
assistant_text("turn 3 response"),
|
||||
assistant_text("child turn 1 response"),
|
||||
assistant_text("recovered turn response"),
|
||||
]));
|
||||
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider.clone())
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.session_memory_backend(store.clone())
|
||||
.build()
|
||||
.expect("RuntimeBundle 装配失败"),
|
||||
);
|
||||
|
||||
let agent: Arc<dyn Agent> = Arc::new(DemoAgent);
|
||||
|
||||
// 2. 构造 SessionManager(auto_checkpoint 默认 true)
|
||||
let sm = SessionManager::new(store.clone());
|
||||
println!("=== SessionManager 创建 ===");
|
||||
|
||||
// 3. create + submit_turn(auto_checkpoint 触发)
|
||||
println!("\n=== 创建根 session + 跑 3 轮 ===");
|
||||
let parent_id = sm
|
||||
.create(agent.clone(), bundle.clone())
|
||||
.await
|
||||
.expect("create 失败");
|
||||
println!("parent_id = {parent_id}");
|
||||
|
||||
for i in 1..=3 {
|
||||
let _resp = sm
|
||||
.submit_turn(&parent_id, format!("turn {i}"))
|
||||
.await
|
||||
.expect("submit_turn 失败");
|
||||
}
|
||||
|
||||
// 写入自定义 session memory 数据(演示持久层往返)
|
||||
sm.get(&parent_id)
|
||||
.await
|
||||
.unwrap()
|
||||
.lock()
|
||||
.await
|
||||
.set_session_data("design", "PostgreSQL")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// 4. 显式 checkpoint(覆盖 auto_checkpoint 的 turn-level,写入额外快照)
|
||||
println!("\n=== Checkpointer 显式 checkpoint ===");
|
||||
let ckpt_id = sm
|
||||
.checkpointer()
|
||||
.checkpoint(&*sm.get(&parent_id).await.unwrap().lock().await)
|
||||
.await
|
||||
.expect("checkpoint 失败");
|
||||
println!("explicit ckpt_id = {ckpt_id}");
|
||||
|
||||
// 5. list_checkpoints
|
||||
let metas = sm
|
||||
.checkpointer()
|
||||
.list_checkpoints(&parent_id)
|
||||
.await
|
||||
.expect("list_checkpoints 失败");
|
||||
println!("parent session 有 {} 个 checkpoint", metas.len());
|
||||
for m in &metas {
|
||||
println!(
|
||||
" - ckpt_id={}, turn_index={}, created_at={}",
|
||||
m.ckpt_id, m.turn_index, m.created_at
|
||||
);
|
||||
}
|
||||
|
||||
// 6. create_child
|
||||
println!("\n=== 创建子 session ===");
|
||||
let child_id = sm
|
||||
.create_child(&parent_id, agent.clone())
|
||||
.await
|
||||
.expect("create_child 失败");
|
||||
println!("child_id = {child_id}");
|
||||
|
||||
let children = sm.children(&parent_id).await.expect("children 失败");
|
||||
assert_eq!(children, vec![child_id.clone()]);
|
||||
println!("children(parent) = {children:?}");
|
||||
|
||||
let parent_of_child = sm.parent(&child_id).await.expect("parent 失败");
|
||||
assert_eq!(parent_of_child, Some(parent_id.clone()));
|
||||
println!("parent({child_id}) = {parent_of_child:?}");
|
||||
|
||||
// 7. recover(模拟"进程重启"——新建 SessionManager 实例,但 store 复用)
|
||||
println!("\n=== 从存储恢复 session(模拟进程重启)===");
|
||||
let sm2 = SessionManager::new(store.clone());
|
||||
let recovered = sm2
|
||||
.recover(&parent_id, agent.clone(), bundle.clone())
|
||||
.await
|
||||
.expect("recover 失败");
|
||||
let recovered_session = recovered.lock().await;
|
||||
let v = recovered_session
|
||||
.get_session_data("design")
|
||||
.await
|
||||
.expect("get_session_data 失败");
|
||||
println!("recovered session_memory['design'] = {v:?}");
|
||||
assert_eq!(v, Some("PostgreSQL".into()));
|
||||
|
||||
// 8. 演示 rollback 流程:先记录当前 turn_index,再 rollback 到一个早期 checkpoint,
|
||||
// 验证 session_memory 和 turn_index 已恢复到 checkpoint 时刻
|
||||
println!("\n=== rollback 流程 ===");
|
||||
let (before_turn, before_cost) = {
|
||||
let s = sm.get(&parent_id).await.unwrap();
|
||||
let g = s.lock().await;
|
||||
(g.turn_index(), g.usage().total().total_tokens)
|
||||
};
|
||||
println!(
|
||||
"rollback 前 turn_index={}, total_tokens={}",
|
||||
before_turn, before_cost
|
||||
);
|
||||
|
||||
let metas = sm
|
||||
.checkpointer()
|
||||
.list_checkpoints(&parent_id)
|
||||
.await
|
||||
.expect("list_checkpoints 失败");
|
||||
assert!(metas.len() >= 2, "至少 2 个 checkpoint 才能演示 rollback");
|
||||
// 取第二个 checkpoint(不是最新的)作为 rollback 目标
|
||||
let rollback_ckpt = &metas[metas.len() - 2];
|
||||
println!("rollback 到 ckpt_id={}", rollback_ckpt.ckpt_id);
|
||||
|
||||
let snapshot = sm
|
||||
.checkpointer()
|
||||
.rollback_load(&parent_id, &rollback_ckpt.ckpt_id)
|
||||
.await
|
||||
.expect("rollback_load 失败");
|
||||
let snapshot_turn = snapshot.turn_index;
|
||||
let snapshot_data_count = snapshot.session_memory_data.len();
|
||||
println!(
|
||||
"checkpoint 时刻 turn_index={}, session_memory 条目数={}",
|
||||
snapshot_turn, snapshot_data_count
|
||||
);
|
||||
|
||||
let mut rolled_back =
|
||||
AgentSession::from_snapshot(snapshot, agent.clone(), bundle.clone()).expect("from_snapshot");
|
||||
rolled_back
|
||||
.restore_memory()
|
||||
.await
|
||||
.expect("restore_memory 失败");
|
||||
sm.replace(&parent_id, rolled_back)
|
||||
.await
|
||||
.expect("replace 失败");
|
||||
|
||||
// 验证 rollback 后状态与 checkpoint 一致
|
||||
let (after_turn, after_cost) = {
|
||||
let s = sm.get(&parent_id).await.unwrap();
|
||||
let g = s.lock().await;
|
||||
(g.turn_index(), g.usage().total().total_tokens)
|
||||
};
|
||||
println!(
|
||||
"rollback 后 turn_index={}, total_tokens={}",
|
||||
after_turn, after_cost
|
||||
);
|
||||
assert!(
|
||||
after_turn <= before_turn,
|
||||
"rollback 后 turn_index({}) 应 ≤ rollback 前({})",
|
||||
after_turn,
|
||||
before_turn
|
||||
);
|
||||
assert_eq!(after_turn, snapshot_turn, "rollback 后 turn_index 应等于 checkpoint 时刻值");
|
||||
assert!(after_cost <= before_cost, "rollback 后 cost 应 ≤ rollback 前");
|
||||
println!("✓ rollback + replace 一致性验证通过");
|
||||
|
||||
// 9. destroy 父子 session
|
||||
println!("\n=== 销毁 session ===");
|
||||
sm.destroy(&child_id).await.expect("destroy child 失败");
|
||||
sm.destroy(&parent_id).await.expect("destroy parent 失败");
|
||||
|
||||
// 验证清理
|
||||
assert!(sm.get(&parent_id).await.is_err());
|
||||
assert!(
|
||||
sm.checkpointer()
|
||||
.list_checkpoints(&parent_id)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
println!("✓ parent 已彻底清理(内存 + meta + checkpoints)");
|
||||
|
||||
// 验证孤儿语义:父被销毁后子仍存在但 parent() 返回 None
|
||||
println!("\n=== 孤儿策略演示(先创建父子,再仅销毁父)===");
|
||||
let p_id = sm.create(agent.clone(), bundle.clone()).await.unwrap();
|
||||
let c_id = sm.create_child(&p_id, agent.clone()).await.unwrap();
|
||||
sm.destroy(&p_id).await.unwrap();
|
||||
let p_of_c = sm.parent(&c_id).await.expect("parent 失败");
|
||||
assert_eq!(p_of_c, None, "父被销毁后 child.parent() 应为 None");
|
||||
println!("✓ child({c_id}) 仍是孤儿 session,parent() = None");
|
||||
|
||||
// 清理孤儿
|
||||
sm.destroy(&c_id).await.unwrap();
|
||||
|
||||
println!("\n✓ engine_demo 完成");
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
//! knowledge_graph_demo -- 知识图谱 + 双通道检索演示。
|
||||
//!
|
||||
//! 演示:
|
||||
//! 1. 构建 KnowledgeGraph(实体 + 关系)
|
||||
//! 2. BFS 图遍历(get_related)
|
||||
//! 3. MemoryRetriever 双通道检索(Hybrid / GraphOnly / KnowledgeOnly)
|
||||
//! 4. 标签管理(set_entity_tags / find_tags)
|
||||
//!
|
||||
//! 运行:`cargo run --example knowledge_graph_demo`
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::memory::{
|
||||
GraphEntity, GraphRelation, InMemoryGraph, InMemoryStore, KnowledgeGraph, KnowledgePage,
|
||||
KnowledgeStore, MemoryRetriever, MemoryStore, RelationDirection, RetrievalItem,
|
||||
RetrievalStrategy, RetrieverConfig,
|
||||
};
|
||||
use time::OffsetDateTime;
|
||||
|
||||
fn make_page(id: &str, title: &str, content: &str) -> KnowledgePage {
|
||||
let now = OffsetDateTime::now_utc();
|
||||
KnowledgePage {
|
||||
id: id.to_string(),
|
||||
title: title.to_string(),
|
||||
summary: content.chars().take(40).collect(),
|
||||
content: content.to_string(),
|
||||
tags: Vec::new(),
|
||||
references: Vec::new(),
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
// ── 1. 构建知识图谱 ──
|
||||
println!("=== 1. 构建知识图谱 ===");
|
||||
let graph = Arc::new(InMemoryGraph::new());
|
||||
|
||||
let mut langchain = GraphEntity::new("langchain", "LangChain", "framework");
|
||||
langchain.description = "LLM application framework".to_string();
|
||||
let mut langgraph = GraphEntity::new("langgraph", "LangGraph", "framework");
|
||||
langgraph.description = "Graph-based agent runtime from LangChain".to_string();
|
||||
let mut langsmith = GraphEntity::new("langsmith", "LangSmith", "tool");
|
||||
langsmith.description = "Tracing and evaluation platform".to_string();
|
||||
let mut python = GraphEntity::new("python", "Python", "language");
|
||||
python.description = "Programming language".to_string();
|
||||
let mut rust = GraphEntity::new("rust", "Rust", "language");
|
||||
rust.description = "Systems programming language".to_string();
|
||||
|
||||
for e in [&langchain, &langgraph, &langsmith, &python, &rust] {
|
||||
graph.add_entity(e.clone()).await.unwrap();
|
||||
}
|
||||
graph
|
||||
.add_relation(GraphRelation::new("langchain", "langgraph", "includes", 0.9))
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
.add_relation(GraphRelation::new("langchain", "langsmith", "includes", 0.7))
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
.add_relation(GraphRelation::new("langchain", "python", "built_with", 0.95))
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
.add_relation(GraphRelation::new("langgraph", "python", "depends_on", 0.8))
|
||||
.await
|
||||
.unwrap();
|
||||
println!("已添加 5 个实体 + 4 条关系");
|
||||
|
||||
// ── 2. BFS 图遍历 ──
|
||||
println!("\n=== 2. BFS 图遍历:从 LangChain 出发,depth=2 ===");
|
||||
let related = graph
|
||||
.get_related("langchain", 2, RelationDirection::Outgoing, None)
|
||||
.await
|
||||
.unwrap();
|
||||
for se in &related {
|
||||
println!(
|
||||
" {} (score={:.3}, path={:?})",
|
||||
se.entity.name, se.score, se.path
|
||||
);
|
||||
}
|
||||
assert!(!related.is_empty(), "应找到关联实体");
|
||||
|
||||
// ── 3. 标签管理 ──
|
||||
println!("\n=== 3. 标签管理 ===");
|
||||
graph
|
||||
.set_entity_tags("langchain", vec!["ai".into(), "framework".into(), "llm".into()])
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
.set_entity_tags("langgraph", vec!["ai".into(), "agent".into()])
|
||||
.await
|
||||
.unwrap();
|
||||
let tags = graph.find_tags("a").await.unwrap();
|
||||
println!("前缀 'a' 查找标签: {:?}", tags);
|
||||
let count = graph.entity_count_by_tag("ai").await.unwrap();
|
||||
println!("标签 'ai' 下实体数: {}", count);
|
||||
|
||||
// ── 4. 双通道检索 ──
|
||||
println!("\n=== 4. 双通道检索 ===");
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
let ks = KnowledgeStore::new(store);
|
||||
ks.add_page(make_page(
|
||||
"p1",
|
||||
"LangChain 框架介绍",
|
||||
"LangChain 是用于构建 LLM 应用的开源框架",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
ks.add_page(make_page(
|
||||
"p2",
|
||||
"Rust 异步编程",
|
||||
"Rust 异步基于 tokio 与 futures 抽象",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Hybrid 策略(默认)
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph.clone());
|
||||
println!("\n--- Hybrid 检索: 'langchain' ---");
|
||||
let result = retriever.retrieve("langchain").await.unwrap();
|
||||
println!("策略: {:?}", result.strategy);
|
||||
for item in &result.items {
|
||||
match item {
|
||||
RetrievalItem::KnowledgePage { page, score } => {
|
||||
println!(" [Store] {} (score={:.3})", page.title, score);
|
||||
}
|
||||
RetrievalItem::GraphEntity {
|
||||
entity, score, path, ..
|
||||
} => {
|
||||
println!(
|
||||
" [Graph] {} (score={:.3}, path={:?})",
|
||||
entity.name, score, path
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
let has_store = result
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::KnowledgePage { .. }));
|
||||
let has_graph = result
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::GraphEntity { .. }));
|
||||
assert!(has_store, "Hybrid 应有 Store 结果");
|
||||
assert!(has_graph, "Hybrid 应有 Graph 结果");
|
||||
|
||||
// GraphOnly 策略
|
||||
let store2: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
let ks2 = KnowledgeStore::new(store2);
|
||||
ks2.add_page(make_page("p1", "LangChain", "LLM framework"))
|
||||
.await
|
||||
.unwrap();
|
||||
let retriever_g = MemoryRetriever::new(ks2, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph.clone())
|
||||
.with_strategy(RetrievalStrategy::GraphOnly);
|
||||
println!("\n--- GraphOnly 检索: 'langchain' ---");
|
||||
let result = retriever_g.retrieve("langchain").await.unwrap();
|
||||
println!("策略: {:?}", result.strategy);
|
||||
for item in &result.items {
|
||||
match item {
|
||||
RetrievalItem::GraphEntity { entity, score, .. } => {
|
||||
println!(" [Graph] {} (score={:.3})", entity.name, score);
|
||||
}
|
||||
RetrievalItem::KnowledgePage { page, score, .. } => {
|
||||
println!(" [Store] {} (score={:.3})", page.title, score);
|
||||
}
|
||||
}
|
||||
}
|
||||
assert!(
|
||||
result.items.iter().all(|i| matches!(i, RetrievalItem::GraphEntity { .. })),
|
||||
"GraphOnly 应只返回 Graph 结果"
|
||||
);
|
||||
|
||||
println!("\n✓ knowledge_graph_demo 完成");
|
||||
}
|
||||
@@ -74,8 +74,15 @@ async fn main() {
|
||||
let result = retriever.retrieve("Rust 异步").await.unwrap();
|
||||
println!("query: {}", result.query);
|
||||
for item in &result.items {
|
||||
println!(" 命中: {} (score={:.3})", item.page.title, item.score);
|
||||
assert!((0.0..=1.0).contains(&item.score), "score 应在 [0, 1] 区间");
|
||||
match item {
|
||||
agcore::memory::RetrievalItem::KnowledgePage { page, score } => {
|
||||
println!(" 命中: {} (score={:.3})", page.title, score);
|
||||
assert!((0.0..=1.0).contains(score), "score 应在 [0, 1] 区间");
|
||||
}
|
||||
agcore::memory::RetrievalItem::GraphEntity { entity, score, .. } => {
|
||||
println!(" 命中实体: {} (score={:.3})", entity.name, score);
|
||||
}
|
||||
}
|
||||
}
|
||||
assert!(!result.items.is_empty(), "应至少命中一个页面");
|
||||
|
||||
@@ -89,6 +96,7 @@ async fn main() {
|
||||
let cfg = RetrieverConfig {
|
||||
max_results: 20,
|
||||
min_score: 0.5,
|
||||
graph_depth: 2,
|
||||
};
|
||||
let retriever2 = MemoryRetriever::new(ks2, cfg);
|
||||
let result = retriever2.retrieve("完全不相关的火锅配方").await.unwrap();
|
||||
@@ -111,6 +119,7 @@ async fn main() {
|
||||
let cfg = RetrieverConfig {
|
||||
max_results: 2,
|
||||
min_score: 0.0,
|
||||
graph_depth: 2,
|
||||
};
|
||||
let retriever3 = MemoryRetriever::new(ks3, cfg);
|
||||
let result = retriever3.retrieve("Rust").await.unwrap();
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
//! sub_agent_dispatch_demo —— SubAgent 并行派发示例。
|
||||
//!
|
||||
//! 演示:
|
||||
//! 1. 创建父 session("主编" agent)
|
||||
//! 2. 并行 dispatch_all 3 个子 agent(研究员 / 写手 / 审校)
|
||||
//! 3. 收集子任务结果
|
||||
//! 4. 验证树形结构:children(parent_id) 应返回 3 个子 ID
|
||||
//!
|
||||
//! 运行:`cargo run --example sub_agent_dispatch_demo`
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use agcore::agent::{Agent, AgentBuilder};
|
||||
use agcore::engine::{SessionManager, SubTaskResult};
|
||||
use agcore::llm::hooks::HookExecutor;
|
||||
use agcore::llm::mock::MockProvider;
|
||||
use agcore::llm::types::Usage;
|
||||
use agcore::llm::types::message::{ContentBlock, Message};
|
||||
use agcore::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use agcore::memory::store::InMemoryStore;
|
||||
use agcore::tools::ToolRegistry;
|
||||
|
||||
struct EditorAgent;
|
||||
struct ResearcherAgent;
|
||||
struct WriterAgent;
|
||||
struct ReviewerAgent;
|
||||
|
||||
impl Agent for EditorAgent {
|
||||
fn name(&self) -> &str {
|
||||
"editor"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are an editor coordinating a team.")
|
||||
}
|
||||
}
|
||||
impl Agent for ResearcherAgent {
|
||||
fn name(&self) -> &str {
|
||||
"researcher"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are a researcher. Provide 3 key findings.")
|
||||
}
|
||||
}
|
||||
impl Agent for WriterAgent {
|
||||
fn name(&self) -> &str {
|
||||
"writer"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are a writer. Draft a section.")
|
||||
}
|
||||
}
|
||||
impl Agent for ReviewerAgent {
|
||||
fn name(&self) -> &str {
|
||||
"reviewer"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some("You are a reviewer. Check for accuracy.")
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_text(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: Usage::from_input_output(8, 4),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
fn print_result(name: &str, r: &Result<SubTaskResult, agcore::engine::EngineError>) {
|
||||
match r {
|
||||
Ok(res) => println!(
|
||||
" ✓ {name} (child_id={}): {} tokens",
|
||||
&res.child_id[..20.min(res.child_id.len())],
|
||||
res.usage.total().total_tokens,
|
||||
),
|
||||
Err(e) => println!(" ✗ {name}: {e}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
println!("=== SubAgent Dispatch Demo ===\n");
|
||||
|
||||
let store: Arc<dyn agcore::memory::store::MemoryStore> = Arc::new(InMemoryStore::new());
|
||||
let provider = Arc::new(MockProvider::new(vec![
|
||||
assistant_text("Researcher: finding 1, 2, 3"),
|
||||
assistant_text("Writer: section drafted"),
|
||||
assistant_text("Reviewer: looks good"),
|
||||
]));
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.session_memory_backend(store.clone())
|
||||
.build()
|
||||
.expect("RuntimeBundle"),
|
||||
);
|
||||
|
||||
let editor: Arc<dyn Agent> = Arc::new(EditorAgent);
|
||||
let researcher: Arc<dyn Agent> = Arc::new(ResearcherAgent);
|
||||
let writer: Arc<dyn Agent> = Arc::new(WriterAgent);
|
||||
let reviewer: Arc<dyn Agent> = Arc::new(ReviewerAgent);
|
||||
|
||||
let sm = Arc::new(SessionManager::new(store));
|
||||
let parent_id = sm.create(editor, bundle.clone()).await.expect("create");
|
||||
println!("[1] parent session created: {parent_id}");
|
||||
|
||||
// dispatch_all 3 个子 agent
|
||||
println!("[2] dispatching 3 sub-agents in parallel...\n");
|
||||
let results = sm
|
||||
.dispatch_all(
|
||||
&parent_id,
|
||||
vec![
|
||||
(researcher, "Research topic X".to_string()),
|
||||
(writer, "Draft intro section".to_string()),
|
||||
(reviewer, "Review draft".to_string()),
|
||||
],
|
||||
agcore::engine::DispatchConfig::default(),
|
||||
)
|
||||
.await;
|
||||
|
||||
print_result("researcher", &results[0]);
|
||||
print_result("writer", &results[1]);
|
||||
print_result("reviewer", &results[2]);
|
||||
|
||||
let success_count = results.iter().filter(|r| r.is_ok()).count();
|
||||
assert_eq!(success_count, 3, "all 3 should succeed");
|
||||
|
||||
// 验证树形
|
||||
let children = sm.children(&parent_id).await.expect("children");
|
||||
println!("\n[3] children(parent) = {} session(s)", children.len());
|
||||
assert_eq!(children.len(), 3);
|
||||
|
||||
println!("\n✓ dispatch_all completed: 3/3 sub-agents succeeded");
|
||||
}
|
||||
+4
-2
@@ -16,18 +16,20 @@ pub mod error;
|
||||
pub mod runtime;
|
||||
pub mod session;
|
||||
pub mod session_memory;
|
||||
pub mod summary;
|
||||
pub mod task;
|
||||
|
||||
// 重导出公共 API(按使用频度排序)
|
||||
pub use agent::Agent;
|
||||
pub use builder::AgentBuilder;
|
||||
pub use context::{
|
||||
ContextBudget, ContextSlot, DeriveStrategy, FocusedConfig, SlotConfig, SlotMeta, SlotMode,
|
||||
SlotSource,
|
||||
ContextBudget, ContextSlot, DeriveStrategy, FocusedConfig, MergeStrategy, SlotConfig,
|
||||
SlotMeta, SlotMode, SlotSource,
|
||||
};
|
||||
pub use error::AgentError;
|
||||
pub use runtime::{AgentConfig, RuntimeBundle};
|
||||
pub use session::AgentSession;
|
||||
pub use session_memory::SessionMemory;
|
||||
pub use summary::SummaryConfig;
|
||||
pub use task::JsonPlanParser;
|
||||
pub use task::{Plan, PlanParser, Step, StepStatus, TaskAgent};
|
||||
|
||||
@@ -11,6 +11,7 @@ use std::sync::Arc;
|
||||
|
||||
use crate::agent::error::AgentError;
|
||||
use crate::agent::runtime::{AgentConfig, RuntimeBundle};
|
||||
use crate::agent::summary::SummaryConfig;
|
||||
use crate::llm::hooks::HookExecutor;
|
||||
use crate::llm::provider::LlmProvider;
|
||||
use crate::memory::retriever::MemoryRetriever;
|
||||
@@ -86,6 +87,15 @@ impl AgentBuilder {
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置摘要自动生成配置(覆盖字段,而非整体覆盖 config)。
|
||||
/// 不传则沿用现有 `config.summary_config`(默认 `None`,即关闭)。
|
||||
pub fn summary_config(mut self, cfg: SummaryConfig) -> Self {
|
||||
let mut config = self.config.take().unwrap_or_default();
|
||||
config.summary_config = Some(cfg);
|
||||
self.config = Some(config);
|
||||
self
|
||||
}
|
||||
|
||||
/// 构造 `RuntimeBundle`,校验必填字段。
|
||||
///
|
||||
/// **错误**:`provider` / `tool_registry` / `hook_executor` 任一缺失则返回
|
||||
|
||||
+227
-3
@@ -17,7 +17,7 @@ use crate::memory::store::MemoryStore;
|
||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||
|
||||
/// 上下文槽 —— 一段带策略配置的消息列表。
|
||||
#[derive(Debug, Clone)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ContextSlot {
|
||||
/// 当前 slot 的唯一标识(同一个 session_id 内唯一)。
|
||||
pub id: String,
|
||||
@@ -74,8 +74,9 @@ pub struct FocusedConfig {
|
||||
pub keep_system: bool,
|
||||
/// 保留的最近消息条数(以消息条数而非对话轮次为单位,因为一轮对话可能包含多条 tool 消息)。
|
||||
pub recent_messages: usize,
|
||||
/// 摘要覆盖(v0.2 仅消费端:手动设置则注入,不自动生成)。
|
||||
/// v0.3 将支持 Hook 驱动的自动摘要生成。
|
||||
/// 摘要覆盖(消费端:手动或自动生成的摘要会注入到消息列表末尾)。
|
||||
/// v0.3 Phase 16 起,`AgentBuilder::summary_config(cfg)` 内联检查点会
|
||||
/// 自动调用 LLM 生成摘要并写入此字段,详见 `docs/22-phase16-summary-auto-generation.md`。
|
||||
pub summary_override: Option<String>,
|
||||
}
|
||||
|
||||
@@ -103,6 +104,18 @@ pub enum DeriveStrategy {
|
||||
Focused(FocusedConfig),
|
||||
}
|
||||
|
||||
/// 合并策略 —— Phase 13 新增,控制 `ContextSlot::merge` 如何将子 slot 消息合入父 slot。
|
||||
///
|
||||
/// `#[non_exhaustive]` 预留未来扩展(如 `Summarize` 变体)。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[non_exhaustive]
|
||||
pub enum MergeStrategy {
|
||||
/// 子 slot 消息追加到父 slot 末尾。
|
||||
Append,
|
||||
/// 用子 slot 消息替换父 slot 内容。
|
||||
Replace,
|
||||
}
|
||||
|
||||
/// 上下文预算(v0.2 纯数据结构,无消费逻辑)。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ContextBudget {
|
||||
@@ -422,6 +435,82 @@ impl ContextSlot {
|
||||
_ => self.messages.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 从当前 slot 派生出独立的子 slot(不持久化,调用方负责 `save`)。
|
||||
///
|
||||
/// 子 slot 的 `meta` 全新创建(`SlotMeta::new()`),不继承父 slot 的 `message_count`。
|
||||
/// 子 slot 的 `source` 标记为 `Derived { parent_id, strategy }`,血缘可追溯。
|
||||
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 的消息合并到当前 slot。
|
||||
///
|
||||
/// **注意**:本方法仅操作内存数据,不自动持久化。
|
||||
/// 调用方需在 merge 后自行调用 `self.save(&store)` 将结果写入后端存储。
|
||||
///
|
||||
/// 防御性检查:
|
||||
/// - 禁止 self-merge(`self.id == child.id`)
|
||||
/// - 禁止跨 session merge
|
||||
/// - 禁止合并到 Readonly slot
|
||||
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(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -870,6 +959,141 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
// ====== Phase 13: fork/merge ======
|
||||
|
||||
#[test]
|
||||
fn fork_full_copies_messages() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::user_text("a")]).unwrap();
|
||||
parent.append_messages(vec![Message::assistant("b")]).unwrap();
|
||||
let child = parent.fork("c".into(), DeriveStrategy::Full);
|
||||
assert_eq!(child.messages.len(), 2);
|
||||
assert!(matches!(child.config.mode, SlotMode::Full));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn fork_focused_filters_messages() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::system("sys")]).unwrap();
|
||||
for i in 0..5 {
|
||||
parent
|
||||
.append_messages(vec![Message::user_text(format!("u{i}"))])
|
||||
.unwrap();
|
||||
}
|
||||
let cfg = FocusedConfig {
|
||||
keep_system: true,
|
||||
recent_messages: 2,
|
||||
summary_override: None,
|
||||
};
|
||||
let child = parent.fork("c".into(), DeriveStrategy::Focused(cfg));
|
||||
// system + 最近 2 条非 system
|
||||
assert_eq!(child.messages.len(), 1 + 2);
|
||||
assert!(matches!(child.config.mode, SlotMode::Focused(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn fork_preserves_independence() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::user_text("a")]).unwrap();
|
||||
let mut child = parent.fork("c".into(), DeriveStrategy::Full);
|
||||
let child_count_at_fork = child.messages.len();
|
||||
|
||||
// 父 slot 追加
|
||||
parent
|
||||
.append_messages(vec![Message::user_text("b")])
|
||||
.unwrap();
|
||||
// 子 slot 追加
|
||||
child.append_messages(vec![Message::user_text("c")]).unwrap();
|
||||
|
||||
assert_eq!(parent.messages.len(), 2);
|
||||
assert_eq!(child.messages.len(), child_count_at_fork + 1);
|
||||
assert_eq!(extract_text(child.messages.last().unwrap()), "c");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn fork_sets_derived_source() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::user_text("a")]).unwrap();
|
||||
let child = parent.fork("c".into(), DeriveStrategy::Full);
|
||||
match &child.config.source {
|
||||
SlotSource::Derived { parent_id, strategy } => {
|
||||
assert_eq!(parent_id, "p");
|
||||
assert!(matches!(strategy, DeriveStrategy::Full));
|
||||
}
|
||||
_ => panic!("子 slot source 应为 Derived"),
|
||||
}
|
||||
// 子 slot 的 meta 全新创建
|
||||
assert_eq!(child.meta.message_count, 0);
|
||||
assert!(child.meta.parent_id.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_append_appends_messages() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::user_text("p1")]).unwrap();
|
||||
let child = {
|
||||
let mut c = parent.fork("c".into(), DeriveStrategy::Full);
|
||||
// fork 时 child 继承父的 "p1";再追加一条 c1
|
||||
c.append_messages(vec![Message::user_text("c1")]).unwrap();
|
||||
c
|
||||
};
|
||||
parent
|
||||
.merge(child, MergeStrategy::Append)
|
||||
.expect("merge ok");
|
||||
// Append 追加 child 全部消息到父:1 (p1) + 2 (p1 + c1) = 3
|
||||
assert_eq!(parent.messages.len(), 3);
|
||||
assert_eq!(parent.meta.message_count, 3);
|
||||
assert_eq!(extract_text(&parent.messages[2]), "c1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_replace_replaces_messages() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::user_text("p1")]).unwrap();
|
||||
parent.append_messages(vec![Message::user_text("p2")]).unwrap();
|
||||
let child = {
|
||||
let mut c = parent.fork("c".into(), DeriveStrategy::Full);
|
||||
// 清空 child 再追加
|
||||
c.messages.clear();
|
||||
c.append_messages(vec![Message::user_text("c-only")])
|
||||
.unwrap();
|
||||
c
|
||||
};
|
||||
parent
|
||||
.merge(child, MergeStrategy::Replace)
|
||||
.expect("merge ok");
|
||||
assert_eq!(parent.messages.len(), 1);
|
||||
assert_eq!(extract_text(&parent.messages[0]), "c-only");
|
||||
assert_eq!(parent.meta.message_count, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_self_rejected() {
|
||||
let mut slot = make_slot("p", "s1");
|
||||
slot.append_messages(vec![Message::user_text("a")]).unwrap();
|
||||
let child = slot.fork("p".into(), DeriveStrategy::Full);
|
||||
let err = slot.merge(child, MergeStrategy::Append).unwrap_err();
|
||||
assert!(matches!(err, AgentError::Config(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_readonly_rejected() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.config.mode = SlotMode::Readonly;
|
||||
let child = ContextSlot::new("s1", "c", SlotConfig::default());
|
||||
let err = parent.merge(child, MergeStrategy::Append).unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotReadonly(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_cross_session_rejected() {
|
||||
let mut parent = make_slot("p", "s1");
|
||||
parent.append_messages(vec![Message::user_text("a")]).unwrap();
|
||||
let child = ContextSlot::new("OTHER_SESSION", "c", SlotConfig::default());
|
||||
let err = parent.merge(child, MergeStrategy::Append).unwrap_err();
|
||||
assert!(matches!(err, AgentError::Config(_)));
|
||||
}
|
||||
|
||||
// ====== Colon 校验(key 格式保护) ======
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -15,6 +15,7 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::agent::summary::SummaryConfig;
|
||||
use crate::llm::compact::CompactConfig;
|
||||
use crate::llm::hooks::HookExecutor;
|
||||
use crate::llm::provider::LlmProvider;
|
||||
@@ -33,6 +34,10 @@ pub struct AgentConfig {
|
||||
pub session_ttl: Option<Duration>,
|
||||
/// 上下文压缩配置(None 表示不启用自动压缩),默认 None。
|
||||
pub compact_config: Option<CompactConfig>,
|
||||
/// 摘要自动生成配置(`None` = 不启用)。
|
||||
/// 设置后 `AgentSession` 每轮 OnTurnEnd 之后进行水位 + 防抖检查,触发时调 LLM
|
||||
/// 生成摘要并写入 `FocusedConfig.summary_override` 与 `SessionMemory["conversation_summary"]`。
|
||||
pub summary_config: Option<SummaryConfig>,
|
||||
}
|
||||
|
||||
impl Default for AgentConfig {
|
||||
@@ -42,6 +47,7 @@ impl Default for AgentConfig {
|
||||
max_tool_turns: 10,
|
||||
session_ttl: None,
|
||||
compact_config: None,
|
||||
summary_config: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+692
-36
@@ -16,13 +16,21 @@ use futures_core::Stream;
|
||||
|
||||
use crate::agent::agent::Agent;
|
||||
use crate::agent::context::{
|
||||
ContextSlot, DeriveStrategy, SlotConfig, SlotMode, SlotSource,
|
||||
ContextSlot, DeriveStrategy, SlotConfig, SlotMode,
|
||||
};
|
||||
// SlotSource 仅在 `mod tests` 中使用(通过 `use super::*;` 引入),lib 主体保留以避免测试 import 变更。
|
||||
#[allow(unused_imports)]
|
||||
use crate::agent::context::SlotSource;
|
||||
use crate::agent::error::AgentError;
|
||||
use crate::agent::runtime::RuntimeBundle;
|
||||
use crate::agent::session_memory::SessionMemory;
|
||||
use crate::agent::summary::{format_messages_as_text, SummaryConfig};
|
||||
use crate::engine::snapshot::{SessionMemoryEntry, SessionSnapshot};
|
||||
use crate::engine::EngineError;
|
||||
use crate::llm::cycle::{CostTracker, CycleConfig, LlmCycle};
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::hooks::{HookContext, HookEvent};
|
||||
use crate::llm::provider::LlmProvider;
|
||||
use crate::llm::stream::StreamEvent;
|
||||
use crate::llm::types::message::Message;
|
||||
use crate::llm::types::response_v2::MessageResponse;
|
||||
@@ -52,6 +60,14 @@ pub struct AgentSession {
|
||||
slots: HashMap<String, ContextSlot>,
|
||||
/// Phase 10 新增:当前活跃 slot 的 id。
|
||||
current_slot_id: String,
|
||||
/// Phase 16 新增:上次摘要生成时的 `turn_index`(用于 `debounce_turns` 防抖)。
|
||||
/// `None` 表示从未生成过摘要(首次触发不受防抖约束)。
|
||||
last_summary_turn: Option<u32>,
|
||||
/// Phase 17 新增:`from_snapshot()` 后暂存的待写回条目。
|
||||
/// `None` 表示无 pending restore(正常状态)。
|
||||
/// 调用 `restore_memory()` 后会被消费并设为 `None`。
|
||||
/// 这是 transient state,不参与序列化(AgentSession 本身不 derive Serialize)。
|
||||
pending_memory_restore: Option<HashMap<String, SessionMemoryEntry>>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for AgentSession {
|
||||
@@ -110,6 +126,8 @@ impl AgentSession {
|
||||
session_memory,
|
||||
slots,
|
||||
current_slot_id: "default".to_string(),
|
||||
last_summary_turn: None,
|
||||
pending_memory_restore: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -128,6 +146,11 @@ impl AgentSession {
|
||||
&self.session_memory
|
||||
}
|
||||
|
||||
/// RuntimeBundle 引用(Phase 17 新增,供 SessionManager::create_child 继承父 bundle)。
|
||||
pub(crate) fn bundle(&self) -> &Arc<RuntimeBundle> {
|
||||
&self.bundle
|
||||
}
|
||||
|
||||
/// 写入一条会话级数据(覆盖同名 key)。
|
||||
pub async fn set_session_data(
|
||||
&mut self,
|
||||
@@ -224,38 +247,9 @@ impl AgentSession {
|
||||
.slots
|
||||
.get(parent_id)
|
||||
.ok_or_else(|| AgentError::SlotNotFound(parent_id.to_string()))?;
|
||||
|
||||
let parent_messages = parent.messages.clone();
|
||||
let (messages, focused_cfg) = match &strategy {
|
||||
DeriveStrategy::Full => (parent_messages, None),
|
||||
DeriveStrategy::Focused(cfg) => {
|
||||
let filtered = ContextSlot::filter_focused(&parent_messages, cfg);
|
||||
(filtered, Some(cfg.clone()))
|
||||
}
|
||||
};
|
||||
|
||||
let mode = match focused_cfg {
|
||||
Some(cfg) => SlotMode::Focused(cfg),
|
||||
None => SlotMode::Full,
|
||||
};
|
||||
|
||||
let slot = ContextSlot::new(
|
||||
&self.session_id,
|
||||
&slot_id,
|
||||
SlotConfig {
|
||||
mode,
|
||||
source: SlotSource::Derived {
|
||||
parent_id: parent_id.to_string(),
|
||||
strategy,
|
||||
},
|
||||
budget: Default::default(),
|
||||
compact: true,
|
||||
},
|
||||
);
|
||||
let mut slot = slot;
|
||||
slot.messages = messages;
|
||||
slot.save(&*self.resolve_store()).await?;
|
||||
self.slots.insert(slot_id, slot);
|
||||
let child = parent.fork(slot_id.clone(), strategy);
|
||||
child.save(&*self.resolve_store()).await?;
|
||||
self.slots.insert(slot_id, child);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -371,6 +365,9 @@ impl AgentSession {
|
||||
let end_ctx = HookContext::new(HookEvent::OnTurnEnd).with_turn_index(turn_index);
|
||||
hook_executor.execute(HookEvent::OnTurnEnd, &end_ctx).await;
|
||||
|
||||
// 7.5 Phase 16: 摘要自动生成检查点
|
||||
self.maybe_summarize(turn_index).await;
|
||||
|
||||
// 8. turn_index 递增
|
||||
self.turn_index += 1;
|
||||
|
||||
@@ -444,9 +441,6 @@ impl AgentSession {
|
||||
// 6. turn_index 递增 —— 配合 finalize_turn 用 (turn_index - 1) 传递正确的 OnTurnEnd 序号
|
||||
self.turn_index += 1;
|
||||
|
||||
// 注:hook_executor 不显式 drop,生命周期由 Arc 自动管理
|
||||
let _ = hook_executor;
|
||||
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
@@ -491,20 +485,277 @@ impl AgentSession {
|
||||
.hook_executor
|
||||
.execute(HookEvent::OnTurnEnd, &end_ctx)
|
||||
.await;
|
||||
|
||||
// Phase 16: 摘要检查点(流式路径 turn_index 已被 submit_turn_stream 提前 ++1)
|
||||
self.maybe_summarize(self.turn_index.saturating_sub(1)).await;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ====== Phase 17: 快照序列化 ======
|
||||
|
||||
/// 将当前状态拍平为 `SessionSnapshot`。
|
||||
///
|
||||
/// **需要 async**:因为 `session_memory` 的条目存储在 `MemoryStore` 中,读取需异步 I/O。
|
||||
/// 通过 `SessionMemory::list_entries()` 获取完整条目(保留 `metadata` 和 `created_at`)。
|
||||
///
|
||||
/// `Arc<dyn Agent>` 和 `Arc<RuntimeBundle>` **不进入快照**——由 `from_snapshot()` 调用方注入。
|
||||
pub async fn to_snapshot(&self) -> SessionSnapshot {
|
||||
// 拍平 session_memory → HashMap<String, SessionMemoryEntry>
|
||||
// 失败时回退到空 map(错误已记录,不阻断 checkpoint 主流程)。
|
||||
let session_memory_data = match self.session_memory.list_entries().await {
|
||||
Ok(entries) => entries
|
||||
.into_iter()
|
||||
.map(|(key, value, metadata, created_at)| {
|
||||
(
|
||||
key,
|
||||
SessionMemoryEntry {
|
||||
value,
|
||||
metadata,
|
||||
created_at: Some(created_at),
|
||||
},
|
||||
)
|
||||
})
|
||||
.collect(),
|
||||
Err(e) => {
|
||||
tracing::error!("session_memory list_entries failed: {}", e);
|
||||
HashMap::new()
|
||||
}
|
||||
};
|
||||
|
||||
SessionSnapshot {
|
||||
session_id: self.session_id.clone(),
|
||||
agent_name: self.agent.name().to_string(),
|
||||
turn_index: self.turn_index,
|
||||
cost_so_far: self.cost_so_far.clone(),
|
||||
slots: self.slots.clone(),
|
||||
current_slot_id: self.current_slot_id.clone(),
|
||||
last_summary_turn: self.last_summary_turn,
|
||||
session_memory_data,
|
||||
}
|
||||
}
|
||||
|
||||
/// 从 `SessionSnapshot` + agent + bundle **纯同步**重建 `AgentSession`。
|
||||
///
|
||||
/// **不执行任何 I/O**:`session_memory_data` 暂存于 `pending_memory_restore` 字段,
|
||||
/// 由调用方显式 `await session.restore_memory()` 写回持久层。
|
||||
///
|
||||
/// 调用方负责提供与 `snapshot.agent_name` 对应的 `Arc<dyn Agent>`(引擎层只保留名字做调试用)。
|
||||
pub fn from_snapshot(
|
||||
snapshot: SessionSnapshot,
|
||||
agent: Arc<dyn Agent>,
|
||||
bundle: Arc<RuntimeBundle>,
|
||||
) -> Result<Self, EngineError> {
|
||||
// 校验 bundle 的 session_memory_backend 与 snapshot 兼容
|
||||
// (v0.3 不强制同 backend——以新构造的 session_memory 所属 backend 为准)
|
||||
let backend = bundle
|
||||
.session_memory_backend
|
||||
.clone()
|
||||
.unwrap_or_else(|| Arc::new(InMemoryStore::new()));
|
||||
let session_memory = SessionMemory::new(backend, &snapshot.session_id);
|
||||
|
||||
// 解析 agent_name 仅供调试(不强制匹配,因为不同进程的 Agent 实现可能不同)
|
||||
let _ = snapshot.agent_name.as_str();
|
||||
|
||||
// 确保至少有一个 slot(与 new() 行为一致)
|
||||
let mut slots = snapshot.slots;
|
||||
if slots.is_empty() {
|
||||
slots.insert(
|
||||
"default".to_string(),
|
||||
ContextSlot::new(
|
||||
&snapshot.session_id,
|
||||
"default",
|
||||
SlotConfig::default(),
|
||||
),
|
||||
);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
session_id: snapshot.session_id,
|
||||
agent,
|
||||
bundle,
|
||||
turn_index: snapshot.turn_index,
|
||||
cost_so_far: snapshot.cost_so_far,
|
||||
session_memory,
|
||||
slots,
|
||||
current_slot_id: snapshot.current_slot_id,
|
||||
last_summary_turn: snapshot.last_summary_turn,
|
||||
pending_memory_restore: if snapshot.session_memory_data.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(snapshot.session_memory_data)
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
/// 将 `from_snapshot()` 暂存的 `session_memory_data` 写回 `SessionMemory` 持久层。
|
||||
///
|
||||
/// **从 `from_snapshot()` 中剥离的异步操作**:确保构造函数是纯同步的。
|
||||
/// 调用方在 `from_snapshot()` 后显式 `await`。
|
||||
///
|
||||
/// **错误处理**:逐条写入。某条失败时返回 `Err` 但**不回滚**已写入条目。
|
||||
/// 调用方可选择重试或忽略——不影响 AgentSession 内存状态。
|
||||
///
|
||||
/// **幂等性**:重复调用安全(首次成功后 `pending_memory_restore` 已被设为 `None`,
|
||||
/// 第二次调用立即返回 `Ok(())`)。
|
||||
///
|
||||
/// **完整恢复**:使用 `SessionMemory::set_with_meta()` 保留原始 `metadata` 和 `created_at`
|
||||
/// ——不像 `set()` 会清空 metadata 并把 created_at 设为当前时间。
|
||||
pub async fn restore_memory(&mut self) -> Result<(), EngineError> {
|
||||
// 取出 pending 并立即清空(避免重复 restore 时二次写入;幂等性保证)
|
||||
let entries = self.pending_memory_restore.take();
|
||||
let entries = match entries {
|
||||
Some(m) if !m.is_empty() => m,
|
||||
_ => return Ok(()), // 无 pending 或已被清空 → 立即返回
|
||||
};
|
||||
|
||||
for (key, entry) in entries {
|
||||
self.session_memory
|
||||
.set_with_meta(
|
||||
&key,
|
||||
&entry.value,
|
||||
entry.metadata.clone(),
|
||||
entry.created_at,
|
||||
)
|
||||
.await
|
||||
.map_err(EngineError::Agent)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 是否有待写回的 `session_memory_data`(`from_snapshot()` 后尚未 `restore_memory()`)。
|
||||
pub fn has_pending_memory_restore(&self) -> bool {
|
||||
self.pending_memory_restore
|
||||
.as_ref()
|
||||
.map(|m| !m.is_empty())
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
// ====== Phase 16: 摘要自动生成 ======
|
||||
|
||||
/// 读取 SessionMemory 中最新的对话摘要(`None` 表示从未生成过)。
|
||||
pub async fn get_conversation_summary(&self) -> Result<Option<String>, AgentError> {
|
||||
self.session_memory.get("conversation_summary").await
|
||||
}
|
||||
|
||||
/// 水位 + 防抖检查:是否应当触发摘要生成。
|
||||
/// 防抖只对"上一轮与本轮之间的间隔"起作用——首次(`last_summary_turn.is_none()`)不阻塞。
|
||||
/// `current_turn` 显式传入而非读 `self.turn_index`,因为流式路径中 `submit_turn_stream` 已提前 ++1,
|
||||
/// `finalize_turn` 会用 `saturating_sub(1)` 修正后的值传入此函数。
|
||||
fn should_summarize(&self, cfg: &SummaryConfig, current_turn: u32) -> bool {
|
||||
let debounce_ok = match self.last_summary_turn {
|
||||
None => true,
|
||||
Some(last) => current_turn.saturating_sub(last) >= cfg.debounce_turns,
|
||||
};
|
||||
debounce_ok
|
||||
&& self.cost_so_far.total().total_tokens as f64
|
||||
>= cfg.max_context_tokens as f64 * cfg.trigger_token_ratio
|
||||
}
|
||||
|
||||
/// 检查点入口:水位超阈值时调 LLM 生成摘要,写入 slot config 与 SessionMemory。
|
||||
/// 所有错误(含 LLM error、save 失败、session_memory 写失败)均静默(`tracing::error!` 后返回)。
|
||||
async fn maybe_summarize(&mut self, current_turn: u32) {
|
||||
let cfg = match self.bundle.config.summary_config.clone() {
|
||||
Some(c) => c,
|
||||
None => return,
|
||||
};
|
||||
if !self.should_summarize(&cfg, current_turn) {
|
||||
return;
|
||||
}
|
||||
|
||||
// 先 clone 出 &self 借用范围内所需数据,后续释放借用再 await/mut
|
||||
let provider = Arc::clone(&self.bundle.provider);
|
||||
let messages = self
|
||||
.slots
|
||||
.get(&self.current_slot_id)
|
||||
.map(|s| s.messages.clone())
|
||||
.unwrap_or_default();
|
||||
if messages.is_empty() {
|
||||
return;
|
||||
}
|
||||
let max_tool_result_chars = cfg.max_tool_result_chars;
|
||||
let model = cfg.summary_model.clone();
|
||||
let prompt = cfg.summary_prompt.clone();
|
||||
|
||||
let result =
|
||||
Self::generate_summary(&provider, &messages, &prompt, model.as_deref(), max_tool_result_chars)
|
||||
.await;
|
||||
|
||||
match result {
|
||||
Ok(text) => {
|
||||
tracing::info!(turn = current_turn, summary_len = text.len(), "摘要自动生成成功");
|
||||
// Resolve store first (immutable borrow on self) before mutable borrow on slots.
|
||||
let store = self.resolve_store();
|
||||
if let Some(slot) = self.slots.get_mut(&self.current_slot_id)
|
||||
&& let SlotMode::Focused(ref mut focused_cfg) = slot.config.mode
|
||||
{
|
||||
focused_cfg.summary_override = Some(text.clone());
|
||||
// Full 模式下 summary_override 未被修改,无需持久化 slot
|
||||
if let Err(e) = slot.save(&*store).await {
|
||||
tracing::error!("summary config persist failed: {}", e);
|
||||
}
|
||||
}
|
||||
if let Err(e) = self.session_memory.set("conversation_summary", &text).await {
|
||||
tracing::error!("summary session_memory write failed: {}", e);
|
||||
}
|
||||
self.last_summary_turn = Some(current_turn);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("摘要自动生成失败 (turn={}): {}", current_turn, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 关联函数:调一次 LLM 生成摘要。空消息列表直接返回空串(不浪费 LLM 调用)。
|
||||
/// `summary_model=None` 时沿用 `CycleConfig::default()` 的默认模型(避免硬编码到非 OpenAI 用户不适配的 `"gpt-4o"`)。
|
||||
async fn generate_summary(
|
||||
provider: &Arc<dyn LlmProvider>,
|
||||
messages: &[Message],
|
||||
prompt_template: &str,
|
||||
summary_model: Option<&str>,
|
||||
max_tool_result_chars: usize,
|
||||
) -> Result<String, LlmError> {
|
||||
if messages.is_empty() {
|
||||
return Ok(String::new());
|
||||
}
|
||||
let messages_text = format_messages_as_text(messages, max_tool_result_chars);
|
||||
let prompt = prompt_template.replace("{messages}", &messages_text);
|
||||
|
||||
let config = CycleConfig {
|
||||
max_tokens: Some(1024),
|
||||
..CycleConfig::default()
|
||||
};
|
||||
let config = if let Some(model) = summary_model {
|
||||
CycleConfig {
|
||||
model: model.to_string(),
|
||||
..config
|
||||
}
|
||||
} else {
|
||||
config
|
||||
};
|
||||
|
||||
let mut cycle = LlmCycle::new_with_arc(Arc::clone(provider), config);
|
||||
// submit_messages 使用自身参数构造 request,不读 self.messages——prompt 必须放在 messages 参数里
|
||||
let response = cycle
|
||||
.submit_messages(vec![Message::user_text(prompt)], vec![])
|
||||
.await?;
|
||||
Ok(response.text())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::agent::builder::AgentBuilder;
|
||||
use crate::agent::FocusedConfig;
|
||||
use crate::llm::hooks::{Hook, HookContext, HookExecutor, HookResult};
|
||||
use crate::llm::mock::MockProvider;
|
||||
use crate::llm::stream::StreamEvent;
|
||||
use crate::llm::types::message::ContentBlock;
|
||||
use crate::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use crate::tools::ToolRegistry;
|
||||
use async_trait::async_trait;
|
||||
use futures_util::StreamExt;
|
||||
use std::sync::atomic::{AtomicU32, Ordering};
|
||||
|
||||
/// 计数 hook —— 每被调用一次 +1。
|
||||
@@ -974,4 +1225,409 @@ mod tests {
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, AgentError::SlotReadonly(_)));
|
||||
}
|
||||
|
||||
// ====== Phase 9 Step 5: 集成测试 ======
|
||||
|
||||
/// Phase 9 Step 5.1 — `submit_turn_stream` 端到端链路。
|
||||
///
|
||||
/// 验证:mock provider → `submit_turn_stream` 消费流 → 收到 TextDelta + MessageComplete
|
||||
/// → `finalize_turn` 后 `cost_so_far` 正确更新,turn_index 递增。
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn submit_turn_stream_end_to_end() {
|
||||
let (mut session, _, _) = build_session(vec![assistant_text("hi back")]);
|
||||
|
||||
let mut stream = session
|
||||
.submit_turn_stream("user msg")
|
||||
.await
|
||||
.expect("submit_turn_stream 应成功");
|
||||
|
||||
// 消费流并提取 MessageComplete
|
||||
let mut final_response: Option<MessageResponse> = None;
|
||||
while let Some(ev) = stream.next().await {
|
||||
if let StreamEvent::MessageComplete { full_response } = &ev {
|
||||
final_response = Some(full_response.clone());
|
||||
}
|
||||
}
|
||||
|
||||
let response = final_response.expect("流中应有 MessageComplete");
|
||||
// consumer 负责构造本轮新增消息列表(user_input + assistant_response)。
|
||||
// submit_turn_stream 不会自动写入 self.slots(流是延迟的),
|
||||
// 消费者需在 finalize_turn 时把 [user_input, ...tool_results, final_response] 一并传入。
|
||||
let new_messages = vec![Message::user_text("user msg"), response.message.clone()];
|
||||
session
|
||||
.finalize_turn(&response, new_messages)
|
||||
.await
|
||||
.expect("finalize_turn 应成功");
|
||||
|
||||
// cost_so_far 已累计(assistant_text 的 usage 是 from_input_output(10, 5))
|
||||
assert_eq!(session.usage().total().prompt_tokens, 10);
|
||||
assert_eq!(session.usage().total().completion_tokens, 5);
|
||||
// turn_index 已递增
|
||||
assert_eq!(session.turn_index(), 1);
|
||||
|
||||
// default slot 应包含 user 输入和 assistant 响应
|
||||
let slot = session.slots.get("default").expect("default slot");
|
||||
let has_user = slot
|
||||
.messages
|
||||
.iter()
|
||||
.any(|m| matches!(m, Message::User { .. }));
|
||||
let has_resp = slot.messages.iter().any(|m| extract_text(m) == "hi back");
|
||||
assert!(has_user && has_resp, "default slot 应包含 user 和 assistant 消息");
|
||||
}
|
||||
|
||||
/// Phase 9 Step 5.2 — `submit_turn_stream` 触发 OnTurnStart / OnTurnEnd hook。
|
||||
///
|
||||
/// 验证:OnTurnStart 在 `submit_turn_stream` 返回流之前已触发;
|
||||
/// OnTurnEnd 在 `finalize_turn` 调用后才触发。
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn submit_turn_stream_triggers_turn_hooks() {
|
||||
let (mut session, start_count, end_count) = build_session(vec![assistant_text("ok")]);
|
||||
|
||||
// 初始状态:两个 hook 计数都是 0
|
||||
assert_eq!(start_count.0.load(Ordering::SeqCst), 0);
|
||||
assert_eq!(end_count.0.load(Ordering::SeqCst), 0);
|
||||
|
||||
// 调 submit_turn_stream
|
||||
let mut stream = session
|
||||
.submit_turn_stream("user msg")
|
||||
.await
|
||||
.expect("submit_turn_stream 应成功");
|
||||
|
||||
// OnTurnStart 应在流返回前已触发
|
||||
assert_eq!(
|
||||
start_count.0.load(Ordering::SeqCst),
|
||||
1,
|
||||
"OnTurnStart 应在 submit_turn_stream 返回流之前触发"
|
||||
);
|
||||
// OnTurnEnd 此时尚未触发
|
||||
assert_eq!(
|
||||
end_count.0.load(Ordering::SeqCst),
|
||||
0,
|
||||
"OnTurnEnd 不应在 submit_turn_stream 阶段触发"
|
||||
);
|
||||
|
||||
// 消费流
|
||||
let mut final_response: Option<MessageResponse> = None;
|
||||
while let Some(ev) = stream.next().await {
|
||||
if let StreamEvent::MessageComplete { full_response } = &ev {
|
||||
final_response = Some(full_response.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// finalize_turn
|
||||
let response = final_response.expect("流中应有 MessageComplete");
|
||||
session
|
||||
.finalize_turn(&response, vec![])
|
||||
.await
|
||||
.expect("finalize_turn 应成功");
|
||||
|
||||
// OnTurnEnd 已触发
|
||||
assert_eq!(
|
||||
end_count.0.load(Ordering::SeqCst),
|
||||
1,
|
||||
"OnTurnEnd 应在 finalize_turn 后触发"
|
||||
);
|
||||
}
|
||||
|
||||
// ====== Phase 16: 摘要自动生成测试 ======
|
||||
|
||||
/// 构造带 `SummaryConfig` 的 session。
|
||||
/// mock provider 队列按 `[conv_1, summary_1, conv_2, summary_2, ...]` 交错排列,
|
||||
/// 因为每轮 `submit_turn` 中 conversation LLM 调用先于 summary LLM 调用。
|
||||
fn build_session_with_summary(
|
||||
provider_responses: Vec<MessageResponse>,
|
||||
summary_responses: Vec<MessageResponse>,
|
||||
cfg: SummaryConfig,
|
||||
) -> AgentSession {
|
||||
let mut interleaved = Vec::new();
|
||||
let max_len = provider_responses.len().max(summary_responses.len());
|
||||
for i in 0..max_len {
|
||||
if let Some(r) = provider_responses.get(i) {
|
||||
interleaved.push(r.clone());
|
||||
}
|
||||
if let Some(r) = summary_responses.get(i) {
|
||||
interleaved.push(r.clone());
|
||||
}
|
||||
}
|
||||
|
||||
let provider = Arc::new(MockProvider::new(interleaved));
|
||||
let agent = Arc::new(StubAgent {
|
||||
name: "stub".into(),
|
||||
prompt: None,
|
||||
});
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.summary_config(cfg)
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
AgentSession::new(agent, "summary-session", bundle)
|
||||
}
|
||||
|
||||
/// 默认用法:token 用量 ~15,远低于默认 32K 窗口的 0.75=24K 阈值 → 不触发摘要。
|
||||
#[tokio::test]
|
||||
async fn summary_not_generated_below_threshold() {
|
||||
let mut session = build_session_with_summary(
|
||||
vec![assistant_text("a"), assistant_text("b"), assistant_text("c")],
|
||||
vec![assistant_text("should_not_appear")],
|
||||
SummaryConfig::default(),
|
||||
);
|
||||
|
||||
for i in 0..3 {
|
||||
session
|
||||
.submit_turn(&format!("msg {}", i))
|
||||
.await
|
||||
.expect("submit_turn 应成功");
|
||||
}
|
||||
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert!(summary.is_none(), "未达阈值时不应生成摘要");
|
||||
}
|
||||
|
||||
/// 设置极低 max_context_tokens=100 + 0.5 比例 → 第一轮触发(usage 为 10+5=15 > 50)。
|
||||
#[tokio::test]
|
||||
async fn summary_generated_above_threshold() {
|
||||
let mut session = build_session_with_summary(
|
||||
vec![assistant_text("a"), assistant_text("b"), assistant_text("c")],
|
||||
vec![
|
||||
assistant_text("summary-1"),
|
||||
assistant_text("summary-2"),
|
||||
assistant_text("summary-3"),
|
||||
],
|
||||
SummaryConfig {
|
||||
max_context_tokens: 20, // 阈值 20 * 0.5 = 10
|
||||
trigger_token_ratio: 0.5,
|
||||
debounce_turns: 0, // 关闭防抖便于测试
|
||||
..SummaryConfig::default()
|
||||
},
|
||||
);
|
||||
|
||||
// 第 1 轮:usage=15 ≥ 10,debounce=0 → 触发
|
||||
session.submit_turn("m1").await.unwrap();
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert!(summary.is_some(), "应触发摘要");
|
||||
}
|
||||
|
||||
/// 防抖:trigger 触发后,debounce_turns=3 内即使再次达阈值也不重复。
|
||||
#[tokio::test]
|
||||
async fn summary_debounce_works() {
|
||||
let mut session = build_session_with_summary(
|
||||
vec![
|
||||
assistant_text("r1"),
|
||||
assistant_text("r2"),
|
||||
assistant_text("r3"),
|
||||
assistant_text("r4"),
|
||||
],
|
||||
vec![assistant_text("sum-1")],
|
||||
SummaryConfig {
|
||||
max_context_tokens: 20,
|
||||
trigger_token_ratio: 0.5,
|
||||
debounce_turns: 3,
|
||||
..SummaryConfig::default()
|
||||
},
|
||||
);
|
||||
|
||||
session.submit_turn("m1").await.unwrap();
|
||||
let first_summary = session.get_conversation_summary().await.unwrap();
|
||||
assert_eq!(first_summary.as_deref(), Some("sum-1"));
|
||||
|
||||
// 第 2、3 轮:即使都超阈值,debounce 阻止再次触发
|
||||
for _ in 0..2 {
|
||||
session.submit_turn("m").await.unwrap();
|
||||
}
|
||||
let still_summary = session.get_conversation_summary().await.unwrap();
|
||||
assert_eq!(
|
||||
still_summary.as_deref(),
|
||||
Some("sum-1"),
|
||||
"debounce 内不应重复生成(Provider 上没有更多预设摘要响应可用)"
|
||||
);
|
||||
}
|
||||
|
||||
/// Full 模式:摘要被生成并写入 session_memory,但 slot config.summary_override 仍为 None。
|
||||
#[tokio::test]
|
||||
async fn summary_written_to_session_memory_but_full_mode_does_not_inject() {
|
||||
let mut session = build_session_with_summary(
|
||||
vec![assistant_text("a"), assistant_text("b"), assistant_text("c")],
|
||||
vec![assistant_text("captured-summary")],
|
||||
SummaryConfig {
|
||||
max_context_tokens: 20,
|
||||
trigger_token_ratio: 0.5,
|
||||
debounce_turns: 0,
|
||||
..SummaryConfig::default()
|
||||
},
|
||||
);
|
||||
|
||||
session.submit_turn("m1").await.unwrap();
|
||||
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert_eq!(summary.as_deref(), Some("captured-summary"));
|
||||
|
||||
// default slot 是 Full 模式 → summary_override 应为 None(filter_focused 不会触发)
|
||||
let slot = session.slots.get("default").unwrap();
|
||||
assert!(matches!(slot.config.mode, SlotMode::Full));
|
||||
}
|
||||
|
||||
/// 摘要生成失败不阻断 submit_turn(Provider 队列只够对话轮次,摘要调用返回 Other 错误)。
|
||||
#[tokio::test]
|
||||
async fn summary_failure_does_not_block_turn() {
|
||||
// 故意只提供 1 个对话响应;摘要调用时队列耗尽,MockProvider 返回 LlmError::Other
|
||||
let mut session = build_session_with_summary(
|
||||
vec![assistant_text("only-one")], // 后续摘要会失败
|
||||
vec![], // 无摘要响应
|
||||
SummaryConfig {
|
||||
max_context_tokens: 20,
|
||||
trigger_token_ratio: 0.5,
|
||||
debounce_turns: 0,
|
||||
..SummaryConfig::default()
|
||||
},
|
||||
);
|
||||
|
||||
let response = session
|
||||
.submit_turn("m1")
|
||||
.await
|
||||
.expect("submit_turn 应成功(即便摘要失败)");
|
||||
assert_eq!(extract_text(&response.message), "only-one");
|
||||
|
||||
// 摘要未生成(Provider 已耗尽)
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert!(summary.is_none());
|
||||
}
|
||||
|
||||
/// 未配置 SummaryConfig 时零影响。
|
||||
#[tokio::test]
|
||||
async fn summary_skipped_when_not_configured() {
|
||||
let (mut session, _, _) = build_session(vec![assistant_text("r1"), assistant_text("r2")]);
|
||||
for _ in 0..2 {
|
||||
session.submit_turn("m").await.unwrap();
|
||||
}
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert!(summary.is_none());
|
||||
}
|
||||
|
||||
/// 流式路径(submit_turn_stream + finalize_turn):摘要检查点正确触发。
|
||||
#[tokio::test(flavor = "multi_thread")]
|
||||
async fn summary_stream_path_triggers_check() {
|
||||
let mut session = build_session_with_summary(
|
||||
vec![assistant_text("stream-resp")],
|
||||
vec![assistant_text("stream-summary")],
|
||||
SummaryConfig {
|
||||
max_context_tokens: 20,
|
||||
trigger_token_ratio: 0.5,
|
||||
debounce_turns: 0,
|
||||
..SummaryConfig::default()
|
||||
},
|
||||
);
|
||||
|
||||
let mut stream = session
|
||||
.submit_turn_stream("user msg")
|
||||
.await
|
||||
.expect("stream ok");
|
||||
let mut response: Option<MessageResponse> = None;
|
||||
while let Some(ev) = stream.next().await {
|
||||
if let StreamEvent::MessageComplete { full_response } = &ev {
|
||||
response = Some(full_response.clone());
|
||||
}
|
||||
}
|
||||
let resp = response.expect("MessageComplete event");
|
||||
// finalize_turn 需要本轮新增消息:用户输入 + assistant 响应。
|
||||
// slot.append_messages 之后才会被 maybe_summarize 看到。
|
||||
let new_messages = vec![Message::user_text("user msg"), resp.message.clone()];
|
||||
session
|
||||
.finalize_turn(&resp, new_messages)
|
||||
.await
|
||||
.expect("finalize_turn ok");
|
||||
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert_eq!(summary.as_deref(), Some("stream-summary"));
|
||||
}
|
||||
|
||||
/// W7:Focused 模式摘要写入 `summary_override` + `slot.save()` 正向验证。
|
||||
#[tokio::test]
|
||||
async fn summary_written_to_focused_slot_config() {
|
||||
let mut session = build_session_with_summary(
|
||||
vec![assistant_text("a"), assistant_text("b")],
|
||||
vec![assistant_text("the-summary")],
|
||||
SummaryConfig {
|
||||
max_context_tokens: 20,
|
||||
trigger_token_ratio: 0.5,
|
||||
debounce_turns: 0,
|
||||
..SummaryConfig::default()
|
||||
},
|
||||
);
|
||||
|
||||
// 1. 把 default slot 切到 Focused 模式
|
||||
session
|
||||
.create_slot(
|
||||
"focused",
|
||||
Some(SlotConfig {
|
||||
mode: SlotMode::Focused(FocusedConfig {
|
||||
keep_system: false,
|
||||
recent_messages: 5,
|
||||
summary_override: None,
|
||||
}),
|
||||
source: SlotSource::New,
|
||||
budget: Default::default(),
|
||||
compact: true,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
session.switch_slot("focused").await.unwrap();
|
||||
|
||||
// 2. 触发摘要
|
||||
session.submit_turn("m1").await.unwrap();
|
||||
|
||||
// 3. SessionMemory 有值
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert_eq!(summary.as_deref(), Some("the-summary"));
|
||||
|
||||
// 4. Focused slot 的 summary_override 也应有值(正向验证)
|
||||
let slot = session.slots.get("focused").unwrap();
|
||||
assert!(
|
||||
matches!(&slot.config.mode, SlotMode::Focused(focused) if focused.summary_override.is_some()),
|
||||
"Focused 模式下 summary_override 应被写入"
|
||||
);
|
||||
}
|
||||
|
||||
/// W1: 空消息守卫——`generate_summary` 空消息直接返回 `""`,不调用 LLM。
|
||||
/// 这里通过构建一个空 slot 触发,第一次 `submit_turn` 后 slot 才有消息。
|
||||
/// 验证:先调用 `format_messages_as_text` 走纯函数路径检查。
|
||||
#[tokio::test]
|
||||
async fn summary_skipped_for_empty_messages() {
|
||||
// 直接走 format_messages_as_text,验证空消息返回空串。
|
||||
// 这等同于 generate_summary 入口守卫(见 session.rs:560-562)。
|
||||
let text = format_messages_as_text(&[], 500);
|
||||
assert_eq!(text, "");
|
||||
}
|
||||
|
||||
/// W1: `max_context_tokens` 设置过大时永不触发摘要。
|
||||
#[tokio::test]
|
||||
async fn summary_not_generated_if_max_context_unreachable() {
|
||||
let mut session = build_session_with_summary(
|
||||
vec![
|
||||
assistant_text("r1"),
|
||||
assistant_text("r2"),
|
||||
assistant_text("r3"),
|
||||
assistant_text("r4"),
|
||||
],
|
||||
vec![assistant_text("should-not-appear")],
|
||||
SummaryConfig {
|
||||
max_context_tokens: 1_000_000, // 远大于任何合理累计 token
|
||||
trigger_token_ratio: 0.75,
|
||||
debounce_turns: 0,
|
||||
..SummaryConfig::default()
|
||||
},
|
||||
);
|
||||
|
||||
// 多轮 submit_turn,全部 15 token/轮,远低于 0.75 * 1M = 750K 阈值
|
||||
for i in 0..4 {
|
||||
session.submit_turn(&format!("m{}", i)).await.unwrap();
|
||||
}
|
||||
|
||||
let summary = session.get_conversation_summary().await.unwrap();
|
||||
assert!(summary.is_none(), "巨型 max_context_tokens 应永不触发");
|
||||
}
|
||||
}
|
||||
@@ -44,12 +44,37 @@ impl SessionMemory {
|
||||
}
|
||||
|
||||
/// 写入一条 key-value 条目(覆盖同名 key)。
|
||||
///
|
||||
/// **不保留 metadata 和 created_at** —— 写入时 metadata 为空 JSON `{}`,created_at 为 `now_utc()`。
|
||||
/// 若需保留这两个字段(如 checkpoint rollback),使用 [`Self::set_with_meta`]。
|
||||
pub async fn set(&self, key: &str, value: &str) -> Result<(), AgentError> {
|
||||
self.set_with_meta(key, value, serde_json::json!({}), None).await
|
||||
}
|
||||
|
||||
/// 写入一条 key-value 条目(含完整 metadata + created_at)。
|
||||
///
|
||||
/// Phase 17 新增:供 `AgentSession::restore_memory()` 使用,保证 checkpoint rollback 时
|
||||
/// 恢复完整的 session_memory 条目(包括原 metadata 和创建时间戳)。
|
||||
///
|
||||
/// - `metadata`: 通常为 `serde_json::Value`(快照中保留的 metadata JSON)
|
||||
/// - `created_at`: 快照中的原始时间戳(Unix 秒);若为 `None` 则用 `now_utc()`(默认行为)
|
||||
pub async fn set_with_meta(
|
||||
&self,
|
||||
key: &str,
|
||||
value: &str,
|
||||
metadata: serde_json::Value,
|
||||
created_at: Option<i64>,
|
||||
) -> Result<(), AgentError> {
|
||||
let created_at_dt = match created_at {
|
||||
Some(secs) => OffsetDateTime::from_unix_timestamp(secs)
|
||||
.unwrap_or_else(|_| OffsetDateTime::now_utc()),
|
||||
None => OffsetDateTime::now_utc(),
|
||||
};
|
||||
let item = MemoryItem {
|
||||
id: self.internal_key(key),
|
||||
content: value.to_string(),
|
||||
metadata: serde_json::json!({}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
metadata,
|
||||
created_at: created_at_dt,
|
||||
};
|
||||
self.store.save(item).await.map_err(AgentError::Memory)
|
||||
}
|
||||
@@ -103,6 +128,34 @@ impl SessionMemory {
|
||||
.map_err(AgentError::Memory)
|
||||
}
|
||||
|
||||
/// 列出当前 namespace 下所有条目(含完整 `MemoryItem`:value / metadata / created_at)。
|
||||
///
|
||||
/// Phase 17 新增:供 `AgentSession::to_snapshot()` 拍平 session_memory 时使用,
|
||||
/// 保留 metadata 和 created_at 时间戳(用 `set/get/remove` 三个 API 会丢字段)。
|
||||
///
|
||||
/// 返回 `Vec<(原始 key, value, metadata, created_at_unix_secs)>`,原始 key 已剥离 namespace 前缀。
|
||||
pub async fn list_entries(
|
||||
&self,
|
||||
) -> Result<Vec<(String, String, serde_json::Value, i64)>, AgentError> {
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some(format!("{}:", self.namespace)),
|
||||
..Default::default()
|
||||
};
|
||||
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
|
||||
let prefix_with_colon = format!("{}:", self.namespace);
|
||||
let mut out = Vec::with_capacity(items.len());
|
||||
for item in items {
|
||||
let key = item
|
||||
.id
|
||||
.strip_prefix(&prefix_with_colon)
|
||||
.unwrap_or(&item.id)
|
||||
.to_string();
|
||||
let created_at_unix = item.created_at.unix_timestamp();
|
||||
out.push((key, item.content, item.metadata, created_at_unix));
|
||||
}
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
/// 清空当前 namespace 下所有条目。
|
||||
pub async fn clear(&self) -> Result<(), AgentError> {
|
||||
let filter = MemoryFilter {
|
||||
|
||||
@@ -0,0 +1,242 @@
|
||||
//! 摘要自动生成 —— 在长对话中自动压缩上下文。
|
||||
//!
|
||||
//! 通过 `AgentSession` 内联检查点检测 token 水位,调用 LLM 生成摘要,
|
||||
//! 写入 `FocusedConfig.summary_override` 与 `SessionMemory["conversation_summary"]`。
|
||||
//!
|
||||
//! 关闭端位于 `FocusedConfig::filter_focused`(见 `agent/context.rs`)。
|
||||
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
|
||||
/// 默认摘要 prompt(含 `{messages}` 占位符,运行期替换为对话历史文本)。
|
||||
pub const DEFAULT_SUMMARY_PROMPT: &str = "请为以下对话生成一个简洁的中文摘要,突出关键结论、用户偏好和重要上下文信息。保持客观,不要添加对话中不存在的信息。\n\n{messages}";
|
||||
|
||||
/// 摘要自动生成配置(opt-in:通过 `AgentBuilder::summary_config(cfg)` 启用)。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SummaryConfig {
|
||||
/// Token 水位触发比例(0.0 ~ 1.0)。
|
||||
pub trigger_token_ratio: f64,
|
||||
|
||||
/// 模型上下文窗口大小(token)。
|
||||
/// ⚠️ 设置为超过模型实际窗口的值会导致摘要永远不触发。
|
||||
pub max_context_tokens: u32,
|
||||
|
||||
/// 摘要 prompt 模板。`{messages}` 将被替换为对话历史纯文本。
|
||||
pub summary_prompt: String,
|
||||
|
||||
/// 摘要间隔防抖(轮次):两次摘要至少间隔这么多次 `submit_turn`。
|
||||
pub debounce_turns: u32,
|
||||
|
||||
/// 摘要生成使用的模型(`None` = 沿用主 provider 默认模型)。
|
||||
/// 推荐设为便宜模型(如 `"gpt-4o-mini"`)以节省摘要成本。
|
||||
pub summary_model: Option<String>,
|
||||
|
||||
/// 单个 `ToolResult` 在摘要输入中保留的最大 Unicode 字符数。
|
||||
/// 超过此值从开头截断(`chars().take(n)`,字符级安全)。
|
||||
pub max_tool_result_chars: usize,
|
||||
}
|
||||
|
||||
impl Default for SummaryConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
trigger_token_ratio: 0.75,
|
||||
max_context_tokens: 32_000,
|
||||
summary_prompt: DEFAULT_SUMMARY_PROMPT.into(),
|
||||
debounce_turns: 3,
|
||||
summary_model: None,
|
||||
max_tool_result_chars: 500,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 把消息列表格式化为摘要 LLM 所需的纯文本(简洁版)。
|
||||
///
|
||||
/// 每行一条消息:
|
||||
/// - `System/User/Assistant` 取首个 `Text` block 拼接
|
||||
/// - `Assistant` 中的 `ToolUse` 标记为 `[Tool: {name}]`
|
||||
/// - `ToolResult` 标记为 `Tool Result [{tool_call_id}]:`(含 tool_call_id 以便多工具场景关联)
|
||||
/// - 长 `ToolResult` 截断到 `max_tool_result_chars` 个字符
|
||||
///
|
||||
/// 整段对话若超过 `30_000` 字符,从前面截断,**优先保留最新消息**,
|
||||
/// 因为新近交互对摘要而言更有信息量。
|
||||
pub fn format_messages_as_text(messages: &[Message], max_tool_result_chars: usize) -> String {
|
||||
let mut lines = Vec::with_capacity(messages.len());
|
||||
for msg in messages {
|
||||
match msg {
|
||||
Message::System { content } => {
|
||||
if let Some(text) = first_text(content) {
|
||||
lines.push(format!("System: {}", text));
|
||||
}
|
||||
}
|
||||
Message::User { content } => {
|
||||
if let Some(text) = first_text(content) {
|
||||
lines.push(format!("User: {}", text));
|
||||
}
|
||||
}
|
||||
Message::Assistant { content } => {
|
||||
let mut parts = Vec::new();
|
||||
for block in content {
|
||||
match block {
|
||||
ContentBlock::Text { text } => parts.push(text.clone()),
|
||||
ContentBlock::ToolUse { name, .. } => {
|
||||
parts.push(format!("[Tool: {}]", name));
|
||||
}
|
||||
ContentBlock::Thinking { text, .. } => {
|
||||
parts.push(format!("[Thinking: {}]", truncate_chars(text, 100)));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
if !parts.is_empty() {
|
||||
lines.push(format!("Assistant: {}", parts.join(" ")));
|
||||
}
|
||||
}
|
||||
Message::UserImage { .. } => {
|
||||
lines.push("User: [image]".to_string());
|
||||
}
|
||||
Message::ToolResult {
|
||||
tool_call_id,
|
||||
content,
|
||||
is_error,
|
||||
} => {
|
||||
let label = if *is_error { "Tool Error" } else { "Tool Result" };
|
||||
if let Some(text) = first_text(content) {
|
||||
let truncated = truncate_chars(text, max_tool_result_chars);
|
||||
lines.push(format!("{} [{}]: {}", label, tool_call_id, truncated));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let joined = lines.join("\n");
|
||||
truncate_total_chars(&joined, MAX_TOTAL_CHARS)
|
||||
}
|
||||
|
||||
/// 整段对话输出字符上限。超过时从前面截断,保留尾部最新消息。
|
||||
const MAX_TOTAL_CHARS: usize = 30_000;
|
||||
|
||||
fn truncate_total_chars(s: &str, max_chars: usize) -> String {
|
||||
let total = s.chars().count();
|
||||
if total <= max_chars {
|
||||
return s.to_string();
|
||||
}
|
||||
// 计算需要从前面丢弃的字符数。保留窗口从 (total - max_chars) 开始。
|
||||
let skip = total - max_chars;
|
||||
let dropped: String = s.chars().take(skip).collect();
|
||||
let mut kept = String::with_capacity(max_chars + 8);
|
||||
kept.push_str("[... earlier messages truncated ...]\n");
|
||||
// 字节切安全:dropped 由 s.chars().take(skip).collect() 构建,
|
||||
// 每个 char 的 UTF-8 字节序列完整保留,故 dropped.len() 恰好是 s 的某个 char 边界字节偏移。
|
||||
kept.push_str(&s[dropped.len()..]);
|
||||
kept
|
||||
}
|
||||
|
||||
fn first_text(content: &[ContentBlock]) -> Option<&str> {
|
||||
content.iter().find_map(|b| match b {
|
||||
ContentBlock::Text { text } => Some(text.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
}
|
||||
|
||||
fn truncate_chars(s: &str, max_chars: usize) -> String {
|
||||
if s.chars().count() <= max_chars {
|
||||
return s.to_string();
|
||||
}
|
||||
let truncated: String = s.chars().take(max_chars).collect();
|
||||
format!("{}...", truncated)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn default_values() {
|
||||
let cfg = SummaryConfig::default();
|
||||
assert_eq!(cfg.trigger_token_ratio, 0.75);
|
||||
assert_eq!(cfg.max_context_tokens, 32_000);
|
||||
assert_eq!(cfg.debounce_turns, 3);
|
||||
assert_eq!(cfg.max_tool_result_chars, 500);
|
||||
assert!(cfg.summary_model.is_none());
|
||||
assert!(cfg.summary_prompt.contains("{messages}"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_skips_empty_input() {
|
||||
let text = format_messages_as_text(&[], 500);
|
||||
assert!(text.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_user_assistant_round_trip() {
|
||||
let msgs = vec![
|
||||
Message::system("you are a translator"),
|
||||
Message::user_text("hello"),
|
||||
Message::assistant("hi"),
|
||||
];
|
||||
let text = format_messages_as_text(&msgs, 500);
|
||||
assert!(text.contains("System: you are a translator"));
|
||||
assert!(text.contains("User: hello"));
|
||||
assert!(text.contains("Assistant: hi"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_tool_result_includes_tool_call_id() {
|
||||
let msgs = vec![Message::tool_result("call_42", "ok", false)];
|
||||
let text = format_messages_as_text(&msgs, 500);
|
||||
assert_eq!(text, "Tool Result [call_42]: ok");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_tool_result_error_label() {
|
||||
let msgs = vec![Message::tool_result("call_9", "boom", true)];
|
||||
let text = format_messages_as_text(&msgs, 500);
|
||||
assert_eq!(text, "Tool Error [call_9]: boom");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_tool_use_in_assistant() {
|
||||
let msgs = vec![Message::Assistant {
|
||||
content: vec![
|
||||
ContentBlock::Text {
|
||||
text: "let me search".into(),
|
||||
},
|
||||
ContentBlock::ToolUse {
|
||||
id: "c1".into(),
|
||||
name: "search".into(),
|
||||
input: serde_json::json!({"q": "rust"}),
|
||||
},
|
||||
],
|
||||
}];
|
||||
let text = format_messages_as_text(&msgs, 500);
|
||||
assert_eq!(text, "Assistant: let me search [Tool: search]");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_truncates_long_tool_result_at_unicode_boundary() {
|
||||
let long = "a".repeat(1000);
|
||||
let msgs = vec![Message::tool_result("c", &long, false)];
|
||||
let text = format_messages_as_text(&msgs, 100);
|
||||
// 100 chars + "..."
|
||||
assert!(text.contains("..."));
|
||||
let truncated_part = text.split("...").next().unwrap();
|
||||
// "Tool Result [c]: " is 18 chars, plus 100 a's
|
||||
let a_count = truncated_part.chars().filter(|c| *c == 'a').count();
|
||||
assert_eq!(a_count, 100);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_total_charset_truncation_keeps_recent() {
|
||||
// 50 段 user 消息,每段 1000 字符 = ~50K,触发 30K 整体截断
|
||||
let mut msgs = Vec::new();
|
||||
for _ in 0..50 {
|
||||
msgs.push(Message::user_text("x".repeat(1000)));
|
||||
}
|
||||
let text = format_messages_as_text(&msgs, 500);
|
||||
// 总字符数 ≤ 30K + prefix "[... earlier messages truncated ...]\n"
|
||||
assert!(text.chars().count() <= 30_000 + 40);
|
||||
// 头部有截断标记
|
||||
assert!(text.contains("[... earlier messages truncated ...]"));
|
||||
// 最后一行的标记字符 (30 个 x) 应保留在末尾
|
||||
assert!(text.ends_with("xxxxxxxxxx"));
|
||||
}
|
||||
}
|
||||
+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"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,377 @@
|
||||
//! Checkpointer —— Time-travel 检查点管理器(Phase 17 Step 4)。
|
||||
//!
|
||||
//! 设计要点:
|
||||
//! - **不依赖 SessionManager**,可独立使用。直接操作 `MemoryStore`。
|
||||
//! - **存储 key 格式**:`ckpt:{session_id}:{ckpt_id}` → `SessionSnapshot` JSON
|
||||
//! - **ckpt_id 生成**:时间戳(纳秒)+ 单调计数器,无外部依赖(ponytail 优先于 uuid)
|
||||
//! - **rollback_load** 两阶段:仅反序列化为 `SessionSnapshot`;不重建 `AgentSession`。
|
||||
//! 调用方拿到 `SessionSnapshot` 后自行 `AgentSession::from_snapshot` + `restore_memory` + `replace`。
|
||||
//!
|
||||
//! 所有持久化错误通过 `EngineError::Memory` 透传。
|
||||
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use time::OffsetDateTime;
|
||||
|
||||
use crate::agent::session::AgentSession;
|
||||
use crate::engine::snapshot::SessionSnapshot;
|
||||
use crate::engine::EngineError;
|
||||
use crate::memory::store::MemoryStore;
|
||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||
|
||||
/// 检查点元数据(公开 API)。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct CkptMeta {
|
||||
pub ckpt_id: String,
|
||||
pub session_id: String,
|
||||
pub turn_index: u32,
|
||||
/// Unix 时间戳秒(人类可读)。
|
||||
pub created_at: u64,
|
||||
/// Unix 时间戳纳秒(用于同秒内的精确排序)。
|
||||
pub created_at_nanos: u128,
|
||||
}
|
||||
|
||||
/// 全局单调计数器(避免同一纳秒内并发 checkpoint 撞 id)。
|
||||
static CKPT_COUNTER: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
/// 生成 ckpt_id:纳秒时间戳 + 单调计数器(避免同纳秒冲突)。
|
||||
fn generate_ckpt_id() -> String {
|
||||
let nanos = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_nanos() as u64)
|
||||
.unwrap_or(0);
|
||||
let counter = CKPT_COUNTER.fetch_add(1, Ordering::Relaxed);
|
||||
format!("{:x}_{:x}", nanos, counter)
|
||||
}
|
||||
|
||||
fn ckpt_key(session_id: &str, ckpt_id: &str) -> String {
|
||||
format!("ckpt:{}:{}", session_id, ckpt_id)
|
||||
}
|
||||
|
||||
fn assert_no_colon(id: &str, field: &str) {
|
||||
if id.contains(':') {
|
||||
panic!(
|
||||
"{field} '{id}' contains ':' which would break key format. \
|
||||
Use only letters, digits, hyphens and underscores."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Time-travel 检查点管理器。
|
||||
pub struct Checkpointer {
|
||||
store: std::sync::Arc<dyn MemoryStore>,
|
||||
}
|
||||
|
||||
impl Checkpointer {
|
||||
/// 构造 Checkpointer。
|
||||
pub fn new(store: std::sync::Arc<dyn MemoryStore>) -> Self {
|
||||
Self { store }
|
||||
}
|
||||
|
||||
/// 创建新检查点。返回生成的 `ckpt_id`。
|
||||
///
|
||||
/// 流程:`session.to_snapshot().await` → 序列化为 JSON → 存 `ckpt:{session_id}:{ckpt_id}`。
|
||||
pub async fn checkpoint(&self, session: &AgentSession) -> Result<String, EngineError> {
|
||||
let snapshot = session.to_snapshot().await;
|
||||
assert_no_colon(&session.session_id, "session_id");
|
||||
|
||||
let ckpt_id = generate_ckpt_id();
|
||||
let key = ckpt_key(&session.session_id, &ckpt_id);
|
||||
|
||||
let json = serde_json::to_string(&snapshot).map_err(|e| {
|
||||
EngineError::Serialization(format!("snapshot serialize failed: {e}"))
|
||||
})?;
|
||||
|
||||
let item = MemoryItem {
|
||||
id: key,
|
||||
content: json,
|
||||
metadata: serde_json::json!({
|
||||
"turn_index": snapshot.turn_index,
|
||||
}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
};
|
||||
self.store.save(item).await?;
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session.session_id,
|
||||
ckpt_id = %ckpt_id,
|
||||
turn_index = snapshot.turn_index,
|
||||
snapshot_size = snapshot.session_memory_data.len(),
|
||||
"checkpoint created"
|
||||
);
|
||||
|
||||
Ok(ckpt_id)
|
||||
}
|
||||
|
||||
/// 反序列化 checkpoint 为 `SessionSnapshot`(不重建 `AgentSession`)。
|
||||
///
|
||||
/// 两阶段 rollback 的第一阶段。调用方拿到 `SessionSnapshot` 后自行:
|
||||
/// `AgentSession::from_snapshot(snapshot, agent, bundle)` → `restore_memory()` → `replace()`
|
||||
pub async fn rollback_load(
|
||||
&self,
|
||||
session_id: &str,
|
||||
ckpt_id: &str,
|
||||
) -> Result<SessionSnapshot, EngineError> {
|
||||
let key = ckpt_key(session_id, ckpt_id);
|
||||
let item = self.store.get(&key).await?.ok_or_else(|| {
|
||||
EngineError::CheckpointNotFound(format!("{} (session={})", ckpt_id, session_id))
|
||||
})?;
|
||||
|
||||
let snapshot: SessionSnapshot = serde_json::from_str(&item.content).map_err(|e| {
|
||||
EngineError::Serialization(format!(
|
||||
"snapshot deserialize failed (ckpt_id={}): {e}",
|
||||
ckpt_id
|
||||
))
|
||||
})?;
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
ckpt_id = %ckpt_id,
|
||||
turn_index = snapshot.turn_index,
|
||||
"checkpoint loaded for rollback"
|
||||
);
|
||||
|
||||
Ok(snapshot)
|
||||
}
|
||||
|
||||
/// 列出某 session 的所有检查点(按创建时间**降序**——最新的在前)。
|
||||
///
|
||||
/// prefix 查询 `ckpt:{session_id}:` → 反序列化 `SessionSnapshot` → 提取元数据。
|
||||
/// 不需要 `CkptMeta` 单独存储——`SessionSnapshot` 已含 `turn_index` 字段,
|
||||
/// `created_at` 用 `MemoryItem.created_at` 转换。
|
||||
pub async fn list_checkpoints(
|
||||
&self,
|
||||
session_id: &str,
|
||||
) -> Result<Vec<CkptMeta>, EngineError> {
|
||||
let prefix = format!("ckpt:{}:", session_id);
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some(prefix),
|
||||
..Default::default()
|
||||
};
|
||||
let items = self.store.list(&filter).await?;
|
||||
|
||||
tracing::debug!(
|
||||
session_id = %session_id,
|
||||
count = items.len(),
|
||||
"checkpoints listed"
|
||||
);
|
||||
|
||||
let mut metas: Vec<CkptMeta> = items
|
||||
.into_iter()
|
||||
.filter_map(|item| {
|
||||
// 从 id 中提取 ckpt_id: "ckpt:{session_id}:{ckpt_id}"
|
||||
let prefix_with_session = format!("ckpt:{}:", session_id);
|
||||
let ckpt_id = item.id.strip_prefix(&prefix_with_session)?.to_string();
|
||||
let snapshot: SessionSnapshot = serde_json::from_str(&item.content).ok()?;
|
||||
let created_at_nanos = item
|
||||
.created_at
|
||||
.unix_timestamp_nanos()
|
||||
.try_into()
|
||||
.unwrap_or(0u128);
|
||||
Some(CkptMeta {
|
||||
ckpt_id,
|
||||
session_id: session_id.to_string(),
|
||||
turn_index: snapshot.turn_index,
|
||||
created_at: item.created_at.unix_timestamp() as u64,
|
||||
created_at_nanos,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
// 按 created_at_nanos 降序(精确排序)
|
||||
metas.sort_by(|a, b| b.created_at_nanos.cmp(&a.created_at_nanos));
|
||||
Ok(metas)
|
||||
}
|
||||
|
||||
/// 删除某 session 的所有检查点(session 被 `destroy` 时调用)。
|
||||
pub async fn delete_all(&self, session_id: &str) -> Result<(), EngineError> {
|
||||
let prefix = format!("ckpt:{}:", session_id);
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some(prefix.clone()),
|
||||
..Default::default()
|
||||
};
|
||||
let items = self.store.list(&filter).await?;
|
||||
let deleted = items.len();
|
||||
for item in items {
|
||||
self.store.delete(&item.id).await?;
|
||||
}
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
deleted_count = deleted,
|
||||
"all checkpoints deleted"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 获取某 session 的最新 checkpoint(按 created_at 降序取第一个)。
|
||||
///
|
||||
/// 供 `SessionManager::recover()` 调用。
|
||||
pub async fn latest_snapshot(
|
||||
&self,
|
||||
session_id: &str,
|
||||
) -> Result<Option<SessionSnapshot>, EngineError> {
|
||||
let metas = self.list_checkpoints(session_id).await?;
|
||||
match metas.first() {
|
||||
Some(meta) => {
|
||||
let snap = self.rollback_load(session_id, &meta.ckpt_id).await?;
|
||||
Ok(Some(snap))
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::agent::builder::AgentBuilder;
|
||||
use crate::agent::summary::SummaryConfig;
|
||||
use crate::llm::mock::MockProvider;
|
||||
use crate::tools::ToolRegistry;
|
||||
use async_trait::async_trait;
|
||||
use std::sync::Arc;
|
||||
|
||||
/// 极简 Agent(无 system_prompt)。
|
||||
struct StubAgent;
|
||||
#[async_trait]
|
||||
impl crate::agent::agent::Agent for StubAgent {
|
||||
fn name(&self) -> &str {
|
||||
"stub"
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn make_bundle() -> Arc<crate::agent::runtime::RuntimeBundle> {
|
||||
Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(Arc::new(MockProvider::new(vec![])))
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(crate::llm::hooks::HookExecutor::new()))
|
||||
.summary_config(SummaryConfig::default())
|
||||
.build()
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn make_session(session_id: &str, agent: Arc<dyn crate::agent::agent::Agent>) -> AgentSession {
|
||||
AgentSession::new(agent, session_id, make_bundle())
|
||||
}
|
||||
|
||||
fn new_session_for_test(session_id: &str) -> AgentSession {
|
||||
let agent: Arc<dyn crate::agent::agent::Agent> = Arc::new(StubAgent);
|
||||
make_session(session_id, agent)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn checkpoint_roundtrip() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let cp = Checkpointer::new(store.clone());
|
||||
|
||||
let mut session = new_session_for_test("ckpt-session");
|
||||
session
|
||||
.set_session_data("design", "PostgreSQL")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let ckpt_id = cp.checkpoint(&session).await.unwrap();
|
||||
assert!(!ckpt_id.is_empty());
|
||||
|
||||
let restored = cp.rollback_load("ckpt-session", &ckpt_id).await.unwrap();
|
||||
assert_eq!(restored.session_id, "ckpt-session");
|
||||
assert_eq!(restored.turn_index, 0);
|
||||
assert_eq!(restored.session_memory_data.len(), 1);
|
||||
let entry = restored.session_memory_data.get("design").unwrap();
|
||||
assert_eq!(entry.value, "PostgreSQL");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn checkpoint_not_found() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let cp = Checkpointer::new(store);
|
||||
|
||||
let err = cp.rollback_load("nonexistent", "ckpt_x").await.unwrap_err();
|
||||
match err {
|
||||
EngineError::CheckpointNotFound(_) => {}
|
||||
other => panic!("expected CheckpointNotFound, got {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_checkpoints_returns_desc() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let cp = Checkpointer::new(store);
|
||||
|
||||
let session = new_session_for_test("list-session");
|
||||
let ckpt_id_1 = cp.checkpoint(&session).await.unwrap();
|
||||
// 短暂 sleep 确保时间戳不同(InMemoryStore 内部用 OffsetDateTime 精度到 ns)
|
||||
tokio::time::sleep(std::time::Duration::from_millis(2)).await;
|
||||
let ckpt_id_2 = cp.checkpoint(&session).await.unwrap();
|
||||
let ckpt_id_3 = cp.checkpoint(&session).await.unwrap();
|
||||
|
||||
let metas = cp.list_checkpoints("list-session").await.unwrap();
|
||||
assert_eq!(metas.len(), 3);
|
||||
// 降序:最新在前
|
||||
let ids: Vec<_> = metas.iter().map(|m| m.ckpt_id.clone()).collect();
|
||||
assert_eq!(ids[0], ckpt_id_3);
|
||||
assert_eq!(ids[1], ckpt_id_2);
|
||||
assert_eq!(ids[2], ckpt_id_1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delete_all_removes_checkpoints() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let cp = Checkpointer::new(store.clone());
|
||||
|
||||
let session = new_session_for_test("del-session");
|
||||
cp.checkpoint(&session).await.unwrap();
|
||||
cp.checkpoint(&session).await.unwrap();
|
||||
assert_eq!(cp.list_checkpoints("del-session").await.unwrap().len(), 2);
|
||||
|
||||
cp.delete_all("del-session").await.unwrap();
|
||||
assert_eq!(cp.list_checkpoints("del-session").await.unwrap().len(), 0);
|
||||
|
||||
// delete_all 幂等:再次调用不报错
|
||||
cp.delete_all("del-session").await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn latest_snapshot_returns_most_recent() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let cp = Checkpointer::new(store);
|
||||
|
||||
let mut session = new_session_for_test("latest-session");
|
||||
cp.checkpoint(&session).await.unwrap();
|
||||
tokio::time::sleep(std::time::Duration::from_millis(2)).await;
|
||||
session
|
||||
.set_session_data("v", "2")
|
||||
.await
|
||||
.unwrap();
|
||||
cp.checkpoint(&session).await.unwrap();
|
||||
|
||||
let latest = cp.latest_snapshot("latest-session").await.unwrap().unwrap();
|
||||
let v = latest.session_memory_data.get("v").unwrap();
|
||||
assert_eq!(v.value, "2");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn checkpoint_isolation_between_sessions() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let cp = Checkpointer::new(store);
|
||||
|
||||
let a = new_session_for_test("iso-a");
|
||||
let b = new_session_for_test("iso-b");
|
||||
cp.checkpoint(&a).await.unwrap();
|
||||
cp.checkpoint(&b).await.unwrap();
|
||||
|
||||
let metas_a = cp.list_checkpoints("iso-a").await.unwrap();
|
||||
let metas_b = cp.list_checkpoints("iso-b").await.unwrap();
|
||||
assert_eq!(metas_a.len(), 1);
|
||||
assert_eq!(metas_b.len(), 1);
|
||||
assert_eq!(metas_a[0].session_id, "iso-a");
|
||||
assert_eq!(metas_b[0].session_id, "iso-b");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
//! Engine 模块统一错误类型。
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::agent::error::AgentError;
|
||||
use crate::memory::error::MemoryError;
|
||||
|
||||
/// Engine 模块错误枚举。
|
||||
///
|
||||
/// - `Session*` / `Checkpoint*`:engine 层独有的错误变体
|
||||
/// - `Memory`:透传 `MemoryError`,与项目既有 `AgentError` 风格一致(对比 `AgentError::Memory`)
|
||||
/// - `Serialization`:快照 JSON 解析失败
|
||||
/// - `Agent`:透传 `AgentError`(后续 Stage 5/6 集成时需要)
|
||||
#[derive(Debug, Error)]
|
||||
#[non_exhaustive]
|
||||
pub enum EngineError {
|
||||
/// 指定 session_id 不存在。
|
||||
/// 适用场景:`get()` 内存未命中、`create_child()` parent 不存在、
|
||||
/// `recover()` 存储中查不到。
|
||||
///
|
||||
/// **不适用** `destroy()`:`destroy()` 对不存在的 session 静默返回 `Ok(())`
|
||||
/// (幂等删除语义,调用方无需先检查)。
|
||||
#[error("Session not found: {0}")]
|
||||
SessionNotFound(String),
|
||||
|
||||
/// 创建 session 时 ID 已存在(自动生成 UUID 时通常不会触发;当前主要在重复 `recover` 已存在 ID 时使用)。
|
||||
#[error("Session already exists: {0}")]
|
||||
SessionAlreadyExists(String),
|
||||
|
||||
/// 指定 ckpt_id 不存在。
|
||||
#[error("Checkpoint not found: {0}")]
|
||||
CheckpointNotFound(String),
|
||||
|
||||
/// 存储错误(透传 `MemoryError`)。
|
||||
/// Checkpointer 和 SessionManager 的所有 `MemoryStore` 操作通过此变体传播错误。
|
||||
#[error("存储错误: {0}")]
|
||||
Memory(#[from] MemoryError),
|
||||
|
||||
/// 序列化/反序列化失败(serde_json / snapshot 格式错误)。
|
||||
#[error("序列化错误: {0}")]
|
||||
Serialization(String),
|
||||
|
||||
/// Agent 错误(透传 `AgentError`,供后续 Stage 5/6 的 `recover`/`replace` 等集成入口使用)。
|
||||
#[error("Agent 错误: {0}")]
|
||||
Agent(#[from] AgentError),
|
||||
|
||||
/// 子代理调度失败(`dispatch` 过程中遇到不可恢复错误,子 session 已被清理)。
|
||||
/// 调用方收到此错误时,子 session 已通过 `destroy()` 清理(SessionMeta + checkpoint 全部清空)。
|
||||
#[error("Dispatch failed: {0}")]
|
||||
DispatchFailed(String),
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
//! Engine 模块 —— Agent 执行引擎。
|
||||
//!
|
||||
//! Phase 17 新增。提供 SessionManager(会话树管理)和 Checkpointer(time-travel 检查点)能力。
|
||||
//!
|
||||
//! ## 子模块
|
||||
//!
|
||||
//! - [`session_manager`]:SessionManager + SessionManagerConfig
|
||||
//! - [`checkpointer`]:Checkpointer(time-travel 检查点)
|
||||
//! - [`snapshot`]:SessionSnapshot + SessionMemoryEntry(可序列化快照)
|
||||
//! - [`error`]:EngineError 枚举
|
||||
|
||||
pub mod checkpointer;
|
||||
pub mod error;
|
||||
pub mod session_manager;
|
||||
pub mod snapshot;
|
||||
pub mod sub_agent;
|
||||
pub mod switch;
|
||||
|
||||
pub use checkpointer::{Checkpointer, CkptMeta};
|
||||
pub use error::EngineError;
|
||||
pub use session_manager::{SessionManager, SessionManagerConfig};
|
||||
pub use snapshot::{SessionMemoryEntry, SessionSnapshot};
|
||||
pub use sub_agent::{DispatchConfig, SubTaskResult, SubTaskStreamEvent};
|
||||
@@ -0,0 +1,908 @@
|
||||
//! SessionManager —— Session 生命周期管理器(Phase 17 Step 5)。
|
||||
//!
|
||||
//! 组合持有 [`Checkpointer`],提供 session 的 CRUD、树形关系查询和检查点集成。
|
||||
//! 内部用 `tokio::sync::RwLock<HashMap>` 管理活跃 session。
|
||||
//!
|
||||
//! ## 锁契约
|
||||
//!
|
||||
//! - 所有写操作(`create`/`destroy`/`replace`)内部**先完成 HashMap 操作**(持写锁),
|
||||
//! 释放 RwLock 后再调用 Checkpointer/MemoryStore 的异步 I/O。
|
||||
//! - `get()` 返回 `Arc<Mutex<AgentSession>>` 后**立即释放 RwLock 读锁**,
|
||||
//! 调用方持有的是 session 级别的 Mutex 锁而非管理器级别的锁。
|
||||
//! - **不持有 RwLock 跨越 `.await`** —— 所有 .await 点必须在 RwLock guard drop 之后。
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use time::OffsetDateTime;
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
|
||||
use crate::agent::agent::Agent;
|
||||
use crate::agent::runtime::RuntimeBundle;
|
||||
use crate::agent::session::AgentSession;
|
||||
use crate::engine::checkpointer::Checkpointer;
|
||||
use crate::engine::error::EngineError;
|
||||
use crate::memory::store::MemoryStore;
|
||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||
|
||||
/// Session 元数据(持久化到 `session:{session_id}:meta`)。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub(crate) struct SessionMeta {
|
||||
pub session_id: String,
|
||||
pub agent_name: String,
|
||||
pub parent_id: Option<String>,
|
||||
pub created_at: u64, // Unix 时间戳秒
|
||||
pub turn_count: u32,
|
||||
}
|
||||
|
||||
impl SessionMeta {
|
||||
fn meta_key(session_id: &str) -> String {
|
||||
format!("session:{}:meta", session_id)
|
||||
}
|
||||
|
||||
fn assert_no_colon(id: &str, field: &str) {
|
||||
if id.contains(':') {
|
||||
panic!(
|
||||
"{field} '{id}' contains ':' which would break key format. \
|
||||
Use only letters, digits, hyphens and underscores."
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn from_session(session: &AgentSession, parent_id: Option<String>) -> Self {
|
||||
Self::assert_no_colon(&session.session_id, "session_id");
|
||||
Self {
|
||||
session_id: session.session_id.clone(),
|
||||
agent_name: session.agent.name().to_string(),
|
||||
parent_id,
|
||||
created_at: SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_secs())
|
||||
.unwrap_or(0),
|
||||
turn_count: session.turn_index(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// SessionManager 配置。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SessionManagerConfig {
|
||||
/// 是否在 submit_turn 后自动 checkpoint(Step 6 集成)。
|
||||
pub auto_checkpoint: bool,
|
||||
/// 可选的默认 bundle,用于 `recover` 时的 bundle 注入。
|
||||
pub default_bundle: Option<Arc<RuntimeBundle>>,
|
||||
}
|
||||
|
||||
impl Default for SessionManagerConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
auto_checkpoint: true,
|
||||
default_bundle: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Session 生命周期管理器。
|
||||
///
|
||||
/// 内部 `RwLock<HashMap>`:读多写少场景优化。写操作先持写锁完成 HashMap 更新后立即释放,
|
||||
/// 再调用 Checkpointer/MemoryStore 的异步 I/O。
|
||||
pub struct SessionManager {
|
||||
pub(crate) sessions: RwLock<HashMap<String, Arc<Mutex<AgentSession>>>>,
|
||||
pub(crate) checkpointer: Checkpointer,
|
||||
pub(crate) store: Arc<dyn MemoryStore>,
|
||||
pub(crate) config: SessionManagerConfig,
|
||||
}
|
||||
|
||||
impl SessionManager {
|
||||
/// 构造 SessionManager(使用默认配置)。
|
||||
pub fn new(store: Arc<dyn MemoryStore>) -> Self {
|
||||
Self {
|
||||
sessions: RwLock::new(HashMap::new()),
|
||||
checkpointer: Checkpointer::new(store.clone()),
|
||||
store,
|
||||
config: SessionManagerConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 构造 SessionManager(带自定义配置)。
|
||||
pub fn with_config(store: Arc<dyn MemoryStore>, config: SessionManagerConfig) -> Self {
|
||||
Self {
|
||||
sessions: RwLock::new(HashMap::new()),
|
||||
checkpointer: Checkpointer::new(store.clone()),
|
||||
store,
|
||||
config,
|
||||
}
|
||||
}
|
||||
|
||||
/// 暴露 Checkpointer 引用(调用方可直接操作检查点)。
|
||||
pub fn checkpointer(&self) -> &Checkpointer {
|
||||
&self.checkpointer
|
||||
}
|
||||
|
||||
/// 暴露 MemoryStore 引用。
|
||||
pub fn store(&self) -> &Arc<dyn MemoryStore> {
|
||||
&self.store
|
||||
}
|
||||
|
||||
// ====== 内部辅助:SessionMeta 持久化 ======
|
||||
|
||||
pub(crate) async fn save_session_meta(&self, meta: &SessionMeta) -> Result<(), EngineError> {
|
||||
let json = serde_json::to_string(meta)
|
||||
.map_err(|e| EngineError::Serialization(format!("SessionMeta serialize: {e}")))?;
|
||||
let item = MemoryItem {
|
||||
id: SessionMeta::meta_key(&meta.session_id),
|
||||
content: json,
|
||||
metadata: serde_json::json!({}),
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
};
|
||||
self.store.save(item).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) async fn load_session_meta(&self, session_id: &str) -> Result<Option<SessionMeta>, EngineError> {
|
||||
let item = self
|
||||
.store
|
||||
.get(&SessionMeta::meta_key(session_id))
|
||||
.await?;
|
||||
match item {
|
||||
Some(item) => {
|
||||
let meta: SessionMeta = serde_json::from_str(&item.content).map_err(|e| {
|
||||
EngineError::Serialization(format!("SessionMeta deserialize: {e}"))
|
||||
})?;
|
||||
Ok(Some(meta))
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
// ====== 公开 API ======
|
||||
|
||||
/// 创建新 session。session_id 内部自动生成(时间戳+计数器,ponytail)。
|
||||
///
|
||||
/// 流程:生成 session_id → `AgentSession::new` → 存 `SessionMeta` → 注册到 HashMap。
|
||||
pub async fn create(
|
||||
&self,
|
||||
agent: Arc<dyn Agent>,
|
||||
bundle: Arc<RuntimeBundle>,
|
||||
) -> Result<String, EngineError> {
|
||||
let session_id = Self::generate_session_id();
|
||||
|
||||
let session = AgentSession::new(agent, &session_id, bundle);
|
||||
let meta = SessionMeta::from_session(&session, None);
|
||||
|
||||
// 1. 存 SessionMeta(持久层)—— 在 lock 外做 I/O
|
||||
self.save_session_meta(&meta).await?;
|
||||
|
||||
// 2. 注册到 HashMap —— 短暂持写锁
|
||||
{
|
||||
let mut sessions = self.sessions.write().await;
|
||||
sessions.insert(session_id.clone(), Arc::new(Mutex::new(session)));
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
agent_name = %meta.agent_name,
|
||||
"session created"
|
||||
);
|
||||
Ok(session_id)
|
||||
}
|
||||
|
||||
/// 从父 session 创建子 session。
|
||||
///
|
||||
/// 继承父的 `RuntimeBundle`(`Arc::clone` 共享引用)。
|
||||
/// session_id 内部自动生成。
|
||||
///
|
||||
/// **限制**:父 session **必须已加载到内存**(通过 `get()` 或 `recover()`)。
|
||||
/// 因为 `RuntimeBundle` 不可序列化,bundle 必须从内存中的父 session 获取。
|
||||
/// 父 session 已被 `destroy` 或冷启动后未加载时,本方法返回 `SessionNotFound`。
|
||||
///
|
||||
/// 如果 `parent_id` 在存储中查不到 SessionMeta,返回 `EngineError::SessionNotFound(parent_id)`。
|
||||
pub async fn create_child(
|
||||
&self,
|
||||
parent_id: &str,
|
||||
agent: Arc<dyn Agent>,
|
||||
) -> Result<String, EngineError> {
|
||||
// 验证 parent 存在(从存储读 SessionMeta,避免依赖内存状态)
|
||||
let parent_meta = self
|
||||
.load_session_meta(parent_id)
|
||||
.await?
|
||||
.ok_or_else(|| EngineError::SessionNotFound(parent_id.to_string()))?;
|
||||
|
||||
// 读取父 session 的 bundle(必须在内存中才能拿到;如果不在内存则要求用户先 get)
|
||||
let parent_bundle = {
|
||||
let sessions = self.sessions.read().await;
|
||||
let parent_arc = sessions
|
||||
.get(parent_id)
|
||||
.ok_or_else(|| EngineError::SessionNotFound(parent_id.to_string()))?;
|
||||
Arc::clone(parent_arc.lock().await.bundle())
|
||||
};
|
||||
|
||||
let session_id = Self::generate_session_id();
|
||||
let session = AgentSession::new(agent, &session_id, parent_bundle);
|
||||
let meta = SessionMeta::from_session(&session, Some(parent_meta.session_id));
|
||||
|
||||
self.save_session_meta(&meta).await?;
|
||||
|
||||
{
|
||||
let mut sessions = self.sessions.write().await;
|
||||
sessions.insert(session_id.clone(), Arc::new(Mutex::new(session)));
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
parent_id = %parent_id,
|
||||
"child session created"
|
||||
);
|
||||
Ok(session_id)
|
||||
}
|
||||
|
||||
/// 按 ID 获取 session(**仅查内存**,不自动从存储恢复)。
|
||||
///
|
||||
/// 冷启动时 `get()` 未命中返回 `SessionNotFound`。如需从存储恢复,使用 `recover()` 方法。
|
||||
pub async fn get(
|
||||
&self,
|
||||
session_id: &str,
|
||||
) -> Result<Arc<Mutex<AgentSession>>, EngineError> {
|
||||
let sessions = self.sessions.read().await;
|
||||
let result = sessions.get(session_id).cloned();
|
||||
tracing::debug!(
|
||||
session_id = %session_id,
|
||||
found = result.is_some(),
|
||||
"session get"
|
||||
);
|
||||
result.ok_or_else(|| EngineError::SessionNotFound(session_id.to_string()))
|
||||
}
|
||||
|
||||
/// 从存储恢复 session。
|
||||
///
|
||||
/// 流程:读 SessionMeta → 从 latest checkpoint 读 SessionSnapshot →
|
||||
/// `AgentSession::from_snapshot(snapshot, agent, bundle)` → `restore_memory` →
|
||||
/// 注册到 HashMap。
|
||||
pub async fn recover(
|
||||
&self,
|
||||
session_id: &str,
|
||||
agent: Arc<dyn Agent>,
|
||||
bundle: Arc<RuntimeBundle>,
|
||||
) -> Result<Arc<Mutex<AgentSession>>, EngineError> {
|
||||
// 内存中已存在 → 拒绝(避免覆盖丢失数据)
|
||||
{
|
||||
let sessions = self.sessions.read().await;
|
||||
if sessions.contains_key(session_id) {
|
||||
return Err(EngineError::SessionAlreadyExists(session_id.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
// 读 SessionMeta(如果不存在则报错)
|
||||
let meta = self
|
||||
.load_session_meta(session_id)
|
||||
.await?
|
||||
.ok_or_else(|| EngineError::SessionNotFound(session_id.to_string()))?;
|
||||
|
||||
// 从 latest checkpoint 读 SessionSnapshot
|
||||
let snapshot = self
|
||||
.checkpointer
|
||||
.latest_snapshot(session_id)
|
||||
.await?
|
||||
.ok_or_else(|| {
|
||||
EngineError::CheckpointNotFound(format!(
|
||||
"no checkpoint for session_id={session_id}"
|
||||
))
|
||||
})?;
|
||||
|
||||
// 同步重建 + 异步写回
|
||||
// 注意:restore_memory 现在是 &mut self(清空 pending_memory_restore),
|
||||
// 需要先 Arc<Mutex<>> 包装后再 lock + 调用
|
||||
let session = AgentSession::from_snapshot(snapshot, agent, bundle)?;
|
||||
let arc = Arc::new(Mutex::new(session));
|
||||
{
|
||||
let mut guard = arc.lock().await;
|
||||
guard.restore_memory().await?;
|
||||
}
|
||||
|
||||
{
|
||||
let mut sessions = self.sessions.write().await;
|
||||
sessions.insert(session_id.to_string(), arc.clone());
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
turn_index = meta.turn_count,
|
||||
"session recovered from storage"
|
||||
);
|
||||
Ok(arc)
|
||||
}
|
||||
|
||||
/// 替换 SessionManager 中指定 session_id 的 AgentSession 实例。
|
||||
///
|
||||
/// 用于 `Checkpointer::rollback_load() + from_snapshot + restore_memory` 后的无缝切换。
|
||||
///
|
||||
/// 内部执行:
|
||||
/// 1. 写回 SessionMeta
|
||||
/// 2. 调用 `session.restore_memory()` 写回持久层(`&mut self` 调用会清空 pending)
|
||||
/// 3. 内存替换
|
||||
pub async fn replace(
|
||||
&self,
|
||||
session_id: &str,
|
||||
mut session: AgentSession,
|
||||
) -> Result<(), EngineError> {
|
||||
// 1. 写回 SessionMeta(取新 session 的 turn_index)
|
||||
let meta = SessionMeta::from_session(&session, None);
|
||||
self.save_session_meta(&meta).await?;
|
||||
|
||||
// 2. restore_memory 写回持久层(pending_memory_restore → None)
|
||||
session.restore_memory().await?;
|
||||
|
||||
// 3. 替换内存中的 session
|
||||
let mut sessions = self.sessions.write().await;
|
||||
sessions.insert(session_id.to_string(), Arc::new(Mutex::new(session)));
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
turn_index = meta.turn_count,
|
||||
"session replaced"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 封装 `AgentSession::submit_turn`:自动加锁 + 可选自动 checkpoint。
|
||||
///
|
||||
/// 流程:
|
||||
/// 1. `get(session_id)` 获取 session
|
||||
/// 2. lock + `session.submit_turn(user_input)`
|
||||
/// 3. 如果 `config.auto_checkpoint == true`,同步调用 `checkpointer.checkpoint(&session).await`
|
||||
/// - checkpoint 失败时通过 `tracing::error!` 记录,不阻断 `Ok` 返回
|
||||
/// - 调用方如需强持久化保证,应显式调用 `checkpointer.checkpoint()` 并处理 `Result`
|
||||
pub async fn submit_turn(
|
||||
&self,
|
||||
session_id: &str,
|
||||
user_input: impl Into<String>,
|
||||
) -> Result<crate::llm::types::response_v2::MessageResponse, EngineError> {
|
||||
let session = self.get(session_id).await?;
|
||||
let response = {
|
||||
let mut guard = session.lock().await;
|
||||
guard
|
||||
.submit_turn(user_input)
|
||||
.await
|
||||
.map_err(EngineError::from)?
|
||||
};
|
||||
|
||||
// 自动 checkpoint(在 lock 外做 I/O)
|
||||
if self.config.auto_checkpoint {
|
||||
let snapshot_session = session.lock().await;
|
||||
if let Err(e) = self.checkpointer.checkpoint(&snapshot_session).await {
|
||||
tracing::error!(
|
||||
session_id = %session_id,
|
||||
error = %e,
|
||||
"auto_checkpoint failed; submit_turn result already returned"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
/// 封装 `AgentSession::submit_turn_stream`:流式 API + 可选自动 checkpoint。
|
||||
///
|
||||
/// 与 `submit_turn` 的差异:自动 checkpoint 推迟到 `finalize_turn_stream` 调用时。
|
||||
/// 流期间不创建 checkpoint,避免客户端断开导致半成品 checkpoint 污染。
|
||||
pub async fn submit_turn_stream(
|
||||
&self,
|
||||
session_id: &str,
|
||||
user_input: impl Into<String>,
|
||||
) -> Result<
|
||||
std::pin::Pin<
|
||||
Box<
|
||||
dyn futures_core::Stream<Item = crate::llm::stream::StreamEvent>
|
||||
+ Send,
|
||||
>,
|
||||
>,
|
||||
EngineError,
|
||||
> {
|
||||
let session = self.get(session_id).await?;
|
||||
let stream = {
|
||||
let mut guard = session.lock().await;
|
||||
guard
|
||||
.submit_turn_stream(user_input)
|
||||
.await
|
||||
.map_err(EngineError::from)?
|
||||
};
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
/// 流消费完成后调用:累计 cost + 触发 OnTurnEnd + 自动 checkpoint(如启用)。
|
||||
///
|
||||
/// 委托给 `AgentSession::finalize_turn`,然后在 lock 外执行 auto_checkpoint。
|
||||
pub async fn finalize_turn_stream(
|
||||
&self,
|
||||
session_id: &str,
|
||||
response: &crate::llm::types::response_v2::MessageResponse,
|
||||
new_messages_from_cycle: Vec<crate::llm::types::message::Message>,
|
||||
) -> Result<(), EngineError> {
|
||||
let session = self.get(session_id).await?;
|
||||
{
|
||||
let mut guard = session.lock().await;
|
||||
guard
|
||||
.finalize_turn(response, new_messages_from_cycle)
|
||||
.await
|
||||
.map_err(EngineError::from)?;
|
||||
}
|
||||
|
||||
if self.config.auto_checkpoint {
|
||||
let snapshot_session = session.lock().await;
|
||||
if let Err(e) = self.checkpointer.checkpoint(&snapshot_session).await {
|
||||
tracing::error!(
|
||||
session_id = %session_id,
|
||||
error = %e,
|
||||
"auto_checkpoint failed after finalize_turn"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 查询某 parent 的所有直接子 session 的 ID 列表。
|
||||
///
|
||||
/// 实现:prefix 查询所有 `session:*:meta`,过滤 `parent_id == parent_id`。
|
||||
// ponytail: O(N) 全表扫描,当前规模(≤10K session)可接受。
|
||||
// 如有性能需求,可维护 `parent:{parent_id}:children` 索引 key 替代扫描。
|
||||
pub async fn children(&self, parent_id: &str) -> Result<Vec<String>, EngineError> {
|
||||
let filter = MemoryFilter {
|
||||
prefix: Some("session:".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let items = self.store.list(&filter).await?;
|
||||
|
||||
let mut child_ids = Vec::new();
|
||||
for item in items {
|
||||
// 解析 SessionMeta JSON,过滤 parent_id
|
||||
if let Ok(meta) = serde_json::from_str::<SessionMeta>(&item.content)
|
||||
&& meta.parent_id.as_deref() == Some(parent_id)
|
||||
{
|
||||
child_ids.push(meta.session_id);
|
||||
}
|
||||
}
|
||||
Ok(child_ids)
|
||||
}
|
||||
|
||||
/// 查询某 child session 的 parent ID。
|
||||
///
|
||||
/// 如果 parent 已被销毁,返回 `Ok(None)`(允许孤儿 session 存在)。
|
||||
pub async fn parent(&self, child_id: &str) -> Result<Option<String>, EngineError> {
|
||||
let meta = self.load_session_meta(child_id).await?;
|
||||
let parent_id = match meta {
|
||||
Some(m) => m.parent_id,
|
||||
None => return Ok(None),
|
||||
};
|
||||
// 如果 parent_id 已被 destroy,load_session_meta 返回 None → 返回 Ok(None)
|
||||
match parent_id {
|
||||
Some(pid) => {
|
||||
let parent_meta = self.load_session_meta(&pid).await?;
|
||||
Ok(parent_meta.map(|_| pid))
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// 销毁 session:从内存移除 + 清理 SessionMeta + 清理检查点。
|
||||
///
|
||||
/// **孤儿策略**:允许孤儿 session 存在(子 session 的 `parent_id` 仍指向已删除的父,
|
||||
/// 但 `parent()` 返回 `None`)。不递归删除子 session。
|
||||
///
|
||||
/// **幂等性**:对不存在的 session 静默返回 `Ok(())`(不报错)。
|
||||
/// `MemoryStore::delete()` 和 `Checkpointer::delete_all()` 本身幂等。
|
||||
/// 调用方无需先 `get()` 检查存在性。
|
||||
pub async fn destroy(&self, session_id: &str) -> Result<(), EngineError> {
|
||||
// 从内存移除
|
||||
{
|
||||
let mut sessions = self.sessions.write().await;
|
||||
sessions.remove(session_id);
|
||||
}
|
||||
|
||||
// 删除 SessionMeta
|
||||
self.store
|
||||
.delete(&SessionMeta::meta_key(session_id))
|
||||
.await?;
|
||||
|
||||
// 删除所有 checkpoints
|
||||
self.checkpointer.delete_all(session_id).await?;
|
||||
|
||||
tracing::info!(session_id = %session_id, "session destroyed");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 内部辅助:生成 session_id(纳秒+计数器)。
|
||||
fn generate_session_id() -> String {
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
static COUNTER: AtomicU64 = AtomicU64::new(0);
|
||||
let nanos = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_nanos() as u64)
|
||||
.unwrap_or(0);
|
||||
let counter = COUNTER.fetch_add(1, Ordering::Relaxed);
|
||||
format!("sess-{:x}-{:x}", nanos, counter)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::agent::builder::AgentBuilder;
|
||||
use crate::llm::hooks::HookExecutor;
|
||||
use crate::llm::mock::MockProvider;
|
||||
use crate::tools::ToolRegistry;
|
||||
|
||||
struct StubAgent(String);
|
||||
#[async_trait::async_trait]
|
||||
impl Agent for StubAgent {
|
||||
fn name(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn make_bundle() -> Arc<RuntimeBundle> {
|
||||
Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(Arc::new(MockProvider::new(vec![])))
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.build()
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_and_get() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a1".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(!id.is_empty());
|
||||
|
||||
let session = sm.get(&id).await.unwrap();
|
||||
assert_eq!(session.lock().await.session_id, id);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_not_found() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let err = sm.get("missing").await.unwrap_err();
|
||||
match err {
|
||||
EngineError::SessionNotFound(_) => {}
|
||||
other => panic!("expected SessionNotFound, got {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn destroy_removes_meta_and_checkpoints() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store.clone());
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a1".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
let session = sm.get(&id).await.unwrap();
|
||||
sm.checkpointer().checkpoint(&*session.lock().await).await.unwrap();
|
||||
assert_eq!(sm.checkpointer().list_checkpoints(&id).await.unwrap().len(), 1);
|
||||
|
||||
sm.destroy(&id).await.unwrap();
|
||||
|
||||
// 内存中查不到
|
||||
assert!(sm.get(&id).await.is_err());
|
||||
// SessionMeta 已删除
|
||||
assert!(sm.load_session_meta(&id).await.unwrap().is_none());
|
||||
// Checkpoint 已删除
|
||||
assert_eq!(sm.checkpointer().list_checkpoints(&id).await.unwrap().len(), 0);
|
||||
|
||||
// destroy 不存在的 session 不报错
|
||||
sm.destroy(&id).await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_child_inherits_parent_bundle() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let parent_id = sm
|
||||
.create(Arc::new(StubAgent("parent".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let child_id = sm
|
||||
.create_child(&parent_id, Arc::new(StubAgent("child".into())))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_ne!(child_id, parent_id);
|
||||
|
||||
let children = sm.children(&parent_id).await.unwrap();
|
||||
assert_eq!(children, vec![child_id.clone()]);
|
||||
|
||||
let parent = sm.parent(&child_id).await.unwrap();
|
||||
assert_eq!(parent, Some(parent_id.clone()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_child_parent_not_found() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let err = sm
|
||||
.create_child("nope", Arc::new(StubAgent("c".into())))
|
||||
.await
|
||||
.unwrap_err();
|
||||
match err {
|
||||
EngineError::SessionNotFound(_) => {}
|
||||
other => panic!("expected SessionNotFound, got {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn parent_returns_none_after_parent_destroyed() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let parent_id = sm
|
||||
.create(Arc::new(StubAgent("p".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
let child_id = sm
|
||||
.create_child(&parent_id, Arc::new(StubAgent("c".into())))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
sm.destroy(&parent_id).await.unwrap();
|
||||
|
||||
// 父被销毁 → parent() 返回 None(孤儿策略)
|
||||
assert_eq!(sm.parent(&child_id).await.unwrap(), None);
|
||||
// 孤儿 session 仍然存在于存储
|
||||
assert!(sm.load_session_meta(&child_id).await.unwrap().is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn replace_after_recover() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store.clone());
|
||||
|
||||
// 1. 创建 session + checkpoint
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
{
|
||||
let s = sm.get(&id).await.unwrap();
|
||||
s.lock().await.set_session_data("k", "v1").await.unwrap();
|
||||
sm.checkpointer().checkpoint(&*s.lock().await).await.unwrap();
|
||||
}
|
||||
|
||||
// 2. 模拟"进程重启"——清空内存但保留 store
|
||||
let sm2 = SessionManager::new(store.clone());
|
||||
let bundle = make_bundle();
|
||||
|
||||
// 3. recover
|
||||
let recovered = sm2
|
||||
.recover(&id, Arc::new(StubAgent("a".into())), bundle)
|
||||
.await
|
||||
.unwrap();
|
||||
let recovered_session = recovered.lock().await;
|
||||
let v = recovered_session.get_session_data("k").await.unwrap();
|
||||
assert_eq!(v, Some("v1".into()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn recover_session_already_in_memory() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store.clone());
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
{
|
||||
let s = sm.get(&id).await.unwrap();
|
||||
sm.checkpointer().checkpoint(&*s.lock().await).await.unwrap();
|
||||
}
|
||||
|
||||
let err = sm
|
||||
.recover(&id, Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap_err();
|
||||
match err {
|
||||
EngineError::SessionAlreadyExists(_) => {}
|
||||
other => panic!("expected SessionAlreadyExists, got {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn children_empty() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
// 无子 session
|
||||
assert_eq!(sm.children(&id).await.unwrap(), Vec::<String>::new());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn submit_turn_no_provider_responses() {
|
||||
// 没有 LLM 响应 → MockProvider 返回 LlmError → EngineError::Llm
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let err = sm.submit_turn(&id, "hello").await.unwrap_err();
|
||||
match err {
|
||||
EngineError::Agent(crate::agent::error::AgentError::Llm(_)) => {}
|
||||
other => panic!("expected EngineError::Agent(AgentError::Llm), got {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn submit_turn_auto_checkpoint_off() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let config = SessionManagerConfig {
|
||||
auto_checkpoint: false,
|
||||
default_bundle: None,
|
||||
};
|
||||
let sm = SessionManager::with_config(store, config);
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
let s = sm.get(&id).await.unwrap();
|
||||
// 配置 auto_checkpoint=false → 即便 submit_turn 失败(MockProvider 空响应)也不会触发 checkpoint
|
||||
let _ = sm.submit_turn(&id, "x").await;
|
||||
|
||||
// 手动 checkpoint 仍可工作
|
||||
sm.checkpointer().checkpoint(&*s.lock().await).await.unwrap();
|
||||
assert_eq!(sm.checkpointer().list_checkpoints(&id).await.unwrap().len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn replace_preserves_session_id() {
|
||||
let store: Arc<dyn MemoryStore> = Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store);
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// 构造一个新 session(同 session_id)然后 replace
|
||||
let mut new_session = AgentSession::new(Arc::new(StubAgent("a".into())), &id, make_bundle());
|
||||
new_session
|
||||
.set_session_data("replaced", "yes")
|
||||
.await
|
||||
.unwrap();
|
||||
sm.replace(&id, new_session).await.unwrap();
|
||||
|
||||
let v = sm
|
||||
.get(&id)
|
||||
.await
|
||||
.unwrap()
|
||||
.lock()
|
||||
.await
|
||||
.get_session_data("replaced")
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(v, Some("yes".into()));
|
||||
}
|
||||
|
||||
// ====== 实施审查补充:3 个边界测试 ======
|
||||
|
||||
/// 序列化向前兼容:旧版 SessionSnapshot 缺少新字段时,`#[serde(default)]` 兜底生效。
|
||||
#[test]
|
||||
fn snapshot_deserialize_with_minimal_fields() {
|
||||
// 构造一个 v0.1 风格的最小 JSON(仅含核心标识字段,缺少 cost_so_far/slots/
|
||||
// session_memory_data/last_summary_turn)
|
||||
let minimal_json = r#"{
|
||||
"session_id": "legacy-session",
|
||||
"agent_name": "legacy",
|
||||
"turn_index": 5,
|
||||
"current_slot_id": "default"
|
||||
}"#;
|
||||
let snapshot: crate::engine::snapshot::SessionSnapshot =
|
||||
serde_json::from_str(minimal_json).expect("应能反序列化最小 JSON");
|
||||
|
||||
// 核心字段保留
|
||||
assert_eq!(snapshot.session_id, "legacy-session");
|
||||
assert_eq!(snapshot.agent_name, "legacy");
|
||||
assert_eq!(snapshot.turn_index, 5);
|
||||
assert_eq!(snapshot.current_slot_id, "default");
|
||||
|
||||
// 可选字段走 #[serde(default)]
|
||||
assert_eq!(snapshot.cost_so_far.total().total_tokens, 0);
|
||||
assert!(snapshot.slots.is_empty());
|
||||
assert!(snapshot.session_memory_data.is_empty());
|
||||
assert_eq!(snapshot.last_summary_turn, None);
|
||||
}
|
||||
|
||||
/// restore_memory 幂等性:第二次调用应立即返回 Ok(())(pending 已被清空)。
|
||||
#[tokio::test]
|
||||
async fn restore_memory_is_idempotent() {
|
||||
let store: Arc<dyn MemoryStore> =
|
||||
Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = SessionManager::new(store.clone());
|
||||
|
||||
let id = sm
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
{
|
||||
let s = sm.get(&id).await.unwrap();
|
||||
s.lock()
|
||||
.await
|
||||
.set_session_data("k", "v")
|
||||
.await
|
||||
.unwrap();
|
||||
sm.checkpointer().checkpoint(&*s.lock().await).await.unwrap();
|
||||
}
|
||||
|
||||
// 模拟"进程重启"——新建 SessionManager,复用 store
|
||||
let sm2 = SessionManager::new(store.clone());
|
||||
let recovered = sm2
|
||||
.recover(&id, Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// 1. recover 已调用 restore_memory → pending_memory_restore 应为 None
|
||||
assert!(!recovered.lock().await.has_pending_memory_restore());
|
||||
|
||||
// 2. 再次 restore_memory → 幂等(不会二次写入,不会 panic)
|
||||
recovered.lock().await.restore_memory().await.unwrap();
|
||||
assert!(!recovered.lock().await.has_pending_memory_restore());
|
||||
|
||||
// 3. 第三次仍然幂等
|
||||
recovered.lock().await.restore_memory().await.unwrap();
|
||||
}
|
||||
|
||||
/// 10 并发 session 创建:验证 RwLock 写锁争用下不冲突,所有 ID 唯一。
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
|
||||
async fn concurrent_create_ten_sessions() {
|
||||
let store: Arc<dyn MemoryStore> =
|
||||
Arc::new(crate::memory::store::InMemoryStore::new());
|
||||
let sm = Arc::new(SessionManager::new(store));
|
||||
|
||||
let mut handles = Vec::with_capacity(10);
|
||||
for _ in 0..10 {
|
||||
let sm_clone = Arc::clone(&sm);
|
||||
handles.push(tokio::spawn(async move {
|
||||
sm_clone
|
||||
.create(Arc::new(StubAgent("a".into())), make_bundle())
|
||||
.await
|
||||
}));
|
||||
}
|
||||
|
||||
let mut ids = Vec::with_capacity(10);
|
||||
for h in handles {
|
||||
ids.push(h.await.expect("task join").expect("create ok"));
|
||||
}
|
||||
|
||||
// 所有 ID 唯一
|
||||
let unique: std::collections::HashSet<_> = ids.iter().collect();
|
||||
assert_eq!(unique.len(), 10, "并发创建应产生 10 个唯一 session_id");
|
||||
|
||||
// 全部可 get
|
||||
for id in &ids {
|
||||
assert!(sm.get(id).await.is_ok());
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
//! SessionSnapshot —— 见下文 doc comment。Step 3 将填充完整实现。
|
||||
//! Step 2 仅占位:定义空 struct + derive,使 `engine` 模块编译通过。
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// SessionMemory 条目的可序列化形式(保留元数据与时间戳)。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct SessionMemoryEntry {
|
||||
/// 原始值(字符串)。
|
||||
pub value: String,
|
||||
/// 元数据(自由 JSON)。
|
||||
#[serde(default)]
|
||||
pub metadata: serde_json::Value,
|
||||
/// 创建时间(Unix 时间戳秒;`None` 兼容旧快照)。
|
||||
///
|
||||
/// 类型为 `i64` 而非 `u64`:`time` crate 的 `OffsetDateTime::from_unix_timestamp` 接收 `i64`,
|
||||
/// 这里保持与 `SessionMemory::set_with_meta` 签名一致,避免来回转换。
|
||||
#[serde(default)]
|
||||
pub created_at: Option<i64>,
|
||||
}
|
||||
|
||||
/// AgentSession 的可序列化快照。
|
||||
///
|
||||
/// Step 2 占位:字段已定义但未实装 to_snapshot/from_snapshot。
|
||||
/// Step 3 将基于 `#[serde(default)]` 宽松反序列化,添加 `agent_name`/`turn_index`/slot 等字段。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct SessionSnapshot {
|
||||
pub session_id: String,
|
||||
pub agent_name: String,
|
||||
pub turn_index: u32,
|
||||
#[serde(default)]
|
||||
pub cost_so_far: crate::llm::types::usage::CostTracker,
|
||||
#[serde(default)]
|
||||
pub slots: std::collections::HashMap<String, crate::agent::context::ContextSlot>,
|
||||
pub current_slot_id: String,
|
||||
pub last_summary_turn: Option<u32>,
|
||||
#[serde(default)]
|
||||
pub session_memory_data: std::collections::HashMap<String, SessionMemoryEntry>,
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,222 @@
|
||||
//! Agent 角色热切换(Phase 18)。
|
||||
//!
|
||||
//! 提供 `SessionManager::switch_agent()`,运行时替换 session 绑定的 Agent,
|
||||
//! 保留 slot 历史 / turn_index / session_memory / cost_so_far。
|
||||
//!
|
||||
//! **未来扩展方向**(roadmap 备注):
|
||||
//! - v0.4 可考虑提供 `switch_agent_with_rollback`,在切换前自动 checkpoint
|
||||
//! - 可在 `SessionMeta` 中记录 `previous_agent_name` 支持审计历史
|
||||
//! - 可新增 `switch_history` API 暴露切换时间序列
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::agent::Agent;
|
||||
use crate::engine::error::EngineError;
|
||||
use crate::engine::session_manager::{SessionManager, SessionMeta};
|
||||
|
||||
impl SessionManager {
|
||||
/// 热切换指定 session 的 Agent 角色。
|
||||
///
|
||||
/// # 行为
|
||||
///
|
||||
/// - **保留上下文**:slot 历史 / turn_index / session_memory / cost_so_far 全部保留
|
||||
/// - **更新 SessionMeta**:`agent_name` 替换为新 agent,`created_at`/`parent_id` 保持原始
|
||||
/// - **不自动 checkpoint**:与 `auto_checkpoint` 语义一致(仅 `submit_turn`/`finalize_turn` 触发)
|
||||
///
|
||||
/// # 注意
|
||||
///
|
||||
/// 切换后**新的 system_prompt 将与已有对话历史共存**。建议在切换后
|
||||
/// 发送一条明确的上下文过渡提示(如"你现在以新角色 X 的身份继续对话")
|
||||
/// 作为切换后的首条输入,以避免 LLM 误解对话历史。
|
||||
///
|
||||
/// # 安全提示
|
||||
///
|
||||
/// `AgentSession.agent` 是 `pub` 字段可直接访问。**绕过 `switch_agent`
|
||||
/// 直接修改会导致 SessionMeta 中的 `agent_name` 与内存状态不一致**,
|
||||
/// 请始终使用此方法。
|
||||
pub async fn switch_agent(
|
||||
&self,
|
||||
session_id: &str,
|
||||
new_agent: Arc<dyn Agent>,
|
||||
) -> Result<(), EngineError> {
|
||||
// 1. get session(RwLock 读锁,返回后释放)
|
||||
let session = self.get(session_id).await?;
|
||||
|
||||
// 2. lock Mutex,替换 agent,读 name + turn_index
|
||||
let (agent_name, turn_index) = {
|
||||
let mut guard = session.lock().await;
|
||||
guard.agent = new_agent;
|
||||
(guard.agent.name().to_string(), guard.turn_index())
|
||||
}; // 释放 Mutex
|
||||
|
||||
// 3. 读取原始 SessionMeta(用于保留 created_at / parent_id)
|
||||
let existing_meta = self
|
||||
.load_session_meta(session_id)
|
||||
.await?
|
||||
.ok_or_else(|| EngineError::SessionNotFound(session_id.to_string()))?;
|
||||
|
||||
// 4. 构造新 meta 并持久化(I/O,无锁)
|
||||
let meta = SessionMeta {
|
||||
session_id: session_id.to_string(),
|
||||
agent_name,
|
||||
parent_id: existing_meta.parent_id,
|
||||
created_at: existing_meta.created_at,
|
||||
turn_count: turn_index,
|
||||
};
|
||||
self.save_session_meta(&meta).await?;
|
||||
|
||||
tracing::info!(
|
||||
session_id = %session_id,
|
||||
agent_name = %meta.agent_name,
|
||||
previous_agent = %existing_meta.agent_name,
|
||||
"agent switched"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::agent::Agent;
|
||||
use crate::agent::builder::AgentBuilder;
|
||||
use crate::agent::runtime::RuntimeBundle;
|
||||
use crate::engine::session_manager::SessionManager;
|
||||
use crate::llm::hooks::HookExecutor;
|
||||
use crate::llm::mock::MockProvider;
|
||||
use crate::llm::types::Usage;
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
use crate::llm::types::response_v2::{MessageResponse, StopReason};
|
||||
use crate::memory::store::InMemoryStore;
|
||||
use crate::tools::ToolRegistry;
|
||||
use std::sync::Arc;
|
||||
|
||||
/// 测试用 MockAgent(name + system_prompt 可控)
|
||||
struct MockAgent {
|
||||
name: String,
|
||||
system_prompt: String,
|
||||
}
|
||||
|
||||
impl MockAgent {
|
||||
fn new(name: &str, system_prompt: &str) -> Self {
|
||||
Self {
|
||||
name: name.to_string(),
|
||||
system_prompt: system_prompt.to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Agent for MockAgent {
|
||||
fn name(&self) -> &str {
|
||||
&self.name
|
||||
}
|
||||
fn system_prompt(&self) -> Option<&str> {
|
||||
Some(&self.system_prompt)
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_text(text: &str) -> MessageResponse {
|
||||
MessageResponse {
|
||||
id: String::new(),
|
||||
model: String::new(),
|
||||
message: Message::Assistant {
|
||||
content: vec![ContentBlock::Text { text: text.into() }],
|
||||
},
|
||||
usage: Usage::from_input_output(8, 4),
|
||||
stop_reason: StopReason::Stop,
|
||||
extra: Default::default(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn make_manager_and_bundle() -> (Arc<SessionManager>, Arc<RuntimeBundle>) {
|
||||
let store: Arc<dyn crate::memory::store::MemoryStore> =
|
||||
Arc::new(InMemoryStore::new());
|
||||
let provider = Arc::new(MockProvider::new(vec![
|
||||
assistant_text("response1"),
|
||||
assistant_text("response2"),
|
||||
assistant_text("response3"),
|
||||
]));
|
||||
let bundle = Arc::new(
|
||||
AgentBuilder::new()
|
||||
.provider(provider)
|
||||
.tool_registry(Arc::new(ToolRegistry::new()))
|
||||
.hook_executor(Arc::new(HookExecutor::new()))
|
||||
.session_memory_backend(store.clone())
|
||||
.build()
|
||||
.expect("RuntimeBundle build"),
|
||||
);
|
||||
let sm = Arc::new(SessionManager::new(store));
|
||||
(sm, bundle)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_switch_agent_basic() {
|
||||
let (sm, bundle) = make_manager_and_bundle().await;
|
||||
let agent_a: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||||
let agent_b: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_b", "I am B"));
|
||||
|
||||
let sid = sm.create(agent_a, bundle).await.unwrap();
|
||||
sm.switch_agent(&sid, agent_b).await.unwrap();
|
||||
|
||||
let session = sm.get(&sid).await.unwrap();
|
||||
let guard = session.lock().await;
|
||||
assert_eq!(guard.agent.name(), "agent_b");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_switch_agent_preserves_context() {
|
||||
let (sm, bundle) = make_manager_and_bundle().await;
|
||||
let agent_a: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||||
let agent_b: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_b", "I am B"));
|
||||
|
||||
let sid = sm.create(agent_a, bundle).await.unwrap();
|
||||
|
||||
// 在 switch 前写入 session_data 并提交一轮
|
||||
{
|
||||
let session = sm.get(&sid).await.unwrap();
|
||||
let mut guard = session.lock().await;
|
||||
guard
|
||||
.set_session_data("key1", "value1")
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
sm.submit_turn(&sid, "hello").await.unwrap();
|
||||
|
||||
// switch
|
||||
sm.switch_agent(&sid, agent_b).await.unwrap();
|
||||
|
||||
// 验证 turn_index 保留
|
||||
let session = sm.get(&sid).await.unwrap();
|
||||
let guard = session.lock().await;
|
||||
assert_eq!(guard.turn_index(), 1, "turn_index should be preserved");
|
||||
// 验证 session_memory 保留
|
||||
let val = guard.session_memory().get("key1").await.unwrap();
|
||||
assert_eq!(val, Some("value1".to_string()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_switch_agent_updates_session_meta() {
|
||||
let (sm, bundle) = make_manager_and_bundle().await;
|
||||
let agent_a: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||||
let agent_b: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_b", "I am B"));
|
||||
|
||||
let sid = sm.create(agent_a, bundle).await.unwrap();
|
||||
sm.switch_agent(&sid, agent_b).await.unwrap();
|
||||
|
||||
// 通过 load_session_meta 验证持久化
|
||||
let meta = sm.load_session_meta(&sid).await.unwrap();
|
||||
assert!(meta.is_some(), "SessionMeta should persist");
|
||||
let meta = meta.unwrap();
|
||||
assert_eq!(meta.agent_name, "agent_b");
|
||||
// created_at / parent_id 保持
|
||||
assert_eq!(meta.parent_id, None, "parent_id should remain None");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_switch_agent_session_not_found() {
|
||||
let (sm, _bundle) = make_manager_and_bundle().await;
|
||||
let agent: Arc<dyn Agent> = Arc::new(MockAgent::new("agent_a", "I am A"));
|
||||
let result = sm.switch_agent("nonexistent_session_id", agent).await;
|
||||
assert!(matches!(result, Err(EngineError::SessionNotFound(_))));
|
||||
}
|
||||
}
|
||||
@@ -1,11 +1,15 @@
|
||||
//! agcore —— 智能体(Agent)核心工具箱。
|
||||
|
||||
pub mod agent;
|
||||
pub mod document;
|
||||
pub mod engine;
|
||||
pub mod llm;
|
||||
pub mod memory;
|
||||
pub mod prompt;
|
||||
pub mod tools;
|
||||
|
||||
pub use document::Document;
|
||||
|
||||
use tracing_subscriber::{EnvFilter, fmt, prelude::*};
|
||||
|
||||
static INIT: std::sync::Once = std::sync::Once::new();
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
pub mod compact;
|
||||
pub mod convert;
|
||||
pub mod cycle;
|
||||
pub mod embedding;
|
||||
pub mod error;
|
||||
pub mod hooks;
|
||||
pub mod mock;
|
||||
|
||||
+16
-3
@@ -18,7 +18,7 @@ use tokio_stream::wrappers::UnboundedReceiverStream;
|
||||
use crate::llm::compact::{CompactConfig, CompactState, microcompact, should_compact};
|
||||
use crate::llm::cycle::retry::should_retry;
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::hooks::{HookContext, HookExecutor};
|
||||
use crate::llm::hooks::{HookContext, HookEvent, HookExecutor};
|
||||
use crate::llm::provider::LlmProvider;
|
||||
use crate::llm::stream::StreamEvent;
|
||||
use crate::llm::types::message::{ContentBlock, Message};
|
||||
@@ -774,8 +774,21 @@ async fn run_tool_loop(
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
// ② PreRequest hook(fire-and-forget,仅占位保留以保持接口对称)
|
||||
let _ = hook_executor.as_ref();
|
||||
// ② PreRequest hook —— 与 `submit_with_tools` / `submit_stream` 行为对齐:
|
||||
// 触发 hook → 检查 should_block → 阻断则事件化 Error + return 结束 task。
|
||||
// 阻断原因透传,让消费者看到完整的拒绝原因。
|
||||
if let Some(ref executor) = hook_executor {
|
||||
let ctx = HookContext::new(HookEvent::PreRequest).with_request(&request);
|
||||
let results = executor.execute(HookEvent::PreRequest, &ctx).await;
|
||||
if let Some(blocking) = results.iter().find(|r| r.should_block) {
|
||||
let reason = blocking
|
||||
.reason
|
||||
.clone()
|
||||
.unwrap_or_else(|| "Blocked by pre-request hook".to_string());
|
||||
let _ = tx.send(StreamEvent::Error { message: reason });
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
// ③ chat_stream —— 第一层错误
|
||||
let mut stream = match provider.chat_stream(request).await {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
+315
-8
@@ -25,15 +25,319 @@ use super::{LlmProvider, ProviderCapabilities, ProviderFeatures};
|
||||
use crate::llm::convert::{from_openai, to_openai};
|
||||
use crate::llm::error::LlmError;
|
||||
use crate::llm::types::message::{ContentBlock, ContentBlockType, Message};
|
||||
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage};
|
||||
use crate::llm::types::request::{OpenaiChatRequest, OpenaiTool, StreamOptions};
|
||||
use crate::llm::types::openai_message::{ContentField, OpenaiChatMessage, OpenaiContentPart};
|
||||
use crate::llm::types::request_v2::MessageRequest;
|
||||
use crate::llm::types::response::{OpenaiChatChunk, OpenaiChatResponse};
|
||||
use crate::llm::types::response_v2::{
|
||||
MessageResponse, PartialMessageResponse, PartialUsage, StopReason, StreamEvent,
|
||||
};
|
||||
use crate::llm::types::shared::FinishReason;
|
||||
use crate::llm::types::tool::OpenaiToolCall;
|
||||
use crate::llm::types::shared::{FinishReason, ResponseFormat, ServiceTier, StopSequence};
|
||||
use crate::llm::types::tool::{OpenaiToolCall, OpenaiToolDefinition, ToolChoice};
|
||||
use serde::Deserialize;
|
||||
|
||||
// =============================================================================
|
||||
// 0. OpenAI wire-format 类型(Phase 13 从 types::request 迁入)
|
||||
// =============================================================================
|
||||
|
||||
/// 流式响应选项。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>,
|
||||
}
|
||||
|
||||
/// OpenAI wire-format 工具定义。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case", tag = "type")]
|
||||
pub(crate) enum OpenaiTool {
|
||||
Function { function: OpenaiToolDefinition },
|
||||
}
|
||||
|
||||
/// 音频输出参数。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct AudioParam {
|
||||
pub format: String,
|
||||
pub voice: String,
|
||||
}
|
||||
|
||||
/// 预测内容(OpenAI `prediction` 字段)。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct PredictionContent {
|
||||
#[serde(rename = "type")]
|
||||
pub pred_type: String,
|
||||
pub content: String,
|
||||
}
|
||||
|
||||
/// 用户位置(web search 用)。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct UserLocation {
|
||||
#[serde(rename = "type")]
|
||||
pub loc_type: String,
|
||||
pub approximate: Approximate,
|
||||
}
|
||||
|
||||
/// 近似位置。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>,
|
||||
}
|
||||
|
||||
/// Web search 选项。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct WebSearchOptions {
|
||||
pub search_context_size: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub user_location: Option<UserLocation>,
|
||||
}
|
||||
|
||||
/// OpenAI Chat Completions 请求体。
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub(crate) struct OpenaiChatRequest {
|
||||
pub model: String,
|
||||
pub messages: Vec<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>,
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// 0b. OpenAI wire-format 响应类型(Phase 13 从 types::response 迁入)
|
||||
// =============================================================================
|
||||
|
||||
/// 单 token logprob。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>>,
|
||||
}
|
||||
|
||||
/// Top-K logprob。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct TopLogprob {
|
||||
pub token: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bytes: Option<Vec<u32>>,
|
||||
pub logprob: f64,
|
||||
}
|
||||
|
||||
/// Logprobs 容器。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>>,
|
||||
}
|
||||
|
||||
/// URL 引用(annotation 用)。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct URLCitation {
|
||||
pub end_index: u32,
|
||||
pub start_index: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub title: Option<String>,
|
||||
pub url: String,
|
||||
}
|
||||
|
||||
/// 注释(response 中可包含)。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct Annotation {
|
||||
#[serde(rename = "type")]
|
||||
pub ann_type: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub url_citation: Option<URLCitation>,
|
||||
}
|
||||
|
||||
/// OpenAI 音频输出。
|
||||
#[allow(dead_code)]
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>,
|
||||
}
|
||||
|
||||
/// 非流式 choice。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>,
|
||||
}
|
||||
|
||||
/// OpenAI Chat Completions 响应。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct OpenaiChatResponse {
|
||||
pub id: String,
|
||||
pub object: String,
|
||||
pub created: u64,
|
||||
pub model: String,
|
||||
pub choices: Vec<Choice>,
|
||||
pub usage: crate::llm::types::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>,
|
||||
}
|
||||
|
||||
/// 流式响应 delta。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>>,
|
||||
}
|
||||
|
||||
/// 流式 chunk 的 choice。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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>,
|
||||
}
|
||||
|
||||
/// OpenAI Chat Completions 流式 chunk。
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) 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<crate::llm::types::usage::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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// 1. GenericOpenaiProvider —— OpenAI-compatible 协议共用实现
|
||||
@@ -210,7 +514,10 @@ impl GenericOpenaiProvider {
|
||||
///
|
||||
/// 实现注意:先在函数顶部抽取出所有 needed 字段(clone 或 move),避免后续
|
||||
/// 部分移动 `request` 后无法借用其它字段。
|
||||
pub fn convert_request(&self, request: MessageRequest) -> Result<OpenaiChatRequest, LlmError> {
|
||||
pub(crate) fn convert_request(
|
||||
&self,
|
||||
request: MessageRequest,
|
||||
) -> Result<OpenaiChatRequest, LlmError> {
|
||||
// ponytail: 先抽取 / clone 所有 owned 字段,再访问 request.extra,
|
||||
// 避免部分移动导致后续 `&self` borrow 失败。
|
||||
let model = request.model.clone();
|
||||
@@ -273,7 +580,7 @@ impl GenericOpenaiProvider {
|
||||
/// `OpenaiChatResponse` → `MessageResponse`。
|
||||
///
|
||||
/// 返回 `Err(LlmError::Other)` 当 `choices` 为空。
|
||||
pub fn convert_response(
|
||||
pub(crate) fn convert_response(
|
||||
&self,
|
||||
response: OpenaiChatResponse,
|
||||
) -> Result<MessageResponse, LlmError> {
|
||||
@@ -986,7 +1293,7 @@ data: [DONE]\n\n";
|
||||
object: "chat.completion".into(),
|
||||
created: 0,
|
||||
model: "gpt-4o".into(),
|
||||
choices: vec![crate::llm::types::response::Choice {
|
||||
choices: vec![Choice {
|
||||
index: 0,
|
||||
message: OpenaiChatMessage::Assistant {
|
||||
content: ContentField::String(String::new()),
|
||||
|
||||
+5
-199
@@ -1,203 +1,9 @@
|
||||
//! 流式事件系统 —— 将 LLM 流式响应解析为语义化事件。
|
||||
//! 流式事件系统 —— 重导出 `StreamEvent` 供向后兼容。
|
||||
//!
|
||||
//! Phase 0 修订(参见 `docs/10a-phase0-types-and-trait.md` §"StreamEvent 命名冲突处理"):
|
||||
//! 历史说明(Phase 0 → Phase 13):
|
||||
//! - 对外暴露的 `StreamEvent` 是高精度 IR 版本(来自 `response_v2::StreamEvent`)。
|
||||
//! - 旧变体(`AssistantTextDelta` / `ToolExecutionStarted` 等)重命名为 `LegacyStreamEvent`
|
||||
//! 放在 `crate::llm::types::old_stream` 模块,本文件内部消费。
|
||||
//! - Phase 1 重写 Provider 时可直接消费新事件流后整体删除 `LegacyStreamEvent` 相关代码。
|
||||
//!
|
||||
//! 当前实现:旧的 `parse_chunk_stream` 内部消费 `OpenaiChatChunk`,映射为
|
||||
//! `LegacyStreamEvent`,再在 `LegacyToIrEventStream` 中映射为新 IR `StreamEvent`
|
||||
//! 后输出。Phase 1 会重写此层(OpenAI Provider 直接产出新事件流)。
|
||||
//! - 旧版 chunk 解析 + LegacyStreamEvent 适配层在 Phase 13 完成后已整体删除。
|
||||
//! - 当前文件仅保留 `pub use` 重导出,保持与既有
|
||||
//! `use crate::llm::stream::StreamEvent` 的代码兼容。
|
||||
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use futures_core::stream::Stream;
|
||||
use futures_util::FutureExt;
|
||||
use futures_util::future::poll_fn;
|
||||
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;
|
||||
|
||||
/// 将原始 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 },
|
||||
}
|
||||
}
|
||||
|
||||
+1
-77
@@ -1,9 +1,6 @@
|
||||
pub mod message;
|
||||
pub mod old_stream;
|
||||
pub mod openai_message;
|
||||
pub mod request;
|
||||
pub mod request_v2;
|
||||
pub mod response;
|
||||
pub mod response_v2;
|
||||
pub mod shared;
|
||||
pub mod tool;
|
||||
@@ -12,12 +9,7 @@ pub mod usage;
|
||||
pub use openai_message::{
|
||||
ContentField, FileData, ImageURL, InputAudio, OpenaiChatMessage, OpenaiContentPart,
|
||||
};
|
||||
pub use request::{OpenaiChatRequest, OpenaiTool, StreamOptions, ToolChoice};
|
||||
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::{
|
||||
ContentBlockBuilder, MessageResponse, PartialMessageResponse, PartialUsage, StopReason,
|
||||
StreamEvent,
|
||||
@@ -26,73 +18,5 @@ pub use shared::{
|
||||
AudioFormat, FinishReason, ImageDetail, Modality, ResponseFormat, Role, ServiceTier,
|
||||
StopSequence,
|
||||
};
|
||||
pub use tool::{FunctionCall, OpenaiToolCall, ToolDef};
|
||||
pub use tool::{FunctionCall, OpenaiToolCall, ToolChoice, ToolDef};
|
||||
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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,187 +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, 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",
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[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>,
|
||||
}
|
||||
@@ -9,7 +9,7 @@ use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::llm::types::message::Message;
|
||||
use crate::llm::types::request::ToolChoice;
|
||||
use crate::llm::types::tool::ToolChoice;
|
||||
use crate::llm::types::tool::ToolDef;
|
||||
|
||||
/// Provider 无关的请求类型。
|
||||
|
||||
@@ -1,177 +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 crate::llm::types::{ContentField, OpenaiContentPart};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[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,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -63,3 +63,93 @@ pub struct FunctionCall {
|
||||
pub enum OpenaiToolCall {
|
||||
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",
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -31,7 +31,7 @@ pub struct PromptTokensDetails {
|
||||
pub cached_tokens: Option<u32>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize)]
|
||||
pub struct CostTracker {
|
||||
accumulated: Usage,
|
||||
}
|
||||
@@ -61,6 +61,14 @@ impl CostTracker {
|
||||
}
|
||||
}
|
||||
|
||||
impl From<Usage> for CostTracker {
|
||||
fn from(usage: Usage) -> Self {
|
||||
CostTracker {
|
||||
accumulated: usage,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Usage {
|
||||
pub fn from_input_output(input: u32, output: u32) -> Self {
|
||||
let total = input.saturating_add(output);
|
||||
|
||||
+9
-1
@@ -2,23 +2,31 @@
|
||||
|
||||
pub mod conversation;
|
||||
pub mod error;
|
||||
pub mod graph;
|
||||
pub mod knowledge;
|
||||
pub mod retriever;
|
||||
pub mod store;
|
||||
pub mod types;
|
||||
pub mod vector;
|
||||
pub mod vector_store;
|
||||
|
||||
// 高频类型(大多数下游需要)
|
||||
pub use conversation::{ConversationMemory, ConversationMemoryConfig};
|
||||
pub use error::MemoryError;
|
||||
pub use graph::{GraphEntity, GraphRelation, InMemoryGraph, KnowledgeGraph, RelationDirection, ScoredEntity};
|
||||
pub use knowledge::KnowledgeStore;
|
||||
pub use retriever::MemoryRetriever;
|
||||
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 graph::TagConstraints;
|
||||
pub use knowledge::{KNOWLEDGE_PREFIX, PageIndexEntry};
|
||||
pub use retriever::{RetrievalResult, RetrieverConfig, ScoredItem};
|
||||
#[allow(deprecated)]
|
||||
pub use retriever::ScoredItem;
|
||||
pub use retriever::{RetrievalItem, RetrievalResult, RetrievalStrategy, RetrieverConfig};
|
||||
pub use store::{EvictionConfig, EvictionPolicy};
|
||||
pub use types::{KnowledgePage, MemoryFilter, MemoryItem};
|
||||
|
||||
+1100
File diff suppressed because it is too large
Load Diff
+410
-31
@@ -1,8 +1,13 @@
|
||||
//! 记忆检索器 —— 基于 TextOverlap (Dice 系数) 的单通道关键词检索。
|
||||
//! 记忆检索器 -- 双通道关键词检索(KnowledgeStore + KnowledgeGraph)。
|
||||
//!
|
||||
//! 单通道模式:仅 KnowledgeStore,基于 TextOverlap (Dice 系数) 评分。
|
||||
//! 双通道模式:并行检索 KnowledgeStore + KnowledgeGraph,结果合并为统一列表。
|
||||
|
||||
use std::collections::HashSet;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::memory::error::MemoryError;
|
||||
use crate::memory::graph::{GraphEntity, KnowledgeGraph, ScoredEntity};
|
||||
use crate::memory::knowledge::KnowledgeStore;
|
||||
use crate::memory::types::KnowledgePage;
|
||||
|
||||
@@ -13,6 +18,8 @@ pub struct RetrieverConfig {
|
||||
pub max_results: usize,
|
||||
/// 最低分数阈值 [0.0, 1.0](默认 0.1)。
|
||||
pub min_score: f32,
|
||||
/// 图遍历的默认深度(默认 2)。
|
||||
pub graph_depth: usize,
|
||||
}
|
||||
|
||||
impl Default for RetrieverConfig {
|
||||
@@ -20,64 +27,193 @@ impl Default for RetrieverConfig {
|
||||
Self {
|
||||
max_results: 20,
|
||||
min_score: 0.1,
|
||||
graph_depth: 2,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 单条带评分的检索结果。
|
||||
/// 检索策略 -- 控制双通道分流。
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
pub enum RetrievalStrategy {
|
||||
/// 并行 KnowledgeStore + KnowledgeGraph,合并排序(默认)。
|
||||
#[default]
|
||||
Hybrid,
|
||||
/// 仅 KnowledgeStore。
|
||||
KnowledgeOnly,
|
||||
/// 仅 KnowledgeGraph。
|
||||
GraphOnly,
|
||||
}
|
||||
|
||||
/// 统一检索条目 -- enum 变体区分类别,两通道分数均在 \[0,1\] 区间。
|
||||
///
|
||||
/// 注意:两个通道的分数维度不同(TextOverlap vs 图距离),
|
||||
/// 合并排序仅用于统一返回,不代表跨通道可比性。
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum RetrievalItem {
|
||||
/// 知识页面(来自 KnowledgeStore)。
|
||||
KnowledgePage {
|
||||
page: KnowledgePage,
|
||||
/// TextOverlap 评分 [0.0, 1.0]。
|
||||
score: f32,
|
||||
},
|
||||
/// 图谱实体(来自 KnowledgeGraph)。
|
||||
GraphEntity {
|
||||
entity: GraphEntity,
|
||||
/// 图距离评分 [0.0, 1.0]。
|
||||
score: f32,
|
||||
/// 从查询实体到当前实体的 ID 路径。
|
||||
path: Vec<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl RetrievalItem {
|
||||
/// 统一分数(用于合并排序)。
|
||||
pub fn score(&self) -> f32 {
|
||||
match self {
|
||||
Self::KnowledgePage { score, .. } => *score,
|
||||
Self::GraphEntity { score, .. } => *score,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 旧版带评分的知识页面检索结果(已废弃)。
|
||||
///
|
||||
/// 请迁移到 [`RetrievalItem::KnowledgePage`]。
|
||||
#[deprecated(since = "0.3.0", note = "使用 RetrievalItem::KnowledgePage 代替")]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ScoredItem {
|
||||
pub page: KnowledgePage,
|
||||
/// TextOverlap 评分 [0.0, 1.0]
|
||||
pub score: f32,
|
||||
}
|
||||
|
||||
/// 检索结果。
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RetrievalResult {
|
||||
pub items: Vec<ScoredItem>,
|
||||
/// 统一条目列表,按分数降序排列。
|
||||
pub items: Vec<RetrievalItem>,
|
||||
pub query: String,
|
||||
/// 本次检索实际执行的策略(可能因 graph 未注入而退化),而非用户通过 `with_strategy()` 配置的值。
|
||||
pub strategy: RetrievalStrategy,
|
||||
}
|
||||
|
||||
/// 记忆检索器 —— 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分。
|
||||
/// 记忆检索器 -- 在 `KnowledgeStore` 中做关键词检索并按 TextOverlap 评分,
|
||||
/// 可选注入 `KnowledgeGraph` 启用双通道检索。
|
||||
pub struct MemoryRetriever {
|
||||
knowledge_store: KnowledgeStore,
|
||||
/// 可选知识图谱(None 时退化为单通道)。
|
||||
knowledge_graph: Option<Arc<dyn KnowledgeGraph>>,
|
||||
/// 检索策略(默认 Hybrid)。
|
||||
strategy: RetrievalStrategy,
|
||||
config: RetrieverConfig,
|
||||
/// 停用词表(用于关键词提取)。
|
||||
stop_words: HashSet<String>,
|
||||
}
|
||||
|
||||
impl MemoryRetriever {
|
||||
/// 创建一个新的 MemoryRetriever。
|
||||
/// 创建一个新的 MemoryRetriever(保持向后兼容)。
|
||||
pub fn new(knowledge_store: KnowledgeStore, config: RetrieverConfig) -> Self {
|
||||
Self {
|
||||
knowledge_store,
|
||||
knowledge_graph: None,
|
||||
strategy: RetrievalStrategy::default(),
|
||||
config,
|
||||
stop_words: default_stop_words(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 注入知识图谱,启用双通道检索。
|
||||
pub fn with_knowledge_graph(mut self, graph: Arc<dyn KnowledgeGraph>) -> Self {
|
||||
self.knowledge_graph = Some(graph);
|
||||
self
|
||||
}
|
||||
|
||||
/// 设置检索策略。
|
||||
pub fn with_strategy(mut self, strategy: RetrievalStrategy) -> Self {
|
||||
self.strategy = strategy;
|
||||
self
|
||||
}
|
||||
|
||||
/// 替换停用词表。
|
||||
pub fn with_stop_words(mut self, stop_words: HashSet<String>) -> Self {
|
||||
self.stop_words = stop_words;
|
||||
self
|
||||
}
|
||||
|
||||
/// 检索相关知识页面。
|
||||
/// 检索相关记忆(双通道)。
|
||||
///
|
||||
/// 根据 `strategy` 和是否注入 `knowledge_graph` 分流:
|
||||
/// - `KnowledgeOnly` 或未注入 graph -> 仅检索 KnowledgeStore
|
||||
/// - `GraphOnly` -> 仅检索 KnowledgeGraph
|
||||
/// - `Hybrid` -> 并行检索两通道,合并排序
|
||||
pub async fn retrieve(&self, query: &str) -> Result<RetrievalResult, MemoryError> {
|
||||
if query.is_empty() {
|
||||
return Ok(RetrievalResult {
|
||||
items: Vec::new(),
|
||||
query: query.to_string(),
|
||||
strategy: self.strategy.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
// 1. 关键词提取
|
||||
let keywords = extract_keywords(query, &self.stop_words);
|
||||
let has_graph = self.knowledge_graph.is_some();
|
||||
|
||||
// 2. 用关键词在 KnowledgeStore 中搜索
|
||||
// 按策略分流
|
||||
match (&self.strategy, has_graph) {
|
||||
// 仅知识页面(或未注入 graph 退化为单通道)
|
||||
(RetrievalStrategy::KnowledgeOnly, _) | (_, false) => {
|
||||
let items = self.search_knowledge_store(query, &keywords).await?;
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
strategy: RetrievalStrategy::KnowledgeOnly,
|
||||
})
|
||||
}
|
||||
// 仅图谱
|
||||
(RetrievalStrategy::GraphOnly, true) => {
|
||||
let graph = self.knowledge_graph.as_ref().unwrap();
|
||||
let items = self.search_graph(&keywords, graph).await?;
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
strategy: self.strategy.clone(),
|
||||
})
|
||||
}
|
||||
// 混合:并行执行,合并排序
|
||||
(RetrievalStrategy::Hybrid, true) => {
|
||||
let graph = self.knowledge_graph.as_ref().unwrap();
|
||||
let (kp_items, g_items) = tokio::join!(
|
||||
self.search_knowledge_store(query, &keywords),
|
||||
self.search_graph(&keywords, graph),
|
||||
);
|
||||
|
||||
let mut items = kp_items?;
|
||||
items.extend(g_items?);
|
||||
// 过滤最低分数
|
||||
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.truncate(self.config.max_results);
|
||||
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
strategy: self.strategy.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 通道 1:KnowledgeStore 关键词检索 + TextOverlap 评分。
|
||||
async fn search_knowledge_store(
|
||||
&self,
|
||||
query: &str,
|
||||
keywords: &[String],
|
||||
) -> Result<Vec<RetrievalItem>, MemoryError> {
|
||||
let mut pages = Vec::new();
|
||||
for keyword in &keywords {
|
||||
for keyword in keywords {
|
||||
let found = self.knowledge_store.search(keyword).await?;
|
||||
for page in found {
|
||||
if !pages.iter().any(|p: &KnowledgePage| p.id == page.id) {
|
||||
@@ -86,32 +222,79 @@ impl MemoryRetriever {
|
||||
}
|
||||
}
|
||||
|
||||
// 3. TextOverlap 评分
|
||||
let mut items: Vec<ScoredItem> = pages
|
||||
let mut items: Vec<RetrievalItem> = pages
|
||||
.into_iter()
|
||||
.map(|page| {
|
||||
let score = text_overlap_score(query, &page);
|
||||
ScoredItem { page, score }
|
||||
RetrievalItem::KnowledgePage { page, score }
|
||||
})
|
||||
.collect();
|
||||
|
||||
// 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)
|
||||
b.score()
|
||||
.partial_cmp(&a.score())
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
items.truncate(self.config.max_results);
|
||||
Ok(items)
|
||||
}
|
||||
|
||||
Ok(RetrievalResult {
|
||||
items,
|
||||
query: query.to_string(),
|
||||
})
|
||||
/// 通道 2:KnowledgeGraph 关键词检索 + BFS 图遍历。
|
||||
///
|
||||
/// 流程:find_by_keywords 找到起始实体 -> 对每个起始实体 BFS 遍历 ->
|
||||
/// 收集 ScoredEntity -> 转为 RetrievalItem::GraphEntity。
|
||||
async fn search_graph(
|
||||
&self,
|
||||
keywords: &[String],
|
||||
graph: &Arc<dyn KnowledgeGraph>,
|
||||
) -> Result<Vec<RetrievalItem>, MemoryError> {
|
||||
let starts = graph.find_by_keywords(keywords).await?;
|
||||
if starts.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let depth = self.config.graph_depth;
|
||||
let mut seen: HashSet<String> = HashSet::new();
|
||||
let mut items: Vec<RetrievalItem> = Vec::new();
|
||||
|
||||
for start in starts {
|
||||
// 起始实体自身也加入结果(score = 1.0)
|
||||
if seen.insert(start.id.clone()) {
|
||||
items.push(RetrievalItem::GraphEntity {
|
||||
entity: start.clone(),
|
||||
score: 1.0,
|
||||
path: vec![start.id.clone()],
|
||||
});
|
||||
}
|
||||
// BFS 找相关实体
|
||||
let related: Vec<ScoredEntity> = graph
|
||||
.get_related(&start.id, depth, crate::memory::graph::RelationDirection::Both, None)
|
||||
.await?;
|
||||
for se in related {
|
||||
if seen.insert(se.entity.id.clone()) {
|
||||
items.push(RetrievalItem::GraphEntity {
|
||||
entity: se.entity,
|
||||
score: se.score,
|
||||
path: se.path,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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.truncate(self.config.max_results);
|
||||
Ok(items)
|
||||
}
|
||||
}
|
||||
|
||||
/// 从 query 中提取关键词:按非字母数字字符分割 → 转小写 → 过滤单字符和停用词。
|
||||
/// 从 query 中提取关键词:按非字母数字字符分割 -> 转小写 -> 过滤单字符和停用词。
|
||||
fn extract_keywords(query: &str, stop_words: &HashSet<String>) -> Vec<String> {
|
||||
query
|
||||
.split(|c: char| !c.is_alphanumeric())
|
||||
@@ -178,6 +361,7 @@ fn default_stop_words() -> HashSet<String> {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::memory::graph::{GraphEntity, InMemoryGraph, GraphRelation};
|
||||
use crate::memory::knowledge::KnowledgeStore;
|
||||
use crate::memory::{InMemoryStore, MemoryStore};
|
||||
use std::sync::Arc;
|
||||
@@ -197,13 +381,26 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn make_store() -> (Arc<InMemoryStore>, KnowledgeStore) {
|
||||
let store = Arc::new(InMemoryStore::new());
|
||||
let ks = KnowledgeStore::new(store.clone());
|
||||
(store, ks)
|
||||
}
|
||||
|
||||
fn make_retriever(ks: KnowledgeStore) -> MemoryRetriever {
|
||||
MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
}
|
||||
|
||||
// ── 单通道(KnowledgeStore)测试 ──
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_empty_query() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default());
|
||||
let (_, ks) = make_store();
|
||||
let retriever = make_retriever(ks);
|
||||
let result = retriever.retrieve("").await.unwrap();
|
||||
assert!(result.items.is_empty());
|
||||
// 空查询返回用户配置的策略
|
||||
assert_eq!(result.strategy, RetrievalStrategy::Hybrid);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -225,19 +422,25 @@ mod tests {
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default());
|
||||
let result = retriever.retrieve("LangGraph state").await.unwrap();
|
||||
assert!(!result.items.is_empty());
|
||||
assert_eq!(result.items[0].page.id, "p1");
|
||||
assert!(result.items[0].score > 0.0);
|
||||
// 第一个结果应该是 KnowledgePage 变体
|
||||
match &result.items[0] {
|
||||
RetrievalItem::KnowledgePage { page, score } => {
|
||||
assert_eq!(page.id, "p1");
|
||||
assert!(*score > 0.0);
|
||||
}
|
||||
_ => panic!("expected KnowledgePage variant"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_respects_min_score() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
let (_, ks) = make_store();
|
||||
ks.add_page(make_page("p1", "X", "Y", "Z")).await.unwrap();
|
||||
|
||||
let config = RetrieverConfig {
|
||||
max_results: 10,
|
||||
min_score: 0.99,
|
||||
graph_depth: 2,
|
||||
};
|
||||
let retriever = MemoryRetriever::new(ks, config);
|
||||
let result = retriever
|
||||
@@ -249,8 +452,7 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_respects_max_results() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
let (_, ks) = make_store();
|
||||
for i in 0..5 {
|
||||
ks.add_page(make_page(
|
||||
&format!("p{i}"),
|
||||
@@ -264,12 +466,160 @@ mod tests {
|
||||
let config = RetrieverConfig {
|
||||
max_results: 2,
|
||||
min_score: 0.0,
|
||||
graph_depth: 2,
|
||||
};
|
||||
let retriever = MemoryRetriever::new(ks, config);
|
||||
let result = retriever.retrieve("LangGraph").await.unwrap();
|
||||
assert_eq!(result.items.len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_without_graph_degrades_to_knowledge_only() {
|
||||
let (_, ks) = make_store();
|
||||
ks.add_page(make_page("p1", "Test", "t", "t"))
|
||||
.await
|
||||
.unwrap();
|
||||
// 默认 Hybrid 策略,但未注入 graph
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default());
|
||||
let result = retriever.retrieve("Test").await.unwrap();
|
||||
// 应退化为 KnowledgeOnly
|
||||
assert_eq!(result.strategy, RetrievalStrategy::KnowledgeOnly);
|
||||
}
|
||||
|
||||
// ── 双通道(KnowledgeGraph)测试 ──
|
||||
|
||||
async fn make_graph_with_data() -> Arc<InMemoryGraph> {
|
||||
let graph = Arc::new(InMemoryGraph::new());
|
||||
// 准备一些实体
|
||||
let mut langchain = GraphEntity::new("langchain", "LangChain", "framework");
|
||||
langchain.description = "LLM application framework".to_string();
|
||||
let mut langgraph = GraphEntity::new("langgraph", "LangGraph", "framework");
|
||||
langgraph.description = "Graph-based agent runtime".to_string();
|
||||
let mut python = GraphEntity::new("python", "Python", "language");
|
||||
python.description = "Programming language".to_string();
|
||||
|
||||
graph.add_entity(langchain).await.unwrap();
|
||||
graph.add_entity(langgraph).await.unwrap();
|
||||
graph.add_entity(python).await.unwrap();
|
||||
graph
|
||||
.add_relation(GraphRelation::new("langchain", "langgraph", "includes", 0.8))
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
.add_relation(GraphRelation::new("langchain", "python", "built_with", 0.9))
|
||||
.await
|
||||
.unwrap();
|
||||
graph
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_graph_only() {
|
||||
let (_, ks) = make_store();
|
||||
let graph = make_graph_with_data().await;
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph)
|
||||
.with_strategy(RetrievalStrategy::GraphOnly);
|
||||
|
||||
let result = retriever.retrieve("langchain").await.unwrap();
|
||||
assert_eq!(result.strategy, RetrievalStrategy::GraphOnly);
|
||||
// 应该全部是 GraphEntity 变体
|
||||
assert!(result.items.iter().all(|i| matches!(i, RetrievalItem::GraphEntity { .. })));
|
||||
// langchain 是起始实体(score=1.0),langgraph 和 python 是 BFS 结果
|
||||
assert!(!result.items.is_empty(), "should find graph entities");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_hybrid_both_channels() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
// KnowledgeStore 有匹配
|
||||
ks.add_page(make_page(
|
||||
"p1",
|
||||
"LangChain framework",
|
||||
"LLM application framework",
|
||||
"LangChain is a framework for building LLM applications",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
// KnowledgeGraph 也有匹配
|
||||
let graph = make_graph_with_data().await;
|
||||
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph);
|
||||
let result = retriever.retrieve("langchain").await.unwrap();
|
||||
assert_eq!(result.strategy, RetrievalStrategy::Hybrid);
|
||||
// 应该同时包含 KnowledgePage 和 GraphEntity
|
||||
let has_page = result
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::KnowledgePage { .. }));
|
||||
let has_entity = result
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::GraphEntity { .. }));
|
||||
assert!(has_page, "should have KnowledgePage results");
|
||||
assert!(has_entity, "should have GraphEntity results");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_hybrid_only_store_has_results() {
|
||||
let (_, ks) = make_store();
|
||||
ks.add_page(make_page("p1", "LangChain", "framework", "LLM app"))
|
||||
.await
|
||||
.unwrap();
|
||||
// 空图
|
||||
let graph = Arc::new(InMemoryGraph::new());
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph);
|
||||
let result = retriever.retrieve("langchain").await.unwrap();
|
||||
assert_eq!(result.strategy, RetrievalStrategy::Hybrid);
|
||||
// 图空,只有 Store 结果
|
||||
let has_page = result
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::KnowledgePage { .. }));
|
||||
assert!(has_page);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_hybrid_only_graph_has_results() {
|
||||
let store = Arc::new(InMemoryStore::new()) as Arc<dyn MemoryStore>;
|
||||
let ks = KnowledgeStore::new(store);
|
||||
// Store 空,图有数据
|
||||
let graph = make_graph_with_data().await;
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph);
|
||||
let result = retriever.retrieve("langchain").await.unwrap();
|
||||
assert_eq!(result.strategy, RetrievalStrategy::Hybrid);
|
||||
let has_entity = result
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::GraphEntity { .. }));
|
||||
assert!(has_entity, "should have graph results even when store is empty");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrieve_strategy_knowledge_only_ignores_graph() {
|
||||
let (_, ks) = make_store();
|
||||
ks.add_page(make_page("p1", "Test", "t", "t"))
|
||||
.await
|
||||
.unwrap();
|
||||
let graph = make_graph_with_data().await;
|
||||
let retriever = MemoryRetriever::new(ks, RetrieverConfig::default())
|
||||
.with_knowledge_graph(graph)
|
||||
.with_strategy(RetrievalStrategy::KnowledgeOnly);
|
||||
let result = retriever.retrieve("langchain").await.unwrap();
|
||||
assert_eq!(result.strategy, RetrievalStrategy::KnowledgeOnly);
|
||||
// 即使图有数据,KnowledgeOnly 也不应该返回 GraphEntity
|
||||
let has_entity = result
|
||||
.items
|
||||
.iter()
|
||||
.any(|i| matches!(i, RetrievalItem::GraphEntity { .. }));
|
||||
assert!(!has_entity, "KnowledgeOnly should not return graph entities");
|
||||
}
|
||||
|
||||
// ── 辅助函数测试 ──
|
||||
|
||||
#[test]
|
||||
fn text_overlap_dice_zero_on_empty() {
|
||||
assert_eq!(text_overlap_dice("hello", ""), 0.0);
|
||||
@@ -299,4 +649,33 @@ mod tests {
|
||||
assert!(!kws.contains(&"a".to_string()));
|
||||
assert!(!kws.contains(&"b".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn retrieval_item_score_accessor() {
|
||||
let page = KnowledgePage {
|
||||
id: "p1".into(),
|
||||
title: "T".into(),
|
||||
summary: "S".into(),
|
||||
content: "C".into(),
|
||||
tags: vec![],
|
||||
references: vec![],
|
||||
created_at: OffsetDateTime::now_utc(),
|
||||
updated_at: OffsetDateTime::now_utc(),
|
||||
};
|
||||
let item = RetrievalItem::KnowledgePage { page, score: 0.5 };
|
||||
assert!((item.score() - 0.5).abs() < 0.001);
|
||||
|
||||
let entity = GraphEntity::new("e1", "E1", "x");
|
||||
let item2 = RetrievalItem::GraphEntity {
|
||||
entity,
|
||||
score: 0.8,
|
||||
path: vec!["e1".into()],
|
||||
};
|
||||
assert!((item2.score() - 0.8).abs() < 0.001);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_strategy_is_hybrid() {
|
||||
assert_eq!(RetrievalStrategy::default(), RetrievalStrategy::Hybrid);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@ use crate::memory::error::MemoryError;
|
||||
///
|
||||
/// **稳定性**:实验性 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` 加入索引。
|
||||
@@ -44,10 +45,12 @@ pub trait VectorRetriever: Send + Sync {
|
||||
/// - 不做向量维度校验(不同维度向量查询结果无意义但不 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 {
|
||||
@@ -57,12 +60,14 @@ impl InMemoryVectorRetriever {
|
||||
}
|
||||
}
|
||||
|
||||
#[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> {
|
||||
@@ -118,6 +123,7 @@ fn dot(a: &[f32], b: &[f32]) -> f32 {
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[allow(deprecated)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -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()
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user