use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EvalMemory {
pub key: String,
pub text: String,
#[serde(default)]
pub valid_to: Option<String>,
#[serde(default)]
pub superseded_by_key: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum CaseKind {
#[default]
Recall,
KnowledgeUpdate,
Contradiction,
Temporal,
MultiSession,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EvalCase {
pub query: String,
pub relevant: Vec<String>,
#[serde(default)]
pub family: String,
#[serde(default)]
pub kind: CaseKind,
#[serde(default)]
pub stale: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EvalFixture {
pub memories: Vec<EvalMemory>,
pub cases: Vec<EvalCase>,
}
pub fn recall_at_k(ranked: &[String], relevant: &[String], k: usize) -> f64 {
if relevant.is_empty() {
return 0.0;
}
if k == 0 || ranked.is_empty() {
return 0.0;
}
let window = &ranked[..k.min(ranked.len())];
let relevant: std::collections::HashSet<_> = relevant.iter().collect();
let found = relevant
.iter()
.filter(|r| window.iter().any(|w| w == **r))
.count();
found as f64 / relevant.len() as f64
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct EvaluationMetrics {
pub positive_count: usize,
pub negative_count: usize,
pub recall_at_2: Option<f64>,
pub recall_at_4: Option<f64>,
pub hit_at_2: Option<f64>,
pub hit_at_4: Option<f64>,
pub mrr: Option<f64>,
pub negative_accuracy: Option<f64>,
pub false_injection_rate: Option<f64>,
pub quality: Option<f64>,
pub mean_final_bound: f64,
}
pub fn summarize_deliveries(
cases: &[&EvalCase],
ranked: &[Vec<String>],
final_bounds: &[u32],
) -> Result<EvaluationMetrics, String> {
if cases.len() != ranked.len() || cases.len() != final_bounds.len() {
return Err("every case requires one delivery and final cost measurement".into());
}
let positives: Vec<_> = cases
.iter()
.zip(ranked)
.filter(|(c, _)| !c.relevant.is_empty())
.collect();
let negatives: Vec<_> = cases
.iter()
.zip(ranked)
.filter(|(c, _)| c.relevant.is_empty())
.collect();
let avg = |values: Vec<f64>| (!values.is_empty()).then(|| mean(&values));
let recall = |k| {
avg(positives
.iter()
.map(|(c, r)| recall_at_k(r, &c.relevant, k))
.collect())
};
let hit = |k| {
avg(positives
.iter()
.map(|(c, r)| f64::from(recall_at_k(r, &c.relevant, k) > 0.0))
.collect())
};
let mrr = avg(positives.iter().map(|(c, r)| mrr(r, &c.relevant)).collect());
let negative_accuracy = avg(negatives
.iter()
.map(|(_, r)| f64::from(r.is_empty()))
.collect());
let quality = avg(mrr.into_iter().chain(negative_accuracy).collect());
Ok(EvaluationMetrics {
positive_count: positives.len(),
negative_count: negatives.len(),
recall_at_2: recall(2),
recall_at_4: recall(4),
hit_at_2: hit(2),
hit_at_4: hit(4),
mrr,
negative_accuracy,
false_injection_rate: negative_accuracy.map(|a| 1.0 - a),
quality,
mean_final_bound: mean(
&final_bounds
.iter()
.map(|b| f64::from(*b))
.collect::<Vec<_>>(),
),
})
}
pub fn mrr(ranked: &[String], relevant: &[String]) -> f64 {
if relevant.is_empty() || ranked.is_empty() {
return 0.0;
}
for (idx, key) in ranked.iter().enumerate() {
if relevant.iter().any(|r| r == key) {
return 1.0 / (idx as f64 + 1.0);
}
}
0.0
}
pub fn mean(values: &[f64]) -> f64 {
if values.is_empty() {
return 0.0;
}
values.iter().sum::<f64>() / values.len() as f64
}
pub fn stale_hit_rate(ranked: &[String], stale: &[String], k: usize) -> f64 {
if stale.is_empty() || k == 0 || ranked.is_empty() {
return 0.0;
}
let window = &ranked[..k.min(ranked.len())];
if stale.iter().any(|s| window.iter().any(|w| w == s)) {
1.0
} else {
0.0
}
}
pub fn resolution_correct(ranked: &[String], relevant: &[String], stale: &[String]) -> bool {
if relevant.is_empty() {
return false;
}
let best_relevant = ranked
.iter()
.enumerate()
.find(|(_, k)| relevant.iter().any(|r| r == *k))
.map(|(i, _)| i);
let best_relevant = match best_relevant {
Some(pos) => pos,
None => return false, };
let best_stale = ranked
.iter()
.enumerate()
.find(|(_, k)| stale.iter().any(|s| s == *k))
.map(|(i, _)| i);
match best_stale {
None => true, Some(stale_pos) => best_relevant < stale_pos,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn s(v: &[&str]) -> Vec<String> {
v.iter().map(|x| x.to_string()).collect()
}
#[test]
fn recall_at_k_negative_has_no_vacuous_quality_credit() {
assert_eq!(recall_at_k(&s(&["a", "b"]), &[], 4), 0.0);
assert_eq!(recall_at_k(&[], &[], 4), 0.0);
}
#[test]
fn delivered_metrics_separate_fraction_hit_and_known_negative_accuracy() {
let positive: EvalCase =
serde_json::from_value(serde_json::json!({"query":"two facts","relevant":["a","b"]}))
.unwrap();
let negative: EvalCase =
serde_json::from_value(serde_json::json!({"query":"unanswerable","relevant":[]}))
.unwrap();
let metrics =
summarize_deliveries(&[&positive, &negative], &[s(&["a"]), vec![]], &[512, 256])
.unwrap();
assert_eq!(metrics.recall_at_2, Some(0.5));
assert_eq!(metrics.hit_at_2, Some(1.0));
assert_eq!(metrics.mrr, Some(1.0));
assert_eq!(metrics.negative_accuracy, Some(1.0));
assert_eq!(metrics.quality, Some(1.0));
assert_eq!(metrics.mean_final_bound, 384.0);
let injecting = summarize_deliveries(
&[&positive, &negative],
&[s(&["a"]), s(&["junk"])],
&[512, 512],
)
.unwrap();
assert_eq!(injecting.quality, Some(0.5));
let only_negative = summarize_deliveries(&[&negative], &[vec![]], &[256]).unwrap();
assert_eq!(only_negative.mrr, None);
assert_eq!(only_negative.recall_at_2, None);
assert!(summarize_deliveries(&[&positive], &[], &[]).is_err());
}
#[test]
fn recall_at_k_zero_k_is_zero() {
assert_eq!(recall_at_k(&s(&["a", "b"]), &s(&["a"]), 0), 0.0);
}
#[test]
fn recall_at_k_k_larger_than_ranked_uses_full_list() {
let ranked = s(&["a", "b"]);
let relevant = s(&["a", "b", "c"]);
let r = recall_at_k(&ranked, &relevant, 100);
assert!((r - 2.0 / 3.0).abs() < 1e-9);
}
#[test]
fn recall_at_k_exact_hits() {
let ranked = s(&["a", "b", "c", "d"]);
let relevant = s(&["b", "d"]);
assert!((recall_at_k(&ranked, &relevant, 2) - 0.5).abs() < 1e-9);
assert_eq!(recall_at_k(&ranked, &relevant, 4), 1.0);
}
#[test]
fn recall_at_k_duplicates_in_ranked_count_once() {
let ranked = s(&["a", "a", "b"]);
let relevant = s(&["a", "b"]);
assert_eq!(recall_at_k(&ranked, &relevant, 3), 1.0);
assert!((recall_at_k(&ranked, &relevant, 1) - 0.5).abs() < 1e-9);
}
#[test]
fn recall_at_k_no_hits_is_zero() {
let ranked = s(&["x", "y", "z"]);
let relevant = s(&["a", "b"]);
assert_eq!(recall_at_k(&ranked, &relevant, 5), 0.0);
}
#[test]
fn mrr_first_position_is_one() {
let ranked = s(&["a", "b", "c"]);
let relevant = s(&["a"]);
assert_eq!(mrr(&ranked, &relevant), 1.0);
}
#[test]
fn mrr_second_position_is_half() {
let ranked = s(&["x", "a", "b"]);
let relevant = s(&["a"]);
assert!((mrr(&ranked, &relevant) - 0.5).abs() < 1e-9);
}
#[test]
fn mrr_third_position_is_one_third() {
let ranked = s(&["x", "y", "a"]);
let relevant = s(&["a"]);
assert!((mrr(&ranked, &relevant) - 1.0 / 3.0).abs() < 1e-9);
}
#[test]
fn mrr_absent_is_zero() {
let ranked = s(&["x", "y", "z"]);
let relevant = s(&["a"]);
assert_eq!(mrr(&ranked, &relevant), 0.0);
}
#[test]
fn mrr_empty_relevant_is_zero() {
let ranked = s(&["a", "b"]);
assert_eq!(mrr(&ranked, &[]), 0.0);
}
#[test]
fn mrr_empty_ranked_is_zero() {
assert_eq!(mrr(&[], &s(&["a"]),), 0.0);
}
#[test]
fn mrr_uses_first_hit_when_multiple_relevant() {
let ranked = s(&["x", "b", "a"]);
let relevant = s(&["a", "b"]);
assert!((mrr(&ranked, &relevant) - 0.5).abs() < 1e-9);
}
#[test]
fn mean_empty_is_zero() {
assert_eq!(mean(&[]), 0.0);
}
#[test]
fn mean_single() {
assert!((mean(&[0.75]) - 0.75).abs() < 1e-9);
}
#[test]
fn mean_normal() {
let v = [0.0, 0.5, 1.0];
assert!((mean(&v) - 0.5).abs() < 1e-9);
}
#[test]
fn mean_all_ones() {
assert!((mean(&[1.0, 1.0, 1.0]) - 1.0).abs() < 1e-9);
}
#[test]
fn stale_hit_rate_no_stale_is_zero() {
assert_eq!(stale_hit_rate(&s(&["a", "b", "c"]), &[], 4), 0.0);
assert_eq!(stale_hit_rate(&[], &[], 4), 0.0);
}
#[test]
fn stale_hit_rate_stale_in_top_k_is_one() {
let ranked = s(&["a", "b", "c", "d"]);
let stale = s(&["b"]);
assert_eq!(stale_hit_rate(&ranked, &stale, 4), 1.0);
}
#[test]
fn stale_hit_rate_stale_beyond_k_is_zero() {
let ranked = s(&["a", "b", "c", "d"]);
let stale = s(&["d"]);
assert_eq!(stale_hit_rate(&ranked, &stale, 2), 0.0);
}
#[test]
fn stale_hit_rate_stale_absent_is_zero() {
let ranked = s(&["a", "b", "c"]);
let stale = s(&["z"]);
assert_eq!(stale_hit_rate(&ranked, &stale, 4), 0.0);
}
#[test]
fn resolution_correct_relevant_above_stale_is_true() {
let ranked = s(&["new", "x", "old"]);
assert!(resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
}
#[test]
fn resolution_correct_stale_above_relevant_is_false() {
let ranked = s(&["old", "x", "new"]);
assert!(!resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
}
#[test]
fn resolution_correct_stale_absent_is_true() {
let ranked = s(&["new", "x", "y"]);
assert!(resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
}
#[test]
fn resolution_correct_relevant_absent_is_false() {
let ranked = s(&["old", "x", "y"]);
assert!(!resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
}
#[test]
fn resolution_correct_empty_relevant_is_false() {
let ranked = s(&["new", "old"]);
assert!(!resolution_correct(&ranked, &[], &s(&["old"])));
}
}