use serde::{Deserialize, Serialize};
use std::path::PathBuf;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SparseEmbeddingConfig {
#[serde(default = "default_sparse_model", deserialize_with = "deserialize_null_model")]
pub model: SparseEmbeddingModelType,
#[serde(default = "default_batch_size")]
pub batch_size: usize,
#[serde(default = "default_max_length")]
pub max_length: usize,
#[serde(default)]
pub show_download_progress: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_dir: Option<PathBuf>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub acceleration: Option<super::acceleration::AccelerationConfig>,
#[serde(default = "default_max_embed_duration_secs", skip_serializing_if = "Option::is_none")]
pub max_embed_duration_secs: Option<u64>,
}
impl Default for SparseEmbeddingConfig {
fn default() -> Self {
Self {
model: default_sparse_model(),
batch_size: default_batch_size(),
max_length: default_max_length(),
show_download_progress: false,
cache_dir: None,
acceleration: None,
max_embed_duration_secs: default_max_embed_duration_secs(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum SparseEmbeddingModelType {
Preset {
name: String,
},
Custom {
model_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
model_file: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
additional_files: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
max_length: Option<i64>,
},
Plugin {
name: String,
},
}
impl Default for SparseEmbeddingModelType {
fn default() -> Self {
Self::Preset {
name: "opensearch-v3-distill".to_string(),
}
}
}
fn default_sparse_model() -> SparseEmbeddingModelType {
SparseEmbeddingModelType::default()
}
fn default_batch_size() -> usize {
16
}
fn default_max_length() -> usize {
256
}
fn default_max_embed_duration_secs() -> Option<u64> {
Some(60)
}
fn deserialize_null_model<'de, D>(deserializer: D) -> Result<SparseEmbeddingModelType, D::Error>
where
D: serde::Deserializer<'de>,
{
let opt = Option::<SparseEmbeddingModelType>::deserialize(deserializer)?;
Ok(opt.unwrap_or_default())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_config_uses_opensearch_preset() {
let config = SparseEmbeddingConfig::default();
assert!(matches!(config.model, SparseEmbeddingModelType::Preset { name } if name == "opensearch-v3-distill"));
assert_eq!(config.batch_size, 16);
assert_eq!(config.max_length, 256);
}
#[test]
fn null_model_deserializes_to_default() {
let json = r#"{"model": null}"#;
let config: SparseEmbeddingConfig = serde_json::from_str(json).unwrap();
assert!(matches!(config.model, SparseEmbeddingModelType::Preset { name } if name == "opensearch-v3-distill"));
}
#[test]
fn custom_model_roundtrips() {
let config = SparseEmbeddingConfig {
model: SparseEmbeddingModelType::Custom {
model_id: "org/splade".to_string(),
model_file: Some("onnx/model.onnx".to_string()),
additional_files: vec![],
max_length: Some(256),
},
..Default::default()
};
let json = serde_json::to_string(&config).unwrap();
let back: SparseEmbeddingConfig = serde_json::from_str(&json).unwrap();
assert!(matches!(back.model, SparseEmbeddingModelType::Custom { model_id, .. } if model_id == "org/splade"));
}
}