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::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.
16pub struct MemorySqlite {
17    /// Short-term conversation memory.
18    pub short_term: Arc<dyn ShortTermMemory>,
19    /// Long-term semantic memory (uses the supplied `Embedder`).
20    pub long_term: Arc<dyn LongTermMemory>,
21    /// Episodic event log.
22    pub episodic: Arc<dyn EpisodicMemory>,
23}
24
25impl MemorySqlite {
26    /// Open a SQLite database at `path` (use `":memory:"` for an
27    /// ephemeral DB) and return all three memory handles wired to that
28    /// shared connection.
29    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        // Short-term
71        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        // Long-term
88        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        // Episodic
108        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}