1use crate::error::Result;
7
8pub trait Reranker: Send {
9 fn id(&self) -> &str;
10 fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>>;
13}
14
15pub struct FakeReranker {
17 needle: String,
18}
19
20impl FakeReranker {
21 pub fn preferring(needle: &str) -> Self {
22 Self {
23 needle: needle.to_lowercase(),
24 }
25 }
26}
27
28impl Reranker for FakeReranker {
29 fn id(&self) -> &str {
30 "fake-reranker"
31 }
32
33 fn rerank(&self, _query: &str, documents: &[&str]) -> Result<Vec<f32>> {
34 Ok(documents
35 .iter()
36 .map(|d| {
37 if d.to_lowercase().contains(&self.needle) {
38 1.0
39 } else {
40 0.0
41 }
42 })
43 .collect())
44 }
45}
46
47#[cfg(feature = "local-embed")]
49pub struct OnnxReranker {
50 model: std::cell::RefCell<fastembed::TextRerank>,
51}
52
53#[cfg(feature = "local-embed")]
54impl OnnxReranker {
55 pub fn new(cache_dir: &std::path::Path) -> Result<Self> {
56 let options = fastembed::RerankInitOptions::new(fastembed::RerankerModel::BGERerankerBase)
57 .with_cache_dir(cache_dir.to_path_buf())
58 .with_show_download_progress(false);
59 let model = fastembed::TextRerank::try_new(options)
60 .map_err(|e| crate::SconeError::Embed(format!("reranker: {e}")))?;
61 Ok(Self {
62 model: std::cell::RefCell::new(model),
63 })
64 }
65}
66
67#[cfg(feature = "local-embed")]
68impl Reranker for OnnxReranker {
69 fn id(&self) -> &str {
70 "bge-reranker-base"
71 }
72
73 fn rerank(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>> {
74 let results = self
75 .model
76 .borrow_mut()
77 .rerank(query, documents, false, None)
78 .map_err(|e| crate::SconeError::Embed(format!("reranker: {e}")))?;
79 let mut scores = vec![0.0f32; documents.len()];
80 for r in results {
81 if let Some(slot) = scores.get_mut(r.index) {
82 *slot = r.score;
83 }
84 }
85 Ok(scores)
86 }
87}