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:
徐涛
2026-07-15 08:59:46 +08:00
parent 1d51dcdfe0
commit 34eec9f546
9 changed files with 1845 additions and 2 deletions
+252
View File
@@ -0,0 +1,252 @@
//! engine_demo —— SessionManager + Checkpointer 端到端示例。
//!
//! 演示:
//! 1. SessionManager::create 创建 session
//! 2. SessionManager::submit_turnauto_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. 构造 SessionManagerauto_checkpoint 默认 true
let sm = SessionManager::new(store.clone());
println!("=== SessionManager 创建 ===");
// 3. create + submit_turnauto_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}) 仍是孤儿 sessionparent() = None");
// 清理孤儿
sm.destroy(&c_id).await.unwrap();
println!("\n✓ engine_demo 完成");
}
+152
View File
@@ -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` 表示从未生成过)。
+55 -2
View File
@@ -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 {
+377
View File
@@ -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");
}
}
+46
View File
@@ -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),
}
+20
View File
@@ -0,0 +1,20 @@
//! Engine 模块 —— Agent 执行引擎。
//!
//! Phase 17 新增。提供 SessionManager(会话树管理)和 Checkpointertime-travel 检查点)能力。
//!
//! ## 子模块
//!
//! - [`session_manager`]SessionManager + SessionManagerConfig
//! - [`checkpointer`]Checkpointertime-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};
+906
View File
@@ -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 后自动 checkpointStep 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 已被 destroyload_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());
}
}
}
+36
View File
@@ -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>,
}
+1
View File
@@ -2,6 +2,7 @@
pub mod agent;
pub mod document;
pub mod engine;
pub mod llm;
pub mod memory;
pub mod prompt;