klieo_memory_sqlite/
factory.rs1use crate::connection::DbHandle;
5use crate::embedder::Embedder;
6use crate::episodic::SqliteEpisodic;
7use crate::long_term::SqliteLongTerm;
8use crate::short_term::SqliteShortTerm;
9use klieo_core::error::MemoryError;
10use klieo_core::memory::{EpisodicMemory, LongTermMemory, MemoryHandles, ShortTermMemory};
11use std::path::Path;
12use std::sync::Arc;
13
14pub struct MemorySqlite {
17 pub short_term: Arc<dyn ShortTermMemory>,
19 pub long_term: Arc<dyn LongTermMemory>,
21 pub episodic: Arc<dyn EpisodicMemory>,
23}
24
25impl MemorySqlite {
26 pub async fn new(
30 path: impl AsRef<Path>,
31 embedder: Arc<dyn Embedder>,
32 ) -> Result<Self, MemoryError> {
33 let db = DbHandle::open(path).await?;
34 #[cfg(feature = "sqlite-vec")]
35 {
36 db.create_vec_table(embedder.dimension()).await?;
37 }
38 let short_term: Arc<dyn ShortTermMemory> = Arc::new(SqliteShortTerm::new(db.clone()));
39 let long_term: Arc<dyn LongTermMemory> =
40 Arc::new(SqliteLongTerm::new(db.clone(), embedder));
41 let episodic: Arc<dyn EpisodicMemory> = Arc::new(SqliteEpisodic::new(db));
42 Ok(Self {
43 short_term,
44 long_term,
45 episodic,
46 })
47 }
48}
49
50impl From<MemorySqlite> for MemoryHandles {
51 fn from(value: MemorySqlite) -> Self {
52 MemoryHandles::new(value.short_term, value.long_term, value.episodic)
53 }
54}
55
56#[cfg(test)]
57mod tests {
58 use super::*;
59 use crate::embedder::DummyEmbedder;
60 use klieo_core::ids::{RunId, ThreadId};
61 use klieo_core::llm::{Message, Role};
62 use klieo_core::memory::{Episode, Fact, Scope};
63
64 #[tokio::test]
65 async fn factory_wires_all_three_traits_against_shared_db() {
66 let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
67 .await
68 .unwrap();
69
70 let thread = ThreadId::new("t1");
72 mem.short_term
73 .append(
74 thread.clone(),
75 Message {
76 role: Role::User,
77 content: "hi".into(),
78 tool_calls: vec![],
79 tool_call_id: None,
80 },
81 )
82 .await
83 .unwrap();
84 let loaded = mem.short_term.load(thread, 1000).await.unwrap();
85 assert_eq!(loaded.len(), 1);
86
87 let id = mem
89 .long_term
90 .remember(
91 Scope::Global,
92 Fact {
93 text: "shared db round trip".into(),
94 metadata: serde_json::Value::Null,
95 },
96 )
97 .await
98 .unwrap();
99 let hits = mem
100 .long_term
101 .recall(Scope::Global, "anything", 10)
102 .await
103 .unwrap();
104 assert_eq!(hits.len(), 1);
105 mem.long_term.forget(id).await.unwrap();
106
107 let run = RunId::new();
109 mem.episodic
110 .record(
111 run,
112 Episode::Started {
113 agent: "test".into(),
114 },
115 )
116 .await
117 .unwrap();
118 mem.episodic.record(run, Episode::Completed).await.unwrap();
119 let replay = mem.episodic.replay(run).await.unwrap();
120 assert_eq!(replay.len(), 2);
121 }
122
123 #[tokio::test]
124 async fn memory_handles_from_factory_preserves_handles() {
125 let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
126 .await
127 .unwrap();
128 let st_ptr = Arc::as_ptr(&mem.short_term);
129 let lt_ptr = Arc::as_ptr(&mem.long_term);
130 let ep_ptr = Arc::as_ptr(&mem.episodic);
131 let handles: MemoryHandles = mem.into();
132 assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.short_term), st_ptr));
133 assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.long_term), lt_ptr));
134 assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.episodic), ep_ptr));
135 }
136}