use crate::model::CodeNode;
use serde::Serialize;
#[derive(Debug, Clone, Serialize)]
pub struct SearchHit {
pub node: CodeNode,
pub score: f64,
pub source: String,
#[serde(default)]
pub callers: Vec<String>,
#[serde(default)]
pub callees: Vec<String>,
}
pub fn rrf_merge(results: &[Vec<SearchHit>], top_k: usize, k: f64) -> Vec<SearchHit> {
use std::collections::HashMap;
let mut scores: HashMap<(String, String), (f64, CodeNode)> = HashMap::new();
for list in results {
for (rank, hit) in list.iter().enumerate() {
let key = (
hit.node.name.clone(),
hit.node.file_path.clone().unwrap_or_default(),
);
let entry = scores.entry(key).or_insert_with(|| (0.0, hit.node.clone()));
entry.0 += k / (k + rank as f64);
}
}
let mut merged: Vec<SearchHit> = scores.into_iter()
.map(|(_key, (score, node))| SearchHit {
node, score, source: "hybrid".into(),
callers: vec![], callees: vec![],
})
.collect();
merged.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(std::cmp::Ordering::Equal));
merged.truncate(top_k);
merged
}
pub fn text_results_to_hits(results: Vec<(CodeNode, f64)>) -> Vec<SearchHit> {
results.into_iter().map(|(node, score)| SearchHit {
node, score, source: "text".into(),
callers: vec![], callees: vec![],
}).collect()
}
pub fn semantic_results_to_hits(results: Vec<(CodeNode, f32)>) -> Vec<SearchHit> {
results.into_iter().map(|(node, score)| SearchHit {
node, score: score as f64, source: "semantic".into(),
callers: vec![], callees: vec![],
}).collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{NodeKind, NodeId};
fn make_node(name: &str) -> CodeNode {
CodeNode {
id: NodeId::new(0), kind: NodeKind::Function, name: name.into(),
file_path: None, line_range: None, doc_comment: None,
signature: None, module_path: vec![], visibility: None,
}
}
#[test]
fn test_rrf_merge_single_source() {
let hits = vec![
SearchHit { node: make_node("foo"), score: 1.0, source: "text".into(), callers: vec![], callees: vec![] },
SearchHit { node: make_node("bar"), score: 0.5, source: "text".into(), callers: vec![], callees: vec![] },
];
let result = rrf_merge(&[hits], 5, 60.0);
assert_eq!(result.len(), 2);
assert_eq!(result[0].node.name, "foo");
}
#[test]
fn test_rrf_merge_dedup() {
let t = vec![
SearchHit { node: make_node("foo"), score: 1.0, source: "text".into(), callers: vec![], callees: vec![] },
];
let s = vec![
SearchHit { node: make_node("foo"), score: 0.9, source: "semantic".into(), callers: vec![], callees: vec![] },
];
let result = rrf_merge(&[t, s], 5, 60.0);
assert_eq!(result.len(), 1);
}
#[test]
fn test_rrf_merge_keeps_same_name_different_file() {
let mut a = make_node("foo");
a.file_path = Some("src/a.rs".into());
let mut b = make_node("foo");
b.file_path = Some("src/b.rs".into());
let text = vec![
SearchHit { node: a.clone(), score: 1.0, source: "text".into(), callers: vec![], callees: vec![] },
];
let semantic = vec![
SearchHit { node: b, score: 0.9, source: "semantic".into(), callers: vec![], callees: vec![] },
];
let result = rrf_merge(&[text, semantic], 5, 60.0);
assert_eq!(result.len(), 2, "同名不同文件的实体不应被折叠");
}
#[test]
fn test_rrf_merge_cross_engine_ranking() {
let text = vec![
SearchHit { node: make_node("a"), score: 1.0, source: "text".into(), callers: vec![], callees: vec![] },
SearchHit { node: make_node("b"), score: 0.9, source: "text".into(), callers: vec![], callees: vec![] },
];
let semantic = vec![
SearchHit { node: make_node("b"), score: 1.0, source: "semantic".into(), callers: vec![], callees: vec![] },
SearchHit { node: make_node("x"), score: 0.8, source: "semantic".into(), callers: vec![], callees: vec![] },
SearchHit { node: make_node("a"), score: 0.6, source: "semantic".into(), callers: vec![], callees: vec![] },
];
let f1 = rrf_merge(&[text.clone(), semantic.clone()], 5, 60.0);
let f2 = rrf_merge(&[semantic, text], 5, 60.0);
assert_eq!(f1.len(), f2.len());
for (h1, h2) in f1.iter().zip(f2.iter()) {
assert_eq!(h1.node.name, h2.node.name);
assert!((h1.score - h2.score).abs() < 1e-9);
}
assert_eq!(f1[0].node.name, "b");
}
#[test]
fn test_text_results_to_hits() {
let r = vec![(make_node("x"), 2.0)];
let hits = text_results_to_hits(r);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].source, "text");
}
#[test]
fn test_semantic_results_to_hits() {
let r = vec![(make_node("x"), 0.9)];
let hits = semantic_results_to_hits(r);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].source, "semantic");
}
}