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:
+3
-1
@@ -19,9 +19,11 @@ futures-core = "0.3"
|
|||||||
bytes = "1"
|
bytes = "1"
|
||||||
async-stream = "0.3"
|
async-stream = "0.3"
|
||||||
tokio-util = { version = "0.7", features = ["rt"] }
|
tokio-util = { version = "0.7", features = ["rt"] }
|
||||||
time = { version = "0.3", features = ["serde"] }
|
time = { version = "0.3", features = ["serde", "parsing", "formatting", "macros"] }
|
||||||
|
rusqlite = { version = "0.32", features = ["bundled"] }
|
||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
dotenvy = "0.15.7"
|
dotenvy = "0.15.7"
|
||||||
wiremock = "0.6"
|
wiremock = "0.6"
|
||||||
temp-env = "0.3"
|
temp-env = "0.3"
|
||||||
|
tempfile = "3"
|
||||||
|
|||||||
+1
-1
@@ -12,7 +12,7 @@ pub use conversation::{ConversationMemory, ConversationMemoryConfig};
|
|||||||
pub use error::MemoryError;
|
pub use error::MemoryError;
|
||||||
pub use knowledge::KnowledgeStore;
|
pub use knowledge::KnowledgeStore;
|
||||||
pub use retriever::MemoryRetriever;
|
pub use retriever::MemoryRetriever;
|
||||||
pub use store::{InMemoryStore, MemoryStore};
|
pub use store::{InMemoryStore, MemoryStore, SqliteStore};
|
||||||
|
|
||||||
// 低频类型(配置/高级使用)
|
// 低频类型(配置/高级使用)
|
||||||
pub use conversation::MemoryStrategy;
|
pub use conversation::MemoryStrategy;
|
||||||
|
|||||||
@@ -6,8 +6,10 @@ use crate::memory::error::MemoryError;
|
|||||||
use crate::memory::types::{MemoryFilter, MemoryItem};
|
use crate::memory::types::{MemoryFilter, MemoryItem};
|
||||||
|
|
||||||
pub mod in_memory;
|
pub mod in_memory;
|
||||||
|
pub mod sqlite_store;
|
||||||
|
|
||||||
pub use in_memory::InMemoryStore;
|
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