use serde::{Deserialize, Serialize};
use crate::dtype::Dtype;
use crate::error::{Error, Result};
use crate::metric::Metric;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
pub enum IndexBackend {
#[default]
Hnsw,
Ivf,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
pub enum IoProfile {
Hdd,
Ssd,
#[default]
Auto,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct StorageParams {
pub dtype: Dtype,
pub block_size: usize,
pub io_profile: IoProfile,
}
impl Default for StorageParams {
fn default() -> Self {
Self {
dtype: Dtype::F32,
block_size: 64 * 1024,
io_profile: IoProfile::Auto,
}
}
}
impl StorageParams {
pub fn validate(&self) -> Result<()> {
if self.block_size < 512 {
return Err(Error::invalid_config(format!(
"block_size must be >= 512 bytes, got {}",
self.block_size
)));
}
if !self.block_size.is_power_of_two() {
return Err(Error::invalid_config(format!(
"block_size must be a power of two, got {}",
self.block_size
)));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct PqParams {
pub num_subquantizers: usize,
pub bits_per_code: u8,
}
impl Default for PqParams {
fn default() -> Self {
Self {
num_subquantizers: 16,
bits_per_code: 8,
}
}
}
impl PqParams {
#[inline]
pub const fn centroids_per_subspace(&self) -> usize {
1usize << self.bits_per_code
}
pub fn validate(&self) -> Result<()> {
if self.num_subquantizers == 0 {
return Err(Error::invalid_config("num_subquantizers must be >= 1"));
}
if self.bits_per_code == 0 || self.bits_per_code > 8 {
return Err(Error::invalid_config(format!(
"bits_per_code must be in 1..=8, got {}",
self.bits_per_code
)));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct HnswParams {
pub m: usize,
pub m_max: usize,
pub ef_construction: usize,
pub hub_fraction: f32,
}
impl Default for HnswParams {
fn default() -> Self {
Self {
m: 8,
m_max: 32,
ef_construction: 200,
hub_fraction: 0.02,
}
}
}
impl HnswParams {
pub fn validate(&self) -> Result<()> {
if self.m == 0 {
return Err(Error::invalid_config("hnsw.m must be >= 1"));
}
if self.m_max < self.m {
return Err(Error::invalid_config(format!(
"hnsw.m_max ({}) must be >= hnsw.m ({})",
self.m_max, self.m
)));
}
if self.ef_construction == 0 {
return Err(Error::invalid_config("hnsw.ef_construction must be >= 1"));
}
if !(0.0..=1.0).contains(&self.hub_fraction) {
return Err(Error::invalid_config(format!(
"hnsw.hub_fraction must be in [0, 1], got {}",
self.hub_fraction
)));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct IvfParams {
pub num_lists: usize,
pub num_probes: usize,
pub soft_assign: usize,
pub max_kmeans_iters: usize,
}
impl Default for IvfParams {
fn default() -> Self {
Self {
num_lists: 256,
num_probes: 16,
soft_assign: 2,
max_kmeans_iters: 25,
}
}
}
impl IvfParams {
pub fn validate(&self) -> Result<()> {
if self.num_lists == 0 {
return Err(Error::invalid_config("ivf.num_lists must be >= 1"));
}
if self.num_probes == 0 || self.num_probes > self.num_lists {
return Err(Error::invalid_config(format!(
"ivf.num_probes must be in 1..={}, got {}",
self.num_lists, self.num_probes
)));
}
if self.soft_assign == 0 || self.soft_assign > self.num_lists {
return Err(Error::invalid_config(format!(
"ivf.soft_assign must be in 1..={}, got {}",
self.num_lists, self.soft_assign
)));
}
if self.max_kmeans_iters == 0 {
return Err(Error::invalid_config("ivf.max_kmeans_iters must be >= 1"));
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct BuildConfig {
pub dimensions: usize,
pub metric: Metric,
pub backend: IndexBackend,
pub storage: StorageParams,
pub pq: PqParams,
pub hnsw: HnswParams,
pub ivf: IvfParams,
pub num_shards: usize,
}
impl Default for BuildConfig {
fn default() -> Self {
Self {
dimensions: 0,
metric: Metric::L2,
backend: IndexBackend::Hnsw,
storage: StorageParams::default(),
pq: PqParams::default(),
hnsw: HnswParams::default(),
ivf: IvfParams::default(),
num_shards: 1,
}
}
}
impl BuildConfig {
pub fn new(dimensions: usize, metric: Metric, backend: IndexBackend) -> Self {
Self {
dimensions,
metric,
backend,
..Self::default()
}
}
pub fn validate(&self) -> Result<()> {
if self.dimensions == 0 {
return Err(Error::invalid_config("dimensions must be set (> 0)"));
}
self.storage.validate()?;
self.pq.validate()?;
self.hnsw.validate()?;
self.ivf.validate()?;
if !self.dimensions.is_multiple_of(self.pq.num_subquantizers) {
return Err(Error::invalid_config(format!(
"dimensions ({}) must be divisible by pq.num_subquantizers ({})",
self.dimensions, self.pq.num_subquantizers
)));
}
if self.num_shards == 0 {
return Err(Error::invalid_config("num_shards must be >= 1"));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct SearchConfig {
pub k: usize,
pub ef_search: usize,
pub rerank_ratio: f32,
pub fetch_batch_size: usize,
}
impl Default for SearchConfig {
fn default() -> Self {
Self {
k: 10,
ef_search: 64,
rerank_ratio: 0.2,
fetch_batch_size: 64,
}
}
}
impl SearchConfig {
pub fn validate(&self) -> Result<()> {
if self.k == 0 {
return Err(Error::invalid_config("k must be >= 1"));
}
if self.ef_search < self.k {
return Err(Error::invalid_config(format!(
"ef_search ({}) must be >= k ({})",
self.ef_search, self.k
)));
}
if !(self.rerank_ratio > 0.0 && self.rerank_ratio <= 1.0) {
return Err(Error::invalid_config(format!(
"rerank_ratio must be in (0, 1], got {}",
self.rerank_ratio
)));
}
if self.fetch_batch_size == 0 {
return Err(Error::invalid_config("fetch_batch_size must be >= 1"));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::field_reassign_with_default)]
use super::*;
#[test]
fn default_build_config_requires_dimensions() {
let cfg = BuildConfig::default();
assert!(cfg.validate().is_err(), "dimensions=0 must be rejected");
}
#[test]
fn valid_build_config_passes() {
let cfg = BuildConfig::new(768, Metric::Cosine, IndexBackend::Hnsw);
assert!(cfg.validate().is_ok(), "{:?}", cfg.validate());
}
#[test]
fn pq_divisibility_enforced() {
let mut cfg = BuildConfig::new(770, Metric::L2, IndexBackend::Hnsw);
cfg.pq.num_subquantizers = 16; assert!(cfg.validate().is_err());
}
#[test]
fn pq_centroid_count() {
assert_eq!(PqParams::default().centroids_per_subspace(), 256);
}
#[test]
fn storage_block_size_must_be_pow2() {
let mut sp = StorageParams::default();
sp.block_size = 1000;
assert!(sp.validate().is_err());
sp.block_size = 4096;
assert!(sp.validate().is_ok());
}
#[test]
fn hnsw_degree_ordering() {
let mut p = HnswParams::default();
p.m = 40;
p.m_max = 32;
assert!(p.validate().is_err());
}
#[test]
fn ivf_probe_bounds() {
let mut p = IvfParams::default();
p.num_probes = p.num_lists + 1;
assert!(p.validate().is_err());
}
#[test]
fn search_config_bounds() {
let mut s = SearchConfig::default();
assert!(s.validate().is_ok());
s.ef_search = 1;
s.k = 10;
assert!(s.validate().is_err());
s = SearchConfig::default();
s.rerank_ratio = 0.0;
assert!(s.validate().is_err());
}
#[test]
fn config_roundtrips_through_json() {
let cfg = BuildConfig::new(128, Metric::InnerProduct, IndexBackend::Ivf);
let json = serde_json::to_string(&cfg).expect("serialize");
let back: BuildConfig = serde_json::from_str(&json).expect("deserialize");
assert_eq!(cfg, back);
}
}