use std::sync::Arc;
use async_trait::async_trait;
use mnemo_core::embedding::EmbeddingProvider;
use mnemo_core::error::Result as MnResult;
use mnemo_core::query::MnemoEngine;
use mnemo_core::query::recall::RecallRequest;
use mnemo_core::query::remember::RememberRequest;
use mnemo_postgres::{PgStorage, PgVectorIndex};
const DIM: usize = 4;
const AGENT_A: &str = "pgann-A";
const AGENT_B: &str = "pgann-B";
static PG_CONNECT_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
async fn connect_storage(url: &str) -> std::sync::Arc<PgStorage> {
let _guard = PG_CONNECT_LOCK.lock().await;
std::sync::Arc::new(
PgStorage::connect(url, DIM)
.await
.expect("connect + run migrations"),
)
}
fn vec_for(text: &str) -> Vec<f32> {
match text {
"alpha" => vec![1.0, 0.0, 0.0, 0.0],
"beta" => vec![0.0, 1.0, 0.0, 0.0],
"gamma" => vec![0.0, 0.0, 1.0, 0.0],
"secret" => vec![0.9, 0.4, 0.1, 0.0],
"query" => vec![0.8, 0.5, 0.2, 0.0],
_ => vec![0.0, 0.0, 0.0, 0.0],
}
}
struct MapEmbedding;
#[async_trait]
impl EmbeddingProvider for MapEmbedding {
async fn embed(&self, text: &str) -> MnResult<Vec<f32>> {
Ok(vec_for(text))
}
async fn embed_batch(&self, texts: &[&str]) -> MnResult<Vec<Vec<f32>>> {
Ok(texts.iter().map(|t| vec_for(t)).collect())
}
fn dimensions(&self) -> usize {
DIM
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn pgvector_ann_semantic_auto_and_permission_filter() {
let Ok(url) = std::env::var("MNEMO_TEST_POSTGRES_URL") else {
eprintln!(
"skipping pgvector ANN test: set MNEMO_TEST_POSTGRES_URL=postgres://... \
(needs the pgvector extension) to run it"
);
return;
};
let storage = connect_storage(&url).await;
let _ = sqlx::query("DELETE FROM memories WHERE agent_id = ANY($1)")
.bind(vec![AGENT_A.to_string(), AGENT_B.to_string()])
.execute(&storage.pool())
.await;
let index = Arc::new(PgVectorIndex::with_pool(storage.pool(), DIM));
let engine = Arc::new(MnemoEngine::new(
storage.clone(),
index,
Arc::new(MapEmbedding),
AGENT_A.to_string(),
None,
));
for word in ["alpha", "beta", "gamma"] {
engine
.remember(RememberRequest::new(word.to_string()))
.await
.expect("remember");
}
let mut secret = RememberRequest::new("secret".to_string());
secret.agent_id = Some(AGENT_B.to_string());
let secret_id = engine.remember(secret).await.expect("remember secret").id;
let mut req = RecallRequest::new("query".to_string());
req.strategy = Some("semantic".to_string());
req.limit = Some(5);
let resp = engine.recall(req).await.expect("semantic recall");
let contents: Vec<String> = resp.memories.iter().map(|m| m.content.clone()).collect();
assert!(
!contents.iter().any(|c| c == "secret"),
"AGENT_B's private record must be filtered from AGENT_A's recall, got {contents:?}"
);
assert_eq!(
contents,
vec!["alpha", "beta", "gamma"],
"semantic recall must return the nearest in rank order"
);
let mut areq = RecallRequest::new("query".to_string());
areq.strategy = Some("auto".to_string());
areq.limit = Some(3);
let aresp = engine.recall(areq).await.expect("auto recall");
assert_eq!(
aresp.memories.first().map(|m| m.content.as_str()),
Some("alpha"),
"auto recall must rank the nearest first"
);
let mut breq = RecallRequest::new("query".to_string());
breq.strategy = Some("semantic".to_string());
breq.agent_id = Some(AGENT_B.to_string());
breq.limit = Some(5);
let bresp = engine.recall(breq).await.expect("agent B recall");
assert!(
bresp.memories.iter().any(|m| m.id == secret_id),
"AGENT_B must see its own private record"
);
}
#[tokio::test(flavor = "current_thread")]
async fn semantic_recall_on_current_thread_runtime() {
let Ok(url) = std::env::var("MNEMO_TEST_POSTGRES_URL") else {
eprintln!(
"skipping current-thread pgvector recall test: set MNEMO_TEST_POSTGRES_URL=postgres://..."
);
return;
};
const AGENT_C: &str = "pgann-C-ct";
let storage = connect_storage(&url).await;
let _ = sqlx::query("DELETE FROM memories WHERE agent_id = $1")
.bind(AGENT_C.to_string())
.execute(&storage.pool())
.await;
let index = Arc::new(PgVectorIndex::with_pool(storage.pool(), DIM));
let engine = Arc::new(MnemoEngine::new(
storage,
index,
Arc::new(MapEmbedding),
AGENT_C.to_string(),
None,
));
for word in ["alpha", "beta", "gamma"] {
engine
.remember(RememberRequest::new(word.to_string()))
.await
.expect("remember");
}
let mut req = RecallRequest::new("query".to_string());
req.strategy = Some("semantic".to_string());
req.limit = Some(3);
let resp = engine
.recall(req)
.await
.expect("semantic recall must not panic");
assert!(
!resp.memories.is_empty(),
"semantic recall on a current_thread runtime returned no hits"
);
assert_eq!(
resp.memories.first().map(|m| m.content.as_str()),
Some("alpha"),
"nearest record must rank first"
);
}