1use 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
25pub type StatsResult<T> = Result<T, StatsError>;
27
28#[derive(Clone, Debug, PartialEq)]
33pub enum StatsError {
34 EmptyInput {
36 metric: &'static str,
38 },
39 InsufficientInput {
41 metric: &'static str,
43 minimum: usize,
45 actual: usize,
47 },
48 NonFinite {
50 metric: &'static str,
52 index: Option<usize>,
54 value: f64,
56 },
57 ProbabilityOutOfRange {
59 metric: &'static str,
61 index: Option<usize>,
63 value: f64,
65 },
66 ProbabilityMass {
68 metric: &'static str,
70 sum: f64,
72 },
73 ZeroEvidence {
75 metric: &'static str,
77 },
78 ZeroTotal {
80 label: &'static str,
82 },
83 ZeroReferenceRate {
85 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#[derive(Clone, Copy, Debug, PartialEq)]
140pub struct BinaryOutcomeCounts {
141 pub selected: u64,
143 pub total: u64,
145}
146
147impl BinaryOutcomeCounts {
148 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 pub fn selection_rate(self) -> f64 {
167 self.selected as f64 / self.total as f64
168 }
169}
170
171#[derive(Clone, Copy, Debug, PartialEq)]
173pub struct DisparateImpact {
174 pub reference_rate: f64,
176 pub comparison_rate: f64,
178 pub ratio: f64,
180 pub passes_four_fifths: bool,
182}
183
184pub 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
211pub 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
224pub 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
247pub fn mean(values: &[f64]) -> StatsResult<f64> {
260 validate_values("mean", values)?;
261 Ok(values.iter().sum::<f64>() / values.len() as f64)
262}
263
264pub fn variance(values: &[f64]) -> StatsResult<f64> {
277 population_variance(values)
278}
279
280pub 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
294pub 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
315pub 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
327pub 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}