- README 添加 feature 组合表 + 模块级 features 清单 + 升级指南 - 18 个 example 顶部添加 Required features 注释 - roadmap.md 和 roadmap-v0.3.2.md 同步 Phase 26-27 完成状态 - cargo fmt 全量格式化(修复预存格式问题,CI format job 可通过)
231 lines
8.3 KiB
Rust
231 lines
8.3 KiB
Rust
//! SessionMemory —— 会话级记忆,用于 context 间的信息桥接。
|
||
//!
|
||
//! 设计要点(参见 `docs/7-agent-runtime.md` §3.2.8):
|
||
//!
|
||
//! - **会话级**:单 session 内共享,跨 context 桥接信息(不是持久层,也不是对话历史)
|
||
//! - **复用 Phase 3 `MemoryStore`**:不引入新的存储后端机制
|
||
//! - **按 `namespace` 隔离**:每个 session 一个独立命名空间,防止跨 session 泄漏
|
||
//! - **`snapshot()` 格式化为标记文本**:专为注入 system prompt 设计
|
||
//! - **所有方法为 `async`**:因为后端可能是跨进程的(Redis / DB)
|
||
|
||
use std::sync::Arc;
|
||
|
||
use time::OffsetDateTime;
|
||
|
||
use crate::agent::error::AgentError;
|
||
use crate::memory::store::MemoryStore;
|
||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||
|
||
/// 会话级记忆实例。
|
||
///
|
||
/// 基于 [`MemoryStore`] 后端,按 `namespace` 隔离键值数据。
|
||
/// 适用于 session 内各 context 之间的信息桥接(如将关键结论传递给后续 context)。
|
||
pub struct SessionMemory {
|
||
store: Arc<dyn MemoryStore>,
|
||
namespace: String,
|
||
}
|
||
|
||
impl SessionMemory {
|
||
/// 创建新的 session 级记忆实例。
|
||
///
|
||
/// - `store`:后端存储(可跨进程共享的 `MemoryStore` 实现)。
|
||
/// - `namespace`:按 session_id 隔离,防止跨 session 泄漏。
|
||
/// 内部会自动添加 `"_session_"` 前缀。
|
||
pub fn new(store: Arc<dyn MemoryStore>, namespace: &str) -> Self {
|
||
Self {
|
||
store,
|
||
namespace: format!("_session_{namespace}"),
|
||
}
|
||
}
|
||
|
||
/// 内部 key 格式:`"{namespace}:{key}"`。
|
||
fn internal_key(&self, key: &str) -> String {
|
||
format!("{}:{}", self.namespace, key)
|
||
}
|
||
|
||
/// 写入一条 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,
|
||
created_at: created_at_dt,
|
||
};
|
||
self.store.save(item).await.map_err(AgentError::Memory)
|
||
}
|
||
|
||
/// 读取指定 key 的值。
|
||
pub async fn get(&self, key: &str) -> Result<Option<String>, AgentError> {
|
||
let item = self
|
||
.store
|
||
.get(&self.internal_key(key))
|
||
.await
|
||
.map_err(AgentError::Memory)?;
|
||
Ok(item.map(|i| i.content))
|
||
}
|
||
|
||
/// 返回所有条目的格式化快照,适合注入 system prompt。
|
||
///
|
||
/// 格式:
|
||
/// ```text
|
||
/// <session-context>
|
||
/// key1: value1
|
||
/// key2: value2
|
||
/// </session-context>
|
||
/// ```
|
||
pub async fn snapshot(&self) -> Result<String, AgentError> {
|
||
let filter = MemoryFilter {
|
||
prefix: Some(format!("{}:", self.namespace)),
|
||
..Default::default()
|
||
};
|
||
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
|
||
|
||
let mut lines = Vec::with_capacity(items.len() + 2);
|
||
lines.push("<session-context>".to_string());
|
||
for item in items {
|
||
// 从 id 中提取原始 key(去掉 namespace 前缀)
|
||
let key = item
|
||
.id
|
||
.strip_prefix(&format!("{}:", self.namespace))
|
||
.unwrap_or(&item.id);
|
||
lines.push(format!("{}: {}", key, item.content));
|
||
}
|
||
lines.push("</session-context>".to_string());
|
||
|
||
Ok(lines.join("\n"))
|
||
}
|
||
|
||
/// 删除指定 key。
|
||
pub async fn remove(&self, key: &str) -> Result<(), AgentError> {
|
||
self.store
|
||
.delete(&self.internal_key(key))
|
||
.await
|
||
.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 {
|
||
prefix: Some(format!("{}:", self.namespace)),
|
||
..Default::default()
|
||
};
|
||
let items = self.store.list(&filter).await.map_err(AgentError::Memory)?;
|
||
|
||
for item in items {
|
||
self.store
|
||
.delete(&item.id)
|
||
.await
|
||
.map_err(AgentError::Memory)?;
|
||
}
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
use crate::memory::store::InMemoryStore;
|
||
|
||
fn make_store() -> Arc<dyn MemoryStore> {
|
||
Arc::new(InMemoryStore::new())
|
||
}
|
||
|
||
/// 烟雾测试 1:set / get / remove 基本读写。
|
||
#[tokio::test]
|
||
async fn set_get_remove() {
|
||
let mem = SessionMemory::new(make_store(), "test-session");
|
||
|
||
assert!(mem.get("k").await.unwrap().is_none());
|
||
|
||
mem.set("k", "v").await.unwrap();
|
||
assert_eq!(mem.get("k").await.unwrap(), Some("v".into()));
|
||
|
||
mem.remove("k").await.unwrap();
|
||
assert!(mem.get("k").await.unwrap().is_none());
|
||
}
|
||
|
||
/// 烟雾测试 2:snapshot 格式化输出。
|
||
#[tokio::test]
|
||
async fn snapshot_format() {
|
||
let mem = SessionMemory::new(make_store(), "s1");
|
||
mem.set("design", "PostgreSQL").await.unwrap();
|
||
mem.set("lang", "Rust").await.unwrap();
|
||
|
||
let snap = mem.snapshot().await.unwrap();
|
||
assert!(snap.contains("<session-context>"));
|
||
assert!(snap.contains("</session-context>"));
|
||
assert!(snap.contains("design: PostgreSQL"));
|
||
assert!(snap.contains("lang: Rust"));
|
||
}
|
||
|
||
/// 烟雾测试 3:clear 清空当前 namespace。
|
||
#[tokio::test]
|
||
async fn clear_only_affects_own_namespace() {
|
||
let store = make_store();
|
||
let mem_a = SessionMemory::new(store.clone(), "a");
|
||
let mem_b = SessionMemory::new(store.clone(), "b");
|
||
|
||
mem_a.set("key", "val_a").await.unwrap();
|
||
mem_b.set("key", "val_b").await.unwrap();
|
||
|
||
mem_a.clear().await.unwrap();
|
||
|
||
assert!(mem_a.get("key").await.unwrap().is_none());
|
||
assert_eq!(mem_b.get("key").await.unwrap(), Some("val_b".into()));
|
||
}
|
||
}
|