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:
+1
-1
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
/// 底层记忆存储抽象接口。
|
||||
///
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user