1use 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
12pub struct QueryExecutor {
14 store: Arc<dyn VectorStore>,
16 embedder: Arc<dyn Embedder>,
18 parser: QueryParser,
20 hybrid: bool,
22 max_limit: usize,
24 scope_prefix: Option<String>,
26}
27
28impl QueryExecutor {
29 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 #[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 #[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 #[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 pub async fn execute(&self, query_str: &str) -> Result<Vec<SearchResult>, ragfs_core::Error> {
74 debug!("Executing query: {}", query_str);
75
76 let mut parsed = self.parser.parse(query_str);
78 parsed.limit = self.clamp_limit(parsed.limit);
79
80 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 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 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 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 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 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 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 #[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 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 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 let query_results = executor.execute("search query limit:2").await.unwrap();
430
431 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 let executor = QueryExecutor::new(Arc::clone(&store), Arc::clone(&embedder), 10, false);
467 assert!(!executor.is_hybrid());
468
469 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}