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