1use lc_vector_stores::{Document, SearchResult};
7use std::collections::HashMap;
8
9#[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
27pub struct RerankingConfig {
29 pub top_n: usize,
31
32 pub min_score: Option<f32>,
34
35 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
70pub trait Reranker: Send + Sync {
72 fn score(&self, query: &str, documents: &[Document]) -> Result<Vec<f32>, RerankingError>;
73}
74
75pub struct KeywordReranker {
77 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
142pub 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 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
273pub 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}