Skip to main content

lawkit_core/generate/
normal.rs

1use super::{DataGenerator, GenerateConfig};
2use crate::error::Result;
3use rand::prelude::*;
4use rand_distr::{Distribution, Normal};
5
6#[derive(Debug, Clone)]
7pub struct NormalGenerator {
8    pub mean: f64,
9    pub stddev: f64,
10}
11
12impl NormalGenerator {
13    pub fn new(mean: f64, stddev: f64) -> Self {
14        Self { mean, stddev }
15    }
16}
17
18impl DataGenerator for NormalGenerator {
19    type Output = Vec<f64>;
20
21    fn generate(&self, config: &GenerateConfig) -> Result<Self::Output> {
22        let mut rng = config.create_rng();
23        let mut numbers = Vec::with_capacity(config.samples);
24
25        let normal = Normal::new(self.mean, self.stddev).map_err(|e| {
26            crate::error::BenfError::ParseError(format!("Invalid normal parameters: {e}"))
27        })?;
28
29        for _ in 0..config.samples {
30            let value = normal.sample(&mut rng);
31            numbers.push(value);
32        }
33
34        // Inject fraud if specified (add non-normal outliers)
35        if config.fraud_rate > 0.0 {
36            inject_normal_fraud(&mut numbers, config.fraud_rate, &mut rng);
37        }
38
39        Ok(numbers)
40    }
41}
42
43fn inject_normal_fraud(numbers: &mut [f64], fraud_rate: f64, rng: &mut impl Rng) {
44    let fraud_count = (numbers.len() as f64 * fraud_rate) as usize;
45    let mean = numbers.iter().sum::<f64>() / numbers.len() as f64;
46    let stddev =
47        (numbers.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / numbers.len() as f64).sqrt();
48
49    // Add outliers beyond 3 standard deviations
50    for _ in 0..fraud_count {
51        let index = rng.gen_range(0..numbers.len());
52        let outlier_multiplier = rng.gen_range(3.5..6.0);
53        let sign = if rng.gen_bool(0.5) { 1.0 } else { -1.0 };
54        numbers[index] = mean + sign * outlier_multiplier * stddev;
55    }
56}
57
58#[cfg(test)]
59mod tests {
60    use super::*;
61
62    #[test]
63    fn test_normal_generator() {
64        let generator = NormalGenerator::new(100.0, 15.0);
65        let config = GenerateConfig::new(1000).with_seed(42);
66
67        let result = generator.generate(&config).unwrap();
68        assert_eq!(result.len(), 1000);
69
70        // Check mean and standard deviation are approximately correct
71        let mean = result.iter().sum::<f64>() / result.len() as f64;
72        let variance = result.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / result.len() as f64;
73        let stddev = variance.sqrt();
74
75        assert!((mean - 100.0).abs() < 5.0); // Within 5 units of target mean
76        assert!((stddev - 15.0).abs() < 3.0); // Within 3 units of target stddev
77    }
78}