From bfb0137e9fa9c1ed225ad48a4a484a4eca56fe43 Mon Sep 17 00:00:00 2001 From: peterc-s Date: Tue, 28 Apr 2026 20:33:14 +0100 Subject: [PATCH] make dataset simulation multithreaded ready for re-benchmarking --- crates/chaff-cli/src/main.rs | 6 +- crates/chaff-cli/src/subcommands/simulate.rs | 93 +++++++++++++++----- 2 files changed, 75 insertions(+), 24 deletions(-) diff --git a/crates/chaff-cli/src/main.rs b/crates/chaff-cli/src/main.rs index ec5c7b2..da95982 100644 --- a/crates/chaff-cli/src/main.rs +++ b/crates/chaff-cli/src/main.rs @@ -12,6 +12,7 @@ use chaff::machine::Machine; use chaff_cli::{ errors::CliError, subcommands::{cap_convert, capture, dataset_convert, dataset_stats, simulate, trace_stats}, + utils::parse_dataset, }; use chaff_machines::constant; @@ -159,7 +160,10 @@ fn run() -> Result<(), CliError> { output, input, dataset_type, - } => simulate::run_dataset(&input, &dataset_type, &output, machine), + } => { + let input_dataset = parse_dataset(&dataset_type, &input)?; + simulate::run_dataset(&input_dataset, &output, &machine) + } } } CliOptions::CapConvert { pcap, trace, mac } => cap_convert::run(mac, &pcap, &trace), diff --git a/crates/chaff-cli/src/subcommands/simulate.rs b/crates/chaff-cli/src/subcommands/simulate.rs index 8c852ff..e2df0d6 100644 --- a/crates/chaff-cli/src/subcommands/simulate.rs +++ b/crates/chaff-cli/src/subcommands/simulate.rs @@ -1,13 +1,18 @@ //! Module for the `chaff-cli sim` subcommand. -use std::{fs, path::PathBuf}; +use std::{ + fs, + path::PathBuf, + sync::{Arc, Mutex, mpsc}, + thread, +}; use chaff::{framework::Framework, machine::Machine}; use chaff_capture::trace::Trace; -use chaff_datasets::dataset::DatasetBuilder; +use chaff_datasets::dataset::{Dataset, DatasetBuilder}; use chaff_sim::{Simulator, SimulatorOverheads}; -use crate::{errors::CliError, utils::parse_dataset}; +use crate::errors::CliError; /// Run the simulator on a singular trace. /// @@ -41,34 +46,76 @@ pub fn run_trace( /// # Errors /// /// If parsing the dataset fails. If output is supplied, an error may be returned if dumping fails. +/// +/// # Panics +/// +/// If a [`std::sync::Mutex::lock`] fails. pub fn run_dataset( - input: &PathBuf, - dataset_type: &str, + dataset: &Dataset, output: &Option, - machine: Machine, + machine: &Machine, ) -> Result<(), CliError> { - let input_dataset = parse_dataset(dataset_type, input)?; - let input_data = input_dataset.get_dataset(); + let input_data = dataset.get_dataset(); + let mut output_dataset_builder = DatasetBuilder::new(dataset.get_pad_to()); - let mut output_dataset_builder = DatasetBuilder::new(input_dataset.get_pad_to()); + let tasks: Vec<_> = input_data + .iter() + .flat_map(|(class, traces)| traces.iter().map(move |trace| (class, trace))) + .collect(); - let framework = Framework::new(machine, rand::rng()); - let mut sim = Simulator::with(framework, Trace::default(), rand::rng()); - let mut overheads = Vec::with_capacity(input_data.len()); - - for (class, traces) in input_data { - for trace in traces { - sim.replace_trace(trace.clone()); - let (trace, overhead) = sim.run(); - output_dataset_builder.push_to_class(class, trace); + let num_tasks = tasks.len(); + let num_threads = thread::available_parallelism().map_or(1, std::num::NonZero::get); + + let work_queue = Arc::new(Mutex::new(tasks.into_iter())); + let (tx, rx) = mpsc::channel(); + + thread::scope(|s| { + for _ in 0..num_threads { + let thread_tx = tx.clone(); + let thread_machine = machine.clone(); + let thread_queue = Arc::clone(&work_queue); + + s.spawn(move || { + let mut sim = Simulator::with( + Framework::new(thread_machine, rand::rng()), + Trace::default(), + rand::rng(), + ); + + loop { + let task = { + #[expect(clippy::expect_used)] + let mut queue = thread_queue + .lock() + .expect("other thread panicked while holding thread queue"); + queue.next() + }; + + match task { + Some((class, trace)) => { + sim.replace_trace(trace.clone()); + let (out_trace, overhead) = sim.run(); + let _ = thread_tx.send((class, out_trace, overhead)); + } + None => break, + } + } + }); + } + + drop(tx); + + let mut overheads = Vec::with_capacity(num_tasks); + while let Ok((class, out_trace, overhead)) = rx.recv() { + output_dataset_builder.push_to_class(class, out_trace); overheads.push(overhead); } - } - let overheads = SimulatorOverheads::total_from(overheads); - if let Some(overheads) = overheads { - println!("{overheads}"); - } + let overheads_total = SimulatorOverheads::total_from(overheads); + if let Some(total) = overheads_total { + println!("{total}"); + } + }); if let Some(output) = output { if !output.exists() {