feat(engine): 实现 Agent 执行引擎(SessionManager + Checkpointer + SessionSnapshot)
Phase 17 主体交付:解决 v0.2 中 session 在变量里、无父子关系、无 checkpoint、
不可序列化的空白。
新增模块 src/engine/(5 文件,约 1300 行纯实现 + 21 个内联测试):
- error.rs(46 行)—— EngineError 枚举(6 变体)
- SessionNotFound / SessionAlreadyExists / CheckpointNotFound
- Memory(#[from] MemoryError) 透传(与 AgentError 风格一致)
- Serialization / Agent(#[from] AgentError)
- snapshot.rs(36 行)—— SessionSnapshot + SessionMemoryEntry
- 独立 struct 避开 Arc<dyn Agent> 不可序列化限制
- SessionMemoryEntry 保留 value/metadata/created_at 完整信息
- 所有字段 #[serde(default)] 宽松反序列化保证前向兼容
- checkpointer.rs(377 行)—— Time-travel 检查点管理器
- checkpoint() / rollback_load() / list_checkpoints() / delete_all() / latest_snapshot()
- 存储 key:ckpt:{session_id}:{ckpt_id}
- ckpt_id 纳秒+计数器(无外部依赖,ponytail)
- CkptMeta.created_at_nanos 字段确保同秒内精确降序
- 6 个内联测试覆盖 roundtrip / 不存在 ckpt / 降序排序 / delete_all 幂等 /
latest_snapshot / 跨 session 隔离
- session_manager.rs(906 行)—— SessionManager 会话树管理
- 内部 RwLock<HashMap> + Arc<tokio::sync::Mutex<AgentSession>> 双重锁
- create / create_child / get / recover / replace / children / parent /
destroy / submit_turn / submit_turn_stream / finalize_turn_stream 共 11 个公开方法
- SessionManagerConfig.auto_checkpoint 默认 true(同步写入 +
tracing::error! 失败不阻断主流程,不提供强持久化保证)
- 孤儿策略:destroy 不递归删除子 session;父被销毁后 parent() 返回 None
- create_child 限制:父 session 必须先 get/recover 到内存(bundle 不可序列化)
- 15 个内联测试覆盖 CRUD / recover / replace / 树形 / 孤儿 / auto_checkpoint 开关 /
序列化兼容性 / 幂等 / 10 并发创建
- mod.rs(20 行)—— 统一 pub use 重导出 EngineError / SessionManager /
SessionManagerConfig / Checkpointer / CkptMeta / SessionSnapshot /
SessionMemoryEntry
AgentSession 扩展(src/agent/session.rs,+152 行):
- to_snapshot() pub async —— 从 MemoryStore 拍平 session_memory 全量数据
- from_snapshot() pub fn Result —— 纯同步构造器(pending_memory_restore 暂存)
- restore_memory() pub async &mut self —— 写回持久层并清空 pending
- has_pending_memory_restore() —— 查询 pending 状态
- pub(crate) fn bundle() —— accessor 供 SessionManager::create_child 继承
SessionMemory 扩展(src/agent/session_memory.rs,+57 行):
- list_entries() —— 返回 Vec<(key, value, metadata, created_at_unix_secs)>
- set_with_meta() —— 保留 metadata/created_at 写入(供 restore_memory 完整恢复)
src/lib.rs —— pub mod engine 声明
examples/engine_demo.rs(+251 行):
端到端演示 create → submit_turn → checkpoint → list_checkpoints →
rollback_load → from_snapshot → restore_memory → replace → destroy 完整链路,
含 rollback 一致性 assert(turn_index 和 cost 恢复到 checkpoint 时刻)。
零新外部依赖(serde_json 已有)。全量 353 → 374(+21 新测试)。
clippy 0 警告,doc 0 warning,example exit 0。
Phase 17: Agent 执行引擎 — Step 2-7(合并提交)
This commit is contained in:
@@ -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 完成");
|
||||
}
|
||||
@@ -25,6 +25,8 @@ 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};
|
||||
@@ -61,6 +63,11 @@ pub struct AgentSession {
|
||||
/// 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 {
|
||||
@@ -120,6 +127,7 @@ impl AgentSession {
|
||||
slots,
|
||||
current_slot_id: "default".to_string(),
|
||||
last_summary_turn: None,
|
||||
pending_memory_restore: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -138,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,
|
||||
@@ -479,6 +492,145 @@ impl AgentSession {
|
||||
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` 表示从未生成过)。
|
||||
|
||||
@@ -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,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,46 @@
|
||||
//! 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),
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
//! 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 use checkpointer::{Checkpointer, CkptMeta};
|
||||
pub use error::EngineError;
|
||||
pub use session_manager::{SessionManager, SessionManagerConfig};
|
||||
pub use snapshot::{SessionMemoryEntry, SessionSnapshot};
|
||||
@@ -0,0 +1,906 @@
|
||||
//! 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 持久化 ======
|
||||
|
||||
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(())
|
||||
}
|
||||
|
||||
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`。
|
||||
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,36 @@
|
||||
//! 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` 兼容旧快照)。
|
||||
#[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>,
|
||||
}
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
pub mod agent;
|
||||
pub mod document;
|
||||
pub mod engine;
|
||||
pub mod llm;
|
||||
pub mod memory;
|
||||
pub mod prompt;
|
||||
|
||||
Reference in New Issue
Block a user