klieo-memory-sqlite 0.4.0

SQLite-backed implementations of klieo-core's memory traits.
Documentation
//! `MemorySqlite` — umbrella factory wiring all three SQLite-backed
//! memory traits to the same database file.

use crate::connection::DbHandle;
use crate::embedder::Embedder;
use crate::episodic::SqliteEpisodic;
use crate::long_term::SqliteLongTerm;
use crate::short_term::SqliteShortTerm;
use klieo_core::error::MemoryError;
use klieo_core::memory::{EpisodicMemory, LongTermMemory, ShortTermMemory};
use std::path::Path;
use std::sync::Arc;

/// Wired-together SQLite-backed memory bundle. Build once, share across
/// agents.
pub struct MemorySqlite {
    /// Short-term conversation memory.
    pub short_term: Arc<dyn ShortTermMemory>,
    /// Long-term semantic memory (uses the supplied `Embedder`).
    pub long_term: Arc<dyn LongTermMemory>,
    /// Episodic event log.
    pub episodic: Arc<dyn EpisodicMemory>,
}

impl MemorySqlite {
    /// Open a SQLite database at `path` (use `":memory:"` for an
    /// ephemeral DB) and return all three memory handles wired to that
    /// shared connection.
    pub async fn new(
        path: impl AsRef<Path>,
        embedder: Arc<dyn Embedder>,
    ) -> Result<Self, MemoryError> {
        let db = DbHandle::open(path).await?;
        #[cfg(feature = "sqlite-vec")]
        {
            db.create_vec_table(embedder.dimension()).await?;
        }
        let short_term: Arc<dyn ShortTermMemory> = Arc::new(SqliteShortTerm::new(db.clone()));
        let long_term: Arc<dyn LongTermMemory> =
            Arc::new(SqliteLongTerm::new(db.clone(), embedder));
        let episodic: Arc<dyn EpisodicMemory> = Arc::new(SqliteEpisodic::new(db));
        Ok(Self {
            short_term,
            long_term,
            episodic,
        })
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::embedder::DummyEmbedder;
    use klieo_core::ids::{RunId, ThreadId};
    use klieo_core::llm::{Message, Role};
    use klieo_core::memory::{Episode, Fact, Scope};

    #[tokio::test]
    async fn factory_wires_all_three_traits_against_shared_db() {
        let mem = MemorySqlite::new(":memory:", Arc::new(DummyEmbedder))
            .await
            .unwrap();

        // Short-term
        let thread = ThreadId::new("t1");
        mem.short_term
            .append(
                thread.clone(),
                Message {
                    role: Role::User,
                    content: "hi".into(),
                    tool_calls: vec![],
                    tool_call_id: None,
                },
            )
            .await
            .unwrap();
        let loaded = mem.short_term.load(thread, 1000).await.unwrap();
        assert_eq!(loaded.len(), 1);

        // Long-term
        let id = mem
            .long_term
            .remember(
                Scope::Global,
                Fact {
                    text: "shared db round trip".into(),
                    metadata: serde_json::Value::Null,
                },
            )
            .await
            .unwrap();
        let hits = mem
            .long_term
            .recall(Scope::Global, "anything", 10)
            .await
            .unwrap();
        assert_eq!(hits.len(), 1);
        mem.long_term.forget(id).await.unwrap();

        // Episodic
        let run = RunId::new();
        mem.episodic
            .record(
                run,
                Episode::Started {
                    agent: "test".into(),
                },
            )
            .await
            .unwrap();
        mem.episodic.record(run, Episode::Completed).await.unwrap();
        let replay = mem.episodic.replay(run).await.unwrap();
        assert_eq!(replay.len(), 2);
    }
}