weavatrix-search-vector 0.3.1

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
use super::codec::Decoder;
use super::format::{HEADER_LEN, MAGIC, VERSION, decode_value, read_u32, read_u64, to_usize};
use super::io::checksum;
use super::validation::{validate_prepared, validate_quantized};
use crate::error::SearchError;
use crate::hnsw::VectorIndex;
use crate::metadata::{Metadata, MetadataIndex};
use crate::mutable::{MutableSnapshot, MutableVectorIndex};
use crate::quantized::QuantizedIndex;
use std::collections::{BTreeMap, BTreeSet};
use std::path::Path;
use std::sync::Arc;

pub(super) fn load_complete(
    path: &Path,
) -> Result<(MutableVectorIndex, Option<QuantizedIndex>), SearchError> {
    let bytes = std::fs::read(path).map_err(|error| SearchError::storage("read bundle", &error))?;
    if bytes.get(..8) != Some(MAGIC.as_slice()) {
        return Ok((
            MutableVectorIndex::from_index(VectorIndex::load(path)?),
            None,
        ));
    }
    validate_header(&bytes)?;
    let base_len = to_usize(read_u64(&bytes, 32)?)?;
    let sealed_len = to_usize(read_u64(&bytes, 40)?)?;
    let pending_count = to_usize(read_u64(&bytes, 48)?)?;
    let deleted_count = to_usize(read_u64(&bytes, 56)?)?;
    let metadata_count = to_usize(read_u64(&bytes, 64)?)?;
    let quantized_len = if read_u32(&bytes, 8)? >= 2 {
        to_usize(read_u64(&bytes, 72)?)?
    } else {
        0
    };
    let mut decoder = Decoder::new(&bytes, HEADER_LEN);
    let base = Arc::new(VectorIndex::from_bytes(decoder.take(base_len)?)?);
    let sealed = decode_sealed(&mut decoder, sealed_len, &base)?;
    let dimensions = base.dimensions();
    let pending = decode_pending(&mut decoder, pending_count, dimensions, &base)?;
    let deleted = decode_deleted(&mut decoder, deleted_count, &pending)?;
    let metadata = decode_metadata(&mut decoder, metadata_count)?;
    let quantized = if quantized_len == 0 {
        None
    } else {
        Some(crate::quantized_io::decode(decoder.take(quantized_len)?)?)
    };
    if !decoder.is_finished() {
        return Err(SearchError::CorruptSnapshot(
            "bundle contains trailing payload bytes",
        ));
    }
    let mutable = MutableVectorIndex::from_snapshot(MutableSnapshot {
        config: base.config().clone(),
        base,
        sealed,
        pending,
        deleted,
        metadata,
    });
    validate_quantized(&mutable.snapshot(), quantized.as_ref())?;
    Ok((mutable, quantized))
}

fn validate_header(bytes: &[u8]) -> Result<(), SearchError> {
    if bytes.len() < HEADER_LEN {
        return Err(SearchError::CorruptSnapshot("bundle header is truncated"));
    }
    let version = read_u32(bytes, 8)?;
    if version != 1 && version != VERSION {
        return Err(SearchError::UnsupportedSnapshotVersion(version));
    }
    if read_u32(bytes, 12)? as usize != HEADER_LEN {
        return Err(SearchError::CorruptSnapshot(
            "bundle header length does not match",
        ));
    }
    if to_usize(read_u64(bytes, 24)?)? != bytes.len() {
        return Err(SearchError::CorruptSnapshot(
            "bundle file length does not match",
        ));
    }
    if checksum(&bytes[HEADER_LEN..]) != read_u64(bytes, 16)? {
        return Err(SearchError::CorruptSnapshot(
            "bundle payload checksum does not match",
        ));
    }
    Ok(())
}

fn decode_sealed(
    decoder: &mut Decoder<'_>,
    sealed_len: usize,
    base: &VectorIndex,
) -> Result<Option<Arc<VectorIndex>>, SearchError> {
    let sealed = if sealed_len == 0 {
        None
    } else {
        Some(Arc::new(VectorIndex::from_bytes(
            decoder.take(sealed_len)?,
        )?))
    };
    if sealed
        .as_ref()
        .is_some_and(|index| index.config() != base.config())
    {
        return Err(SearchError::CorruptSnapshot(
            "sealed delta config differs from base",
        ));
    }
    Ok(sealed)
}

fn decode_pending(
    decoder: &mut Decoder<'_>,
    count: usize,
    dimensions: usize,
    base: &VectorIndex,
) -> Result<BTreeMap<u64, Vec<f32>>, SearchError> {
    let mut pending = BTreeMap::new();
    for position in 0..count {
        let key = decoder.u64()?;
        let vector = decoder.f32s(dimensions)?;
        validate_prepared(base.config().metric, &vector, position)?;
        if pending.insert(key, vector).is_some() {
            return Err(SearchError::CorruptSnapshot(
                "bundle pending keys are duplicated",
            ));
        }
    }
    Ok(pending)
}

fn decode_deleted(
    decoder: &mut Decoder<'_>,
    count: usize,
    pending: &BTreeMap<u64, Vec<f32>>,
) -> Result<BTreeSet<u64>, SearchError> {
    let mut deleted = BTreeSet::new();
    for _ in 0..count {
        if !deleted.insert(decoder.u64()?) {
            return Err(SearchError::CorruptSnapshot(
                "bundle tombstones are duplicated",
            ));
        }
    }
    if pending.keys().any(|key| deleted.contains(key)) {
        return Err(SearchError::CorruptSnapshot(
            "bundle key is both pending and deleted",
        ));
    }
    Ok(deleted)
}

fn decode_metadata(decoder: &mut Decoder<'_>, count: usize) -> Result<MetadataIndex, SearchError> {
    let mut metadata = MetadataIndex::new();
    for _ in 0..count {
        let key = decoder.u64()?;
        let field_count = decoder.u32()? as usize;
        let mut record = Metadata::new();
        for _ in 0..field_count {
            let field = String::from_utf8(decoder.length_prefixed()?.to_vec())
                .map_err(|_| SearchError::CorruptSnapshot("metadata field is not UTF-8"))?;
            if record.insert(field, decode_value(decoder)?).is_some() {
                return Err(SearchError::CorruptSnapshot("metadata field is duplicated"));
            }
        }
        if metadata.insert(key, record).is_some() {
            return Err(SearchError::CorruptSnapshot(
                "metadata record key is duplicated",
            ));
        }
    }
    Ok(metadata)
}