use serde::{Deserialize, Serialize};
use uuid::Uuid;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RetrievalLabel {
Decisive,
Supporting,
Background,
Irrelevant,
AdjacentWrong,
}
impl RetrievalLabel {
pub fn gain(self) -> f64 {
match self {
Self::Decisive => 3.0,
Self::Supporting => 2.0,
Self::Background => 0.5,
Self::Irrelevant => 0.0,
Self::AdjacentWrong => -2.0,
}
}
pub fn is_relevant(self) -> bool {
matches!(self, Self::Decisive | Self::Supporting)
}
pub fn is_distractor(self) -> bool {
matches!(self, Self::AdjacentWrong)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LabeledResult {
pub section_id: Uuid,
pub score: f64,
pub label: RetrievalLabel,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RetrievalMetrics {
pub recall_at_k: Vec<(usize, f64)>,
pub ndcg_at_10: f64,
pub precision_at_5: f64,
pub precision_at_10: f64,
pub distractor_at_10: f64,
pub net_evidence_at_10: f64,
pub mrr: f64,
pub flip_ratio: Option<f64>,
}
pub fn recall_at_k(results: &[LabeledResult], k: usize) -> f64 {
let total_relevant: usize = results.iter().filter(|r| r.label.is_relevant()).count();
if total_relevant == 0 {
return 1.0;
}
let k = k.min(results.len());
let found: usize = results[..k]
.iter()
.filter(|r| r.label.is_relevant())
.count();
found as f64 / total_relevant as f64
}
pub fn precision_at_k(results: &[LabeledResult], k: usize) -> f64 {
if k == 0 || results.is_empty() {
return 0.0;
}
let k = k.min(results.len());
let relevant: usize = results[..k]
.iter()
.filter(|r| r.label.is_relevant())
.count();
relevant as f64 / k as f64
}
pub fn distractor_at_k(results: &[LabeledResult], k: usize) -> f64 {
if k == 0 || results.is_empty() {
return 0.0;
}
let k = k.min(results.len());
let distractors: usize = results[..k]
.iter()
.filter(|r| r.label.is_distractor())
.count();
distractors as f64 / k as f64
}
pub fn net_evidence_at_k(results: &[LabeledResult], k: usize) -> f64 {
if k == 0 || results.is_empty() {
return 0.0;
}
let k = k.min(results.len());
results[..k]
.iter()
.enumerate()
.map(|(i, r)| r.label.gain() / (i as f64 + 2.0).log2())
.sum()
}
pub fn ndcg_at_k(results: &[LabeledResult], k: usize) -> f64 {
if k == 0 || results.is_empty() {
return 0.0;
}
let k = k.min(results.len());
let dcg = results[..k]
.iter()
.enumerate()
.map(|(i, r)| r.label.gain() / (i as f64 + 2.0).log2())
.sum::<f64>();
let mut gains: Vec<f64> = results.iter().map(|r| r.label.gain()).collect();
gains.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
let idcg = gains[..k]
.iter()
.enumerate()
.map(|(i, &g)| g / (i as f64 + 2.0).log2())
.sum::<f64>();
if idcg == 0.0 {
return 1.0;
}
if idcg < 0.0 {
return 0.0;
}
(dcg / idcg).clamp(0.0, 1.0)
}
pub fn mrr(results: &[LabeledResult]) -> f64 {
for (i, r) in results.iter().enumerate() {
if r.label == RetrievalLabel::Decisive {
return 1.0 / (i as f64 + 1.0);
}
}
0.0
}
pub fn compute_all(results: &[LabeledResult]) -> RetrievalMetrics {
let recall_at_k_vals = vec![
(1, recall_at_k(results, 1)),
(3, recall_at_k(results, 3)),
(5, recall_at_k(results, 5)),
(10, recall_at_k(results, 10)),
];
RetrievalMetrics {
recall_at_k: recall_at_k_vals,
ndcg_at_10: ndcg_at_k(results, 10),
precision_at_5: precision_at_k(results, 5),
precision_at_10: precision_at_k(results, 10),
distractor_at_10: distractor_at_k(results, 10),
net_evidence_at_10: net_evidence_at_k(results, 10),
mrr: mrr(results),
flip_ratio: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn uuid(n: u64) -> Uuid {
Uuid::from_u64_pair(0, n)
}
fn make_result(n: u64, label: RetrievalLabel) -> LabeledResult {
LabeledResult {
section_id: uuid(n),
score: 1.0 / (n as f64 + 1.0),
label,
}
}
#[test]
fn label_gain_values() {
assert_eq!(RetrievalLabel::Decisive.gain(), 3.0);
assert_eq!(RetrievalLabel::Supporting.gain(), 2.0);
assert_eq!(RetrievalLabel::Background.gain(), 0.5);
assert_eq!(RetrievalLabel::Irrelevant.gain(), 0.0);
assert_eq!(RetrievalLabel::AdjacentWrong.gain(), -2.0);
}
#[test]
fn label_is_relevant() {
assert!(RetrievalLabel::Decisive.is_relevant());
assert!(RetrievalLabel::Supporting.is_relevant());
assert!(!RetrievalLabel::Background.is_relevant());
assert!(!RetrievalLabel::Irrelevant.is_relevant());
assert!(!RetrievalLabel::AdjacentWrong.is_relevant());
}
#[test]
fn label_is_distractor() {
assert!(RetrievalLabel::AdjacentWrong.is_distractor());
assert!(!RetrievalLabel::Decisive.is_distractor());
assert!(!RetrievalLabel::Irrelevant.is_distractor());
}
#[test]
fn recall_at_k_all_relevant() {
let results: Vec<LabeledResult> = (0..3)
.map(|i| make_result(i, RetrievalLabel::Decisive))
.collect();
assert!((recall_at_k(&results, 3) - 1.0).abs() < 1e-9);
}
#[test]
fn recall_at_k_partial() {
let results = vec![
make_result(0, RetrievalLabel::Decisive),
make_result(1, RetrievalLabel::Decisive),
make_result(2, RetrievalLabel::Irrelevant),
make_result(3, RetrievalLabel::Irrelevant),
];
assert!((recall_at_k(&results, 1) - 0.5).abs() < 1e-9);
assert!((recall_at_k(&results, 2) - 1.0).abs() < 1e-9);
}
#[test]
fn recall_at_k_none_relevant_vacuously_one() {
let results = vec![
make_result(0, RetrievalLabel::Irrelevant),
make_result(1, RetrievalLabel::Background),
];
assert!((recall_at_k(&results, 5) - 1.0).abs() < 1e-9);
}
#[test]
fn recall_at_k_k_exceeds_length() {
let results = vec![make_result(0, RetrievalLabel::Decisive)];
assert!((recall_at_k(&results, 100) - 1.0).abs() < 1e-9);
}
#[test]
fn precision_at_k_perfect() {
let results: Vec<LabeledResult> = (0..5)
.map(|i| make_result(i, RetrievalLabel::Decisive))
.collect();
assert!((precision_at_k(&results, 5) - 1.0).abs() < 1e-9);
}
#[test]
fn precision_at_k_half_relevant() {
let results = vec![
make_result(0, RetrievalLabel::Decisive),
make_result(1, RetrievalLabel::Irrelevant),
make_result(2, RetrievalLabel::Supporting),
make_result(3, RetrievalLabel::Irrelevant),
];
assert!((precision_at_k(&results, 4) - 0.5).abs() < 1e-9);
}
#[test]
fn precision_at_k_zero_when_k_zero() {
let results = vec![make_result(0, RetrievalLabel::Decisive)];
assert_eq!(precision_at_k(&results, 0), 0.0);
}
#[test]
fn precision_at_k_zero_when_empty() {
assert_eq!(precision_at_k(&[], 5), 0.0);
}
#[test]
fn precision_at_k_adjacent_wrong_not_counted() {
let results = vec![
make_result(0, RetrievalLabel::AdjacentWrong),
make_result(1, RetrievalLabel::AdjacentWrong),
];
assert_eq!(precision_at_k(&results, 2), 0.0);
}
#[test]
fn distractor_at_k_all_wrong() {
let results: Vec<LabeledResult> = (0..4)
.map(|i| make_result(i, RetrievalLabel::AdjacentWrong))
.collect();
assert!((distractor_at_k(&results, 4) - 1.0).abs() < 1e-9);
}
#[test]
fn distractor_at_k_none_wrong() {
let results = vec![
make_result(0, RetrievalLabel::Decisive),
make_result(1, RetrievalLabel::Irrelevant),
];
assert_eq!(distractor_at_k(&results, 2), 0.0);
}
#[test]
fn distractor_at_k_mixed() {
let results = vec![
make_result(0, RetrievalLabel::Decisive),
make_result(1, RetrievalLabel::AdjacentWrong),
make_result(2, RetrievalLabel::Irrelevant),
make_result(3, RetrievalLabel::Background),
];
assert!((distractor_at_k(&results, 4) - 0.25).abs() < 1e-9);
}
#[test]
fn distractor_at_k_zero_when_k_zero() {
let results = vec![make_result(0, RetrievalLabel::AdjacentWrong)];
assert_eq!(distractor_at_k(&results, 0), 0.0);
}
#[test]
fn net_evidence_at_k_single_decisive_rank1() {
let results = vec![make_result(0, RetrievalLabel::Decisive)];
assert!((net_evidence_at_k(&results, 1) - 3.0).abs() < 1e-9);
}
#[test]
fn net_evidence_at_k_negative_for_all_wrong() {
let results: Vec<LabeledResult> = (0..3)
.map(|i| make_result(i as u64, RetrievalLabel::AdjacentWrong))
.collect();
let score = net_evidence_at_k(&results, 3);
assert!(
score < 0.0,
"all distractors should produce negative net evidence"
);
}
#[test]
fn net_evidence_at_k_zero_for_empty() {
assert_eq!(net_evidence_at_k(&[], 5), 0.0);
}
#[test]
fn net_evidence_at_k_zero_for_k_zero() {
let results = vec![make_result(0, RetrievalLabel::Decisive)];
assert_eq!(net_evidence_at_k(&results, 0), 0.0);
}
#[test]
fn net_evidence_at_k_mixed_sums_correctly() {
let results = vec![
make_result(0, RetrievalLabel::Decisive),
make_result(1, RetrievalLabel::Supporting),
];
let expected = 3.0 / 2.0_f64.log2() + 2.0 / 3.0_f64.log2();
let actual = net_evidence_at_k(&results, 2);
assert!(
(actual - expected).abs() < 1e-9,
"expected {expected}, got {actual}"
);
}
#[test]
fn ndcg_at_k_perfect_ranking() {
let results = vec![
make_result(0, RetrievalLabel::Decisive),
make_result(1, RetrievalLabel::Supporting),
make_result(2, RetrievalLabel::Irrelevant),
];
let score = ndcg_at_k(&results, 3);
assert!(
(score - 1.0).abs() < 1e-9,
"perfect ranking should yield nDCG=1.0, got {score}"
);
}
#[test]
fn ndcg_at_k_suboptimal_ranking() {
let results = vec![
make_result(0, RetrievalLabel::Irrelevant),
make_result(1, RetrievalLabel::Decisive),
];
let score = ndcg_at_k(&results, 2);
assert!(
score < 1.0 && score > 0.0,
"suboptimal ranking should yield 0 < nDCG < 1.0, got {score}"
);
}
#[test]
fn ndcg_at_k_all_irrelevant_vacuously_one() {
let results = vec![
make_result(0, RetrievalLabel::Irrelevant),
make_result(1, RetrievalLabel::Irrelevant),
];
let score = ndcg_at_k(&results, 2);
assert!((score - 1.0).abs() < 1e-9);
}
#[test]
fn ndcg_at_k_zero_for_zero_k() {
let results = vec![make_result(0, RetrievalLabel::Decisive)];
assert_eq!(ndcg_at_k(&results, 0), 0.0);
}
#[test]
fn ndcg_at_k_clamped_not_above_one() {
let results: Vec<LabeledResult> = (0..10)
.map(|i| make_result(i, RetrievalLabel::Decisive))
.collect();
let score = ndcg_at_k(&results, 10);
assert!(
score <= 1.0 + 1e-12,
"nDCG must not exceed 1.0, got {score}"
);
}
#[test]
fn mrr_decisive_at_rank1() {
let results = vec![
make_result(0, RetrievalLabel::Decisive),
make_result(1, RetrievalLabel::Irrelevant),
];
assert!((mrr(&results) - 1.0).abs() < 1e-9);
}
#[test]
fn mrr_decisive_at_rank3() {
let results = vec![
make_result(0, RetrievalLabel::Irrelevant),
make_result(1, RetrievalLabel::Supporting),
make_result(2, RetrievalLabel::Decisive),
];
assert!((mrr(&results) - 1.0 / 3.0).abs() < 1e-9);
}
#[test]
fn mrr_no_decisive() {
let results = vec![
make_result(0, RetrievalLabel::Supporting),
make_result(1, RetrievalLabel::Irrelevant),
];
assert_eq!(mrr(&results), 0.0);
}
#[test]
fn mrr_empty() {
assert_eq!(mrr(&[]), 0.0);
}
#[test]
fn compute_all_returns_correct_structure() {
let results: Vec<LabeledResult> = (0..10)
.map(|i| {
let label = if i < 3 {
RetrievalLabel::Decisive
} else {
RetrievalLabel::Irrelevant
};
make_result(i, label)
})
.collect();
let metrics = compute_all(&results);
assert_eq!(metrics.recall_at_k.len(), 4);
assert_eq!(metrics.recall_at_k[0].0, 1);
assert_eq!(metrics.recall_at_k[1].0, 3);
assert_eq!(metrics.recall_at_k[2].0, 5);
assert_eq!(metrics.recall_at_k[3].0, 10);
assert!((metrics.recall_at_k[1].1 - 1.0).abs() < 1e-9);
assert!((metrics.mrr - 1.0).abs() < 1e-9);
assert!(metrics.flip_ratio.is_none());
}
#[test]
fn compute_all_distractor_metric() {
let results: Vec<LabeledResult> = (0..10)
.map(|i| {
let label = if i < 5 {
RetrievalLabel::AdjacentWrong
} else {
RetrievalLabel::Irrelevant
};
make_result(i, label)
})
.collect();
let metrics = compute_all(&results);
assert!(
(metrics.distractor_at_10 - 0.5).abs() < 1e-9,
"got {}",
metrics.distractor_at_10
);
assert_eq!(metrics.mrr, 0.0);
}
#[test]
fn label_serde_roundtrip() {
for label in [
RetrievalLabel::Decisive,
RetrievalLabel::Supporting,
RetrievalLabel::Background,
RetrievalLabel::Irrelevant,
RetrievalLabel::AdjacentWrong,
] {
let json = serde_json::to_string(&label).unwrap();
let back: RetrievalLabel = serde_json::from_str(&json).unwrap();
assert_eq!(label, back);
}
}
#[test]
fn metrics_serde_roundtrip() {
let results = vec![make_result(0, RetrievalLabel::Decisive)];
let m = compute_all(&results);
let json = serde_json::to_string(&m).unwrap();
let back: RetrievalMetrics = serde_json::from_str(&json).unwrap();
assert_eq!(back.recall_at_k.len(), 4);
assert!((back.mrr - 1.0).abs() < 1e-9);
}
}