Skip to main content

ragfs_query/
executor.rs

1//! Query execution.
2
3use ragfs_core::{
4    DirectoryScope, DistanceMetric, Embedder, EmbeddingConfig, SearchQuery, SearchResult,
5    VectorStore,
6};
7use std::sync::Arc;
8use tracing::debug;
9
10use crate::parser::{ParsedQuery, QueryParser};
11
12/// Query executor.
13pub struct QueryExecutor {
14    /// Vector store
15    store: Arc<dyn VectorStore>,
16    /// Embedder for query embedding
17    embedder: Arc<dyn Embedder>,
18    /// Query parser
19    parser: QueryParser,
20    /// Whether to use hybrid search
21    hybrid: bool,
22    /// Upper bound applied to parsed / CLI limits
23    max_limit: usize,
24    /// Optional directory scope, relative to the index root
25    scope_prefix: Option<String>,
26}
27
28impl QueryExecutor {
29    /// Create a new query executor.
30    pub fn new(
31        store: Arc<dyn VectorStore>,
32        embedder: Arc<dyn Embedder>,
33        default_limit: usize,
34        hybrid: bool,
35    ) -> Self {
36        Self {
37            store,
38            embedder,
39            parser: QueryParser::new(default_limit),
40            hybrid,
41            max_limit: usize::MAX,
42            scope_prefix: None,
43        }
44    }
45
46    /// Clamp result counts to `max_limit` from `[query].max_limit`.
47    #[must_use]
48    pub fn with_max_limit(mut self, max_limit: usize) -> Self {
49        self.max_limit = max_limit.max(1);
50        self
51    }
52
53    /// Restrict search to a directory (relative to the index root) and its subdirectories.
54    #[must_use]
55    pub fn with_scope(mut self, scope: Option<String>) -> Self {
56        self.scope_prefix = scope
57            .map(|s| DirectoryScope::normalize_prefix(&s))
58            .filter(|s| !s.is_empty() && s != ".");
59        self
60    }
61
62    /// Whether this executor uses hybrid (vector + FTS) search.
63    #[must_use]
64    pub fn is_hybrid(&self) -> bool {
65        self.hybrid
66    }
67
68    fn clamp_limit(&self, limit: usize) -> usize {
69        limit.min(self.max_limit).max(1)
70    }
71
72    /// Execute a query string.
73    pub async fn execute(&self, query_str: &str) -> Result<Vec<SearchResult>, ragfs_core::Error> {
74        debug!("Executing query: {}", query_str);
75
76        // Parse query
77        let mut parsed = self.parser.parse(query_str);
78        parsed.limit = self.clamp_limit(parsed.limit);
79
80        // Embed query text
81        let config = EmbeddingConfig::default();
82        let embedding = self
83            .embedder
84            .embed_query(&parsed.text, &config)
85            .await
86            .map_err(ragfs_core::Error::Embedding)?;
87
88        // Build search query
89        let search_query = SearchQuery {
90            embedding: embedding.embedding,
91            text: if self.hybrid {
92                Some(parsed.text.clone())
93            } else {
94                None
95            },
96            limit: parsed.limit,
97            filters: parsed.filters,
98            metric: DistanceMetric::Cosine,
99            scope_prefix: self.scope_prefix.clone().or(parsed.scope_prefix),
100        };
101
102        // Execute search
103        let results = if self.hybrid {
104            self.store.hybrid_search(search_query).await
105        } else {
106            self.store.search(search_query).await
107        }
108        .map_err(ragfs_core::Error::Store)?;
109
110        debug!("Found {} results", results.len());
111        Ok(results)
112    }
113
114    /// Execute with a pre-parsed query.
115    pub async fn execute_parsed(
116        &self,
117        parsed: ParsedQuery,
118    ) -> Result<Vec<SearchResult>, ragfs_core::Error> {
119        let mut parsed = parsed;
120        parsed.limit = self.clamp_limit(parsed.limit);
121
122        let config = EmbeddingConfig::default();
123        let embedding = self
124            .embedder
125            .embed_query(&parsed.text, &config)
126            .await
127            .map_err(ragfs_core::Error::Embedding)?;
128
129        let search_query = SearchQuery {
130            embedding: embedding.embedding,
131            text: if self.hybrid {
132                Some(parsed.text.clone())
133            } else {
134                None
135            },
136            limit: parsed.limit,
137            filters: parsed.filters,
138            metric: DistanceMetric::Cosine,
139            scope_prefix: self.scope_prefix.clone().or(parsed.scope_prefix),
140        };
141
142        let results = if self.hybrid {
143            self.store.hybrid_search(search_query).await
144        } else {
145            self.store.search(search_query).await
146        }
147        .map_err(ragfs_core::Error::Store)?;
148
149        Ok(results)
150    }
151}
152
153#[cfg(test)]
154mod tests {
155    use super::*;
156    use async_trait::async_trait;
157    use ragfs_core::{
158        Chunk, EmbedError, EmbeddingOutput, FileRecord, Modality, StoreError, StoreStats,
159    };
160    use std::collections::HashMap;
161    use std::path::{Path, PathBuf};
162    use tokio::sync::RwLock;
163    use uuid::Uuid;
164
165    const TEST_DIM: usize = 384;
166
167    // ==================== Mock Embedder ====================
168
169    struct MockEmbedder {
170        dimension: usize,
171    }
172
173    impl MockEmbedder {
174        fn new(dimension: usize) -> Self {
175            Self { dimension }
176        }
177    }
178
179    #[async_trait]
180    impl Embedder for MockEmbedder {
181        fn model_name(&self) -> &'static str {
182            "mock-embedder"
183        }
184
185        fn dimension(&self) -> usize {
186            self.dimension
187        }
188
189        fn max_tokens(&self) -> usize {
190            512
191        }
192
193        fn modalities(&self) -> &[Modality] {
194            &[Modality::Text]
195        }
196
197        async fn embed_text(
198            &self,
199            texts: &[&str],
200            _config: &EmbeddingConfig,
201        ) -> Result<Vec<EmbeddingOutput>, EmbedError> {
202            Ok(texts
203                .iter()
204                .map(|_| EmbeddingOutput {
205                    embedding: vec![0.1; self.dimension],
206                    token_count: 10,
207                })
208                .collect())
209        }
210
211        async fn embed_query(
212            &self,
213            _query: &str,
214            _config: &EmbeddingConfig,
215        ) -> Result<EmbeddingOutput, EmbedError> {
216            Ok(EmbeddingOutput {
217                embedding: vec![0.1; self.dimension],
218                token_count: 10,
219            })
220        }
221    }
222
223    // ==================== Mock VectorStore ====================
224
225    struct MockStore {
226        results: Arc<RwLock<Vec<SearchResult>>>,
227        hybrid_results: Arc<RwLock<Vec<SearchResult>>>,
228        last_query: Arc<RwLock<Option<SearchQuery>>>,
229    }
230
231    impl MockStore {
232        fn new() -> Self {
233            Self {
234                results: Arc::new(RwLock::new(Vec::new())),
235                hybrid_results: Arc::new(RwLock::new(Vec::new())),
236                last_query: Arc::new(RwLock::new(None)),
237            }
238        }
239
240        fn with_results(results: Vec<SearchResult>) -> Self {
241            Self {
242                results: Arc::new(RwLock::new(results)),
243                hybrid_results: Arc::new(RwLock::new(Vec::new())),
244                last_query: Arc::new(RwLock::new(None)),
245            }
246        }
247
248        fn with_hybrid_results(results: Vec<SearchResult>, hybrid: Vec<SearchResult>) -> Self {
249            Self {
250                results: Arc::new(RwLock::new(results)),
251                hybrid_results: Arc::new(RwLock::new(hybrid)),
252                last_query: Arc::new(RwLock::new(None)),
253            }
254        }
255    }
256
257    #[async_trait]
258    impl VectorStore for MockStore {
259        async fn init(&self) -> Result<(), StoreError> {
260            Ok(())
261        }
262
263        async fn upsert_chunks(&self, _chunks: &[Chunk]) -> Result<(), StoreError> {
264            Ok(())
265        }
266
267        async fn search(&self, query: SearchQuery) -> Result<Vec<SearchResult>, StoreError> {
268            *self.last_query.write().await = Some(query);
269            let results = self.results.read().await;
270            Ok(results.clone())
271        }
272
273        async fn hybrid_search(&self, query: SearchQuery) -> Result<Vec<SearchResult>, StoreError> {
274            *self.last_query.write().await = Some(query);
275            let results = self.hybrid_results.read().await;
276            Ok(results.clone())
277        }
278
279        async fn delete_by_file_path(&self, _path: &Path) -> Result<u64, StoreError> {
280            Ok(0)
281        }
282
283        async fn get_file(&self, _path: &Path) -> Result<Option<FileRecord>, StoreError> {
284            Ok(None)
285        }
286
287        async fn upsert_file(&self, _record: &FileRecord) -> Result<(), StoreError> {
288            Ok(())
289        }
290
291        async fn stats(&self) -> Result<StoreStats, StoreError> {
292            Ok(StoreStats {
293                total_chunks: 0,
294                total_files: 0,
295                index_size_bytes: 0,
296                last_updated: None,
297            })
298        }
299
300        async fn update_file_path(&self, _from: &Path, _to: &Path) -> Result<u64, StoreError> {
301            Ok(0)
302        }
303
304        async fn get_chunks_for_file(&self, _path: &Path) -> Result<Vec<Chunk>, StoreError> {
305            Ok(vec![])
306        }
307
308        async fn get_all_chunks(&self) -> Result<Vec<Chunk>, StoreError> {
309            Ok(vec![])
310        }
311
312        async fn get_all_files(&self) -> Result<Vec<FileRecord>, StoreError> {
313            Ok(vec![])
314        }
315    }
316
317    // ==================== Helper functions ====================
318
319    fn create_test_result(path: &str, content: &str, score: f32) -> SearchResult {
320        SearchResult {
321            chunk_id: Uuid::new_v4(),
322            file_path: PathBuf::from(path),
323            content: content.to_string(),
324            score,
325            byte_range: 0..content.len() as u64,
326            line_range: Some(0..1),
327            metadata: HashMap::new(),
328        }
329    }
330
331    // ==================== Tests ====================
332
333    #[tokio::test]
334    async fn test_execute_simple_query() {
335        let results = vec![
336            create_test_result("/test/file1.txt", "Authentication module", 0.9),
337            create_test_result("/test/file2.txt", "Auth config", 0.8),
338        ];
339
340        let store = Arc::new(MockStore::with_results(results.clone()));
341        let embedder = Arc::new(MockEmbedder::new(TEST_DIM));
342
343        let executor = QueryExecutor::new(store, embedder, 10, false);
344
345        let query_results = executor.execute("authentication").await.unwrap();
346
347        assert_eq!(query_results.len(), 2);
348        assert_eq!(query_results[0].content, "Authentication module");
349        assert_eq!(query_results[1].content, "Auth config");
350    }
351
352    #[tokio::test]
353    async fn test_execute_with_hybrid_search() {
354        let vector_results = vec![create_test_result("/test/vector.txt", "Vector result", 0.8)];
355        let hybrid_results = vec![
356            create_test_result("/test/hybrid1.txt", "Hybrid result 1", 0.95),
357            create_test_result("/test/hybrid2.txt", "Hybrid result 2", 0.85),
358        ];
359
360        let store = Arc::new(MockStore::with_hybrid_results(
361            vector_results,
362            hybrid_results.clone(),
363        ));
364        let embedder = Arc::new(MockEmbedder::new(TEST_DIM));
365
366        // hybrid=true should use hybrid_search
367        let executor = QueryExecutor::new(store, embedder, 10, true);
368
369        let query_results = executor.execute("search query").await.unwrap();
370
371        assert_eq!(query_results.len(), 2);
372        assert_eq!(query_results[0].content, "Hybrid result 1");
373    }
374
375    #[tokio::test]
376    async fn test_execute_vector_only() {
377        let vector_results = vec![create_test_result(
378            "/test/vector.txt",
379            "Vector only result",
380            0.9,
381        )];
382        let hybrid_results = vec![create_test_result(
383            "/test/hybrid.txt",
384            "Hybrid result",
385            0.95,
386        )];
387
388        let store = Arc::new(MockStore::with_hybrid_results(
389            vector_results.clone(),
390            hybrid_results,
391        ));
392        let embedder = Arc::new(MockEmbedder::new(TEST_DIM));
393
394        // hybrid=false should use regular search
395        let executor = QueryExecutor::new(store, embedder, 10, false);
396
397        let query_results = executor.execute("search query").await.unwrap();
398
399        assert_eq!(query_results.len(), 1);
400        assert_eq!(query_results[0].content, "Vector only result");
401    }
402
403    #[tokio::test]
404    async fn test_execute_empty_results() {
405        let store = Arc::new(MockStore::new());
406        let embedder = Arc::new(MockEmbedder::new(TEST_DIM));
407
408        let executor = QueryExecutor::new(store, embedder, 10, false);
409
410        let query_results = executor.execute("no results query").await.unwrap();
411
412        assert!(query_results.is_empty());
413    }
414
415    #[tokio::test]
416    async fn test_execute_with_limit_in_query() {
417        let results = vec![
418            create_test_result("/test/file1.txt", "Result 1", 0.9),
419            create_test_result("/test/file2.txt", "Result 2", 0.8),
420            create_test_result("/test/file3.txt", "Result 3", 0.7),
421        ];
422
423        let store = Arc::new(MockStore::with_results(results));
424        let embedder = Arc::new(MockEmbedder::new(TEST_DIM));
425
426        let executor = QueryExecutor::new(store, embedder, 10, false);
427
428        // Query with limit:2 filter
429        let query_results = executor.execute("search query limit:2").await.unwrap();
430
431        // Note: The mock store returns all results regardless of limit.
432        // In a real test with actual store, we'd verify the limit is applied.
433        assert!(!query_results.is_empty());
434    }
435
436    #[tokio::test]
437    async fn test_execute_parsed_query() {
438        use crate::parser::ParsedQuery;
439
440        let results = vec![create_test_result("/test/file.txt", "Parsed result", 0.9)];
441
442        let store = Arc::new(MockStore::with_results(results));
443        let embedder = Arc::new(MockEmbedder::new(TEST_DIM));
444
445        let executor = QueryExecutor::new(store, embedder, 10, false);
446
447        let parsed = ParsedQuery {
448            text: "pre-parsed query".to_string(),
449            limit: 5,
450            filters: vec![],
451            scope_prefix: None,
452        };
453
454        let query_results = executor.execute_parsed(parsed).await.unwrap();
455
456        assert_eq!(query_results.len(), 1);
457        assert_eq!(query_results[0].content, "Parsed result");
458    }
459
460    #[test]
461    fn test_query_executor_creation() {
462        let store: Arc<dyn VectorStore> = Arc::new(MockStore::new());
463        let embedder: Arc<dyn Embedder> = Arc::new(MockEmbedder::new(TEST_DIM));
464
465        // Test with hybrid=false
466        let executor = QueryExecutor::new(Arc::clone(&store), Arc::clone(&embedder), 10, false);
467        assert!(!executor.is_hybrid());
468
469        // Test with hybrid=true
470        let executor2 = QueryExecutor::new(store, embedder, 20, true);
471        assert!(executor2.is_hybrid());
472    }
473
474    #[test]
475    fn test_max_limit_from_config_is_stored() {
476        let store: Arc<dyn VectorStore> = Arc::new(MockStore::new());
477        let embedder: Arc<dyn Embedder> = Arc::new(MockEmbedder::new(TEST_DIM));
478        let executor = QueryExecutor::new(store, embedder, 10, true).with_max_limit(7);
479        assert_eq!(executor.clamp_limit(100), 7);
480        assert!(executor.is_hybrid());
481    }
482
483    #[tokio::test]
484    async fn test_execute_forwards_scope_prefix() {
485        let store = Arc::new(MockStore::with_results(vec![create_test_result(
486            "/project/src/auth/login.rs",
487            "login",
488            0.9,
489        )]));
490        let last_query = store.last_query.clone();
491        let embedder = Arc::new(MockEmbedder::new(TEST_DIM));
492        let executor = QueryExecutor::new(store, embedder, 10, false)
493            .with_scope(Some("src/auth/".to_string()));
494
495        executor.execute("authentication").await.unwrap();
496
497        let query = last_query.read().await.clone().expect("search was called");
498        assert_eq!(query.scope_prefix.as_deref(), Some("src/auth"));
499    }
500}