use std::collections::HashMap;
use std::sync::{Arc, Mutex, OnceLock};
use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use crate::types::SearchResult;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default)]
pub struct RerankerConfig {
pub model_path: String,
pub tokenizer_path: String,
pub max_seq_len: usize,
pub top_k: usize,
pub rerank_weight: f32,
}
impl Default for RerankerConfig {
fn default() -> Self {
Self {
model_path: String::new(),
tokenizer_path: String::new(),
max_seq_len: 512,
top_k: 50,
rerank_weight: 1.0,
}
}
}
pub trait Reranker: Send + Sync {
fn score(&self, query: &str, passages: &[&str]) -> Result<Vec<f32>>;
}
pub fn validate_config(cfg: &Option<RerankerConfig>) -> Result<()> {
let Some(cfg) = cfg else { return Ok(()) };
if cfg.model_path.is_empty() || !std::path::Path::new(&cfg.model_path).exists() {
anyhow::bail!(
"[search.reranker] is configured but model_path {:?} does not exist. \
Bobbin does not download reranker models — supply a cross-encoder \
ONNX file, or remove the [search.reranker] section.",
cfg.model_path
);
}
if cfg.tokenizer_path.is_empty() || !std::path::Path::new(&cfg.tokenizer_path).exists() {
anyhow::bail!(
"[search.reranker] is configured but tokenizer_path {:?} does not exist. \
Supply the model's tokenizer.json, or remove the [search.reranker] section.",
cfg.tokenizer_path
);
}
Ok(())
}
static RERANKERS: OnceLock<Mutex<HashMap<String, Arc<dyn Reranker>>>> = OnceLock::new();
pub fn for_config(cfg: &RerankerConfig) -> Result<Arc<dyn Reranker>> {
validate_config(&Some(cfg.clone()))?;
let key = format!(
"{}\x1f{}\x1f{}",
cfg.model_path, cfg.tokenizer_path, cfg.max_seq_len
);
let cache = RERANKERS.get_or_init(|| Mutex::new(HashMap::new()));
let mut cache = cache
.lock()
.map_err(|e| anyhow::anyhow!("reranker cache lock poisoned: {e}"))?;
if let Some(r) = cache.get(&key) {
return Ok(r.clone());
}
let loaded: Arc<dyn Reranker> = Arc::new(
super::rerank_onnx::OnnxCrossEncoder::load(cfg)
.context("failed to load [search.reranker] cross-encoder")?,
);
cache.insert(key, loaded.clone());
Ok(loaded)
}
fn sigmoid(x: f32) -> f32 {
if x >= 0.0 {
1.0 / (1.0 + (-x).exp())
} else {
let e = x.exp();
e / (1.0 + e)
}
}
pub fn apply_rerank(
reranker: &dyn Reranker,
query: &str,
mut results: Vec<SearchResult>,
top_k: usize,
rerank_weight: f32,
) -> Result<Vec<SearchResult>> {
let k = top_k.min(results.len());
if k == 0 {
return Ok(results);
}
let passages: Vec<&str> = results[..k]
.iter()
.map(|r| r.chunk.content.as_str())
.collect();
let raw = reranker.score(query, &passages)?;
if raw.len() != k {
anyhow::bail!("reranker returned {} scores for {} passages", raw.len(), k);
}
let fused_min = results[..k]
.iter()
.map(|r| r.score)
.fold(f32::INFINITY, f32::min);
let fused_max = results[..k]
.iter()
.map(|r| r.score)
.fold(f32::NEG_INFINITY, f32::max);
let span = fused_max - fused_min;
let mut head: Vec<(SearchResult, f32)> = results
.drain(..k)
.zip(raw.iter())
.map(|(r, &logit)| {
let fused_norm = if span > 0.0 {
(r.score - fused_min) / span
} else {
0.5
};
let blended = rerank_weight * sigmoid(logit) + (1.0 - rerank_weight) * fused_norm;
(r, blended)
})
.collect();
head.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let mut out = Vec::with_capacity(head.len() + results.len());
for (mut r, blended) in head {
r.score = blended;
out.push(r);
}
out.append(&mut results);
Ok(out)
}
#[cfg(test)]
#[path = "rerank_tests.rs"]
mod tests;