use crate::connection::DbHandle;
use crate::embedder::{DummyEmbedder, Embedder};
use crate::episodic::SqliteEpisodic;
use crate::long_term::SqliteLongTerm;
use crate::short_term::SqliteShortTerm;
use klieo_core::error::MemoryError;
use klieo_core::memory::{EpisodicMemory, LongTermMemory, MemoryHandles, ShortTermMemory};
use std::path::Path;
use std::sync::Arc;
#[non_exhaustive]
pub struct MemorySqlite {
pub short_term: Arc<dyn ShortTermMemory>,
pub long_term: Arc<dyn LongTermMemory>,
pub episodic: Arc<dyn EpisodicMemory>,
}
impl MemorySqlite {
pub async fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
Self::new(path, Arc::new(DummyEmbedder)).await
}
pub async fn new(
path: impl AsRef<Path>,
embedder: Arc<dyn Embedder>,
) -> Result<Self, MemoryError> {
let db = DbHandle::open(path).await?;
#[cfg(feature = "sqlite-vec")]
{
db.create_vec_table(embedder.dimension()).await?;
}
let short_term: Arc<dyn ShortTermMemory> = Arc::new(SqliteShortTerm::new(db.clone()));
let long_term: Arc<dyn LongTermMemory> =
Arc::new(SqliteLongTerm::new(db.clone(), embedder));
let episodic: Arc<dyn EpisodicMemory> = Arc::new(SqliteEpisodic::new(db));
Ok(Self {
short_term,
long_term,
episodic,
})
}
}
impl From<MemorySqlite> for MemoryHandles {
fn from(value: MemorySqlite) -> Self {
MemoryHandles::new(value.short_term, value.long_term, value.episodic)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::embedder::DummyEmbedder;
use klieo_core::ids::{RunId, ThreadId};
use klieo_core::llm::{Message, Role};
use klieo_core::memory::{Episode, Fact, Scope};
#[tokio::test]
async fn factory_wires_all_three_traits_against_shared_db() {
let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
.await
.unwrap();
let thread = ThreadId::new("t1");
mem.short_term
.append(
thread.clone(),
Message {
role: Role::User,
content: "hi".into(),
tool_calls: vec![],
tool_call_id: None,
},
)
.await
.unwrap();
let loaded = mem.short_term.load(thread, 1000).await.unwrap();
assert_eq!(loaded.len(), 1);
let id = mem
.long_term
.remember(Scope::Global, Fact::new("shared db round trip"))
.await
.unwrap();
let hits = mem
.long_term
.recall(Scope::Global, "anything", 10)
.await
.unwrap();
assert_eq!(hits.len(), 1);
mem.long_term.forget(id).await.unwrap();
let run = RunId::new();
mem.episodic
.record(
run,
Episode::Started {
agent: "test".into(),
},
)
.await
.unwrap();
mem.episodic.record(run, Episode::Completed).await.unwrap();
let replay = mem.episodic.replay(run).await.unwrap();
assert_eq!(replay.len(), 2);
}
#[tokio::test]
async fn open_propagates_storage_error_on_unwritable_path() {
let result = MemorySqlite::open("/this/path/does/not/exist/db.sqlite").await;
assert!(
result.is_err(),
"opening under a non-existent parent directory must fail"
);
}
#[tokio::test]
async fn open_uses_dummy_embedder_by_default() {
let mem = MemorySqlite::open(":memory:").await.unwrap();
let id = mem
.long_term
.remember(Scope::Global, Fact::new("open-default round trip"))
.await
.unwrap();
let hits = mem
.long_term
.recall(Scope::Global, "anything", 10)
.await
.unwrap();
assert_eq!(hits.len(), 1);
mem.long_term.forget(id).await.unwrap();
}
#[tokio::test]
async fn memory_handles_from_factory_preserves_handles() {
let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
.await
.unwrap();
let st_ptr = Arc::as_ptr(&mem.short_term);
let lt_ptr = Arc::as_ptr(&mem.long_term);
let ep_ptr = Arc::as_ptr(&mem.episodic);
let handles: MemoryHandles = mem.into();
assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.short_term), st_ptr));
assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.long_term), lt_ptr));
assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.episodic), ep_ptr));
}
}