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(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        // Episodic
117        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}