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