use crate::search::dense_backend::DenseBackendKind;
use anyhow::Result;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct ModelCapabilities {
#[serde(default)]
pub multi_vector: bool,
#[serde(default)]
pub matryoshka_dims: Option<Vec<usize>>,
#[serde(default)]
pub produces_sparse: bool,
#[serde(default)]
pub instruction_aware: bool,
#[serde(default)]
pub max_batch: Option<usize>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BackendKind {
ColbertPlaid,
CoderankHnsw,
}
pub fn backend_for(caps: &ModelCapabilities) -> Result<BackendKind> {
if caps.multi_vector {
Ok(BackendKind::ColbertPlaid)
} else {
Ok(BackendKind::CoderankHnsw)
}
}
impl BackendKind {
pub fn dense_kind(self) -> Result<DenseBackendKind> {
match self {
BackendKind::ColbertPlaid => Ok(DenseBackendKind::ColbertPlaid),
BackendKind::CoderankHnsw => Ok(DenseBackendKind::CoderankHnsw),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn single_vector_routes_to_hnsw() {
let sv = ModelCapabilities {
multi_vector: false,
..Default::default()
};
assert_eq!(backend_for(&sv).unwrap(), BackendKind::CoderankHnsw);
}
#[test]
fn multi_vector_routes_to_plaid() {
let mv = ModelCapabilities {
multi_vector: true,
..Default::default()
};
assert_eq!(backend_for(&mv).unwrap(), BackendKind::ColbertPlaid);
}
#[test]
fn coderank_hnsw_dense_kind_maps_to_s1_dense_kind() {
assert_eq!(
BackendKind::CoderankHnsw.dense_kind().unwrap(),
DenseBackendKind::CoderankHnsw
);
}
#[test]
fn colbert_plaid_dense_kind_maps_to_s1_dense_kind() {
assert_eq!(
BackendKind::ColbertPlaid.dense_kind().unwrap(),
DenseBackendKind::ColbertPlaid
);
}
#[test]
fn capabilities_default_is_single_vector_profile() {
let c = ModelCapabilities::default();
assert!(!c.multi_vector);
assert!(c.matryoshka_dims.is_none());
assert!(!c.produces_sparse);
assert!(c.max_batch.is_none());
}
#[test]
fn partial_toml_keeps_unset_capabilities_off() {
let c: ModelCapabilities = toml::from_str("multi_vector = true\n").unwrap();
assert!(c.multi_vector);
assert!(!c.instruction_aware);
assert!(c.matryoshka_dims.is_none());
}
}