use super::segment::Segment;
use crate::core::embeddings::cosine_similarity;
use crate::core::surprise::ScoringCtx;
const REDUNDANCY_COSINE: f64 = 0.92;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RankMode {
Centrality,
Query,
}
fn quantize(score: f32) -> i64 {
(f64::from(score) * 10_000.0).round() as i64
}
pub(super) fn centrality_scores(embs: &[Vec<f32>]) -> Vec<f32> {
let n = embs.len();
if n <= 1 {
return vec![0.0; n];
}
let mut scores = vec![0.0f32; n];
for i in 0..n {
let mut sum = 0.0f32;
for j in 0..n {
if i != j {
sum += cosine_similarity(&embs[i], &embs[j]);
}
}
scores[i] = sum / (n as f32 - 1.0);
}
scores
}
pub(super) fn query_scores(embs: &[Vec<f32>], anchor: &[f32]) -> Vec<f32> {
embs.iter().map(|e| cosine_similarity(e, anchor)).collect()
}
fn cost_of(seg: &Segment) -> usize {
seg.text.len() + 1
}
pub(super) fn select(
segs: &[Segment],
embs: &[Vec<f32>],
mode: RankMode,
anchor: Option<&[f32]>,
budget_chars: usize,
) -> Vec<usize> {
debug_assert_eq!(segs.len(), embs.len());
let n = segs.len();
let scores = match mode {
RankMode::Centrality => centrality_scores(embs),
RankMode::Query => {
let Some(a) = anchor else {
return Vec::new();
};
query_scores(embs, a)
}
};
let mut kept: Vec<usize> = Vec::new();
let mut ctx = ScoringCtx::new();
let mut used = 0usize;
for (i, seg) in segs.iter().enumerate() {
if seg.protected {
kept.push(i);
used += cost_of(seg);
ctx.push_kept(embs[i].clone());
}
}
let mut ranked: Vec<usize> = (0..n).filter(|&i| !segs[i].protected).collect();
ranked.sort_by(|&a, &b| {
quantize(scores[b])
.cmp(&quantize(scores[a]))
.then(a.cmp(&b))
});
for i in ranked {
let cost = cost_of(&segs[i]);
if used + cost > budget_chars && !kept.is_empty() {
continue;
}
if ctx.max_cosine(&embs[i]) >= REDUNDANCY_COSINE {
continue;
}
kept.push(i);
used += cost;
ctx.push_kept(embs[i].clone());
}
kept.sort_unstable();
kept
}
#[cfg(test)]
mod tests {
use super::*;
fn unit(v: [f32; 3]) -> Vec<f32> {
let norm = (v[0] * v[0] + v[1] * v[1] + v[2] * v[2]).sqrt();
vec![v[0] / norm, v[1] / norm, v[2] / norm]
}
fn seg(idx: usize, text: &str, protected: bool) -> Segment {
Segment {
idx,
para: 0,
text: text.to_string(),
protected,
}
}
#[test]
fn centrality_ranks_the_outlier_lowest() {
let embs = vec![
unit([1.0, 0.0, 0.0]),
unit([0.98, 0.02, 0.0]),
unit([0.97, 0.0, 0.03]),
unit([0.0, 0.0, 1.0]),
];
let scores = centrality_scores(&embs);
let outlier = scores
.iter()
.enumerate()
.min_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.unwrap()
.0;
assert_eq!(outlier, 3, "the orthogonal vector is least central");
}
#[test]
fn query_mode_ranks_by_anchor() {
let embs = vec![unit([1.0, 0.0, 0.0]), unit([0.0, 1.0, 0.0])];
let anchor = unit([0.9, 0.1, 0.0]);
let scores = query_scores(&embs, &anchor);
assert!(scores[0] > scores[1], "segment aligned with anchor wins");
}
#[test]
fn select_keeps_protected_and_central_drops_peripheral() {
let segs = vec![
seg(0, "PROTECTED", true),
seg(1, "central a", false),
seg(2, "central b", false),
seg(3, "peripheral", false),
];
let embs = vec![
unit([1.0, 0.0, 0.0]),
unit([0.6, 0.8, 0.0]),
unit([0.6, -0.8, 0.0]),
unit([0.0, 0.0, 1.0]),
];
let budget = "PROTECTED".len() + "central a".len() + "central b".len() + 3;
let kept = select(&segs, &embs, RankMode::Centrality, None, budget);
assert!(kept.contains(&0), "protected always kept");
assert!(kept.contains(&1) && kept.contains(&2), "central kept");
assert!(!kept.contains(&3), "peripheral dropped under budget");
let mut sorted = kept.clone();
sorted.sort_unstable();
assert_eq!(kept, sorted);
}
#[test]
fn mmr_drops_near_duplicate_segment() {
let segs = vec![
seg(0, "unique sentence one", false),
seg(1, "duplicate", false),
seg(2, "duplicate copy", false),
];
let embs = vec![
unit([1.0, 0.0, 0.0]),
unit([0.0, 1.0, 0.0]),
unit([0.0, 1.0, 0.0]),
];
let kept = select(&segs, &embs, RankMode::Centrality, None, 10_000);
let dups = [1usize, 2].iter().filter(|i| kept.contains(i)).count();
assert_eq!(dups, 1, "MMR keeps only one of the duplicates");
}
#[test]
fn select_is_deterministic() {
let segs = vec![
seg(0, "alpha", false),
seg(1, "beta", false),
seg(2, "gamma", false),
];
let embs = vec![
unit([1.0, 0.1, 0.0]),
unit([0.9, 0.2, 0.0]),
unit([0.0, 0.0, 1.0]),
];
let a = select(&segs, &embs, RankMode::Centrality, None, 12);
let b = select(&segs, &embs, RankMode::Centrality, None, 12);
assert_eq!(a, b);
}
}