diff --git a/Cargo.toml b/Cargo.toml index bf448bc9..05567431 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -25,7 +25,7 @@ warnings = "deny" [workspace.dependencies] libc = "0.2" log = "0.4" -rand = { version = "0.10", default-features = false } +rand = { version = "0.10", default-features = false, features = ["alloc"] } simple_logger = { version = "5", default-features = false } postcard = { version = "1.1", default-features = false, features = ["alloc"] } bitcoin = "0.32" diff --git a/smite-ir-mutator/src/lib.rs b/smite-ir-mutator/src/lib.rs index 30acdd3b..67775a2c 100644 --- a/smite-ir-mutator/src/lib.rs +++ b/smite-ir-mutator/src/lib.rs @@ -37,7 +37,7 @@ use std::os::raw::{c_char, c_uint, c_void}; use std::slice; use rand::rngs::SmallRng; -use rand::{RngExt, SeedableRng, seq::IteratorRandom}; +use rand::{RngExt, SeedableRng}; use smite_ir::generators::AnyGenerator; use smite_ir::minimizers::{CommonSubexpressionEliminator, DeadCodeEliminator, Minimizer}; @@ -72,15 +72,11 @@ impl MutatorState { } } - /// Generates a fresh program from scratch by randomly delegating to one of - /// the registered generators. + /// Generates a fresh program from scratch by delegating to one of the + /// registered generators, picked by weight. fn generate_fresh(&mut self) -> Program { let mut builder = ProgramBuilder::new(); - AnyGenerator::ALL - .iter() - .choose(&mut self.rng) - .expect("AnyGenerator::ALL is non-empty") - .generate(&mut builder, &mut self.rng); + AnyGenerator::choose(&mut self.rng).generate(&mut builder, &mut self.rng); self.last_sequence.clear(); self.last_sequence.push("fresh"); builder.build() @@ -114,11 +110,8 @@ impl MutatorState { "instr-reorder" } 4 => { - let generator = *AnyGenerator::ALL - .iter() - .choose(&mut self.rng) - .expect("AnyGenerator::ALL is non-empty"); - let mutator = GeneratorInsertionMutator::new(generator); + let mutator = + GeneratorInsertionMutator::new(AnyGenerator::choose(&mut self.rng)); mutator.mutate(program, &mut self.rng); "gen-insert" } diff --git a/smite-ir/src/generators.rs b/smite-ir/src/generators.rs index 58d1b2df..0c6c19f0 100644 --- a/smite-ir/src/generators.rs +++ b/smite-ir/src/generators.rs @@ -22,6 +22,7 @@ pub use node_announcement::NodeAnnouncementGenerator; pub use open_channel::OpenChannelGenerator; use rand::Rng; +use rand::seq::IndexedRandom; use super::builder::ProgramBuilder; @@ -55,6 +56,32 @@ impl AnyGenerator { Self::ChannelReady(ChannelReadyGenerator), Self::FundingFlow(FundingFlowGenerator), ]; + + /// Relative pick weight for [`Self::choose`]; 0 disables. + #[must_use] + pub fn weight(&self) -> u32 { + match self { + Self::ChannelAnnouncement(_) + | Self::ChannelUpdate(_) + | Self::NodeAnnouncement(_) + | Self::OpenChannel(_) + | Self::FundingCreated(_) + | Self::ChannelReady(_) + | Self::FundingFlow(_) => 10, + } + } + + /// Picks a generator from `ALL` with probability proportional to its + /// [`Self::weight`]. + /// + /// # Panics + /// + /// Panics if every generator has a weight of zero. + pub fn choose(rng: &mut impl Rng) -> Self { + *Self::ALL + .choose_weighted(rng, Self::weight) + .expect("at least one generator must have non-zero weight") + } } impl Generator for AnyGenerator { diff --git a/smite-ir/src/tests.rs b/smite-ir/src/tests.rs index b627d63c..dfc6e34a 100644 --- a/smite-ir/src/tests.rs +++ b/smite-ir/src/tests.rs @@ -921,6 +921,38 @@ fn any_generator_all_is_complete() { assert_eq!(AnyGenerator::ALL.len(), variant_count(AnyGenerator::ALL[0])); } +// Weighted selection must cover every generator with non-zero weight and +// skew towards heavier ones in proportion to their weight. +#[test] +fn any_generator_choose_respects_weights() { + let mut rng = SmallRng::seed_from_u64(7); + let mut counts = vec![0u32; AnyGenerator::ALL.len()]; + let rounds = 20_000; + for _ in 0..rounds { + let picked = AnyGenerator::choose(&mut rng); + let idx = AnyGenerator::ALL + .iter() + .position(|g| std::mem::discriminant(g) == std::mem::discriminant(&picked)) + .expect("chosen generator is in ALL"); + counts[idx] += 1; + } + + let total: u32 = AnyGenerator::ALL.iter().map(AnyGenerator::weight).sum(); + for (generator, &count) in AnyGenerator::ALL.iter().zip(&counts) { + let weight = generator.weight(); + if weight == 0 { + assert_eq!(count, 0, "zero-weight generator must never be picked"); + continue; + } + let expected = f64::from(rounds) * f64::from(weight) / f64::from(total); + let ratio = f64::from(count) / expected; + assert!( + (0.85..=1.15).contains(&ratio), + "generator picked {count} times, expected ~{expected:.0} (ratio {ratio:.2})" + ); + } +} + // -- ShutdownScriptVariant tests -- // Ensure ShutdownScriptVariant and ShutdownScriptVariant::VARIANT_COUNT stay in