use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use super::llm::LlmConfig;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RerankerConfig {
#[serde(default = "default_reranker_model", deserialize_with = "deserialize_null_model")]
pub model: RerankerModelType,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub top_k: Option<usize>,
#[serde(default = "default_batch_size")]
pub batch_size: 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_rerank_duration_secs",
skip_serializing_if = "Option::is_none"
)]
pub max_rerank_duration_secs: Option<u64>,
}
impl Default for RerankerConfig {
fn default() -> Self {
Self {
model: RerankerModelType::Preset {
name: "balanced".to_string(),
},
top_k: None,
batch_size: 32,
show_download_progress: false,
cache_dir: None,
acceleration: None,
max_rerank_duration_secs: Some(60),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RerankerHead {
CrossEncoder,
Qwen3Generative,
}
impl Default for RerankerHead {
fn default() -> Self {
Self::CrossEncoder
}
}
impl From<RerankerHead> for String {
fn from(head: RerankerHead) -> Self {
match head {
RerankerHead::CrossEncoder => "cross_encoder".to_string(),
RerankerHead::Qwen3Generative => "qwen3_generative".to_string(),
}
}
}
impl From<String> for RerankerHead {
fn from(value: String) -> Self {
match value.as_str() {
"qwen3_generative" => RerankerHead::Qwen3Generative,
_ => RerankerHead::CrossEncoder,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum RerankerModelType {
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>,
#[serde(default)]
head: RerankerHead,
},
Llm {
llm: LlmConfig,
},
Plugin {
name: String,
},
}
impl Default for RerankerModelType {
fn default() -> Self {
Self::Preset {
name: "balanced".to_string(),
}
}
}
fn default_batch_size() -> usize {
32
}
fn default_reranker_model() -> RerankerModelType {
RerankerModelType::Preset {
name: "balanced".to_string(),
}
}
fn default_max_rerank_duration_secs() -> Option<u64> {
Some(60)
}
fn deserialize_null_model<'de, D>(deserializer: D) -> Result<RerankerModelType, D::Error>
where
D: serde::Deserializer<'de>,
{
let opt = Option::<RerankerModelType>::deserialize(deserializer)?;
Ok(opt.unwrap_or_else(default_reranker_model))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_config_is_balanced_preset() {
let config = RerankerConfig::default();
assert!(matches!(
config.model,
RerankerModelType::Preset { ref name } if name == "balanced"
));
assert_eq!(config.batch_size, 32);
assert!(config.top_k.is_none());
assert_eq!(config.max_rerank_duration_secs, Some(60));
}
#[test]
fn default_model_type_is_balanced() {
let model = RerankerModelType::default();
assert!(matches!(model, RerankerModelType::Preset { ref name } if name == "balanced"));
}
#[test]
fn serde_roundtrip_preset() {
let config = RerankerConfig {
model: RerankerModelType::Preset {
name: "fast".to_string(),
},
top_k: Some(5),
..Default::default()
};
let json = serde_json::to_string(&config).unwrap();
let back: RerankerConfig = serde_json::from_str(&json).unwrap();
assert!(matches!(back.model, RerankerModelType::Preset { ref name } if name == "fast"));
assert_eq!(back.top_k, Some(5));
}
#[test]
fn serde_roundtrip_custom() {
let config = RerankerConfig {
model: RerankerModelType::Custom {
model_id: "cross-encoder/ms-marco-MiniLM-L6-v2".to_string(),
model_file: None,
additional_files: Vec::new(),
max_length: Some(512),
head: RerankerHead::CrossEncoder,
},
..Default::default()
};
let json = serde_json::to_string(&config).unwrap();
let back: RerankerConfig = serde_json::from_str(&json).unwrap();
assert!(matches!(
back.model,
RerankerModelType::Custom { ref model_id, .. } if model_id.contains("ms-marco")
));
}
#[test]
fn reranker_head_defaults_to_cross_encoder() {
assert_eq!(RerankerHead::default(), RerankerHead::CrossEncoder);
}
#[test]
fn reranker_head_serde_roundtrip() {
for head in [RerankerHead::CrossEncoder, RerankerHead::Qwen3Generative] {
let json = serde_json::to_string(&head).unwrap();
let back: RerankerHead = serde_json::from_str(&json).unwrap();
assert_eq!(back, head);
}
assert_eq!(
serde_json::to_string(&RerankerHead::CrossEncoder).unwrap(),
"\"cross_encoder\""
);
assert_eq!(
serde_json::to_string(&RerankerHead::Qwen3Generative).unwrap(),
"\"qwen3_generative\""
);
}
#[test]
fn custom_model_type_head_defaults_when_absent_from_json() {
let json = r#"{"type": "custom", "model_id": "cross-encoder/ms-marco-MiniLM-L6-v2"}"#;
let model: RerankerModelType = serde_json::from_str(json).unwrap();
assert!(matches!(
model,
RerankerModelType::Custom {
head: RerankerHead::CrossEncoder,
..
}
));
}
#[test]
fn null_model_field_deserializes_to_balanced() {
let json = r#"{"model": null}"#;
let config: RerankerConfig = serde_json::from_str(json).unwrap();
assert!(matches!(config.model, RerankerModelType::Preset { ref name } if name == "balanced"));
}
}