use std::path::Path;
use std::sync::atomic::{AtomicUsize, Ordering};
use mnemo_core::error::{Error, Result};
use mnemo_core::index::VectorIndex;
use pgvector::Vector;
use sqlx::Row;
use uuid::Uuid;
pub struct PgVectorIndex {
pool: Option<sqlx::PgPool>,
dimensions: usize,
count: AtomicUsize,
}
impl PgVectorIndex {
pub fn new() -> Self {
Self {
pool: None,
dimensions: 0,
count: AtomicUsize::new(0),
}
}
pub fn with_pool(pool: sqlx::PgPool, dimensions: usize) -> Self {
Self {
pool: Some(pool),
dimensions,
count: AtomicUsize::new(0),
}
}
async fn ann_query(
pool: &sqlx::PgPool,
query: &Vector,
limit: usize,
) -> Result<Vec<(Uuid, f32)>> {
let rows = sqlx::query(
"SELECT id, (embedding <=> $1) AS dist \
FROM memories \
WHERE embedding IS NOT NULL AND deleted_at IS NULL \
ORDER BY embedding <=> $1 \
LIMIT $2",
)
.bind(query)
.bind(limit as i64)
.fetch_all(pool)
.await
.map_err(map_ann_error)?;
let mut out = Vec::with_capacity(rows.len());
for row in &rows {
let id: Uuid = row.try_get("id").map_err(|e| Error::Index(e.to_string()))?;
let dist: f64 = row
.try_get("dist")
.map_err(|e| Error::Index(e.to_string()))?;
out.push((id, dist as f32));
}
Ok(out)
}
fn pool_for(&self, query: &[f32]) -> Result<&sqlx::PgPool> {
let pool = self.pool.as_ref().ok_or_else(ann_unsupported)?;
if self.dimensions != 0 && query.len() != self.dimensions {
return Err(Error::Index(format!(
"query embedding has {} dims but the pgvector column is {} — \
re-embed with the configured model",
query.len(),
self.dimensions
)));
}
Ok(pool)
}
}
impl Default for PgVectorIndex {
fn default() -> Self {
Self::new()
}
}
fn ann_unsupported() -> Error {
Error::BackendUnsupported {
backend: "postgres".to_string(),
capability: "semantic_recall".to_string(),
detail: "pgvector ANN search is unavailable: the index has no database \
pool, or the pgvector extension / `<=>` operator is not \
installed. Ensure the `vector` extension and the \
`idx_memories_embedding_hnsw` index exist (created by \
migrations), or use strategy=\"lexical\"/\"exact\". \
Tracking: https://github.com/sattyamjjain/mnemo/issues/99"
.to_string(),
}
}
fn map_ann_error(e: sqlx::Error) -> Error {
let msg = e.to_string();
let lower = msg.to_lowercase();
let capability_absent = (lower.contains("operator does not exist") && lower.contains("<=>"))
|| lower.contains("type \"vector\" does not exist")
|| lower.contains("extension \"vector\"");
if capability_absent {
ann_unsupported()
} else {
Error::Index(format!("pgvector ANN query failed: {msg}"))
}
}
#[async_trait::async_trait]
impl VectorIndex for PgVectorIndex {
fn add(&self, _id: Uuid, _vector: &[f32]) -> Result<()> {
self.count.fetch_add(1, Ordering::Relaxed);
Ok(())
}
fn remove(&self, _id: Uuid) -> Result<()> {
let _ = self
.count
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |n| {
Some(n.saturating_sub(1))
});
Ok(())
}
async fn search(&self, query: &[f32], limit: usize) -> Result<Vec<(Uuid, f32)>> {
let pool = self.pool_for(query)?;
let vec = Vector::from(query.to_vec());
Self::ann_query(pool, &vec, limit).await
}
async fn filtered_search(
&self,
query: &[f32],
limit: usize,
filter: &(dyn Fn(Uuid) -> bool + Send + Sync),
) -> Result<Vec<(Uuid, f32)>> {
let pool = self.pool_for(query)?;
let vec = Vector::from(query.to_vec());
if limit == 0 {
return Ok(Vec::new());
}
let mut oversample = limit.saturating_mul(3).max(1);
loop {
let candidates = Self::ann_query(pool, &vec, oversample).await?;
let exhausted = candidates.len() < oversample;
let filtered: Vec<(Uuid, f32)> = candidates
.into_iter()
.filter(|(id, _)| filter(*id))
.take(limit)
.collect();
if filtered.len() >= limit || exhausted {
return Ok(filtered);
}
oversample = oversample.saturating_mul(2);
}
}
fn save(&self, _path: &Path) -> Result<()> {
Ok(())
}
fn load(&self, _path: &Path) -> Result<()> {
Ok(())
}
fn len(&self) -> usize {
self.count.load(Ordering::Relaxed)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn ann_search_fails_loud_not_silent_empty() {
let idx = PgVectorIndex::new();
idx.add(Uuid::nil(), &[0.1, 0.2, 0.3]).unwrap();
assert_eq!(idx.len(), 1);
idx.remove(Uuid::nil()).unwrap();
assert_eq!(idx.len(), 0);
assert!(
idx.search(&[0.1, 0.2, 0.3], 5).await.is_err(),
"search must fail loud, not return Ok(empty)"
);
assert!(
idx.filtered_search(&[0.1, 0.2, 0.3], 5, &|_| true)
.await
.is_err(),
"filtered_search must fail loud, not return Ok(empty)"
);
match idx.search(&[0.0], 1).await.unwrap_err() {
Error::BackendUnsupported {
backend,
capability,
detail,
} => {
assert_eq!(backend, "postgres");
assert_eq!(capability, "semantic_recall");
assert!(
detail.contains("issues/99"),
"detail should reference the tracking issue: {detail}"
);
}
other => panic!("expected BackendUnsupported, got: {other}"),
}
}
#[tokio::test]
async fn dimension_mismatch_is_loud() {
let idx = PgVectorIndex::new();
assert!(idx.search(&[0.1; 4], 3).await.is_err());
}
}