use crate::tests_stat::Alternative;
#[must_use]
pub fn u_distribution(n1: usize, n2: usize) -> Vec<f64> {
let max_u = n1 * n2;
let mut poly = vec![0.0_f64; max_u + 1];
if let Some(first) = poly.first_mut() {
*first = 1.0;
}
for i in 1..=n1 {
let shift = n2 + i;
for u in (0..=max_u).rev() {
let sub = u.checked_sub(shift).and_then(|k| poly.get(k)).copied();
if let (Some(s), Some(slot)) = (sub, poly.get_mut(u)) {
*slot -= s;
}
}
for u in i..=max_u {
let prev = poly.get(u - i).copied().unwrap_or(0.0);
if let Some(slot) = poly.get_mut(u) {
*slot += prev;
}
}
}
poly
}
#[must_use]
pub fn p_value(u1: f64, n1: usize, n2: usize, alternative: Alternative) -> f64 {
let dist = u_distribution(n1, n2);
let total: f64 = dist.iter().sum();
let u_idx = u1.round();
let le: f64 = dist
.iter()
.enumerate()
.filter(|(u, _)| f64_from_index(*u) <= u_idx)
.map(|(_, &c)| c)
.sum();
let ge: f64 = dist
.iter()
.enumerate()
.filter(|(u, _)| f64_from_index(*u) >= u_idx)
.map(|(_, &c)| c)
.sum();
let p_less = le / total;
let p_greater = ge / total;
match alternative {
Alternative::Less => p_less.clamp(0.0, 1.0),
Alternative::Greater => p_greater.clamp(0.0, 1.0),
Alternative::TwoSided => (2.0 * p_less.min(p_greater)).clamp(0.0, 1.0),
}
}
fn f64_from_index(i: usize) -> f64 {
u32::try_from(i).map_or(f64::INFINITY, f64::from)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn distribution_sums_to_choose() {
let dist = u_distribution(3, 4);
let total: f64 = dist.iter().sum();
assert!((total - 35.0).abs() < 1e-9, "total was {total}");
assert_eq!(dist.len(), 13, "len was {}", dist.len());
let expected = [
1.0, 1.0, 2.0, 3.0, 4.0, 4.0, 5.0, 4.0, 4.0, 3.0, 2.0, 1.0, 1.0,
];
for (u, (&got, &want)) in dist.iter().zip(&expected).enumerate() {
assert!((got - want).abs() < 1e-9, "coeff {u}: {got} vs {want}");
}
}
#[test]
fn tail_at_zero_matches_reference() {
let less = p_value(0.0, 3, 4, Alternative::Less);
let greater = p_value(0.0, 3, 4, Alternative::Greater);
let two = p_value(0.0, 3, 4, Alternative::TwoSided);
assert!((less - 1.0 / 35.0).abs() < 1e-12, "less was {less}");
assert!((greater - 1.0).abs() < 1e-12, "greater was {greater}");
assert!((two - 2.0 / 35.0).abs() < 1e-12, "two was {two}");
}
}