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 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 1.0;
}
if k == 0 || ranked.is_empty() {
return 0.0;
}
let window = &ranked[..k.min(ranked.len())];
let found = relevant
.iter()
.filter(|r| window.iter().any(|w| w == *r))
.count();
found as f64 / relevant.len() as f64
}
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_empty_relevant_is_one() {
assert_eq!(recall_at_k(&s(&["a", "b"]), &[], 4), 1.0);
assert_eq!(recall_at_k(&[], &[], 4), 1.0);
}
#[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"])));
}
}