1use 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
32fn 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
48fn 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
56pub 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 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 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
210fn 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}