klieo-memory-pgvector 2.3.0

PostgreSQL + pgvector implementation of klieo-core's LongTermMemory.
Documentation
//! `PgvectorLongTerm` — `LongTermMemory` over one Postgres table with a
//! `vector(dim)` column (HNSW, cosine) and `scope_kind` / `scope_value` filter
//! columns. Metadata is stored as `jsonb`.
//!
//! The table is provisioned lazily on first `remember`; a read that races
//! ahead of it maps the `undefined_table` error to an empty result.

use crate::client::PgvectorHandle;
use crate::embedder::Embedder;
use crate::error::{is_undefined_table, store_err};
use async_trait::async_trait;
use klieo_core::error::MemoryError;
use klieo_core::ids::FactId;
use klieo_core::memory::{Fact, LongTermMemory, Scope};
use klieo_memory_graph::FilterableLongTermMemory;
use sqlx::Row;
use std::sync::Arc;
use uuid::Uuid;

fn scope_columns(scope: &Scope) -> (&'static str, String) {
    match scope {
        Scope::Workspace(s) => ("workspace", s.clone()),
        Scope::Agent(s) => ("agent", s.clone()),
        Scope::Global => ("global", String::new()),
    }
}

fn row_to_fact(text: String, metadata: serde_json::Value) -> Fact {
    Fact::new(text).with_metadata(metadata)
}

/// Render an embedding as a pgvector text literal (`[0.1,0.2,...]`). Bound as a
/// plain `text` parameter and cast with `::vector` in SQL, which avoids a
/// dedicated pgvector client crate (and the sqlx-driver version pin it forces).
fn vector_literal(embedding: &[f32]) -> String {
    let mut out = String::with_capacity(embedding.len() * 8 + 2);
    out.push('[');
    for (i, value) in embedding.iter().enumerate() {
        if i > 0 {
            out.push(',');
        }
        out.push_str(&value.to_string());
    }
    out.push(']');
    out
}

/// Parse a caller-facing [`FactId`] back into the `uuid` column type. Ids are
/// minted as UUIDs by [`PgvectorLongTerm::remember`], so a parse failure means
/// a foreign id and is surfaced as a typed error rather than silently dropped.
fn parse_fact_id(id: &FactId) -> Result<Uuid, MemoryError> {
    Uuid::parse_str(&id.to_string())
        .map_err(|e| MemoryError::Store(format!("fact id is not a uuid: {e}")))
}

/// `LongTermMemory` + `FilterableLongTermMemory` over the single table
/// described in the module docs; cosine recall via the HNSW index, scope
/// enforced by the `scope_kind`/`scope_value` columns.
pub struct PgvectorLongTerm {
    handle: PgvectorHandle,
    embedder: Arc<dyn Embedder>,
    embedder_id: String,
}

impl PgvectorLongTerm {
    pub(crate) fn new(
        handle: PgvectorHandle,
        embedder: Arc<dyn Embedder>,
        embedder_id: String,
    ) -> Self {
        Self {
            handle,
            embedder,
            embedder_id,
        }
    }

    /// Embed one text and validate it against the embedder's declared
    /// dimension before it reaches the `vector(dim)` column.
    async fn embed_one(&self, text: &str) -> Result<Vec<f32>, MemoryError> {
        let dim = self.embedder.dimension();
        let vector = self
            .embedder
            .embed(&[text.to_string()])
            .await?
            .into_iter()
            .next()
            .ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
        if vector.len() != dim {
            return Err(MemoryError::Embedding(format!(
                "embedder produced {}-dim vector, expected {}",
                vector.len(),
                dim
            )));
        }
        Ok(vector)
    }
}

#[async_trait]
impl LongTermMemory for PgvectorLongTerm {
    async fn remember(&self, scope: Scope, fact: Fact) -> Result<FactId, MemoryError> {
        self.handle
            .ensure_table(self.embedder.dimension() as u64)
            .await?;
        let vector = self.embed_one(&fact.text).await?;
        let id = Uuid::new_v4();
        let (kind, value) = scope_columns(&scope);
        let sql = format!(
            "INSERT INTO {} (fact_id, text, metadata, embedding, scope_kind, scope_value) \
             VALUES ($1, $2, $3, $4::vector, $5, $6)",
            self.handle.table
        );
        sqlx::query(&sql)
            .bind(id)
            .bind(&fact.text)
            .bind(&fact.metadata)
            .bind(vector_literal(&vector))
            .bind(kind)
            .bind(value)
            .execute(&self.handle.pool)
            .await
            .map_err(store_err)?;
        Ok(FactId::new(id.to_string()))
    }

    async fn recall(&self, scope: Scope, query: &str, k: usize) -> Result<Vec<Fact>, MemoryError> {
        if k == 0 {
            return Ok(Vec::new());
        }
        let vector = self.embed_one(query).await?;
        let (kind, value) = scope_columns(&scope);
        let sql = format!(
            "SELECT text, metadata FROM {} \
             WHERE scope_kind = $1 AND scope_value = $2 \
             ORDER BY embedding <=> $3::vector LIMIT $4",
            self.handle.table
        );
        let rows = sqlx::query(&sql)
            .bind(kind)
            .bind(value)
            .bind(vector_literal(&vector))
            .bind(k as i64)
            .fetch_all(&self.handle.pool)
            .await;
        rows_to_facts(rows)
    }

    async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
        let uuid = parse_fact_id(&id)?;
        let sql = format!("DELETE FROM {} WHERE fact_id = $1", self.handle.table);
        match sqlx::query(&sql)
            .bind(uuid)
            .execute(&self.handle.pool)
            .await
        {
            Ok(_) => Ok(()),
            Err(e) if is_undefined_table(&e) => Ok(()),
            Err(e) => Err(store_err(e)),
        }
    }
}

#[async_trait]
impl FilterableLongTermMemory for PgvectorLongTerm {
    async fn recall_filtered(
        &self,
        scope: Scope,
        query: &str,
        k: usize,
        candidate_ids: &[FactId],
    ) -> Result<Vec<Fact>, MemoryError> {
        if k == 0 || candidate_ids.is_empty() {
            return Ok(Vec::new());
        }
        let candidates: Vec<Uuid> = candidate_ids
            .iter()
            .map(parse_fact_id)
            .collect::<Result<_, _>>()?;
        let vector = self.embed_one(query).await?;
        let (kind, value) = scope_columns(&scope);
        let sql = format!(
            "SELECT text, metadata FROM {} \
             WHERE scope_kind = $1 AND scope_value = $2 AND fact_id = ANY($3) \
             ORDER BY embedding <=> $4::vector LIMIT $5",
            self.handle.table
        );
        let rows = sqlx::query(&sql)
            .bind(kind)
            .bind(value)
            .bind(candidates)
            .bind(vector_literal(&vector))
            .bind(k as i64)
            .fetch_all(&self.handle.pool)
            .await;
        rows_to_facts(rows)
    }

    fn embedder_id(&self) -> &str {
        &self.embedder_id
    }
}

/// Decode a query result into `Vec<Fact>`, mapping `undefined_table` (a read
/// that raced ahead of the first `remember`) to an empty result.
fn rows_to_facts(
    rows: Result<Vec<sqlx::postgres::PgRow>, sqlx::Error>,
) -> Result<Vec<Fact>, MemoryError> {
    let rows = match rows {
        Ok(rows) => rows,
        Err(e) if is_undefined_table(&e) => return Ok(Vec::new()),
        Err(e) => return Err(store_err(e)),
    };
    rows.into_iter()
        .map(|row| {
            let text: String = row.try_get("text").map_err(store_err)?;
            let metadata: serde_json::Value = row.try_get("metadata").map_err(store_err)?;
            Ok(row_to_fact(text, metadata))
        })
        .collect()
}

#[cfg(test)]
mod tests {
    use super::*;

    #[allow(dead_code)]
    fn _trait_object_safe(_: &dyn LongTermMemory) {}

    #[test]
    fn scope_columns_maps_each_variant() {
        assert_eq!(
            scope_columns(&Scope::Workspace("ws".into())),
            ("workspace", "ws".to_string())
        );
        assert_eq!(
            scope_columns(&Scope::Agent("a".into())),
            ("agent", "a".to_string())
        );
        assert_eq!(scope_columns(&Scope::Global), ("global", String::new()));
    }

    #[test]
    fn row_to_fact_carries_text_and_metadata() {
        let fact = row_to_fact("hi".to_string(), serde_json::json!({"k": 1}));
        assert_eq!(fact.text, "hi");
        assert_eq!(fact.metadata, serde_json::json!({"k": 1}));
    }

    #[test]
    fn parse_fact_id_rejects_non_uuid() {
        let err = parse_fact_id(&FactId::new("not-a-uuid".to_string())).unwrap_err();
        assert!(matches!(err, MemoryError::Store(_)));
    }

    #[test]
    fn vector_literal_renders_pgvector_text_form() {
        assert_eq!(vector_literal(&[]), "[]");
        assert_eq!(vector_literal(&[1.0]), "[1]");
        assert_eq!(vector_literal(&[0.5, -1.25, 2.0]), "[0.5,-1.25,2]");
    }

    #[test]
    fn parse_fact_id_accepts_minted_uuid() {
        let id = uuid::Uuid::new_v4().to_string();
        let parsed = parse_fact_id(&FactId::new(id.clone())).unwrap();
        assert_eq!(parsed.to_string(), id);
    }
}