klieo_memory_sqlite/
factory.rs1use crate::connection::DbHandle;
5use crate::embedder::{DummyEmbedder, 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
14#[non_exhaustive]
17pub struct MemorySqlite {
18 pub short_term: Arc<dyn ShortTermMemory>,
20 pub long_term: Arc<dyn LongTermMemory>,
22 pub episodic: Arc<dyn EpisodicMemory>,
24}
25
26impl MemorySqlite {
27 pub async fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
35 Self::new(path, Arc::new(DummyEmbedder)).await
36 }
37
38 pub async fn new(
45 path: impl AsRef<Path>,
46 embedder: Arc<dyn Embedder>,
47 ) -> Result<Self, MemoryError> {
48 let db = DbHandle::open(path).await?;
49 #[cfg(feature = "sqlite-vec")]
50 {
51 db.create_vec_table(embedder.dimension()).await?;
52 }
53 let short_term: Arc<dyn ShortTermMemory> = Arc::new(SqliteShortTerm::new(db.clone()));
54 let long_term: Arc<dyn LongTermMemory> =
55 Arc::new(SqliteLongTerm::new(db.clone(), embedder));
56 let episodic: Arc<dyn EpisodicMemory> = Arc::new(SqliteEpisodic::new(db));
57 Ok(Self {
58 short_term,
59 long_term,
60 episodic,
61 })
62 }
63}
64
65impl From<MemorySqlite> for MemoryHandles {
66 fn from(value: MemorySqlite) -> Self {
67 MemoryHandles::new(value.short_term, value.long_term, value.episodic)
68 }
69}
70
71#[cfg(test)]
72mod tests {
73 use super::*;
74 use crate::embedder::DummyEmbedder;
75 use klieo_core::ids::{RunId, ThreadId};
76 use klieo_core::llm::{Message, Role};
77 use klieo_core::memory::{Episode, Fact, Scope};
78
79 #[tokio::test]
80 async fn factory_wires_all_three_traits_against_shared_db() {
81 let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
82 .await
83 .unwrap();
84
85 let thread = ThreadId::new("t1");
87 mem.short_term
88 .append(
89 thread.clone(),
90 Message {
91 role: Role::User,
92 content: "hi".into(),
93 tool_calls: vec![],
94 tool_call_id: None,
95 },
96 )
97 .await
98 .unwrap();
99 let loaded = mem.short_term.load(thread, 1000).await.unwrap();
100 assert_eq!(loaded.len(), 1);
101
102 let id = mem
104 .long_term
105 .remember(Scope::Global, Fact::new("shared db round trip"))
106 .await
107 .unwrap();
108 let hits = mem
109 .long_term
110 .recall(Scope::Global, "anything", 10)
111 .await
112 .unwrap();
113 assert_eq!(hits.len(), 1);
114 mem.long_term.forget(id).await.unwrap();
115
116 let run = RunId::new();
118 mem.episodic
119 .record(
120 run,
121 Episode::Started {
122 agent: "test".into(),
123 },
124 )
125 .await
126 .unwrap();
127 mem.episodic.record(run, Episode::Completed).await.unwrap();
128 let replay = mem.episodic.replay(run).await.unwrap();
129 assert_eq!(replay.len(), 2);
130 }
131
132 #[tokio::test]
133 async fn open_propagates_storage_error_on_unwritable_path() {
134 let result = MemorySqlite::open("/this/path/does/not/exist/db.sqlite").await;
135 assert!(
136 result.is_err(),
137 "opening under a non-existent parent directory must fail"
138 );
139 }
140
141 #[tokio::test]
142 async fn open_uses_dummy_embedder_by_default() {
143 let mem = MemorySqlite::open(":memory:").await.unwrap();
144 let id = mem
145 .long_term
146 .remember(Scope::Global, Fact::new("open-default round trip"))
147 .await
148 .unwrap();
149 let hits = mem
150 .long_term
151 .recall(Scope::Global, "anything", 10)
152 .await
153 .unwrap();
154 assert_eq!(hits.len(), 1);
155 mem.long_term.forget(id).await.unwrap();
156 }
157
158 #[tokio::test]
159 async fn memory_handles_from_factory_preserves_handles() {
160 let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
161 .await
162 .unwrap();
163 let st_ptr = Arc::as_ptr(&mem.short_term);
164 let lt_ptr = Arc::as_ptr(&mem.long_term);
165 let ep_ptr = Arc::as_ptr(&mem.episodic);
166 let handles: MemoryHandles = mem.into();
167 assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.short_term), st_ptr));
168 assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.long_term), lt_ptr));
169 assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.episodic), ep_ptr));
170 }
171}