diff --git a/Cargo.lock b/Cargo.lock index 8665c168..21e0ad42 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -613,8 +613,10 @@ dependencies = [ name = "smite-ir-mutator" version = "0.0.0" dependencies = [ + "log", "postcard", "rand", + "simple_logger", "smite-ir", ] diff --git a/README.md b/README.md index bcc8a0c1..1041f703 100644 --- a/README.md +++ b/README.md @@ -66,13 +66,11 @@ printf '\x00' > /tmp/smite-seeds/empty AFL_CUSTOM_MUTATOR_LIBRARY=target/release/libsmite_ir_mutator.so \ AFL_CUSTOM_MUTATOR_ONLY=1 \ AFL_FRAMESHIFT_DISABLE=1 \ -AFL_DISABLE_TRIM=1 \ ~/AFLplusplus/afl-fuzz -X -i /tmp/smite-seeds -o /tmp/smite-out -- /tmp/smite-nyx ``` `AFL_CUSTOM_MUTATOR_ONLY=1` disables AFL++'s built-in mutators (which would -corrupt the postcard encoding). `AFL_DISABLE_TRIM=1` prevents AFL++ from -trimming inputs (which would also corrupt the encoding). +corrupt the postcard encoding). ## Running Modes diff --git a/smite-ir-mutator/Cargo.toml b/smite-ir-mutator/Cargo.toml index dd924818..b75a317d 100644 --- a/smite-ir-mutator/Cargo.toml +++ b/smite-ir-mutator/Cargo.toml @@ -14,3 +14,5 @@ crate-type = ["cdylib"] smite-ir = { path = "../smite-ir" } rand.workspace = true postcard.workspace = true +log.workspace = true +simple_logger.workspace = true diff --git a/smite-ir-mutator/src/lib.rs b/smite-ir-mutator/src/lib.rs index 856102e8..55db77e5 100644 --- a/smite-ir-mutator/src/lib.rs +++ b/smite-ir-mutator/src/lib.rs @@ -16,8 +16,15 @@ //! - `AFL_FRAMESHIFT_DISABLE=1` -- disable AFL++'s `FrameShift` analysis that //! bypasses our custom mutators. This was an AFL++ bug fixed upstream in //! commit eddb2701b022351fb34b696ccf923bb856e9d953. -//! - `AFL_DISABLE_TRIM=1` -- this library does not implement custom trim and -//! AFL++'s default byte-level trim would corrupt our structured programs. +//! +//! # Logging +//! +//! Logging is opt-in: [`afl_custom_init`] installs a logger only when +//! `RUST_LOG` is set. With `RUST_LOG` unset, every `log::*!` callsite is a +//! no-op and AFL's stderr stays clean. Useful filters: +//! - `RUST_LOG=smite_ir_mutator::trim=debug` -- print the decoded `Program` +//! before and after each successful trim. +//! - `RUST_LOG=debug` -- everything this crate emits. //! //! # Buffer ownership //! @@ -33,6 +40,7 @@ use rand::rngs::SmallRng; use rand::{RngExt, SeedableRng}; use smite_ir::generators::{NodeAnnouncementGenerator, OpenChannelGenerator}; +use smite_ir::minimizers::{CommonSubexpressionEliminator, DeadCodeEliminator, Minimizer}; use smite_ir::mutators::{InputSwapMutator, OperationParamMutator}; use smite_ir::{Generator, Mutator, Program, ProgramBuilder}; @@ -152,14 +160,26 @@ fn warn_on_unset_afl_env() { /// Allocates a new [`MutatorState`] and returns an opaque pointer to it. AFL++ /// passes this pointer back on every function call as the `data` argument. /// +/// Also installs `simple_logger` if `RUST_LOG` is set, so `log::*!` callsites +/// in this crate become live. Without `RUST_LOG` no logger is installed and +/// every callsite is a no-op. Set +/// `RUST_LOG=smite_ir_mutator::trim=debug` to see trim before/after dumps. +/// /// # Safety /// /// The returned pointer is heap-allocated via `Box::into_raw` and must be freed /// by a matching call to [`afl_custom_deinit`]. +/// +/// # Panics +/// +/// Panics if `RUST_LOG` is set and `simple_logger` fails to initialize. #[unsafe(no_mangle)] pub unsafe extern "C" fn afl_custom_init(_afl: *const c_void, seed: c_uint) -> *mut c_void { #[cfg(not(test))] warn_on_unset_afl_env(); + if std::env::var_os("RUST_LOG").is_some() { + simple_logger::init_with_env().expect("logger initializes"); + } Box::into_raw(Box::new(MutatorState::new(seed))).cast::() } @@ -228,6 +248,115 @@ pub unsafe extern "C" fn afl_custom_fuzz( len } +/// Runs the full minimizer pipeline (`DeadCodeEliminator` then +/// `CommonSubexpressionEliminator`) on the corpus entry and stages the +/// resulting candidate for [`afl_custom_trim`] to hand back. +/// +/// Both minimizers are deterministic in-process transforms safe in IR +/// semantics, so we don't need iterative AFL feedback. We compose them +/// once and offer a single candidate. AFL still gets to verify it (its +/// coverage cksum is the source of truth); on rejection AFL silently +/// discards the candidate and keeps the original corpus entry. +/// +/// AFL drives the trim loop with `while (stage_cur < stage_max)`, where +/// `stage_max` is this function's return value and `stage_cur` is updated +/// from [`afl_custom_post_trim`]'s return. +/// +/// # Returns +/// +/// - `1` if there's a candidate to offer (decode succeeded, validate +/// passed, and the trim actually shrank the program). AFL enters the +/// trim loop for one iteration. +/// - `0` if there's nothing to do (decode/validate failed, or the trim +/// was a no-op). AFL skips trim entirely. +/// - Negative would signal a fatal error to AFL; we never produce one. +/// +/// # Safety +/// +/// - `data` must be a pointer returned by [`afl_custom_init`]. +/// - `buf` must point to `buf_size` readable bytes. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn afl_custom_init_trim( + data: *mut c_void, + buf: *mut u8, + buf_size: usize, +) -> i32 { + let state = unsafe { &mut *data.cast::() }; + + let input = unsafe { slice::from_raw_parts(buf, buf_size) }; + let Some(program) = decode_and_validate(input) else { + return 0; + }; + + let before = log::log_enabled!(target: "smite_ir_mutator::trim", log::Level::Debug) + .then(|| program.clone()); + + let mut trimmed = program; + let dce_changed = DeadCodeEliminator.minimize(&mut trimmed); + let cse_changed = CommonSubexpressionEliminator.minimize(&mut trimmed); + if (!dce_changed && !cse_changed) || !state.serialize(&trimmed, buf_size) { + return 0; + } + + if let Some(before) = before { + log::debug!( + target: "smite_ir_mutator::trim", + "dce={dce_changed} cse={cse_changed}\n--- before ---\n{before}--- after ---\n{trimmed}---" + ); + } + + 1 +} + +/// Hands the pre-serialized trimmed candidate back to AFL. +/// +/// The pointer written into `*out_buf` borrows from `MutatorState::out_buf` +/// and is valid until the next call into this library; AFL copies the +/// bytes before re-entering us. We always write a non-null pointer (even +/// on the zero-length path) to satisfy AFL's `if (unlikely(!retbuf)) +/// FATAL(...)` check. +/// +/// # Returns +/// +/// - `> 0` on the first call after [`afl_custom_init_trim`]: the byte +/// length of the candidate at `*out_buf`. +/// - `0` afterwards. AFL treats this as "skip this iteration" rather than +/// a stop signal; the loop terminates via [`afl_custom_post_trim`]'s +/// return. +/// +/// # Safety +/// +/// - `data` must be a pointer returned by [`afl_custom_init`]. +/// - `out_buf` must be a valid, writable pointer to a `*const u8` slot. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn afl_custom_trim(data: *mut c_void, out_buf: *mut *const u8) -> usize { + let state = unsafe { &mut *data.cast::() }; + unsafe { *out_buf = state.out_buf.as_ptr() }; + state.out_buf.len() +} + +/// Always returns `1` to terminate AFL's trim loop after a single +/// iteration. +/// +/// AFL drives trim with `while (stage_cur < stage_max)` and assigns +/// `stage_cur` from this function's return value. With `stage_max = 1` +/// (set by [`afl_custom_init_trim`]), returning `1` makes the condition +/// `1 < 1` false and breaks the loop. +/// +/// `success` indicates whether the candidate's coverage cksum matched the +/// original. We don't need to act on it: AFL itself either persists the +/// trimmed buffer (on success) or keeps the original corpus entry (on +/// failure), and we don't track partial state across iterations because +/// there's only one. +/// +/// # Safety +/// +/// - `data` must be a pointer returned by [`afl_custom_init`]. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn afl_custom_post_trim(_data: *mut c_void, _success: u8) -> i32 { + 1 +} + /// Marker symbol that tells AFL++ not to populate `add_buf` for /// [`afl_custom_fuzz`]. AFL++ never actually calls this function -- it only /// checks for the symbol's presence via `dlsym` and, if found, skips picking a @@ -338,6 +467,19 @@ mod tests { postcard::to_allocvec(&builder.build()).expect("postcard serialization") } + /// `seed_program_bytes()` plus an unreferenced `LoadAmount`, so the + /// pipeline has something for DCE to drop (and thus `init_trim` + /// returns `1`). + fn reducible_seed_bytes() -> Vec { + let bytes = seed_program_bytes(); + let mut program: Program = postcard::from_bytes(&bytes).expect("decode"); + program.instructions.push(smite_ir::Instruction { + operation: smite_ir::Operation::LoadAmount(0xdead_beef), + inputs: vec![], + }); + postcard::to_allocvec(&program).expect("encode") + } + #[test] fn init_returns_nonnull() { let state = State::new(0); @@ -435,4 +577,106 @@ mod tests { // crash either. unsafe { afl_custom_splice_optout(ptr::null_mut()) }; } + + // -- Trim tests -- + + fn init_trim_via_ffi(state: &State, mut input: Vec) -> i32 { + unsafe { afl_custom_init_trim(state.0, input.as_mut_ptr(), input.len()) } + } + + fn trim_via_ffi(state: &State) -> (*const u8, usize) { + let mut out: *const u8 = ptr::null(); + let len = unsafe { afl_custom_trim(state.0, &raw mut out) }; + (out, len) + } + + fn post_trim_via_ffi(state: &State, success: bool) -> i32 { + unsafe { afl_custom_post_trim(state.0, u8::from(success)) } + } + + #[test] + fn trim_init_returns_1_when_reduction_possible() { + let state = State::new(0); + let rv = init_trim_via_ffi(&state, reducible_seed_bytes()); + assert_eq!(rv, 1); + } + + #[test] + fn trim_init_returns_0_when_no_reduction_possible() { + // Generator output has no dead code or duplicate loads; the + // pipeline is a no-op, so we tell AFL to skip trim entirely. + let state = State::new(0); + let rv = init_trim_via_ffi(&state, seed_program_bytes()); + assert_eq!(rv, 0); + } + + #[test] + fn trim_init_returns_0_for_garbage() { + let state = State::new(0); + let rv = init_trim_via_ffi(&state, vec![0xFF; 16]); + assert_eq!(rv, 0); + } + + #[test] + fn trim_yields_candidate_after_init() { + let state = State::new(0); + init_trim_via_ffi(&state, reducible_seed_bytes()); + let (out, len) = trim_via_ffi(&state); + assert!(len > 0); + decode_and_validate(out, len); + } + + #[test] + fn trim_post_trim_returns_1_to_terminate_loop() { + let state = State::new(0); + init_trim_via_ffi(&state, reducible_seed_bytes()); + let _ = trim_via_ffi(&state); + // post_trim returns 1 unconditionally — it's the load-bearing + // termination signal that pushes AFL's `stage_cur` to `stage_max`. + assert_eq!(post_trim_via_ffi(&state, true), 1); + assert_eq!(post_trim_via_ffi(&state, false), 1); + } + + #[test] + fn trim_init_does_not_overwrite_sequence() { + // Trim is not a mutation; `last_sequence` (used by `describe` to + // name queue entries from fuzz) must survive both the no-op and + // successful trim paths. + for (label, input, expected_rv) in [ + ("no-op", seed_program_bytes(), 0), + ("success", reducible_seed_bytes(), 1), + ] { + let state = State::new(0); + // Run a fuzz call so last_sequence has known contents. + let _ = fuzz_via_ffi(&state, Vec::new(), 1 << 16); + let before = unsafe { CStr::from_ptr(afl_custom_describe(state.0, 256)) } + .to_str() + .expect("valid utf-8") + .to_string(); + let rv = init_trim_via_ffi(&state, input); + assert_eq!(rv, expected_rv, "{label}"); + let after = unsafe { CStr::from_ptr(afl_custom_describe(state.0, 256)) } + .to_str() + .expect("valid utf-8") + .to_string(); + assert_eq!(before, after, "{label}"); + } + } + + #[test] + fn trim_candidate_is_smaller_than_input() { + let original_bytes = reducible_seed_bytes(); + let original_program: Program = postcard::from_bytes(&original_bytes).expect("decode"); + + let state = State::new(0); + init_trim_via_ffi(&state, original_bytes); + + let (out, len) = trim_via_ffi(&state); + assert!(len > 0, "trim should yield a candidate"); + let trimmed = decode_and_validate(out, len); + assert!( + trimmed.instructions.len() < original_program.instructions.len(), + "trim should shrink instruction count" + ); + } } diff --git a/smite-ir/src/instruction.rs b/smite-ir/src/instruction.rs index 8c5b9c30..c4a18aff 100644 --- a/smite-ir/src/instruction.rs +++ b/smite-ir/src/instruction.rs @@ -12,7 +12,7 @@ use super::Operation; /// /// In SSA form, each instruction produces at most one variable (at the index /// equal to the instruction's position in the program). -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] pub struct Instruction { /// The operation to perform. pub operation: Operation, diff --git a/smite-ir/src/lib.rs b/smite-ir/src/lib.rs index 72d0874a..70992c4b 100644 --- a/smite-ir/src/lib.rs +++ b/smite-ir/src/lib.rs @@ -7,6 +7,7 @@ //! //! # Modules //! - [`instruction`] - Single IR instruction (operation + input references). +//! - [`minimizers`] - Shrink a program while preserving interesting behaviour. //! - [`operation`] - Operations that load, compute, build or act. //! - [`program`] - Ordered list of instructions. //! - [`variable`] - Typed runtime values and lightweight type tags. @@ -14,6 +15,7 @@ pub mod builder; pub mod generators; pub mod instruction; +pub mod minimizers; pub mod mutators; pub mod operation; pub mod program; @@ -22,6 +24,7 @@ pub mod variable; pub use builder::ProgramBuilder; pub use generators::Generator; pub use instruction::Instruction; +pub use minimizers::Minimizer; pub use mutators::Mutator; pub use operation::Operation; pub use program::Program; diff --git a/smite-ir/src/minimizers.rs b/smite-ir/src/minimizers.rs new file mode 100644 index 00000000..7cf7cd7e --- /dev/null +++ b/smite-ir/src/minimizers.rs @@ -0,0 +1,24 @@ +//! IR program minimizers. +//! +//! A [`Minimizer`] reduces a [`Program`] to a smaller, behaviourally +//! equivalent version in a single pass. Both transforms are safe in IR +//! semantics, so they don't need an oracle to drive the search. +//! +//! Run them in pipeline order for best results: +//! 1. [`DeadCodeEliminator`] — drop dead instructions and reindex +//! 2. [`CommonSubexpressionEliminator`] — merge equivalent pure expressions + +mod cse; +mod dead_code; + +pub use cse::CommonSubexpressionEliminator; +pub use dead_code::DeadCodeEliminator; + +use super::Program; + +/// A minimizer that reduces an IR program in one call. +pub trait Minimizer { + /// Reduces `program` in place to a smaller, behaviourally equivalent + /// version. Returns `true` if the program was modified. + fn minimize(&self, program: &mut Program) -> bool; +} diff --git a/smite-ir/src/minimizers/cse.rs b/smite-ir/src/minimizers/cse.rs new file mode 100644 index 00000000..e900a6d9 --- /dev/null +++ b/smite-ir/src/minimizers/cse.rs @@ -0,0 +1,52 @@ +//! Common-subexpression elimination minimizer. + +use std::collections::HashMap; +use std::collections::hash_map::Entry; + +use super::Minimizer; +use crate::{Instruction, Program}; + +/// Merges instructions that compute the same pure expression. +/// +/// Two pure instructions are equivalent when they share the same operation +/// and the same canonicalized inputs. Walking the program in order makes +/// the merge transitive: by the time we reach instruction `i`, SSA +/// guarantees every input it references is already canonicalized, so two +/// compute ops whose inputs collapsed to the same canonical loads are +/// themselves recognized as equivalent. +pub struct CommonSubexpressionEliminator; + +impl Minimizer for CommonSubexpressionEliminator { + fn minimize(&self, program: &mut Program) -> bool { + let n = program.instructions.len(); + let mut canonical: HashMap = HashMap::new(); + let mut new_idx = vec![0usize; n]; + let mut instructions = Vec::with_capacity(n); + + for (i, mut instr) in std::mem::take(&mut program.instructions) + .into_iter() + .enumerate() + { + for input in &mut instr.inputs { + *input = new_idx[*input]; + } + if instr.operation.has_side_effects() { + new_idx[i] = instructions.len(); + instructions.push(instr); + continue; + } + match canonical.entry(instr.clone()) { + Entry::Occupied(e) => new_idx[i] = *e.get(), + Entry::Vacant(e) => { + e.insert(instructions.len()); + new_idx[i] = instructions.len(); + instructions.push(instr); + } + } + } + + let changed = instructions.len() < n; + program.instructions = instructions; + changed + } +} diff --git a/smite-ir/src/minimizers/dead_code.rs b/smite-ir/src/minimizers/dead_code.rs new file mode 100644 index 00000000..be2709fa --- /dev/null +++ b/smite-ir/src/minimizers/dead_code.rs @@ -0,0 +1,49 @@ +//! Dead-code elimination minimizer. + +use super::Minimizer; +use crate::Program; + +/// Removes unreferenced instructions and reindexes the remaining inputs. +/// +/// An instruction is removed when (a) its operation has no side effects +/// and (b) no later instruction references its output. The reverse +/// traversal lets a chain of dead instructions collapse, once we drop the +/// user of some load, that load's reference count falls to zero and the +/// load itself becomes eligible. +pub struct DeadCodeEliminator; + +impl Minimizer for DeadCodeEliminator { + fn minimize(&self, program: &mut Program) -> bool { + let n = program.instructions.len(); + let mut keep = vec![false; n]; + for idx in (0..n).rev() { + if !keep[idx] && !program.instructions[idx].operation.has_side_effects() { + continue; + } + keep[idx] = true; + for &input in &program.instructions[idx].inputs { + keep[input] = true; + } + } + + let mut remap = vec![0usize; n]; + let mut instructions = Vec::with_capacity(n); + for (old, mut instr) in std::mem::take(&mut program.instructions) + .into_iter() + .enumerate() + { + if !keep[old] { + continue; + } + for input in &mut instr.inputs { + *input = remap[*input]; + } + remap[old] = instructions.len(); + instructions.push(instr); + } + + let changed = instructions.len() < n; + program.instructions = instructions; + changed + } +} diff --git a/smite-ir/src/operation.rs b/smite-ir/src/operation.rs index bb9f3de5..c6abe4db 100644 --- a/smite-ir/src/operation.rs +++ b/smite-ir/src/operation.rs @@ -20,7 +20,7 @@ use super::VariableType; /// An IR operation. Each instruction in a program contains one operation plus /// input variable indices. -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum Operation { // -- Load: produce a variable from an embedded literal or the context -- /// Load a satoshi or millisatoshi amount. @@ -156,7 +156,7 @@ pub enum Operation { /// Each variant encodes to a script matching one of the formats required by /// BOLT 2 for the upfront shutdown TLV. `Empty` opts out of upfront shutdown /// entirely and is accepted regardless of feature negotiation. -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum ShutdownScriptVariant { /// Zero-length script. Opts out of upfront shutdown. Empty, @@ -327,7 +327,7 @@ impl fmt::Display for ShutdownScriptVariant { /// Additionally, the following bits can be added to any channel type: /// - `option_scid_alias` (bit 46) /// - `option_zeroconf` (bit 50) -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum ChannelTypeVariant { /// bit 12 StaticRemoteKey, @@ -470,7 +470,7 @@ impl fmt::Display for ChannelTypeVariant { } /// Fields that can be extracted from an `AcceptChannel` compound variable. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum AcceptChannelField { TemporaryChannelId, DustLimitSatoshis, @@ -745,6 +745,40 @@ impl Operation { } } + /// Returns `true` if this operation has I/O side effects and therefore + /// cannot be dropped by DCE or deduplicated by CSE. + #[must_use] + pub fn has_side_effects(&self) -> bool { + match self { + Self::SendMessage + | Self::RecvAcceptChannel + | Self::MineBlocks(_) + | Self::CreateFundingTransaction + | Self::BroadcastTransaction => true, + Self::LoadAmount(_) + | Self::LoadShortChannelId(_) + | Self::BuildChannelAnnouncement + | Self::LoadFeeratePerKw(_) + | Self::LoadForwardingFee(_) + | Self::LoadBlockHeight(_) + | Self::LoadTimestamp(_) + | Self::LoadU16(_) + | Self::LoadU8(_) + | Self::LoadBytes(_) + | Self::LoadFeatures(_) + | Self::LoadPrivateKey(_) + | Self::LoadChannelId(_) + | Self::LoadShutdownScript(_) + | Self::LoadChannelType(_) + | Self::LoadTargetPubkeyFromContext + | Self::LoadChainHashFromContext + | Self::DerivePoint + | Self::ExtractAcceptChannel(_) + | Self::BuildOpenChannel + | Self::BuildNodeAnnouncement { .. } => false, + } + } + /// Returns true if this operation has parameters that can be mutated /// by `OperationParamMutator`. #[must_use] diff --git a/smite-ir/src/tests.rs b/smite-ir/src/tests.rs index 8740a1e7..549d0bb6 100644 --- a/smite-ir/src/tests.rs +++ b/smite-ir/src/tests.rs @@ -7,6 +7,7 @@ use smite::bolt::{MAX_MESSAGE_SIZE, ShortChannelId}; use super::*; use generators::{NodeAnnouncementGenerator, OpenChannelGenerator}; +use minimizers::{CommonSubexpressionEliminator, DeadCodeEliminator, Minimizer}; use mutators::{InputSwapMutator, OperationParamMutator}; use operation::{AcceptChannelField, ChannelTypeVariant, ShutdownScriptVariant}; use program::ValidateError; @@ -1675,3 +1676,339 @@ fn input_swap_preserves_types() { } } } + +// -- DeadCodeEliminator tests -- + +#[test] +fn dead_code_removes_dead_instructions() { + // All three LoadAmount instructions are unreferenced; all three are dropped. + let mut program = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadAmount(1), + inputs: vec![], + }, + Instruction { + operation: Operation::LoadAmount(2), + inputs: vec![], + }, + Instruction { + operation: Operation::LoadAmount(3), + inputs: vec![], + }, + ], + }; + assert!(DeadCodeEliminator.minimize(&mut program)); + assert!( + program.instructions.is_empty(), + "all dead instructions should be removed" + ); + program.validate().expect("trimmed program should validate"); +} + +#[test] +fn dead_code_returns_false_on_empty_program() { + let mut program = Program { + instructions: vec![], + }; + assert!(!DeadCodeEliminator.minimize(&mut program)); + assert!(program.instructions.is_empty()); +} + +/// Build a program with a dead load appended after the generated program. +/// This gives the `DeadCodeEliminator` at least one candidate to try. +fn program_with_dead_load() -> Program { + let mut p = generate_open_channel_program(0); + p.instructions.push(Instruction { + operation: Operation::LoadAmount(42), + inputs: vec![], + }); + p +} + +#[test] +fn dead_code_keeps_send_message() { + let mut program = program_with_dead_load(); + DeadCodeEliminator.minimize(&mut program); + let has_send = program + .instructions + .iter() + .any(|i| matches!(i.operation, Operation::SendMessage)); + assert!(has_send, "DeadCodeEliminator must not remove SendMessage"); +} + +#[test] +fn dead_code_keeps_recv_accept_channel() { + let mut program = program_with_dead_load(); + DeadCodeEliminator.minimize(&mut program); + let has_recv = program + .instructions + .iter() + .any(|i| matches!(i.operation, Operation::RecvAcceptChannel)); + assert!( + has_recv, + "DeadCodeEliminator must not remove RecvAcceptChannel" + ); +} + +#[test] +fn dead_code_result_validates() { + let mut program = program_with_dead_load(); + DeadCodeEliminator.minimize(&mut program); + program.validate().expect("final program should validate"); +} + +#[test] +fn dead_code_reindexes_remaining_inputs() { + // Indexes 0 and 1 are dead loads; 2 is a referenced load; 3 references 2. + // After dropping 0 and 1, the surviving load shifts to index 0 and the + // DerivePoint must be rewritten to reference it. + let mut program = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadAmount(1), + inputs: vec![], + }, + Instruction { + operation: Operation::LoadAmount(2), + inputs: vec![], + }, + Instruction { + operation: Operation::LoadPrivateKey(key(1)), + inputs: vec![], + }, + Instruction { + operation: Operation::SendMessage, + inputs: vec![2], + }, + ], + }; + assert!(DeadCodeEliminator.minimize(&mut program)); + assert_eq!(program.instructions.len(), 2); + assert!(matches!( + program.instructions[0].operation, + Operation::LoadPrivateKey(_) + )); + assert!(matches!( + program.instructions[1].operation, + Operation::SendMessage + )); + assert_eq!(program.instructions[1].inputs, vec![0]); +} + +#[test] +fn dead_code_chains_collapse() { + // Two chains share a root LoadPrivateKey. One DerivePoint feeds an + // impure SendMessage (alive); the other is unreferenced (dead). DCE + // drops the dead DerivePoint, but the shared root must survive because + // the alive chain still references it. + // + // Note: this program is type-invalid (SendMessage expects Message, not + // Point), but the minimizer doesn't typecheck so it's fine for the test. + let mut program = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadPrivateKey(key(1)), + inputs: vec![], + }, + Instruction { + operation: Operation::DerivePoint, // alive + inputs: vec![0], + }, + Instruction { + operation: Operation::DerivePoint, // dead + inputs: vec![0], + }, + Instruction { + operation: Operation::SendMessage, + inputs: vec![1], + }, + ], + }; + let expected = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadPrivateKey(key(1)), + inputs: vec![], + }, + Instruction { + operation: Operation::DerivePoint, + inputs: vec![0], + }, + Instruction { + operation: Operation::SendMessage, + inputs: vec![1], + }, + ], + }; + assert!(DeadCodeEliminator.minimize(&mut program)); + assert_eq!(program, expected); +} + +#[test] +fn dead_code_idempotent() { + let mut once = program_with_dead_load(); + DeadCodeEliminator.minimize(&mut once); + let mut twice = once.clone(); + assert!( + !DeadCodeEliminator.minimize(&mut twice), + "second pass must report unchanged" + ); + assert_eq!(once, twice, "elimination is idempotent"); +} + +// -- CommonSubexpressionEliminator tests -- + +#[test] +fn cse_returns_false_on_empty_program() { + let mut program = Program { + instructions: vec![], + }; + assert!(!CommonSubexpressionEliminator.minimize(&mut program)); + assert!(program.instructions.is_empty()); +} + +#[test] +fn cse_rewires_references() { + // A downstream DerivePoint consumes the duplicate load. After CSE, its + // input must be rewired from the dropped duplicate (index 1) to the + // surviving canonical load (index 0). + let mut program = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadPrivateKey(key(7)), + inputs: vec![], + }, + Instruction { + operation: Operation::LoadPrivateKey(key(7)), // duplicate of index 0 + inputs: vec![], + }, + Instruction { + operation: Operation::DerivePoint, + inputs: vec![1], // must be rewired to 0 + }, + ], + }; + let expected = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadPrivateKey(key(7)), + inputs: vec![], + }, + Instruction { + operation: Operation::DerivePoint, + inputs: vec![0], + }, + ], + }; + assert!(CommonSubexpressionEliminator.minimize(&mut program)); + assert_eq!(program, expected); + program.validate().expect("program should still validate"); +} + +#[test] +fn cse_result_validates() { + let mut program = generate_open_channel_program(0); + CommonSubexpressionEliminator.minimize(&mut program); + program.validate().expect("merged program should validate"); +} + +#[test] +fn cse_idempotent() { + let mut once = generate_open_channel_program(0); + CommonSubexpressionEliminator.minimize(&mut once); + let mut twice = once.clone(); + assert!( + !CommonSubexpressionEliminator.minimize(&mut twice), + "second pass must report unchanged" + ); + assert_eq!(once, twice, "merging is idempotent"); +} + +#[test] +fn cse_merges_compute_ops_through_canonicalized_inputs() { + // Two LoadPrivateKey duplicates feed two DerivePoint instructions. + // CSE first merges the loads, which canonicalizes the DerivePoint + // inputs to the same index, which in turn lets CSE merge the + // DerivePoints themselves. + let mut program = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadPrivateKey(key(7)), + inputs: vec![], + }, + Instruction { + operation: Operation::DerivePoint, + inputs: vec![0], + }, + Instruction { + operation: Operation::LoadPrivateKey(key(7)), // duplicate of 0 + inputs: vec![], + }, + Instruction { + operation: Operation::DerivePoint, + inputs: vec![2], // canonicalizes to 0 -> matches index 1 + }, + ], + }; + program.validate().expect("input program should validate"); + assert!(CommonSubexpressionEliminator.minimize(&mut program)); + assert_eq!(program.instructions.len(), 2); + assert!(matches!( + program.instructions[0].operation, + Operation::LoadPrivateKey(_) + )); + assert!(matches!( + program.instructions[1].operation, + Operation::DerivePoint + )); + assert_eq!(program.instructions[1].inputs, vec![0]); +} + +#[test] +fn cse_does_not_merge_send_message() { + // SendMessage is not pure (network side-effect): two with the same + // input must both survive. The duplicate LoadBytes upstream should + // be merged, and both SendMessages remapped to the surviving load. + // + // Note: this program is type-invalid (SendMessage expects Message, not + // Bytes), but the minimizer doesn't typecheck so it's fine for the test. + let mut program = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadBytes(vec![0xab]), + inputs: vec![], + }, + Instruction { + operation: Operation::SendMessage, + inputs: vec![0], + }, + Instruction { + operation: Operation::LoadBytes(vec![0xab]), // duplicate of 0 + inputs: vec![], + }, + Instruction { + operation: Operation::SendMessage, + inputs: vec![2], // canonicalizes to 0 + }, + ], + }; + let expected = Program { + instructions: vec![ + Instruction { + operation: Operation::LoadBytes(vec![0xab]), + inputs: vec![], + }, + Instruction { + operation: Operation::SendMessage, + inputs: vec![0], + }, + Instruction { + operation: Operation::SendMessage, + inputs: vec![0], + }, + ], + }; + assert!(CommonSubexpressionEliminator.minimize(&mut program)); + assert_eq!(program, expected, "SendMessage must not be deduplicated"); +}