use crate::model::spec::{ModelRole, ModelSource, ModelSpec, RoleData, ScoreStrategyKind};
use crate::search::fastembed_reranker::{ENV_MODEL, select_model_from_env};
use crate::search::onnx_reranker::{PromptTemplate, ScoreStrategy};
use crate::search::reranker_download::ModelFiles;
use anyhow::{Context, Result};
#[derive(Debug, Clone)]
pub struct OnnxModelSpec {
pub files: ModelFiles,
pub session_file: String,
pub max_context: usize,
pub strategy: ScoreStrategy,
}
#[derive(Debug, Clone)]
pub enum RerankerChoice {
Fastembed(fastembed::RerankerModel),
Onnx(OnnxModelSpec),
}
impl RerankerChoice {
pub fn from_spec(spec: &ModelSpec) -> Result<Self> {
anyhow::ensure!(
spec.role == ModelRole::Reranker,
"model `{}` is not a reranker (role {:?})",
spec.id,
spec.role
);
let RoleData::Reranker(rspec) = &spec.role_data else {
anyhow::bail!("reranker `{}` has no reranker role data", spec.id);
};
if let Some(model) = fastembed_model_for_id(&spec.id) {
return Ok(Self::Fastembed(model));
}
let (subdir, base_url, files) = source_to_coords(&spec.source, &spec.id)?;
let session_file = files
.first()
.cloned()
.with_context(|| format!("reranker `{}` lists no files", spec.id))?;
let strategy = match rspec.score_strategy {
ScoreStrategyKind::ClassifierHead => ScoreStrategy::ClassifierLogit,
ScoreStrategyKind::YesNoLogit => {
let yes_id = rspec.yes_token_id.with_context(|| {
format!("yes_no reranker `{}` is missing `yes_token_id`", spec.id)
})?;
let no_id = rspec.no_token_id.with_context(|| {
format!("yes_no reranker `{}` is missing `no_token_id`", spec.id)
})?;
ScoreStrategy::YesNoLogit {
yes_id,
no_id,
prompt: PromptTemplate {
prefix: rspec.prompt_prefix.clone(),
middle: rspec.prompt_middle.clone(),
suffix: rspec.prompt_suffix.clone(),
},
}
}
};
Ok(Self::Onnx(OnnxModelSpec {
files: ModelFiles {
subdir,
base_url,
files,
},
session_file,
max_context: rspec.max_context,
strategy,
}))
}
}
fn fastembed_model_for_id(id: &str) -> Option<fastembed::RerankerModel> {
use fastembed::RerankerModel as M;
match id {
"bge-reranker-v2-m3" | "bge-v2-m3" => Some(M::BGERerankerV2M3),
"bge-reranker-base" | "bge-base" => Some(M::BGERerankerBase),
"jina-reranker-v1-turbo-en" | "jina-v1" => Some(M::JINARerankerV1TurboEn),
"jina-reranker-v2-base-multilingual" | "jina-v2" => Some(M::JINARerankerV2BaseMultiligual),
_ => None,
}
}
fn source_to_coords(source: &ModelSource, id: &str) -> Result<(String, String, Vec<String>)> {
match source {
ModelSource::Hf { repo, files } => {
let subdir = repo.rsplit('/').next().unwrap_or(repo).to_string();
let base_url = format!("https://huggingface.co/{repo}/resolve/main");
Ok((subdir, base_url, files.clone()))
}
ModelSource::Url { base, files } => {
Ok((id.to_string(), base.clone(), files.clone()))
}
ModelSource::Local { dir } => {
anyhow::bail!(
"reranker `{id}`: ModelSource::Local ({dir}) is not yet supported by the \
generic ONNX reranker download path"
)
}
}
}
#[must_use]
pub fn select_reranker_choice_from_env() -> RerankerChoice {
let raw = std::env::var(ENV_MODEL).unwrap_or_default();
match raw.to_ascii_lowercase().as_str() {
"qwen3-reranker-0.6b" | "qwen3-reranker" | "qwen3" => {
choice_from_builtin_id("qwen3-reranker-0.6b")
.unwrap_or_else(|| RerankerChoice::Fastembed(select_model_from_env()))
}
"bge-reranker-v2-m3-onnx" | "bge-v2-m3-onnx" | "bge-onnx" => {
RerankerChoice::Onnx(bge_v2_m3_onnx())
}
_ => RerankerChoice::Fastembed(select_model_from_env()),
}
}
fn choice_from_builtin_id(id: &str) -> Option<RerankerChoice> {
crate::model::manifest::builtin_specs()
.into_iter()
.find(|s| s.id == id)
.and_then(|spec| RerankerChoice::from_spec(&spec).ok())
}
fn bge_v2_m3_onnx() -> OnnxModelSpec {
OnnxModelSpec {
files: ModelFiles {
subdir: "bge-reranker-v2-m3-onnx".to_string(),
base_url: "https://huggingface.co/BAAI/bge-reranker-v2-m3/resolve/main/onnx"
.to_string(),
files: vec!["model.onnx".to_string(), "tokenizer.json".to_string()],
},
session_file: "model.onnx".to_string(),
max_context: 512,
strategy: ScoreStrategy::ClassifierLogit,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::SemantexConfig;
use crate::model::ModelRegistry;
fn with_env<F: FnOnce()>(key: &str, val: Option<&str>, f: F) {
let _g = crate::search::RERANKER_TEST_ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let prior = std::env::var(key).ok();
unsafe {
match val {
Some(v) => std::env::set_var(key, v),
None => std::env::remove_var(key),
}
}
let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(f));
unsafe {
match prior {
Some(v) => std::env::set_var(key, v),
None => std::env::remove_var(key),
}
}
if let Err(e) = r {
std::panic::resume_unwind(e);
}
}
#[test]
fn qwen3_aliases_route_to_onnx_yesno() {
for alias in ["qwen3", "Qwen3-Reranker-0.6B", "qwen3-reranker"] {
with_env(
ENV_MODEL,
Some(alias),
|| match select_reranker_choice_from_env() {
RerankerChoice::Onnx(spec) => {
assert!(matches!(spec.strategy, ScoreStrategy::YesNoLogit { .. }));
assert_eq!(spec.session_file, "model.onnx");
assert_eq!(spec.files.subdir, "Qwen3-Reranker-0.6B-onnx");
}
other @ RerankerChoice::Fastembed(_) => {
panic!("expected Onnx for {alias}, got {other:?}")
}
},
);
}
}
#[test]
fn bge_onnx_alias_routes_to_onnx_classifier() {
with_env(
ENV_MODEL,
Some("bge-onnx"),
|| match select_reranker_choice_from_env() {
RerankerChoice::Onnx(spec) => {
assert!(matches!(spec.strategy, ScoreStrategy::ClassifierLogit));
}
other @ RerankerChoice::Fastembed(_) => {
panic!("expected Onnx classifier, got {other:?}")
}
},
);
}
#[test]
fn default_and_unknown_route_to_fastembed_bge() {
for v in [None, Some(""), Some("garbage"), Some("bge-reranker-v2-m3")] {
with_env(ENV_MODEL, v, || match select_reranker_choice_from_env() {
RerankerChoice::Fastembed(m) => {
assert_eq!(m, fastembed::RerankerModel::BGERerankerV2M3);
}
other @ RerankerChoice::Onnx(_) => {
panic!("expected Fastembed bge for {v:?}, got {other:?}")
}
});
}
}
#[test]
fn jina_v3_is_not_selectable() {
with_env(ENV_MODEL, Some("jina-reranker-v3"), || {
assert!(matches!(
select_reranker_choice_from_env(),
RerankerChoice::Fastembed(fastembed::RerankerModel::BGERerankerV2M3)
));
});
}
#[test]
fn from_spec_resolves_bge_to_fastembed() {
let reg = ModelRegistry::from_config(&SemantexConfig::default(), None).unwrap();
let spec = reg.active_reranker().unwrap();
match RerankerChoice::from_spec(spec).unwrap() {
RerankerChoice::Fastembed(m) => {
assert_eq!(m, fastembed::RerankerModel::BGERerankerV2M3);
}
other @ RerankerChoice::Onnx(_) => {
panic!("expected Fastembed bge from registry spec, got {other:?}")
}
}
}
#[test]
fn from_spec_resolves_qwen3_to_onnx_yesno() {
let cfg = SemantexConfig {
reranker_model: "qwen3-reranker-0.6b".to_string(),
..Default::default()
};
let reg = ModelRegistry::from_config(&cfg, None).unwrap();
let spec = reg.active_reranker().unwrap();
match RerankerChoice::from_spec(spec).unwrap() {
RerankerChoice::Onnx(o) => {
match &o.strategy {
ScoreStrategy::YesNoLogit {
yes_id,
no_id,
prompt,
} => {
assert_eq!(*yes_id, 9693);
assert_eq!(*no_id, 2152);
let rendered = prompt.render("binary search", "fn bsearch() {}");
assert!(
rendered.contains("<|im_start|>system"),
"missing system im_start; got:\n{rendered}"
);
assert!(
rendered.contains("<|im_start|>user"),
"missing user im_start"
);
assert!(
rendered.contains("<|im_start|>assistant"),
"missing assistant im_start"
);
assert!(rendered.contains("<|im_end|>"), "missing im_end");
assert!(rendered.contains("<think>"), "missing <think> terminator");
assert!(
rendered.contains("binary search") && rendered.contains("fn bsearch"),
"query/doc not injected"
);
}
other @ ScoreStrategy::ClassifierLogit => {
panic!("expected YesNoLogit, got {other:?}")
}
}
assert_eq!(o.session_file, "model.onnx");
assert_eq!(o.max_context, 2048);
assert_eq!(o.files.subdir, "Qwen3-Reranker-0.6B-onnx");
assert!(
o.files
.base_url
.contains("MisterTK/Qwen3-Reranker-0.6B-onnx")
);
}
other @ RerankerChoice::Fastembed(_) => {
panic!("expected Onnx for qwen3 registry spec, got {other:?}")
}
}
}
#[test]
fn qwen3_env_alias_renders_full_chat_template() {
with_env(
ENV_MODEL,
Some("qwen3"),
|| match select_reranker_choice_from_env() {
RerankerChoice::Onnx(o) => match &o.strategy {
ScoreStrategy::YesNoLogit { prompt, .. } => {
let rendered = prompt.render("q", "d");
assert!(rendered.contains("<|im_start|>"), "missing im_start");
assert!(rendered.contains("<think>"), "missing <think>");
assert_eq!(o.max_context, 2048);
}
other @ ScoreStrategy::ClassifierLogit => {
panic!("expected YesNoLogit, got {other:?}")
}
},
other @ RerankerChoice::Fastembed(_) => {
panic!("expected Onnx for qwen3 alias, got {other:?}")
}
},
);
}
}