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