1use lc_vector_stores::Document;
7use std::collections::HashMap;
8
9pub const RRF_K: usize = 60;
11
12fn doc_content_hash(doc: &Document) -> String {
18 use std::hash::{Hash, Hasher};
19 let mut hasher = fnv::FnvHasher::default();
20 doc.content.hash(&mut hasher);
21 format!("{:016x}", hasher.finish())
22}
23
24#[derive(Debug, Clone)]
26pub struct RetrievedDocument {
27 pub document: Document,
29 pub score: f64,
31 pub source: RetrievalSource,
33}
34
35#[derive(Debug, Clone, Copy, PartialEq)]
37pub enum RetrievalSource {
38 BM25,
40 Vector,
42 Hybrid,
44}
45
46pub fn filter_by_score<T, S: PartialOrd>(scored: Vec<(T, S)>, min_score: S) -> Vec<(T, S)> {
53 scored.into_iter().filter(|(_, s)| *s > min_score).collect()
54}
55
56pub fn reciprocal_rank_fusion(
68 bm25_results: Vec<Document>,
69 vector_results: Vec<Document>,
70 k: usize,
71) -> Vec<RetrievedDocument> {
72 let mut rrf_scores: HashMap<String, (f64, Document)> = HashMap::new();
73
74 for (rank, doc) in bm25_results.iter().enumerate() {
76 let doc_id = doc.id.clone().unwrap_or_else(|| doc_content_hash(doc));
77 let rrf_contribution = 1.0 / (k as f64 + (rank + 1) as f64);
78
79 rrf_scores
80 .entry(doc_id.clone())
81 .and_modify(|(score, _existing_doc)| {
82 *score += rrf_contribution;
83 })
84 .or_insert((rrf_contribution, doc.clone()));
85 }
86
87 for (rank, doc) in vector_results.iter().enumerate() {
89 let doc_id = doc.id.clone().unwrap_or_else(|| doc_content_hash(doc));
90 let rrf_contribution = 1.0 / (k as f64 + (rank + 1) as f64);
91
92 rrf_scores
93 .entry(doc_id.clone())
94 .and_modify(|(score, _)| {
95 *score += rrf_contribution;
96 })
97 .or_insert((rrf_contribution, doc.clone()));
98 }
99
100 let mut results: Vec<RetrievedDocument> = rrf_scores
102 .into_iter()
103 .map(|(_, (score, doc))| RetrievedDocument {
104 document: doc,
105 score,
106 source: RetrievalSource::Hybrid,
107 })
108 .collect();
109
110 results.sort_by(|a, b| {
111 b.score
112 .partial_cmp(&a.score)
113 .unwrap_or(std::cmp::Ordering::Equal)
114 });
115
116 results
117}
118
119#[cfg(test)]
120mod tests {
121 use super::*;
122
123 #[test]
124 fn test_rrf_basic() {
125 let bm25_docs = vec![
126 Document::new("Rust系统编程").with_id("doc1"),
127 Document::new("Python数据科学").with_id("doc2"),
128 Document::new("Go并发编程").with_id("doc3"),
129 ];
130
131 let vector_docs = vec![
132 Document::new("Rust系统编程").with_id("doc1"),
133 Document::new("JavaScript前端").with_id("doc4"),
134 Document::new("Python数据科学").with_id("doc2"),
135 ];
136
137 let results = reciprocal_rank_fusion(bm25_docs, vector_docs, 60);
138
139 println!("RRF 融合结果:");
140 for (i, r) in results.iter().enumerate() {
141 println!(
142 " [{}] doc_id={}, score={:.4}",
143 i,
144 r.document.id.clone().unwrap_or_default(),
145 r.score
146 );
147 }
148
149 let first_doc_id = results[0].document.id.clone().unwrap_or_default();
151 println!("最高分文档: {}", first_doc_id);
152 }
153
154 #[test]
157 fn test_filter_by_score() {
158 let scored = vec![("a", 0.9_f32), ("b", 0.2), ("c", -0.3), ("d", 0.0)];
159
160 let filtered = filter_by_score(scored.clone(), 0.0);
162 let ids: Vec<&str> = filtered.iter().map(|(id, _)| *id).collect();
163 assert_eq!(ids, vec!["a", "b"]);
164
165 let relaxed = filter_by_score(scored.clone(), -0.5);
167 assert_eq!(relaxed.len(), 4);
168
169 let strict = filter_by_score(scored.clone(), 0.5);
171 let ids: Vec<&str> = strict.iter().map(|(id, _)| *id).collect();
172 assert_eq!(ids, vec!["a"]);
173 }
174
175 #[test]
177 fn test_filter_by_score_f64() {
178 let scored = vec![("x", 0.8_f64), ("y", 0.0), ("z", -0.5)];
179 let filtered = filter_by_score(scored, 0.0);
180 let ids: Vec<&str> = filtered.iter().map(|(id, _)| *id).collect();
181 assert_eq!(ids, vec!["x"]);
182 }
183
184 #[test]
187 fn test_doc_content_hash_stable() {
188 let content = "Rust 系统编程与并发";
189 let doc_a = Document::new(content.to_string());
190 let doc_b = Document::new(content.to_string());
191 let doc_c = Document::new("Python 数据科学");
192
193 let hash_a1 = doc_content_hash(&doc_a);
194 let hash_a2 = doc_content_hash(&doc_b);
195 assert_eq!(hash_a1, hash_a2, "相同内容应产生相同哈希");
196
197 let hash_c = doc_content_hash(&doc_c);
198 assert_ne!(hash_a1, hash_c, "不同内容应产生不同哈希");
199 assert_eq!(hash_a1.len(), 16, "应为 64 位哈希的 16 位十六进制表示");
200 }
201}