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())
}
}