use super::hnsw::{HnswIndex, HnswMetric};
use crate::graph::schema::{CurrentSelection, DirGraph, EmbeddingStore};
use crate::graph::storage::GraphRead;
use petgraph::graph::NodeIndex;
use std::borrow::Cow;
use std::collections::{BTreeSet, BinaryHeap, HashSet};
#[derive(Clone, Copy, Debug)]
pub enum DistanceMetric {
Cosine,
DotProduct,
Euclidean,
Poincare,
}
impl DistanceMetric {
pub fn from_name(name: &str) -> Option<Self> {
match name {
"cosine" => Some(DistanceMetric::Cosine),
"dot_product" => Some(DistanceMetric::DotProduct),
"euclidean" => Some(DistanceMetric::Euclidean),
"poincare" => Some(DistanceMetric::Poincare),
_ => None,
}
}
}
#[derive(Clone, Debug)]
pub struct VectorSearchResult {
pub node_idx: NodeIndex,
pub score: f32,
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct VectorSearchOptions {
pub top_k: usize,
pub metric: DistanceMetric,
pub exact: bool,
use_stored_metric: bool,
}
impl Default for VectorSearchOptions {
fn default() -> Self {
Self {
top_k: 10,
metric: DistanceMetric::Cosine,
exact: false,
use_stored_metric: false,
}
}
}
impl VectorSearchOptions {
pub fn with_top_k(mut self, top_k: usize) -> Self {
self.top_k = top_k;
self
}
pub fn with_metric(mut self, metric: DistanceMetric) -> Self {
self.metric = metric;
self.use_stored_metric = false;
self
}
pub fn with_stored_metric(mut self) -> Self {
self.use_stored_metric = true;
self
}
pub fn with_exact(mut self, exact: bool) -> Self {
self.exact = exact;
self
}
}
const PARALLEL_THRESHOLD: usize = 10_000;
const HNSW_AUTO_MIN: usize = 256;
const HNSW_OVERSAMPLE: usize = 4;
fn parse_stored_metric(metric: &str) -> Result<DistanceMetric, String> {
DistanceMetric::from_name(metric).ok_or_else(|| {
format!(
"Embedding store uses unknown metric '{metric}'. Expected cosine, dot_product, euclidean, or poincare."
)
})
}
fn resolve_metric_from_nodes(
graph: &DirGraph,
nodes: impl Iterator<Item = NodeIndex>,
embedding_property: &str,
) -> Result<DistanceMetric, String> {
let mut seen_types = BTreeSet::new();
let mut metrics = BTreeSet::new();
for node in nodes {
let Some(node_type_key) = GraphRead::node_type_of(&graph.graph, node) else {
continue;
};
if seen_types.contains(&node_type_key) {
continue;
}
let node_type = graph.interner.resolve(node_type_key);
let key = (node_type.to_string(), embedding_property.to_string());
let Some(store) = graph.embeddings.get(&key) else {
seen_types.insert(node_type_key);
continue;
};
if store.get_embedding(node.index()).is_none() {
continue;
}
seen_types.insert(node_type_key);
metrics.insert(store.metric.as_deref().unwrap_or("cosine"));
}
if metrics.len() > 1 {
return Err(format!(
"Selected embedding stores use multiple stored metrics ({}); pass metric= explicitly",
metrics.into_iter().collect::<Vec<_>>().join(", ")
));
}
parse_stored_metric(metrics.into_iter().next().unwrap_or("cosine"))
}
pub fn vector_search(
graph: &DirGraph,
selection: &CurrentSelection,
embedding_property: &str,
query_vector: &[f32],
options: &VectorSearchOptions,
) -> Result<Vec<VectorSearchResult>, String> {
let _arena_guard = graph.graph.begin_query(); let VectorSearchOptions {
top_k,
metric,
exact,
use_stored_metric,
} = *options;
let _arena_guard = graph.graph.begin_query();
let level_count = selection.get_level_count();
if level_count == 0 {
return Ok(Vec::new());
}
let level = match selection.get_level(level_count - 1) {
Some(level) => level,
None => return Ok(Vec::new()),
};
if level.node_count() == 0 || top_k == 0 {
return Ok(Vec::new());
}
let candidates: Cow<'_, [NodeIndex]> = {
let mut groups = level.iter_groups();
match (groups.next(), groups.next()) {
(Some((_, nodes)), None) => Cow::Borrowed(nodes.as_slice()),
_ => Cow::Owned(level.get_all_nodes()),
}
};
let first_candidate = candidates[0];
let first_type = GraphRead::node_type_of(&graph.graph, first_candidate);
let tentative_single_type = first_type.and_then(|node_type| {
let node_type = graph.interner.resolve(node_type);
let key = (node_type.to_string(), embedding_property.to_string());
graph.embeddings.get(&key).map(|store| (node_type, store))
});
let ordered_whole_store = tentative_single_type.is_some_and(|(_, store)| {
candidates.len() == store.len()
&& candidates
.iter()
.zip(&store.slot_to_node)
.all(|(candidate, &stored)| candidate.index() == stored)
});
let single_type = tentative_single_type.filter(|_| {
ordered_whole_store
|| first_type.is_some_and(|expected| {
candidates.iter().all(|&candidate| {
GraphRead::node_type_of(&graph.graph, candidate)
.is_none_or(|node_type| node_type == expected)
})
})
});
let metric = if use_stored_metric {
match single_type {
Some((_, store)) => parse_stored_metric(store.metric.as_deref().unwrap_or("cosine"))?,
None => {
resolve_metric_from_nodes(graph, candidates.iter().copied(), embedding_property)?
}
}
} else {
metric
};
let results = if let Some((node_type, store)) = single_type {
if query_vector.len() != store.dimension {
return Err(format!(
"Query vector dimension {} does not match embedding dimension {} for '{}.{}'",
query_vector.len(),
store.dimension,
node_type,
embedding_property
));
}
let scorer = Scorer::new(metric, query_vector);
let hnsw_result = if exact {
None
} else {
store.index.as_ref().and_then(|idx| {
let eligible = HnswMetric::from_distance(metric) == Some(idx.metric())
&& candidates.len() >= HNSW_AUTO_MIN
&& candidates.len().saturating_mul(2) >= store.len();
if eligible {
hnsw_search(
store,
idx,
candidates.as_ref(),
ordered_whole_store,
query_vector,
top_k,
&scorer,
)
} else {
None
}
})
};
match hnsw_result {
Some(r) => r,
None if candidates.len() > PARALLEL_THRESHOLD => {
parallel_search(&candidates, store, query_vector, top_k, &scorer)
}
None => sequential_search(&candidates, store, query_vector, top_k, &scorer),
}
} else {
let scorer = Scorer::new(metric, query_vector);
let mut heap = MinHeap::with_capacity(top_k);
let mut cached_type = None;
let mut cached_store = None;
for &node_idx in candidates.iter() {
let node_type = match GraphRead::node_type_of(&graph.graph, node_idx) {
Some(node_type) => node_type,
None => continue,
};
if cached_type != Some(node_type) {
let node_type_name = graph.interner.resolve(node_type);
let key = (node_type_name.to_string(), embedding_property.to_string());
cached_store = graph.embeddings.get(&key);
cached_type = Some(node_type);
}
let store = match cached_store {
Some(s) => s,
None => continue,
};
if query_vector.len() != store.dimension {
let node_type = graph.interner.resolve(node_type);
return Err(format!(
"Query vector dimension {} does not match embedding dimension {} for '{}.{}'",
query_vector.len(),
store.dimension,
node_type,
embedding_property
));
}
if let Some((embedding, norm)) = store.get_embedding_with_norm(node_idx.index()) {
let score = scorer.score(query_vector, embedding, norm);
heap.push_if_better(node_idx, score, top_k);
}
}
heap.into_sorted_results()
};
Ok(results)
}
fn hnsw_search(
store: &EmbeddingStore,
idx: &HnswIndex,
candidates: &[NodeIndex],
ordered_whole_store: bool,
query: &[f32],
top_k: usize,
scorer: &Scorer,
) -> Option<Vec<VectorSearchResult>> {
let membership: Option<HashSet<usize>> = if ordered_whole_store {
None
} else {
let selected: HashSet<usize> = candidates.iter().map(|n| n.index()).collect();
if store_is_fully_selected(store, |node| selected.contains(&node)) {
None
} else {
Some(selected)
}
};
let whole_store = membership.is_none();
let query_norm = dot_product(query, query).sqrt();
let k_fetch = top_k
.saturating_mul(HNSW_OVERSAMPLE)
.min(store.len())
.max(top_k);
let ef = k_fetch.max(idx.params().ef_search);
let raw = idx.search(
query,
query_norm,
k_fetch,
Some(ef),
&store.data,
&store.norms,
);
let mut heap = MinHeap::with_capacity(top_k);
for (slot, _dist) in raw {
let node_raw = store.slot_to_node[slot as usize];
if let Some(set) = &membership {
if !set.contains(&node_raw) {
continue;
}
}
let start = slot as usize * store.dimension;
let emb = &store.data[start..start + store.dimension];
let norm = store.norms[slot as usize];
let score = scorer.score(query, emb, norm);
heap.push_if_better(NodeIndex::new(node_raw), score, top_k);
}
let results = heap.into_sorted_results();
if !whole_store && results.len() < top_k {
return None;
}
Some(results)
}
pub(crate) fn store_is_fully_selected(
store: &EmbeddingStore,
contains_node: impl Fn(usize) -> bool,
) -> bool {
store.slot_to_node.iter().copied().all(contains_node)
}
type SimilarityFn = fn(&[f32], &[f32]) -> f32;
#[derive(Clone, Copy)]
pub struct Scorer {
kind: ScorerKind,
}
#[derive(Clone, Copy)]
enum ScorerKind {
Cosine { query_norm: f32 },
Generic(SimilarityFn),
}
impl Scorer {
pub fn new(metric: DistanceMetric, query: &[f32]) -> Self {
let kind = match metric {
DistanceMetric::Cosine => ScorerKind::Cosine {
query_norm: dot_product(query, query).sqrt(),
},
DistanceMetric::DotProduct => ScorerKind::Generic(dot_product),
DistanceMetric::Euclidean => ScorerKind::Generic(neg_euclidean_distance),
DistanceMetric::Poincare => ScorerKind::Generic(neg_poincare_distance),
};
Scorer { kind }
}
#[inline]
pub fn score(&self, query: &[f32], emb: &[f32], emb_norm: f32) -> f32 {
match self.kind {
ScorerKind::Cosine { query_norm } => {
let denom = query_norm * emb_norm;
if denom > 0.0 {
dot_product(query, emb) / denom
} else {
0.0
}
}
ScorerKind::Generic(f) => f(query, emb),
}
}
}
#[allow(dead_code)]
#[inline]
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
let (mut dot0, mut dot1, mut dot2, mut dot3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let (mut na0, mut na1, mut na2, mut na3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let (mut nb0, mut nb1, mut nb2, mut nb3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let a_chunks = a.chunks_exact(8);
let b_chunks = b.chunks_exact(8);
let a_rem = a_chunks.remainder();
let b_rem = b_chunks.remainder();
for (ac, bc) in a_chunks.zip(b_chunks) {
dot0 += ac[0] * bc[0];
dot1 += ac[1] * bc[1];
dot2 += ac[2] * bc[2];
dot3 += ac[3] * bc[3];
na0 += ac[0] * ac[0];
na1 += ac[1] * ac[1];
na2 += ac[2] * ac[2];
na3 += ac[3] * ac[3];
nb0 += bc[0] * bc[0];
nb1 += bc[1] * bc[1];
nb2 += bc[2] * bc[2];
nb3 += bc[3] * bc[3];
dot0 += ac[4] * bc[4];
dot1 += ac[5] * bc[5];
dot2 += ac[6] * bc[6];
dot3 += ac[7] * bc[7];
na0 += ac[4] * ac[4];
na1 += ac[5] * ac[5];
na2 += ac[6] * ac[6];
na3 += ac[7] * ac[7];
nb0 += bc[4] * bc[4];
nb1 += bc[5] * bc[5];
nb2 += bc[6] * bc[6];
nb3 += bc[7] * bc[7];
}
for (av, bv) in a_rem.iter().zip(b_rem.iter()) {
dot0 += av * bv;
na0 += av * av;
nb0 += bv * bv;
}
let dot = (dot0 + dot1) + (dot2 + dot3);
let norm_a = (na0 + na1) + (na2 + na3);
let norm_b = (nb0 + nb1) + (nb2 + nb3);
let denom = (norm_a * norm_b).sqrt();
if denom > 0.0 {
dot / denom
} else {
0.0
}
}
#[inline]
pub fn dot_product(a: &[f32], b: &[f32]) -> f32 {
let (mut s0, mut s1, mut s2, mut s3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let a_chunks = a.chunks_exact(8);
let b_chunks = b.chunks_exact(8);
let a_rem = a_chunks.remainder();
let b_rem = b_chunks.remainder();
for (ac, bc) in a_chunks.zip(b_chunks) {
s0 += ac[0] * bc[0];
s1 += ac[1] * bc[1];
s2 += ac[2] * bc[2];
s3 += ac[3] * bc[3];
s0 += ac[4] * bc[4];
s1 += ac[5] * bc[5];
s2 += ac[6] * bc[6];
s3 += ac[7] * bc[7];
}
for (av, bv) in a_rem.iter().zip(b_rem.iter()) {
s0 += av * bv;
}
(s0 + s1) + (s2 + s3)
}
#[inline]
pub fn neg_euclidean_distance(a: &[f32], b: &[f32]) -> f32 {
let (mut s0, mut s1, mut s2, mut s3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let a_chunks = a.chunks_exact(8);
let b_chunks = b.chunks_exact(8);
let a_rem = a_chunks.remainder();
let b_rem = b_chunks.remainder();
for (ac, bc) in a_chunks.zip(b_chunks) {
let d0 = ac[0] - bc[0];
let d1 = ac[1] - bc[1];
let d2 = ac[2] - bc[2];
let d3 = ac[3] - bc[3];
s0 += d0 * d0;
s1 += d1 * d1;
s2 += d2 * d2;
s3 += d3 * d3;
let d4 = ac[4] - bc[4];
let d5 = ac[5] - bc[5];
let d6 = ac[6] - bc[6];
let d7 = ac[7] - bc[7];
s0 += d4 * d4;
s1 += d5 * d5;
s2 += d6 * d6;
s3 += d7 * d7;
}
for (av, bv) in a_rem.iter().zip(b_rem.iter()) {
let d = av - bv;
s0 += d * d;
}
-((s0 + s1) + (s2 + s3)).sqrt()
}
#[inline]
pub fn neg_poincare_distance(a: &[f32], b: &[f32]) -> f32 {
let (mut na0, mut na1, mut na2, mut na3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let (mut nb0, mut nb1, mut nb2, mut nb3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let (mut d0, mut d1, mut d2, mut d3) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
let a_chunks = a.chunks_exact(8);
let b_chunks = b.chunks_exact(8);
let a_rem = a_chunks.remainder();
let b_rem = b_chunks.remainder();
for (ac, bc) in a_chunks.zip(b_chunks) {
na0 += ac[0] * ac[0];
na1 += ac[1] * ac[1];
na2 += ac[2] * ac[2];
na3 += ac[3] * ac[3];
nb0 += bc[0] * bc[0];
nb1 += bc[1] * bc[1];
nb2 += bc[2] * bc[2];
nb3 += bc[3] * bc[3];
let dd0 = ac[0] - bc[0];
let dd1 = ac[1] - bc[1];
let dd2 = ac[2] - bc[2];
let dd3 = ac[3] - bc[3];
d0 += dd0 * dd0;
d1 += dd1 * dd1;
d2 += dd2 * dd2;
d3 += dd3 * dd3;
na0 += ac[4] * ac[4];
na1 += ac[5] * ac[5];
na2 += ac[6] * ac[6];
na3 += ac[7] * ac[7];
nb0 += bc[4] * bc[4];
nb1 += bc[5] * bc[5];
nb2 += bc[6] * bc[6];
nb3 += bc[7] * bc[7];
let dd4 = ac[4] - bc[4];
let dd5 = ac[5] - bc[5];
let dd6 = ac[6] - bc[6];
let dd7 = ac[7] - bc[7];
d0 += dd4 * dd4;
d1 += dd5 * dd5;
d2 += dd6 * dd6;
d3 += dd7 * dd7;
}
for (av, bv) in a_rem.iter().zip(b_rem.iter()) {
na0 += av * av;
nb0 += bv * bv;
let dd = av - bv;
d0 += dd * dd;
}
let norm_a_sq = (na0 + na1) + (na2 + na3);
let norm_b_sq = (nb0 + nb1) + (nb2 + nb3);
let diff_sq = (d0 + d1) + (d2 + d3);
let alpha = (1.0 - norm_a_sq).max(1e-7); let beta = (1.0 - norm_b_sq).max(1e-7);
let gamma = 1.0 + 2.0 * diff_sq / (alpha * beta);
let gamma = gamma.max(1.0);
let dist = (gamma + (gamma * gamma - 1.0).sqrt()).ln();
-dist
}
struct MinHeap {
heap: BinaryHeap<ScoredNode>,
}
struct ScoredNode {
score: f32,
node_idx: NodeIndex,
}
impl PartialEq for ScoredNode {
fn eq(&self, other: &Self) -> bool {
self.score == other.score
}
}
impl Eq for ScoredNode {}
impl PartialOrd for ScoredNode {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for ScoredNode {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other
.score
.partial_cmp(&self.score)
.unwrap_or(std::cmp::Ordering::Equal)
}
}
impl MinHeap {
fn with_capacity(cap: usize) -> Self {
MinHeap {
heap: BinaryHeap::with_capacity(cap + 1),
}
}
#[inline]
fn push_if_better(&mut self, node_idx: NodeIndex, score: f32, top_k: usize) {
if self.heap.len() < top_k {
self.heap.push(ScoredNode { score, node_idx });
} else if let Some(min) = self.heap.peek() {
if score > min.score {
self.heap.pop();
self.heap.push(ScoredNode { score, node_idx });
}
}
}
fn into_sorted_results(self) -> Vec<VectorSearchResult> {
let mut results: Vec<VectorSearchResult> = self
.heap
.into_vec()
.into_iter()
.map(|sn| VectorSearchResult {
node_idx: sn.node_idx,
score: sn.score,
})
.collect();
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
results
}
}
fn sequential_search(
candidates: &[NodeIndex],
store: &EmbeddingStore,
query: &[f32],
top_k: usize,
scorer: &Scorer,
) -> Vec<VectorSearchResult> {
let mut heap = MinHeap::with_capacity(top_k);
for &node_idx in candidates {
if let Some((embedding, norm)) = store.get_embedding_with_norm(node_idx.index()) {
let score = scorer.score(query, embedding, norm);
heap.push_if_better(node_idx, score, top_k);
}
}
heap.into_sorted_results()
}
fn parallel_search(
candidates: &[NodeIndex],
store: &EmbeddingStore,
query: &[f32],
top_k: usize,
scorer: &Scorer,
) -> Vec<VectorSearchResult> {
use rayon::prelude::*;
let chunk_size = (candidates.len() / rayon::current_num_threads()).max(1024);
let per_thread_results: Vec<Vec<VectorSearchResult>> = candidates
.par_chunks(chunk_size)
.map(|chunk| sequential_search(chunk, store, query, top_k, scorer))
.collect();
let mut heap = MinHeap::with_capacity(top_k);
for thread_results in per_thread_results {
for result in thread_results {
heap.push_if_better(result.node_idx, result.score, top_k);
}
}
heap.into_sorted_results()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::datatypes::Value;
use crate::graph::algorithms::hnsw::HnswParams;
use crate::graph::schema::NodeData;
use crate::graph::storage::GraphWrite;
use std::collections::HashMap;
fn selection_of(nodes: Vec<NodeIndex>) -> CurrentSelection {
let mut selection = CurrentSelection::new();
selection
.get_level_mut(0)
.expect("initial selection level")
.add_selection(None, nodes);
selection
}
#[test]
fn hnsw_mixed_selection_never_returns_unselected_store_nodes() {
const DOCS: usize = 320;
const OMITTED_DOCS: usize = 32;
const UNEMBEDDED: usize = 320;
let mut graph = DirGraph::new();
let mut docs = Vec::with_capacity(DOCS);
for id in 0..DOCS {
let node = NodeData::new(
Value::Int64(id as i64),
Value::String(format!("Doc {id}")),
"Doc".to_string(),
HashMap::new(),
&mut graph.interner,
);
let idx = GraphWrite::add_node(&mut graph.graph, node);
graph
.type_indices
.entry_or_default("Doc".to_string())
.push(idx);
docs.push(idx);
}
let mut unembedded = Vec::with_capacity(UNEMBEDDED);
for id in 0..UNEMBEDDED {
let node = NodeData::new(
Value::Int64(id as i64),
Value::String(format!("Other {id}")),
"Doc".to_string(),
HashMap::new(),
&mut graph.interner,
);
let idx = GraphWrite::add_node(&mut graph.graph, node);
graph
.type_indices
.entry_or_default("Doc".to_string())
.push(idx);
unembedded.push(idx);
}
let mut store = EmbeddingStore::with_metric(2, "euclidean");
for (id, &node) in docs.iter().enumerate() {
store.set_embedding(node.index(), &[id as f32, 0.0]);
}
store
.build_index(DistanceMetric::Euclidean, HnswParams::default(), 7)
.expect("build HNSW index");
graph
.embeddings
.insert(("Doc".to_string(), "summary_emb".to_string()), store);
let selected_docs = docs[OMITTED_DOCS..].to_vec();
let mut mixed_nodes = selected_docs.clone();
mixed_nodes.extend(unembedded);
assert!(mixed_nodes.len() >= DOCS);
let store = graph
.embeddings
.get(&("Doc".to_string(), "summary_emb".to_string()))
.expect("inserted embedding store");
let mixed_store_members: HashSet<usize> =
mixed_nodes.iter().map(|node| node.index()).collect();
assert!(!store_is_fully_selected(store, |node| {
mixed_store_members.contains(&node)
}));
let whole_store_members: HashSet<usize> = docs.iter().map(|node| node.index()).collect();
assert!(store_is_fully_selected(store, |node| {
whole_store_members.contains(&node)
}));
let duplicate_candidates = vec![docs[OMITTED_DOCS]; DOCS];
let duplicate_nodes: HashSet<usize> = duplicate_candidates
.iter()
.map(|node| node.index())
.collect();
assert!(
!store_is_fully_selected(store, |node| duplicate_nodes.contains(&node)),
"duplicate candidates cannot stand in for omitted store nodes"
);
let mixed_members: HashSet<NodeIndex> = mixed_nodes.iter().copied().collect();
let query = [0.0, 0.0];
let options = VectorSearchOptions::default()
.with_top_k(5)
.with_metric(DistanceMetric::Euclidean);
let mixed = vector_search(
&graph,
&selection_of(mixed_nodes.clone()),
"summary_emb",
&query,
&options,
)
.expect("mixed approximate search");
assert_eq!(mixed.len(), 5);
assert!(
mixed
.iter()
.all(|result| mixed_members.contains(&result.node_idx)),
"mixed HNSW results escaped the current selection: {:?}",
mixed
.iter()
.map(|result| result.node_idx)
.collect::<Vec<_>>()
);
let mixed_exact = vector_search(
&graph,
&selection_of(mixed_nodes),
"summary_emb",
&query,
&options.clone().with_exact(true),
)
.expect("mixed exact search");
assert_eq!(mixed_exact.len(), 5);
assert_eq!(mixed_exact[0].node_idx, docs[OMITTED_DOCS]);
assert!(mixed_exact
.iter()
.all(|result| mixed_members.contains(&result.node_idx)));
let filtered_members: HashSet<NodeIndex> = selected_docs.iter().copied().collect();
let filtered = vector_search(
&graph,
&selection_of(selected_docs.clone()),
"summary_emb",
&query,
&options,
)
.expect("filtered approximate search");
assert_eq!(filtered.len(), 5);
assert!(filtered
.iter()
.all(|result| filtered_members.contains(&result.node_idx)));
let exact = vector_search(
&graph,
&selection_of(selected_docs),
"summary_emb",
&query,
&options.clone().with_exact(true),
)
.expect("filtered exact search");
assert_eq!(exact.len(), 5);
assert_eq!(exact[0].node_idx, docs[OMITTED_DOCS]);
let whole = vector_search(
&graph,
&selection_of(docs.clone()),
"summary_emb",
&query,
&options,
)
.expect("whole-store approximate search");
let whole_members: HashSet<NodeIndex> = docs.iter().copied().collect();
assert_eq!(whole.len(), 5);
assert!(whole
.iter()
.all(|result| whole_members.contains(&result.node_idx)));
let whole_exact = vector_search(
&graph,
&selection_of(docs.clone()),
"summary_emb",
&query,
&options.with_exact(true),
)
.expect("whole-store exact search");
assert_eq!(whole_exact.len(), 5);
assert_eq!(whole_exact[0].node_idx, docs[0]);
assert!(whole_exact
.iter()
.all(|result| whole_members.contains(&result.node_idx)));
}
#[test]
fn hnsw_index_is_not_used_for_a_different_requested_metric() {
const DOCS: usize = 320;
let mut graph = DirGraph::new();
let mut docs = Vec::with_capacity(DOCS);
let mut store = EmbeddingStore::with_metric(2, "cosine");
for id in 0..DOCS {
let node = NodeData::new(
Value::Int64(id as i64),
Value::String(format!("Doc {id}")),
"Doc".to_string(),
HashMap::new(),
&mut graph.interner,
);
let idx = GraphWrite::add_node(&mut graph.graph, node);
graph
.type_indices
.entry_or_default("Doc".to_string())
.push(idx);
docs.push(idx);
let embedding = match id {
0 => [1.0, 0.0],
1 => [100.0, 100.0],
_ => [1.0, 0.1],
};
store.set_embedding(idx.index(), &embedding);
}
store
.build_index(DistanceMetric::Cosine, HnswParams::default(), 7)
.unwrap();
graph
.embeddings
.insert(("Doc".to_string(), "summary_emb".to_string()), store);
let options = VectorSearchOptions::default()
.with_top_k(1)
.with_metric(DistanceMetric::DotProduct);
let automatic = vector_search(
&graph,
&selection_of(docs.clone()),
"summary_emb",
&[1.0, 0.0],
&options,
)
.unwrap();
let exact = vector_search(
&graph,
&selection_of(docs),
"summary_emb",
&[1.0, 0.0],
&options.with_exact(true),
)
.unwrap();
assert_eq!(automatic[0].node_idx, exact[0].node_idx);
assert_eq!(exact[0].node_idx, NodeIndex::new(1));
}
#[test]
fn test_cosine_similarity_identical() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![1.0, 2.0, 3.0, 4.0];
let sim = cosine_similarity(&a, &b);
assert!((sim - 1.0).abs() < 1e-6);
}
#[test]
fn test_cosine_similarity_orthogonal() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![0.0, 1.0, 0.0];
let sim = cosine_similarity(&a, &b);
assert!(sim.abs() < 1e-6);
}
#[test]
fn test_cosine_similarity_opposite() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![-1.0, -2.0, -3.0];
let sim = cosine_similarity(&a, &b);
assert!((sim + 1.0).abs() < 1e-6);
}
#[test]
fn test_cosine_similarity_large_vector() {
let a: Vec<f32> = (0..100).map(|i| i as f32).collect();
let b: Vec<f32> = (0..100).map(|i| (i * 2) as f32).collect();
let sim = cosine_similarity(&a, &b);
assert!(sim > 0.99); }
#[test]
fn test_dot_product_basic() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
let dp = dot_product(&a, &b);
assert!((dp - 32.0).abs() < 1e-6); }
#[test]
fn test_neg_euclidean_distance_identical() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![1.0, 2.0, 3.0];
let d = neg_euclidean_distance(&a, &b);
assert!(d.abs() < 1e-6); }
#[test]
fn test_neg_euclidean_distance_basic() {
let a = vec![0.0, 0.0, 0.0];
let b = vec![3.0, 4.0, 0.0];
let d = neg_euclidean_distance(&a, &b);
assert!((d + 5.0).abs() < 1e-6); }
#[test]
fn test_min_heap_top_k() {
let mut heap = MinHeap::with_capacity(3);
let scores = [0.5, 0.9, 0.1, 0.8, 0.3, 0.95, 0.2];
for (i, &score) in scores.iter().enumerate() {
heap.push_if_better(NodeIndex::new(i), score, 3);
}
let results = heap.into_sorted_results();
assert_eq!(results.len(), 3);
assert!((results[0].score - 0.95).abs() < 1e-6);
assert!((results[1].score - 0.9).abs() < 1e-6);
assert!((results[2].score - 0.8).abs() < 1e-6);
}
#[test]
fn test_embedding_store_basic() {
let mut store = EmbeddingStore::new(3);
store.set_embedding(0, &[1.0, 2.0, 3.0]);
store.set_embedding(5, &[4.0, 5.0, 6.0]);
assert_eq!(store.len(), 2);
assert_eq!(store.get_embedding(0), Some([1.0, 2.0, 3.0].as_slice()));
assert_eq!(store.get_embedding(5), Some([4.0, 5.0, 6.0].as_slice()));
assert_eq!(store.get_embedding(1), None);
}
#[test]
fn test_embedding_store_replace() {
let mut store = EmbeddingStore::new(2);
store.set_embedding(0, &[1.0, 2.0]);
store.set_embedding(0, &[3.0, 4.0]);
assert_eq!(store.len(), 1);
assert_eq!(store.get_embedding(0), Some([3.0, 4.0].as_slice()));
}
#[test]
fn test_cosine_similarity_zero_vector() {
let a = vec![0.0, 0.0, 0.0];
let b = vec![1.0, 2.0, 3.0];
let sim = cosine_similarity(&a, &b);
assert_eq!(sim, 0.0);
}
#[test]
fn test_poincare_identical_vectors() {
let a = vec![0.3, 0.2, 0.1];
let score = neg_poincare_distance(&a, &a);
assert!(
(score - 0.0).abs() < 1e-5,
"identical vectors should have distance 0, got {}",
score
);
}
#[test]
fn test_poincare_origin_to_point() {
let origin = vec![0.0, 0.0, 0.0];
let point = vec![0.5, 0.0, 0.0];
let score = neg_poincare_distance(&origin, &point);
let expected = -((1.6667f32 + (1.6667f32 * 1.6667f32 - 1.0).sqrt()).ln());
assert!(
(score - expected).abs() < 0.01,
"got {}, expected {}",
score,
expected
);
}
#[test]
fn test_poincare_distance_increases_near_boundary() {
let origin = vec![0.0, 0.0, 0.0];
let near = vec![0.1, 0.0, 0.0];
let mid = vec![0.5, 0.0, 0.0];
let far = vec![0.9, 0.0, 0.0];
let score_near = neg_poincare_distance(&origin, &near);
let score_mid = neg_poincare_distance(&origin, &mid);
let score_far = neg_poincare_distance(&origin, &far);
assert!(
score_near > score_mid,
"near {} should > mid {}",
score_near,
score_mid
);
assert!(
score_mid > score_far,
"mid {} should > far {}",
score_mid,
score_far
);
}
#[test]
fn test_poincare_symmetry() {
let a = vec![0.3, 0.2, 0.1];
let b = vec![0.1, 0.4, 0.2];
let d_ab = neg_poincare_distance(&a, &b);
let d_ba = neg_poincare_distance(&b, &a);
assert!(
(d_ab - d_ba).abs() < 1e-6,
"should be symmetric: {} vs {}",
d_ab,
d_ba
);
}
#[test]
fn test_poincare_large_vector() {
let a = vec![0.1; 16];
let b = vec![0.2; 16];
let score = neg_poincare_distance(&a, &b);
assert!(score < 0.0, "different vectors should have negative score");
assert!(score.is_finite(), "score should be finite");
}
#[test]
fn test_scorer_cosine_matches_kernel() {
let cases: Vec<(Vec<f32>, Vec<f32>)> = vec![
(vec![1.0, 2.0, 3.0, 4.0], vec![4.0, 3.0, 2.0, 1.0]),
(vec![0.1; 16], vec![0.2; 16]),
(
(0..100).map(|i| i as f32).collect(),
(0..100).map(|i| (i as f32 * 0.37).sin()).collect(),
),
(vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0]),
(vec![1.0, 2.0, 3.0], vec![-1.0, -2.0, -3.0]),
(vec![0.5, 1.5, 2.5, 3.5, 4.5], vec![5.5, 4.5, 3.5, 2.5, 1.5]),
];
for (q, v) in cases {
let mut store = EmbeddingStore::new(q.len());
store.set_embedding(0, &v);
let (emb, norm) = store.get_embedding_with_norm(0).unwrap();
let scorer = Scorer::new(DistanceMetric::Cosine, &q);
let got = scorer.score(&q, emb, norm);
let expected = cosine_similarity(&q, &v);
assert!(
(got - expected).abs() < 1e-5,
"cosine parity failed: scorer={}, kernel={}",
got,
expected
);
}
}
#[test]
fn test_scorer_cosine_zero_vectors() {
let mut store = EmbeddingStore::new(3);
store.set_embedding(0, &[0.0, 0.0, 0.0]);
let (emb, norm) = store.get_embedding_with_norm(0).unwrap();
let scorer = Scorer::new(DistanceMetric::Cosine, &[1.0, 2.0, 3.0]);
assert_eq!(scorer.score(&[1.0, 2.0, 3.0], emb, norm), 0.0);
store.set_embedding(0, &[1.0, 2.0, 3.0]);
let (emb, norm) = store.get_embedding_with_norm(0).unwrap();
let scorer = Scorer::new(DistanceMetric::Cosine, &[0.0, 0.0, 0.0]);
assert_eq!(scorer.score(&[0.0, 0.0, 0.0], emb, norm), 0.0);
}
#[test]
fn test_scorer_generic_metrics_match_kernels() {
let q = vec![1.0, 2.0, 3.0, 4.0];
let v = vec![4.0, 3.0, 2.0, 1.0];
let mut store = EmbeddingStore::new(4);
store.set_embedding(0, &v);
let (emb, norm) = store.get_embedding_with_norm(0).unwrap();
let dot = Scorer::new(DistanceMetric::DotProduct, &q);
assert!((dot.score(&q, emb, norm) - dot_product(&q, &v)).abs() < 1e-6);
let euc = Scorer::new(DistanceMetric::Euclidean, &q);
assert!((euc.score(&q, emb, norm) - neg_euclidean_distance(&q, &v)).abs() < 1e-6);
let poi = Scorer::new(DistanceMetric::Poincare, &q);
assert!((poi.score(&q, emb, norm) - neg_poincare_distance(&q, &v)).abs() < 1e-6);
}
#[test]
fn test_embedding_store_norm_cache() {
let mut store = EmbeddingStore::new(3);
store.set_embedding(0, &[3.0, 4.0, 0.0]); store.set_embedding(7, &[0.0, 0.0, 0.0]); let (_, n0) = store.get_embedding_with_norm(0).unwrap();
let (_, n7) = store.get_embedding_with_norm(7).unwrap();
assert!((n0 - 5.0).abs() < 1e-6);
assert_eq!(n7, 0.0);
store.set_embedding(0, &[5.0, 12.0, 0.0]); let (_, n0b) = store.get_embedding_with_norm(0).unwrap();
assert!((n0b - 13.0).abs() < 1e-6);
store.norms.clear();
store.rebuild_norms();
let (_, n0c) = store.get_embedding_with_norm(0).unwrap();
let (_, n7c) = store.get_embedding_with_norm(7).unwrap();
assert!((n0c - 13.0).abs() < 1e-6);
assert_eq!(n7c, 0.0);
}
#[test]
fn test_poincare_numerical_stability_near_boundary() {
let a = vec![0.999, 0.0, 0.0];
let b = vec![0.0, 0.999, 0.0];
let score = neg_poincare_distance(&a, &b);
assert!(
score.is_finite(),
"should not produce infinity near boundary"
);
assert!(score < 0.0, "should be negative");
}
}