//! HNSW index implementation with HDF5 serialization. use std::collections::{BinaryHeap, HashSet}; use clawhdf5_format::attribute::extract_attributes_full; use clawhdf5_format::data_layout::DataLayout; use clawhdf5_format::data_read::{read_as_f32, read_as_i32, read_raw_data_full}; use clawhdf5_format::dataspace::Dataspace; use clawhdf5_format::datatype::Datatype; use clawhdf5_format::error::FormatError; use clawhdf5_format::file_writer::{AttrValue, FileWriter as FmtWriter}; use clawhdf5_format::filter_pipeline::FilterPipeline; use clawhdf5_format::group_v2::resolve_path_any; use clawhdf5_format::message_type::MessageType; use clawhdf5_format::object_header::ObjectHeader; use clawhdf5_format::signature::find_signature; use clawhdf5_format::superblock::Superblock; use clawhdf5_io::FileWriter as IoFileWriter; /// Distance metric for the HNSW index. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum DistanceMetric { /// L2 (Euclidean) distance. L2, /// Cosine distance (1 - cosine_similarity). Cosine, } impl DistanceMetric { fn as_str(self) -> &'static str { match self { DistanceMetric::L2 => "l2", DistanceMetric::Cosine => "cosine", } } fn from_str(s: &str) -> Option { match s { "l2" => Some(DistanceMetric::L2), "cosine" => Some(DistanceMetric::Cosine), _ => None, } } } /// Compute distance between two vectors using the given metric. fn compute_distance(a: &[f32], b: &[f32], metric: DistanceMetric) -> f32 { match metric { DistanceMetric::L2 => { let mut sum = 0.0f32; for i in 0..a.len() { let d = a[i] - b[i]; sum += d * d; } sum.sqrt() } DistanceMetric::Cosine => { let mut dot = 0.0f32; let mut norm_a = 0.0f32; let mut norm_b = 0.0f32; for i in 0..a.len() { dot += a[i] * b[i]; norm_a += a[i] * a[i]; norm_b += b[i] * b[i]; } let denom = norm_a.sqrt() * norm_b.sqrt(); if denom < f32::EPSILON { 1.0 } else { 1.0 - (dot / denom) } } } } /// Assign a random level to a new node based on the HNSW probability distribution. /// /// Uses a deterministic approach based on the node index for reproducibility. fn assign_level(node_id: usize, m: usize) -> usize { let ml = 1.0 / (m as f64).ln(); // Use a simple hash-based pseudo-random for reproducibility let hash = splitmix64(node_id as u64); let uniform = (hash >> 11) as f64 / (1u64 << 53) as f64; (-uniform.ln() * ml).floor() as usize } /// Simple splitmix64 hash for deterministic level assignment. fn splitmix64(mut x: u64) -> u64 { x = x.wrapping_add(0x9e3779b97f4a7c15); x = (x ^ (x >> 30)).wrapping_mul(0xbf58476d1ce4e5b9); x = (x ^ (x >> 27)).wrapping_mul(0x94d049bb133111eb); x ^ (x >> 31) } /// Candidate neighbor for priority queue operations. #[derive(Debug, Clone)] struct Candidate { id: usize, distance: f32, } impl PartialEq for Candidate { fn eq(&self, other: &Self) -> bool { self.distance.to_bits() == other.distance.to_bits() && self.id == other.id } } impl Eq for Candidate {} impl PartialOrd for Candidate { fn partial_cmp(&self, other: &Self) -> Option { Some(self.cmp(other)) } } impl Ord for Candidate { fn cmp(&self, other: &Self) -> std::cmp::Ordering { // Reverse ordering for min-heap behavior other .distance .partial_cmp(&self.distance) .unwrap_or(std::cmp::Ordering::Equal) } } /// Max-heap candidate (furthest first). #[derive(Debug, Clone)] struct FarCandidate { id: usize, distance: f32, } impl PartialEq for FarCandidate { fn eq(&self, other: &Self) -> bool { self.distance.to_bits() == other.distance.to_bits() && self.id == other.id } } impl Eq for FarCandidate {} impl PartialOrd for FarCandidate { fn partial_cmp(&self, other: &Self) -> Option { Some(self.cmp(other)) } } impl Ord for FarCandidate { fn cmp(&self, other: &Self) -> std::cmp::Ordering { self.distance .partial_cmp(&other.distance) .unwrap_or(std::cmp::Ordering::Equal) } } /// On-disk format version for the serialized HNSW index. /// /// - Version 1: original layout (`vectors`, `graph_layer_*`, `config`), no /// deletion support and no explicit version tag. /// - Version 2: adds a `format_version` attribute and a `deleted` bitset dataset /// so live insert/delete state survives a save/load round-trip. /// /// Files written before this constant existed are treated as version 1 on load. pub const HNSW_FORMAT_VERSION: i64 = 2; /// HNSW (Hierarchical Navigable Small World) approximate nearest neighbor index. /// /// Supports building an index from vectors, incremental insertion and soft /// deletion, searching for nearest neighbors, and serializing/deserializing to /// HDF5 format. #[derive(Debug, Clone)] pub struct HnswIndex { /// All vectors in the index. vectors: Vec>, /// Adjacency lists per layer. `graph[layer][node]` = list of neighbor IDs. graph: Vec>>, /// Soft-deletion flags, one per node. Deleted nodes remain in the graph for /// connectivity but are never returned from [`HnswIndex::search`]. deleted: Vec, /// Entry point node ID. entry_point: usize, /// Maximum number of connections per node (per layer). m: usize, /// Maximum connections for layer 0 (typically 2*m). m_max0: usize, /// ef parameter used during construction. ef_construction: usize, /// Maximum layer assigned to each node. node_levels: Vec, /// Distance metric. metric: DistanceMetric, } impl HnswIndex { /// Build an HNSW index from a set of vectors. /// /// # Parameters /// - `vectors`: The vectors to index. All must have the same dimension. /// - `m`: Maximum number of connections per node (higher = more accurate, more memory). /// - `ef_construction`: Size of the dynamic candidate list during construction. /// /// Uses L2 distance by default. Use [`build_with_metric`] to specify the metric. pub fn build(vectors: &[Vec], m: usize, ef_construction: usize) -> Self { Self::build_with_metric(vectors, m, ef_construction, DistanceMetric::L2) } /// Build an HNSW index with a specific distance metric. pub fn build_with_metric( vectors: &[Vec], m: usize, ef_construction: usize, metric: DistanceMetric, ) -> Self { assert!(!vectors.is_empty(), "cannot build index from empty vectors"); assert!(m >= 2, "m must be at least 2"); let dim = vectors[0].len(); for v in vectors { assert_eq!(v.len(), dim, "all vectors must have the same dimension"); } let m_max0 = m * 2; let n = vectors.len(); // Assign levels to all nodes let mut node_levels = Vec::with_capacity(n); let mut max_level = 0; for i in 0..n { let level = assign_level(i, m); if level > max_level { max_level = level; } node_levels.push(level); } // Initialize graph layers let num_layers = max_level + 1; let mut graph: Vec>> = Vec::with_capacity(num_layers); for _ in 0..num_layers { graph.push(vec![Vec::new(); n]); } let mut entry_point = 0; let mut ep_level = node_levels[0]; // Insert nodes one by one for i in 1..n { let node_level = node_levels[i]; let mut ep = entry_point; // Phase 1: greedy search from top layer down to node_level + 1 let start_layer = ep_level; for layer in (node_level + 1..=start_layer).rev() { ep = greedy_closest(vectors, &graph[layer], &vectors[i], ep, metric); } // Phase 2: search and connect at layers node_level down to 0 let bottom = if node_level < start_layer { node_level } else { start_layer }; for layer in (0..=bottom).rev() { let max_conn = if layer == 0 { m_max0 } else { m }; let neighbors = search_layer( vectors, &graph[layer], &vectors[i], ep, ef_construction, metric, ); // Select up to m closest neighbors let selected: Vec = neighbors.iter().take(max_conn).map(|c| c.id).collect(); // Add bidirectional connections graph[layer][i] = selected.clone(); for &neighbor in &selected { graph[layer][neighbor].push(i); // Prune if over limit if graph[layer][neighbor].len() > max_conn { prune_connections( vectors, &mut graph[layer][neighbor], neighbor, max_conn, metric, ); } } if !selected.is_empty() { ep = selected[0]; } } // Update entry point if this node has a higher level if node_level > ep_level { entry_point = i; ep_level = node_level; } } Self { vectors: vectors.to_vec(), graph, deleted: vec![false; n], entry_point, m, m_max0, ef_construction, node_levels, metric, } } /// Create an empty index with the given parameters. Used as the starting /// point for incremental [`HnswIndex::insert`] and as the result of /// [`HnswIndex::compact`] when every vector has been deleted. pub fn new(m: usize, ef_construction: usize, metric: DistanceMetric) -> Self { assert!(m >= 2, "m must be at least 2"); Self { vectors: Vec::new(), graph: Vec::new(), deleted: Vec::new(), entry_point: 0, m, m_max0: m * 2, ef_construction, node_levels: Vec::new(), metric, } } /// Insert a single vector into the index incrementally and return its id. /// /// The id is the vector's position in insertion order and is stable for the /// life of the index (until [`HnswIndex::compact`] renumbers survivors). /// Inserting into an empty index seeds the entry point. /// /// # Panics /// Panics if `vector`'s dimension does not match the existing vectors. pub fn insert(&mut self, vector: Vec) -> usize { let id = self.vectors.len(); // Seed an empty index. if id == 0 { let node_level = assign_level(0, self.m); self.vectors.push(vector); self.deleted.push(false); self.node_levels.push(node_level); self.graph = (0..=node_level).map(|_| vec![Vec::new(); 1]).collect(); self.entry_point = 0; return 0; } assert_eq!( vector.len(), self.vectors[0].len(), "insert dimension mismatch" ); let node_level = assign_level(id, self.m); self.vectors.push(vector); self.deleted.push(false); self.node_levels.push(node_level); // Grow every existing layer with an empty adjacency slot for `id`, and // add any brand-new top layers this node introduces. for layer in self.graph.iter_mut() { layer.push(Vec::new()); } while self.graph.len() <= node_level { self.graph.push(vec![Vec::new(); id + 1]); } let ep_level = self.node_levels[self.entry_point]; let mut ep = self.entry_point; // Phase 1: greedy descent from the top down to node_level + 1. for layer in (node_level + 1..=ep_level).rev() { ep = greedy_closest(&self.vectors, &self.graph[layer], &self.vectors[id], ep, self.metric); } // Phase 2: search and connect from min(node_level, ep_level) down to 0. let bottom = node_level.min(ep_level); for layer in (0..=bottom).rev() { let max_conn = if layer == 0 { self.m_max0 } else { self.m }; let neighbors = search_layer( &self.vectors, &self.graph[layer], &self.vectors[id], ep, self.ef_construction, self.metric, ); let selected: Vec = neighbors.iter().take(max_conn).map(|c| c.id).collect(); self.graph[layer][id] = selected.clone(); for &neighbor in &selected { self.graph[layer][neighbor].push(id); if self.graph[layer][neighbor].len() > max_conn { prune_connections( &self.vectors, &mut self.graph[layer][neighbor], neighbor, max_conn, self.metric, ); } } if !selected.is_empty() { ep = selected[0]; } } // Promote the entry point if this node sits on a higher layer. if node_level > ep_level { self.entry_point = id; } id } /// Soft-delete the vector with the given id. The node stays in the graph so /// traversal/connectivity is preserved, but it will never be returned from /// [`HnswIndex::search`]. Idempotent; out-of-range ids are ignored. /// /// Returns `true` if the id existed and was not already deleted. pub fn mark_deleted(&mut self, id: usize) -> bool { if id >= self.deleted.len() || self.deleted[id] { return false; } self.deleted[id] = true; true } /// Returns whether the vector with the given id is soft-deleted. pub fn is_deleted(&self, id: usize) -> bool { self.deleted.get(id).copied().unwrap_or(false) } /// Number of soft-deleted vectors still occupying the index. pub fn deleted_count(&self) -> usize { self.deleted.iter().filter(|&&d| d).count() } /// Number of live (non-deleted) vectors. pub fn active_len(&self) -> usize { self.vectors.len() - self.deleted_count() } /// Rebuild the index from scratch, dropping all soft-deleted vectors and /// renumbering the survivors into a compact `0..active_len` id space. /// /// Returns a mapping from old id to new id (`None` for dropped vectors) so /// callers can rewrite any external id references they keep. pub fn compact(&mut self) -> Vec> { let mut mapping = vec![None; self.vectors.len()]; let mut surviving: Vec> = Vec::with_capacity(self.active_len()); for (old, v) in self.vectors.iter().enumerate() { if !self.deleted[old] { mapping[old] = Some(surviving.len()); surviving.push(v.clone()); } } *self = if surviving.is_empty() { Self::new(self.m, self.ef_construction, self.metric) } else { Self::build_with_metric(&surviving, self.m, self.ef_construction, self.metric) }; mapping } /// Search the index for the `k` nearest neighbors to the query vector. /// /// # Parameters /// - `query`: The query vector. /// - `k`: Number of nearest neighbors to return. /// - `ef`: Size of the dynamic candidate list during search (must be >= k). /// /// # Returns /// A vector of `(id, distance)` pairs sorted by distance (closest first). pub fn search(&self, query: &[f32], k: usize, ef: usize) -> Vec<(usize, f32)> { if self.vectors.is_empty() { return Vec::new(); } assert_eq!( query.len(), self.vectors[0].len(), "query dimension mismatch" ); let ef = ef.max(k); let mut ep = self.entry_point; let top_layer = self.graph.len().saturating_sub(1); // Greedy search from top layer down to layer 1 for layer in (1..=top_layer).rev() { ep = greedy_closest(&self.vectors, &self.graph[layer], query, ep, self.metric); } // Search layer 0 with ef candidates. Deleted nodes are still traversed // (they remain valid graph waypoints) but are filtered from the result. let candidates = search_layer(&self.vectors, &self.graph[0], query, ep, ef, self.metric); candidates .into_iter() .filter(|c| !self.deleted[c.id]) .take(k) .map(|c| (c.id, c.distance)) .collect() } /// Save the index to an HDF5 file via the given writer. pub fn save_to_hdf5(&self, writer: &mut IoFileWriter) -> Result<(), FormatError> { let bytes = self.to_hdf5_bytes()?; writer .write_bytes_owned(bytes) .map_err(|e| FormatError::SerializationError(e.to_string()))?; Ok(()) } /// Serialize the index to HDF5 bytes. pub fn to_hdf5_bytes(&self) -> Result, FormatError> { let mut fw = FmtWriter::new(); let n = self.vectors.len(); let dim = if n > 0 { self.vectors[0].len() } else { 0 }; // Flatten vectors into a 1D array for storage let flat_vectors: Vec = self .vectors .iter() .flat_map(|v| v.iter().copied()) .collect(); let mut group = fw.create_group("ann"); group .create_dataset("vectors") .with_f32_data(&flat_vectors) .with_shape(&[n as u64, dim as u64]) .set_attr("rows", AttrValue::I64(n as i64)) .set_attr("cols", AttrValue::I64(dim as i64)); // Serialize graph layers: store as flat i32 arrays with metadata let num_layers = self.graph.len(); for (layer_idx, layer) in self.graph.iter().enumerate() { // Flatten: for each node, store [count, neighbor1, neighbor2, ...] let mut flat: Vec = Vec::new(); for neighbors in layer { flat.push(neighbors.len() as i32); for &n_id in neighbors { flat.push(n_id as i32); } } let ds_name = format!("graph_layer_{layer_idx}"); group .create_dataset(&ds_name) .with_i32_data(&flat) .set_attr("layer", AttrValue::I64(layer_idx as i64)); } // Soft-deletion bitset (format version 2+): 0 = live, 1 = deleted. let deleted_i32: Vec = self.deleted.iter().map(|&d| d as i32).collect(); group.create_dataset("deleted").with_i32_data(&deleted_i32); // Store config as attributes on a small dataset let node_levels_i32: Vec = self.node_levels.iter().map(|&l| l as i32).collect(); group .create_dataset("config") .with_i32_data(&node_levels_i32) .set_attr("format_version", AttrValue::I64(HNSW_FORMAT_VERSION)) .set_attr("m", AttrValue::I64(self.m as i64)) .set_attr( "ef_construction", AttrValue::I64(self.ef_construction as i64), ) .set_attr("entry_point", AttrValue::I64(self.entry_point as i64)) .set_attr("num_layers", AttrValue::I64(num_layers as i64)) .set_attr( "metric", AttrValue::String(self.metric.as_str().to_string()), ) .set_attr("num_vectors", AttrValue::I64(n as i64)) .set_attr("dimension", AttrValue::I64(dim as i64)); let finished = group.finish(); fw.add_group(finished); fw.finish() } /// Load an HNSW index from HDF5 bytes. /// /// The HDF5 data must contain the `/ann/vectors`, `/ann/graph_layer_*`, /// and `/ann/config` datasets as produced by [`to_hdf5_bytes`]. pub fn load_from_hdf5(data: &[u8]) -> Result { let sig_offset = find_signature(data)?; let sb = Superblock::parse(data, sig_offset)?; // Read config dataset and its attributes let config_attrs = read_dataset_attrs(data, &sb, "ann/config")?; let config_raw = read_dataset_raw(data, &sb, "ann/config")?; let config_dt = read_dataset_datatype(data, &sb, "ann/config")?; let node_levels_i32 = read_as_i32(&config_raw, &config_dt)?; // Files written before format version 2 have no version attribute; treat // them as version 1. Reject anything newer than we understand. let format_version = get_attr_i64_opt(&config_attrs, "format_version").unwrap_or(1); if format_version > HNSW_FORMAT_VERSION { return Err(FormatError::SerializationError(format!( "unsupported HNSW format version {format_version} (this build understands up to {HNSW_FORMAT_VERSION})" ))); } let m = get_attr_i64(&config_attrs, "m")? as usize; let ef_construction = get_attr_i64(&config_attrs, "ef_construction")? as usize; let entry_point = get_attr_i64(&config_attrs, "entry_point")? as usize; let num_layers = get_attr_i64(&config_attrs, "num_layers")? as usize; let n = get_attr_i64(&config_attrs, "num_vectors")? as usize; let dim = get_attr_i64(&config_attrs, "dimension")? as usize; let metric_str = get_attr_string(&config_attrs, "metric")?; let metric = DistanceMetric::from_str(&metric_str).ok_or_else(|| { FormatError::SerializationError(format!("unknown metric: {metric_str}")) })?; let node_levels: Vec = node_levels_i32.iter().map(|&l| l as usize).collect(); // Read vectors let vectors_raw = read_dataset_raw(data, &sb, "ann/vectors")?; let vectors_dt = read_dataset_datatype(data, &sb, "ann/vectors")?; let flat_vectors = read_as_f32(&vectors_raw, &vectors_dt)?; let mut vectors = Vec::with_capacity(n); for i in 0..n { let start = i * dim; let end = start + dim; if end > flat_vectors.len() { return Err(FormatError::DataSizeMismatch { expected: end, actual: flat_vectors.len(), }); } vectors.push(flat_vectors[start..end].to_vec()); } // Read graph layers let mut graph = Vec::with_capacity(num_layers); for layer_idx in 0..num_layers { let ds_name = format!("ann/graph_layer_{layer_idx}"); let layer_raw = read_dataset_raw(data, &sb, &ds_name)?; let layer_dt = read_dataset_datatype(data, &sb, &ds_name)?; let flat = read_as_i32(&layer_raw, &layer_dt)?; let mut layer_graph = Vec::with_capacity(n); let mut pos = 0; while pos < flat.len() { let count = flat[pos] as usize; pos += 1; let mut neighbors = Vec::with_capacity(count); for _ in 0..count { if pos >= flat.len() { return Err(FormatError::SerializationError( "truncated graph data".into(), )); } neighbors.push(flat[pos] as usize); pos += 1; } layer_graph.push(neighbors); } // Pad with empty if needed (nodes not present at this layer) while layer_graph.len() < n { layer_graph.push(Vec::new()); } graph.push(layer_graph); } // Deleted bitset (version 2+). Older files default every node to live. let deleted = if format_version >= 2 { let deleted_raw = read_dataset_raw(data, &sb, "ann/deleted")?; let deleted_dt = read_dataset_datatype(data, &sb, "ann/deleted")?; let deleted_i32 = read_as_i32(&deleted_raw, &deleted_dt)?; let mut deleted: Vec = deleted_i32.iter().map(|&d| d != 0).collect(); deleted.resize(n, false); deleted } else { vec![false; n] }; Ok(Self { vectors, graph, deleted, entry_point, m, m_max0: m * 2, ef_construction, node_levels, metric, }) } /// Returns the number of vectors in the index. pub fn len(&self) -> usize { self.vectors.len() } /// Returns true if the index is empty. pub fn is_empty(&self) -> bool { self.vectors.is_empty() } /// Returns the dimension of vectors in the index. pub fn dimension(&self) -> usize { if self.vectors.is_empty() { 0 } else { self.vectors[0].len() } } /// Returns the number of layers in the graph. pub fn num_layers(&self) -> usize { self.graph.len() } /// Returns the distance metric used by this index. pub fn metric(&self) -> DistanceMetric { self.metric } /// Returns the maximum number of connections at layer 0. pub fn m_max0(&self) -> usize { self.m_max0 } } // --------------------------------------------------------------------------- // Internal HNSW algorithms // --------------------------------------------------------------------------- /// Greedy search: find the single closest node to `query` starting from `ep`. fn greedy_closest( vectors: &[Vec], layer: &[Vec], query: &[f32], mut ep: usize, metric: DistanceMetric, ) -> usize { let mut best_dist = compute_distance(query, &vectors[ep], metric); loop { let mut changed = false; for &neighbor in &layer[ep] { let d = compute_distance(query, &vectors[neighbor], metric); if d < best_dist { best_dist = d; ep = neighbor; changed = true; } } if !changed { break; } } ep } /// Search a single layer for the ef closest nodes to `query`. fn search_layer( vectors: &[Vec], layer: &[Vec], query: &[f32], ep: usize, ef: usize, metric: DistanceMetric, ) -> Vec { let ep_dist = compute_distance(query, &vectors[ep], metric); // Min-heap of candidates to explore let mut candidates = BinaryHeap::new(); candidates.push(Candidate { id: ep, distance: ep_dist, }); // Max-heap of current results (furthest first) let mut results = BinaryHeap::new(); results.push(FarCandidate { id: ep, distance: ep_dist, }); let mut visited = HashSet::new(); visited.insert(ep); while let Some(closest) = candidates.pop() { let furthest_dist = results.peek().map_or(f32::MAX, |f| f.distance); if closest.distance > furthest_dist && results.len() >= ef { break; } for &neighbor in &layer[closest.id] { if visited.contains(&neighbor) { continue; } visited.insert(neighbor); let d = compute_distance(query, &vectors[neighbor], metric); let furthest_dist = results.peek().map_or(f32::MAX, |f| f.distance); if d < furthest_dist || results.len() < ef { candidates.push(Candidate { id: neighbor, distance: d, }); results.push(FarCandidate { id: neighbor, distance: d, }); if results.len() > ef { results.pop(); } } } } // Convert to sorted vec let mut result: Vec = results .into_iter() .map(|f| Candidate { id: f.id, distance: f.distance, }) .collect(); result.sort_by(|a, b| { a.distance .partial_cmp(&b.distance) .unwrap_or(std::cmp::Ordering::Equal) }); result } /// Prune connections for a node to keep only the closest `max_conn` neighbors. fn prune_connections( vectors: &[Vec], neighbors: &mut Vec, node: usize, max_conn: usize, metric: DistanceMetric, ) { if neighbors.len() <= max_conn { return; } let mut scored: Vec<(usize, f32)> = neighbors .iter() .map(|&n| (n, compute_distance(&vectors[node], &vectors[n], metric))) .collect(); scored.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal)); scored.truncate(max_conn); *neighbors = scored.into_iter().map(|(id, _)| id).collect(); } // --------------------------------------------------------------------------- // HDF5 reading helpers // --------------------------------------------------------------------------- fn read_dataset_raw(data: &[u8], sb: &Superblock, path: &str) -> Result, FormatError> { let addr = resolve_path_any(data, sb, path)?; let header = ObjectHeader::parse(data, addr as usize, sb.offset_size, sb.length_size)?; let dt_msg = header .messages .iter() .find(|m| m.msg_type == MessageType::Datatype) .ok_or(FormatError::DatasetMissingData)?; let (datatype, _) = Datatype::parse(&dt_msg.data)?; let ds_msg = header .messages .iter() .find(|m| m.msg_type == MessageType::Dataspace) .ok_or(FormatError::DatasetMissingShape)?; let dataspace = Dataspace::parse(&ds_msg.data, sb.length_size)?; let dl_msg = header .messages .iter() .find(|m| m.msg_type == MessageType::DataLayout) .ok_or(FormatError::DatasetMissingData)?; let layout = DataLayout::parse(&dl_msg.data, sb.offset_size, sb.length_size)?; let pipeline = header .messages .iter() .find(|m| m.msg_type == MessageType::FilterPipeline) .and_then(|msg| FilterPipeline::parse(&msg.data).ok()); read_raw_data_full( data, &layout, &dataspace, &datatype, pipeline.as_ref(), sb.offset_size, sb.length_size, ) } fn read_dataset_datatype( data: &[u8], sb: &Superblock, path: &str, ) -> Result { let addr = resolve_path_any(data, sb, path)?; let header = ObjectHeader::parse(data, addr as usize, sb.offset_size, sb.length_size)?; let dt_msg = header .messages .iter() .find(|m| m.msg_type == MessageType::Datatype) .ok_or(FormatError::DatasetMissingData)?; let (datatype, _) = Datatype::parse(&dt_msg.data)?; Ok(datatype) } fn read_dataset_attrs( data: &[u8], sb: &Superblock, path: &str, ) -> Result, FormatError> { let addr = resolve_path_any(data, sb, path)?; let header = ObjectHeader::parse(data, addr as usize, sb.offset_size, sb.length_size)?; let attr_msgs = extract_attributes_full(data, &header, sb.offset_size, sb.length_size)?; let mut result = Vec::new(); for attr in &attr_msgs { let name = attr.name.clone(); if let Some(val) = decode_simple_attr(attr) { result.push((name, val)); } } Ok(result) } fn decode_simple_attr(attr: &clawhdf5_format::attribute::AttributeMessage) -> Option { let raw = &attr.raw_data; match &attr.datatype { Datatype::FixedPoint { size, signed, .. } => { if *signed { match size { 8 => { if raw.len() >= 8 { let val = i64::from_le_bytes(raw[..8].try_into().ok()?); Some(AttrValue::I64(val)) } else { None } } 4 => { if raw.len() >= 4 { let val = i32::from_le_bytes(raw[..4].try_into().ok()?) as i64; Some(AttrValue::I64(val)) } else { None } } _ => None, } } else if raw.len() >= *size as usize { let val = u64::from_le_bytes({ let mut buf = [0u8; 8]; buf[..*size as usize].copy_from_slice(&raw[..*size as usize]); buf }); Some(AttrValue::U64(val)) } else { None } } Datatype::FloatingPoint { size: 8, .. } => { if raw.len() >= 8 { let val = f64::from_le_bytes(raw[..8].try_into().ok()?); Some(AttrValue::F64(val)) } else { None } } Datatype::String { size, .. } => { let s = std::str::from_utf8(&raw[..*size as usize]).ok()?; Some(AttrValue::String(s.trim_end_matches('\0').to_string())) } _ => None, } } fn get_attr_i64(attrs: &[(String, AttrValue)], name: &str) -> Result { for (n, v) in attrs { if n == name { return match v { AttrValue::I64(val) => Ok(*val), AttrValue::U64(val) => Ok(*val as i64), _ => Err(FormatError::SerializationError(format!( "attribute {name} is not an integer" ))), }; } } Err(FormatError::SerializationError(format!( "missing attribute: {name}" ))) } /// Like [`get_attr_i64`] but returns `None` when the attribute is absent or not /// an integer, instead of erroring. Used for optional/back-compat attributes. fn get_attr_i64_opt(attrs: &[(String, AttrValue)], name: &str) -> Option { attrs.iter().find(|(n, _)| n == name).and_then(|(_, v)| match v { AttrValue::I64(val) => Some(*val), AttrValue::U64(val) => Some(*val as i64), _ => None, }) } fn get_attr_string(attrs: &[(String, AttrValue)], name: &str) -> Result { for (n, v) in attrs { if n == name { return match v { AttrValue::String(s) => Ok(s.clone()), _ => Err(FormatError::SerializationError(format!( "attribute {name} is not a string" ))), }; } } Err(FormatError::SerializationError(format!( "missing attribute: {name}" ))) } // --------------------------------------------------------------------------- // Tests // --------------------------------------------------------------------------- #[cfg(test)] mod tests { use super::*; fn make_random_vectors(n: usize, dim: usize, seed: u64) -> Vec> { let mut vectors = Vec::with_capacity(n); let mut state = seed; for _ in 0..n { let mut v = Vec::with_capacity(dim); for _ in 0..dim { state = splitmix64(state); let val = (state >> 40) as f32 / 16777216.0 - 0.5; v.push(val); } vectors.push(v); } vectors } #[test] fn build_small_index() { let vectors = vec![ vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0], vec![0.0, 0.0, 1.0], vec![1.0, 1.0, 0.0], ]; let index = HnswIndex::build(&vectors, 4, 16); assert_eq!(index.len(), 4); assert_eq!(index.dimension(), 3); assert!(!index.is_empty()); } #[test] fn search_exact_match() { let vectors = vec![ vec![1.0, 0.0], vec![0.0, 1.0], vec![1.0, 1.0], vec![0.5, 0.5], ]; let index = HnswIndex::build(&vectors, 4, 16); let results = index.search(&[1.0, 0.0], 1, 16); assert_eq!(results.len(), 1); assert_eq!(results[0].0, 0); // exact match assert!(results[0].1 < 1e-6); // distance ~0 } #[test] fn search_k_neighbors() { let vectors = make_random_vectors(50, 8, 42); let index = HnswIndex::build(&vectors, 8, 32); let results = index.search(&vectors[0], 5, 32); assert_eq!(results.len(), 5); // First result should be the query itself assert_eq!(results[0].0, 0); assert!(results[0].1 < 1e-6); // Results should be sorted by distance for i in 1..results.len() { assert!(results[i].1 >= results[i - 1].1); } } #[test] fn search_returns_correct_count() { let vectors = make_random_vectors(20, 4, 123); let index = HnswIndex::build(&vectors, 4, 16); let results = index.search(&vectors[5], 3, 16); assert_eq!(results.len(), 3); } #[test] fn cosine_distance() { let a = vec![1.0, 0.0]; let b = vec![0.0, 1.0]; let d = compute_distance(&a, &b, DistanceMetric::Cosine); assert!((d - 1.0).abs() < 1e-6); // orthogonal vectors = cosine distance 1 let c = vec![1.0, 0.0]; let d_same = compute_distance(&a, &c, DistanceMetric::Cosine); assert!(d_same < 1e-6); // same direction = cosine distance 0 } #[test] fn l2_distance() { let a = vec![0.0, 0.0]; let b = vec![3.0, 4.0]; let d = compute_distance(&a, &b, DistanceMetric::L2); assert!((d - 5.0).abs() < 1e-6); // 3-4-5 triangle } #[test] fn cosine_index_build_and_search() { let vectors = vec![ vec![1.0, 0.0, 0.0], vec![0.9, 0.1, 0.0], vec![0.0, 1.0, 0.0], vec![0.0, 0.0, 1.0], ]; let index = HnswIndex::build_with_metric(&vectors, 4, 16, DistanceMetric::Cosine); let results = index.search(&[1.0, 0.0, 0.0], 2, 16); assert_eq!(results.len(), 2); // The closest should be vector 0 (exact) or vector 1 (very similar) assert!(results[0].0 == 0 || results[0].0 == 1); } #[test] fn save_and_load_roundtrip() { let vectors = make_random_vectors(30, 4, 999); let index = HnswIndex::build(&vectors, 4, 16); let bytes = index.to_hdf5_bytes().unwrap(); assert!(!bytes.is_empty()); assert_eq!(&bytes[..8], b"\x89HDF\r\n\x1a\n"); let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap(); assert_eq!(loaded.len(), index.len()); assert_eq!(loaded.dimension(), index.dimension()); assert_eq!(loaded.metric(), DistanceMetric::L2); assert_eq!(loaded.m, index.m); assert_eq!(loaded.ef_construction, index.ef_construction); assert_eq!(loaded.entry_point, index.entry_point); // Verify vectors match for i in 0..loaded.len() { assert_eq!(loaded.vectors[i], index.vectors[i]); } } #[test] fn save_and_load_preserves_search_results() { let vectors = make_random_vectors(40, 6, 777); let index = HnswIndex::build(&vectors, 6, 24); let query = &vectors[10]; let original_results = index.search(query, 5, 24); let bytes = index.to_hdf5_bytes().unwrap(); let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap(); let loaded_results = loaded.search(query, 5, 24); assert_eq!(original_results.len(), loaded_results.len()); for (orig, load) in original_results.iter().zip(loaded_results.iter()) { assert_eq!(orig.0, load.0); assert!((orig.1 - load.1).abs() < 1e-6); } } #[test] fn cosine_roundtrip() { let vectors = make_random_vectors(20, 3, 555); let index = HnswIndex::build_with_metric(&vectors, 4, 16, DistanceMetric::Cosine); let bytes = index.to_hdf5_bytes().unwrap(); let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap(); assert_eq!(loaded.metric(), DistanceMetric::Cosine); } #[test] fn index_metadata() { let vectors = make_random_vectors(10, 5, 111); let index = HnswIndex::build(&vectors, 3, 12); assert_eq!(index.len(), 10); assert_eq!(index.dimension(), 5); assert!(index.num_layers() >= 1); assert_eq!(index.metric(), DistanceMetric::L2); } #[test] fn save_to_file_writer() { let vectors = make_random_vectors(15, 3, 333); let index = HnswIndex::build(&vectors, 4, 16); let dir = std::env::temp_dir(); let path = dir.join("clawhdf5_ann_test_save.h5"); let mut writer = IoFileWriter::create(&path).unwrap(); index.save_to_hdf5(&mut writer).unwrap(); // Verify file exists and has HDF5 signature let data = std::fs::read(&path).unwrap(); assert_eq!(&data[..8], b"\x89HDF\r\n\x1a\n"); // Load it back let loaded = HnswIndex::load_from_hdf5(&data).unwrap(); assert_eq!(loaded.len(), 15); std::fs::remove_file(&path).ok(); } #[test] fn search_accuracy_l2() { // Build a small index and verify that brute-force search agrees let vectors = make_random_vectors(100, 8, 42); let index = HnswIndex::build(&vectors, 16, 64); let query = &vectors[50]; let k = 10; let results = index.search(query, k, 64); // Brute-force nearest neighbors let mut brute: Vec<(usize, f32)> = vectors .iter() .enumerate() .map(|(i, v)| (i, compute_distance(query, v, DistanceMetric::L2))) .collect(); brute.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap()); brute.truncate(k); // With high ef and m, HNSW should find at least 80% of true neighbors let hnsw_ids: HashSet = results.iter().map(|r| r.0).collect(); let brute_ids: HashSet = brute.iter().map(|r| r.0).collect(); let overlap = hnsw_ids.intersection(&brute_ids).count(); assert!(overlap >= k * 8 / 10, "HNSW recall too low: {overlap}/{k}"); } #[test] fn search_accuracy_cosine() { let vectors = make_random_vectors(100, 8, 99); let index = HnswIndex::build_with_metric(&vectors, 16, 64, DistanceMetric::Cosine); let query = &vectors[25]; let k = 10; let results = index.search(query, k, 64); let mut brute: Vec<(usize, f32)> = vectors .iter() .enumerate() .map(|(i, v)| (i, compute_distance(query, v, DistanceMetric::Cosine))) .collect(); brute.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap()); brute.truncate(k); let hnsw_ids: HashSet = results.iter().map(|r| r.0).collect(); let brute_ids: HashSet = brute.iter().map(|r| r.0).collect(); let overlap = hnsw_ids.intersection(&brute_ids).count(); assert!( overlap >= k * 8 / 10, "HNSW cosine recall too low: {overlap}/{k}" ); } #[test] fn distance_metric_str_roundtrip() { assert_eq!(DistanceMetric::from_str("l2"), Some(DistanceMetric::L2)); assert_eq!( DistanceMetric::from_str("cosine"), Some(DistanceMetric::Cosine) ); assert_eq!(DistanceMetric::from_str("unknown"), None); assert_eq!(DistanceMetric::L2.as_str(), "l2"); assert_eq!(DistanceMetric::Cosine.as_str(), "cosine"); } #[test] fn cosine_zero_vector() { let a = vec![0.0, 0.0]; let b = vec![1.0, 0.0]; let d = compute_distance(&a, &b, DistanceMetric::Cosine); assert!((d - 1.0).abs() < 1e-6); // zero vector -> distance 1 } #[test] fn insert_into_empty_index() { let mut index = HnswIndex::new(4, 16, DistanceMetric::L2); assert!(index.is_empty()); let id = index.insert(vec![1.0, 0.0, 0.0]); assert_eq!(id, 0); assert_eq!(index.len(), 1); let results = index.search(&[1.0, 0.0, 0.0], 1, 16); assert_eq!(results, vec![(0, 0.0)]); } #[test] fn incremental_insert_matches_batch_recall() { // Build one index incrementally and one in batch from the same vectors, // then confirm the incremental index has acceptable recall vs brute force. let vectors = make_random_vectors(120, 8, 2024); let mut incremental = HnswIndex::new(16, 64, DistanceMetric::L2); for v in &vectors { incremental.insert(v.clone()); } assert_eq!(incremental.len(), vectors.len()); let query = &vectors[60]; let k = 10; let results = incremental.search(query, k, 64); let mut brute: Vec<(usize, f32)> = vectors .iter() .enumerate() .map(|(i, v)| (i, compute_distance(query, v, DistanceMetric::L2))) .collect(); brute.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap()); brute.truncate(k); let hnsw_ids: HashSet = results.iter().map(|r| r.0).collect(); let brute_ids: HashSet = brute.iter().map(|r| r.0).collect(); let overlap = hnsw_ids.intersection(&brute_ids).count(); assert!( overlap >= k * 8 / 10, "incremental HNSW recall too low: {overlap}/{k}" ); } #[test] fn mark_deleted_excludes_from_search() { let vectors = vec![ vec![1.0, 0.0], vec![0.9, 0.1], vec![0.0, 1.0], vec![0.1, 0.9], ]; let mut index = HnswIndex::build(&vectors, 4, 16); // Exact match on vector 0 before deletion. let before = index.search(&[1.0, 0.0], 1, 16); assert_eq!(before[0].0, 0); assert!(index.mark_deleted(0)); assert!(index.is_deleted(0)); assert!(!index.mark_deleted(0)); // idempotent assert_eq!(index.deleted_count(), 1); assert_eq!(index.active_len(), 3); // Vector 0 must no longer be returned; nearest is now vector 1. let after = index.search(&[1.0, 0.0], 2, 16); assert!(after.iter().all(|(id, _)| *id != 0)); assert_eq!(after[0].0, 1); } #[test] fn compact_drops_deleted_and_renumbers() { let vectors = make_random_vectors(20, 4, 4242); let mut index = HnswIndex::build(&vectors, 8, 32); index.mark_deleted(3); index.mark_deleted(7); index.mark_deleted(11); let mapping = index.compact(); assert_eq!(mapping.len(), 20); assert_eq!(index.len(), 17); assert_eq!(index.deleted_count(), 0); // Deleted ids map to None; survivors map to a dense 0..17 range. assert!(mapping[3].is_none() && mapping[7].is_none() && mapping[11].is_none()); let mut new_ids: Vec = mapping.iter().filter_map(|m| *m).collect(); new_ids.sort_unstable(); assert_eq!(new_ids, (0..17).collect::>()); } #[test] fn compact_all_deleted_yields_empty_index() { let vectors = make_random_vectors(5, 3, 1); let mut index = HnswIndex::build(&vectors, 4, 16); for i in 0..5 { index.mark_deleted(i); } index.compact(); assert!(index.is_empty()); assert!(index.search(&[0.0, 0.0, 0.0], 3, 16).is_empty()); } #[test] fn versioned_roundtrip_preserves_deletions() { let vectors = make_random_vectors(30, 4, 8080); let mut index = HnswIndex::build(&vectors, 6, 24); index.mark_deleted(5); index.mark_deleted(12); let bytes = index.to_hdf5_bytes().unwrap(); let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap(); assert_eq!(loaded.len(), index.len()); assert!(loaded.is_deleted(5)); assert!(loaded.is_deleted(12)); assert_eq!(loaded.deleted_count(), 2); // A deleted vector is not returned even after a round-trip. let results = loaded.search(&vectors[5], 5, 24); assert!(results.iter().all(|(id, _)| *id != 5)); } #[test] fn insert_then_save_load_search() { let mut index = HnswIndex::new(8, 32, DistanceMetric::Cosine); let vectors = make_random_vectors(25, 5, 31337); for v in &vectors { index.insert(v.clone()); } let bytes = index.to_hdf5_bytes().unwrap(); let loaded = HnswIndex::load_from_hdf5(&bytes).unwrap(); assert_eq!(loaded.len(), 25); assert_eq!(loaded.metric(), DistanceMetric::Cosine); let results = loaded.search(&vectors[0], 3, 32); assert_eq!(results.len(), 3); assert_eq!(results[0].0, 0); } }