use crate::connection::DbHandle;
use crate::embedder::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, ShortTermMemory};
use std::path::Path;
use std::sync::Arc;
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 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,
})
}
}
#[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 {
text: "shared db round trip".into(),
metadata: serde_json::Value::Null,
},
)
.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);
}
}