use crate::classic::trees::persistence::{read_json, validate_vector_shape, write_json_atomic};
use crate::simd;
use crate::RetrieveError;
use serde::{Deserialize, Serialize};
use std::path::Path;
const RP_FOREST_FORMAT_VERSION: u32 = 1;
const DEFAULT_RP_FOREST_SEED: u64 = 0xA24B_AED4_963E_E407;
#[derive(Deserialize, Serialize)]
pub struct RpForestIndex {
pub(crate) vectors: Vec<f32>,
pub(crate) dimension: usize,
pub(crate) num_vectors: usize,
doc_ids: Vec<u32>,
params: RpForestParams,
built: bool,
pub(crate) trees: Vec<RPTree>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct RpForestParams {
pub num_trees: usize,
pub tree_params: RPTreeParams,
}
impl Default for RpForestParams {
fn default() -> Self {
Self {
num_trees: 10,
tree_params: RPTreeParams::default(),
}
}
}
#[derive(Clone, Deserialize, Serialize)]
pub(crate) struct RPTree {
root: Option<TreeNode>,
}
#[derive(Clone, Deserialize, Serialize)]
enum TreeNode {
Leaf {
indices: Vec<u32>,
},
Internal {
hyperplane: Vec<f32>, #[allow(dead_code)]
threshold: f32, left: Box<TreeNode>,
right: Box<TreeNode>,
},
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct RPTreeParams {
pub max_leaf_size: usize,
}
#[derive(Deserialize, Serialize)]
struct RpForestSnapshot {
version: u32,
index: RpForestIndex,
}
impl Default for RPTreeParams {
fn default() -> Self {
Self { max_leaf_size: 10 }
}
}
impl RpForestIndex {
pub fn new(dimension: usize, params: RpForestParams) -> Result<Self, RetrieveError> {
if dimension == 0 {
return Err(RetrieveError::InvalidParameter(
"dimension must be > 0".into(),
));
}
Ok(Self {
vectors: Vec::new(),
dimension,
num_vectors: 0,
doc_ids: Vec::new(),
params,
built: false,
trees: Vec::new(),
})
}
pub fn add(&mut self, doc_id: u32, vector: Vec<f32>) -> Result<(), RetrieveError> {
if self.built {
return Err(RetrieveError::InvalidParameter(
"Cannot add vectors after index is built".to_string(),
));
}
if vector.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: vector.len(),
doc_dim: self.dimension,
});
}
self.vectors.extend_from_slice(&vector);
self.doc_ids.push(doc_id);
self.num_vectors += 1;
Ok(())
}
pub fn build(&mut self) -> Result<(), RetrieveError> {
if self.built {
return Ok(());
}
if self.num_vectors == 0 {
return Err(RetrieveError::EmptyIndex);
}
self.trees = Vec::new();
for tree_idx in 0..self.params.num_trees {
let tree = self.build_tree(tree_seed(tree_idx))?;
self.trees.push(tree);
}
self.built = true;
Ok(())
}
pub fn save_to_dir(&self, output_dir: impl AsRef<Path>) -> Result<(), RetrieveError> {
if !self.built {
return Err(RetrieveError::InvalidParameter(
"cannot save unbuilt RP-forest index".into(),
));
}
let output_dir = output_dir.as_ref();
std::fs::create_dir_all(output_dir)?;
write_json_atomic(
&output_dir.join("index.json"),
&RpForestSnapshot {
version: RP_FOREST_FORMAT_VERSION,
index: self.clone_for_snapshot(),
},
)
}
pub fn load_from_dir(input_dir: impl AsRef<Path>) -> Result<Self, RetrieveError> {
let snapshot: RpForestSnapshot = read_json(&input_dir.as_ref().join("index.json"))?;
if snapshot.version != RP_FOREST_FORMAT_VERSION {
return Err(RetrieveError::FormatError(format!(
"unsupported RP-forest format version {}",
snapshot.version
)));
}
let index = snapshot.index;
validate_vector_shape(
"RP-forest",
index.dimension,
index.num_vectors,
&index.vectors,
&index.doc_ids,
)?;
if !index.built || index.trees.is_empty() {
return Err(RetrieveError::FormatError(
"RP-forest snapshot is not built".into(),
));
}
Ok(index)
}
fn clone_for_snapshot(&self) -> Self {
Self {
vectors: self.vectors.clone(),
dimension: self.dimension,
num_vectors: self.num_vectors,
doc_ids: self.doc_ids.clone(),
params: self.params.clone(),
built: self.built,
trees: self.trees.clone(),
}
}
pub fn memory_usage(&self) -> crate::memory::MemoryReport {
let tree_storage = self.trees.capacity() * std::mem::size_of::<RPTree>();
let node_storage: usize = self.trees.iter().map(RPTree::owned_bytes).sum();
crate::memory::MemoryReport {
vectors_bytes: self.vectors.capacity() * std::mem::size_of::<f32>(),
graph_bytes: tree_storage + node_storage,
quantized_bytes: 0,
metadata_bytes: self.doc_ids.capacity() * std::mem::size_of::<u32>(),
}
}
fn build_tree(&self, seed: u64) -> Result<RPTree, RetrieveError> {
let indices: Vec<u32> = (0..self.num_vectors as u32).collect();
let root = self.build_tree_recursive(&indices, seed)?;
Ok(RPTree { root })
}
fn build_tree_recursive(
&self,
indices: &[u32],
seed: u64,
) -> Result<Option<TreeNode>, RetrieveError> {
if indices.is_empty() {
return Ok(None);
}
if indices.len() <= self.params.tree_params.max_leaf_size {
return Ok(Some(TreeNode::Leaf {
indices: indices.to_vec(),
}));
}
use rand::rngs::StdRng;
use rand::Rng;
use rand::SeedableRng;
let mut rng = StdRng::seed_from_u64(seed);
let mut hyperplane = Vec::with_capacity(self.dimension);
let mut norm = 0.0f32;
for _ in 0..self.dimension {
let val = rng.random::<f32>() * 2.0 - 1.0;
norm += val * val;
hyperplane.push(val);
}
let norm = norm.sqrt();
if norm > 0.0 {
for val in &mut hyperplane {
*val /= norm;
}
}
let mut left_indices = Vec::new();
let mut right_indices = Vec::new();
for &idx in indices {
let vec = self.get_vector(idx as usize);
let projection = simd::dot(vec, &hyperplane);
if projection < 0.0 {
left_indices.push(idx);
} else {
right_indices.push(idx);
}
}
if left_indices.is_empty() || right_indices.is_empty() {
return Ok(Some(TreeNode::Leaf {
indices: indices.to_vec(),
}));
}
let left = self.build_tree_recursive(&left_indices, child_seed(seed, 0))?;
let right = self.build_tree_recursive(&right_indices, child_seed(seed, 1))?;
Ok(Some(TreeNode::Internal {
hyperplane,
threshold: 0.0,
left: Box::new(left.unwrap_or(TreeNode::Leaf {
indices: Vec::new(),
})),
right: Box::new(right.unwrap_or(TreeNode::Leaf {
indices: Vec::new(),
})),
}))
}
pub fn search(&self, query: &[f32], k: usize) -> Result<Vec<(u32, f32)>, RetrieveError> {
if !self.built {
return Err(RetrieveError::InvalidParameter(
"Index must be built before search".to_string(),
));
}
if query.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: query.len(),
doc_dim: self.dimension,
});
}
let mut candidate_set = std::collections::HashSet::with_capacity(
self.params.num_trees * self.params.tree_params.max_leaf_size,
);
for tree in &self.trees {
if let Some(ref root) = tree.root {
self.collect_leaf_candidates(root, query, &mut candidate_set);
}
}
let mut results: Vec<(u32, f32)> = candidate_set
.iter()
.map(|&idx| {
let vec = self.get_vector(idx as usize);
let dist = 1.0 - simd::dot(query, vec);
(self.doc_ids[idx as usize], dist)
})
.collect();
results.sort_unstable_by(|a, b| a.1.total_cmp(&b.1).then_with(|| a.0.cmp(&b.0)));
Ok(results.into_iter().take(k).collect())
}
fn collect_leaf_candidates(
&self,
node: &TreeNode,
query: &[f32],
candidates: &mut std::collections::HashSet<u32>,
) {
match node {
TreeNode::Leaf { indices } => candidates.extend(indices.iter().copied()),
TreeNode::Internal {
hyperplane,
threshold: _,
left,
right,
} => {
let projection = simd::dot(query, hyperplane);
if projection < 0.0 {
self.collect_leaf_candidates(left, query, candidates);
} else {
self.collect_leaf_candidates(right, query, candidates);
}
}
}
}
fn get_vector(&self, idx: usize) -> &[f32] {
let start = idx * self.dimension;
let end = start + self.dimension;
&self.vectors[start..end]
}
}
impl RPTree {
fn owned_bytes(&self) -> usize {
self.root.as_ref().map(TreeNode::owned_bytes).unwrap_or(0)
}
}
impl TreeNode {
fn owned_bytes(&self) -> usize {
match self {
TreeNode::Leaf { indices } => indices.capacity() * std::mem::size_of::<u32>(),
TreeNode::Internal {
hyperplane,
left,
right,
..
} => {
hyperplane.capacity() * std::mem::size_of::<f32>()
+ boxed_node_bytes(left)
+ boxed_node_bytes(right)
}
}
}
}
fn boxed_node_bytes(node: &TreeNode) -> usize {
std::mem::size_of::<TreeNode>() + node.owned_bytes()
}
fn tree_seed(tree_idx: usize) -> u64 {
DEFAULT_RP_FOREST_SEED ^ (tree_idx as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15)
}
fn child_seed(seed: u64, branch: u64) -> u64 {
seed.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407 ^ branch)
}
#[cfg(test)]
mod tests {
use super::*;
fn build_index(n: usize, dim: usize) -> RpForestIndex {
let params = RpForestParams {
num_trees: 5,
tree_params: RPTreeParams { max_leaf_size: 10 },
};
let mut index = RpForestIndex::new(dim, params).unwrap();
for i in 0..n {
let mut v = vec![0.0f32; dim];
v[i % dim] = 1.0;
index.add(5000 + i as u32, v).unwrap();
}
index.build().unwrap();
index
}
#[test]
fn build_is_deterministic() {
let first = build_index(50, 4);
let second = build_index(50, 4);
assert_eq!(
serde_json::to_value(&first.trees).unwrap(),
serde_json::to_value(&second.trees).unwrap()
);
}
#[test]
fn test_basic_search_returns_results() {
let index = build_index(50, 4);
let query = vec![1.0, 0.0, 0.0, 0.0];
let results = index.search(&query, 5).unwrap();
assert!(!results.is_empty());
assert!(results.len() <= 5);
}
#[test]
fn test_search_returns_at_most_k() {
let index = build_index(50, 4);
let query = vec![1.0, 0.0, 0.0, 0.0];
for k in [1, 3, 5, 10] {
let results = index.search(&query, k).unwrap();
assert!(results.len() <= k);
}
}
#[test]
fn test_results_sorted_by_distance() {
let index = build_index(50, 4);
let query = vec![1.0, 0.0, 0.0, 0.0];
let results = index.search(&query, 10).unwrap();
for w in results.windows(2) {
assert!(w[0].1 <= w[1].1, "results not sorted: {:?}", results);
}
}
#[test]
fn test_ids_in_bounds() {
let n = 50usize;
let index = build_index(n, 4);
let query = vec![1.0, 0.0, 0.0, 0.0];
let results = index.search(&query, 10).unwrap();
for (id, _) in results {
assert!(
(5000..5000 + n as u32).contains(&id),
"id {} out of shifted doc-id bounds (n={})",
id,
n
);
}
}
#[test]
fn test_multiple_trees_improve_coverage() {
let dim = 8;
let n = 100;
let params = RpForestParams {
num_trees: 20,
tree_params: RPTreeParams { max_leaf_size: 5 },
};
let mut index = RpForestIndex::new(dim, params).unwrap();
for i in 0..n {
let mut v = vec![0.0f32; dim];
if i < n / 2 {
v[0] = 1.0;
v[1] = 0.01 * (i as f32);
} else {
v[1] = 1.0;
v[0] = 0.01 * (i as f32);
}
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
for x in &mut v {
*x /= norm;
}
index.add(6000 + i as u32, v).unwrap();
}
index.build().unwrap();
let query = vec![1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let results = index.search(&query, 10).unwrap();
let correct = results
.iter()
.filter(|(id, _)| (6000..6000 + (n / 2) as u32).contains(id))
.count();
assert!(
correct >= 5,
"only {correct}/10 results from correct cluster — hyperplane independence may be broken"
);
}
#[test]
fn test_build_errors_on_empty_index() {
let mut index = RpForestIndex::new(4, RpForestParams::default()).unwrap();
assert!(index.build().is_err());
}
#[test]
fn test_add_after_build_errors() {
let mut index = RpForestIndex::new(4, RpForestParams::default()).unwrap();
index.add(0, vec![1.0, 0.0, 0.0, 0.0]).unwrap();
index.build().unwrap();
assert!(index.add(1, vec![0.0, 1.0, 0.0, 0.0]).is_err());
}
#[test]
fn test_dimension_mismatch_errors() {
let mut index = RpForestIndex::new(4, RpForestParams::default()).unwrap();
assert!(index.add(0, vec![1.0, 0.0]).is_err());
}
#[test]
fn save_load_roundtrip_preserves_sampled_forest() {
let index = build_index(50, 4);
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = RpForestIndex::load_from_dir(dir.path()).unwrap();
let query = [1.0, 0.0, 0.0, 0.0];
assert_eq!(
index.search(&query, 8).unwrap(),
loaded.search(&query, 8).unwrap()
);
}
#[test]
fn test_degenerate_split_does_not_recurse_infinitely() {
let params = RpForestParams {
num_trees: 3,
tree_params: RPTreeParams { max_leaf_size: 2 },
};
let mut index = RpForestIndex::new(4, params).unwrap();
for i in 0..20u32 {
index.add(i, vec![1.0, 0.0, 0.0, 0.0]).unwrap();
}
index.build().unwrap();
let results = index.search(&[1.0, 0.0, 0.0, 0.0], 5).unwrap();
assert!(!results.is_empty());
}
}