use parking_lot::RwLock;
use std::collections::BinaryHeap;
use crate::expression::Expression;
use crate::traits::Index;
use radixdb_core::{DataType, Error, IndexEntry, IndexType, Operator, Result, RowIdVec, Value};
use radixdb_core::{I64Map, I64Set};
mod distance;
use distance::{
as_f32_slice, bytes_to_f32_vec, cosine_distance_f32, ip_distance_f32, l2_distance_sq_f32,
prefetch_read,
};
const HNSW_GRAPH_HEADER_LEN: usize = 25;
const HNSW_MAX_GRAPH_BYTES: usize = 1 << 30;
const HNSW_MAX_M: usize = (u16::MAX as usize) / 2;
const HNSW_IDENTITY_MAGIC: &[u8; 4] = b"HNSI";
const HNSW_IDENTITY_VERSION: u32 = 1;
const HNSW_IDENTITY_HEADER_LEN: usize = 49;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum HnswDistanceMetric {
L2 = 0,
Cosine = 1,
InnerProduct = 2,
}
impl HnswDistanceMetric {
pub fn as_u8(self) -> u8 {
self as u8
}
pub fn from_u8(v: u8) -> Option<Self> {
match v {
0 => Some(Self::L2),
1 => Some(Self::Cosine),
2 => Some(Self::InnerProduct),
_ => None,
}
}
pub fn from_name(name: &str) -> Option<Self> {
match name {
"l2" | "euclidean" => Some(Self::L2),
"cosine" => Some(Self::Cosine),
"ip" | "inner_product" | "dot" => Some(Self::InnerProduct),
_ => None,
}
}
}
struct MaxEntry {
distance: f32,
node: u32,
}
impl PartialEq for MaxEntry {
fn eq(&self, other: &Self) -> bool {
self.distance == other.distance
}
}
impl Eq for MaxEntry {}
impl PartialOrd for MaxEntry {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for MaxEntry {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.distance.total_cmp(&other.distance)
}
}
struct MinEntry {
distance: f32,
node: u32,
}
impl PartialEq for MinEntry {
fn eq(&self, other: &Self) -> bool {
self.distance == other.distance
}
}
impl Eq for MinEntry {}
impl PartialOrd for MinEntry {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for MinEntry {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other.distance.total_cmp(&self.distance)
}
}
struct HnswNode {
neighbors: Vec<Vec<(u32, f32)>>,
}
struct SearchScratch {
visited: Vec<u64>,
candidates: BinaryHeap<MinEntry>,
result: BinaryHeap<MaxEntry>,
}
impl SearchScratch {
fn new() -> Self {
Self {
visited: Vec::new(),
candidates: BinaryHeap::new(),
result: BinaryHeap::new(),
}
}
#[inline]
fn reset(&mut self, node_count: usize) {
let words = (node_count + 63) >> 6;
if self.visited.len() < words {
self.visited.resize(words, 0);
} else {
self.visited[..words].fill(0);
}
self.candidates.clear();
self.result.clear();
}
}
struct QueryScratch {
visited: Vec<u64>,
candidates: BinaryHeap<MinEntry>,
result: BinaryHeap<MaxEntry>,
}
impl QueryScratch {
fn new() -> Self {
Self {
visited: Vec::new(),
candidates: BinaryHeap::new(),
result: BinaryHeap::new(),
}
}
#[inline]
fn reset(&mut self, node_count: usize) {
let words = (node_count + 63) >> 6;
if self.visited.len() < words {
self.visited.resize(words, 0);
} else {
self.visited[..words].fill(0);
}
self.candidates.clear();
self.result.clear();
}
}
thread_local! {
static QUERY_SCRATCH: std::cell::RefCell<QueryScratch> =
std::cell::RefCell::new(QueryScratch::new());
}
#[cfg(feature = "parallel")]
thread_local! {
static BUILD_SCRATCH: std::cell::RefCell<SearchScratch> =
std::cell::RefCell::new(SearchScratch::new());
}
struct HnswInner {
nodes: Vec<HnswNode>,
entry_point: Option<u32>,
max_layer: usize,
vectors: Vec<u8>,
dims_bytes: usize,
node_to_row_id: Vec<i64>,
row_id_to_node: ahash::AHashMap<i64, u32>,
metric: HnswDistanceMetric,
scratch: SearchScratch,
deleted_bits: Vec<u64>,
unique_map: Option<ahash::AHashMap<u64, Vec<i64>>>,
build_seed: u64,
}
impl HnswInner {
#[cfg(test)]
fn new(dims: usize, metric: HnswDistanceMetric) -> Self {
Self::with_build_seed(dims, metric, 0x7261_6469_7864_6231)
}
fn with_build_seed(dims: usize, metric: HnswDistanceMetric, build_seed: u64) -> Self {
Self {
nodes: Vec::new(),
entry_point: None,
max_layer: 0,
vectors: Vec::new(),
dims_bytes: dims
.checked_mul(std::mem::size_of::<f32>())
.expect("validated HNSW dimensions"),
node_to_row_id: Vec::new(),
row_id_to_node: ahash::AHashMap::new(),
metric,
scratch: SearchScratch::new(),
deleted_bits: Vec::new(),
unique_map: None,
build_seed,
}
}
#[inline]
fn hash_vec_bytes(bytes: &[u8]) -> u64 {
use std::hash::{Hash, Hasher};
let mut hasher = ahash::AHasher::default();
bytes.hash(&mut hasher);
hasher.finish()
}
fn build_unique_map(&mut self) {
let mut map: ahash::AHashMap<u64, Vec<i64>> =
ahash::AHashMap::with_capacity(self.nodes.len());
for (node_idx, &row_id) in self.node_to_row_id.iter().enumerate() {
if self.is_deleted(node_idx as u32) {
continue;
}
let offset = node_idx * self.dims_bytes;
if let Some(vec_bytes) = self.vectors.get(offset..offset + self.dims_bytes) {
let hash = Self::hash_vec_bytes(vec_bytes);
map.entry(hash).or_default().push(row_id);
}
}
self.unique_map = Some(map);
}
#[inline]
fn unique_map_insert(&mut self, node_id: u32) {
if let Some(ref mut map) = self.unique_map {
let row_id = self.node_to_row_id[node_id as usize];
let offset = node_id as usize * self.dims_bytes;
let hash = Self::hash_vec_bytes(&self.vectors[offset..offset + self.dims_bytes]);
map.entry(hash).or_default().push(row_id);
}
}
#[inline]
fn unique_map_remove(&mut self, node_id: u32) {
if let Some(ref mut map) = self.unique_map {
let row_id = self.node_to_row_id[node_id as usize];
let offset = node_id as usize * self.dims_bytes;
let hash = Self::hash_vec_bytes(&self.vectors[offset..offset + self.dims_bytes]);
if let Some(row_ids) = map.get_mut(&hash) {
row_ids.retain(|&r| r != row_id);
if row_ids.is_empty() {
map.remove(&hash);
}
}
}
}
#[inline(always)]
fn is_deleted(&self, node: u32) -> bool {
let idx = node as usize;
unsafe { (*self.deleted_bits.get_unchecked(idx >> 6) >> (idx & 63)) & 1 != 0 }
}
#[inline(always)]
fn set_deleted(&mut self, node: u32) {
let idx = node as usize;
unsafe {
*self.deleted_bits.get_unchecked_mut(idx >> 6) |= 1u64 << (idx & 63);
}
}
#[inline(always)]
fn clear_deleted(&mut self, node: u32) {
let idx = node as usize;
unsafe {
*self.deleted_bits.get_unchecked_mut(idx >> 6) &= !(1u64 << (idx & 63));
}
}
#[inline]
fn push_node_alive(&mut self) {
let node_count = self.nodes.len(); let words_needed = (node_count + 63) >> 6;
if self.deleted_bits.len() < words_needed {
self.deleted_bits.resize(words_needed, 0);
}
}
#[inline]
fn vector_bytes(&self, node: u32) -> &[u8] {
let start = node as usize * self.dims_bytes;
&self.vectors[start..start + self.dims_bytes]
}
#[inline]
fn vector_f32(&self, node: u32) -> &[f32] {
as_f32_slice(self.vector_bytes(node))
}
#[inline]
fn distance_f32(&self, node: u32, query: &[f32]) -> f32 {
let v = self.vector_f32(node);
match self.metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(v, query),
HnswDistanceMetric::Cosine => cosine_distance_f32(v, query),
HnswDistanceMetric::InnerProduct => ip_distance_f32(v, query),
}
}
fn search_layer(
&self,
query: &[f32],
entry_points: &[u32],
ef: usize,
layer: usize,
) -> Vec<MaxEntry> {
QUERY_SCRATCH.with(|cell| {
let mut scratch = cell.borrow_mut();
let node_count = self.nodes.len();
scratch.reset(node_count);
let visited_ptr = scratch.visited.as_mut_ptr();
let deleted_ptr = self.deleted_bits.as_ptr();
let vectors_ptr = self.vectors.as_ptr();
let dims_bytes = self.dims_bytes;
let dims = dims_bytes / 4;
let metric = self.metric;
for &ep in entry_points {
if (ep as usize) >= node_count {
continue;
}
let n_idx = ep as usize;
unsafe {
*visited_ptr.add(n_idx >> 6) |= 1u64 << (n_idx & 63);
}
let v_ptr = unsafe { vectors_ptr.add(n_idx * dims_bytes) as *const f32 };
let v_slice = unsafe { std::slice::from_raw_parts(v_ptr, dims) };
let d = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(v_slice, query),
HnswDistanceMetric::Cosine => cosine_distance_f32(v_slice, query),
HnswDistanceMetric::InnerProduct => ip_distance_f32(v_slice, query),
};
scratch.candidates.push(MinEntry {
distance: d,
node: ep,
});
scratch.result.push(MaxEntry {
distance: d,
node: ep,
});
}
let mut farthest_dist = scratch.result.peek().map_or(f32::MAX, |e| e.distance);
let mut result_len = scratch.result.len();
let nodes_ptr = self.nodes.as_ptr();
let nodes_len = self.nodes.len();
while let Some(MinEntry {
distance: c_dist,
node: c_id,
}) = scratch.candidates.pop()
{
if c_dist > farthest_dist && result_len >= ef {
break;
}
if let Some(next) = scratch.candidates.peek() {
let next_id = next.node as usize;
if next_id < nodes_len {
let next_node_ptr = unsafe { nodes_ptr.add(next_id) } as *const u8;
prefetch_read(next_node_ptr);
}
}
let node = unsafe { &*nodes_ptr.add(c_id as usize) };
if layer < node.neighbors.len() {
let neighbors = &node.neighbors[layer];
let nlen = neighbors.len();
let nptr = neighbors.as_ptr();
let mut ni = 0usize;
while ni < nlen {
let (neighbor, _) = unsafe { *nptr.add(ni) };
ni += 1;
let n_idx = neighbor as usize;
if ni < nlen {
let next_nb = unsafe { (*nptr.add(ni)).0 } as usize;
let next_v_word = next_nb >> 6;
prefetch_read(unsafe { visited_ptr.add(next_v_word) } as *const u8);
}
let v_word_idx = n_idx >> 6;
let v_bit_mask = 1u64 << (n_idx & 63);
let v_word_ptr = unsafe { visited_ptr.add(v_word_idx) };
let v_word = unsafe { *v_word_ptr };
if (v_word & v_bit_mask) != 0 {
continue;
}
unsafe {
*v_word_ptr = v_word | v_bit_mask;
}
let is_deleted =
unsafe { (*deleted_ptr.add(v_word_idx)) & v_bit_mask != 0 };
if is_deleted {
continue;
}
if ni < nlen {
let next_nb = unsafe { (*nptr.add(ni)).0 } as usize;
let vec_addr = unsafe { vectors_ptr.add(next_nb * dims_bytes) };
prefetch_read(vec_addr);
}
let v_ptr = unsafe { vectors_ptr.add(n_idx * dims_bytes) as *const f32 };
let v_slice = unsafe { std::slice::from_raw_parts(v_ptr, dims) };
let d = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(v_slice, query),
HnswDistanceMetric::Cosine => cosine_distance_f32(v_slice, query),
HnswDistanceMetric::InnerProduct => ip_distance_f32(v_slice, query),
};
if d < farthest_dist || result_len < ef {
scratch.candidates.push(MinEntry {
distance: d,
node: neighbor,
});
scratch.result.push(MaxEntry {
distance: d,
node: neighbor,
});
result_len += 1;
if result_len > ef {
scratch.result.pop();
result_len -= 1;
farthest_dist =
scratch.result.peek().map_or(f32::MAX, |e| e.distance);
} else {
farthest_dist = farthest_dist.max(d);
}
}
}
}
}
scratch.result.drain().collect()
})
}
fn select_neighbors(&self, mut candidates: Vec<MaxEntry>, m: usize) -> Vec<(u32, f32)> {
if candidates.is_empty() || m == 0 {
return Vec::new();
}
candidates.sort_unstable_by(|a, b| a.distance.total_cmp(&b.distance));
let mut selected: Vec<(u32, f32)> = Vec::with_capacity(m);
let mut pruned: Vec<(u32, f32)> = Vec::with_capacity(candidates.len());
let metric = self.metric;
for entry in &candidates {
if selected.len() >= m {
break;
}
let dist_to_query = entry.distance;
let entry_vec = self.vector_f32(entry.node);
let mut is_diverse = true;
for &(sel_node, _) in &selected {
let sel_vec = self.vector_f32(sel_node);
let dist_to_selected = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(entry_vec, sel_vec),
HnswDistanceMetric::Cosine => cosine_distance_f32(entry_vec, sel_vec),
HnswDistanceMetric::InnerProduct => ip_distance_f32(entry_vec, sel_vec),
};
if dist_to_selected < dist_to_query {
is_diverse = false;
break;
}
}
if is_diverse {
selected.push((entry.node, entry.distance));
} else {
pruned.push((entry.node, entry.distance));
}
}
for entry in pruned {
if selected.len() >= m {
break;
}
selected.push(entry);
}
selected
}
fn search_layer_mut(
&mut self,
query: &[f32],
entry_points: &[u32],
ef: usize,
layer: usize,
) -> Vec<MaxEntry> {
let node_count = self.nodes.len();
let mut scratch = std::mem::replace(&mut self.scratch, SearchScratch::new());
scratch.reset(node_count);
let visited_ptr = scratch.visited.as_mut_ptr();
let deleted_ptr = self.deleted_bits.as_ptr();
let vectors_ptr = self.vectors.as_ptr();
let dims_bytes = self.dims_bytes;
let dims = dims_bytes / 4;
let metric = self.metric;
for &ep in entry_points {
if (ep as usize) >= node_count {
continue;
}
let n_idx = ep as usize;
unsafe {
*visited_ptr.add(n_idx >> 6) |= 1u64 << (n_idx & 63);
}
let v_ptr = unsafe { vectors_ptr.add(n_idx * dims_bytes) as *const f32 };
let v_slice = unsafe { std::slice::from_raw_parts(v_ptr, dims) };
let d = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(v_slice, query),
HnswDistanceMetric::Cosine => cosine_distance_f32(v_slice, query),
HnswDistanceMetric::InnerProduct => ip_distance_f32(v_slice, query),
};
scratch.candidates.push(MinEntry {
distance: d,
node: ep,
});
scratch.result.push(MaxEntry {
distance: d,
node: ep,
});
}
let mut farthest_dist = scratch.result.peek().map_or(f32::MAX, |e| e.distance);
let mut result_len = scratch.result.len();
let nodes_ptr = self.nodes.as_ptr();
let nodes_len = self.nodes.len();
while let Some(MinEntry {
distance: c_dist,
node: c_id,
}) = scratch.candidates.pop()
{
if c_dist > farthest_dist && result_len >= ef {
break;
}
if let Some(next) = scratch.candidates.peek() {
let next_id = next.node as usize;
if next_id < nodes_len {
let next_node_ptr = unsafe { nodes_ptr.add(next_id) } as *const u8;
prefetch_read(next_node_ptr);
}
}
let node = unsafe { &*nodes_ptr.add(c_id as usize) };
if layer < node.neighbors.len() {
let neighbors = &node.neighbors[layer];
let nlen = neighbors.len();
let nptr = neighbors.as_ptr();
let mut ni = 0usize;
while ni < nlen {
let (neighbor, _) = unsafe { *nptr.add(ni) };
ni += 1;
let n_idx = neighbor as usize;
if ni < nlen {
let next_nb = unsafe { (*nptr.add(ni)).0 } as usize;
let next_v_word = next_nb >> 6;
prefetch_read(unsafe { visited_ptr.add(next_v_word) } as *const u8);
}
let v_word_idx = n_idx >> 6;
let v_bit_mask = 1u64 << (n_idx & 63);
let v_word_ptr = unsafe { visited_ptr.add(v_word_idx) };
let v_word = unsafe { *v_word_ptr };
if (v_word & v_bit_mask) != 0 {
continue;
}
unsafe {
*v_word_ptr = v_word | v_bit_mask;
}
let is_deleted = unsafe { (*deleted_ptr.add(v_word_idx)) & v_bit_mask != 0 };
if is_deleted {
continue;
}
if ni < nlen {
let next_nb = unsafe { (*nptr.add(ni)).0 } as usize;
let vec_addr = unsafe { vectors_ptr.add(next_nb * dims_bytes) };
prefetch_read(vec_addr);
}
let v_ptr = unsafe { vectors_ptr.add(n_idx * dims_bytes) as *const f32 };
let v_slice = unsafe { std::slice::from_raw_parts(v_ptr, dims) };
let d = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(v_slice, query),
HnswDistanceMetric::Cosine => cosine_distance_f32(v_slice, query),
HnswDistanceMetric::InnerProduct => ip_distance_f32(v_slice, query),
};
if d < farthest_dist || result_len < ef {
scratch.candidates.push(MinEntry {
distance: d,
node: neighbor,
});
scratch.result.push(MaxEntry {
distance: d,
node: neighbor,
});
result_len += 1;
if result_len > ef {
scratch.result.pop();
result_len -= 1;
farthest_dist = scratch.result.peek().map_or(f32::MAX, |e| e.distance);
} else {
farthest_dist = farthest_dist.max(d);
}
}
}
}
}
let result = scratch.result.drain().collect();
self.scratch = scratch;
result
}
#[inline]
fn greedy_closest(&self, query: &[f32], mut ep: u32, layer: usize) -> u32 {
let mut ep_dist = self.distance_f32(ep, query);
loop {
let mut best_dist = ep_dist;
let mut best_node = ep;
let node = &self.nodes[ep as usize];
if layer < node.neighbors.len() {
for &(neighbor, _) in &node.neighbors[layer] {
if self.is_deleted(neighbor) {
continue;
}
let d = self.distance_f32(neighbor, query);
if d < best_dist {
best_dist = d;
best_node = neighbor;
}
}
}
if best_node == ep {
return ep;
}
ep = best_node;
ep_dist = best_dist;
}
}
fn update_connection_fast(
&mut self,
target: u32,
new_nb: u32,
new_dist: f32,
layer: usize,
max_conn: usize,
) {
if max_conn == 0 || target == new_nb {
return;
}
let neighbors = &self.nodes[target as usize].neighbors[layer];
if let Some(existing_idx) = neighbors.iter().position(|&(nid, _)| nid == new_nb) {
if new_dist < neighbors[existing_idx].1 {
self.nodes[target as usize].neighbors[layer][existing_idx].1 = new_dist;
}
return;
}
let cur_len = neighbors.len();
if cur_len < max_conn {
self.nodes[target as usize].neighbors[layer].push((new_nb, new_dist));
return;
}
let mut farthest_idx = 0;
let mut farthest_dist = f32::MIN;
for (i, &(_, d)) in neighbors.iter().enumerate() {
if d > farthest_dist {
farthest_dist = d;
farthest_idx = i;
}
}
if new_dist < farthest_dist {
self.nodes[target as usize].neighbors[layer][farthest_idx] = (new_nb, new_dist);
}
}
fn update_connection(
&mut self,
target: u32,
new_nb: u32,
new_dist: f32,
layer: usize,
max_conn: usize,
prefer_diversity: bool,
) {
if !prefer_diversity {
self.update_connection_fast(target, new_nb, new_dist, layer, max_conn);
return;
}
if layer > 0 {
self.update_connection_fast(target, new_nb, new_dist, layer, max_conn);
return;
}
const DIVERSITY_OVERFLOW_MAX_NODES: usize = 200_000;
if self.nodes.len() > DIVERSITY_OVERFLOW_MAX_NODES {
self.update_connection_fast(target, new_nb, new_dist, layer, max_conn);
return;
}
if max_conn == 0 || target == new_nb {
return;
}
let neighbors = &self.nodes[target as usize].neighbors[layer];
if let Some(existing_idx) = neighbors.iter().position(|&(nid, _)| nid == new_nb) {
if new_dist < neighbors[existing_idx].1 {
self.nodes[target as usize].neighbors[layer][existing_idx].1 = new_dist;
}
return;
}
let cur_len = neighbors.len();
if cur_len < max_conn {
self.nodes[target as usize].neighbors[layer].push((new_nb, new_dist));
return;
}
let mut farthest_idx = 0;
let mut farthest_dist = f32::MIN;
for (i, &(_, d)) in neighbors.iter().enumerate() {
if d > farthest_dist {
farthest_dist = d;
farthest_idx = i;
}
}
if new_dist >= farthest_dist {
return;
}
if new_dist > farthest_dist * 0.90 {
self.nodes[target as usize].neighbors[layer][farthest_idx] = (new_nb, new_dist);
return;
}
let mut candidates: Vec<MaxEntry> = Vec::with_capacity(cur_len + 1);
for &(nid, dist) in neighbors {
candidates.push(MaxEntry {
distance: dist,
node: nid,
});
}
candidates.push(MaxEntry {
distance: new_dist,
node: new_nb,
});
let pruned = self.select_neighbors(candidates, max_conn);
self.nodes[target as usize].neighbors[layer] = pruned;
}
fn insert(
&mut self,
vector_bytes: &[u8],
row_id: i64,
m: usize,
m0: usize,
ef_construction: usize,
ml: f64,
) {
if let Some(&existing_node) = self.row_id_to_node.get(&row_id) {
if self.is_deleted(existing_node) {
let offset = existing_node as usize * self.dims_bytes;
self.vectors[offset..offset + self.dims_bytes].copy_from_slice(vector_bytes);
self.clear_deleted(existing_node);
self.unique_map_insert(existing_node);
let query_vec = bytes_to_f32_vec(vector_bytes);
let query = &query_vec;
let level = self.nodes[existing_node as usize].neighbors.len() - 1;
if let Some(ep) = self.entry_point {
let mut cur_ep = ep;
if level < self.max_layer {
for l in ((level + 1)..=self.max_layer).rev() {
cur_ep = self.greedy_closest(query, cur_ep, l);
}
}
let build_ef = self.effective_ef_construction(ef_construction, m);
let mut entry_points: Vec<u32> = vec![cur_ep];
let start_layer = level.min(self.max_layer);
for l in (0..=start_layer).rev() {
let candidates = self.search_layer_mut(query, &entry_points, build_ef, l);
let max_conn = if l == 0 { m0 } else { m };
let neighbors = self.select_neighbors(candidates, m);
for &(neighbor, dist) in &neighbors {
if l >= self.nodes[neighbor as usize].neighbors.len() {
continue;
}
self.update_connection(
neighbor,
existing_node,
dist,
l,
max_conn,
true,
);
}
entry_points.clear();
entry_points.extend(neighbors.iter().map(|&(n, _)| n));
self.nodes[existing_node as usize].neighbors[l] = neighbors;
}
}
return;
}
return;
}
let node_id = self.nodes.len() as u32;
let level = deterministic_level(self.build_seed, row_id, ml);
self.vectors.extend_from_slice(vector_bytes);
self.nodes.push(HnswNode {
neighbors: vec![Vec::new(); level + 1],
});
self.push_node_alive();
self.node_to_row_id.push(row_id);
self.row_id_to_node.insert(row_id, node_id);
self.unique_map_insert(node_id);
if node_id == 0 {
self.entry_point = Some(0);
self.max_layer = level;
return;
}
let ep = match self.entry_point {
Some(ep) => ep,
None => return,
};
let query_vec = bytes_to_f32_vec(vector_bytes);
let query = &query_vec;
let mut cur_ep = ep;
if level < self.max_layer {
for l in ((level + 1)..=self.max_layer).rev() {
cur_ep = self.greedy_closest(query, cur_ep, l);
}
}
let mut entry_points: Vec<u32> = Vec::with_capacity(m.max(1));
entry_points.push(cur_ep);
let build_ef = self.effective_ef_construction(ef_construction, m);
let start_layer = level.min(self.max_layer);
for l in (0..=start_layer).rev() {
let candidates = self.search_layer_mut(query, &entry_points, build_ef, l);
let max_conn = if l == 0 { m0 } else { m };
let neighbors = self.select_neighbors(candidates, m);
for &(neighbor, dist) in &neighbors {
if l >= self.nodes[neighbor as usize].neighbors.len() {
continue;
}
self.update_connection(neighbor, node_id, dist, l, max_conn, true);
}
entry_points.clear();
entry_points.extend(neighbors.iter().map(|&(n, _)| n));
self.nodes[node_id as usize].neighbors[l] = neighbors;
}
if level > self.max_layer {
self.max_layer = level;
self.entry_point = Some(node_id);
}
}
fn search(&self, query_bytes: &[u8], k: usize, ef_search: usize) -> Vec<(i64, f64)> {
if self.nodes.is_empty() {
return Vec::new();
}
let ep = match self.entry_point {
Some(ep) => ep,
None => return Vec::new(),
};
let query_vec = bytes_to_f32_vec(query_bytes);
let query = &query_vec;
let dynamic_ef_floor = if self.nodes.len() >= 1_000_000 {
768
} else if self.nodes.len() >= 500_000 {
640
} else if self.nodes.len() >= 100_000 {
512
} else {
0
};
let ef = ef_search.max(k).max(dynamic_ef_floor);
let mut cur_ep = ep;
for l in (1..=self.max_layer).rev() {
cur_ep = self.greedy_closest(query, cur_ep, l);
}
let mut results = self.search_layer(query, std::slice::from_ref(&cur_ep), ef, 0);
results.sort_unstable_by(|a, b| a.distance.total_cmp(&b.distance));
let metric = self.metric;
let mut exact_results: Vec<(i64, f64)> = results
.into_iter()
.filter(|e| !self.is_deleted(e.node))
.map(|e| {
let row_id = self.node_to_row_id[e.node as usize];
let offset = e.node as usize * self.dims_bytes;
let vector = &self.vectors[offset..offset + self.dims_bytes];
let final_dist = match metric {
HnswDistanceMetric::L2 => {
radixdb_core::vector::l2_distance_bytes(vector, query_bytes)
}
HnswDistanceMetric::Cosine => {
radixdb_core::vector::cosine_distance_bytes(vector, query_bytes)
}
HnswDistanceMetric::InnerProduct => {
radixdb_core::vector::ip_distance_bytes(vector, query_bytes)
}
};
(row_id, final_dist.unwrap_or(f64::INFINITY))
})
.collect();
exact_results.sort_unstable_by(|a, b| a.1.total_cmp(&b.1).then_with(|| a.0.cmp(&b.0)));
exact_results.truncate(k);
exact_results
}
#[inline]
fn effective_ef_construction(&self, requested: usize, m: usize) -> usize {
let n = self.nodes.len();
let floor = if n >= 1_000_000 {
m.saturating_mul(10)
} else if n >= 300_000 {
m.saturating_mul(8)
} else if n >= 100_000 {
m.saturating_mul(6)
} else {
requested
};
requested.max(floor).min(requested.saturating_mul(2))
}
#[cfg(feature = "parallel")]
#[inline]
fn parallel_seed_count(total: usize) -> usize {
let sqrt_seed = (total as f64).sqrt() as usize;
let frac_seed = total / 400;
sqrt_seed.max(frac_seed).clamp(256, 4_000).min(total)
}
#[cfg(feature = "parallel")]
#[inline]
fn parallel_batch_size(node_count: usize) -> usize {
if node_count < 100_000 {
512
} else if node_count < 400_000 {
1024
} else {
1536
}
}
#[cfg(feature = "parallel")]
fn insert_batch_parallel(
&mut self,
entries: &[(&[u8], i64)],
m: usize,
m0: usize,
ef_construction: usize,
ml: f64,
) {
const PARALLEL_THRESHOLD: usize = 5000;
if entries.len() < PARALLEL_THRESHOLD {
for &(vec_bytes, row_id) in entries {
self.insert(vec_bytes, row_id, m, m0, ef_construction, ml);
}
return;
}
let seed_count = Self::parallel_seed_count(entries.len());
for &(vec_bytes, row_id) in &entries[..seed_count] {
self.insert(vec_bytes, row_id, m, m0, ef_construction, ml);
}
let mut offset = seed_count;
while offset < entries.len() {
let batch_size = Self::parallel_batch_size(self.nodes.len());
let end = (offset + batch_size).min(entries.len());
self.insert_batch_inner(&entries[offset..end], m, m0, ef_construction, ml);
offset = end;
}
}
#[cfg(feature = "parallel")]
fn insert_batch_inner(
&mut self,
batch: &[(&[u8], i64)],
m: usize,
m0: usize,
ef_construction: usize,
ml: f64,
) {
use rayon::prelude::*;
let mut batch_nodes: Vec<(u32, usize)> = Vec::with_capacity(batch.len());
for &(vec_bytes, row_id) in batch {
if let Some(&existing_node) = self.row_id_to_node.get(&row_id) {
if self.is_deleted(existing_node) {
let offset = existing_node as usize * self.dims_bytes;
self.vectors[offset..offset + self.dims_bytes].copy_from_slice(vec_bytes);
self.clear_deleted(existing_node);
self.unique_map_insert(existing_node);
let level = self.nodes[existing_node as usize].neighbors.len() - 1;
if self.entry_point.is_some() {
batch_nodes.push((existing_node, level));
}
}
continue;
}
let node_id = self.nodes.len() as u32;
let level = deterministic_level(self.build_seed, row_id, ml);
self.vectors.extend_from_slice(vec_bytes);
self.nodes.push(HnswNode {
neighbors: vec![Vec::new(); level + 1],
});
self.push_node_alive();
self.node_to_row_id.push(row_id);
self.row_id_to_node.insert(row_id, node_id);
self.unique_map_insert(node_id);
if self.entry_point.is_none() {
self.entry_point = Some(node_id);
self.max_layer = level;
continue;
}
batch_nodes.push((node_id, level));
}
let ep = match self.entry_point {
Some(ep) => ep,
None => return,
};
let fast_bulk_mode = self.nodes.len() >= 300_000;
let dims = self.dims_bytes / 4;
let mut query_buf: Vec<f32> = Vec::with_capacity(dims);
let build_ef = self.effective_ef_construction(ef_construction, m);
let build_ef = if self.nodes.len() >= 100_000 {
build_ef.min(128)
} else {
build_ef
};
struct SearchTask {
node_id: u32,
level: usize,
entry_point: u32,
}
let mut tasks: Vec<SearchTask> = Vec::with_capacity(batch_nodes.len());
for &(node_id, level) in &batch_nodes {
let entry_point = if fast_bulk_mode {
ep
} else {
let mut entry_point = ep;
if self.max_layer > 0 {
let query = self.vector_f32(node_id);
for l in (1..=self.max_layer).rev() {
entry_point = self.greedy_closest(query, entry_point, l);
}
}
entry_point
};
tasks.push(SearchTask {
node_id,
level,
entry_point,
});
}
if !fast_bulk_mode {
for task in &tasks {
if task.level == 0 {
continue;
}
let node_id = task.node_id;
let level = task.level;
query_buf.clear();
query_buf.extend_from_slice(self.vector_f32(node_id));
let mut cur_ep = ep;
if level < self.max_layer {
for l in ((level + 1)..=self.max_layer).rev() {
cur_ep = self.greedy_closest(&query_buf, cur_ep, l);
}
}
let mut entry_points: Vec<u32> = Vec::with_capacity(m.max(1));
entry_points.push(cur_ep);
let start_layer = level.min(self.max_layer);
for l in (1..=start_layer).rev() {
let candidates = self.search_layer_mut(&query_buf, &entry_points, build_ef, l);
let max_conn = m;
let neighbors = self.select_neighbors(candidates, max_conn);
for &(neighbor, dist) in &neighbors {
if l >= self.nodes[neighbor as usize].neighbors.len() {
continue;
}
self.update_connection(neighbor, node_id, dist, l, max_conn, false);
}
entry_points.clear();
entry_points.extend(neighbors.iter().map(|&(n, _)| n));
self.nodes[node_id as usize].neighbors[l] = neighbors;
}
if level > self.max_layer {
self.max_layer = level;
self.entry_point = Some(node_id);
}
}
}
let batch_ef = build_ef;
let results: Vec<(u32, Vec<(u32, f32)>)>;
{
let nodes = &self.nodes;
let vectors = &self.vectors;
let dims_bytes = self.dims_bytes;
let metric = self.metric;
let deleted_bits = &self.deleted_bits;
results = tasks
.par_iter()
.map(|task| {
let start = task.node_id as usize * dims_bytes;
let query = as_f32_slice(&vectors[start..start + dims_bytes]);
let candidates = search_layer_shared(
nodes,
deleted_bits,
vectors,
dims_bytes,
metric,
query,
std::slice::from_ref(&task.entry_point),
batch_ef,
0,
);
let neighbors =
select_neighbors_shared(vectors, dims_bytes, metric, candidates, m);
(task.node_id, neighbors)
})
.collect();
}
for (node_id, neighbors) in results {
for &(nb, dist) in &neighbors {
if self.nodes[nb as usize].neighbors.is_empty() {
continue;
}
self.update_connection(nb, node_id, dist, 0, m0, false);
}
self.nodes[node_id as usize].neighbors[0] = neighbors;
}
}
fn serialize_graph(&self) -> std::result::Result<Vec<u8>, String> {
let node_count = self.nodes.len();
let dims_width = u32::try_from(self.dims_bytes)
.map_err(|_| "HNSW dimensions exceed persisted u32 width".to_string())?;
let node_width = u32::try_from(node_count)
.map_err(|_| "HNSW node count exceeds persisted u32 width".to_string())?;
let max_layer_width = u32::try_from(self.max_layer)
.map_err(|_| "HNSW max layer exceeds persisted u32 width".to_string())?;
let expected_vector_bytes = node_count
.checked_mul(self.dims_bytes)
.ok_or_else(|| "HNSW vector byte count overflow".to_string())?;
if self.vectors.len() != expected_vector_bytes || self.node_to_row_id.len() != node_count {
return Err("HNSW graph authorities have inconsistent lengths".to_string());
}
let mut seen_row_ids = ahash::AHashSet::with_capacity(node_count);
if self
.node_to_row_id
.iter()
.any(|row_id| !seen_row_ids.insert(*row_id))
{
return Err("HNSW graph contains duplicate row IDs".to_string());
}
let estimated = HNSW_GRAPH_HEADER_LEN
.checked_add(self.vectors.len())
.and_then(|size| size.checked_add(node_count.checked_mul(8)?))
.and_then(|size| size.checked_add(node_count.checked_mul(20)?))
.ok_or_else(|| "HNSW serialized size overflow".to_string())?;
if estimated > HNSW_MAX_GRAPH_BYTES {
return Err("HNSW graph exceeds persisted size budget".to_string());
}
let mut buf = Vec::with_capacity(estimated);
buf.extend_from_slice(b"HNSW");
buf.extend_from_slice(&2u32.to_le_bytes()); buf.push(self.metric.as_u8());
buf.extend_from_slice(&dims_width.to_le_bytes());
buf.extend_from_slice(&node_width.to_le_bytes());
buf.extend_from_slice(&max_layer_width.to_le_bytes());
buf.extend_from_slice(&self.entry_point.unwrap_or(u32::MAX).to_le_bytes());
buf.extend_from_slice(&self.vectors);
for &rid in &self.node_to_row_id {
buf.extend_from_slice(&rid.to_le_bytes());
}
for (i, node) in self.nodes.iter().enumerate() {
buf.push(self.is_deleted(i as u32) as u8);
let layer_count = u8::try_from(node.neighbors.len())
.map_err(|_| format!("HNSW node {i} exceeds persisted layer width"))?;
buf.push(layer_count);
for layer in &node.neighbors {
let neighbor_count = u16::try_from(layer.len())
.map_err(|_| format!("HNSW node {i} exceeds persisted neighbor width"))?;
buf.extend_from_slice(&neighbor_count.to_le_bytes());
for &(nbr, dist) in layer {
buf.extend_from_slice(&nbr.to_le_bytes());
buf.extend_from_slice(&dist.to_le_bytes());
}
}
}
if buf.len() > HNSW_MAX_GRAPH_BYTES {
return Err("HNSW graph exceeds persisted size budget".to_string());
}
Ok(buf)
}
fn deserialize_graph(data: &[u8]) -> std::result::Result<Self, String> {
if data.len() < HNSW_GRAPH_HEADER_LEN {
return Err("HNSW data too short for header".to_string());
}
if data.len() > HNSW_MAX_GRAPH_BYTES {
return Err("HNSW graph exceeds persisted size budget".to_string());
}
if &data[0..4] != b"HNSW" {
return Err("Invalid HNSW magic bytes".to_string());
}
let version = u32::from_le_bytes([data[4], data[5], data[6], data[7]]);
if version != 2 {
return Err(format!("Unsupported HNSW version: {}", version));
}
let metric = HnswDistanceMetric::from_u8(data[8])
.ok_or_else(|| format!("Invalid HNSW metric: {}", data[8]))?;
let dims_bytes = u32::from_le_bytes([data[9], data[10], data[11], data[12]]) as usize;
let node_count = u32::from_le_bytes([data[13], data[14], data[15], data[16]]) as usize;
let max_layer = u32::from_le_bytes([data[17], data[18], data[19], data[20]]) as usize;
if dims_bytes == 0 || !dims_bytes.is_multiple_of(std::mem::size_of::<f32>()) {
return Err("HNSW vector width is zero or not f32-aligned".to_string());
}
if max_layer > u8::MAX as usize {
return Err("HNSW max layer exceeds node layer width".to_string());
}
let ep_raw = u32::from_le_bytes([data[21], data[22], data[23], data[24]]);
let entry_point = if ep_raw == u32::MAX {
None
} else {
Some(ep_raw)
};
let mut pos = HNSW_GRAPH_HEADER_LEN;
let vec_size = node_count
.checked_mul(dims_bytes)
.ok_or_else(|| "HNSW vector byte count overflow".to_string())?;
let vector_end = pos
.checked_add(vec_size)
.ok_or_else(|| "HNSW vector end offset overflow".to_string())?;
if vector_end > data.len() {
return Err("HNSW data truncated at vectors".to_string());
}
let vectors = data[pos..vector_end].to_vec();
pos = vector_end;
let rid_size = node_count
.checked_mul(8)
.ok_or_else(|| "HNSW row ID byte count overflow".to_string())?;
let row_ids_end = pos
.checked_add(rid_size)
.ok_or_else(|| "HNSW row ID end offset overflow".to_string())?;
let minimum_node_bytes = node_count
.checked_mul(2)
.ok_or_else(|| "HNSW node header byte count overflow".to_string())?;
if row_ids_end > data.len() || data.len().saturating_sub(row_ids_end) < minimum_node_bytes {
return Err("HNSW data truncated at row_ids".to_string());
}
let mut node_to_row_id = Vec::with_capacity(node_count);
let mut row_id_to_node = ahash::AHashMap::with_capacity(node_count);
for i in 0..node_count {
let off = pos
.checked_add(
i.checked_mul(8)
.ok_or_else(|| "HNSW row ID offset overflow".to_string())?,
)
.ok_or_else(|| "HNSW row ID offset overflow".to_string())?;
let rid = i64::from_le_bytes([
data[off],
data[off + 1],
data[off + 2],
data[off + 3],
data[off + 4],
data[off + 5],
data[off + 6],
data[off + 7],
]);
if row_id_to_node.contains_key(&rid) {
return Err(format!("HNSW graph contains duplicate row ID {rid}"));
}
node_to_row_id.push(rid);
row_id_to_node.insert(rid, i as u32);
}
pos = row_ids_end;
let mut nodes = Vec::with_capacity(node_count);
let deleted_words = node_count
.checked_add(63)
.ok_or_else(|| "HNSW deleted bitset size overflow".to_string())?
>> 6;
let mut deleted_bits = vec![0u64; deleted_words];
for node_idx in 0..node_count {
let node_header_end = pos
.checked_add(2)
.ok_or_else(|| "HNSW node header offset overflow".to_string())?;
if node_header_end > data.len() {
return Err("HNSW data truncated at node header".to_string());
}
let deleted = data[pos] != 0;
if deleted {
deleted_bits[node_idx >> 6] |= 1u64 << (node_idx & 63);
}
let num_layers = data[pos + 1] as usize;
if num_layers > max_layer.saturating_add(1) {
return Err(format!(
"HNSW node {node_idx} has {num_layers} layers above max layer {max_layer}"
));
}
pos = node_header_end;
let mut neighbors = Vec::with_capacity(num_layers);
for _ in 0..num_layers {
let layer_header_end = pos
.checked_add(2)
.ok_or_else(|| "HNSW layer header offset overflow".to_string())?;
if layer_header_end > data.len() {
return Err("HNSW data truncated at layer header".to_string());
}
let count = u16::from_le_bytes([data[pos], data[pos + 1]]) as usize;
pos = layer_header_end;
let neighbor_bytes = count
.checked_mul(8)
.ok_or_else(|| "HNSW neighbor byte count overflow".to_string())?;
let neighbor_end = pos
.checked_add(neighbor_bytes)
.ok_or_else(|| "HNSW neighbor end offset overflow".to_string())?;
if neighbor_end > data.len() {
return Err("HNSW data truncated at neighbor list".to_string());
}
let mut layer = Vec::with_capacity(count);
for j in 0..count {
let off = pos
.checked_add(
j.checked_mul(8)
.ok_or_else(|| "HNSW neighbor offset overflow".to_string())?,
)
.ok_or_else(|| "HNSW neighbor offset overflow".to_string())?;
let nid = u32::from_le_bytes([
data[off],
data[off + 1],
data[off + 2],
data[off + 3],
]);
let dist = f32::from_le_bytes([
data[off + 4],
data[off + 5],
data[off + 6],
data[off + 7],
]);
layer.push((nid, dist));
}
pos = neighbor_end;
neighbors.push(layer);
}
nodes.push(HnswNode { neighbors });
}
if pos != data.len() {
return Err(format!(
"HNSW graph has {} trailing bytes",
data.len() - pos
));
}
let node_count = nodes.len() as u32;
for (i, node) in nodes.iter().enumerate() {
for (l, layer) in node.neighbors.iter().enumerate() {
for &(nid, _) in layer {
if nid >= node_count {
return Err(format!(
"HNSW corrupted: node {} layer {} has neighbor {} but only {} nodes exist",
i, l, nid, node_count
));
}
}
}
}
if let Some(ep) = entry_point {
if ep >= node_count {
return Err(format!(
"HNSW corrupted: entry_point {} but only {} nodes exist",
ep, node_count
));
}
}
let inner = Self {
nodes,
entry_point,
max_layer,
vectors,
dims_bytes,
node_to_row_id,
row_id_to_node,
metric,
scratch: SearchScratch::new(),
deleted_bits,
unique_map: None,
build_seed: 0,
};
Ok(inner)
}
}
pub fn default_m_for_dims(dims: usize) -> usize {
if dims >= 256 {
48
} else if dims >= 64 {
32
} else {
16
}
}
pub fn default_ef_construction(m: usize) -> usize {
if m >= 48 {
256
} else if m >= 32 {
200
} else {
128
}
}
pub fn default_ef_search(m: usize) -> usize {
if m >= 48 {
256
} else if m >= 32 {
200
} else {
128
}
}
fn deterministic_level(build_seed: u64, row_id: i64, ml: f64) -> usize {
let mut sample = build_seed ^ (row_id as u64).wrapping_mul(0x9e37_79b9_7f4a_7c15);
sample = sample.wrapping_add(0x9e37_79b9_7f4a_7c15);
sample = (sample ^ (sample >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
sample = (sample ^ (sample >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
sample ^= sample >> 31;
let unit = (sample >> 11) as f64 * (1.0 / ((1u64 << 53) as f64));
let r = unit.max(f64::MIN_POSITIVE);
(-r.ln() * ml).floor() as usize
}
fn stable_hnsw_build_seed(
name: &str,
table_name: &str,
column_name: &str,
column_id: i32,
dims: usize,
m: usize,
metric: HnswDistanceMetric,
) -> u64 {
let mut hash = 0xcbf2_9ce4_8422_2325u64;
for bytes in [
name.as_bytes(),
table_name.as_bytes(),
column_name.as_bytes(),
] {
for byte in bytes {
hash ^= u64::from(*byte);
hash = hash.wrapping_mul(0x100_0000_01b3);
}
hash ^= 0xff;
hash = hash.wrapping_mul(0x100_0000_01b3);
}
for byte in column_id
.to_le_bytes()
.into_iter()
.chain(dims.to_le_bytes())
.chain(m.to_le_bytes())
.chain([metric.as_u8()])
{
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x100_0000_01b3);
}
hash
}
#[cfg(feature = "parallel")]
#[allow(clippy::too_many_arguments)]
fn search_layer_shared(
nodes: &[HnswNode],
deleted_bits: &[u64],
vectors: &[u8],
dims_bytes: usize,
metric: HnswDistanceMetric,
query: &[f32],
entry_points: &[u32],
ef: usize,
layer: usize,
) -> Vec<MaxEntry> {
BUILD_SCRATCH.with(|cell| {
let mut scratch = cell.borrow_mut();
scratch.reset(nodes.len());
let dims = dims_bytes / 4;
let vectors_ptr = vectors.as_ptr();
let visited_ptr = scratch.visited.as_mut_ptr();
let deleted_ptr = deleted_bits.as_ptr();
for &ep in entry_points {
if (ep as usize) >= nodes.len() {
continue;
}
let n_idx = ep as usize;
unsafe {
*visited_ptr.add(n_idx >> 6) |= 1u64 << (n_idx & 63);
}
let v_ptr = unsafe { vectors_ptr.add(n_idx * dims_bytes) as *const f32 };
let v_slice = unsafe { std::slice::from_raw_parts(v_ptr, dims) };
let d = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(v_slice, query),
HnswDistanceMetric::Cosine => cosine_distance_f32(v_slice, query),
HnswDistanceMetric::InnerProduct => ip_distance_f32(v_slice, query),
};
scratch.candidates.push(MinEntry {
distance: d,
node: ep,
});
scratch.result.push(MaxEntry {
distance: d,
node: ep,
});
}
let mut farthest_dist = scratch.result.peek().map_or(f32::MAX, |e| e.distance);
let mut result_len = scratch.result.len();
let nodes_ptr = nodes.as_ptr();
let nodes_len = nodes.len();
while let Some(MinEntry {
distance: c_dist,
node: c_id,
}) = scratch.candidates.pop()
{
if c_dist > farthest_dist && result_len >= ef {
break;
}
if let Some(next) = scratch.candidates.peek() {
let next_id = next.node as usize;
if next_id < nodes_len {
let next_node_ptr = unsafe { nodes_ptr.add(next_id) } as *const u8;
prefetch_read(next_node_ptr);
}
}
let node = unsafe { &*nodes_ptr.add(c_id as usize) };
if layer < node.neighbors.len() {
let neighbors = &node.neighbors[layer];
let nlen = neighbors.len();
let nptr = neighbors.as_ptr();
let mut ni = 0usize;
while ni < nlen {
let (neighbor, _) = unsafe { *nptr.add(ni) };
ni += 1;
let n_idx = neighbor as usize;
if ni < nlen {
let next_nb = unsafe { (*nptr.add(ni)).0 } as usize;
let next_v_word = next_nb >> 6;
prefetch_read(unsafe { visited_ptr.add(next_v_word) } as *const u8);
}
let v_word_idx = n_idx >> 6;
let v_bit_mask = 1u64 << (n_idx & 63);
let v_word_ptr = unsafe { visited_ptr.add(v_word_idx) };
let v_word = unsafe { *v_word_ptr };
if (v_word & v_bit_mask) != 0 {
continue;
}
unsafe {
*v_word_ptr = v_word | v_bit_mask;
}
let is_deleted = unsafe { (*deleted_ptr.add(v_word_idx)) & v_bit_mask != 0 };
if is_deleted {
continue;
}
if ni < nlen {
let next_nb = unsafe { (*nptr.add(ni)).0 } as usize;
let vec_addr = unsafe { vectors_ptr.add(next_nb * dims_bytes) };
prefetch_read(vec_addr);
}
let v_ptr = unsafe { vectors_ptr.add(n_idx * dims_bytes) as *const f32 };
let v_slice = unsafe { std::slice::from_raw_parts(v_ptr, dims) };
let d = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(v_slice, query),
HnswDistanceMetric::Cosine => cosine_distance_f32(v_slice, query),
HnswDistanceMetric::InnerProduct => ip_distance_f32(v_slice, query),
};
if d < farthest_dist || result_len < ef {
scratch.candidates.push(MinEntry {
distance: d,
node: neighbor,
});
scratch.result.push(MaxEntry {
distance: d,
node: neighbor,
});
result_len += 1;
if result_len > ef {
scratch.result.pop();
result_len -= 1;
farthest_dist = scratch.result.peek().map_or(f32::MAX, |e| e.distance);
} else {
farthest_dist = farthest_dist.max(d);
}
}
}
}
}
scratch.result.drain().collect()
})
}
#[cfg(feature = "parallel")]
fn select_neighbors_shared(
vectors: &[u8],
dims_bytes: usize,
metric: HnswDistanceMetric,
mut candidates: Vec<MaxEntry>,
m: usize,
) -> Vec<(u32, f32)> {
if candidates.is_empty() || m == 0 {
return Vec::new();
}
candidates.sort_unstable_by(|a, b| a.distance.total_cmp(&b.distance));
let mut selected: Vec<(u32, f32)> = Vec::with_capacity(m);
let mut pruned: Vec<(u32, f32)> = Vec::with_capacity(candidates.len());
for entry in &candidates {
if selected.len() >= m {
break;
}
let dist_to_query = entry.distance;
let start_a = entry.node as usize * dims_bytes;
let va = as_f32_slice(&vectors[start_a..start_a + dims_bytes]);
let mut is_diverse = true;
for &(sel_node, _) in &selected {
let start_b = sel_node as usize * dims_bytes;
let vb = as_f32_slice(&vectors[start_b..start_b + dims_bytes]);
let dist_to_selected = match metric {
HnswDistanceMetric::L2 => l2_distance_sq_f32(va, vb),
HnswDistanceMetric::Cosine => cosine_distance_f32(va, vb),
HnswDistanceMetric::InnerProduct => ip_distance_f32(va, vb),
};
if dist_to_selected < dist_to_query {
is_diverse = false;
break;
}
}
if is_diverse {
selected.push((entry.node, entry.distance));
} else {
pruned.push((entry.node, entry.distance));
}
}
for entry in pruned {
if selected.len() >= m {
break;
}
selected.push(entry);
}
selected
}
pub struct HnswIndex {
inner: RwLock<HnswInner>,
name: String,
table_name: String,
column_ids: Vec<i32>,
column_names: Vec<String>,
data_types: Vec<DataType>,
dims: usize,
m: usize,
m0: usize,
ef_construction: usize,
ef_search: usize,
ml: f64,
metric: HnswDistanceMetric,
is_unique: bool,
build_seed: u64,
}
impl HnswIndex {
#[allow(clippy::too_many_arguments)]
pub fn new(
name: String,
table_name: String,
column_name: String,
column_id: i32,
dims: usize,
m: usize,
ef_construction: usize,
ef_search: usize,
metric: HnswDistanceMetric,
) -> Result<Self> {
if name.is_empty() || table_name.is_empty() || column_name.is_empty() || column_id < 0 {
return Err(Error::invalid_argument("invalid HNSW index identity"));
}
if dims == 0 || dims.checked_mul(std::mem::size_of::<f32>()).is_none() {
return Err(Error::invalid_argument("invalid HNSW dimensions"));
}
if !(2..=HNSW_MAX_M).contains(&m) {
return Err(Error::invalid_argument(format!(
"HNSW m must be in 2..={HNSW_MAX_M}"
)));
}
if ef_construction == 0
|| ef_search == 0
|| u16::try_from(ef_construction).is_err()
|| u16::try_from(ef_search).is_err()
{
return Err(Error::invalid_argument(
"HNSW ef_construction/ef_search must fit non-zero u16",
));
}
let m0 = m
.checked_mul(2)
.ok_or_else(|| Error::invalid_argument("HNSW m0 overflow"))?;
let ml = 1.0 / (m as f64).ln();
let build_seed =
stable_hnsw_build_seed(&name, &table_name, &column_name, column_id, dims, m, metric);
Ok(Self {
inner: RwLock::new(HnswInner::with_build_seed(dims, metric, build_seed)),
name,
table_name,
column_ids: vec![column_id],
column_names: vec![column_name],
data_types: vec![DataType::Vector],
dims,
m,
m0,
ef_construction,
ef_search,
ml,
metric,
is_unique: false,
build_seed,
})
}
pub fn set_unique(&mut self, unique: bool) -> Result<()> {
if !unique {
self.inner.write().unique_map = None;
self.is_unique = false;
return Ok(());
}
let mut inner = self.inner.write();
let mut exact = ahash::AHashMap::<Vec<u8>, i64>::with_capacity(inner.nodes.len());
let mut hashed = ahash::AHashMap::<u64, Vec<i64>>::with_capacity(inner.nodes.len());
for (node_idx, &row_id) in inner.node_to_row_id.iter().enumerate() {
if inner.is_deleted(node_idx as u32) {
continue;
}
let offset = node_idx.checked_mul(inner.dims_bytes).ok_or_else(|| {
radixdb_core::Error::invalid_argument("HNSW vector offset overflow")
})?;
let end = offset.checked_add(inner.dims_bytes).ok_or_else(|| {
radixdb_core::Error::invalid_argument("HNSW vector end offset overflow")
})?;
let bytes = inner.vectors.get(offset..end).ok_or_else(|| {
radixdb_core::Error::invalid_argument("HNSW vector data truncated")
})?;
if let Some(existing_row_id) = exact.insert(bytes.to_vec(), row_id) {
if existing_row_id != row_id {
return Err(radixdb_core::Error::unique_constraint(
&self.name,
self.column_names.join(", "),
format!("<vector({} dims)>", self.dims),
));
}
}
hashed
.entry(HnswInner::hash_vec_bytes(bytes))
.or_default()
.push(row_id);
}
inner.unique_map = Some(hashed);
drop(inner);
self.is_unique = true;
Ok(())
}
pub fn distance_metric(&self) -> HnswDistanceMetric {
self.metric
}
pub fn params(&self) -> (usize, usize, usize, HnswDistanceMetric) {
(self.m, self.ef_construction, self.ef_search, self.metric)
}
pub fn search_nearest(
&self,
query_bytes: &[u8],
k: usize,
ef_search: usize,
) -> Vec<(i64, f64)> {
if query_bytes.len() != self.dims * 4 {
return Vec::new();
}
let inner = self.inner.read();
inner.search(query_bytes, k, ef_search)
}
fn extract_vector_bytes(value: &Value) -> Option<&[u8]> {
match value {
Value::Extension(data) if data.first() == Some(&(DataType::Vector as u8)) => {
Some(&data[1..])
}
_ => None,
}
}
fn validate_vector_value<'a>(&self, values: &'a [Value]) -> Result<&'a [u8]> {
let value = values.first().ok_or_else(|| {
radixdb_core::Error::invalid_argument("HNSW index requires one VECTOR value")
})?;
let bytes = Self::extract_vector_bytes(value).ok_or_else(|| {
radixdb_core::Error::invalid_argument("HNSW index requires a VECTOR value")
})?;
let got = u16::try_from(bytes.len() / 4).unwrap_or(u16::MAX);
let expected = u16::try_from(self.dims).unwrap_or(u16::MAX);
if bytes.len() != self.dims * 4 {
return Err(radixdb_core::Error::VectorDimensionMismatch { expected, got });
}
Ok(bytes)
}
fn find_exact_duplicate_in_inner(
inner: &HnswInner,
vec_bytes: &[u8],
exclude_row_id: i64,
ignored_row_ids: Option<&I64Set>,
) -> Option<i64> {
if let Some(ref map) = inner.unique_map {
let hash = HnswInner::hash_vec_bytes(vec_bytes);
if let Some(row_ids) = map.get(&hash) {
for &candidate_row_id in row_ids {
if candidate_row_id == exclude_row_id {
continue;
}
if ignored_row_ids.is_some_and(|ignored| ignored.contains(candidate_row_id)) {
continue;
}
if let Some(&node_id) = inner.row_id_to_node.get(&candidate_row_id) {
if !inner.is_deleted(node_id) {
let offset = node_id as usize * inner.dims_bytes;
if let Some(existing_bytes) =
inner.vectors.get(offset..offset + inner.dims_bytes)
{
if existing_bytes == vec_bytes {
return Some(candidate_row_id);
}
}
}
}
}
}
return None;
}
for (node_idx, &existing_row_id) in inner.node_to_row_id.iter().enumerate() {
if existing_row_id == exclude_row_id {
continue;
}
if ignored_row_ids.is_some_and(|ignored| ignored.contains(existing_row_id)) {
continue;
}
if inner.is_deleted(node_idx as u32) {
continue;
}
let offset = node_idx * inner.dims_bytes;
if let Some(existing_bytes) = inner.vectors.get(offset..offset + inner.dims_bytes) {
if existing_bytes == vec_bytes {
return Some(existing_row_id);
}
}
}
None
}
pub fn find_exact_duplicate(
&self,
value: &Value,
exclude_row_id: i64,
ignored_row_ids: Option<&I64Set>,
) -> Option<i64> {
let vec_bytes = Self::extract_vector_bytes(value)?;
if vec_bytes.len() != self.dims * 4 {
return None;
}
let inner = self.inner.read();
Self::find_exact_duplicate_in_inner(&inner, vec_bytes, exclude_row_id, ignored_row_ids)
}
fn insert_prepared(&self, inner: &mut HnswInner, prepared: &[(&[u8], i64)]) {
#[cfg(feature = "parallel")]
{
inner.insert_batch_parallel(prepared, self.m, self.m0, self.ef_construction, self.ml);
}
#[cfg(not(feature = "parallel"))]
{
for &(vec_bytes, row_id) in prepared {
inner.insert(
vec_bytes,
row_id,
self.m,
self.m0,
self.ef_construction,
self.ml,
);
}
}
}
pub fn save_graph(&self, path: &std::path::Path) -> std::io::Result<()> {
let data = self.serialize_graph_bytes().map_err(|error| {
std::io::Error::new(std::io::ErrorKind::InvalidData, error.to_string())
})?;
let tmp_path = path.with_extension("bin.tmp");
std::fs::write(&tmp_path, data)?;
std::fs::rename(&tmp_path, path)
}
#[allow(clippy::too_many_arguments)]
pub fn load_graph(
path: &std::path::Path,
name: String,
table_name: String,
column_name: String,
column_id: i32,
dims: usize,
m: usize,
ef_construction: usize,
ef_search: usize,
) -> std::io::Result<Option<Self>> {
if !path.exists() {
return Ok(None);
}
let file_len = std::fs::metadata(path)?.len();
if file_len > HNSW_MAX_GRAPH_BYTES as u64 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW graph exceeds persisted size budget",
));
}
let data = std::fs::read(path)?;
if data.len() < HNSW_IDENTITY_HEADER_LEN || &data[..4] != HNSW_IDENTITY_MAGIC {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"missing HNSW identity envelope",
));
}
let read_u32 = |offset: usize| {
u32::from_le_bytes(data[offset..offset + 4].try_into().expect("bounded header"))
};
let version = read_u32(4);
if version != HNSW_IDENTITY_VERSION {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("unsupported HNSW identity version {version}"),
));
}
let name_len = read_u32(8) as usize;
let table_len = read_u32(12) as usize;
let column_len = read_u32(16) as usize;
let stored_column_id = i32::from_le_bytes(data[20..24].try_into().unwrap());
let stored_dims = read_u32(24) as usize;
let stored_m = read_u32(28) as usize;
let stored_ef_construction = read_u32(32) as usize;
let stored_ef_search = read_u32(36) as usize;
let stored_metric = HnswDistanceMetric::from_u8(data[40]).ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
"invalid HNSW identity metric",
)
})?;
let graph_len = u64::from_le_bytes(data[41..49].try_into().unwrap());
let identity_len = name_len
.checked_add(table_len)
.and_then(|size| size.checked_add(column_len))
.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW identity length overflow",
)
})?;
let graph_start = HNSW_IDENTITY_HEADER_LEN
.checked_add(identity_len)
.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW identity end offset overflow",
)
})?;
let graph_len = usize::try_from(graph_len).map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW graph length does not fit this platform",
)
})?;
let graph_end = graph_start.checked_add(graph_len).ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW graph end offset overflow",
)
})?;
if graph_end != data.len() {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW identity envelope length mismatch",
));
}
let name_end = HNSW_IDENTITY_HEADER_LEN + name_len;
let table_end = name_end + table_len;
let column_end = table_end + column_len;
let stored_name =
std::str::from_utf8(&data[HNSW_IDENTITY_HEADER_LEN..name_end]).map_err(|_| {
std::io::Error::new(std::io::ErrorKind::InvalidData, "invalid HNSW name")
})?;
let stored_table = std::str::from_utf8(&data[name_end..table_end]).map_err(|_| {
std::io::Error::new(std::io::ErrorKind::InvalidData, "invalid HNSW table name")
})?;
let stored_column = std::str::from_utf8(&data[table_end..column_end]).map_err(|_| {
std::io::Error::new(std::io::ErrorKind::InvalidData, "invalid HNSW column name")
})?;
if stored_name != name
|| stored_table != table_name
|| stored_column != column_name
|| stored_column_id != column_id
|| stored_dims != dims
|| stored_m != m
|| stored_ef_construction != ef_construction
|| stored_ef_search != ef_search
{
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW graph identity does not match requested index",
));
}
match HnswInner::deserialize_graph(&data[graph_start..graph_end]) {
Ok(mut inner) => {
let expected_dims_bytes = dims.checked_mul(4).ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"HNSW expected dimensions overflow",
)
})?;
if inner.dims_bytes != expected_dims_bytes {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"HNSW graph dimension mismatch: file has {} bytes/vector ({} dims) \
but schema expects {} bytes/vector ({} dims)",
inner.dims_bytes,
inner.dims_bytes / 4,
expected_dims_bytes,
dims,
),
));
}
if inner.metric != stored_metric {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"HNSW identity metric does not match graph metric",
));
}
if !(2..=HNSW_MAX_M).contains(&m)
|| ef_construction == 0
|| ef_search == 0
|| u16::try_from(ef_construction).is_err()
|| u16::try_from(ef_search).is_err()
{
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"invalid HNSW construction parameters",
));
}
let metric = inner.metric;
let ml = 1.0 / (m as f64).ln();
let m0 = m.checked_mul(2).ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::InvalidInput, "HNSW m0 overflow")
})?;
let build_seed = stable_hnsw_build_seed(
&name,
&table_name,
&column_name,
column_id,
dims,
m,
metric,
);
inner.build_seed = build_seed;
Ok(Some(Self {
inner: RwLock::new(inner),
name,
table_name,
column_ids: vec![column_id],
column_names: vec![column_name],
data_types: vec![DataType::Vector],
dims,
m,
m0,
ef_construction,
ef_search,
ml,
metric,
is_unique: false,
build_seed,
}))
}
Err(e) => Err(std::io::Error::new(std::io::ErrorKind::InvalidData, e)),
}
}
pub fn serialize_graph_bytes(&self) -> Result<Vec<u8>> {
let inner = self.inner.read();
let graph = inner.serialize_graph().map_err(Error::invalid_argument)?;
let name_len = u32::try_from(self.name.len())
.map_err(|_| Error::invalid_argument("HNSW index name is too long"))?;
let table_len = u32::try_from(self.table_name.len())
.map_err(|_| Error::invalid_argument("HNSW table name is too long"))?;
let column_name = &self.column_names[0];
let column_len = u32::try_from(column_name.len())
.map_err(|_| Error::invalid_argument("HNSW column name is too long"))?;
let dims = u32::try_from(self.dims)
.map_err(|_| Error::invalid_argument("HNSW dimensions exceed persisted width"))?;
let m = u32::try_from(self.m)
.map_err(|_| Error::invalid_argument("HNSW m exceeds persisted width"))?;
let ef_construction = u32::try_from(self.ef_construction)
.map_err(|_| Error::invalid_argument("HNSW ef_construction exceeds persisted width"))?;
let ef_search = u32::try_from(self.ef_search)
.map_err(|_| Error::invalid_argument("HNSW ef_search exceeds persisted width"))?;
let graph_len = u64::try_from(graph.len())
.map_err(|_| Error::invalid_argument("HNSW graph exceeds persisted width"))?;
let capacity = HNSW_IDENTITY_HEADER_LEN
.checked_add(self.name.len())
.and_then(|size| size.checked_add(self.table_name.len()))
.and_then(|size| size.checked_add(column_name.len()))
.and_then(|size| size.checked_add(graph.len()))
.ok_or_else(|| Error::invalid_argument("HNSW envelope size overflow"))?;
if capacity > HNSW_MAX_GRAPH_BYTES {
return Err(Error::invalid_argument(
"HNSW identity envelope exceeds persisted size budget",
));
}
let mut encoded = Vec::with_capacity(capacity);
encoded.extend_from_slice(HNSW_IDENTITY_MAGIC);
encoded.extend_from_slice(&HNSW_IDENTITY_VERSION.to_le_bytes());
encoded.extend_from_slice(&name_len.to_le_bytes());
encoded.extend_from_slice(&table_len.to_le_bytes());
encoded.extend_from_slice(&column_len.to_le_bytes());
encoded.extend_from_slice(&self.column_ids[0].to_le_bytes());
encoded.extend_from_slice(&dims.to_le_bytes());
encoded.extend_from_slice(&m.to_le_bytes());
encoded.extend_from_slice(&ef_construction.to_le_bytes());
encoded.extend_from_slice(&ef_search.to_le_bytes());
encoded.push(self.metric.as_u8());
encoded.extend_from_slice(&graph_len.to_le_bytes());
encoded.extend_from_slice(self.name.as_bytes());
encoded.extend_from_slice(self.table_name.as_bytes());
encoded.extend_from_slice(column_name.as_bytes());
encoded.extend_from_slice(&graph);
Ok(encoded)
}
pub fn build_seed(&self) -> u64 {
self.build_seed
}
pub fn node_count(&self) -> usize {
let inner = self.inner.read();
inner.nodes.len()
}
}
impl Index for HnswIndex {
fn name(&self) -> &str {
&self.name
}
fn table_name(&self) -> &str {
&self.table_name
}
fn build(&mut self) -> Result<()> {
Ok(())
}
fn add(&self, values: &[Value], row_id: i64, _ref_id: i64) -> Result<()> {
let vec_bytes = self.validate_vector_value(values)?;
let mut inner = self.inner.write();
if self.is_unique
&& Self::find_exact_duplicate_in_inner(&inner, vec_bytes, row_id, None).is_some()
{
return Err(radixdb_core::Error::unique_constraint(
&self.name,
self.column_names.join(", "),
format!("<vector({} dims)>", self.dims),
));
}
inner.insert(
vec_bytes,
row_id,
self.m,
self.m0,
self.ef_construction,
self.ml,
);
Ok(())
}
fn add_batch(&self, entries: &I64Map<Vec<Value>>) -> Result<()> {
let mut inner = self.inner.write();
let dims_bytes = inner.dims_bytes;
let expected_vec_len = self.dims * 4;
let mut prepared: Vec<(&[u8], i64)> = Vec::with_capacity(entries.len());
for (row_id, values) in entries.iter() {
let vec_bytes = self.validate_vector_value(values)?;
debug_assert_eq!(vec_bytes.len(), expected_vec_len);
prepared.push((vec_bytes, row_id));
}
if self.is_unique {
let mut seen: ahash::AHashMap<&[u8], i64> =
ahash::AHashMap::with_capacity(prepared.len());
for &(vec_bytes, row_id) in &prepared {
if let Some(&existing_row_id) = seen.get(vec_bytes) {
if existing_row_id != row_id {
return Err(radixdb_core::Error::unique_constraint(
&self.name,
self.column_names.join(", "),
format!("<vector({} dims)>", self.dims),
));
}
} else {
seen.insert(vec_bytes, row_id);
}
if Self::find_exact_duplicate_in_inner(&inner, vec_bytes, row_id, None).is_some() {
return Err(radixdb_core::Error::unique_constraint(
&self.name,
self.column_names.join(", "),
format!("<vector({} dims)>", self.dims),
));
}
}
}
inner.vectors.reserve(prepared.len() * dims_bytes);
inner.nodes.reserve(prepared.len());
inner.node_to_row_id.reserve(prepared.len());
inner.row_id_to_node.reserve(prepared.len());
self.insert_prepared(&mut inner, &prepared);
Ok(())
}
fn remove(&self, _values: &[Value], row_id: i64, _ref_id: i64) -> Result<()> {
let mut inner = self.inner.write();
if let Some(&node_id) = inner.row_id_to_node.get(&row_id) {
inner.unique_map_remove(node_id);
inner.set_deleted(node_id);
}
Ok(())
}
fn remove_batch(&self, entries: &I64Map<Vec<Value>>) -> Result<()> {
let mut inner = self.inner.write();
for row_id in entries.keys() {
if let Some(&node_id) = inner.row_id_to_node.get(&row_id) {
inner.unique_map_remove(node_id);
inner.set_deleted(node_id);
}
}
Ok(())
}
fn add_batch_slice(&self, entries: &[(i64, &[Value])]) -> Result<()> {
let mut inner = self.inner.write();
let dims_bytes = inner.dims_bytes;
let expected_vec_len = self.dims * 4;
let mut prepared: Vec<(&[u8], i64)> = Vec::with_capacity(entries.len());
for &(row_id, values) in entries {
let vec_bytes = self.validate_vector_value(values)?;
debug_assert_eq!(vec_bytes.len(), expected_vec_len);
prepared.push((vec_bytes, row_id));
}
if self.is_unique {
let mut seen: ahash::AHashMap<&[u8], i64> =
ahash::AHashMap::with_capacity(prepared.len());
for &(vec_bytes, row_id) in &prepared {
if let Some(&existing_row_id) = seen.get(vec_bytes) {
if existing_row_id != row_id {
return Err(radixdb_core::Error::unique_constraint(
&self.name,
self.column_names.join(", "),
format!("<vector({} dims)>", self.dims),
));
}
} else {
seen.insert(vec_bytes, row_id);
}
if Self::find_exact_duplicate_in_inner(&inner, vec_bytes, row_id, None).is_some() {
return Err(radixdb_core::Error::unique_constraint(
&self.name,
self.column_names.join(", "),
format!("<vector({} dims)>", self.dims),
));
}
}
}
inner.vectors.reserve(prepared.len() * dims_bytes);
inner.nodes.reserve(prepared.len());
inner.node_to_row_id.reserve(prepared.len());
inner.row_id_to_node.reserve(prepared.len());
self.insert_prepared(&mut inner, &prepared);
Ok(())
}
fn remove_batch_slice(&self, entries: &[(i64, &[Value])]) -> Result<()> {
let mut inner = self.inner.write();
for &(row_id, _) in entries {
if let Some(&node_id) = inner.row_id_to_node.get(&row_id) {
inner.unique_map_remove(node_id);
inner.set_deleted(node_id);
}
}
Ok(())
}
fn column_ids(&self) -> &[i32] {
&self.column_ids
}
fn column_names(&self) -> &[String] {
&self.column_names
}
fn data_types(&self) -> &[DataType] {
&self.data_types
}
fn index_type(&self) -> IndexType {
IndexType::Hnsw
}
fn is_unique(&self) -> bool {
self.is_unique
}
fn find(&self, _values: &[Value]) -> Result<Vec<IndexEntry>> {
Err(radixdb_core::Error::NotSupported(
"HNSW does not support equality lookup".to_string(),
))
}
fn find_range(
&self,
_min: &[Value],
_max: &[Value],
_min_inclusive: bool,
_max_inclusive: bool,
) -> Result<Vec<IndexEntry>> {
Err(radixdb_core::Error::NotSupported(
"HNSW does not support range lookup".to_string(),
))
}
fn find_with_operator(&self, _op: Operator, _values: &[Value]) -> Result<Vec<IndexEntry>> {
Err(radixdb_core::Error::NotSupported(
"HNSW does not support scalar operator lookup".to_string(),
))
}
fn get_filtered_row_ids(&self, _expr: &dyn Expression) -> Result<RowIdVec> {
Err(radixdb_core::Error::NotSupported(
"HNSW does not support expression-based filtering".to_string(),
))
}
fn search_nearest(&self, query: &Value, k: usize, ef_search: usize) -> Option<Vec<(i64, f64)>> {
let query_bytes = Self::extract_vector_bytes(query)?;
if query_bytes.len() != self.dims * 4 {
return None;
}
Some(self.search_nearest(query_bytes, k, ef_search))
}
fn hnsw_distance_metric(&self) -> Option<u8> {
Some(self.metric.as_u8())
}
fn hnsw_m(&self) -> Option<u16> {
Some(self.m as u16)
}
fn hnsw_ef_construction(&self) -> Option<u16> {
Some(self.ef_construction as u16)
}
fn default_ef_search(&self) -> Option<usize> {
Some(self.ef_search)
}
fn hnsw_graph_bytes(&self) -> Result<Option<Vec<u8>>> {
self.serialize_graph_bytes().map(Some)
}
fn hnsw_indexed_row_ids(&self) -> Option<Vec<i64>> {
let inner = self.inner.read();
Some(
inner
.node_to_row_id
.iter()
.enumerate()
.filter_map(|(node_idx, &row_id)| {
(!inner.is_deleted(node_idx as u32)).then_some(row_id)
})
.collect(),
)
}
fn clear(&self) {
let mut inner = self.inner.write();
*inner = HnswInner::with_build_seed(self.dims, self.metric, self.build_seed);
if self.is_unique {
inner.unique_map = Some(ahash::AHashMap::new());
}
}
fn cleanup(&self) -> Result<()> {
let mut inner = self.inner.write();
let total = inner.nodes.len();
if total == 0 {
return Ok(());
}
let deleted: usize = inner
.deleted_bits
.iter()
.map(|w| w.count_ones() as usize)
.sum();
if deleted == 0 {
return Ok(());
}
if deleted * 5 < total {
return Ok(());
}
let dims_bytes = inner.dims_bytes;
let live_count = total - deleted;
let mut live_entries: Vec<(usize, i64)> = Vec::with_capacity(live_count);
for node_idx in 0..total {
if !inner.is_deleted(node_idx as u32) {
live_entries.push((node_idx * dims_bytes, inner.node_to_row_id[node_idx]));
}
}
let mut fresh = HnswInner::with_build_seed(self.dims, self.metric, self.build_seed);
if !live_entries.is_empty() {
let prepared: Vec<(&[u8], i64)> = live_entries
.iter()
.map(|&(offset, row_id)| (&inner.vectors[offset..offset + dims_bytes], row_id))
.collect();
fresh.vectors.reserve(live_count * dims_bytes);
fresh.nodes.reserve(live_count);
fresh.node_to_row_id.reserve(live_count);
fresh.row_id_to_node.reserve(live_count);
#[cfg(feature = "parallel")]
{
fresh.insert_batch_parallel(
&prepared,
self.m,
self.m0,
self.ef_construction,
self.ml,
);
}
#[cfg(not(feature = "parallel"))]
{
for &(vec_bytes, row_id) in &prepared {
fresh.insert(
vec_bytes,
row_id,
self.m,
self.m0,
self.ef_construction,
self.ml,
);
}
}
}
if self.is_unique {
fresh.build_unique_map();
}
*inner = fresh;
Ok(())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn close(&mut self) -> Result<()> {
Err(radixdb_core::Error::NotSupported(
"HNSW index lifecycle is owned by its table".to_string(),
))
}
}
#[cfg(test)]
#[path = "hnsw/tests.rs"]
mod tests;