use std::io::{Read, Write};
use crate::error::{LaurusError, Result};
use crate::vector::core::quantization::{PqParams, ScalarQuantParams};
pub const VECTOR_SEGMENT_MAGIC: [u8; 4] = *b"LVS1";
const PQ_PREFIX_SIZE: u64 = 24;
pub const CURRENT_VERSION: u16 = 1;
pub const VERSION_ORDINAL_GRAPH: u16 = 2;
pub const VERSION_FIELD_DICT: u16 = 3;
pub const MAX_SUPPORTED_VERSION: u16 = VERSION_FIELD_DICT;
pub mod quant_kind {
pub const NONE: u16 = 0;
pub const SCALAR_8BIT: u16 = 1;
pub const PRODUCT_QUANTIZATION: u16 = 2;
#[cfg(feature = "pq-fastscan")]
pub const PRODUCT_QUANTIZATION_FASTSCAN: u16 = 3;
}
pub const FIXED_HEADER_SIZE: usize = 16;
pub const SCALAR_8BIT_METADATA_SIZE: usize = 8;
pub const PQ_FIXED_METADATA_SIZE: usize = 8;
pub const PQ_CENTROIDS_PER_SUBVECTOR: u16 = 256;
#[derive(Debug, Clone, PartialEq)]
pub struct VectorSegmentHeader {
pub version: u16,
pub quant: QuantHeader,
pub field_dict: Vec<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum QuantHeader {
Scalar8Bit(ScalarQuantParams),
ProductQuantization {
params: PqParams,
codebook: Vec<f32>,
},
#[cfg(feature = "pq-fastscan")]
ProductQuantizationFastScan {
params: PqParams,
codebook: Vec<f32>,
},
}
impl QuantHeader {
pub fn kind_code(&self) -> u16 {
match self {
Self::Scalar8Bit(_) => quant_kind::SCALAR_8BIT,
Self::ProductQuantization { .. } => quant_kind::PRODUCT_QUANTIZATION,
#[cfg(feature = "pq-fastscan")]
Self::ProductQuantizationFastScan { .. } => quant_kind::PRODUCT_QUANTIZATION_FASTSCAN,
}
}
pub fn metadata_size(&self) -> usize {
match self {
Self::Scalar8Bit(_) => SCALAR_8BIT_METADATA_SIZE,
Self::ProductQuantization { params, .. } => {
PQ_FIXED_METADATA_SIZE + params.codebook_byte_size()
}
#[cfg(feature = "pq-fastscan")]
Self::ProductQuantizationFastScan { params, .. } => {
PQ_FIXED_METADATA_SIZE + params.codebook_byte_size()
}
}
}
}
impl VectorSegmentHeader {
pub fn scalar_8bit(params: ScalarQuantParams) -> Self {
Self {
version: CURRENT_VERSION,
quant: QuantHeader::Scalar8Bit(params),
field_dict: Vec::new(),
}
}
pub fn product_quantization(params: PqParams, codebook: Vec<f32>) -> Self {
debug_assert_eq!(codebook.len(), params.codebook_len());
Self {
version: CURRENT_VERSION,
quant: QuantHeader::ProductQuantization { params, codebook },
field_dict: Vec::new(),
}
}
#[cfg(feature = "pq-fastscan")]
pub fn product_quantization_fastscan(params: PqParams, codebook: Vec<f32>) -> Self {
debug_assert_eq!(params.k, 16, "FastScan requires k == 16");
debug_assert_eq!(codebook.len(), params.codebook_len());
Self {
version: CURRENT_VERSION,
quant: QuantHeader::ProductQuantizationFastScan { params, codebook },
field_dict: Vec::new(),
}
}
pub fn with_version(mut self, version: u16) -> Self {
self.version = version;
self
}
pub fn with_field_dict(mut self, field_dict: Vec<String>) -> Self {
self.field_dict = field_dict;
self
}
fn dict_size(&self) -> usize {
if self.version >= VERSION_FIELD_DICT {
2 + self
.field_dict
.iter()
.map(|name| 2 + name.len())
.sum::<usize>()
} else {
0
}
}
pub fn serialized_size(&self) -> usize {
FIXED_HEADER_SIZE + self.quant.metadata_size() + self.dict_size()
}
pub fn write_to<W: Write>(&self, writer: &mut W) -> Result<()> {
writer.write_all(&VECTOR_SEGMENT_MAGIC)?;
writer.write_all(&self.version.to_le_bytes())?;
writer.write_all(&self.quant.kind_code().to_le_bytes())?;
writer.write_all(&[0u8; 8])?; match &self.quant {
QuantHeader::Scalar8Bit(params) => {
writer.write_all(¶ms.offset.to_le_bytes())?;
writer.write_all(¶ms.scale.to_le_bytes())?;
}
QuantHeader::ProductQuantization { params, codebook } => {
writer.write_all(¶ms.m.to_le_bytes())?;
writer.write_all(¶ms.k.to_le_bytes())?;
writer.write_all(¶ms.sub_dim.to_le_bytes())?;
writer.write_all(&0u16.to_le_bytes())?; debug_assert_eq!(codebook.len(), params.codebook_len());
for &f in codebook {
writer.write_all(&f.to_le_bytes())?;
}
}
#[cfg(feature = "pq-fastscan")]
QuantHeader::ProductQuantizationFastScan { params, codebook } => {
writer.write_all(¶ms.m.to_le_bytes())?;
writer.write_all(¶ms.k.to_le_bytes())?;
writer.write_all(¶ms.sub_dim.to_le_bytes())?;
writer.write_all(&0u16.to_le_bytes())?; debug_assert_eq!(codebook.len(), params.codebook_len());
for &f in codebook {
writer.write_all(&f.to_le_bytes())?;
}
}
}
if self.version >= VERSION_FIELD_DICT {
if self.field_dict.len() > u16::MAX as usize {
return Err(LaurusError::InvalidOperation(format!(
"vector segment field dictionary has {} entries; the v3 \
format supports at most {}",
self.field_dict.len(),
u16::MAX
)));
}
writer.write_all(&(self.field_dict.len() as u16).to_le_bytes())?;
for name in &self.field_dict {
if name.len() > u16::MAX as usize {
return Err(LaurusError::InvalidOperation(format!(
"vector field name is {} bytes; the v3 format supports \
at most {} bytes per name",
name.len(),
u16::MAX
)));
}
writer.write_all(&(name.len() as u16).to_le_bytes())?;
writer.write_all(name.as_bytes())?;
}
} else {
debug_assert!(
self.field_dict.is_empty(),
"field_dict is only serialized at version >= VERSION_FIELD_DICT"
);
}
Ok(())
}
pub fn read_from<R: Read>(reader: &mut R, available: u64) -> Result<Self> {
let mut magic = [0u8; 4];
reader.read_exact(&mut magic)?;
if magic != VECTOR_SEGMENT_MAGIC {
return Err(LaurusError::IncompatibleFormat(format!(
"expected vector segment magic {:?} (\"LVS1\"), found {:?}. \
Pre-quantization (f32) segments must be rebuilt — \
Issue #481 Stage 1 introduced int8 scalar quantization \
with a new on-disk format.",
VECTOR_SEGMENT_MAGIC, magic
)));
}
let mut version_bytes = [0u8; 2];
reader.read_exact(&mut version_bytes)?;
let version = u16::from_le_bytes(version_bytes);
if !(CURRENT_VERSION..=MAX_SUPPORTED_VERSION).contains(&version) {
return Err(LaurusError::IncompatibleFormat(format!(
"unsupported vector segment header version {version} \
(this build supports {CURRENT_VERSION}..={MAX_SUPPORTED_VERSION})"
)));
}
let mut kind_bytes = [0u8; 2];
reader.read_exact(&mut kind_bytes)?;
let kind = u16::from_le_bytes(kind_bytes);
let mut reserved = [0u8; 8];
reader.read_exact(&mut reserved)?;
let quant = match kind {
quant_kind::SCALAR_8BIT => {
let mut offset_bytes = [0u8; 4];
let mut scale_bytes = [0u8; 4];
reader.read_exact(&mut offset_bytes)?;
reader.read_exact(&mut scale_bytes)?;
QuantHeader::Scalar8Bit(ScalarQuantParams {
offset: f32::from_le_bytes(offset_bytes),
scale: f32::from_le_bytes(scale_bytes),
})
}
quant_kind::PRODUCT_QUANTIZATION => {
let mut buf2 = [0u8; 2];
reader.read_exact(&mut buf2)?;
let m = u16::from_le_bytes(buf2);
reader.read_exact(&mut buf2)?;
let k = u16::from_le_bytes(buf2);
reader.read_exact(&mut buf2)?;
let sub_dim = u16::from_le_bytes(buf2);
reader.read_exact(&mut buf2)?; let params = PqParams::new(m, k, sub_dim).map_err(|e| {
LaurusError::IncompatibleFormat(format!("invalid PQ params: {e}"))
})?;
let codebook_len = params.codebook_len();
crate::vector::index::alloc_bounds::checked_capacity(
codebook_len,
4,
available.saturating_sub(PQ_PREFIX_SIZE),
"vector segment PQ codebook length",
)?;
let mut codebook = Vec::with_capacity(codebook_len);
let mut fbuf = [0u8; 4];
for _ in 0..codebook_len {
reader.read_exact(&mut fbuf)?;
codebook.push(f32::from_le_bytes(fbuf));
}
QuantHeader::ProductQuantization { params, codebook }
}
#[cfg(feature = "pq-fastscan")]
quant_kind::PRODUCT_QUANTIZATION_FASTSCAN => {
let mut buf2 = [0u8; 2];
reader.read_exact(&mut buf2)?;
let m = u16::from_le_bytes(buf2);
reader.read_exact(&mut buf2)?;
let k = u16::from_le_bytes(buf2);
reader.read_exact(&mut buf2)?;
let sub_dim = u16::from_le_bytes(buf2);
reader.read_exact(&mut buf2)?; let params = PqParams::new(m, k, sub_dim).map_err(|e| {
LaurusError::IncompatibleFormat(format!("invalid PQ FastScan params: {e}"))
})?;
if params.k != 16 {
return Err(LaurusError::IncompatibleFormat(format!(
"PQ FastScan segment must declare k == 16, got k = {}",
params.k
)));
}
let codebook_len = params.codebook_len();
crate::vector::index::alloc_bounds::checked_capacity(
codebook_len,
4,
available.saturating_sub(PQ_PREFIX_SIZE),
"vector segment PQ FastScan codebook length",
)?;
let mut codebook = Vec::with_capacity(codebook_len);
let mut fbuf = [0u8; 4];
for _ in 0..codebook_len {
reader.read_exact(&mut fbuf)?;
codebook.push(f32::from_le_bytes(fbuf));
}
QuantHeader::ProductQuantizationFastScan { params, codebook }
}
quant_kind::NONE => {
return Err(LaurusError::IncompatibleFormat(
"vector segment header reports quant_kind = 0 (no quantization), \
but the Rust-side QuantizationMethod has no None variant; \
this build cannot read unquantized segments"
.to_string(),
));
}
other => {
return Err(LaurusError::IncompatibleFormat(format!(
"unknown vector segment quant_kind = {other}"
)));
}
};
let field_dict = if version >= VERSION_FIELD_DICT {
let mut count_bytes = [0u8; 2];
reader.read_exact(&mut count_bytes)?;
let count = u16::from_le_bytes(count_bytes) as usize;
let mut dict = Vec::with_capacity(count);
for _ in 0..count {
let mut len_bytes = [0u8; 2];
reader.read_exact(&mut len_bytes)?;
let len = u16::from_le_bytes(len_bytes) as usize;
let mut name_bytes = vec![0u8; len];
reader.read_exact(&mut name_bytes)?;
let name = String::from_utf8(name_bytes).map_err(|e| {
LaurusError::IncompatibleFormat(format!(
"invalid UTF-8 in vector segment field dictionary: {e}"
))
})?;
dict.push(name);
}
dict
} else {
Vec::new()
};
Ok(Self {
version,
quant,
field_dict,
})
}
pub(crate) fn read_record_field<R: Read>(
&self,
reader: &mut R,
records_remaining: u64,
what: &str,
) -> Result<String> {
if self.version >= VERSION_FIELD_DICT {
let mut id_bytes = [0u8; 2];
reader.read_exact(&mut id_bytes)?;
let field_id = u16::from_le_bytes(id_bytes);
self.field_dict
.get(field_id as usize)
.cloned()
.ok_or_else(|| {
LaurusError::index(format!(
"vector segment corrupt: record field_id {field_id} out \
of dictionary range ({} entries)",
self.field_dict.len()
))
})
} else {
use crate::vector::index::alloc_bounds::checked_len;
let mut len_bytes = [0u8; 4];
reader.read_exact(&mut len_bytes)?;
let len = u32::from_le_bytes(len_bytes) as usize;
checked_len(len, records_remaining, what)?;
let mut name_bytes = vec![0u8; len];
reader.read_exact(&mut name_bytes)?;
String::from_utf8(name_bytes).map_err(|e| {
LaurusError::InvalidOperation(format!("Invalid UTF-8 in field name: {}", e))
})
}
}
}
pub(crate) struct FieldInterner {
names: Vec<std::sync::Arc<str>>,
fixed: bool,
}
impl FieldInterner {
pub(crate) fn from_header(header: &VectorSegmentHeader) -> Self {
if header.version >= VERSION_FIELD_DICT {
Self {
names: header
.field_dict
.iter()
.map(|s| std::sync::Arc::from(s.as_str()))
.collect(),
fixed: true,
}
} else {
Self {
names: Vec::new(),
fixed: false,
}
}
}
pub(crate) fn read_record_field_id<R: Read>(
&mut self,
header: &VectorSegmentHeader,
reader: &mut R,
records_remaining: u64,
what: &str,
) -> Result<u16> {
if self.fixed {
let mut id_bytes = [0u8; 2];
reader.read_exact(&mut id_bytes)?;
let field_id = u16::from_le_bytes(id_bytes);
if field_id as usize >= self.names.len() {
return Err(LaurusError::index(format!(
"vector segment corrupt: record field_id {field_id} out of \
dictionary range ({} entries)",
self.names.len()
)));
}
Ok(field_id)
} else {
let name = header.read_record_field(reader, records_remaining, what)?;
if let Some(pos) = self.names.iter().position(|n| **n == *name) {
return Ok(pos as u16);
}
if self.names.len() >= u16::MAX as usize {
return Err(LaurusError::index(format!(
"vector segment has more than {} distinct field names",
u16::MAX
)));
}
self.names.push(std::sync::Arc::from(name.as_str()));
Ok((self.names.len() - 1) as u16)
}
}
pub(crate) fn name(&self, fid: u16) -> &std::sync::Arc<str> {
&self.names[fid as usize]
}
pub(crate) fn into_dict(self) -> std::sync::Arc<[std::sync::Arc<str>]> {
std::sync::Arc::from(self.names)
}
}
#[inline]
pub(crate) fn resolve_field_id(dict: &[std::sync::Arc<str>], field_name: &str) -> Option<u16> {
dict.iter()
.position(|n| **n == *field_name)
.map(|pos| pos as u16)
}
pub(crate) fn record_prefix_size(version: u16) -> u64 {
if version >= VERSION_FIELD_DICT {
10
} else {
12
}
}
pub(crate) fn build_field_dict<'a>(
names: impl Iterator<Item = &'a str>,
) -> Result<(Vec<String>, std::collections::HashMap<String, u16>)> {
let mut dict: Vec<String> = Vec::new();
let mut ids: std::collections::HashMap<String, u16> = std::collections::HashMap::new();
for name in names {
if !ids.contains_key(name) {
if dict.len() >= u16::MAX as usize {
return Err(LaurusError::InvalidOperation(format!(
"vector segment has more than {} distinct field names; \
the v3 format supports at most that many",
u16::MAX
)));
}
ids.insert(name.to_string(), dict.len() as u16);
dict.push(name.to_string());
}
}
Ok((dict, ids))
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
fn sample_params() -> ScalarQuantParams {
ScalarQuantParams {
offset: -1.5,
scale: 3.0 / 255.0,
}
}
#[test]
fn roundtrip_scalar_8bit_header() {
let header = VectorSegmentHeader::scalar_8bit(sample_params());
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
assert_eq!(buf.len(), header.serialized_size());
assert_eq!(buf.len(), FIXED_HEADER_SIZE + SCALAR_8BIT_METADATA_SIZE);
let mut cursor = Cursor::new(&buf);
let parsed = VectorSegmentHeader::read_from(&mut cursor, buf.len() as u64).unwrap();
assert_eq!(parsed, header);
}
#[test]
fn roundtrip_v2_ordinal_graph_header() {
let header =
VectorSegmentHeader::scalar_8bit(sample_params()).with_version(VERSION_ORDINAL_GRAPH);
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
assert_eq!(
u16::from_le_bytes([buf[4], buf[5]]),
VERSION_ORDINAL_GRAPH,
"with_version must stamp the version bytes at offset 4"
);
let mut cursor = Cursor::new(&buf);
let parsed = VectorSegmentHeader::read_from(&mut cursor, buf.len() as u64).unwrap();
assert_eq!(parsed.version, VERSION_ORDINAL_GRAPH);
assert_eq!(parsed, header);
}
#[test]
fn rejects_header_version_above_max_supported() {
let header = VectorSegmentHeader::scalar_8bit(sample_params())
.with_version(MAX_SUPPORTED_VERSION + 1);
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
let mut cursor = Cursor::new(&buf);
let err = VectorSegmentHeader::read_from(&mut cursor, buf.len() as u64).unwrap_err();
assert!(
err.to_string().contains("unsupported vector segment"),
"unexpected error: {err}"
);
}
#[test]
fn roundtrip_v3_header_with_field_dict() {
for dict in [
Vec::new(),
vec!["embedding".to_string()],
vec![
"a".to_string(),
"b".to_string(),
"long_field_name".to_string(),
],
] {
let header = VectorSegmentHeader::scalar_8bit(sample_params())
.with_version(VERSION_FIELD_DICT)
.with_field_dict(dict.clone());
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
assert_eq!(
buf.len(),
header.serialized_size(),
"serialized_size must match written bytes for dict {dict:?}"
);
let mut cursor = Cursor::new(&buf);
let parsed = VectorSegmentHeader::read_from(&mut cursor, buf.len() as u64).unwrap();
assert_eq!(parsed.field_dict, dict);
assert_eq!(parsed, header);
}
}
#[test]
fn v1_header_serialization_is_unchanged_by_the_dict_field() {
let header = VectorSegmentHeader::scalar_8bit(sample_params());
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
assert_eq!(buf.len(), 24, "v1 SQ header must stay 24 bytes");
assert_eq!(header.serialized_size(), 24);
let parsed =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap();
assert!(parsed.field_dict.is_empty());
}
#[test]
fn truncated_field_dict_returns_io_error() {
let header = VectorSegmentHeader::scalar_8bit(sample_params())
.with_version(VERSION_FIELD_DICT)
.with_field_dict(vec!["embedding".to_string()]);
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
buf.truncate(buf.len() - 3);
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
assert!(matches!(err, LaurusError::Io(_)), "got: {err:?}");
}
#[test]
fn write_rejects_field_name_longer_than_u16_max() {
let header = VectorSegmentHeader::scalar_8bit(sample_params())
.with_version(VERSION_FIELD_DICT)
.with_field_dict(vec!["x".repeat(u16::MAX as usize + 1)]);
let mut buf: Vec<u8> = Vec::new();
let err = header.write_to(&mut buf).unwrap_err();
assert!(err.to_string().contains("at most"), "got: {err}");
}
#[test]
fn read_record_field_resolves_v3_ids_and_rejects_out_of_range() {
let header = VectorSegmentHeader::scalar_8bit(sample_params())
.with_version(VERSION_FIELD_DICT)
.with_field_dict(vec!["a".to_string(), "b".to_string()]);
let mut ok = Cursor::new(1u16.to_le_bytes().to_vec());
assert_eq!(header.read_record_field(&mut ok, 0, "test").unwrap(), "b");
let mut bad = Cursor::new(7u16.to_le_bytes().to_vec());
let err = header.read_record_field(&mut bad, 0, "test").unwrap_err();
assert!(
err.to_string().contains("out of dictionary range"),
"got: {err}"
);
}
#[test]
fn read_record_field_parses_inline_names_below_v3() {
let header = VectorSegmentHeader::scalar_8bit(sample_params());
let mut bytes = (5u32).to_le_bytes().to_vec();
bytes.extend_from_slice(b"field");
let mut cursor = Cursor::new(bytes);
assert_eq!(
header.read_record_field(&mut cursor, 64, "test").unwrap(),
"field"
);
}
#[test]
fn interner_synthesizes_dict_for_legacy_versions() {
let header = VectorSegmentHeader::scalar_8bit(sample_params());
let mut interner = FieldInterner::from_header(&header);
let mut rec = |name: &str| {
let mut bytes = (name.len() as u32).to_le_bytes().to_vec();
bytes.extend_from_slice(name.as_bytes());
let mut cursor = Cursor::new(bytes);
interner
.read_record_field_id(&header, &mut cursor, 64, "test")
.unwrap()
};
assert_eq!(rec("b"), 0);
assert_eq!(rec("a"), 1);
assert_eq!(rec("b"), 0, "repeat name must intern to the same id");
let dict = interner.into_dict();
assert_eq!(dict.len(), 2);
assert_eq!(&*dict[0], "b");
assert_eq!(&*dict[1], "a");
}
#[test]
fn interner_uses_fixed_dict_for_v3_and_rejects_out_of_range() {
let header = VectorSegmentHeader::scalar_8bit(sample_params())
.with_version(VERSION_FIELD_DICT)
.with_field_dict(vec!["f".to_string()]);
let mut interner = FieldInterner::from_header(&header);
let mut ok = Cursor::new(0u16.to_le_bytes().to_vec());
assert_eq!(
interner
.read_record_field_id(&header, &mut ok, 0, "test")
.unwrap(),
0
);
let mut bad = Cursor::new(3u16.to_le_bytes().to_vec());
let err = interner
.read_record_field_id(&header, &mut bad, 0, "test")
.unwrap_err();
assert!(err.to_string().contains("out of"), "got: {err}");
}
#[test]
fn resolve_field_id_scans_the_dict() {
let dict: Vec<std::sync::Arc<str>> =
vec![std::sync::Arc::from("a"), std::sync::Arc::from("b")];
assert_eq!(resolve_field_id(&dict, "b"), Some(1));
assert_eq!(resolve_field_id(&dict, "missing"), None);
}
#[test]
fn record_prefix_size_matches_the_version_ladder() {
assert_eq!(record_prefix_size(CURRENT_VERSION), 12);
assert_eq!(record_prefix_size(VERSION_ORDINAL_GRAPH), 12);
assert_eq!(record_prefix_size(VERSION_FIELD_DICT), 10);
}
#[test]
fn build_field_dict_assigns_first_appearance_order() {
let names = ["b", "a", "b", "c", "a"];
let (dict, ids) = build_field_dict(names.into_iter()).unwrap();
assert_eq!(
dict,
vec!["b".to_string(), "a".to_string(), "c".to_string()]
);
assert_eq!(ids["b"], 0);
assert_eq!(ids["a"], 1);
assert_eq!(ids["c"], 2);
}
#[test]
fn header_starts_with_magic_then_version_then_kind() {
let header = VectorSegmentHeader::scalar_8bit(sample_params());
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
assert_eq!(&buf[0..4], b"LVS1");
assert_eq!(u16::from_le_bytes([buf[4], buf[5]]), CURRENT_VERSION);
assert_eq!(
u16::from_le_bytes([buf[6], buf[7]]),
quant_kind::SCALAR_8BIT
);
assert_eq!(&buf[8..16], &[0u8; 8]);
}
#[test]
fn missing_magic_returns_incompatible_format() {
let f32_bytes = [
0xCD, 0xCC, 0x4C, 0x3F, 0x00, 0x00, 0x80, 0x3F, 0x00, 0x00, 0x00, 0x40, 0x00, 0x00, 0x40, 0x40, ];
let mut cursor = Cursor::new(&f32_bytes[..]);
let err = VectorSegmentHeader::read_from(&mut cursor, f32_bytes.len() as u64).unwrap_err();
match err {
LaurusError::IncompatibleFormat(msg) => {
assert!(msg.contains("LVS1"), "message should mention LVS1");
assert!(
msg.contains("rebuilt"),
"message should instruct to rebuild"
);
}
other => panic!("expected IncompatibleFormat, got {other:?}"),
}
}
#[test]
fn unsupported_version_returns_incompatible_format() {
let mut buf = Vec::new();
buf.extend_from_slice(b"LVS1");
buf.extend_from_slice(&99u16.to_le_bytes()); buf.extend_from_slice(&quant_kind::SCALAR_8BIT.to_le_bytes());
buf.extend_from_slice(&[0u8; 8]);
buf.extend_from_slice(&0.0_f32.to_le_bytes());
buf.extend_from_slice(&1.0_f32.to_le_bytes());
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
match err {
LaurusError::IncompatibleFormat(msg) => {
assert!(msg.contains("99"), "message should mention the version");
}
other => panic!("expected IncompatibleFormat, got {other:?}"),
}
}
#[test]
fn quant_kind_zero_is_rejected_for_now() {
let mut buf = Vec::new();
buf.extend_from_slice(b"LVS1");
buf.extend_from_slice(&CURRENT_VERSION.to_le_bytes());
buf.extend_from_slice(&quant_kind::NONE.to_le_bytes());
buf.extend_from_slice(&[0u8; 8]);
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
assert!(matches!(err, LaurusError::IncompatibleFormat(_)));
}
#[test]
fn oversized_pq_codebook_declaration_is_rejected_cleanly() {
let mut buf = Vec::new();
buf.extend_from_slice(b"LVS1");
buf.extend_from_slice(&CURRENT_VERSION.to_le_bytes());
buf.extend_from_slice(&quant_kind::PRODUCT_QUANTIZATION.to_le_bytes());
buf.extend_from_slice(&[0u8; 8]);
buf.extend_from_slice(&u16::MAX.to_le_bytes());
buf.extend_from_slice(&256u16.to_le_bytes());
buf.extend_from_slice(&u16::MAX.to_le_bytes());
buf.extend_from_slice(&0u16.to_le_bytes());
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
match err {
LaurusError::Index(msg) => {
assert!(msg.contains("corrupted"), "corruption must be named: {msg}");
}
other => panic!("expected LaurusError::Index, got {other:?}"),
}
}
#[cfg(feature = "pq-fastscan")]
#[test]
fn oversized_pq_fastscan_codebook_declaration_is_rejected_cleanly() {
let mut buf = Vec::new();
buf.extend_from_slice(b"LVS1");
buf.extend_from_slice(&CURRENT_VERSION.to_le_bytes());
buf.extend_from_slice(&quant_kind::PRODUCT_QUANTIZATION_FASTSCAN.to_le_bytes());
buf.extend_from_slice(&[0u8; 8]);
buf.extend_from_slice(&u16::MAX.to_le_bytes());
buf.extend_from_slice(&16u16.to_le_bytes());
buf.extend_from_slice(&u16::MAX.to_le_bytes());
buf.extend_from_slice(&0u16.to_le_bytes());
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
match err {
LaurusError::Index(msg) => {
assert!(msg.contains("corrupted"), "corruption must be named: {msg}");
}
other => panic!("expected LaurusError::Index, got {other:?}"),
}
}
fn sample_pq_params() -> PqParams {
PqParams::new(4, 256, 2).expect("valid PQ params")
}
fn sample_pq_codebook(params: &PqParams) -> Vec<f32> {
let mut cb = Vec::with_capacity(params.codebook_len());
for m in 0..params.m as usize {
for k in 0..params.k as usize {
for d in 0..params.sub_dim as usize {
cb.push((m * 100 + k) as f32 + d as f32 * 0.01);
}
}
}
cb
}
#[test]
fn roundtrip_product_quantization_header() {
let params = sample_pq_params();
let codebook = sample_pq_codebook(¶ms);
let header = VectorSegmentHeader::product_quantization(params, codebook);
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
assert_eq!(buf.len(), header.serialized_size());
assert_eq!(
buf.len(),
FIXED_HEADER_SIZE + PQ_FIXED_METADATA_SIZE + 4 * 256 * 2 * 4
);
let parsed =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap();
assert_eq!(parsed, header);
}
#[test]
fn pq_header_serialised_starts_with_magic_and_kind_two() {
let params = sample_pq_params();
let codebook = sample_pq_codebook(¶ms);
let header = VectorSegmentHeader::product_quantization(params, codebook);
let mut buf: Vec<u8> = Vec::new();
header.write_to(&mut buf).unwrap();
assert_eq!(&buf[0..4], b"LVS1");
assert_eq!(u16::from_le_bytes([buf[4], buf[5]]), CURRENT_VERSION);
assert_eq!(
u16::from_le_bytes([buf[6], buf[7]]),
quant_kind::PRODUCT_QUANTIZATION
);
assert_eq!(u16::from_le_bytes([buf[16], buf[17]]), 4); assert_eq!(u16::from_le_bytes([buf[18], buf[19]]), 256); assert_eq!(u16::from_le_bytes([buf[20], buf[21]]), 2); assert_eq!(u16::from_le_bytes([buf[22], buf[23]]), 0);
}
#[test]
fn pq_header_metadata_size_matches_codebook() {
let params = sample_pq_params();
let header = QuantHeader::ProductQuantization {
params,
codebook: vec![0.0; params.codebook_len()],
};
assert_eq!(
header.metadata_size(),
PQ_FIXED_METADATA_SIZE + params.codebook_byte_size()
);
}
#[test]
fn pq_header_rejects_invalid_params() {
let mut buf = Vec::new();
buf.extend_from_slice(b"LVS1");
buf.extend_from_slice(&CURRENT_VERSION.to_le_bytes());
buf.extend_from_slice(&quant_kind::PRODUCT_QUANTIZATION.to_le_bytes());
buf.extend_from_slice(&[0u8; 8]);
buf.extend_from_slice(&0u16.to_le_bytes()); buf.extend_from_slice(&256u16.to_le_bytes());
buf.extend_from_slice(&2u16.to_le_bytes());
buf.extend_from_slice(&0u16.to_le_bytes());
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
assert!(matches!(err, LaurusError::IncompatibleFormat(_)));
}
#[test]
fn unknown_quant_kind_is_rejected() {
let mut buf = Vec::new();
buf.extend_from_slice(b"LVS1");
buf.extend_from_slice(&CURRENT_VERSION.to_le_bytes());
buf.extend_from_slice(&999u16.to_le_bytes()); buf.extend_from_slice(&[0u8; 8]);
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
match err {
LaurusError::IncompatibleFormat(msg) => {
assert!(msg.contains("999"));
}
other => panic!("expected IncompatibleFormat, got {other:?}"),
}
}
#[test]
fn truncated_header_returns_io_error() {
let buf = b"LVS1".to_vec();
let err =
VectorSegmentHeader::read_from(&mut Cursor::new(&buf), buf.len() as u64).unwrap_err();
assert!(matches!(err, LaurusError::Io(_)));
}
#[test]
fn serialized_size_constants_are_consistent() {
let h = VectorSegmentHeader::scalar_8bit(sample_params());
assert_eq!(h.serialized_size(), 24);
assert_eq!(FIXED_HEADER_SIZE + SCALAR_8BIT_METADATA_SIZE, 24);
}
#[test]
fn fixed_header_size_is_sixteen() {
assert_eq!(FIXED_HEADER_SIZE, 16);
}
#[test]
fn quant_header_kind_code_matches_constants() {
let h = QuantHeader::Scalar8Bit(sample_params());
assert_eq!(h.kind_code(), quant_kind::SCALAR_8BIT);
assert_eq!(h.metadata_size(), SCALAR_8BIT_METADATA_SIZE);
}
}