use std::sync::Arc;
use roaring::RoaringTreemap;
use serde::{Deserialize, Serialize};
use crate::error::Result;
use crate::vector::core::vector::Vector;
use crate::vector::search::filter_set::FilterSet;
#[cfg_attr(not(feature = "native"), allow(dead_code))]
pub(crate) const PARALLEL_SCAN_THRESHOLD: usize = 2048;
pub(crate) fn parallel_scan<I, T, F>(items: &[I], compute: F) -> Result<Vec<T>>
where
I: Sync,
T: Send,
F: Fn(&I) -> Result<Option<T>> + Sync + Send,
{
#[cfg(feature = "native")]
{
if items.len() >= PARALLEL_SCAN_THRESHOLD {
use rayon::prelude::*;
return Ok(items
.par_iter()
.map(&compute)
.collect::<Result<Vec<_>>>()?
.into_iter()
.flatten()
.collect());
}
}
let mut out = Vec::with_capacity(items.len());
for item in items {
if let Some(t) = compute(item)? {
out.push(t);
}
}
Ok(out)
}
#[derive(Debug, Clone)]
pub struct VectorIndexQuery {
pub query: Vector,
pub params: VectorIndexQueryParams,
pub field_name: Option<String>,
pub filter: Option<Arc<FilterSet>>,
}
impl VectorIndexQuery {
pub fn new(query: Vector) -> Self {
VectorIndexQuery {
query,
params: VectorIndexQueryParams::default(),
field_name: None,
filter: None,
}
}
pub fn filter(mut self, filter: Arc<FilterSet>) -> Self {
self.filter = Some(filter);
self
}
pub fn top_k(mut self, top_k: usize) -> Self {
self.params.top_k = top_k;
self
}
pub fn min_similarity(mut self, threshold: f32) -> Self {
self.params.min_similarity = threshold;
self
}
pub fn include_scores(mut self, include: bool) -> Self {
self.params.include_scores = include;
self
}
pub fn include_vectors(mut self, include: bool) -> Self {
self.params.include_vectors = include;
self
}
pub fn timeout_ms(mut self, timeout: u64) -> Self {
self.params.timeout_ms = Some(timeout);
self
}
pub fn field_name(mut self, field_name: String) -> Self {
self.field_name = Some(field_name);
self
}
pub fn rerank_factor(mut self, factor: usize) -> Self {
self.params.rerank_factor = Some(factor);
self
}
pub fn ef_search(mut self, ef: usize) -> Self {
self.params.ef_search = Some(ef);
self
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VectorIndexQueryParams {
pub top_k: usize,
pub min_similarity: f32,
pub include_scores: bool,
pub include_vectors: bool,
pub timeout_ms: Option<u64>,
#[serde(default)]
pub rerank_factor: Option<usize>,
#[serde(default)]
pub ef_search: Option<usize>,
}
impl Default for VectorIndexQueryParams {
fn default() -> Self {
Self {
top_k: 10,
min_similarity: 0.0,
include_scores: true,
include_vectors: false,
timeout_ms: None,
rerank_factor: None,
ef_search: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VectorIndexQueryResult {
pub doc_id: u64,
pub field_name: String,
pub similarity: f32,
pub distance: f32,
pub vector: Option<Vector>,
}
pub(crate) const SCORE_BASIS_METADATA_KEY: &str = "score_basis";
pub(crate) const SCORE_BASIS_F32_RERANK: &str = "f32-rerank";
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VectorIndexQueryResults {
pub results: Vec<VectorIndexQueryResult>,
pub candidates_examined: usize,
pub search_time_ms: f64,
pub query_metadata: std::collections::HashMap<String, String>,
}
impl VectorIndexQueryResults {
pub fn new() -> Self {
Self {
results: Vec::new(),
candidates_examined: 0,
search_time_ms: 0.0,
query_metadata: std::collections::HashMap::new(),
}
}
pub fn is_empty(&self) -> bool {
self.results.is_empty()
}
pub fn len(&self) -> usize {
self.results.len()
}
pub fn sort_by_similarity(&mut self) {
self.results
.sort_by(|a, b| b.similarity.total_cmp(&a.similarity));
}
pub fn sort_by_distance(&mut self) {
self.results
.sort_by(|a, b| a.distance.total_cmp(&b.distance));
}
pub fn take_top_k(&mut self, k: usize) {
if self.results.len() > k {
self.results.truncate(k);
}
}
pub fn filter_by_similarity(&mut self, min_similarity: f32) {
self.results
.retain(|result| result.similarity >= min_similarity);
}
pub fn best_result(&self) -> Option<&VectorIndexQueryResult> {
self.results
.iter()
.max_by(|a, b| a.similarity.total_cmp(&b.similarity))
}
}
impl Default for VectorIndexQueryResults {
fn default() -> Self {
Self::new()
}
}
pub trait VectorIndexSearcher: Send + Sync + std::fmt::Debug {
fn search(&self, request: &VectorIndexQuery) -> Result<VectorIndexQueryResults>;
fn count(&self, request: VectorIndexQuery) -> Result<u64>;
fn warmup(&mut self) -> Result<()> {
Ok(())
}
fn parallel_threshold(&self) -> usize {
4
}
fn search_batch(&self, queries: &[VectorIndexQuery]) -> Result<Vec<VectorIndexQueryResults>> {
self.search_batch_with_threshold(queries, self.parallel_threshold())
}
#[doc(hidden)]
fn search_batch_with_threshold(
&self,
queries: &[VectorIndexQuery],
parallel_threshold: usize,
) -> Result<Vec<VectorIndexQueryResults>> {
#[cfg(feature = "native")]
{
use rayon::prelude::*;
if queries.len() >= parallel_threshold {
return queries
.par_iter()
.map(|q| self.search(q))
.collect::<Result<Vec<_>>>();
}
}
let _ = parallel_threshold;
queries.iter().map(|q| self.search(q)).collect()
}
}
#[derive(Debug, Clone)]
pub enum VectorSearchQuery {
Payloads(Vec<crate::vector::store::request::QueryPayload>),
Vectors(Vec<crate::vector::store::request::QueryVector>),
}
fn default_query_limit() -> usize {
10
}
fn default_overfetch() -> f32 {
2.0
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VectorSearchParams {
#[serde(default)]
pub fields: Option<Vec<crate::vector::store::request::FieldSelector>>,
#[serde(default = "default_query_limit")]
pub limit: usize,
#[serde(default)]
pub score_mode: crate::vector::store::request::VectorScoreMode,
#[serde(default = "default_overfetch")]
pub overfetch: f32,
#[serde(default)]
pub min_score: f32,
#[serde(skip)]
pub allowed_ids: Option<Vec<u64>>,
#[serde(skip)]
pub allowed_filter: Option<Arc<RoaringTreemap>>,
#[serde(default)]
pub rerank_factor: Option<usize>,
#[serde(default)]
pub ef_search: Option<usize>,
}
impl Default for VectorSearchParams {
fn default() -> Self {
Self {
fields: None,
limit: default_query_limit(),
score_mode: crate::vector::store::request::VectorScoreMode::default(),
overfetch: default_overfetch(),
min_score: 0.0,
allowed_ids: None,
allowed_filter: None,
rerank_factor: None,
ef_search: None,
}
}
}
impl VectorSearchParams {
pub(crate) fn overfetch_top_k(&self) -> usize {
if !self.overfetch.is_finite() || self.overfetch <= 1.0 {
return self.limit;
}
let scaled = (self.limit as f32 * self.overfetch).ceil();
if scaled >= usize::MAX as f32 {
usize::MAX
} else {
(scaled as usize).max(self.limit)
}
}
}
#[derive(Debug, Clone)]
pub struct VectorSearchRequest {
pub query: VectorSearchQuery,
pub params: VectorSearchParams,
}
impl Default for VectorSearchRequest {
fn default() -> Self {
Self {
query: VectorSearchQuery::Vectors(Vec::new()),
params: VectorSearchParams::default(),
}
}
}
pub trait VectorSearcher: Send + Sync + std::fmt::Debug {
fn search(
&self,
request: &VectorSearchRequest,
) -> crate::error::Result<crate::vector::store::response::VectorSearchResults>;
fn count(&self, request: &VectorSearchRequest) -> crate::error::Result<u64>;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parallel_scan_parallel_matches_serial() {
for n in [100usize, PARALLEL_SCAN_THRESHOLD + 500] {
let items: Vec<u64> = (0..n as u64).collect();
let mut got = parallel_scan(&items[..], |&x| Ok(Some(x * 2))).unwrap();
got.sort_unstable();
let expected: Vec<u64> = (0..n as u64).map(|x| x * 2).collect();
assert_eq!(got, expected, "n = {n}");
}
}
#[test]
fn parallel_scan_skips_none() {
for n in [100usize, PARALLEL_SCAN_THRESHOLD + 500] {
let items: Vec<u64> = (0..n as u64).collect();
let mut got =
parallel_scan(&items[..], |&x| Ok(if x % 2 == 0 { Some(x) } else { None }))
.unwrap();
got.sort_unstable();
let expected: Vec<u64> = (0..n as u64).filter(|x| x % 2 == 0).collect();
assert_eq!(got, expected, "n = {n}");
}
}
#[test]
fn parallel_scan_propagates_error() {
for n in [100usize, PARALLEL_SCAN_THRESHOLD + 500] {
let items: Vec<u64> = (0..n as u64).collect();
let result: Result<Vec<u64>> = parallel_scan(&items[..], |&x| {
if x == (n as u64 / 2) {
Err(crate::error::LaurusError::internal("boom"))
} else {
Ok(Some(x))
}
});
assert!(result.is_err(), "n = {n}");
}
}
fn params_with_overfetch(limit: usize, overfetch: f32) -> VectorSearchParams {
VectorSearchParams {
limit,
overfetch,
..Default::default()
}
}
#[test]
fn overfetch_top_k_scales_limit() {
assert_eq!(params_with_overfetch(10, 2.0).overfetch_top_k(), 20);
assert_eq!(params_with_overfetch(10, 3.0).overfetch_top_k(), 30);
assert_eq!(params_with_overfetch(10, 1.5).overfetch_top_k(), 15);
assert_eq!(params_with_overfetch(3, 1.5).overfetch_top_k(), 5);
}
#[test]
fn overfetch_top_k_default_is_2x() {
assert_eq!(VectorSearchParams::default().overfetch, 2.0);
assert_eq!(
params_with_overfetch(7, default_overfetch()).overfetch_top_k(),
14
);
}
#[test]
fn overfetch_top_k_clamps_low_and_degenerate_factors() {
assert_eq!(params_with_overfetch(10, 1.0).overfetch_top_k(), 10);
assert_eq!(params_with_overfetch(10, 0.5).overfetch_top_k(), 10);
assert_eq!(params_with_overfetch(10, 0.0).overfetch_top_k(), 10);
assert_eq!(params_with_overfetch(10, -1.0).overfetch_top_k(), 10);
assert_eq!(params_with_overfetch(10, f32::NAN).overfetch_top_k(), 10);
assert_eq!(
params_with_overfetch(10, f32::INFINITY).overfetch_top_k(),
10
);
assert_eq!(params_with_overfetch(0, 4.0).overfetch_top_k(), 0);
}
}