1use 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
27pub type StatsResult<T> = Result<T, StatsError>;
29
30#[derive(Clone, Debug, PartialEq)]
35pub enum StatsError {
36 EmptyInput {
38 metric: &'static str,
40 },
41 InsufficientInput {
43 metric: &'static str,
45 minimum: usize,
47 actual: usize,
49 },
50 NonFinite {
52 metric: &'static str,
54 index: Option<usize>,
56 value: f64,
58 },
59 ProbabilityOutOfRange {
61 metric: &'static str,
63 index: Option<usize>,
65 value: f64,
67 },
68 ProbabilityMass {
70 metric: &'static str,
72 sum: f64,
74 },
75 ZeroEvidence {
77 metric: &'static str,
79 },
80 ZeroTotal {
82 label: &'static str,
84 },
85 ZeroReferenceRate {
87 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#[derive(Clone, Copy, Debug, PartialEq)]
142pub struct BinaryOutcomeCounts {
143 pub selected: u64,
145 pub total: u64,
147}
148
149impl BinaryOutcomeCounts {
150 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 pub fn selection_rate(self) -> f64 {
169 self.selected as f64 / self.total as f64
170 }
171}
172
173#[derive(Clone, Copy, Debug, PartialEq)]
175pub struct DisparateImpact {
176 pub reference_rate: f64,
178 pub comparison_rate: f64,
180 pub ratio: f64,
182 pub passes_four_fifths: bool,
184}
185
186pub 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
213pub 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
226pub 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
249pub fn mean(values: &[f64]) -> StatsResult<f64> {
262 validate_values("mean", values)?;
263 Ok(values.iter().sum::<f64>() / values.len() as f64)
264}
265
266pub fn variance(values: &[f64]) -> StatsResult<f64> {
279 population_variance(values)
280}
281
282pub 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
296pub 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
317pub 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
329pub 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}