use std::path::Path;
use std::sync::Arc;
use std::time::Instant;
use rustc_hash::{FxHashMap, FxHashSet};
use crate::config::bm25::BM25;
use crate::config::limits::{LIMITS, PPR};
use crate::config::scoring::{EGO, pit, rrf};
use crate::config::tokenization::TOKENIZATION;
use crate::edges;
use crate::filtering;
use crate::graph::{self, Graph};
use crate::mode::{PipelineConfig, ScoringKind};
use crate::ppr::personalized_pagerank;
use crate::types::{DiffHunk, Fragment, FragmentId, extract_identifier_list};
pub struct ScoringResult {
pub rel_scores: FxHashMap<FragmentId, f64>,
pub filtered_fragments: Vec<Fragment>,
pub graph: Graph,
pub graph_build_ms: f64,
pub ppr_truncated: bool,
pub ppr_forward_pushes: usize,
pub ppr_backward_pushes: usize,
}
fn finish_scoring(
fragments: &[Fragment],
core_ids: &FxHashSet<FragmentId>,
rel_scores: &FxHashMap<FragmentId, f64>,
graph: &Graph,
) -> Vec<Fragment> {
let filtered = filtering::filter_unrelated_fragments(fragments, core_ids, graph);
let filtered = filtering::filter_positive_relevance(filtered, core_ids, rel_scores);
let filtered = filtering::filter_core_slice_context(filtered, core_ids);
filtering::cap_context_fragments(filtered, core_ids, rel_scores)
}
impl Default for ScoringResult {
fn default() -> Self {
Self {
rel_scores: FxHashMap::default(),
filtered_fragments: Vec::new(),
graph: Graph::new(),
graph_build_ms: 0.0,
ppr_truncated: false,
ppr_forward_pushes: 0,
ppr_backward_pushes: 0,
}
}
}
pub fn create_scoring_strategy(config: &PipelineConfig) -> Box<dyn ScoringStrategy> {
match config.scoring {
ScoringKind::Ego => Box::new(EgoGraphScoring::new(config.ego_depth)),
ScoringKind::Ppr => Box::new(PPRScoring::new(config.ppr_alpha)),
ScoringKind::Bm25 => Box::new(BM25Scoring),
ScoringKind::Rrf => Box::new(RrfFusionScoring::new(config.ego_depth)),
ScoringKind::Pit => Box::new(PitFusionScoring::new(config.ego_depth)),
}
}
pub trait ScoringStrategy: Send + Sync {
fn score_and_filter(
&self,
all_fragments: &[Fragment],
core_ids: &FxHashSet<FragmentId>,
hunks: &[DiffHunk],
repo_root: Option<&Path>,
seed_weights: Option<&FxHashMap<FragmentId, f64>>,
discovered_paths: Option<&FxHashSet<Arc<str>>>,
) -> ScoringResult;
}
pub struct PPRScoring {
pub alpha: f64,
}
impl PPRScoring {
pub fn new(alpha: f64) -> Self {
Self { alpha }
}
}
impl ScoringStrategy for PPRScoring {
fn score_and_filter(
&self,
all_fragments: &[Fragment],
core_ids: &FxHashSet<FragmentId>,
hunks: &[DiffHunk],
repo_root: Option<&Path>,
seed_weights: Option<&FxHashMap<FragmentId, f64>>,
_discovered_paths: Option<&FxHashSet<Arc<str>>>,
) -> ScoringResult {
let skip_expensive = all_fragments.len() > LIMITS.skip_expensive_threshold;
let t_graph = Instant::now();
let capped = edges::collect_capped_edges(all_fragments, repo_root, skip_expensive);
let mut g = graph::build_graph_capped(all_fragments, capped);
let graph_build_ms = t_graph.elapsed().as_secs_f64() * 1000.0;
let ppr = personalized_pagerank(
&mut g,
core_ids,
self.alpha,
PPR.convergence_tolerance,
PPR.forward_blend,
seed_weights,
);
let mut rel_scores = ppr.scores;
if ppr.truncated {
tracing::warn!(
"PPR push-cap hit on {} nodes (fwd_pushes={}, bwd_pushes={}); rel_scores biased",
g.node_count(),
ppr.forward_pushes,
ppr.backward_pushes,
);
}
filtering::apply_hunk_proximity_bonus(&mut rel_scores, core_ids, all_fragments, hunks);
let filtered = finish_scoring(all_fragments, core_ids, &rel_scores, &g);
ScoringResult {
rel_scores,
filtered_fragments: filtered,
graph: g,
graph_build_ms,
ppr_truncated: ppr.truncated,
ppr_forward_pushes: ppr.forward_pushes,
ppr_backward_pushes: ppr.backward_pushes,
}
}
}
pub struct EgoGraphScoring {
pub max_depth: usize,
}
impl EgoGraphScoring {
pub fn new(max_depth: usize) -> Self {
Self { max_depth }
}
}
impl ScoringStrategy for EgoGraphScoring {
fn score_and_filter(
&self,
all_fragments: &[Fragment],
core_ids: &FxHashSet<FragmentId>,
_hunks: &[DiffHunk],
repo_root: Option<&Path>,
_seed_weights: Option<&FxHashMap<FragmentId, f64>>,
_discovered_paths: Option<&FxHashSet<Arc<str>>>,
) -> ScoringResult {
let skip_expensive = all_fragments.len() > LIMITS.skip_expensive_threshold;
let t_graph = Instant::now();
let capped = edges::collect_capped_edges(all_fragments, repo_root, skip_expensive);
let g = graph::build_graph_capped(all_fragments, capped);
let graph_build_ms = t_graph.elapsed().as_secs_f64() * 1000.0;
let mut rel_scores = g.ego_graph(core_ids, self.max_depth);
let diff_idents: FxHashSet<String> = all_fragments
.iter()
.filter(|f| core_ids.contains(&f.id))
.flat_map(|f| f.identifiers.iter().cloned())
.collect();
if !diff_idents.is_empty() {
for frag in all_fragments {
if core_ids.contains(&frag.id) || !rel_scores.contains_key(&frag.id) {
continue;
}
let overlap = frag.identifiers.intersection(&diff_idents).count();
if overlap > 0 {
let bonus = EGO.identifier_overlap_epsilon
* overlap.min(EGO.identifier_overlap_cap) as f64
/ EGO.identifier_overlap_cap as f64;
*rel_scores.get_mut(&frag.id).unwrap() += bonus;
}
}
}
let filtered = finish_scoring(all_fragments, core_ids, &rel_scores, &g);
ScoringResult {
rel_scores,
filtered_fragments: filtered,
graph: g,
graph_build_ms,
..Default::default()
}
}
}
pub struct BM25Scoring;
impl ScoringStrategy for BM25Scoring {
fn score_and_filter(
&self,
all_fragments: &[Fragment],
core_ids: &FxHashSet<FragmentId>,
_hunks: &[DiffHunk],
_repo_root: Option<&Path>,
_seed_weights: Option<&FxHashMap<FragmentId, f64>>,
_discovered_paths: Option<&FxHashSet<Arc<str>>>,
) -> ScoringResult {
let query_tokens: Vec<String> = all_fragments
.iter()
.filter(|f| core_ids.contains(&f.id))
.flat_map(|f| {
extract_identifier_list(&f.content, TOKENIZATION.query_min_identifier_length)
})
.collect();
let query_set: FxHashSet<String> = query_tokens.into_iter().collect();
let docs: Vec<(FragmentId, Vec<String>)> = all_fragments
.iter()
.filter(|f| !core_ids.contains(&f.id))
.map(|f| {
(
f.id.clone(),
extract_identifier_list(&f.content, TOKENIZATION.query_min_identifier_length),
)
})
.collect();
let n_docs = docs.len().max(1);
let avgdl = docs.iter().map(|(_, d)| d.len()).sum::<usize>() as f64 / n_docs as f64;
let mut df: FxHashMap<String, usize> = FxHashMap::default();
for (_, doc) in &docs {
let unique: FxHashSet<&str> = doc.iter().map(|s| s.as_str()).collect();
for term in unique {
*df.entry(term.to_string()).or_insert(0) += 1;
}
}
let idf: FxHashMap<String, f64> = query_set
.iter()
.map(|t| {
let d = df.get(t).copied().unwrap_or(0) as f64;
let val =
((n_docs as f64 - d + BM25.idf_smoothing) / (d + BM25.idf_smoothing)).ln_1p();
(t.clone(), val)
})
.collect();
let mut rel_scores: FxHashMap<FragmentId, f64> = FxHashMap::default();
for frag in all_fragments {
if core_ids.contains(&frag.id) {
rel_scores.insert(frag.id.clone(), 1.0);
}
}
for (fid, doc) in &docs {
let dl = doc.len() as f64;
let mut tf: FxHashMap<&str, u32> = FxHashMap::default();
for t in doc {
*tf.entry(t.as_str()).or_insert(0) += 1;
}
let mut score = 0.0;
for t in &query_set {
let freq = tf.get(t.as_str()).copied().unwrap_or(0) as f64;
if freq == 0.0 {
continue;
}
let idf_val = idf.get(t).copied().unwrap_or(0.0);
score += idf_val * (freq * BM25.k1)
/ (freq + BM25.k1 * (1.0 - BM25.b + BM25.b * dl / avgdl));
}
if score > 0.0 {
rel_scores.insert(fid.clone(), score);
}
}
let max_score = rel_scores.values().copied().fold(0.0f64, f64::max);
if max_score > 0.0 {
for v in rel_scores.values_mut() {
*v /= max_score;
}
}
let filtered =
filtering::filter_positive_relevance(all_fragments.to_vec(), core_ids, &rel_scores);
let filtered = filtering::cap_context_fragments(filtered, core_ids, &rel_scores);
ScoringResult {
rel_scores,
filtered_fragments: filtered,
..Default::default()
}
}
}
pub struct RrfFusionScoring {
pub ego_depth: usize,
pub k: f64,
}
impl RrfFusionScoring {
pub fn new(ego_depth: usize) -> Self {
Self {
ego_depth,
k: rrf().k,
}
}
}
fn rank_positions(
rel: &FxHashMap<FragmentId, f64>,
admitted: &FxHashSet<FragmentId>,
core_ids: &FxHashSet<FragmentId>,
) -> FxHashMap<FragmentId, usize> {
let mut ranked: Vec<(&FragmentId, f64)> = rel
.iter()
.filter(|(fid, score)| **score > 0.0 && !core_ids.contains(*fid) && admitted.contains(*fid))
.map(|(fid, score)| (fid, *score))
.collect();
ranked.sort_by(|(ida, sa), (idb, sb)| sb.total_cmp(sa).then_with(|| ida.cmp(idb)));
ranked
.into_iter()
.enumerate()
.map(|(i, (fid, _))| (fid.clone(), i + 1))
.collect()
}
fn fuse_reciprocal_ranks(
components: &[(&FxHashMap<FragmentId, f64>, &FxHashSet<FragmentId>)],
core_ids: &FxHashSet<FragmentId>,
k: f64,
) -> FxHashMap<FragmentId, f64> {
let mut fused: FxHashMap<FragmentId, f64> = FxHashMap::default();
for (rel, admitted) in components {
for (fid, rank) in rank_positions(rel, admitted, core_ids) {
*fused.entry(fid).or_insert(0.0) += 1.0 / (k + rank as f64);
}
}
let max_fused = fused.values().copied().fold(0.0f64, f64::max);
if max_fused > 0.0 {
for v in fused.values_mut() {
*v /= max_fused;
}
}
for fid in core_ids {
fused.insert(fid.clone(), 1.0);
}
fused
}
impl ScoringStrategy for RrfFusionScoring {
fn score_and_filter(
&self,
all_fragments: &[Fragment],
core_ids: &FxHashSet<FragmentId>,
hunks: &[DiffHunk],
repo_root: Option<&Path>,
seed_weights: Option<&FxHashMap<FragmentId, f64>>,
discovered_paths: Option<&FxHashSet<Arc<str>>>,
) -> ScoringResult {
let ego = EgoGraphScoring::new(self.ego_depth).score_and_filter(
all_fragments,
core_ids,
hunks,
repo_root,
seed_weights,
discovered_paths,
);
let lexical = BM25Scoring.score_and_filter(
all_fragments,
core_ids,
hunks,
repo_root,
seed_weights,
discovered_paths,
);
let ego_admitted: FxHashSet<FragmentId> = ego
.filtered_fragments
.iter()
.map(|f| f.id.clone())
.collect();
let lexical_admitted: FxHashSet<FragmentId> = lexical
.filtered_fragments
.iter()
.map(|f| f.id.clone())
.collect();
let rel_scores = fuse_reciprocal_ranks(
&[
(&ego.rel_scores, &ego_admitted),
(&lexical.rel_scores, &lexical_admitted),
],
core_ids,
self.k,
);
let union_ids: FxHashSet<FragmentId> =
ego_admitted.union(&lexical_admitted).cloned().collect();
let union: Vec<Fragment> = all_fragments
.iter()
.filter(|f| union_ids.contains(&f.id))
.cloned()
.collect();
let filtered = finish_scoring(&union, core_ids, &rel_scores, &ego.graph);
ScoringResult {
rel_scores,
filtered_fragments: filtered,
graph: ego.graph,
graph_build_ms: ego.graph_build_ms,
..Default::default()
}
}
}
pub struct PitFusionScoring {
pub ego_depth: usize,
pub blend: f64,
pub agreement_bonus: f64,
pub agreement_top_k: usize,
}
impl PitFusionScoring {
pub fn new(ego_depth: usize) -> Self {
let cfg = pit();
Self {
ego_depth,
blend: cfg.blend,
agreement_bonus: cfg.agreement_bonus,
agreement_top_k: cfg.agreement_top_k,
}
}
}
fn percentiles(
rel: &FxHashMap<FragmentId, f64>,
admitted: &FxHashSet<FragmentId>,
core_ids: &FxHashSet<FragmentId>,
top_k: usize,
) -> (FxHashMap<FragmentId, f64>, FxHashSet<FragmentId>) {
let mut ranked: Vec<(&FragmentId, f64)> = rel
.iter()
.filter(|(fid, score)| **score > 0.0 && !core_ids.contains(*fid))
.map(|(fid, score)| (fid, *score))
.collect();
ranked.sort_by(|(ida, sa), (idb, sb)| sb.total_cmp(sa).then_with(|| ida.cmp(idb)));
let n = ranked.len();
let mut out: FxHashMap<FragmentId, f64> = FxHashMap::default();
let mut top: FxHashSet<FragmentId> = FxHashSet::default();
if n == 0 {
return (out, top);
}
if std::env::var("DIFFCTX_PIT_TRANSFORM").as_deref() == Ok("maxnorm") {
let denom = ranked
.iter()
.map(|(_, s)| *s)
.fold(0.0f64, f64::max)
.max(f64::MIN_POSITIVE);
for (fid, score) in &ranked {
if admitted.contains(*fid) {
out.insert((*fid).clone(), *score / denom);
}
}
} else {
let mut i = 0;
while i < n {
let mut j = i;
while j + 1 < n && ranked[j + 1].1.to_bits() == ranked[i].1.to_bits() {
j += 1;
}
let mean_pos = (i + j) as f64 / 2.0;
let percentile = 1.0 - mean_pos / n as f64;
for (fid, _) in &ranked[i..=j] {
if admitted.contains(*fid) {
out.insert((*fid).clone(), percentile);
}
}
i = j + 1;
}
}
for (fid, _) in ranked
.iter()
.filter(|(fid, _)| admitted.contains(*fid))
.take(top_k)
{
top.insert((*fid).clone());
}
(out, top)
}
fn quantile_map_to(
reference: &[f64],
fused: &FxHashMap<FragmentId, f64>,
) -> FxHashMap<FragmentId, f64> {
let mut order: Vec<(&FragmentId, f64)> = fused
.iter()
.filter(|(_, s)| **s > 0.0)
.map(|(f, s)| (f, *s))
.collect();
if reference.is_empty() || order.is_empty() {
return order.into_iter().map(|(f, s)| (f.clone(), s)).collect();
}
order.sort_by(|(ida, sa), (idb, sb)| sb.total_cmp(sa).then_with(|| ida.cmp(idb)));
let m = order.len();
let n = reference.len();
let mut out = FxHashMap::default();
let mut i = 0;
while i < m {
let mut j = i;
while j + 1 < m && order[j + 1].1.to_bits() == order[i].1.to_bits() {
j += 1;
}
let mid_rank = (i + j) as f64 / 2.0;
let idx = if m == 1 {
n - 1
} else {
let pos = ((m - 1) as f64 - mid_rank) / (m - 1) as f64;
((pos * (n - 1) as f64).round() as usize).min(n - 1)
};
for (fid, _) in &order[i..=j] {
out.insert((*fid).clone(), reference[idx]);
}
i = j + 1;
}
out
}
impl ScoringStrategy for PitFusionScoring {
fn score_and_filter(
&self,
all_fragments: &[Fragment],
core_ids: &FxHashSet<FragmentId>,
hunks: &[DiffHunk],
repo_root: Option<&Path>,
seed_weights: Option<&FxHashMap<FragmentId, f64>>,
discovered_paths: Option<&FxHashSet<Arc<str>>>,
) -> ScoringResult {
let ego = EgoGraphScoring::new(self.ego_depth).score_and_filter(
all_fragments,
core_ids,
hunks,
repo_root,
seed_weights,
discovered_paths,
);
let lexical = BM25Scoring.score_and_filter(
all_fragments,
core_ids,
hunks,
repo_root,
seed_weights,
discovered_paths,
);
let ego_admitted: FxHashSet<FragmentId> = ego
.filtered_fragments
.iter()
.map(|f| f.id.clone())
.collect();
let lexical_admitted: FxHashSet<FragmentId> = lexical
.filtered_fragments
.iter()
.map(|f| f.id.clone())
.collect();
let (ego_pct, ego_top) = percentiles(
&ego.rel_scores,
&ego_admitted,
core_ids,
self.agreement_top_k,
);
let (lex_pct, lex_top) = percentiles(
&lexical.rel_scores,
&lexical_admitted,
core_ids,
self.agreement_top_k,
);
let mut rel_scores: FxHashMap<FragmentId, f64> = FxHashMap::default();
for fid in ego_pct.keys().chain(lex_pct.keys()) {
if rel_scores.contains_key(fid) {
continue;
}
let e = ego_pct.get(fid).copied().unwrap_or(0.0);
let l = lex_pct.get(fid).copied().unwrap_or(0.0);
let mut score = self.blend * e + (1.0 - self.blend) * l;
if ego_top.contains(fid) && lex_top.contains(fid) {
score += self.agreement_bonus;
}
rel_scores.insert(fid.clone(), score);
}
if std::env::var("DIFFCTX_PIT_SHAPE").as_deref() == Ok("flat") {
let max_fused = rel_scores.values().copied().fold(0.0f64, f64::max);
if max_fused > 0.0 {
for v in rel_scores.values_mut() {
*v /= max_fused;
}
}
for fid in core_ids {
rel_scores.insert(fid.clone(), 1.0);
}
} else {
let mut reference: Vec<f64> = ego
.rel_scores
.iter()
.filter(|(fid, s)| {
**s > 0.0 && !core_ids.contains(*fid) && ego_admitted.contains(*fid)
})
.map(|(_, s)| *s)
.collect();
reference.sort_by(f64::total_cmp);
rel_scores = quantile_map_to(&reference, &rel_scores);
for fid in core_ids {
if let Some(s) = ego.rel_scores.get(fid) {
rel_scores.insert(fid.clone(), *s);
}
}
}
let union_ids: FxHashSet<FragmentId> =
ego_admitted.union(&lexical_admitted).cloned().collect();
let union: Vec<Fragment> = all_fragments
.iter()
.filter(|f| union_ids.contains(&f.id))
.cloned()
.collect();
let filtered = finish_scoring(&union, core_ids, &rel_scores, &ego.graph);
ScoringResult {
rel_scores,
filtered_fragments: filtered,
graph: ego.graph,
graph_build_ms: ego.graph_build_ms,
..Default::default()
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::FragmentKind;
fn fid(path: &str, start: u32) -> FragmentId {
FragmentId::new(Arc::from(path), start, start + 4)
}
fn scores(entries: &[(FragmentId, f64)]) -> FxHashMap<FragmentId, f64> {
entries.iter().cloned().collect()
}
fn ballot(rel: &FxHashMap<FragmentId, f64>) -> FxHashSet<FragmentId> {
rel.keys().cloned().collect()
}
#[test]
fn agreement_between_both_signals_outranks_a_single_signal_top_hit() {
let agreed = fid("agreed.rs", 1);
let ego_only = fid("ego_only.rs", 1);
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let ego = scores(&[(ego_only.clone(), 0.9), (agreed.clone(), 0.5)]);
let lexical = scores(&[(agreed.clone(), 0.5)]);
let fused = fuse_reciprocal_ranks(
&[(&ego, &ballot(&ego)), (&lexical, &ballot(&lexical))],
&cores,
60.0,
);
assert!(
fused[&agreed] > fused[&ego_only],
"agreement lost to a single-signal top hit: {:?} vs {:?}",
fused[&agreed],
fused[&ego_only]
);
}
#[test]
fn only_the_rank_order_of_a_component_matters_not_its_scale() {
let a = fid("a.rs", 1);
let b = fid("b.rs", 1);
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let lexical = scores(&[(a.clone(), 0.4), (b.clone(), 0.1)]);
let modest = scores(&[(a.clone(), 0.6), (b.clone(), 0.4)]);
let enormous = scores(&[(a.clone(), 6_000.0), (b.clone(), 4_000.0)]);
assert_eq!(
fuse_reciprocal_ranks(
&[(&modest, &ballot(&modest)), (&lexical, &ballot(&lexical))],
&cores,
60.0
),
fuse_reciprocal_ranks(
&[
(&enormous, &ballot(&enormous)),
(&lexical, &ballot(&lexical))
],
&cores,
60.0
),
"rescaling one component changed the fused scores"
);
}
#[test]
fn fused_scores_are_normalised_and_cores_anchor_the_top() {
let core = fid("changed.rs", 1);
let ctx = fid("ctx.rs", 1);
let cores: FxHashSet<FragmentId> = std::iter::once(core.clone()).collect();
let ego = scores(&[(core.clone(), 1.0), (ctx.clone(), 0.3)]);
let lexical = scores(&[(ctx.clone(), 0.2)]);
let fused = fuse_reciprocal_ranks(
&[(&ego, &ballot(&ego)), (&lexical, &ballot(&lexical))],
&cores,
60.0,
);
assert_eq!(fused[&core], 1.0, "core is not anchored at the top");
for (id, score) in &fused {
assert!(
(0.0..=1.0).contains(score),
"{id:?} scored {score} outside [0, 1]"
);
}
}
#[test]
fn cores_do_not_occupy_rank_positions() {
let core = fid("changed.rs", 1);
let ctx = fid("ctx.rs", 1);
let cores: FxHashSet<FragmentId> = std::iter::once(core.clone()).collect();
let with_core = scores(&[(core.clone(), 1.0), (ctx.clone(), 0.3)]);
let without_core = scores(&[(ctx.clone(), 0.3)]);
assert_eq!(
rank_positions(&with_core, &ballot(&with_core), &cores).get(&ctx),
rank_positions(&without_core, &ballot(&without_core), &cores).get(&ctx),
"a core shifted the rank of a context fragment"
);
}
#[test]
fn equal_component_scores_rank_deterministically() {
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let ids: Vec<FragmentId> = (0..8).map(|i| fid(&format!("f{i}.rs"), 1)).collect();
let tied: FxHashMap<FragmentId, f64> = ids.iter().map(|i| (i.clone(), 0.5)).collect();
let baseline = rank_positions(&tied, &ballot(&tied), &cores);
for _ in 0..8 {
let again: FxHashMap<FragmentId, f64> =
ids.iter().rev().map(|i| (i.clone(), 0.5)).collect();
assert_eq!(rank_positions(&again, &ballot(&again), &cores), baseline);
}
let mut sorted = ids.clone();
sorted.sort();
for (expected_rank, id) in sorted.iter().enumerate() {
assert_eq!(baseline[id], expected_rank + 1);
}
}
#[test]
fn non_positive_component_scores_are_not_ranked() {
let kept = fid("kept.rs", 1);
let zero = fid("zero.rs", 1);
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let component = scores(&[(kept.clone(), 0.3), (zero.clone(), 0.0)]);
let ranks = rank_positions(&component, &ballot(&component), &cores);
assert!(ranks.contains_key(&kept));
assert!(
!ranks.contains_key(&zero),
"a zero-scored fragment was ranked"
);
let fused = fuse_reciprocal_ranks(&[(&component, &ballot(&component))], &cores, 60.0);
assert!(
!fused.contains_key(&zero),
"a zero-scored fragment was fused"
);
}
#[test]
fn a_larger_k_flattens_the_gap_between_adjacent_ranks() {
let first = fid("first.rs", 1);
let second = fid("second.rs", 1);
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let component = scores(&[(first.clone(), 0.9), (second.clone(), 0.8)]);
let sharp = fuse_reciprocal_ranks(&[(&component, &ballot(&component))], &cores, 1.0);
let flat = fuse_reciprocal_ranks(&[(&component, &ballot(&component))], &cores, 60.0);
assert!(
flat[&second] > sharp[&second],
"k did not flatten adjacent ranks: {} vs {}",
flat[&second],
sharp[&second]
);
}
#[test]
fn a_component_cannot_vote_for_what_its_own_filters_rejected() {
let good = fid("good.rs", 1);
let rejected = fid("garbage.rs", 1);
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let component = scores(&[(rejected.clone(), 0.9), (good.clone(), 0.1)]);
let admitted: FxHashSet<FragmentId> = std::iter::once(good.clone()).collect();
let fused = fuse_reciprocal_ranks(&[(&component, &admitted)], &cores, 60.0);
assert!(
!fused.contains_key(&rejected),
"a fragment the component filtered out still earned fused mass {:?}",
fused.get(&rejected)
);
assert!(
fused.contains_key(&good),
"the admitted fragment lost its vote"
);
}
#[test]
fn a_percentile_is_a_position_in_the_full_population_not_the_admitted_subset() {
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let mut entries: Vec<(FragmentId, f64)> = (0..9)
.map(|i| (fid("strong.rs", i + 1), 1.0 - i as f64 * 0.05))
.collect();
let weak = fid("weak.rs", 100);
entries.push((weak.clone(), 0.01));
let rel = scores(&entries);
let admitted: FxHashSet<FragmentId> = std::iter::once(weak.clone()).collect();
let (pct, _) = percentiles(&rel, &admitted, &cores, 5);
assert_eq!(pct.len(), 1, "only admitted fragments may receive a value");
let p = pct[&weak];
assert!(
p <= 0.2,
"the weakest of ten scored {p}, reading as strong because the CDF \
was estimated over the admitted subset"
);
}
#[test]
fn a_rejected_fragment_gets_no_percentile_at_all() {
let cores: FxHashSet<FragmentId> = FxHashSet::default();
let good = fid("good.rs", 1);
let rejected = fid("garbage.rs", 1);
let rel = scores(&[(rejected.clone(), 0.9), (good.clone(), 0.1)]);
let admitted: FxHashSet<FragmentId> = std::iter::once(good.clone()).collect();
let (pct, top) = percentiles(&rel, &admitted, &cores, 5);
assert!(
!pct.contains_key(&rejected),
"the component's veto was lost"
);
assert!(
!top.contains(&rejected),
"a rejected fragment cannot sit in the component's top-k"
);
assert!(pct.contains_key(&good));
}
#[test]
fn quantile_map_restores_the_reference_distribution_in_fused_order() {
let reference = vec![0.01, 0.02, 0.05, 0.4, 1.9];
let a = fid("a.rs", 1);
let b = fid("b.rs", 1);
let c = fid("c.rs", 1);
let fused = scores(&[(a.clone(), 0.9), (b.clone(), 0.5), (c.clone(), 0.1)]);
let mapped = quantile_map_to(&reference, &fused);
assert_eq!(mapped[&a], 1.9, "the fused top must take the reference max");
assert_eq!(
mapped[&c], 0.01,
"the fused bottom must take the reference min"
);
assert!(
mapped[&a] > mapped[&b] && mapped[&b] > mapped[&c],
"the fused order was not preserved"
);
}
#[test]
fn quantile_map_over_the_same_population_is_the_identity_on_values() {
let ids: Vec<FragmentId> = (0..5).map(|i| fid("f.rs", i + 1)).collect();
let ego_scores = [0.02, 0.07, 0.11, 0.55, 0.9];
let mut reference: Vec<f64> = ego_scores.to_vec();
reference.sort_by(f64::total_cmp);
let fused = scores(
&ids.iter()
.zip([0.2, 0.4, 0.6, 0.8, 1.0])
.map(|(id, p)| (id.clone(), p))
.collect::<Vec<_>>(),
);
let mapped = quantile_map_to(&reference, &fused);
for (id, expected) in ids.iter().zip(ego_scores) {
assert_eq!(
mapped[id], expected,
"same-population quantile map must reproduce the component's own values"
);
}
}
#[test]
fn a_strategy_is_created_for_every_scoring_kind() {
for mode in [
crate::mode::ScoringMode::Ego,
crate::mode::ScoringMode::Ppr,
crate::mode::ScoringMode::Bm25,
crate::mode::ScoringMode::Rrf,
] {
let config = PipelineConfig::from_mode(mode);
let strategy = create_scoring_strategy(&config);
let empty: Vec<Fragment> = Vec::new();
let result =
strategy.score_and_filter(&empty, &FxHashSet::default(), &[], None, None, None);
assert!(
result.filtered_fragments.is_empty(),
"{mode:?} invented fragments from an empty universe"
);
}
let _ = FragmentKind::Function;
}
}