use crate::model::Scored;
#[derive(Debug, Clone, Copy, Default)]
pub struct QueryMetrics {
pub recall: f32,
pub mrr: f32,
pub ndcg: f32,
}
pub fn is_relevant(text: &str, expected: &[String]) -> bool {
let hay = text.to_ascii_lowercase();
expected
.iter()
.any(|e| hay.contains(&e.to_ascii_lowercase()))
}
pub fn evaluate(hits: &[Scored], expected: &[String], k: usize) -> QueryMetrics {
if expected.is_empty() {
return QueryMetrics::default();
}
let top = &hits[..hits.len().min(k)];
let matched = expected
.iter()
.filter(|e| {
let el = e.to_ascii_lowercase();
top.iter()
.any(|h| h.chunk.text.to_ascii_lowercase().contains(&el))
})
.count();
let recall = matched as f32 / expected.len() as f32;
let mut mrr = 0.0;
for (i, h) in top.iter().enumerate() {
if is_relevant(&h.chunk.text, expected) {
mrr = 1.0 / (i as f32 + 1.0);
break;
}
}
let mut dcg = 0.0;
let mut relevant_found = 0usize;
for (i, h) in top.iter().enumerate() {
if is_relevant(&h.chunk.text, expected) {
dcg += 1.0 / ((i as f32 + 2.0).log2());
relevant_found += 1;
}
}
let idcg: f32 = (0..relevant_found)
.map(|i| 1.0 / ((i as f32 + 2.0).log2()))
.sum();
let ndcg = if idcg > 0.0 { dcg / idcg } else { 0.0 };
QueryMetrics { recall, mrr, ndcg }
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::Chunk;
fn hit(text: &str) -> Scored {
Scored::new(Chunk::new("d", 0, text, 0), 1.0)
}
#[test]
fn perfect_ranking_scores_one() {
let hits = vec![hit("the answer is chunking"), hit("unrelated")];
let m = evaluate(&hits, &["chunking".into()], 5);
assert!((m.recall - 1.0).abs() < 1e-6);
assert!((m.mrr - 1.0).abs() < 1e-6);
assert!((m.ndcg - 1.0).abs() < 1e-6);
}
#[test]
fn relevant_lower_down_lowers_mrr() {
let hits = vec![hit("noise"), hit("noise"), hit("real chunking answer")];
let m = evaluate(&hits, &["chunking".into()], 5);
assert!((m.mrr - (1.0 / 3.0)).abs() < 1e-6);
assert!((m.recall - 1.0).abs() < 1e-6); }
#[test]
fn no_match_scores_zero() {
let hits = vec![hit("a"), hit("b")];
let m = evaluate(&hits, &["missing".into()], 5);
assert_eq!(m.recall, 0.0);
assert_eq!(m.mrr, 0.0);
assert_eq!(m.ndcg, 0.0);
}
}