use crate::platform::memset;
use crate::types::{DistanceType, VectorMetadata};
use crate::{RemDbError, Result};
use alloc::vec::Vec;
use core::cmp::Ordering;
use core::ptr::NonNull;
#[cfg(not(feature = "std"))]
struct XorShiftRng {
state: u64,
}
#[cfg(not(feature = "std"))]
impl XorShiftRng {
fn new(seed: u64) -> Self {
XorShiftRng {
state: if seed == 0 { 0x123456789abcdef } else { seed },
}
}
fn next_u64(&mut self) -> u64 {
let mut x = self.state;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.state = x;
x
}
fn next_f64(&mut self) -> f64 {
(self.next_u64() as f64) / (u64::MAX as f64)
}
}
#[repr(C)]
pub struct HNSWNode {
pub vector_offset: usize,
pub record_id: u16,
pub neighbor_counts: Vec<u8>,
pub neighbors: Vec<NonNull<HNSWNode>>,
}
impl HNSWNode {
pub unsafe fn new(vector_offset: usize, record_id: u16, max_level: usize) -> Self {
let mut neighbor_counts = Vec::with_capacity(max_level + 1);
for _ in 0..=max_level {
neighbor_counts.push(0);
}
let neighbors = alloc::vec![NonNull::dangling(); (max_level + 1) * 32];
HNSWNode {
vector_offset,
record_id,
neighbor_counts,
neighbors,
}
}
pub unsafe fn get_neighbors_at_level(&self, level: usize) -> &[NonNull<HNSWNode>] {
let start_offset = level * 32; let count = match self.neighbor_counts.get(level) {
Some(&c) => c as usize,
None => return &[],
};
let end = start_offset + count;
if end > self.neighbors.len() {
return &[];
}
&self.neighbors[start_offset..end]
}
pub unsafe fn add_neighbor_at_level(
&mut self,
level: usize,
neighbor: NonNull<HNSWNode>,
) -> Result<()> {
let start_offset = level * 32;
let count = match self.neighbor_counts.get(level) {
Some(&c) => c as usize,
None => return Err(RemDbError::OutOfMemory),
};
if count >= 32 {
return Err(RemDbError::OutOfMemory);
}
let idx = start_offset + count;
if idx >= self.neighbors.len() {
return Err(RemDbError::OutOfMemory);
}
self.neighbors[idx] = neighbor;
self.neighbor_counts[level] += 1;
Ok(())
}
}
pub struct HNSWIndex {
pub meta: VectorMetadata,
pub vectors: *mut f32,
pub max_level: usize,
pub enter_point: Option<NonNull<HNSWNode>>,
pub layer_enter_points: Vec<Option<NonNull<HNSWNode>>>,
pub nodes: NonNull<HNSWNode>,
pub free_nodes: Option<NonNull<HNSWNode>>,
pub max_nodes: usize,
pub node_count: usize,
pub lock: u32,
}
impl HNSWIndex {
pub unsafe fn new(
meta: VectorMetadata,
vectors: *mut f32,
memory_start: *mut u8,
max_nodes: usize,
) -> Result<Self> {
let max_level = if max_nodes > 0 {
(max_nodes as f64).ln() as usize
} else {
0
};
let nodes = NonNull::new_unchecked(memory_start as *mut HNSWNode);
let mut free_nodes = None;
for i in (0..max_nodes).rev() {
let node_ptr = nodes.as_ptr().add(i);
let node = HNSWNode::new(0, 0, max_level);
core::ptr::write(node_ptr, node);
free_nodes = Some(NonNull::new_unchecked(node_ptr));
}
let mut layer_enter_points = Vec::with_capacity(max_level + 1);
for _ in 0..=max_level {
layer_enter_points.push(None);
}
Ok(HNSWIndex {
meta,
vectors,
max_level,
enter_point: None,
layer_enter_points,
nodes,
free_nodes,
max_nodes,
node_count: 0,
lock: 0,
})
}
unsafe fn calculate_distance(&self, vec1: *const f32, vec2: *const f32) -> f32 {
match self.meta.distance_type {
DistanceType::L2 => {
let mut sum = 0.0;
for i in 0..self.meta.dimension {
let diff = *vec1.add(i as usize) - *vec2.add(i as usize);
sum += diff * diff;
}
sum.sqrt()
}
DistanceType::InnerProduct => {
let mut sum = 0.0;
for i in 0..self.meta.dimension {
sum += *vec1.add(i as usize) * *vec2.add(i as usize);
}
-sum }
DistanceType::Cosine => {
let mut dot = 0.0;
let mut norm1 = 0.0;
let mut norm2 = 0.0;
for i in 0..self.meta.dimension {
let v1 = *vec1.add(i as usize);
let v2 = *vec2.add(i as usize);
dot += v1 * v2;
norm1 += v1 * v1;
norm2 += v2 * v2;
}
let norm1 = norm1.sqrt();
let norm2 = norm2.sqrt();
if norm1 == 0.0 || norm2 == 0.0 {
-1.0 } else {
-(dot / (norm1 * norm2)) }
}
}
}
fn generate_random_level(&self) -> usize {
let mut level = 0;
let p = 0.5; #[cfg(feature = "std")]
while level < self.max_level && rand::random::<f64>() < p {
level += 1;
}
#[cfg(not(feature = "std"))]
{
let seed = core::time::Instant::now().elapsed().as_nanos() as u64;
let mut rng = XorShiftRng::new(seed);
while level < self.max_level && rng.next_f64() < p {
level += 1;
}
}
level
}
#[cfg(feature = "std")]
pub fn save<W: std::io::Write>(&self, writer: &mut W) -> Result<()> {
writer
.write_all(&self.meta.hnsw_m.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.meta.hnsw_ef_construction.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.meta.hnsw_ef_search.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.max_level.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.node_count.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
let enter_point_offset = match self.enter_point {
Some(point) => {
let offset = unsafe { point.as_ptr().offset_from(self.nodes.as_ptr()) } as usize;
offset
}
None => usize::MAX,
};
writer
.write_all(&enter_point_offset.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&self.layer_enter_points.len().to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
for &point in &self.layer_enter_points {
let offset = match point {
Some(p) => {
let offset = unsafe { p.as_ptr().offset_from(self.nodes.as_ptr()) } as usize;
offset
}
None => usize::MAX,
};
writer
.write_all(&offset.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
}
for i in 0..self.node_count {
let node_ptr = unsafe { self.nodes.as_ptr().add(i) };
let node = unsafe { &*node_ptr };
writer
.write_all(&node.vector_offset.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&node.record_id.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
writer
.write_all(&node.neighbor_counts.len().to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
for &count in &node.neighbor_counts {
writer
.write_all(&count.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
}
for &neighbor in &node.neighbors {
let offset = unsafe { neighbor.as_ptr().offset_from(self.nodes.as_ptr()) } as usize;
writer
.write_all(&offset.to_le_bytes())
.map_err(|_| RemDbError::FileIoError)?;
}
}
Ok(())
}
#[cfg(feature = "std")]
pub unsafe fn load<R: std::io::Read>(
mut meta: VectorMetadata,
vectors: *mut f32,
memory_start: *mut u8,
max_nodes: usize,
reader: &mut R,
) -> Result<Self> {
let mut m_bytes = [0u8; 1];
reader
.read_exact(&mut m_bytes)
.map_err(|_| RemDbError::FileIoError)?;
meta.hnsw_m = m_bytes[0];
let mut efc_bytes = [0u8; 4];
reader
.read_exact(&mut efc_bytes)
.map_err(|_| RemDbError::FileIoError)?;
meta.hnsw_ef_construction = u32::from_le_bytes(efc_bytes);
let mut efs_bytes = [0u8; 4];
reader
.read_exact(&mut efs_bytes)
.map_err(|_| RemDbError::FileIoError)?;
meta.hnsw_ef_search = u32::from_le_bytes(efs_bytes);
let mut max_level_bytes = [0u8; 8];
reader
.read_exact(&mut max_level_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let max_level = usize::from_le_bytes(max_level_bytes);
let mut node_count_bytes = [0u8; 8];
reader
.read_exact(&mut node_count_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let node_count = usize::from_le_bytes(node_count_bytes);
let _node_size = core::mem::size_of::<HNSWNode>();
let nodes = NonNull::new_unchecked(memory_start as *mut HNSWNode);
let mut enter_point_offset_bytes = [0u8; 8];
reader
.read_exact(&mut enter_point_offset_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let enter_point_offset = usize::from_le_bytes(enter_point_offset_bytes);
let enter_point = if enter_point_offset == usize::MAX {
None
} else {
Some(NonNull::new_unchecked(
nodes.as_ptr().add(enter_point_offset),
))
};
let mut layer_enter_points_len_bytes = [0u8; 8];
reader
.read_exact(&mut layer_enter_points_len_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let layer_enter_points_len = usize::from_le_bytes(layer_enter_points_len_bytes);
let mut layer_enter_points = Vec::with_capacity(layer_enter_points_len);
for _ in 0..layer_enter_points_len {
let mut offset_bytes = [0u8; 8];
reader
.read_exact(&mut offset_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let offset = usize::from_le_bytes(offset_bytes);
let point = if offset == usize::MAX {
None
} else {
Some(NonNull::new_unchecked(nodes.as_ptr().add(offset)))
};
layer_enter_points.push(point);
}
for i in 0..node_count {
let node_ptr = nodes.as_ptr().add(i);
let mut vector_offset_bytes = [0u8; 8];
reader
.read_exact(&mut vector_offset_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let vector_offset = usize::from_le_bytes(vector_offset_bytes);
let mut record_id_bytes = [0u8; 2];
reader
.read_exact(&mut record_id_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let record_id = u16::from_le_bytes(record_id_bytes);
let mut neighbor_counts_len_bytes = [0u8; 8];
reader
.read_exact(&mut neighbor_counts_len_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let neighbor_counts_len = usize::from_le_bytes(neighbor_counts_len_bytes);
let mut neighbor_counts = Vec::with_capacity(neighbor_counts_len);
for _ in 0..neighbor_counts_len {
let mut count_bytes = [0u8; 1];
reader
.read_exact(&mut count_bytes)
.map_err(|_| RemDbError::FileIoError)?;
neighbor_counts.push(count_bytes[0]);
}
let total_neighbors = neighbor_counts.iter().sum::<u8>() as usize;
let mut neighbors = Vec::with_capacity(total_neighbors);
for _ in 0..total_neighbors {
let mut offset_bytes = [0u8; 8];
reader
.read_exact(&mut offset_bytes)
.map_err(|_| RemDbError::FileIoError)?;
let offset = usize::from_le_bytes(offset_bytes);
let neighbor = NonNull::new_unchecked(nodes.as_ptr().add(offset));
neighbors.push(neighbor);
}
let node = HNSWNode {
vector_offset,
record_id,
neighbor_counts,
neighbors,
};
*node_ptr = node;
}
Ok(HNSWIndex {
meta,
vectors,
max_level,
enter_point,
layer_enter_points,
nodes,
free_nodes: None, max_nodes,
node_count,
lock: 0,
})
}
unsafe fn search_layer(
&self,
query_vec: *const f32,
entry_point: NonNull<HNSWNode>,
ef: usize,
level: usize,
) -> Vec<(f32, NonNull<HNSWNode>)> {
let mut visited = Vec::new();
let mut candidates = Vec::new();
let mut results = Vec::new();
const MAX_ITERATIONS: usize = 10000;
let mut iteration_count = 0;
let entry_node = entry_point.as_ref();
let entry_vec = self.vectors.add(entry_node.vector_offset);
let distance = self.calculate_distance(query_vec, entry_vec);
candidates.push((distance, entry_point));
results.push((distance, entry_point));
visited.push(entry_point);
while let Some((current_dist, current_node)) = candidates.pop() {
iteration_count += 1;
if iteration_count > MAX_ITERATIONS {
#[cfg(feature = "log")]
crate::log::warn!(
"HNSW search_layer reached max iterations, graph may be malformed"
);
break;
}
if results.len() < ef || current_dist < results.last().map(|r| r.0).unwrap_or(f32::MAX)
{
let neighbors = current_node.as_ref().get_neighbors_at_level(level);
for &neighbor in neighbors {
if !visited.contains(&neighbor) {
visited.push(neighbor);
let neighbor_node = neighbor.as_ref();
let neighbor_vec = self.vectors.add(neighbor_node.vector_offset);
let neighbor_dist = self.calculate_distance(query_vec, neighbor_vec);
if results.len() < ef
|| neighbor_dist < results.last().map(|r| r.0).unwrap_or(f32::MAX)
{
candidates.push((neighbor_dist, neighbor));
results.push((neighbor_dist, neighbor));
results
.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(Ordering::Equal));
if results.len() > ef {
results.pop();
}
}
}
}
}
}
results
}
pub unsafe fn search(&self, query_vec: *const f32, _k: usize) -> Result<Vec<(f32, u16)>> {
if self.enter_point.is_none() {
return Err(RemDbError::RecordNotFound);
}
let mut current_point = self.enter_point.ok_or(RemDbError::InvalidState)?;
let mut current_level = self.max_level;
while current_level > 0 {
let results = self.search_layer(query_vec, current_point, 1, current_level);
if let Some(&(_dist, point)) = results.get(0) {
current_point = point;
}
current_level -= 1;
}
let results = self.search_layer(
query_vec,
current_point,
self.meta.hnsw_ef_search as usize,
0,
);
let mut final_results = Vec::new();
for (distance, node) in results {
final_results.push((distance, node.as_ref().record_id));
}
Ok(final_results)
}
pub unsafe fn insert(&mut self, vector_offset: usize, record_id: u16) -> Result<()> {
let vec_ptr = self.vectors.add(vector_offset);
let new_level = self.generate_random_level();
if self.node_count >= self.max_nodes {
return Err(RemDbError::OutOfMemory);
}
let node_ptr = self.nodes.as_ptr().add(self.node_count);
core::ptr::write(
node_ptr,
HNSWNode::new(vector_offset, record_id, self.max_level),
);
let pool_node = NonNull::new_unchecked(node_ptr);
let mut entry_point = match self.enter_point {
Some(point) => point,
None => {
self.enter_point = Some(pool_node);
for i in 0..=new_level {
if let Some(ep) = self.layer_enter_points.get_mut(i) {
*ep = Some(pool_node);
}
}
self.node_count += 1;
return Ok(());
}
};
let mut current_level = self.max_level;
while current_level > new_level {
let results = self.search_layer(vec_ptr, entry_point, 1, current_level);
if let Some(&(_dist, point)) = results.get(0) {
entry_point = point;
}
current_level -= 1;
}
while current_level <= new_level {
let ef_construction = self.meta.hnsw_ef_construction as usize;
let neighbors = self.search_layer(vec_ptr, entry_point, ef_construction, current_level);
let m = self.meta.hnsw_m as usize;
let selected_neighbors = neighbors
.iter()
.take(m)
.map(|&(_d, n)| n)
.collect::<Vec<_>>();
let node_mut = &mut *node_ptr; for &neighbor in &selected_neighbors {
node_mut.add_neighbor_at_level(current_level, neighbor)?;
let neighbor_ptr = neighbor.as_ptr();
let neighbor_mut = &mut *neighbor_ptr;
neighbor_mut.add_neighbor_at_level(current_level, pool_node)?;
}
let current_ep = self.layer_enter_points.get_mut(current_level);
if let Some(ep) = current_ep {
if ep.is_none() {
*ep = Some(pool_node);
}
}
current_level += 1;
}
self.node_count += 1;
if new_level >= self.max_level {
self.enter_point = Some(pool_node);
}
Ok(())
}
pub unsafe fn delete(&mut self, vector_offset: usize) -> Result<()> {
let mut target_node = None;
let mut target_node_idx = None;
for i in 0..self.node_count {
let node_ptr = self.nodes.as_ptr().add(i);
let node = unsafe { &*node_ptr };
if node.vector_offset == vector_offset {
target_node = Some(NonNull::new_unchecked(node_ptr));
target_node_idx = Some(i);
break;
}
}
if let Some(target_node) = target_node {
for i in 0..self.node_count {
let node_ptr = self.nodes.as_ptr().add(i);
let node = unsafe { &mut *node_ptr };
for level in 0..=self.max_level {
let start_offset = level * 32;
let count = match node.neighbor_counts.get(level) {
Some(&c) => c as usize,
None => continue,
};
let end = start_offset + count;
if end > node.neighbors.len() {
continue;
}
let mut new_neighbors = Vec::new();
let mut new_count = 0;
for j in 0..count {
let idx = start_offset + j;
let neighbor = match node.neighbors.get(idx) {
Some(&n) => n,
None => continue,
};
if neighbor != target_node {
new_neighbors.push(neighbor);
new_count += 1;
}
}
if new_count < count {
for (j, &neighbor) in new_neighbors.iter().enumerate() {
let idx = start_offset + j;
if idx < node.neighbors.len() {
node.neighbors[idx] = neighbor;
}
}
if let Some(nc) = node.neighbor_counts.get_mut(level) {
*nc = new_count as u8;
}
}
}
}
memset(
target_node.as_ptr() as *mut u8,
0,
core::mem::size_of::<HNSWNode>(),
);
let mut node = target_node;
node.as_mut().neighbors.clear();
node.as_mut().neighbor_counts.clear();
let _next_free = self.free_nodes;
self.free_nodes = Some(node);
Ok(())
} else {
Err(RemDbError::RecordNotFound)
}
}
}