dstest 0.1.5

Deterministic Simulation Testing for containerised services
use std::collections::HashMap;

#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum AccumulationMode {
    #[default]
    Single,
    Accumulate,
}

impl std::fmt::Display for AccumulationMode {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        match self {
            AccumulationMode::Single => write!(f, "single"),
            AccumulationMode::Accumulate => write!(f, "accumulate"),
        }
    }
}

impl std::str::FromStr for AccumulationMode {
    type Err = String;

    fn from_str(s: &str) -> Result<Self, Self::Err> {
        if s.eq_ignore_ascii_case("single") {
            Ok(AccumulationMode::Single)
        } else if s.eq_ignore_ascii_case("accumulate") {
            Ok(AccumulationMode::Accumulate)
        } else {
            Err(format!("invalid accumulation mode: {}", s))
        }
    }
}

#[derive(Clone, Debug)]
pub struct Config {
    pub substrate: Option<String>,
    pub seed: Option<u64>,
    pub fault_weights: HashMap<String, f32>,
    pub accumulation_mode: AccumulationMode,
    pub http_timeout_secs: u64,
    pub http_retries: u32,
    pub http_retry_delay_ms: u64,
    pub step_delay_ms: u64,
    pub require_seed: bool,
}

impl Default for Config {
    fn default() -> Self {
        let mut fault_weights = HashMap::new();
        fault_weights.insert(String::from("pause"), 0.35);
        fault_weights.insert(String::from("kill"), 0.25);
        fault_weights.insert(String::from("deprive:disk"), 0.10);
        fault_weights.insert(String::from("deprive:network"), 0.10);
        fault_weights.insert(String::from("deprive:memory"), 0.10);
        fault_weights.insert(String::from("deprive:cpu"), 0.10);

        Self {
            substrate: None,
            seed: None,
            fault_weights,
            accumulation_mode: AccumulationMode::Single,
            http_timeout_secs: 5,
            http_retries: 30,
            http_retry_delay_ms: 500,
            step_delay_ms: 1000,
            require_seed: true,
        }
    }
}

impl Config {
    pub fn validate(&self) -> Result<(), String> {
        let total: f32 = self.fault_weights.values().sum();
        if (total - 1.0).abs() > 0.01 {
            return Err(format!("fault weights must sum to 1.0, got {}", total));
        }

        for (name, weight) in &self.fault_weights {
            if *weight < 0.0 {
                return Err(format!("weight for '{}' is negative: {}", name, weight));
            }
        }

        Ok(())
    }

    pub fn normalize_weights(&mut self) {
        let total: f32 = self.fault_weights.values().sum();
        if total > 0.0 {
            for weight in self.fault_weights.values_mut() {
                *weight /= total;
            }
        }
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use std::str::FromStr;

    #[test]
    fn test_accumulation_mode_display() {
        assert_eq!(AccumulationMode::Single.to_string(), "single");
        assert_eq!(AccumulationMode::Accumulate.to_string(), "accumulate");
    }

    #[test]
    fn test_accumulation_mode_from_str() {
        assert!(matches!(
            AccumulationMode::from_str("single"),
            Ok(AccumulationMode::Single)
        ));
        assert!(matches!(
            AccumulationMode::from_str("SINGLE"),
            Ok(AccumulationMode::Single)
        ));
        assert!(matches!(
            AccumulationMode::from_str("accumulate"),
            Ok(AccumulationMode::Accumulate)
        ));
        assert!(AccumulationMode::from_str("invalid").is_err());
    }

    #[test]
    fn test_config_default() {
        let config = Config::default();
        assert!(config.fault_weights.contains_key("pause"));
        assert!(config.fault_weights.contains_key("kill"));
        assert!(config.require_seed);
    }

    #[test]
    fn test_config_validate() {
        let config = Config::default();
        assert!(config.validate().is_ok());
    }

    #[test]
    fn test_config_validate_negative_weight() {
        let mut config = Config::default();
        config.fault_weights.insert("test".to_string(), -0.5);
        assert!(config.validate().is_err());
    }

    #[test]
    fn test_config_normalize_weights() {
        let mut config = Config::default();
        config.fault_weights.insert("extra".to_string(), 1.0);
        config.normalize_weights();
        let total: f32 = config.fault_weights.values().sum();
        assert!((total - 1.0).abs() < 0.01);
    }
}