weavatrix-search-vector 0.3.1

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
use super::{Candidate, VectorStore, distance};
use crate::config::DistanceMetric;
use crate::error::SearchError;
use crate::hit::SearchHit;
use std::collections::BinaryHeap;

impl VectorStore {
    pub(crate) fn len(&self) -> usize {
        self.keys.len()
    }

    pub(crate) fn is_empty(&self) -> bool {
        self.keys.is_empty()
    }

    pub(crate) fn estimated_bytes(&self) -> usize {
        self.keys
            .capacity()
            .saturating_mul(std::mem::size_of::<u64>())
            .saturating_add(
                self.values
                    .capacity()
                    .saturating_mul(std::mem::size_of::<f32>()),
            )
            .saturating_add(
                self.routing_signs
                    .capacity()
                    .saturating_mul(std::mem::size_of::<u16>()),
            )
            .saturating_add(
                self.squared_norms
                    .capacity()
                    .saturating_mul(std::mem::size_of::<f32>()),
            )
    }

    pub(crate) fn key(&self, index: usize) -> u64 {
        self.keys[index]
    }

    pub(crate) fn keys(&self) -> &[u64] {
        &self.keys
    }

    pub(crate) fn values(&self) -> &[f32] {
        &self.values
    }

    pub(crate) fn find_index(&self, key: u64) -> Option<usize> {
        self.keys.binary_search(&key).ok()
    }

    pub(crate) fn vector(&self, index: usize) -> &[f32] {
        let start = index * self.dimensions;
        &self.values[start..start + self.dimensions]
    }

    pub(crate) fn distance_indices(&self, left: usize, right: usize) -> f32 {
        distance(
            self.distance_kernel,
            self.metric,
            self.vector(left),
            self.squared_norms[left],
            self.vector(right),
            self.squared_norms[right],
        )
    }

    pub(crate) fn query_squared_norm(&self, query: &[f32]) -> Result<f32, SearchError> {
        if query.len() != self.dimensions {
            return Err(SearchError::DimensionMismatch {
                expected: self.dimensions,
                actual: query.len(),
                vector: None,
            });
        }
        let norm = super::squared_norm(query, None)?;
        if self.metric == DistanceMetric::Cosine && norm == 0.0 {
            return Err(SearchError::ZeroVector { vector: None });
        }
        Ok(norm)
    }

    pub(crate) fn distance_query(&self, index: usize, query: &[f32], query_norm: f32) -> f32 {
        distance(
            self.distance_kernel,
            self.metric,
            self.vector(index),
            self.squared_norms[index],
            query,
            query_norm,
        )
    }

    pub(crate) fn exact(&self, query: &[f32], count: usize) -> Result<Vec<SearchHit>, SearchError> {
        self.exact_filtered(query, count, |_| true)
    }

    pub(crate) fn exact_filtered<F>(
        &self,
        query: &[f32],
        count: usize,
        mut accepts: F,
    ) -> Result<Vec<SearchHit>, SearchError>
    where
        F: FnMut(u64) -> bool,
    {
        let query_norm = self.query_squared_norm(query)?;
        let limit = count.min(self.len());
        if limit == 0 {
            return Ok(Vec::new());
        }
        let mut best = BinaryHeap::with_capacity(limit);
        for index in 0..self.len() {
            if !accepts(self.key(index)) {
                continue;
            }
            let candidate = Candidate::new(self.distance_query(index, query, query_norm), index);
            if best.len() < limit {
                best.push(candidate);
            } else if best
                .peek()
                .is_some_and(|worst| candidate.cmp(worst).is_lt())
            {
                best.pop();
                best.push(candidate);
            }
        }
        let mut candidates = best.into_vec();
        candidates.sort_unstable();
        Ok(candidates
            .into_iter()
            .map(|candidate| SearchHit {
                key: self.key(candidate.index()),
                distance: candidate.distance,
            })
            .collect())
    }
}