use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use super::llm::LlmConfig;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
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: default_reranker_model(),
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, Hash, 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 {
Self::try_from(value.as_str()).unwrap_or_default()
}
}
impl TryFrom<&str> for RerankerHead {
type Error = crate::XbergError;
fn try_from(value: &str) -> Result<Self, Self::Error> {
match value {
"cross_encoder" => Ok(Self::CrossEncoder),
"qwen3_generative" => Ok(Self::Qwen3Generative),
_ => Err(crate::XbergError::validation(format!(
"invalid RerankerHead value `{value}`; expected one of: cross_encoder, qwen3_generative"
))),
}
}
}
impl std::str::FromStr for RerankerHead {
type Err = crate::XbergError;
fn from_str(value: &str) -> Result<Self, Self::Err> {
Self::try_from(value)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
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: Box<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 config_rejects_unknown_fields() {
let json = r#"{"model":{"type":"preset","name":"balanced"},"batch_limit":16}"#;
assert!(serde_json::from_str::<RerankerConfig>(json).is_err());
}
#[test]
fn model_type_rejects_unknown_fields() {
let json = r#"{"type":"preset","name":"balanced","extra_name":"other"}"#;
assert!(serde_json::from_str::<RerankerModelType>(json).is_err());
}
#[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 reranker_head_parsing_accepts_every_wire_value() {
assert_eq!(
"cross_encoder"
.parse::<RerankerHead>()
.expect("the cross-encoder head must parse"),
RerankerHead::CrossEncoder
);
assert_eq!(
"qwen3_generative"
.parse::<RerankerHead>()
.expect("the Qwen3 head must parse"),
RerankerHead::Qwen3Generative
);
}
#[test]
fn reranker_head_parsing_rejects_unknown_values() {
let error = "classification"
.parse::<RerankerHead>()
.expect_err("unknown heads must be rejected");
assert_eq!(
error.to_string(),
concat!(
"Validation error: invalid RerankerHead value `classification`; ",
"expected one of: cross_encoder, qwen3_generative"
)
);
}
#[test]
fn custom_model_deserialization_rejects_unknown_head() {
let error = serde_json::from_str::<RerankerModelType>(
r#"{"type":"custom","model_id":"example/model","head":"classification"}"#,
)
.expect_err("unknown heads must not cross the JSON configuration boundary");
assert!(error.to_string().contains("unknown variant `classification`"));
}
#[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"));
}
}