weavatrix-search-vector 0.3.1

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
use super::format::{align8, bytes_for, to_u64, update_checksum};
use super::{FNV_OFFSET, GRAPH_HEADER_LEN, HEADER_LEN, Header, NONE_ENTRY};
use crate::error::SearchError;
use crate::hnsw::{Graph, VectorIndex};
use std::io::{Seek, SeekFrom, Write};

pub(super) fn write_snapshot_stream<W: Write + Seek>(
    index: &VectorIndex,
    stream: &mut W,
) -> Result<(), SearchError> {
    let mut header = layout(index)?;
    stream
        .seek(SeekFrom::Start(0))
        .map_err(|error| SearchError::storage("seek snapshot stream", &error))?;
    stream
        .write_all(&header.encode()?)
        .map_err(|error| SearchError::storage("write header", &error))?;
    let checksum = write_payload(index, stream, &header)?;
    header.checksum = checksum;
    stream
        .seek(SeekFrom::Start(16))
        .map_err(|error| SearchError::storage("seek checksum", &error))?;
    stream
        .write_all(&checksum.to_le_bytes())
        .map_err(|error| SearchError::storage("write checksum", &error))?;
    stream
        .seek(SeekFrom::Start(to_u64(header.file_len)?))
        .map_err(|error| SearchError::storage("finish snapshot stream", &error))?;
    stream
        .flush()
        .map_err(|error| SearchError::storage("flush snapshot stream", &error))
}

fn write_payload<W: Write>(
    index: &VectorIndex,
    stream: &mut W,
    header: &Header,
) -> Result<u64, SearchError> {
    let mut writer = PayloadWriter::new(stream);
    writer.pad_to(header.keys_offset)?;
    writer.write_u64s(index.vectors().keys())?;
    writer.pad_to(header.vectors_offset)?;
    writer.write_f32s(index.vectors().values())?;
    writer.pad_to(header.routing_codes_offset)?;
    let routing_codes = index
        .routing()
        .entries
        .iter()
        .map(|entry| entry.0)
        .collect::<Vec<_>>();
    writer.write_u16s(&routing_codes)?;
    writer.pad_to(header.routing_nodes_offset)?;
    let routing_positions = index
        .routing()
        .entries
        .iter()
        .map(|entry| entry.1)
        .collect::<Vec<_>>();
    writer.write_u32s(&routing_positions)?;
    writer.pad_to(header.graphs_offset)?;
    for graph in index.graphs() {
        write_graph(&mut writer, graph, index.len())?;
    }
    if writer.position != header.file_len {
        return Err(SearchError::CorruptSnapshot(
            "calculated and written snapshot lengths differ",
        ));
    }
    Ok(writer.hash)
}

fn write_graph<W: Write>(
    writer: &mut PayloadWriter<'_, W>,
    graph: &Graph,
    count: usize,
) -> Result<(), SearchError> {
    let total_layers = graph
        .nodes
        .iter()
        .try_fold(0_usize, |total, node| total.checked_add(node.layers.len()))
        .ok_or(SearchError::CapacityOverflow)?;
    let neighbor_count = graph
        .nodes
        .iter()
        .try_fold(0_usize, |total, node| {
            node.layers
                .iter()
                .try_fold(total, |sum, layer| sum.checked_add(layer.len()))
        })
        .ok_or(SearchError::CapacityOverflow)?;
    writer.write_u64s(&[
        graph.entry.map_or(NONE_ENTRY, |entry| entry as u64),
        graph.max_level as u64,
        count as u64,
        total_layers as u64,
        neighbor_count as u64,
    ])?;
    let mut layer_total = 0_u64;
    writer.write_u64s(&[0])?;
    for node in &graph.nodes {
        layer_total = layer_total
            .checked_add(node.layers.len() as u64)
            .ok_or(SearchError::CapacityOverflow)?;
        writer.write_u64s(&[layer_total])?;
    }
    let mut neighbor_total = 0_u64;
    writer.write_u64s(&[0])?;
    for node in &graph.nodes {
        for layer in &node.layers {
            neighbor_total = neighbor_total
                .checked_add(layer.len() as u64)
                .ok_or(SearchError::CapacityOverflow)?;
            writer.write_u64s(&[neighbor_total])?;
        }
    }
    for node in &graph.nodes {
        for layer in &node.layers {
            writer.write_u32s(layer)?;
        }
    }
    writer.pad_to(align8(writer.position)?)
}

fn layout(index: &VectorIndex) -> Result<Header, SearchError> {
    let count = index.len();
    let vector_count = count
        .checked_mul(index.dimensions())
        .ok_or(SearchError::CapacityOverflow)?;
    let keys_offset = HEADER_LEN;
    let vectors_offset = align8(
        keys_offset
            .checked_add(bytes_for::<u64>(count)?)
            .ok_or(SearchError::CapacityOverflow)?,
    )?;
    let routing_codes_offset = align8(
        vectors_offset
            .checked_add(bytes_for::<f32>(vector_count)?)
            .ok_or(SearchError::CapacityOverflow)?,
    )?;
    let routing_nodes_offset = align8(
        routing_codes_offset
            .checked_add(bytes_for::<u16>(count)?)
            .ok_or(SearchError::CapacityOverflow)?,
    )?;
    let graphs_offset = align8(
        routing_nodes_offset
            .checked_add(bytes_for::<u32>(count)?)
            .ok_or(SearchError::CapacityOverflow)?,
    )?;
    let mut cursor = graphs_offset;
    for graph in index.graphs() {
        let total_layers = graph
            .nodes
            .iter()
            .try_fold(0_usize, |total, node| total.checked_add(node.layers.len()))
            .ok_or(SearchError::CapacityOverflow)?;
        let neighbor_count = graph
            .nodes
            .iter()
            .flat_map(|node| &node.layers)
            .try_fold(0_usize, |total, layer| total.checked_add(layer.len()))
            .ok_or(SearchError::CapacityOverflow)?;
        cursor = cursor
            .checked_add(GRAPH_HEADER_LEN)
            .and_then(|value| value.checked_add(bytes_for::<u64>(count + 1).ok()?))
            .and_then(|value| value.checked_add(bytes_for::<u64>(total_layers + 1).ok()?))
            .and_then(|value| value.checked_add(bytes_for::<u32>(neighbor_count).ok()?))
            .ok_or(SearchError::CapacityOverflow)?;
        cursor = align8(cursor)?;
    }
    Ok(Header {
        checksum: 0,
        file_len: cursor,
        count,
        config: index.config().clone(),
        keys_offset,
        vectors_offset,
        routing_codes_offset,
        routing_nodes_offset,
        graphs_offset,
    })
}

struct PayloadWriter<'a, W> {
    stream: &'a mut W,
    position: usize,
    hash: u64,
    buffer: Vec<u8>,
}

impl<'a, W: Write> PayloadWriter<'a, W> {
    fn new(stream: &'a mut W) -> Self {
        Self {
            stream,
            position: HEADER_LEN,
            hash: FNV_OFFSET,
            buffer: Vec::with_capacity(64 * 1024),
        }
    }

    fn pad_to(&mut self, offset: usize) -> Result<(), SearchError> {
        if offset < self.position {
            return Err(SearchError::CorruptSnapshot(
                "snapshot writer moved backwards",
            ));
        }
        let mut remaining = offset - self.position;
        let zeroes = [0_u8; 64];
        while remaining != 0 {
            let count = remaining.min(zeroes.len());
            self.write_bytes(&zeroes[..count])?;
            remaining -= count;
        }
        Ok(())
    }

    fn write_u16s(&mut self, values: &[u16]) -> Result<(), SearchError> {
        self.write_encoded(values, |value| value.to_le_bytes())
    }

    fn write_u32s(&mut self, values: &[u32]) -> Result<(), SearchError> {
        self.write_encoded(values, |value| value.to_le_bytes())
    }

    fn write_u64s(&mut self, values: &[u64]) -> Result<(), SearchError> {
        self.write_encoded(values, |value| value.to_le_bytes())
    }

    fn write_f32s(&mut self, values: &[f32]) -> Result<(), SearchError> {
        self.write_encoded(values, |value| value.to_le_bytes())
    }

    fn write_encoded<T, const N: usize>(
        &mut self,
        values: &[T],
        encode: impl Fn(&T) -> [u8; N],
    ) -> Result<(), SearchError> {
        let chunk_len = (64 * 1024 / N).max(1);
        for chunk in values.chunks(chunk_len) {
            self.buffer.clear();
            for value in chunk {
                self.buffer.extend_from_slice(&encode(value));
            }
            let bytes = std::mem::take(&mut self.buffer);
            self.write_bytes(&bytes)?;
            self.buffer = bytes;
        }
        Ok(())
    }

    fn write_bytes(&mut self, bytes: &[u8]) -> Result<(), SearchError> {
        self.stream
            .write_all(bytes)
            .map_err(|error| SearchError::storage("write payload", &error))?;
        self.hash = update_checksum(self.hash, bytes);
        self.position = self
            .position
            .checked_add(bytes.len())
            .ok_or(SearchError::CapacityOverflow)?;
        Ok(())
    }
}