use laurus::Result;
use laurus::vector::Vector;
use laurus::vector::search::searcher::{VectorIndexQueryResult, VectorIndexQueryResults};
use laurus::vector::{VectorIndexQuery, VectorIndexSearcher};
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug, Default)]
struct CountingSearcher {
call_count: AtomicUsize,
}
impl VectorIndexSearcher for CountingSearcher {
fn search(&self, request: &VectorIndexQuery) -> Result<VectorIndexQueryResults> {
let idx = self.call_count.fetch_add(1, Ordering::Relaxed);
Ok(VectorIndexQueryResults {
results: vec![VectorIndexQueryResult {
doc_id: idx as u64,
field_name: request
.field_name
.clone()
.unwrap_or_else(|| "test".to_string()),
similarity: 1.0 - (idx as f32) * 0.01,
distance: (idx as f32) * 0.01,
vector: None,
}],
candidates_examined: 1,
search_time_ms: 0.0,
query_metadata: std::collections::HashMap::new(),
})
}
fn count(&self, _request: VectorIndexQuery) -> Result<u64> {
Ok(0)
}
}
fn dummy_query(seed: u64) -> VectorIndexQuery {
VectorIndexQuery::new(Vector::new(vec![seed as f32, 0.0, 0.0, 0.0])).top_k(1)
}
#[test]
fn test_search_batch_default_impl_preserves_order() {
let searcher = CountingSearcher::default();
let queries: Vec<VectorIndexQuery> = (0..5).map(dummy_query).collect();
let results = searcher
.search_batch_with_threshold(&queries, usize::MAX)
.expect("search_batch_with_threshold");
assert_eq!(results.len(), queries.len());
assert_eq!(searcher.call_count.load(Ordering::Relaxed), 5);
for (i, r) in results.iter().enumerate() {
assert_eq!(r.results.len(), 1, "each query should produce one hit");
assert_eq!(
r.results[0].doc_id, i as u64,
"serial dispatch order should match input order"
);
}
}
#[test]
fn test_search_batch_threshold_zero_runs_parallel_but_results_match() {
let queries: Vec<VectorIndexQuery> = (0..8).map(dummy_query).collect();
let serial_searcher = CountingSearcher::default();
let serial = serial_searcher
.search_batch_with_threshold(&queries, usize::MAX)
.expect("serial");
let parallel_searcher = CountingSearcher::default();
let parallel = parallel_searcher
.search_batch_with_threshold(&queries, 0)
.expect("parallel");
assert_eq!(serial_searcher.call_count.load(Ordering::Relaxed), 8);
assert_eq!(parallel_searcher.call_count.load(Ordering::Relaxed), 8);
assert_eq!(serial.len(), parallel.len());
let mut serial_ids: Vec<u64> = serial.iter().map(|r| r.results[0].doc_id).collect();
let mut parallel_ids: Vec<u64> = parallel.iter().map(|r| r.results[0].doc_id).collect();
serial_ids.sort_unstable();
parallel_ids.sort_unstable();
assert_eq!(
serial_ids, parallel_ids,
"parallel and serial must produce the same doc_id set"
);
}
#[test]
fn test_search_batch_empty_queries() {
let searcher = CountingSearcher::default();
let queries: Vec<VectorIndexQuery> = Vec::new();
let results = searcher.search_batch(&queries).expect("empty batch");
assert!(results.is_empty());
assert_eq!(
searcher.call_count.load(Ordering::Relaxed),
0,
"empty input must not invoke search"
);
let results = searcher
.search_batch_with_threshold(&queries, 0)
.expect("empty batch with threshold=0");
assert!(results.is_empty());
assert_eq!(searcher.call_count.load(Ordering::Relaxed), 0);
}
#[test]
fn test_search_batch_default_threshold_value() {
let searcher = CountingSearcher::default();
assert_eq!(searcher.parallel_threshold(), 4);
}