dstest 0.1.5

Deterministic Simulation Testing for containerised services
use std::sync::Arc;

use mlua::{Lua, Result, Table};

use crate::config::AccumulationMode;
use crate::engine::context::BindingContext;
use crate::fault::FaultTree;
use crate::substrate::Substrate;

#[allow(clippy::await_holding_lock)]
pub fn register<S: Substrate>(lua: &Lua, dstest: &Table, ctx: &BindingContext<S>) -> Result<()> {
    let state = Arc::clone(&ctx.state);
    let config = Arc::clone(&ctx.config);
    let fault_tree = Arc::clone(&ctx.fault_tree);
    let substrate = Arc::clone(&ctx.substrate);
    let oracle = Arc::clone(&ctx.oracle);

    let step_fn = lua.create_async_function(move |lua, ()| {
        let state = Arc::clone(&state);
        let config = Arc::clone(&config);
        let fault_tree = Arc::clone(&fault_tree);
        let substrate = Arc::clone(&substrate);
        let oracle = Arc::clone(&oracle);

        async move {
            let (require_seed, accumulation_mode, step_delay) = {
                let cfg = config.lock().expect("poisoned config lock");
                (cfg.require_seed, cfg.accumulation_mode, cfg.step_delay_ms)
            };

            if require_seed {
                let s = state.lock().expect("poisoned engine state lock");
                if s.seed.is_none() {
                    return Err(mlua::Error::RuntimeError(
                        "dstest.config({ seed = n }) must be called before dstest.step()".into(),
                    ));
                }
            }

            let mut tree_guard = fault_tree.lock().expect("poisoned fault tree lock");

            if tree_guard.is_none() {
                let s = state.lock().expect("poisoned engine state lock");
                if let Some(seed) = s.seed {
                    let subject_ids: Vec<String> =
                        s.subjects.iter().map(|(id, _)| id.clone()).collect();
                    if subject_ids.is_empty() {
                        return Err(mlua::Error::RuntimeError(
                            "no subjects available for fault injection".into(),
                        ));
                    }
                    let cfg = config.lock().expect("poisoned config lock");
                    *tree_guard = Some(FaultTree::new(seed, subject_ids, &cfg));
                } else {
                    return Err(mlua::Error::RuntimeError(
                        "dstest.config({ seed = n }) must be called before dstest.step()".into(),
                    ));
                }
            }

            let tree = tree_guard
                .as_mut()
                .ok_or_else(|| mlua::Error::RuntimeError("fault tree not initialized".into()))?;
            let result = tree.step();

            let Some(step_result) = result else {
                let t = lua.create_table()?;
                t.set("more", false)?;
                return Ok(t);
            };

            let subject = crate::substrate::Subject::new(step_result.subject_id.clone());

            match accumulation_mode {
                AccumulationMode::Single => {
                    substrate
                        .clear_faults(&subject)
                        .map_err(mlua::Error::RuntimeError)?;
                    tokio::time::sleep(std::time::Duration::from_millis(step_delay)).await;
                }
                AccumulationMode::Accumulate => {}
            }

            substrate
                .affect(&subject, &step_result.fault)
                .map_err(mlua::Error::RuntimeError)?;

            let t = lua.create_table()?;
            t.set("fault", step_result.fault.to_string())?;
            t.set("subject", step_result.subject_id.clone())?;
            t.set("round", step_result.round)?;
            t.set("total_rounds", step_result.total_rounds)?;
            t.set("remaining", step_result.remaining)?;
            t.set("more", step_result.more)?;

            {
                let mut o = oracle.lock().expect("poisoned oracle lock");
                if o.enabled {
                    let report = o
                        .check_all(
                            &lua,
                            &step_result.subject_id,
                            &step_result.fault.to_string(),
                            step_result.round,
                        )
                        .await;
                    o.report.merge(&report);
                    t.set("oracle", {
                        let ot = lua.create_table()?;
                        ot.set("passed", report.passed)?;
                        ot.set("total_checks", report.total_checks)?;
                        ot.set("passed_checks", report.passed_checks)?;
                        ot.set("failed_checks", report.failed_checks)?;
                        ot
                    })?;
                }
            }

            Ok(t)
        }
    })?;

    dstest.set("step", step_fn)?;
    Ok(())
}