use std::collections::HashSet;
pub const DEFAULT_DIVERSITY_LAMBDA: f64 = 0.5;
pub const DEFAULT_MIN_DISTINCT_FILES: usize = 3;
#[derive(Debug, Clone, PartialEq)]
pub struct GreedyMmrItem {
pub id: String,
pub score: f64,
pub file_path: String,
}
pub fn greedy_mmr(items: &[GreedyMmrItem], k: usize, lambda: f64) -> Vec<GreedyMmrItem> {
if items.is_empty() || k == 0 {
return Vec::new();
}
let k = k.min(items.len());
let lambda = lambda.clamp(0.0, 1.0);
let mut picked: Vec<usize> = Vec::with_capacity(k);
let mut picked_files: HashSet<&str> = HashSet::new();
let first = items
.iter()
.enumerate()
.max_by(|(ia, a), (ib, b)| {
a.score
.partial_cmp(&b.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.id.cmp(&b.id))
.then_with(|| ib.cmp(ia)) })
.map(|(i, _)| i)
.expect("non-empty items");
picked.push(first);
picked_files.insert(items[first].file_path.as_str());
while picked.len() < k {
let mut best: Option<(usize, f64)> = None;
for (i, it) in items.iter().enumerate() {
if picked.contains(&i) {
continue;
}
let sim = if picked_files.contains(it.file_path.as_str()) {
1.0
} else {
0.0
};
let marginal = lambda * it.score - (1.0 - lambda) * sim;
let better = match best {
None => true,
Some((_, best_m)) => {
marginal > best_m
|| (marginal == best_m
&& (it.id < items[best.unwrap().0].id
|| (it.id == items[best.unwrap().0].id && i < best.unwrap().0)))
}
};
if better {
best = Some((i, marginal));
}
}
let (idx, _) = best.expect("remaining items exist while loop runs");
picked_files.insert(items[idx].file_path.as_str());
picked.push(idx);
}
picked.into_iter().map(|i| items[i].clone()).collect()
}
pub fn apply_mmr_diversity(
items: &[GreedyMmrItem],
top_k: usize,
lambda: f64,
min_distinct_files: usize,
) -> Vec<GreedyMmrItem> {
if items.is_empty() {
return Vec::new();
}
let k = if top_k == 0 {
items.len()
} else {
top_k.min(items.len())
};
if lambda <= 0.0 {
return items.iter().take(k).cloned().collect();
}
let picked = greedy_mmr(items, k, lambda);
let distinct: HashSet<&str> = picked.iter().map(|i| i.file_path.as_str()).collect();
if distinct.len() >= min_distinct_files {
picked
} else {
let mut sorted = items.iter().take(k).cloned().collect::<Vec<_>>();
sorted.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.id.cmp(&b.id))
});
sorted
}
}