use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use synaptic_core::{Document, Embeddings, Retriever, SynapticError};
use tokio::sync::RwLock;
use crate::VectorStore;
pub struct MultiVectorRetriever<S: VectorStore> {
vectorstore: Arc<S>,
embeddings: Arc<dyn Embeddings>,
docstore: Arc<RwLock<HashMap<String, Document>>>,
id_key: String,
k: usize,
}
impl<S: VectorStore + 'static> MultiVectorRetriever<S> {
pub fn new(vectorstore: Arc<S>, embeddings: Arc<dyn Embeddings>, k: usize) -> Self {
Self {
vectorstore,
embeddings,
docstore: Arc::new(RwLock::new(HashMap::new())),
id_key: "parent_id".to_string(),
k,
}
}
pub fn with_id_key(mut self, key: impl Into<String>) -> Self {
self.id_key = key.into();
self
}
pub async fn add_documents(
&self,
parent_docs: Vec<Document>,
child_docs: Vec<Document>,
) -> Result<(), SynapticError> {
{
let mut store = self.docstore.write().await;
for doc in parent_docs {
store.insert(doc.id.clone(), doc);
}
}
self.vectorstore
.add_documents(child_docs, self.embeddings.as_ref())
.await?;
Ok(())
}
}
#[async_trait]
impl<S: VectorStore + 'static> Retriever for MultiVectorRetriever<S> {
async fn retrieve(&self, query: &str, top_k: usize) -> Result<Vec<Document>, SynapticError> {
let k = if top_k > 0 { top_k } else { self.k };
let children = self
.vectorstore
.similarity_search(query, k, self.embeddings.as_ref())
.await?;
let docstore = self.docstore.read().await;
let mut seen = std::collections::HashSet::new();
let mut parents = Vec::new();
for child in &children {
if let Some(parent_id_value) = child.metadata.get(&self.id_key) {
if let Some(parent_id) = parent_id_value.as_str() {
if seen.insert(parent_id.to_string()) {
if let Some(parent) = docstore.get(parent_id) {
parents.push(parent.clone());
}
}
}
}
}
Ok(parents)
}
}