use anyhow::{Context, Result};
use fastembed::{RerankInitOptions, RerankerModel, TextRerank};
use ort::ep;
use std::path::PathBuf;
pub const ENV_ENABLE: &str = "SEMANTEX_RERANKER";
pub const ENV_MODEL: &str = "SEMANTEX_RERANKER_MODEL";
#[must_use]
pub fn reranker_enabled() -> bool {
matches!(
std::env::var(ENV_ENABLE).ok().map(|v| v.to_ascii_lowercase()),
Some(ref s) if matches!(s.as_str(), "on" | "1" | "true" | "yes")
)
}
#[must_use]
pub fn select_model_from_env() -> RerankerModel {
let raw = std::env::var(ENV_MODEL).unwrap_or_default();
match raw.to_ascii_lowercase().as_str() {
"" | "bge-reranker-v2-m3" | "bge-v2-m3" | "bge-v2" | "default" => {
RerankerModel::BGERerankerV2M3
}
"bge-reranker-base" | "bge-base" => RerankerModel::BGERerankerBase,
"jina-reranker-v1-turbo-en" | "jina-v1" | "jina-v1-turbo" => {
RerankerModel::JINARerankerV1TurboEn
}
"jina-reranker-v2-base-multilingual" | "jina-v2" | "jina-v2-base" => {
RerankerModel::JINARerankerV2BaseMultiligual
}
other => {
tracing::warn!(
model = other,
"Unknown {ENV_MODEL} value; falling back to bge-reranker-v2-m3"
);
RerankerModel::BGERerankerV2M3
}
}
}
#[must_use]
pub fn cache_dir() -> PathBuf {
if let Ok(v) = std::env::var("FASTEMBED_CACHE_DIR") {
return PathBuf::from(v);
}
dirs::home_dir().map_or_else(
|| PathBuf::from(".fastembed_cache"),
|h| h.join(".fastembed_cache"),
)
}
pub struct FastembedReranker {
model: TextRerank,
}
impl FastembedReranker {
pub fn new(model: RerankerModel, show_download_progress: bool) -> Result<Self> {
let execution_providers = Self::configure_execution_providers();
let options = RerankInitOptions::new(model)
.with_cache_dir(cache_dir())
.with_show_download_progress(show_download_progress)
.with_execution_providers(execution_providers);
let model =
TextRerank::try_new(options).context("Failed to initialize fastembed reranker")?;
Ok(Self { model })
}
pub fn new_default(show_download_progress: bool) -> Result<Self> {
if !reranker_enabled() {
anyhow::bail!(
"Refusing to construct reranker: SEMANTEX_RERANKER is not enabled. \
Set SEMANTEX_RERANKER=on to load model weights."
);
}
Self::new(select_model_from_env(), show_download_progress)
}
#[allow(clippy::vec_init_then_push)] fn configure_execution_providers() -> Vec<ort::ep::ExecutionProviderDispatch> {
let mut providers = Vec::new();
#[cfg(target_os = "macos")]
if std::env::var("SEMANTEX_COREML").is_ok() {
tracing::debug!("Reranker: CoreML execution provider enabled via SEMANTEX_COREML");
providers.push(ep::CoreML::default().build());
} else {
tracing::debug!(
"Reranker: CoreML disabled by default (set SEMANTEX_COREML=1 to enable). \
CPU-only reranking uses ~50-200 MB vs ~10 GB with CoreML."
);
}
#[cfg(feature = "cuda")]
{
providers.push(ep::CUDA::default().build());
}
providers.push(ep::CPU::default().build());
providers
}
pub fn rerank(
&mut self,
query: &str,
documents: &[&str],
top_k: usize,
) -> Result<Vec<(usize, f32)>> {
if documents.is_empty() {
return Ok(Vec::new());
}
if !reranker_enabled() {
tracing::debug!(
"Reranker disabled (SEMANTEX_RERANKER!=on); returning identity ordering"
);
let n = documents.len().min(top_k);
return Ok((0..n).map(|i| (i, 0.0_f32)).collect());
}
let results = self
.model
.rerank(query, documents, false, None)
.context("Fastembed reranking failed")?;
let mut scored: Vec<(usize, f32)> = results.iter().map(|r| (r.index, r.score)).collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(top_k);
Ok(scored)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn with_env<F: FnOnce()>(vars: &[(&str, Option<&str>)], f: F) {
let _guard = crate::search::RERANKER_TEST_ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let prior: Vec<(String, Option<String>)> = vars
.iter()
.map(|(k, _)| ((*k).to_string(), std::env::var(*k).ok()))
.collect();
unsafe {
for (k, v) in vars {
match v {
Some(val) => std::env::set_var(k, val),
None => std::env::remove_var(k),
}
}
}
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(f));
unsafe {
for (k, v) in &prior {
match v {
Some(val) => std::env::set_var(k, val),
None => std::env::remove_var(k),
}
}
}
if let Err(e) = result {
std::panic::resume_unwind(e);
}
}
#[test]
fn enabled_only_for_truthy_values() {
with_env(&[(ENV_ENABLE, Some("on"))], || {
assert!(reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("ON"))], || {
assert!(reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("1"))], || {
assert!(reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("true"))], || {
assert!(reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("True"))], || {
assert!(reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("yes"))], || {
assert!(reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("off"))], || {
assert!(!reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("0"))], || {
assert!(!reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("false"))], || {
assert!(!reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some(""))], || {
assert!(!reranker_enabled());
});
with_env(&[(ENV_ENABLE, None)], || {
assert!(!reranker_enabled());
});
}
#[test]
fn model_selection_defaults_to_bge_v2_m3() {
with_env(&[(ENV_MODEL, None)], || {
assert_eq!(select_model_from_env(), RerankerModel::BGERerankerV2M3);
});
with_env(&[(ENV_MODEL, Some(""))], || {
assert_eq!(select_model_from_env(), RerankerModel::BGERerankerV2M3);
});
with_env(&[(ENV_MODEL, Some("default"))], || {
assert_eq!(select_model_from_env(), RerankerModel::BGERerankerV2M3);
});
with_env(&[(ENV_MODEL, Some("garbage-not-a-model"))], || {
assert_eq!(select_model_from_env(), RerankerModel::BGERerankerV2M3);
});
}
#[test]
fn model_selection_recognizes_explicit_choices() {
with_env(&[(ENV_MODEL, Some("bge-reranker-v2-m3"))], || {
assert_eq!(select_model_from_env(), RerankerModel::BGERerankerV2M3);
});
with_env(&[(ENV_MODEL, Some("BGE-V2-M3"))], || {
assert_eq!(select_model_from_env(), RerankerModel::BGERerankerV2M3);
});
with_env(&[(ENV_MODEL, Some("bge-reranker-base"))], || {
assert_eq!(select_model_from_env(), RerankerModel::BGERerankerBase);
});
with_env(&[(ENV_MODEL, Some("jina-reranker-v1-turbo-en"))], || {
assert_eq!(
select_model_from_env(),
RerankerModel::JINARerankerV1TurboEn
);
});
with_env(&[(ENV_MODEL, Some("jina-v2"))], || {
assert_eq!(
select_model_from_env(),
RerankerModel::JINARerankerV2BaseMultiligual
);
});
}
#[test]
fn cache_dir_resolves_to_home_or_override() {
with_env(
&[("FASTEMBED_CACHE_DIR", Some("/tmp/semantex-test-cache"))],
|| {
assert_eq!(cache_dir(), PathBuf::from("/tmp/semantex-test-cache"));
},
);
with_env(&[("FASTEMBED_CACHE_DIR", None)], || {
let p = cache_dir();
assert!(p.ends_with(".fastembed_cache"), "got {p:?}");
});
}
#[test]
fn new_default_refuses_when_disabled() {
with_env(&[(ENV_ENABLE, Some("off")), (ENV_MODEL, None)], || {
match FastembedReranker::new_default(false) {
Ok(_) => panic!("new_default must not load weights when disabled"),
Err(e) => {
let msg = format!("{e}");
assert!(
msg.contains("SEMANTEX_RERANKER"),
"error must point at the env var; got: {msg}"
);
}
}
});
}
#[test]
#[ignore = "downloads ~600 MB of fastembed reranker weights; run with --ignored"]
fn reranks_when_enabled() {
with_env(&[(ENV_ENABLE, Some("on")), (ENV_MODEL, None)], || {
let mut reranker =
FastembedReranker::new_default(false).expect("model load failed (offline?)");
let docs = [
"The giant panda is a bear endemic to China.",
"Binary search is an efficient algorithm for finding an item in a sorted slice.",
"Pizza is a popular Italian dish.",
"Rust's slice::binary_search returns the index of a matching element.",
];
let doc_slices: Vec<&str> = docs.to_vec();
let results = reranker
.rerank("how does binary search work in Rust?", &doc_slices, 4)
.expect("rerank failed");
assert!(!results.is_empty());
let top_idx = results[0].0;
assert!(
top_idx == 1 || top_idx == 3,
"expected on-topic doc at top, got idx {top_idx}"
);
});
}
#[test]
fn rerank_is_identity_when_disabled() {
with_env(&[(ENV_ENABLE, None)], || {
assert!(!reranker_enabled());
});
with_env(&[(ENV_ENABLE, Some("off"))], || {
assert!(!reranker_enabled());
});
}
}