use lc_vector_stores::Document;
use std::collections::HashMap;
pub const RRF_K: usize = 60;
fn doc_content_hash(doc: &Document) -> String {
use std::hash::{Hash, Hasher};
let mut hasher = fnv::FnvHasher::default();
doc.content.hash(&mut hasher);
format!("{:016x}", hasher.finish())
}
#[derive(Debug, Clone)]
pub struct RetrievedDocument {
pub document: Document,
pub score: f64,
pub source: RetrievalSource,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum RetrievalSource {
BM25,
Vector,
Hybrid,
}
pub fn filter_by_score<T, S: PartialOrd>(scored: Vec<(T, S)>, min_score: S) -> Vec<(T, S)> {
scored.into_iter().filter(|(_, s)| *s > min_score).collect()
}
pub fn reciprocal_rank_fusion(
bm25_results: Vec<Document>,
vector_results: Vec<Document>,
k: usize,
) -> Vec<RetrievedDocument> {
let mut rrf_scores: HashMap<String, (f64, Document)> = HashMap::new();
for (rank, doc) in bm25_results.iter().enumerate() {
let doc_id = doc.id.clone().unwrap_or_else(|| doc_content_hash(doc));
let rrf_contribution = 1.0 / (k as f64 + (rank + 1) as f64);
rrf_scores
.entry(doc_id.clone())
.and_modify(|(score, _existing_doc)| {
*score += rrf_contribution;
})
.or_insert((rrf_contribution, doc.clone()));
}
for (rank, doc) in vector_results.iter().enumerate() {
let doc_id = doc.id.clone().unwrap_or_else(|| doc_content_hash(doc));
let rrf_contribution = 1.0 / (k as f64 + (rank + 1) as f64);
rrf_scores
.entry(doc_id.clone())
.and_modify(|(score, _)| {
*score += rrf_contribution;
})
.or_insert((rrf_contribution, doc.clone()));
}
let mut results: Vec<RetrievedDocument> = rrf_scores
.into_iter()
.map(|(_, (score, doc))| RetrievedDocument {
document: doc,
score,
source: RetrievalSource::Hybrid,
})
.collect();
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
results
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rrf_basic() {
let bm25_docs = vec![
Document::new("Rust系统编程").with_id("doc1"),
Document::new("Python数据科学").with_id("doc2"),
Document::new("Go并发编程").with_id("doc3"),
];
let vector_docs = vec![
Document::new("Rust系统编程").with_id("doc1"),
Document::new("JavaScript前端").with_id("doc4"),
Document::new("Python数据科学").with_id("doc2"),
];
let results = reciprocal_rank_fusion(bm25_docs, vector_docs, 60);
println!("RRF 融合结果:");
for (i, r) in results.iter().enumerate() {
println!(
" [{}] doc_id={}, score={:.4}",
i,
r.document.id.clone().unwrap_or_default(),
r.score
);
}
let first_doc_id = results[0].document.id.clone().unwrap_or_default();
println!("最高分文档: {}", first_doc_id);
}
#[test]
fn test_filter_by_score() {
let scored = vec![("a", 0.9_f32), ("b", 0.2), ("c", -0.3), ("d", 0.0)];
let filtered = filter_by_score(scored.clone(), 0.0);
let ids: Vec<&str> = filtered.iter().map(|(id, _)| *id).collect();
assert_eq!(ids, vec!["a", "b"]);
let relaxed = filter_by_score(scored.clone(), -0.5);
assert_eq!(relaxed.len(), 4);
let strict = filter_by_score(scored.clone(), 0.5);
let ids: Vec<&str> = strict.iter().map(|(id, _)| *id).collect();
assert_eq!(ids, vec!["a"]);
}
#[test]
fn test_filter_by_score_f64() {
let scored = vec![("x", 0.8_f64), ("y", 0.0), ("z", -0.5)];
let filtered = filter_by_score(scored, 0.0);
let ids: Vec<&str> = filtered.iter().map(|(id, _)| *id).collect();
assert_eq!(ids, vec!["x"]);
}
#[test]
fn test_doc_content_hash_stable() {
let content = "Rust 系统编程与并发";
let doc_a = Document::new(content.to_string());
let doc_b = Document::new(content.to_string());
let doc_c = Document::new("Python 数据科学");
let hash_a1 = doc_content_hash(&doc_a);
let hash_a2 = doc_content_hash(&doc_b);
assert_eq!(hash_a1, hash_a2, "相同内容应产生相同哈希");
let hash_c = doc_content_hash(&doc_c);
assert_ne!(hash_a1, hash_c, "不同内容应产生不同哈希");
assert_eq!(hash_a1.len(), 16, "应为 64 位哈希的 16 位十六进制表示");
}
}