weavatrix-search-vector 0.3.1

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
use super::format::{align8, bytes_for, read_u64, to_usize};
use super::validation::validate_canonical_offsets;
use super::{GRAPH_HEADER_LEN, Header, MappedGraph, NONE_ENTRY, SnapshotValidation};
use crate::error::SearchError;
use crate::mmap::Mapping;

pub(super) fn validate_sections(
    mapping: &Mapping,
    header: &Header,
    validation: SnapshotValidation,
) -> Result<Vec<MappedGraph>, SearchError> {
    validate_canonical_offsets(header)?;
    validate_vector_sections(mapping, header, validation)?;
    let mut cursor = header.graphs_offset;
    let mut graphs = Vec::new();
    graphs
        .try_reserve_exact(header.config.replicas)
        .map_err(|_| SearchError::AllocationFailed)?;
    for _ in 0..header.config.replicas {
        let (graph, next) = parse_mapped_graph(mapping, cursor, header.count)?;
        validate_graph(mapping, &graph, header.count)?;
        graphs.push(graph);
        cursor = next;
    }
    if cursor != header.file_len {
        return Err(SearchError::CorruptSnapshot(
            "snapshot has trailing or missing graph bytes",
        ));
    }
    Ok(graphs)
}

fn validate_vector_sections(
    mapping: &Mapping,
    header: &Header,
    validation: SnapshotValidation,
) -> Result<(), SearchError> {
    let count = header.count;
    let vector_count = count
        .checked_mul(header.config.dimensions)
        .ok_or(SearchError::CapacityOverflow)?;
    let keys = mapping
        .u64_slice(header.keys_offset, count)
        .map_err(SearchError::CorruptSnapshot)?;
    if keys.windows(2).any(|pair| pair[0] >= pair[1]) {
        return Err(SearchError::CorruptSnapshot(
            "mapped keys are not strictly increasing",
        ));
    }
    let vectors = mapping
        .f32_slice(header.vectors_offset, vector_count)
        .map_err(SearchError::CorruptSnapshot)?;
    if validation == SnapshotValidation::Full && vectors.iter().any(|value| !value.is_finite()) {
        return Err(SearchError::CorruptSnapshot(
            "mapped vectors contain a non-finite value",
        ));
    }
    let codes = mapping
        .u16_slice(header.routing_codes_offset, count)
        .map_err(SearchError::CorruptSnapshot)?;
    let nodes = mapping
        .u32_slice(header.routing_nodes_offset, count)
        .map_err(SearchError::CorruptSnapshot)?;
    if codes.windows(2).any(|pair| pair[0] > pair[1]) {
        return Err(SearchError::CorruptSnapshot("routing codes are not sorted"));
    }
    if nodes.iter().any(|node| *node as usize >= count) {
        return Err(SearchError::CorruptSnapshot(
            "routing node is outside the vector range",
        ));
    }
    Ok(())
}

fn parse_mapped_graph(
    mapping: &Mapping,
    cursor: usize,
    count: usize,
) -> Result<(MappedGraph, usize), SearchError> {
    let graph_end = cursor
        .checked_add(GRAPH_HEADER_LEN)
        .ok_or(SearchError::CapacityOverflow)?;
    if graph_end > mapping.len() {
        return Err(SearchError::CorruptSnapshot("graph header is truncated"));
    }
    let entry_raw = read_u64(mapping.bytes(), cursor)?;
    let max_level = to_usize(read_u64(mapping.bytes(), cursor + 8)?)?;
    let node_count = to_usize(read_u64(mapping.bytes(), cursor + 16)?)?;
    let total_layers = to_usize(read_u64(mapping.bytes(), cursor + 24)?)?;
    let neighbor_count = to_usize(read_u64(mapping.bytes(), cursor + 32)?)?;
    if node_count != count {
        return Err(SearchError::CorruptSnapshot(
            "graph node count does not match vector count",
        ));
    }
    let entry = parse_entry(entry_raw, count)?;
    let node_layers_offset = graph_end;
    let layer_neighbors_offset = node_layers_offset
        .checked_add(bytes_for::<u64>(count + 1)?)
        .ok_or(SearchError::CapacityOverflow)?;
    let neighbors_offset = layer_neighbors_offset
        .checked_add(bytes_for::<u64>(total_layers + 1)?)
        .ok_or(SearchError::CapacityOverflow)?;
    let next = align8(
        neighbors_offset
            .checked_add(bytes_for::<u32>(neighbor_count)?)
            .ok_or(SearchError::CapacityOverflow)?,
    )?;
    if next > mapping.len() {
        return Err(SearchError::CorruptSnapshot("graph sections are truncated"));
    }
    Ok((
        MappedGraph {
            entry,
            max_level,
            node_count,
            node_layers_offset,
            total_layers,
            layer_neighbors_offset,
            neighbors_offset,
            neighbor_count,
        },
        next,
    ))
}

fn parse_entry(raw: u64, count: usize) -> Result<Option<usize>, SearchError> {
    if raw == NONE_ENTRY {
        return Ok(None);
    }
    let entry = to_usize(raw)?;
    if entry >= count {
        return Err(SearchError::CorruptSnapshot(
            "graph entry is outside the vector range",
        ));
    }
    Ok(Some(entry))
}

fn validate_graph(mapping: &Mapping, graph: &MappedGraph, count: usize) -> Result<(), SearchError> {
    let node_layers = mapping
        .u64_slice(graph.node_layers_offset, count + 1)
        .map_err(SearchError::CorruptSnapshot)?;
    validate_mapped_offsets(
        node_layers,
        graph.total_layers,
        "graph node-layer offsets are invalid",
    )?;
    if count > 0 && node_layers.windows(2).any(|pair| pair[0] == pair[1]) {
        return Err(SearchError::CorruptSnapshot("graph node has no layer zero"));
    }
    let layer_neighbors = mapping
        .u64_slice(graph.layer_neighbors_offset, graph.total_layers + 1)
        .map_err(SearchError::CorruptSnapshot)?;
    validate_mapped_offsets(
        layer_neighbors,
        graph.neighbor_count,
        "graph neighbor offsets are invalid",
    )?;
    let neighbors = mapping
        .u32_slice(graph.neighbors_offset, graph.neighbor_count)
        .map_err(SearchError::CorruptSnapshot)?;
    if neighbors.iter().any(|node| *node as usize >= count) {
        return Err(SearchError::CorruptSnapshot(
            "graph neighbor is outside the vector range",
        ));
    }
    validate_neighbor_levels(node_layers, layer_neighbors, neighbors, count)?;
    validate_entry_levels(graph, node_layers, count)
}

fn validate_mapped_offsets(
    offsets: &[u64],
    last: usize,
    label: &'static str,
) -> Result<(), SearchError> {
    if offsets.first() != Some(&0)
        || offsets.last().copied() != Some(last as u64)
        || offsets.windows(2).any(|pair| pair[0] > pair[1])
    {
        return Err(SearchError::CorruptSnapshot(label));
    }
    Ok(())
}

fn validate_neighbor_levels(
    node_layers: &[u64],
    layer_neighbors: &[u64],
    neighbors: &[u32],
    count: usize,
) -> Result<(), SearchError> {
    for node in 0..count {
        let first_layer = to_usize(node_layers[node])?;
        let layer_end = to_usize(node_layers[node + 1])?;
        for (level, layer) in (first_layer..layer_end).enumerate() {
            let first_neighbor = to_usize(layer_neighbors[layer])?;
            let neighbor_end = to_usize(layer_neighbors[layer + 1])?;
            for neighbor in &neighbors[first_neighbor..neighbor_end] {
                let neighbor = *neighbor as usize;
                let neighbor_levels = node_layers[neighbor + 1] - node_layers[neighbor];
                if neighbor_levels <= level as u64 {
                    return Err(SearchError::CorruptSnapshot(
                        "graph layer points to a node without that layer",
                    ));
                }
            }
        }
    }
    Ok(())
}

fn validate_entry_levels(
    graph: &MappedGraph,
    node_layers: &[u64],
    count: usize,
) -> Result<(), SearchError> {
    match graph.entry {
        Some(entry) => {
            let levels = node_layers[entry + 1] - node_layers[entry];
            let levels = usize::try_from(levels).expect("validated level count");
            if levels == 0 || graph.max_level >= levels {
                return Err(SearchError::CorruptSnapshot(
                    "graph entry does not contain max_level",
                ));
            }
        }
        None if count != 0 => {
            return Err(SearchError::CorruptSnapshot("non-empty graph has no entry"));
        }
        None => {}
    }
    Ok(())
}