Skip to main content

klieo_memory_sqlite/
factory.rs

1//! `MemorySqlite` — umbrella factory wiring all three SQLite-backed
2//! memory traits to the same database file.
3
4use 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/// Wired-together SQLite-backed memory bundle. Build once, share across
15/// agents.
16#[non_exhaustive]
17pub struct MemorySqlite {
18    /// Short-term conversation memory.
19    pub short_term: Arc<dyn ShortTermMemory>,
20    /// Long-term semantic memory (uses the supplied `Embedder`).
21    pub long_term: Arc<dyn LongTermMemory>,
22    /// Episodic event log.
23    pub episodic: Arc<dyn EpisodicMemory>,
24}
25
26impl MemorySqlite {
27    /// Capability-shaped default — open a SQLite database at `path` with
28    /// [`DummyEmbedder`] (zero-vector, FIFO recall) baked in.
29    ///
30    /// Use `":memory:"` for an ephemeral DB. Equivalent to
31    /// [`Self::new`]`(path, Arc::new(DummyEmbedder))`. Callers needing a
32    /// real embedding model (e.g. `FastEmbedEmbedder` under the
33    /// `fastembed` feature) reach for [`Self::new`] directly.
34    pub async fn open(path: impl AsRef<Path>) -> Result<Self, MemoryError> {
35        Self::new(path, Arc::new(DummyEmbedder)).await
36    }
37
38    /// Open a SQLite database at `path` (use `":memory:"` for an
39    /// ephemeral DB) with a caller-supplied [`Embedder`], and return
40    /// all three memory handles wired to that shared connection.
41    ///
42    /// Prefer [`Self::open`] when you have no specific embedding model
43    /// in mind.
44    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        // Short-term
86        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        // Long-term
103        let id = mem
104            .long_term
105            .remember(
106                Scope::Global,
107                Fact {
108                    text: "shared db round trip".into(),
109                    metadata: serde_json::Value::Null,
110                },
111            )
112            .await
113            .unwrap();
114        let hits = mem
115            .long_term
116            .recall(Scope::Global, "anything", 10)
117            .await
118            .unwrap();
119        assert_eq!(hits.len(), 1);
120        mem.long_term.forget(id).await.unwrap();
121
122        // Episodic
123        let run = RunId::new();
124        mem.episodic
125            .record(
126                run,
127                Episode::Started {
128                    agent: "test".into(),
129                },
130            )
131            .await
132            .unwrap();
133        mem.episodic.record(run, Episode::Completed).await.unwrap();
134        let replay = mem.episodic.replay(run).await.unwrap();
135        assert_eq!(replay.len(), 2);
136    }
137
138    #[tokio::test]
139    async fn open_propagates_storage_error_on_unwritable_path() {
140        let result = MemorySqlite::open("/this/path/does/not/exist/db.sqlite").await;
141        assert!(
142            result.is_err(),
143            "opening under a non-existent parent directory must fail"
144        );
145    }
146
147    #[tokio::test]
148    async fn open_uses_dummy_embedder_by_default() {
149        let mem = MemorySqlite::open(":memory:").await.unwrap();
150        let id = mem
151            .long_term
152            .remember(
153                Scope::Global,
154                Fact {
155                    text: "open-default round trip".into(),
156                    metadata: serde_json::Value::Null,
157                },
158            )
159            .await
160            .unwrap();
161        let hits = mem
162            .long_term
163            .recall(Scope::Global, "anything", 10)
164            .await
165            .unwrap();
166        assert_eq!(hits.len(), 1);
167        mem.long_term.forget(id).await.unwrap();
168    }
169
170    #[tokio::test]
171    async fn memory_handles_from_factory_preserves_handles() {
172        let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
173            .await
174            .unwrap();
175        let st_ptr = Arc::as_ptr(&mem.short_term);
176        let lt_ptr = Arc::as_ptr(&mem.long_term);
177        let ep_ptr = Arc::as_ptr(&mem.episodic);
178        let handles: MemoryHandles = mem.into();
179        assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.short_term), st_ptr));
180        assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.long_term), lt_ptr));
181        assert!(std::ptr::addr_eq(Arc::as_ptr(&handles.episodic), ep_ptr));
182    }
183}