Skip to main content

sim_lib_numbers_stats/
implementation.rs

1//! Implementation of the statistics helpers: descriptive statistics,
2//! probability, and fairness metrics over f64 data, with their error type.
3
4use std::{error::Error, fmt};
5
6#[path = "claim.rs"]
7mod claim;
8#[path = "function.rs"]
9mod function;
10#[path = "runtime.rs"]
11mod runtime;
12
13pub use claim::{
14    FairnessClaimValue, FairnessEvidence, StatsClaimEvidence, StatsClaimValue, fairness_claim,
15    fairness_claim_value, stats_result_claim, stats_result_claim_value,
16};
17pub use function::{
18    StatsNumbersLib, stats_claims_symbol, stats_disparate_impact_claim_symbol,
19    stats_entropy_claim_symbol, stats_mean_claim_symbol, stats_variance_claim_symbol,
20};
21
22const FOUR_FIFTHS_THRESHOLD: f64 = 0.8;
23const PROBABILITY_TOLERANCE: f64 = 1.0e-12;
24
25/// Result alias for the statistics helpers, fixing the error to [`StatsError`].
26pub type StatsResult<T> = Result<T, StatsError>;
27
28/// Errors returned by probability, statistics, and fairness helpers.
29///
30/// Each variant carries the `metric` name of the helper that rejected the
31/// input, so the failure can be reported without the caller tracking context.
32#[derive(Clone, Debug, PartialEq)]
33pub enum StatsError {
34    /// A helper requiring at least one value was given an empty slice.
35    EmptyInput {
36        /// Name of the helper that rejected the input.
37        metric: &'static str,
38    },
39    /// Fewer values were supplied than the helper requires.
40    InsufficientInput {
41        /// Name of the helper that rejected the input.
42        metric: &'static str,
43        /// Smallest number of values the helper accepts.
44        minimum: usize,
45        /// Number of values actually supplied.
46        actual: usize,
47    },
48    /// A value was not finite (`NaN` or infinite).
49    NonFinite {
50        /// Name of the helper that rejected the input.
51        metric: &'static str,
52        /// Position of the offending value, if it came from a slice.
53        index: Option<usize>,
54        /// The offending value.
55        value: f64,
56    },
57    /// A probability fell outside the closed range `0.0..=1.0`.
58    ProbabilityOutOfRange {
59        /// Name of the helper that rejected the input.
60        metric: &'static str,
61        /// Position of the offending value, if it came from a slice.
62        index: Option<usize>,
63        /// The offending value.
64        value: f64,
65    },
66    /// A probability vector did not sum to one within tolerance.
67    ProbabilityMass {
68        /// Name of the helper that rejected the input.
69        metric: &'static str,
70        /// The actual sum of the probabilities.
71        sum: f64,
72    },
73    /// A Bayesian update was given zero total evidence.
74    ZeroEvidence {
75        /// Name of the helper that rejected the input.
76        metric: &'static str,
77    },
78    /// A counts pair was given a zero total.
79    ZeroTotal {
80        /// Label of the count whose total was zero.
81        label: &'static str,
82    },
83    /// A fairness ratio was given a zero reference rate to divide by.
84    ZeroReferenceRate {
85        /// Name of the helper that rejected the input.
86        metric: &'static str,
87    },
88}
89
90impl fmt::Display for StatsError {
91    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
92        match self {
93            Self::EmptyInput { metric } => write!(f, "{metric} requires at least one value"),
94            Self::InsufficientInput {
95                metric,
96                minimum,
97                actual,
98            } => write!(
99                f,
100                "{metric} requires at least {minimum} values, got {actual}"
101            ),
102            Self::NonFinite {
103                metric,
104                index,
105                value,
106            } => match index {
107                Some(index) => write!(f, "{metric} value {index} is not finite: {value}"),
108                None => write!(f, "{metric} value is not finite: {value}"),
109            },
110            Self::ProbabilityOutOfRange {
111                metric,
112                index,
113                value,
114            } => match index {
115                Some(index) => write!(
116                    f,
117                    "{metric} probability {index} must be between 0 and 1, got {value}"
118                ),
119                None => write!(
120                    f,
121                    "{metric} probability must be between 0 and 1, got {value}"
122                ),
123            },
124            Self::ProbabilityMass { metric, sum } => {
125                write!(f, "{metric} probabilities must sum to 1, got {sum}")
126            }
127            Self::ZeroEvidence { metric } => write!(f, "{metric} evidence must be nonzero"),
128            Self::ZeroTotal { label } => write!(f, "{label} total must be nonzero"),
129            Self::ZeroReferenceRate { metric } => {
130                write!(f, "{metric} reference rate must be nonzero")
131            }
132        }
133    }
134}
135
136impl Error for StatsError {}
137
138/// Counts for a binary selection or outcome table.
139#[derive(Clone, Copy, Debug, PartialEq)]
140pub struct BinaryOutcomeCounts {
141    /// Number of selected (positive-outcome) cases.
142    pub selected: u64,
143    /// Total number of cases; the denominator of the selection rate.
144    pub total: u64,
145}
146
147impl BinaryOutcomeCounts {
148    /// Builds a count pair and rejects impossible or empty totals.
149    pub fn new(selected: u64, total: u64) -> StatsResult<Self> {
150        if total == 0 {
151            return Err(StatsError::ZeroTotal {
152                label: "outcome counts",
153            });
154        }
155        if selected > total {
156            return Err(StatsError::ProbabilityOutOfRange {
157                metric: "outcome counts",
158                index: None,
159                value: selected as f64 / total as f64,
160            });
161        }
162        Ok(Self { selected, total })
163    }
164
165    /// Returns `selected / total` as an f64 rate.
166    pub fn selection_rate(self) -> f64 {
167        self.selected as f64 / self.total as f64
168    }
169}
170
171/// Disparate-impact summary for two binary-outcome groups.
172#[derive(Clone, Copy, Debug, PartialEq)]
173pub struct DisparateImpact {
174    /// Selection rate of the reference group.
175    pub reference_rate: f64,
176    /// Selection rate of the comparison group.
177    pub comparison_rate: f64,
178    /// Comparison rate divided by reference rate.
179    pub ratio: f64,
180    /// Whether the ratio meets the four-fifths (0.8) threshold.
181    pub passes_four_fifths: bool,
182}
183
184/// Computes `prior * likelihood / evidence` and validates the posterior.
185///
186/// Every argument and the result must be a probability in `0.0..=1.0`, and
187/// `evidence` must be nonzero.
188///
189/// # Examples
190///
191/// ```
192/// use sim_lib_numbers_stats::bayesian_update;
193///
194/// let posterior = bayesian_update(0.2, 0.75, 0.3).unwrap();
195/// assert!((posterior - 0.5).abs() < 1e-12);
196/// ```
197pub fn bayesian_update(prior: f64, likelihood: f64, evidence: f64) -> StatsResult<f64> {
198    validate_probability("bayesian_update", None, prior)?;
199    validate_probability("bayesian_update", None, likelihood)?;
200    validate_probability("bayesian_update", None, evidence)?;
201    if evidence == 0.0 {
202        return Err(StatsError::ZeroEvidence {
203            metric: "bayesian_update",
204        });
205    }
206    let posterior = (prior * likelihood) / evidence;
207    validate_probability("bayesian_update", None, posterior)?;
208    Ok(posterior)
209}
210
211/// Computes a binary-test posterior from prior, true-positive, and false-positive rates.
212pub fn bayesian_update_binary(
213    prior: f64,
214    true_positive_rate: f64,
215    false_positive_rate: f64,
216) -> StatsResult<f64> {
217    validate_probability("bayesian_update_binary", None, prior)?;
218    validate_probability("bayesian_update_binary", None, true_positive_rate)?;
219    validate_probability("bayesian_update_binary", None, false_positive_rate)?;
220    let evidence = prior * true_positive_rate + (1.0 - prior) * false_positive_rate;
221    bayesian_update(prior, true_positive_rate, evidence)
222}
223
224/// Computes Shannon entropy in bits for a probability vector that sums to one.
225pub fn entropy(probabilities: &[f64]) -> StatsResult<f64> {
226    if probabilities.is_empty() {
227        return Err(StatsError::EmptyInput { metric: "entropy" });
228    }
229    let mut sum = 0.0;
230    let mut bits = 0.0;
231    for (index, probability) in probabilities.iter().copied().enumerate() {
232        validate_probability("entropy", Some(index), probability)?;
233        sum += probability;
234        if probability > 0.0 {
235            bits -= probability * probability.log2();
236        }
237    }
238    if (sum - 1.0).abs() > PROBABILITY_TOLERANCE {
239        return Err(StatsError::ProbabilityMass {
240            metric: "entropy",
241            sum,
242        });
243    }
244    Ok(bits)
245}
246
247/// Computes the arithmetic mean of finite values.
248///
249/// Returns [`StatsError::EmptyInput`] for an empty slice and
250/// [`StatsError::NonFinite`] if any value is `NaN` or infinite.
251///
252/// # Examples
253///
254/// ```
255/// use sim_lib_numbers_stats::mean;
256///
257/// assert_eq!(mean(&[2.0, 4.0, 6.0]).unwrap(), 4.0);
258/// ```
259pub fn mean(values: &[f64]) -> StatsResult<f64> {
260    validate_values("mean", values)?;
261    Ok(values.iter().sum::<f64>() / values.len() as f64)
262}
263
264/// Computes population variance.
265///
266/// Alias for [`population_variance`].
267///
268/// # Examples
269///
270/// ```
271/// use sim_lib_numbers_stats::variance;
272///
273/// let values = [2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0];
274/// assert!((variance(&values).unwrap() - 4.0).abs() < 1e-12);
275/// ```
276pub fn variance(values: &[f64]) -> StatsResult<f64> {
277    population_variance(values)
278}
279
280/// Computes population variance with divisor `n`.
281pub fn population_variance(values: &[f64]) -> StatsResult<f64> {
282    validate_values("population_variance", values)?;
283    let mean = mean(values)?;
284    Ok(values
285        .iter()
286        .map(|value| {
287            let delta = value - mean;
288            delta * delta
289        })
290        .sum::<f64>()
291        / values.len() as f64)
292}
293
294/// Computes sample variance with divisor `n - 1`.
295pub fn sample_variance(values: &[f64]) -> StatsResult<f64> {
296    validate_values("sample_variance", values)?;
297    if values.len() < 2 {
298        return Err(StatsError::InsufficientInput {
299            metric: "sample_variance",
300            minimum: 2,
301            actual: values.len(),
302        });
303    }
304    let mean = mean(values)?;
305    Ok(values
306        .iter()
307        .map(|value| {
308            let delta = value - mean;
309            delta * delta
310        })
311        .sum::<f64>()
312        / (values.len() - 1) as f64)
313}
314
315/// Computes the comparison/reference rate ratio used by the four-fifths rule.
316pub fn four_fifths_ratio(reference_rate: f64, comparison_rate: f64) -> StatsResult<f64> {
317    validate_probability("four_fifths_ratio", None, reference_rate)?;
318    validate_probability("four_fifths_ratio", None, comparison_rate)?;
319    if reference_rate == 0.0 {
320        return Err(StatsError::ZeroReferenceRate {
321            metric: "four_fifths_ratio",
322        });
323    }
324    Ok(comparison_rate / reference_rate)
325}
326
327/// Computes disparate impact and the four-fifths pass/fail flag.
328///
329/// Divides the comparison group's selection rate by the reference group's and
330/// flags whether the resulting ratio clears the four-fifths (0.8) threshold.
331///
332/// # Examples
333///
334/// ```
335/// use sim_lib_numbers_stats::{BinaryOutcomeCounts, disparate_impact};
336///
337/// let reference = BinaryOutcomeCounts::new(80, 100).unwrap();
338/// let comparison = BinaryOutcomeCounts::new(60, 100).unwrap();
339/// let impact = disparate_impact(reference, comparison).unwrap();
340///
341/// assert!((impact.ratio - 0.75).abs() < 1e-12);
342/// assert!(!impact.passes_four_fifths);
343/// ```
344pub fn disparate_impact(
345    reference: BinaryOutcomeCounts,
346    comparison: BinaryOutcomeCounts,
347) -> StatsResult<DisparateImpact> {
348    let reference_rate = reference.selection_rate();
349    let comparison_rate = comparison.selection_rate();
350    let ratio = four_fifths_ratio(reference_rate, comparison_rate)?;
351    Ok(DisparateImpact {
352        reference_rate,
353        comparison_rate,
354        ratio,
355        passes_four_fifths: ratio >= FOUR_FIFTHS_THRESHOLD,
356    })
357}
358
359fn validate_values(metric: &'static str, values: &[f64]) -> StatsResult<()> {
360    if values.is_empty() {
361        return Err(StatsError::EmptyInput { metric });
362    }
363    for (index, value) in values.iter().copied().enumerate() {
364        validate_finite(metric, Some(index), value)?;
365    }
366    Ok(())
367}
368
369fn validate_probability(metric: &'static str, index: Option<usize>, value: f64) -> StatsResult<()> {
370    validate_finite(metric, index, value)?;
371    if !(0.0..=1.0).contains(&value) {
372        return Err(StatsError::ProbabilityOutOfRange {
373            metric,
374            index,
375            value,
376        });
377    }
378    Ok(())
379}
380
381fn validate_finite(metric: &'static str, index: Option<usize>, value: f64) -> StatsResult<()> {
382    if value.is_finite() {
383        Ok(())
384    } else {
385        Err(StatsError::NonFinite {
386            metric,
387            index,
388            value,
389        })
390    }
391}