use scc_core::ContextItem;
use std::collections::HashMap;
pub fn select_with_budget(items: &[ContextItem], budget: usize, hard_max: usize) -> Vec<usize> {
let mut selected: Vec<usize> = Vec::new();
let mut spent: usize = 0;
for (i, item) in items.iter().enumerate() {
if item.required {
selected.push(i);
spent = spent.saturating_add(item.token_cost);
}
}
let mut rest: Vec<usize> = (0..items.len()).filter(|i| !items[*i].required).collect();
rest.sort_by(|a, b| {
let va = items[*a].value / items[*a].token_cost.max(1) as f64;
let vb = items[*b].value / items[*b].token_cost.max(1) as f64;
vb.partial_cmp(&va)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.cmp(b))
});
for i in rest {
let cost = items[i].token_cost;
let next = spent.saturating_add(cost);
if next <= budget && next <= hard_max {
selected.push(i);
spent = next;
}
}
selected
}
pub fn select_in_order(items: &[ContextItem], budget: usize, hard_max: usize) -> Vec<usize> {
let mut selected: Vec<usize> = Vec::new();
let mut spent: usize = 0;
for (i, item) in items.iter().enumerate() {
if item.required {
selected.push(i);
spent = spent.saturating_add(item.token_cost);
continue;
}
let next = spent.saturating_add(item.token_cost);
if next <= budget && next <= hard_max {
selected.push(i);
spent = next;
}
}
selected
}
pub fn mmr_diversify(
ranked: &[(String, f64)],
similarity: impl Fn(&str, &str) -> f64,
lambda: f64,
budget: usize,
) -> Vec<String> {
let n = ranked.len();
let budget = budget.min(n);
let lambda = lambda.clamp(0.0, 1.0);
let mut selected: Vec<usize> = Vec::with_capacity(budget);
let mut picked = vec![false; n];
let mut max_sim = vec![0.0_f64; n];
while selected.len() < budget {
if let Some(&last) = selected.last() {
for i in 0..n {
if picked[i] {
continue;
}
let s = similarity(&ranked[i].0, &ranked[last].0);
if s > max_sim[i] {
max_sim[i] = s;
}
}
}
let mut best: Option<(usize, f64)> = None;
for i in 0..n {
if picked[i] {
continue;
}
let v = lambda * ranked[i].1 - (1.0 - lambda) * max_sim[i];
if best.is_none_or(|(_, bv)| v > bv) {
best = Some((i, v));
}
}
match best {
Some((i, _)) => {
selected.push(i);
picked[i] = true;
}
None => break,
}
}
selected.into_iter().map(|i| ranked[i].0.clone()).collect()
}
pub fn mmr_diversify_default(
ranked: &[(String, f64)],
similarity: impl Fn(&str, &str) -> f64,
budget: usize,
) -> Vec<String> {
mmr_diversify(ranked, similarity, 0.5, budget)
}
pub fn enforce_quotas(
ranked: &[(String, f64)],
kind_of: impl Fn(&str) -> &str,
quotas: &[(String, f64)],
available_tokens: usize,
token_cost: impl Fn(&str) -> usize,
) -> Vec<String> {
let mut caps: HashMap<&str, usize> = HashMap::new();
for (kind, frac) in quotas {
let cap = (frac.clamp(0.0, 1.0) * available_tokens as f64).round() as usize;
caps.insert(kind.as_str(), cap);
}
let mut spent: HashMap<&str, usize> = HashMap::new();
let mut out: Vec<String> = Vec::new();
for (id, _) in ranked {
let k = kind_of(id);
let take = match caps.get(k) {
None => true,
Some(&cap) => {
let s = spent.entry(k).or_insert(0);
let next = s.saturating_add(token_cost(id));
if next <= cap {
*s = next;
true
} else {
false
}
}
};
if take {
out.push(id.clone());
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn item(id: &str, value: f64, token_cost: usize, required: bool) -> ContextItem {
ContextItem {
id: id.to_string(),
value,
token_cost,
required,
group: None,
}
}
#[test]
fn budget_selection_drops_low_value_and_keeps_required() {
let items = vec![
item("required", 0.1, 120, true),
item("high", 0.9, 50, false),
item("mid", 0.5, 50, false),
item("low", 0.1, 50, false),
];
let sel = select_with_budget(&items, 170, 204);
assert_eq!(sel, vec![0, 1]);
assert!(sel.contains(&0));
assert!(!sel.contains(&2) && !sel.contains(&3));
let sel = select_with_budget(&items, 50, 204);
assert_eq!(sel, vec![0]);
let sel = select_with_budget(&items, 0, 204);
assert_eq!(sel, vec![0]);
assert!(select_with_budget(&[], 100, 100).is_empty());
let no_req = vec![item("a", 0.1, 100, false), item("b", 0.9, 10, false)];
let sel = select_with_budget(&no_req, 100, 100);
assert_eq!(sel, vec![1]);
let sel = select_with_budget(&items, 50, 60);
assert_eq!(sel, vec![0]);
let sel = select_with_budget(&items, 10, 5);
assert_eq!(sel, vec![0]);
}
#[test]
fn select_in_order_keeps_rank_order() {
let items = vec![
item("required", 0.1, 120, true),
item("high", 0.9, 50, false),
item("mid", 0.5, 50, false),
item("low", 0.1, 50, false),
];
let sel = select_in_order(&items, 170, 204);
assert_eq!(sel, vec![0, 1]);
let sel = select_in_order(&items, 169, 204);
assert_eq!(sel, vec![0]);
let sel = select_in_order(&items, 0, 204);
assert_eq!(sel, vec![0]);
}
#[test]
fn mmr_caps_same_owner_dtos() {
let ranked: Vec<(String, f64)> = (0..5)
.map(|i| (format!("Order.dto{i}"), 0.9 - 0.1 * i as f64))
.collect();
let same_owner = |a: &str, b: &str| -> f64 {
if a.starts_with("Order.") && b.starts_with("Order.") {
1.0
} else {
0.0
}
};
let sel = mmr_diversify_default(&ranked, same_owner, 2);
assert_eq!(sel.len(), 2);
assert_eq!(sel[0], "Order.dto0");
assert_eq!(sel[1], "Order.dto1");
let all = mmr_diversify_default(&ranked, same_owner, 10);
assert_eq!(all.len(), 5);
}
#[test]
fn mmr_keeps_diverse_items() {
let ranked = vec![
("a".to_string(), 0.9),
("b".to_string(), 0.8),
("c".to_string(), 0.7),
];
let distinct = |_: &str, _: &str| 0.0;
let sel = mmr_diversify_default(&ranked, distinct, 3);
assert_eq!(sel, vec!["a".to_string(), "b".to_string(), "c".to_string()]);
let sel = mmr_diversify(&ranked, distinct, 0.0, 3);
assert_eq!(sel.len(), 3);
assert!(mmr_diversify_default(&[], distinct, 5).is_empty());
assert!(mmr_diversify_default(&ranked, distinct, 0).is_empty());
}
#[test]
fn quotas_are_token_aware() {
let mut ranked: Vec<(String, f64)> = Vec::new();
for i in 0..100 {
ranked.push((format!("pub:{i}"), 1.0 - i as f64 / 200.0));
}
for i in 0..20 {
ranked.push((format!("core:{i}"), 1.0 - i as f64 / 200.0));
}
for i in 0..20 {
ranked.push((format!("types:{i}"), 1.0 - i as f64 / 200.0));
}
fn kind_of(id: &str) -> &str {
if id.starts_with("pub:") {
"public"
} else if id.starts_with("core:") {
"core"
} else {
"types"
}
}
fn other_kind(_: &str) -> &str {
"other"
}
let quotas = vec![
("public".to_string(), 0.50),
("core".to_string(), 0.25),
("types".to_string(), 0.25),
];
let cost10 = |_: &str| 10usize;
let sel = enforce_quotas(&ranked, kind_of, "as, 1400, cost10);
let count = |k: &str| sel.iter().filter(|id| kind_of(id.as_str()) == k).count();
assert_eq!(count("public"), 70);
assert_eq!(count("core"), 20);
assert_eq!(count("types"), 20);
assert_eq!(sel[0], "pub:0");
assert_eq!(sel[70], "core:0");
let cost50 = |_: &str| 50usize;
let sel50 = enforce_quotas(&ranked, kind_of, "as, 1400, cost50);
assert_eq!(
sel50.iter().filter(|id| kind_of(id.as_str()) == "public").count(),
14
);
let sel = enforce_quotas(&ranked, kind_of, "as, 1400, cost10);
let core_first = sel.iter().position(|id| kind_of(id.as_str()) == "core").unwrap();
assert_eq!(sel[core_first], "core:0");
let sel_other = enforce_quotas(
&[("x".to_string(), 0.5), ("y".to_string(), 0.5)],
other_kind,
&[],
100,
|_: &str| 10,
);
assert_eq!(sel_other.len(), 2);
}
}