use super::format::{put_u32, put_u64, read_u32, read_u64, to_u64, to_usize};
use super::{FORMAT_VERSION, HEADER_LEN, Header, MAGIC};
use crate::config::{DistanceMetric, IndexConfig};
use crate::error::SearchError;
impl Header {
pub(super) fn parse(bytes: &[u8]) -> Result<Self, SearchError> {
if bytes.len() < HEADER_LEN {
return Err(SearchError::CorruptSnapshot("header is truncated"));
}
if &bytes[..8] != MAGIC {
return Err(SearchError::CorruptSnapshot("magic does not match"));
}
let version = read_u32(bytes, 8)?;
if version != FORMAT_VERSION {
return Err(SearchError::UnsupportedSnapshotVersion(version));
}
if read_u32(bytes, 12)? as usize != HEADER_LEN {
return Err(SearchError::CorruptSnapshot("header length does not match"));
}
let metric = match read_u64(bytes, 48)? {
0 => DistanceMetric::Cosine,
1 => DistanceMetric::Dot,
2 => DistanceMetric::SquaredEuclidean,
_ => return Err(SearchError::CorruptSnapshot("unknown distance metric")),
};
let config = IndexConfig {
dimensions: to_usize(read_u64(bytes, 32)?)?,
metric,
connectivity: to_usize(read_u64(bytes, 56)?)?,
expansion_build: to_usize(read_u64(bytes, 64)?)?,
expansion_query: to_usize(read_u64(bytes, 72)?)?,
replicas: to_usize(read_u64(bytes, 80)?)?,
build_threads: to_usize(read_u64(bytes, 88)?)?,
query_threads: to_usize(read_u64(bytes, 96)?)?,
seed: read_u64(bytes, 104)?,
};
config.validate()?;
Ok(Self {
checksum: read_u64(bytes, 16)?,
file_len: to_usize(read_u64(bytes, 24)?)?,
count: to_usize(read_u64(bytes, 40)?)?,
config,
keys_offset: to_usize(read_u64(bytes, 112)?)?,
vectors_offset: to_usize(read_u64(bytes, 120)?)?,
routing_codes_offset: to_usize(read_u64(bytes, 128)?)?,
routing_nodes_offset: to_usize(read_u64(bytes, 136)?)?,
graphs_offset: to_usize(read_u64(bytes, 144)?)?,
})
}
pub(super) fn encode(&self) -> Result<[u8; HEADER_LEN], SearchError> {
let mut bytes = [0_u8; HEADER_LEN];
bytes[..8].copy_from_slice(MAGIC);
put_u32(&mut bytes, 8, FORMAT_VERSION);
put_u32(
&mut bytes,
12,
u32::try_from(HEADER_LEN).expect("fixed header length fits u32"),
);
put_u64(&mut bytes, 16, self.checksum);
put_u64(&mut bytes, 24, to_u64(self.file_len)?);
put_u64(&mut bytes, 32, to_u64(self.config.dimensions)?);
put_u64(&mut bytes, 40, to_u64(self.count)?);
put_u64(
&mut bytes,
48,
match self.config.metric {
DistanceMetric::Cosine => 0,
DistanceMetric::Dot => 1,
DistanceMetric::SquaredEuclidean => 2,
},
);
put_u64(&mut bytes, 56, to_u64(self.config.connectivity)?);
put_u64(&mut bytes, 64, to_u64(self.config.expansion_build)?);
put_u64(&mut bytes, 72, to_u64(self.config.expansion_query)?);
put_u64(&mut bytes, 80, to_u64(self.config.replicas)?);
put_u64(&mut bytes, 88, to_u64(self.config.build_threads)?);
put_u64(&mut bytes, 96, to_u64(self.config.query_threads)?);
put_u64(&mut bytes, 104, self.config.seed);
put_u64(&mut bytes, 112, to_u64(self.keys_offset)?);
put_u64(&mut bytes, 120, to_u64(self.vectors_offset)?);
put_u64(&mut bytes, 128, to_u64(self.routing_codes_offset)?);
put_u64(&mut bytes, 136, to_u64(self.routing_nodes_offset)?);
put_u64(&mut bytes, 144, to_u64(self.graphs_offset)?);
Ok(bytes)
}
}