use super::*;
fn ids(items: &[&str]) -> Vec<String> {
items.iter().map(|s| (*s).to_string()).collect()
}
#[test]
fn empty_samples_yield_zero_metrics() {
let report = score_recall(&[], 5);
assert_eq!(report.k, 5);
assert_eq!(report.hit_rate(), 0.0);
assert_eq!(report.mean_reciprocal_rank(), 0.0);
assert_eq!(report.mean_precision_at_k(), 0.0);
assert_eq!(report.mean_recall_at_k(), 0.0);
assert_eq!(report.mean_ndcg_at_k(), 0.0);
}
#[test]
fn perfect_recall_at_rank_one() {
let samples = vec![RecallSample::new(
"q",
ids(&["a", "b"]),
ids(&["a", "b", "c", "d", "e"]),
)];
let report = score_recall(&samples, 5);
assert_eq!(report.per_query[0].first_hit_rank, Some(1));
assert!((report.hit_rate() - 1.0).abs() < 1e-9);
assert!((report.mean_reciprocal_rank() - 1.0).abs() < 1e-9);
assert!((report.mean_precision_at_k() - 0.4).abs() < 1e-9);
assert!((report.mean_recall_at_k() - 1.0).abs() < 1e-9);
assert!((report.mean_ndcg_at_k() - 1.0).abs() < 1e-9);
}
#[test]
fn no_hits_anywhere() {
let samples = vec![RecallSample::new("q", ids(&["a"]), ids(&["x", "y", "z"]))];
let report = score_recall(&samples, 5);
assert_eq!(report.per_query[0].first_hit_rank, None);
assert_eq!(report.hit_rate(), 0.0);
assert_eq!(report.mean_reciprocal_rank(), 0.0);
assert_eq!(report.mean_precision_at_k(), 0.0);
assert_eq!(report.mean_recall_at_k(), 0.0);
assert_eq!(report.mean_ndcg_at_k(), 0.0);
}
#[test]
fn hit_at_rank_three_lowers_mrr() {
let samples = vec![RecallSample::new(
"q",
ids(&["a"]),
ids(&["x", "y", "a", "z"]),
)];
let report = score_recall(&samples, 5);
assert_eq!(report.per_query[0].first_hit_rank, Some(3));
let expected_mrr = 1.0 / 3.0;
assert!((report.mean_reciprocal_rank() - expected_mrr).abs() < 1e-9);
}
#[test]
fn truncates_window_to_k() {
let samples = vec![RecallSample::new(
"q",
ids(&["target"]),
ids(&["a", "b", "c", "d", "e", "target"]),
)];
let report = score_recall(&samples, 5);
assert_eq!(report.per_query[0].first_hit_rank, None);
}
#[test]
fn hit_at_boundary_rank_k() {
let samples = vec![RecallSample::new(
"q",
ids(&["target"]),
ids(&["a", "b", "c", "d", "target"]),
)];
let report = score_recall(&samples, 5);
assert_eq!(report.per_query[0].first_hit_rank, Some(5));
assert!((report.hit_rate() - 1.0).abs() < 1e-9);
}
#[test]
fn k_zero_yields_zero_metrics() {
let samples = vec![RecallSample::new("q", ids(&["a"]), ids(&["a", "b", "c"]))];
let report = score_recall(&samples, 0);
assert_eq!(report.k, 0);
assert_eq!(report.per_query[0].first_hit_rank, None);
assert_eq!(report.hit_rate(), 0.0);
assert_eq!(report.mean_reciprocal_rank(), 0.0);
assert_eq!(report.mean_precision_at_k(), 0.0);
assert_eq!(report.mean_recall_at_k(), 0.0);
assert_eq!(report.mean_ndcg_at_k(), 0.0);
}
#[test]
fn ndcg_discounts_lower_ranks() {
let rank1 = vec![RecallSample::new("q1", ids(&["t"]), ids(&["t", "x", "y"]))];
let rank3 = vec![RecallSample::new("q3", ids(&["t"]), ids(&["x", "y", "t"]))];
let r1 = score_recall(&rank1, 5);
let r3 = score_recall(&rank3, 5);
assert!(r1.mean_ndcg_at_k() > r3.mean_ndcg_at_k());
assert!((r1.mean_ndcg_at_k() - 1.0).abs() < 1e-9);
}
#[test]
fn precision_recall_when_expected_is_subset() {
let samples = vec![RecallSample::new(
"q",
ids(&["a", "b"]),
ids(&["a", "x", "y", "z", "w"]),
)];
let report = score_recall(&samples, 5);
assert!((report.mean_recall_at_k() - 0.5).abs() < 1e-9);
assert!((report.mean_precision_at_k() - 0.2).abs() < 1e-9);
}
#[test]
fn empty_expected_excluded_from_recall_and_ndcg_mean() {
let samples = vec![
RecallSample::new("q-no-expected", ids(&[]), ids(&["x", "y", "z"])),
RecallSample::new("q-perfect", ids(&["a"]), ids(&["a", "b", "c"])),
];
let report = score_recall(&samples, 5);
assert!((report.hit_rate() - 0.5).abs() < 1e-9);
assert!((report.mean_recall_at_k() - 1.0).abs() < 1e-9);
assert!((report.mean_ndcg_at_k() - 1.0).abs() < 1e-9);
let no_expected = &report.per_query[0];
assert_eq!(no_expected.expected_count, 0);
assert_eq!(no_expected.recall_at_k, 0.0);
assert_eq!(no_expected.ndcg_at_k, 0.0);
}
#[test]
fn multi_query_averaging() {
let samples = vec![
RecallSample::new("hit", ids(&["a"]), ids(&["a", "x", "y"])),
RecallSample::new("miss", ids(&["b"]), ids(&["x", "y", "z"])),
];
let report = score_recall(&samples, 5);
assert!((report.hit_rate() - 0.5).abs() < 1e-9);
assert!((report.mean_reciprocal_rank() - 0.5).abs() < 1e-9);
}
#[test]
fn report_serialises_with_metric_values() {
let samples = vec![RecallSample::new("q", ids(&["a"]), ids(&["a", "b", "c"]))];
let report = score_recall(&samples, 5);
let json = serde_json::to_string(&report).expect("serialisable");
assert!(json.contains("\"k\":5"));
assert!(json.contains("\"first_hit_rank\":1"));
assert!(json.contains("\"precision_at_k\":0.2"));
assert!(json.contains("\"recall_at_k\":1.0"));
assert!(json.contains("\"ndcg_at_k\":1.0"));
assert!(json.contains("\"expected_count\":1"));
}