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, 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
50#[cfg(test)]
51mod tests {
52    use super::*;
53    use crate::embedder::DummyEmbedder;
54    use klieo_core::ids::{RunId, ThreadId};
55    use klieo_core::llm::{Message, Role};
56    use klieo_core::memory::{Episode, Fact, Scope};
57
58    #[tokio::test]
59    async fn factory_wires_all_three_traits_against_shared_db() {
60        let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
61            .await
62            .unwrap();
63
64        // Short-term
65        let thread = ThreadId::new("t1");
66        mem.short_term
67            .append(
68                thread.clone(),
69                Message {
70                    role: Role::User,
71                    content: "hi".into(),
72                    tool_calls: vec![],
73                    tool_call_id: None,
74                },
75            )
76            .await
77            .unwrap();
78        let loaded = mem.short_term.load(thread, 1000).await.unwrap();
79        assert_eq!(loaded.len(), 1);
80
81        // Long-term
82        let id = mem
83            .long_term
84            .remember(
85                Scope::Global,
86                Fact {
87                    text: "shared db round trip".into(),
88                    metadata: serde_json::Value::Null,
89                },
90            )
91            .await
92            .unwrap();
93        let hits = mem
94            .long_term
95            .recall(Scope::Global, "anything", 10)
96            .await
97            .unwrap();
98        assert_eq!(hits.len(), 1);
99        mem.long_term.forget(id).await.unwrap();
100
101        // Episodic
102        let run = RunId::new();
103        mem.episodic
104            .record(
105                run,
106                Episode::Started {
107                    agent: "test".into(),
108                },
109            )
110            .await
111            .unwrap();
112        mem.episodic.record(run, Episode::Completed).await.unwrap();
113        let replay = mem.episodic.replay(run).await.unwrap();
114        assert_eq!(replay.len(), 2);
115    }
116}