use std::io::{Read, Write};
use crate::error::{LaurusError, Result};
use crate::storage::Storage;
use crate::storage::checksum::{CrcReader, CrcWriter};
use crate::vector::core::quantization::{PqParams, pq_train_codebook};
use crate::vector::core::vector::Vector;
use crate::vector::index::format::{QuantHeader, VectorSegmentHeader};
const PQ_CODEBOOK_FOOTER_MAGIC: u32 = 0x5051_4342;
const FOOTER_SIZE: usize = 8;
#[cfg(test)]
const PQ_HEADER_PREFIX_SIZE: usize = 24;
#[derive(Debug, Clone)]
pub struct SharedPqCodebook {
pub params: PqParams,
pub codebook: Vec<f32>,
}
impl SharedPqCodebook {
pub fn validate_for(
&self,
dimension: usize,
subvector_count: usize,
expected_k: u16,
) -> Result<()> {
if self.params.original_dim() != dimension {
return Err(LaurusError::InvalidOperation(format!(
"shared PQ codebook dimension {} does not match the configured \
dimension {dimension}",
self.params.original_dim()
)));
}
if self.params.m as usize != subvector_count {
return Err(LaurusError::InvalidOperation(format!(
"shared PQ codebook subvector_count {} does not match the configured \
subvector_count {subvector_count}",
self.params.m
)));
}
if self.params.k != expected_k {
return Err(LaurusError::InvalidOperation(format!(
"shared PQ codebook has k = {} centroids per sub-quantizer but the \
field's quantizer variant requires k = {expected_k} (16 = FastScan, \
256 = standard PQ); retrain the codebook for this variant",
self.params.k
)));
}
if self.codebook.len() != self.params.codebook_len() {
return Err(LaurusError::InvalidOperation(format!(
"shared PQ codebook has {} entries, expected {} for params {:?}",
self.codebook.len(),
self.params.codebook_len(),
self.params
)));
}
Ok(())
}
}
pub fn default_codebook_name(field: &str) -> String {
format!("{field}.pqcb")
}
pub fn write_pq_codebook(
storage: &dyn Storage,
name: &str,
params: PqParams,
codebook: &[f32],
) -> Result<()> {
let tmp_name = format!("{name}.tmp");
let mut output = CrcWriter::new(storage.create_output(&tmp_name)?);
VectorSegmentHeader::product_quantization(params, codebook.to_vec()).write_to(&mut output)?;
let content_crc = output.checksum();
let mut inner = output.into_inner();
inner.write_all(&PQ_CODEBOOK_FOOTER_MAGIC.to_le_bytes())?;
inner.write_all(&content_crc.to_le_bytes())?;
inner.close()?;
storage.rename_file(&tmp_name, name)?;
Ok(())
}
pub fn read_pq_codebook(storage: &dyn Storage, name: &str) -> Result<SharedPqCodebook> {
let file_size = storage.file_size(name)?;
let mut crc_reader = CrcReader::new(storage.open_input(name)?);
let header = VectorSegmentHeader::read_from(&mut crc_reader, file_size)
.map_err(|e| LaurusError::index(format!("shared PQ codebook '{name}': {e}")))?;
let (params, codebook) = match header.quant {
QuantHeader::ProductQuantization { params, codebook } => (params, codebook),
other => {
return Err(LaurusError::index(format!(
"shared PQ codebook '{name}' has an unexpected quantization kind: {other:?}"
)));
}
};
let computed = crc_reader.checksum();
let inner = crc_reader.get_mut();
let mut footer = [0u8; FOOTER_SIZE];
inner.read_exact(&mut footer)?;
let magic = u32::from_le_bytes([footer[0], footer[1], footer[2], footer[3]]);
if magic != PQ_CODEBOOK_FOOTER_MAGIC {
return Err(LaurusError::index(format!(
"shared PQ codebook '{name}' footer magic mismatch: file is corrupted"
)));
}
let stored_crc = u32::from_le_bytes([footer[4], footer[5], footer[6], footer[7]]);
if stored_crc != computed {
return Err(LaurusError::index(format!(
"shared PQ codebook '{name}' checksum mismatch: file is corrupted"
)));
}
Ok(SharedPqCodebook { params, codebook })
}
pub fn train_and_write_pq_codebook(
storage: &dyn Storage,
name: &str,
dimension: usize,
subvector_count: usize,
k: u16,
normalize: bool,
vectors: &[Vector],
) -> Result<SharedPqCodebook> {
let params = PqParams::from_dim_and_m_k(dimension, subvector_count, k)?;
let normalized;
let training_set: &[Vector] = if normalize {
let mut owned = vectors.to_vec();
for v in &mut owned {
v.normalize();
}
normalized = owned;
&normalized
} else {
vectors
};
let codebook = pq_train_codebook(dimension, params, training_set)?;
write_pq_codebook(storage, name, params, &codebook)?;
Ok(SharedPqCodebook { params, codebook })
}
#[cfg(test)]
mod tests {
use super::*;
use crate::storage::memory::{MemoryStorage, MemoryStorageConfig};
fn storage() -> MemoryStorage {
MemoryStorage::new(MemoryStorageConfig::default())
}
fn sample_vectors(count: usize, dim: usize) -> Vec<Vector> {
let mut state: u64 = 0x1234_5678_9ABC_DEF0;
(0..count)
.map(|_| {
let data: Vec<f32> = (0..dim)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((state >> 33) as f32 / u32::MAX as f32) * 2.0 - 1.0
})
.collect();
Vector::new(data)
})
.collect()
}
#[test]
fn write_read_roundtrip_preserves_params_and_codebook() {
let storage = storage();
let vectors = sample_vectors(300, 32);
let trained =
train_and_write_pq_codebook(&storage, "field.pqcb", 32, 4, 256, false, &vectors)
.unwrap();
let loaded = read_pq_codebook(&storage, "field.pqcb").unwrap();
assert_eq!(loaded.params, trained.params);
assert_eq!(loaded.codebook, trained.codebook);
}
#[test]
fn corrupted_payload_byte_fails_checksum() {
let storage = storage();
let vectors = sample_vectors(300, 32);
train_and_write_pq_codebook(&storage, "field.pqcb", 32, 4, 256, false, &vectors).unwrap();
let mut input = storage.open_input("field.pqcb").unwrap();
let mut bytes = Vec::new();
input.read_to_end(&mut bytes).unwrap();
bytes[30] ^= 0xFF;
let mut output = storage.create_output("field.pqcb").unwrap();
output.write_all(&bytes).unwrap();
output.close().unwrap();
let err = read_pq_codebook(&storage, "field.pqcb").unwrap_err();
assert!(
matches!(&err, LaurusError::Index(msg) if msg.contains("checksum")),
"expected a checksum-mismatch Index error, got {err:?}"
);
}
#[test]
fn truncated_file_is_rejected_before_allocating() {
let storage = storage();
let vectors = sample_vectors(300, 32);
train_and_write_pq_codebook(&storage, "field.pqcb", 32, 4, 256, false, &vectors).unwrap();
let mut input = storage.open_input("field.pqcb").unwrap();
let mut bytes = Vec::new();
input.read_to_end(&mut bytes).unwrap();
bytes.truncate(PQ_HEADER_PREFIX_SIZE + 4);
let mut output = storage.create_output("field.pqcb").unwrap();
output.write_all(&bytes).unwrap();
output.close().unwrap();
let err = read_pq_codebook(&storage, "field.pqcb").unwrap_err();
assert!(
matches!(&err, LaurusError::Index(msg) if msg.contains("corrupted")),
"expected a size-mismatch Index error, got {err:?}"
);
}
#[test]
fn validate_for_rejects_dimension_and_subvector_mismatch() {
let storage = storage();
let vectors = sample_vectors(300, 32);
let cb = train_and_write_pq_codebook(&storage, "field.pqcb", 32, 4, 256, false, &vectors)
.unwrap();
assert!(cb.validate_for(32, 4, 256).is_ok());
assert!(
cb.validate_for(64, 4, 256).is_err(),
"dimension mismatch must be rejected"
);
assert!(
cb.validate_for(32, 8, 256).is_err(),
"subvector_count mismatch must be rejected"
);
assert!(
cb.validate_for(32, 4, 16).is_err(),
"k mismatch must be rejected (Issue #920: a k=256 codebook must not \
encode a FastScan field)"
);
}
#[test]
fn k16_codebook_round_trips_through_the_same_format() {
let storage = storage();
let vectors = sample_vectors(300, 32);
let trained =
train_and_write_pq_codebook(&storage, "fs.pqcb", 32, 4, 16, false, &vectors).unwrap();
assert_eq!(trained.params.k, 16);
let loaded = read_pq_codebook(&storage, "fs.pqcb").unwrap();
assert_eq!(loaded.params, trained.params);
assert_eq!(loaded.codebook, trained.codebook);
assert!(loaded.validate_for(32, 4, 16).is_ok());
assert!(
loaded.validate_for(32, 4, 256).is_err(),
"a k=16 codebook must not validate for the standard-PQ variant"
);
}
#[test]
fn train_and_write_normalizes_training_sample_when_requested() {
let storage = storage();
let vectors = sample_vectors(300, 32);
let cb_raw =
train_and_write_pq_codebook(&storage, "raw.pqcb", 32, 4, 256, false, &vectors).unwrap();
let cb_normalized =
train_and_write_pq_codebook(&storage, "norm.pqcb", 32, 4, 256, true, &vectors).unwrap();
assert_ne!(cb_raw.codebook, cb_normalized.codebook);
}
}