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_core::row::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)
}
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
}
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}")))
}
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,
}
}
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_core::query::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_core::query::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_core::query::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_core::query::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
}
}
fn rows_to_facts(
rows: Result<Vec<sqlx_postgres::PgRow>, sqlx_core::error::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);
}
}