Skip to main content

lc_rag/
retriever.rs

1// lc-rag/src/retriever.rs
2//! Retriever implementations
3//!
4//! Provides similarity-based document retrieval.
5
6use async_trait::async_trait;
7use lc_embeddings::Embeddings;
8use lc_vector_stores::{Document, SearchResult, VectorStore, VectorStoreError};
9use std::sync::Arc;
10
11/// Retriever error type
12#[derive(Debug)]
13#[non_exhaustive]
14pub enum RetrieverError {
15    /// Vector store error
16    StoreError(VectorStoreError),
17
18    /// Embedding error
19    EmbeddingError(String),
20
21    /// LLM breakdown failure (call failed / output unparseable). Used by SelfQuery (S4).
22    LlmError(String),
23
24    /// The filter references a field outside the `allowed_attributes` whitelist (SelfQuery).
25    /// Errors out explicitly, never silently falls back to unfiltered retrieval — that would
26    /// return data that should have been filtered out (data-plane over-exposure).
27    InvalidFilter(String),
28
29    /// No results
30    NoResults,
31
32    /// A document handed to a retriever is malformed for its modality
33    /// (e.g. an image document missing its reference). B7 multimodal retrieval.
34    InvalidDocument(String),
35}
36
37impl std::fmt::Display for RetrieverError {
38    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
39        match self {
40            RetrieverError::StoreError(e) => write!(f, "storage error: {}", e),
41            RetrieverError::EmbeddingError(msg) => write!(f, "embedding error: {}", msg),
42            RetrieverError::LlmError(msg) => write!(f, "LLM error: {}", msg),
43            RetrieverError::InvalidFilter(msg) => write!(f, "invalid filter: {}", msg),
44            RetrieverError::NoResults => write!(f, "no relevant documents found"),
45            RetrieverError::InvalidDocument(msg) => write!(f, "invalid document: {msg}"),
46        }
47    }
48}
49
50impl std::error::Error for RetrieverError {}
51
52impl From<VectorStoreError> for RetrieverError {
53    fn from(e: VectorStoreError) -> Self {
54        RetrieverError::StoreError(e)
55    }
56}
57
58/// Retriever trait
59#[async_trait]
60pub trait RetrieverTrait: Send + Sync {
61    /// Retrieves relevant documents
62    ///
63    /// # Arguments
64    /// * `query` - the query text
65    /// * `k` - the number of documents to return
66    ///
67    /// # Returns
68    /// The list of relevant documents
69    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError>;
70
71    /// Retrieves relevant documents (with scores)
72    async fn retrieve_with_scores(
73        &self,
74        query: &str,
75        k: usize,
76    ) -> Result<Vec<SearchResult>, RetrieverError>;
77
78    /// Adds documents
79    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError>;
80}
81
82/// Similarity-based retriever
83pub struct SimilarityRetriever {
84    /// Vector store
85    store: Arc<dyn VectorStore>,
86
87    /// Embedding model
88    embeddings: Arc<dyn Embeddings>,
89}
90
91impl SimilarityRetriever {
92    /// Creates a new similarity retriever
93    pub fn new(store: Arc<dyn VectorStore>, embeddings: Arc<dyn Embeddings>) -> Self {
94        Self { store, embeddings }
95    }
96}
97
98#[async_trait]
99impl RetrieverTrait for SimilarityRetriever {
100    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError> {
101        let results = self.retrieve_with_scores(query, k).await?;
102        Ok(results.into_iter().map(|r| r.document).collect())
103    }
104
105    async fn retrieve_with_scores(
106        &self,
107        query: &str,
108        k: usize,
109    ) -> Result<Vec<SearchResult>, RetrieverError> {
110        // Generate the query vector
111        let query_embedding = self
112            .embeddings
113            .embed_query(query)
114            .await
115            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
116
117        // Retrieve similar documents
118        let results = self.store.similarity_search(&query_embedding, k).await?;
119
120        Ok(results)
121    }
122
123    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
124        // Generate document embeddings
125        let texts: Vec<&str> = documents.iter().map(|d| d.content.as_str()).collect();
126        let embeddings = self
127            .embeddings
128            .embed_documents(&texts)
129            .await
130            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
131
132        // Add to storage
133        self.store.add_documents(documents, embeddings).await?;
134
135        Ok(())
136    }
137}
138
139/// A simplified Retriever type alias (for quick use)
140pub type Retriever = SimilarityRetriever;
141
142#[cfg(test)]
143mod tests {
144    use super::*;
145    use crate::bm25::BM25Retriever;
146    use crate::unified_hybrid::UnifiedHybridIndex;
147    use lc_embeddings::MockEmbeddings;
148    use lc_vector_stores::InMemoryVectorStore;
149
150    /// P0-1: Verifies both BM25 / UnifiedHybrid work as
151    /// `Arc<dyn RetrieverTrait>`, completing the full add + retrieve flow.
152    #[tokio::test]
153    async fn test_retriever_trait_object_hybrid_retrievers() {
154        // BM25Retriever as a trait object
155        let bm25: Arc<dyn RetrieverTrait> = Arc::new(BM25Retriever::new());
156        bm25.add_documents(vec![Document::new(
157            "Rust is a systems programming language",
158        )])
159        .await
160        .unwrap();
161        let results = bm25.retrieve("systems", 1).await.unwrap();
162        assert!(!results.is_empty());
163
164        // UnifiedHybridIndex as a trait object
165        let embeddings = Arc::new(MockEmbeddings::new(128));
166        let vector_store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
167        let unified: Arc<dyn RetrieverTrait> = Arc::new(UnifiedHybridIndex::new(
168            embeddings.clone(),
169            vector_store,
170            128,
171        ));
172        unified
173            .add_documents(vec![Document::new(
174                "Rust is a systems programming language",
175            )])
176            .await
177            .unwrap();
178        let results = unified.retrieve("systems", 1).await.unwrap();
179        assert!(!results.is_empty());
180    }
181
182    #[tokio::test]
183    async fn test_retriever() {
184        let store = Arc::new(InMemoryVectorStore::new());
185        let embeddings = Arc::new(MockEmbeddings::new(128));
186
187        let retriever = SimilarityRetriever::new(store.clone(), embeddings.clone());
188
189        // Add documents
190        let docs = vec![
191            Document::new("Rust is a systems programming language"),
192            Document::new("Python is a scripting language"),
193            Document::new("JavaScript is used for web development"),
194        ];
195
196        retriever.add_documents(docs).await.unwrap();
197        assert_eq!(store.count().await, 3);
198
199        // Retrieve documents
200        let results = retriever.retrieve("programming language", 2).await.unwrap();
201        assert!(
202            !results.is_empty(),
203            "expected at least 1 result, got {}",
204            results.len()
205        );
206    }
207
208    #[tokio::test]
209    async fn test_retriever_with_scores() {
210        let store = Arc::new(InMemoryVectorStore::new());
211        let embeddings = Arc::new(MockEmbeddings::new(64));
212
213        let retriever = SimilarityRetriever::new(store, embeddings);
214
215        let docs = vec![Document::new("Document A"), Document::new("Document B")];
216
217        retriever.add_documents(docs).await.unwrap();
218
219        let results = retriever.retrieve_with_scores("query", 2).await.unwrap();
220        assert_eq!(results.len(), 2);
221
222        // Results should include scores
223        assert!(results[0].score >= -1.0 && results[0].score <= 1.0);
224    }
225}