Skip to main content

klieo_memory_pgvector/
long_term.rs

1//! `PgvectorLongTerm` — `LongTermMemory` over one Postgres table with a
2//! `vector(dim)` column (HNSW, cosine) and `scope_kind` / `scope_value` filter
3//! columns. Metadata is stored as `jsonb`.
4//!
5//! The table is provisioned lazily on first `remember`; a read that races
6//! ahead of it maps the `undefined_table` error to an empty result.
7
8use crate::client::PgvectorHandle;
9use crate::embedder::Embedder;
10use crate::error::{is_undefined_table, store_err};
11use async_trait::async_trait;
12use klieo_core::error::MemoryError;
13use klieo_core::ids::FactId;
14use klieo_core::memory::{Fact, LongTermMemory, RecallSemantics, Scope};
15use klieo_memory_graph::FilterableLongTermMemory;
16use sqlx_core::row::Row;
17use std::sync::Arc;
18use uuid::Uuid;
19
20fn scope_columns(scope: &Scope) -> (&'static str, String) {
21    match scope {
22        Scope::Workspace(s) => ("workspace", s.clone()),
23        Scope::Agent(s) => ("agent", s.clone()),
24        Scope::Global => ("global", String::new()),
25    }
26}
27
28fn row_to_fact(text: String, metadata: serde_json::Value) -> Fact {
29    Fact::new(text).with_metadata(metadata)
30}
31
32/// Render an embedding as a pgvector text literal (`[0.1,0.2,...]`). Bound as a
33/// plain `text` parameter and cast with `::vector` in SQL, which avoids a
34/// dedicated pgvector client crate (and the sqlx-driver version pin it forces).
35fn vector_literal(embedding: &[f32]) -> String {
36    let mut out = String::with_capacity(embedding.len() * 8 + 2);
37    out.push('[');
38    for (i, value) in embedding.iter().enumerate() {
39        if i > 0 {
40            out.push(',');
41        }
42        out.push_str(&value.to_string());
43    }
44    out.push(']');
45    out
46}
47
48/// Parse a caller-facing [`FactId`] back into the `uuid` column type. Ids are
49/// minted as UUIDs by [`PgvectorLongTerm::remember`], so a parse failure means
50/// a foreign id and is surfaced as a typed error rather than silently dropped.
51fn parse_fact_id(id: &FactId) -> Result<Uuid, MemoryError> {
52    Uuid::parse_str(&id.to_string())
53        .map_err(|e| MemoryError::Store(format!("fact id is not a uuid: {e}")))
54}
55
56/// `LongTermMemory` + `FilterableLongTermMemory` over the single table
57/// described in the module docs; cosine recall via the HNSW index, scope
58/// enforced by the `scope_kind`/`scope_value` columns.
59pub struct PgvectorLongTerm {
60    handle: PgvectorHandle,
61    embedder: Arc<dyn Embedder>,
62    embedder_id: String,
63}
64
65impl PgvectorLongTerm {
66    pub(crate) fn new(
67        handle: PgvectorHandle,
68        embedder: Arc<dyn Embedder>,
69        embedder_id: String,
70    ) -> Self {
71        Self {
72            handle,
73            embedder,
74            embedder_id,
75        }
76    }
77
78    /// Embed one text and validate it against the embedder's declared
79    /// dimension before it reaches the `vector(dim)` column.
80    async fn embed_one(&self, text: &str) -> Result<Vec<f32>, MemoryError> {
81        let dim = self.embedder.dimension();
82        let vector = self
83            .embedder
84            .embed(&[text.to_string()])
85            .await?
86            .into_iter()
87            .next()
88            .ok_or_else(|| MemoryError::Embedding("embedder returned empty vec".into()))?;
89        if vector.len() != dim {
90            return Err(MemoryError::Embedding(format!(
91                "embedder produced {}-dim vector, expected {}",
92                vector.len(),
93                dim
94            )));
95        }
96        Ok(vector)
97    }
98}
99
100#[async_trait]
101impl LongTermMemory for PgvectorLongTerm {
102    async fn remember(&self, scope: Scope, fact: Fact) -> Result<FactId, MemoryError> {
103        self.handle
104            .ensure_table(self.embedder.dimension() as u64)
105            .await?;
106        let vector = self.embed_one(&fact.text).await?;
107        let id = Uuid::new_v4();
108        let (kind, value) = scope_columns(&scope);
109        let sql = format!(
110            "INSERT INTO {} (fact_id, text, metadata, embedding, scope_kind, scope_value) \
111             VALUES ($1, $2, $3, $4::vector, $5, $6)",
112            self.handle.table
113        );
114        sqlx_core::query::query(&sql)
115            .bind(id)
116            .bind(&fact.text)
117            .bind(&fact.metadata)
118            .bind(vector_literal(&vector))
119            .bind(kind)
120            .bind(value)
121            .execute(&self.handle.pool)
122            .await
123            .map_err(store_err)?;
124        Ok(FactId::new(id.to_string()))
125    }
126
127    async fn recall(&self, scope: Scope, query: &str, k: usize) -> Result<Vec<Fact>, MemoryError> {
128        if k == 0 {
129            return Ok(Vec::new());
130        }
131        let vector = self.embed_one(query).await?;
132        let (kind, value) = scope_columns(&scope);
133        let sql = format!(
134            "SELECT text, metadata FROM {} \
135             WHERE scope_kind = $1 AND scope_value = $2 \
136             ORDER BY embedding <=> $3::vector LIMIT $4",
137            self.handle.table
138        );
139        let rows = sqlx_core::query::query(&sql)
140            .bind(kind)
141            .bind(value)
142            .bind(vector_literal(&vector))
143            .bind(k as i64)
144            .fetch_all(&self.handle.pool)
145            .await;
146        rows_to_facts(rows)
147    }
148
149    /// Vector recall when the injected embedder actually separates inputs;
150    /// `NonRanking` when it does not (see `NonRankingEmbedder`).
151    fn recall_semantics(&self) -> RecallSemantics {
152        klieo_embed_common::vector_recall_semantics(self.embedder.as_ref())
153    }
154
155    async fn forget(&self, id: FactId) -> Result<(), MemoryError> {
156        let uuid = parse_fact_id(&id)?;
157        let sql = format!("DELETE FROM {} WHERE fact_id = $1", self.handle.table);
158        match sqlx_core::query::query(&sql)
159            .bind(uuid)
160            .execute(&self.handle.pool)
161            .await
162        {
163            Ok(_) => Ok(()),
164            Err(e) if is_undefined_table(&e) => Ok(()),
165            Err(e) => Err(store_err(e)),
166        }
167    }
168}
169
170#[async_trait]
171impl FilterableLongTermMemory for PgvectorLongTerm {
172    async fn recall_filtered(
173        &self,
174        scope: Scope,
175        query: &str,
176        k: usize,
177        candidate_ids: &[FactId],
178    ) -> Result<Vec<Fact>, MemoryError> {
179        if k == 0 || candidate_ids.is_empty() {
180            return Ok(Vec::new());
181        }
182        let candidates: Vec<Uuid> = candidate_ids
183            .iter()
184            .map(parse_fact_id)
185            .collect::<Result<_, _>>()?;
186        let vector = self.embed_one(query).await?;
187        let (kind, value) = scope_columns(&scope);
188        let sql = format!(
189            "SELECT text, metadata FROM {} \
190             WHERE scope_kind = $1 AND scope_value = $2 AND fact_id = ANY($3) \
191             ORDER BY embedding <=> $4::vector LIMIT $5",
192            self.handle.table
193        );
194        let rows = sqlx_core::query::query(&sql)
195            .bind(kind)
196            .bind(value)
197            .bind(candidates)
198            .bind(vector_literal(&vector))
199            .bind(k as i64)
200            .fetch_all(&self.handle.pool)
201            .await;
202        rows_to_facts(rows)
203    }
204
205    fn embedder_id(&self) -> &str {
206        &self.embedder_id
207    }
208}
209
210/// Decode a query result into `Vec<Fact>`, mapping `undefined_table` (a read
211/// that raced ahead of the first `remember`) to an empty result.
212fn rows_to_facts(
213    rows: Result<Vec<sqlx_postgres::PgRow>, sqlx_core::error::Error>,
214) -> Result<Vec<Fact>, MemoryError> {
215    let rows = match rows {
216        Ok(rows) => rows,
217        Err(e) if is_undefined_table(&e) => return Ok(Vec::new()),
218        Err(e) => return Err(store_err(e)),
219    };
220    rows.into_iter()
221        .map(|row| {
222            let text: String = row.try_get("text").map_err(store_err)?;
223            let metadata: serde_json::Value = row.try_get("metadata").map_err(store_err)?;
224            Ok(row_to_fact(text, metadata))
225        })
226        .collect()
227}
228
229#[cfg(test)]
230mod tests {
231    use super::*;
232
233    #[allow(dead_code)]
234    fn _trait_object_safe(_: &dyn LongTermMemory) {}
235
236    #[test]
237    fn scope_columns_maps_each_variant() {
238        assert_eq!(
239            scope_columns(&Scope::Workspace("ws".into())),
240            ("workspace", "ws".to_string())
241        );
242        assert_eq!(
243            scope_columns(&Scope::Agent("a".into())),
244            ("agent", "a".to_string())
245        );
246        assert_eq!(scope_columns(&Scope::Global), ("global", String::new()));
247    }
248
249    #[test]
250    fn row_to_fact_carries_text_and_metadata() {
251        let fact = row_to_fact("hi".to_string(), serde_json::json!({"k": 1}));
252        assert_eq!(fact.text, "hi");
253        assert_eq!(fact.metadata, serde_json::json!({"k": 1}));
254    }
255
256    #[test]
257    fn parse_fact_id_rejects_non_uuid() {
258        let err = parse_fact_id(&FactId::new("not-a-uuid".to_string())).unwrap_err();
259        assert!(matches!(err, MemoryError::Store(_)));
260    }
261
262    #[test]
263    fn vector_literal_renders_pgvector_text_form() {
264        assert_eq!(vector_literal(&[]), "[]");
265        assert_eq!(vector_literal(&[1.0]), "[1]");
266        assert_eq!(vector_literal(&[0.5, -1.25, 2.0]), "[0.5,-1.25,2]");
267    }
268
269    #[test]
270    fn parse_fact_id_accepts_minted_uuid() {
271        let id = uuid::Uuid::new_v4().to_string();
272        let parsed = parse_fact_id(&FactId::new(id.clone())).unwrap();
273        assert_eq!(parsed.to_string(), id);
274    }
275}