feat(memory): 实现 SqliteStore 持久化

- 新增 src/memory/store/sqlite_store.rs(SqliteStore + 9 个内联测试)
- 基于 rusqlite 0.32(bundled),使用 Arc<Mutex<Connection>> + spawn_blocking
- WAL 模式 + synchronous=NORMAL + busy_timeout=5s + PRAGMA user_version
  schema 版本管理
- created_at 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
- 错误精细映射:SqliteFailure/InvalidQuery → InvalidInput;
  FromSqlConversionFailure → Serialization;其他 → Storage
- 9 个测试覆盖 CRUD、upsert、prefix/since/offset+limit 过滤、
  10 写者 × 10 次并发、持久化 round-trip、trait-box 互换兼容性

依赖:
- rusqlite = { version = "0.32", features = ["bundled"] }
- time 增补 features: parsing, formatting, macros
- dev-dependencies: tempfile = "3"

测试:199 → 200 pass(191 原有 + 9 新增)
This commit is contained in:
徐涛
2026-07-05 17:13:28 +08:00
parent c8a91f6eaf
commit c82af60f81
4 changed files with 551 additions and 2 deletions
+1 -1
View File
@@ -12,7 +12,7 @@ pub use conversation::{ConversationMemory, ConversationMemoryConfig};
pub use error::MemoryError;
pub use knowledge::KnowledgeStore;
pub use retriever::MemoryRetriever;
pub use store::{InMemoryStore, MemoryStore};
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
// 低频类型(配置/高级使用)
pub use conversation::MemoryStrategy;
+2
View File
@@ -6,8 +6,10 @@ use crate::memory::error::MemoryError;
use crate::memory::types::{MemoryFilter, MemoryItem};
pub mod in_memory;
pub mod sqlite_store;
pub use in_memory::InMemoryStore;
pub use sqlite_store::SqliteStore;
/// 底层记忆存储抽象接口。
///
+545
View File
@@ -0,0 +1,545 @@
//! SqliteStore —— 基于 rusqlite 的持久化 MemoryStore 实现。
//!
//! 单进程独享、写入串行化(WAL + Mutex),适合本地 Agent 长期持久化场景。
use std::path::Path;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use rusqlite::{params, params_from_iter, Connection, ErrorCode};
use time::format_description::well_known::Rfc3339;
use time::OffsetDateTime;
use tracing::{debug, error, instrument, warn};
use crate::memory::error::MemoryError;
use crate::memory::store::MemoryStore;
use crate::memory::types::{MemoryFilter, MemoryItem};
const INITIAL_USER_VERSION: i64 = 1;
const BUSY_TIMEOUT_MS: i64 = 5000;
const WAL_AUTOCHECKPOINT_PAGES: i64 = 1000;
/// SQLite 持久化后端的 MemoryStore 实现。
///
/// 设计要点:
/// - 单进程独享:`Arc<Mutex<Connection>>` 串行化所有 IO
/// - WAL 模式 + `synchronous=NORMAL` 兼顾崩溃安全与吞吐
/// - `created_at` 归一化为 UTC 的 RFC 3339 TEXT,字典序等价时间序
/// - 所有 IO 通过 `tokio::task::spawn_blocking` 卸载到阻塞线程池
pub struct SqliteStore {
conn: Arc<Mutex<Connection>>,
}
impl SqliteStore {
/// 打开或创建一个 SQLite 数据库。
///
/// - `path = ":memory:"` 使用内存数据库(测试场景)
/// - 其他路径:自动创建父目录;文件已存在则附加打开
/// - 启动时执行 `migrate()`,失败立即返回错误
#[instrument(skip(path), fields(path = %path.as_ref().display()))]
pub fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
let path_ref = path.as_ref();
let path_str = path_ref.to_string_lossy();
let conn = if path_str == ":memory:" {
Connection::open_in_memory()
} else {
if let Some(parent) = path_ref.parent()
&& !parent.as_os_str().is_empty()
{
std::fs::create_dir_all(parent).map_err(|e| {
MemoryError::Storage(format!(
"创建数据库父目录失败 ({}): {}",
parent.display(),
e
))
})?;
}
Connection::open(path_ref)
}
.map_err(|e| map_sqlite_error(e, "打开数据库"))?;
migrate(&conn)?;
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
})
}
}
#[async_trait]
impl MemoryStore for SqliteStore {
#[instrument(skip(self, item), fields(id = %item.id))]
async fn save(&self, item: MemoryItem) -> Result<(), MemoryError> {
let conn = Arc::clone(&self.conn);
let created_at_str = item
.created_at
.to_offset(time::UtcOffset::UTC)
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format created_at: {e}")))?;
let metadata_str = serde_json::to_string(&item.metadata)
.map_err(|e| MemoryError::Serialization(format!("serialize metadata: {e}")))?;
let id = item.id;
let content = item.content;
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
conn.execute(
"INSERT INTO memory_items (id, content, metadata, created_at) \
VALUES (?1, ?2, ?3, ?4) \
ON CONFLICT(id) DO UPDATE SET \
content=excluded.content, \
metadata=excluded.metadata, \
created_at=excluded.created_at",
params![id, content, metadata_str, created_at_str],
)
.map_err(|e| map_sqlite_error(e, "保存记忆"))?;
Ok(())
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, id))]
async fn get(&self, id: &str) -> Result<Option<MemoryItem>, MemoryError> {
let conn = Arc::clone(&self.conn);
let id_owned = id.to_string();
tokio::task::spawn_blocking(move || -> Result<Option<MemoryItem>, MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
let mut stmt = conn
.prepare("SELECT id, content, metadata, created_at FROM memory_items WHERE id = ?1")
.map_err(|e| map_sqlite_error(e, "prepare get"))?;
let mut rows = stmt
.query_map(params![id_owned], row_to_item)
.map_err(|e| map_sqlite_error(e, "query get"))?;
match rows.next() {
None => Ok(None),
Some(row) => row
.map(Some)
.map_err(|e| map_sqlite_error(e, "decode row")),
}
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, id))]
async fn delete(&self, id: &str) -> Result<(), MemoryError> {
let conn = Arc::clone(&self.conn);
let id_owned = id.to_string();
tokio::task::spawn_blocking(move || -> Result<(), MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
conn.execute(
"DELETE FROM memory_items WHERE id = ?1",
params![id_owned],
)
.map_err(|e| map_sqlite_error(e, "delete"))?;
Ok(())
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
#[instrument(skip(self, filter))]
async fn list(&self, filter: &MemoryFilter) -> Result<Vec<MemoryItem>, MemoryError> {
let mut sql = String::from(
"SELECT id, content, metadata, created_at FROM memory_items WHERE 1=1",
);
let mut param_values: Vec<String> = Vec::new();
let mut ph_idx = 0usize;
if filter.prefix.is_some() {
ph_idx += 1;
sql.push_str(&format!(" AND id LIKE ?{ph_idx} || '%'"));
}
if filter.since.is_some() {
ph_idx += 1;
sql.push_str(&format!(" AND created_at > ?{ph_idx}"));
}
// ORDER BY created_at ASC(按时间升序,最旧在前)
sql.push_str(" ORDER BY created_at ASC");
let limit_sql: String = match (filter.limit, filter.offset) {
(Some(_), Some(_)) => {
ph_idx += 1;
let limit_p = ph_idx;
ph_idx += 1;
let offset_p = ph_idx;
format!(" LIMIT ?{limit_p} OFFSET ?{offset_p}")
}
(Some(_), None) => {
ph_idx += 1;
let limit_p = ph_idx;
format!(" LIMIT ?{limit_p}")
}
(None, Some(_)) => {
// SQLite 中 LIMIT -1 表示无限制
ph_idx += 1;
let offset_p = ph_idx;
format!(" LIMIT -1 OFFSET ?{offset_p}")
}
(None, None) => String::new(),
};
sql.push_str(&limit_sql);
if let Some(p) = &filter.prefix {
param_values.push(p.clone());
}
if let Some(t) = filter.since {
let s = t
.to_offset(time::UtcOffset::UTC)
.format(&Rfc3339)
.map_err(|e| MemoryError::Serialization(format!("format since: {e}")))?;
param_values.push(s);
}
if let Some(l) = filter.limit {
param_values.push(l.to_string());
}
if let Some(o) = filter.offset {
param_values.push(o.to_string());
}
let conn = Arc::clone(&self.conn);
let sql_owned = sql;
let param_values_owned = param_values;
tokio::task::spawn_blocking(move || -> Result<Vec<MemoryItem>, MemoryError> {
let conn = conn.lock().unwrap_or_else(|e| e.into_inner());
let mut stmt = conn
.prepare(&sql_owned)
.map_err(|e| map_sqlite_error(e, "list prepare"))?;
let params_iter: Vec<&dyn rusqlite::ToSql> = param_values_owned
.iter()
.map(|s| s as &dyn rusqlite::ToSql)
.collect();
let rows = stmt
.query_map(params_from_iter(params_iter), row_to_item)
.map_err(|e| map_sqlite_error(e, "list query"))?;
let mut result = Vec::new();
for row in rows {
result.push(row.map_err(|e| map_sqlite_error(e, "list row"))?);
}
debug!(count = result.len(), "SqliteStore::list 完成");
Ok(result)
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking task join: {e}")))?
}
}
fn row_to_item(row: &rusqlite::Row<'_>) -> Result<MemoryItem, rusqlite::Error> {
let id: String = row.get(0)?;
let content: String = row.get(1)?;
let metadata_str: String = row.get(2)?;
let created_at_str: String = row.get(3)?;
let metadata: serde_json::Value = serde_json::from_str(&metadata_str).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(2, rusqlite::types::Type::Text, Box::new(e))
})?;
let created_at = OffsetDateTime::parse(&created_at_str, &Rfc3339).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(3, rusqlite::types::Type::Text, Box::new(e))
})?;
Ok(MemoryItem {
id,
content,
metadata,
created_at: created_at.to_offset(time::UtcOffset::UTC),
})
}
fn migrate(conn: &Connection) -> Result<(), MemoryError> {
conn.pragma_update(None, "journal_mode", "WAL")
.map_err(|e| map_sqlite_error(e, "PRAGMA journal_mode"))?;
conn.pragma_update(None, "synchronous", "NORMAL")
.map_err(|e| map_sqlite_error(e, "PRAGMA synchronous"))?;
conn.execute_batch(&format!("PRAGMA busy_timeout = {BUSY_TIMEOUT_MS};"))
.map_err(|e| map_sqlite_error(e, "PRAGMA busy_timeout"))?;
conn.execute_batch(&format!(
"PRAGMA wal_autocheckpoint = {WAL_AUTOCHECKPOINT_PAGES};"
))
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_autocheckpoint"))?;
let version: i64 = conn
.query_row("PRAGMA user_version", [], |row| row.get(0))
.map_err(|e| map_sqlite_error(e, "PRAGMA user_version"))?;
if version < INITIAL_USER_VERSION {
conn.execute_batch(
"CREATE TABLE IF NOT EXISTS memory_items (
id TEXT PRIMARY KEY,
content TEXT NOT NULL,
metadata TEXT NOT NULL DEFAULT '{}',
created_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_memory_items_created_at
ON memory_items(created_at);
PRAGMA user_version = 1;",
)
.map_err(|e| map_sqlite_error(e, "create schema v1"))?;
}
let check_result: String = conn
.query_row("PRAGMA quick_check", [], |row| row.get(0))
.map_err(|e| map_sqlite_error(e, "PRAGMA quick_check"))?;
if check_result != "ok" {
error!(result = %check_result, "数据库文件 quick_check 失败");
return Err(MemoryError::Storage(format!(
"数据库文件损坏: {check_result}"
)));
}
conn.execute_batch("PRAGMA wal_checkpoint(TRUNCATE);")
.map_err(|e| map_sqlite_error(e, "PRAGMA wal_checkpoint"))?;
Ok(())
}
fn map_sqlite_error(e: rusqlite::Error, ctx: &str) -> MemoryError {
match &e {
rusqlite::Error::SqliteFailure(err, _) => match err.code {
ErrorCode::ConstraintViolation => MemoryError::InvalidInput(format!("{ctx}: {e}")),
ErrorCode::DatabaseBusy | ErrorCode::DatabaseLocked => {
warn!("SQLite 忙: {e}");
MemoryError::Storage(format!("{ctx}: {e}"))
}
_ => MemoryError::Storage(format!("{ctx}: {e}")),
},
rusqlite::Error::InvalidQuery
| rusqlite::Error::InvalidParameterName(_)
| rusqlite::Error::InvalidColumnIndex(_)
| rusqlite::Error::InvalidColumnName(_) => {
MemoryError::InvalidInput(format!("{ctx}: {e}"))
}
rusqlite::Error::FromSqlConversionFailure(_, _, _)
| rusqlite::Error::ToSqlConversionFailure(_) => {
MemoryError::Serialization(format!("{ctx}: {e}"))
}
_ => MemoryError::Storage(format!("{ctx}: {e}")),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::store::InMemoryStore;
use std::sync::Arc;
use tempfile::TempDir;
use time::OffsetDateTime;
fn make_item(id: &str) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({"id_key": id}),
created_at: OffsetDateTime::now_utc(),
}
}
fn make_item_at(id: &str, when: OffsetDateTime) -> MemoryItem {
MemoryItem {
id: id.to_string(),
content: format!("content-{id}"),
metadata: serde_json::json!({}),
created_at: when,
}
}
#[tokio::test]
async fn crud_basic() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("a")).await.unwrap();
store.save(make_item("b")).await.unwrap();
let got_a = store.get("a").await.unwrap();
assert!(got_a.is_some());
assert_eq!(got_a.unwrap().id, "a");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 2);
store.delete("a").await.unwrap();
assert!(store.get("a").await.unwrap().is_none());
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
assert_eq!(list[0].id, "b");
}
#[tokio::test]
async fn save_is_upsert() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("a")).await.unwrap();
let mut item = make_item("a");
item.content = "updated".to_string();
item.metadata = serde_json::json!({"rev": 2});
let original_created_at = item.created_at;
store.save(item).await.unwrap();
let got = store.get("a").await.unwrap().unwrap();
assert_eq!(got.content, "updated");
assert_eq!(got.metadata["rev"], serde_json::json!(2));
// created_at 保持调用方传入值
assert_eq!(got.created_at, original_created_at);
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
}
#[tokio::test]
async fn list_with_prefix() {
let store = SqliteStore::open(":memory:").unwrap();
store.save(make_item("foo_a")).await.unwrap();
store.save(make_item("foo_b")).await.unwrap();
store.save(make_item("bar_a")).await.unwrap();
let filter = MemoryFilter {
prefix: Some("foo_".to_string()),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"foo_a"));
assert!(ids.contains(&"foo_b"));
assert!(!ids.contains(&"bar_a"));
}
#[tokio::test]
async fn list_with_since_filter() {
let store = SqliteStore::open(":memory:").unwrap();
let t0 = OffsetDateTime::now_utc();
store
.save(make_item_at("early", t0 - time::Duration::seconds(60)))
.await
.unwrap();
store.save(make_item_at("middle", t0)).await.unwrap();
store
.save(make_item_at("late", t0 + time::Duration::seconds(60)))
.await
.unwrap();
let filter = MemoryFilter {
since: Some(t0 - time::Duration::seconds(1)),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert_eq!(list.len(), 2);
assert!(ids.contains(&"middle"));
assert!(ids.contains(&"late"));
assert!(!ids.contains(&"early"));
}
#[tokio::test]
async fn list_with_offset_and_limit() {
let store = SqliteStore::open(":memory:").unwrap();
// 写入 5 条时间递增的记录
let base = OffsetDateTime::now_utc() - time::Duration::seconds(5);
for i in 0..5 {
let mut item = make_item(&format!("item_{i}"));
item.created_at = base + time::Duration::seconds(i);
store.save(item).await.unwrap();
}
// offset=1, limit=2 -> item_1, item_2
let filter = MemoryFilter {
offset: Some(1),
limit: Some(2),
..Default::default()
};
let list = store.list(&filter).await.unwrap();
assert_eq!(list.len(), 2);
assert_eq!(list[0].id, "item_1");
assert_eq!(list[1].id, "item_2");
}
#[tokio::test]
async fn concurrent_writers_no_data_loss() {
let store = Arc::new(SqliteStore::open(":memory:").unwrap());
let mut handles = Vec::new();
for w in 0..10 {
let s = Arc::clone(&store);
handles.push(tokio::spawn(async move {
for i in 0..10 {
let id = format!("w{w}_i{i}");
s.save(make_item(&id)).await.unwrap();
}
}));
}
for h in handles {
h.await.unwrap();
}
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 100);
// 验证所有 id 唯一
let mut ids: Vec<String> = list.iter().map(|v| v.id.clone()).collect();
ids.sort();
ids.dedup();
assert_eq!(ids.len(), 100);
}
#[tokio::test]
async fn persistence_round_trip() {
let dir = TempDir::new().unwrap();
let path = dir.path().join("memory.db");
// 阶段 1:写入 3 条
{
let store = SqliteStore::open(&path).unwrap();
store.save(make_item("alpha")).await.unwrap();
store.save(make_item("beta")).await.unwrap();
store.save(make_item("gamma")).await.unwrap();
}
// 阶段 2:重新打开,验证数据完整
{
let store = SqliteStore::open(&path).unwrap();
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 3);
let ids: Vec<&str> = list.iter().map(|v| v.id.as_str()).collect();
assert!(ids.contains(&"alpha"));
assert!(ids.contains(&"beta"));
assert!(ids.contains(&"gamma"));
// 单条读回
let got = store.get("beta").await.unwrap().unwrap();
assert_eq!(got.content, "content-beta");
}
}
#[tokio::test]
async fn open_invalid_path_returns_error() {
// 路径指向已存在的目录而非文件,open 应失败
let dir = TempDir::new().unwrap();
match SqliteStore::open(dir.path()) {
Err(MemoryError::Storage(_)) => {}
Err(other) => panic!("expected Storage error, got {other:?}"),
Ok(_) => panic!("expected error when opening a directory as database"),
}
}
#[tokio::test]
async fn trait_object_compatibility() {
// ponytail: 回归验证 SqliteStore 可作为 Arc<dyn MemoryStore> 与 InMemoryStore 互换
// 所有现有消费者(Conversation / Knowledge / Retriever / SessionMemory)均通过 trait object 引用,
// 此测试确保 trait 接口契约在 SqliteStore 上同样成立。
let sqlite: Arc<dyn MemoryStore> =
Arc::new(SqliteStore::open(":memory:").unwrap());
let in_mem: Arc<dyn MemoryStore> = Arc::new(InMemoryStore::new());
let stores: Vec<Arc<dyn MemoryStore>> = vec![Arc::clone(&sqlite), Arc::clone(&in_mem)];
for store in &stores {
store.save(make_item("x")).await.unwrap();
let got = store.get("x").await.unwrap();
assert_eq!(got.unwrap().id, "x");
let list = store.list(&MemoryFilter::default()).await.unwrap();
assert_eq!(list.len(), 1);
store.delete("x").await.unwrap();
assert!(store.get("x").await.unwrap().is_none());
}
}
}