Skip to main content

lc_rag/
reranking.rs

1// src/retrieval/reranking.rs
2//! Reranking(重排序)实现
3//!
4//! 使用评分函数对检索结果重新排序,提升检索精确度。
5
6use lc_vector_stores::{Document, SearchResult};
7use std::collections::HashMap;
8
9/// Reranking 错误类型
10#[derive(Debug)]
11pub enum RerankingError {
12    ScoringError(String),
13    InvalidInput(String),
14}
15
16impl std::fmt::Display for RerankingError {
17    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
18        match self {
19            RerankingError::ScoringError(msg) => write!(f, "评分错误: {}", msg),
20            RerankingError::InvalidInput(msg) => write!(f, "输入无效: {}", msg),
21        }
22    }
23}
24
25impl std::error::Error for RerankingError {}
26
27/// Reranking 配置
28pub struct RerankingConfig {
29    /// 最终返回的文档数量
30    pub top_n: usize,
31
32    /// 最小分数阈值(可选)
33    pub min_score: Option<f32>,
34
35    /// 是否保留原始分数
36    pub preserve_original_score: bool,
37}
38
39impl Default for RerankingConfig {
40    fn default() -> Self {
41        Self {
42            top_n: 5,
43            min_score: None,
44            preserve_original_score: true,
45        }
46    }
47}
48
49impl RerankingConfig {
50    pub fn new() -> Self {
51        Self::default()
52    }
53
54    pub fn with_top_n(mut self, n: usize) -> Self {
55        self.top_n = n;
56        self
57    }
58
59    pub fn with_min_score(mut self, score: f32) -> Self {
60        self.min_score = Some(score);
61        self
62    }
63
64    pub fn with_preserve_original_score(mut self, preserve: bool) -> Self {
65        self.preserve_original_score = preserve;
66        self
67    }
68}
69
70/// Reranking 评分器 trait
71pub trait Reranker: Send + Sync {
72    fn score(&self, query: &str, documents: &[Document]) -> Result<Vec<f32>, RerankingError>;
73}
74
75/// 基于关键词匹配的简单 Reranker
76pub struct KeywordReranker {
77    /// 关键词权重(可选)
78    keyword_weights: HashMap<String, f32>,
79}
80
81impl KeywordReranker {
82    pub fn new() -> Self {
83        Self {
84            keyword_weights: HashMap::new(),
85        }
86    }
87
88    pub fn with_keyword_weights(mut self, weights: HashMap<String, f32>) -> Self {
89        self.keyword_weights = weights;
90        self
91    }
92
93    fn extract_keywords(&self, query: &str) -> Vec<String> {
94        query
95            .split_whitespace()
96            .filter(|w| w.len() > 1)
97            .map(|w| w.to_lowercase())
98            .collect()
99    }
100
101    fn count_keyword_matches(&self, keywords: &[String], document: &Document) -> f32 {
102        let doc_lower = document.content.to_lowercase();
103        let mut score = 0.0;
104
105        for keyword in keywords {
106            let count = doc_lower.matches(keyword).count() as f32;
107            let weight = self.keyword_weights.get(keyword).unwrap_or(&1.0);
108            score += count * weight;
109        }
110
111        score
112    }
113}
114
115impl Default for KeywordReranker {
116    fn default() -> Self {
117        Self::new()
118    }
119}
120
121impl Reranker for KeywordReranker {
122    fn score(&self, query: &str, documents: &[Document]) -> Result<Vec<f32>, RerankingError> {
123        if documents.is_empty() {
124            return Ok(Vec::new());
125        }
126
127        let keywords = self.extract_keywords(query);
128
129        if keywords.is_empty() {
130            return Ok(documents.iter().map(|_| 0.0).collect());
131        }
132
133        let scores: Vec<f32> = documents
134            .iter()
135            .map(|doc| self.count_keyword_matches(&keywords, doc))
136            .collect();
137
138        Ok(scores)
139    }
140}
141
142/// Reranker 执行器
143pub struct RerankingExecutor {
144    reranker: Box<dyn Reranker>,
145    config: RerankingConfig,
146}
147
148impl RerankingExecutor {
149    pub fn new(reranker: Box<dyn Reranker>) -> Self {
150        Self {
151            reranker,
152            config: RerankingConfig::default(),
153        }
154    }
155
156    pub fn with_config(mut self, config: RerankingConfig) -> Self {
157        self.config = config;
158        self
159    }
160
161    pub fn with_top_n(mut self, n: usize) -> Self {
162        self.config.top_n = n;
163        self
164    }
165
166    pub fn with_min_score(mut self, score: f32) -> Self {
167        self.config.min_score = Some(score);
168        self
169    }
170
171    pub fn with_preserve_original_score(mut self, preserve: bool) -> Self {
172        self.config.preserve_original_score = preserve;
173        self
174    }
175
176    pub fn rerank(
177        &self,
178        query: &str,
179        results: Vec<SearchResult>,
180    ) -> Result<Vec<SearchResult>, RerankingError> {
181        if results.is_empty() {
182            return Ok(Vec::new());
183        }
184
185        let documents: Vec<Document> = results.iter().map(|r| r.document.clone()).collect();
186        let scores = self.reranker.score(query, &documents)?;
187
188        // Normalize scores to [0, 1] range before combining (H51)
189        let max_original = results
190            .iter()
191            .map(|r| r.score.abs())
192            .fold(0.0_f32, f32::max);
193        let max_rerank = scores.iter().map(|s| s.abs()).fold(0.0_f32, f32::max);
194
195        let mut reranked: Vec<SearchResult> = results
196            .iter()
197            .enumerate()
198            .map(|(idx, r)| {
199                let new_score = if self.config.preserve_original_score {
200                    let norm_original = if max_original > 0.0 {
201                        r.score / max_original
202                    } else {
203                        0.0
204                    };
205                    let norm_rerank = if max_rerank > 0.0 {
206                        scores[idx] / max_rerank
207                    } else {
208                        0.0
209                    };
210                    norm_original + norm_rerank
211                } else {
212                    scores[idx]
213                };
214
215                SearchResult {
216                    document: r.document.clone(),
217                    score: new_score,
218                }
219            })
220            .collect();
221
222        if let Some(min_score) = self.config.min_score {
223            reranked.retain(|r| r.score >= min_score);
224        }
225
226        reranked.sort_by(|a, b| {
227            b.score
228                .partial_cmp(&a.score)
229                .unwrap_or(std::cmp::Ordering::Equal)
230        });
231
232        reranked.truncate(self.config.top_n);
233
234        Ok(reranked)
235    }
236
237    pub fn rerank_documents(
238        &self,
239        query: &str,
240        documents: Vec<Document>,
241    ) -> Result<Vec<SearchResult>, RerankingError> {
242        if documents.is_empty() {
243            return Ok(Vec::new());
244        }
245
246        let scores = self.reranker.score(query, &documents)?;
247
248        let mut results: Vec<SearchResult> = documents
249            .iter()
250            .enumerate()
251            .map(|(idx, doc)| SearchResult {
252                document: doc.clone(),
253                score: scores[idx],
254            })
255            .collect();
256
257        if let Some(min_score) = self.config.min_score {
258            results.retain(|r| r.score >= min_score);
259        }
260
261        results.sort_by(|a, b| {
262            b.score
263                .partial_cmp(&a.score)
264                .unwrap_or(std::cmp::Ordering::Equal)
265        });
266
267        results.truncate(self.config.top_n);
268
269        Ok(results)
270    }
271}
272
273/// BM25-style Reranker(简化版)
274pub struct BM25Reranker {
275    k1: f32,
276    b: f32,
277}
278
279impl BM25Reranker {
280    pub fn new() -> Self {
281        Self { k1: 1.5, b: 0.75 }
282    }
283
284    pub fn with_params(mut self, k1: f32, b: f32) -> Self {
285        self.k1 = k1;
286        self.b = b;
287        self
288    }
289
290    fn tokenize(&self, text: &str) -> Vec<String> {
291        text.split_whitespace()
292            .filter(|w| w.len() > 1)
293            .map(|w| w.to_lowercase())
294            .collect()
295    }
296}
297
298impl Default for BM25Reranker {
299    fn default() -> Self {
300        Self::new()
301    }
302}
303
304impl Reranker for BM25Reranker {
305    fn score(&self, query: &str, documents: &[Document]) -> Result<Vec<f32>, RerankingError> {
306        if documents.is_empty() {
307            return Ok(Vec::new());
308        }
309
310        let query_terms = self.tokenize(query);
311
312        if query_terms.is_empty() {
313            return Ok(documents.iter().map(|_| 0.0).collect());
314        }
315
316        let avgdl = documents
317            .iter()
318            .map(|d| d.content.split_whitespace().count() as f32)
319            .sum::<f32>()
320            / documents.len() as f32;
321
322        let scores: Vec<f32> = documents
323            .iter()
324            .map(|doc| {
325                let doc_len = doc.content.split_whitespace().count() as f32;
326                let doc_lower = doc.content.to_lowercase();
327                query_terms
328                    .iter()
329                    .map(|term| {
330                        let freq = doc_lower.matches(term.as_str()).count() as f32;
331                        let tf =
332                            freq / (freq + self.k1 * (1.0 - self.b + self.b * doc_len / avgdl));
333                        tf * (1.0 + self.k1)
334                            / (tf + self.k1 * (1.0 - self.b + self.b * doc_len / avgdl))
335                    })
336                    .sum()
337            })
338            .collect();
339
340        Ok(scores)
341    }
342}
343
344#[cfg(test)]
345mod tests {
346    use super::*;
347
348    #[test]
349    fn test_reranking_config_default() {
350        let config = RerankingConfig::default();
351
352        assert_eq!(config.top_n, 5);
353        assert!(config.min_score.is_none());
354        assert!(config.preserve_original_score);
355    }
356
357    #[test]
358    fn test_reranking_config_custom() {
359        let config = RerankingConfig::new()
360            .with_top_n(10)
361            .with_min_score(0.5)
362            .with_preserve_original_score(false);
363
364        assert_eq!(config.top_n, 10);
365        assert_eq!(config.min_score, Some(0.5));
366        assert!(!config.preserve_original_score);
367    }
368
369    #[test]
370    fn test_keyword_reranker_basic() {
371        let reranker = KeywordReranker::new();
372
373        let query = "Rust programming";
374        let documents = vec![
375            Document::new("Rust is a programming language"),
376            Document::new("Python is also a programming language"),
377            Document::new("JavaScript for web"),
378        ];
379
380        let scores = reranker.score(query, &documents).unwrap();
381
382        assert_eq!(scores.len(), 3);
383        assert!(scores[0] > 0.0);
384        assert!(scores[1] > 0.0);
385    }
386
387    #[test]
388    fn test_keyword_reranker_empty_query() {
389        let reranker = KeywordReranker::new();
390
391        let documents = vec![Document::new("Some content")];
392
393        let scores = reranker.score("", &documents).unwrap();
394
395        assert_eq!(scores[0], 0.0);
396    }
397
398    #[test]
399    fn test_reranking_executor_basic() {
400        let reranker = Box::new(KeywordReranker::new());
401        let executor = RerankingExecutor::new(reranker).with_top_n(2);
402
403        let results = vec![
404            SearchResult {
405                document: Document::new("Rust programming language"),
406                score: 0.5,
407            },
408            SearchResult {
409                document: Document::new("Python scripting"),
410                score: 0.4,
411            },
412            SearchResult {
413                document: Document::new("JavaScript web"),
414                score: 0.3,
415            },
416        ];
417
418        let reranked = executor.rerank("Rust programming", results).unwrap();
419
420        assert_eq!(reranked.len(), 2);
421    }
422
423    #[test]
424    fn test_reranking_executor_min_score() {
425        let reranker = Box::new(KeywordReranker::new());
426        let executor = RerankingExecutor::new(reranker)
427            .with_top_n(5)
428            .with_min_score(1.0);
429
430        let results = vec![
431            SearchResult {
432                document: Document::new("Rust Rust Rust"),
433                score: 0.0,
434            },
435            SearchResult {
436                document: Document::new("No match"),
437                score: 0.0,
438            },
439        ];
440
441        let reranked = executor.rerank("Rust", results).unwrap();
442
443        assert!(reranked.len() <= 1);
444    }
445
446    #[test]
447    fn test_bm25_reranker_basic() {
448        let reranker = BM25Reranker::new();
449
450        let query = "programming language";
451        let documents = vec![
452            Document::new("Rust is a programming language"),
453            Document::new("Python is a programming language too"),
454            Document::new("Web development"),
455        ];
456
457        let scores = reranker.score(query, &documents).unwrap();
458
459        assert_eq!(scores.len(), 3);
460        assert!(scores[0] > scores[2]);
461    }
462
463    #[test]
464    fn test_bm25_reranker_params() {
465        let reranker = BM25Reranker::new().with_params(2.0, 0.5);
466
467        let documents = vec![Document::new("test content")];
468
469        let scores = reranker.score("test", &documents).unwrap();
470
471        assert!(scores[0] > 0.0);
472    }
473
474    #[test]
475    fn test_rerank_documents() {
476        let reranker = Box::new(KeywordReranker::new());
477        let executor = RerankingExecutor::new(reranker).with_top_n(2);
478
479        let documents = vec![
480            Document::new("Rust programming"),
481            Document::new("Python scripting"),
482            Document::new("JavaScript web"),
483        ];
484
485        let results = executor.rerank_documents("Rust", documents).unwrap();
486
487        assert_eq!(results.len(), 2);
488    }
489}