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)
}