From e6691d035b6948452d6073317451b3f1ff0a24f8 Mon Sep 17 00:00:00 2001 From: Henry Date: Sat, 25 Apr 2026 16:47:54 +0200 Subject: fix: remove unsafe rewrites Signed-off-by: Henry --- crates/parser/src/lib.rs | 13 ------ crates/parser/src/macros.rs | 17 ++------ crates/parser/src/optimize.rs | 98 ------------------------------------------- crates/parser/src/visit.rs | 54 ++++++++++++++++++++---- 4 files changed, 50 insertions(+), 132 deletions(-) (limited to 'crates/parser') diff --git a/crates/parser/src/lib.rs b/crates/parser/src/lib.rs index 949afa2..684e79b 100644 --- a/crates/parser/src/lib.rs +++ b/crates/parser/src/lib.rs @@ -51,8 +51,6 @@ pub struct ParserOptions { pub optimize_rewrite: bool, /// Whether to remove `Nop` and `MergeBarrier` instructions after rewriting. pub optimize_remove_nop: bool, - /// Whether to invert conditional branches over an unconditional jump. - pub optimize_branch_inversion: bool, } impl Default for ParserOptions { @@ -61,7 +59,6 @@ impl Default for ParserOptions { optimize_local_memory_allocation: true, optimize_rewrite: true, optimize_remove_nop: true, - optimize_branch_inversion: false, } } } @@ -100,16 +97,6 @@ impl ParserOptions { self.optimize_remove_nop } - /// Enable or disable the optimization that inverts conditional branches over an unconditional jump. - pub const fn with_branch_inversion_optimization(mut self, enabled: bool) -> Self { - self.optimize_branch_inversion = enabled; - self - } - - /// Returns whether conditional branch inversion optimization is enabled. - pub const fn optimize_branch_inversion(&self) -> bool { - self.optimize_branch_inversion - } } /// A WebAssembly parser diff --git a/crates/parser/src/macros.rs b/crates/parser/src/macros.rs index a6eab61..ebab704 100644 --- a/crates/parser/src/macros.rs +++ b/crates/parser/src/macros.rs @@ -222,25 +222,16 @@ pub(crate) mod optimize { get = $get:ident, tee = $tee:ident, set = $set:ident, - copy = $copy:ident, binop_local_local_tee = $lltee:ident, binop_local_local_set = $llset:ident, binop_local_const_tee = $lctee:ident, binop_local_const_set = $lcset:ident $(, load_local_tee = $loadtee:ident, load_local_set = $loadset:ident)? ) => { - fn $name(instrs: &mut [Instruction], read: usize, instr: Instruction) -> Option<(Instruction, u16)> { + fn $name(instr: Instruction) -> Option<(Instruction, u16)> { Some(match instr { Instruction::$get(local) => (Instruction::Nop, local), - Instruction::$tee(local) => { - let replacement = if let Some([(prev_idx, Instruction::$get(src))]) = previous_non_nop::<1>(instrs, read) { - instrs[prev_idx] = Instruction::Nop; - if src == local { Instruction::Nop } else { Instruction::$copy(src, local) } - } else { - Instruction::$set(local) - }; - (replacement, local) - } + Instruction::$tee(local) => (Instruction::$set(local), local), Instruction::$lltee(op, a, b, local) => (Instruction::$llset(op, a, b, local), local), Instruction::$lctee(op, src, c, local) => (Instruction::$lcset(op, src, c, local), local), $(Instruction::$loadtee(memarg, addr, local) => (Instruction::$loadset(memarg, addr, local), local.into()),)? @@ -261,10 +252,10 @@ pub(crate) mod optimize { ) => {{ if let Some([(lhs_idx, lhs_src), (rhs_idx, rhs_src), (op_idx, raw_op)]) = previous_non_nop::<3>($instrs, $read) - && let Some((lhs_instr, lhs)) = $source($instrs, lhs_idx, lhs_src) + && let Some((lhs_instr, lhs)) = $source(lhs_src) && let Some(op) = $op(raw_op) { - if let Some((rhs_instr, rhs)) = $source($instrs, rhs_idx, rhs_src) { + if let Some((rhs_instr, rhs)) = $source(rhs_src) { $instrs[lhs_idx] = lhs_instr; $instrs[rhs_idx] = rhs_instr; $instrs[op_idx] = Instruction::Nop; diff --git a/crates/parser/src/optimize.rs b/crates/parser/src/optimize.rs index baaaf66..19abd7c 100644 --- a/crates/parser/src/optimize.rs +++ b/crates/parser/src/optimize.rs @@ -22,7 +22,6 @@ pub(crate) fn optimize_instructions( self_func_addr, imported_memory_count, track_local_memory_usage, - options.optimize_branch_inversion(), ) } else { track_local_memory_usage @@ -40,7 +39,6 @@ fn rewrite( self_func_addr: u32, imported_memory_count: u32, track_local_memory_usage: bool, - optimize_branch_inversion: bool, ) -> bool { use Instruction::*; let mut uses_local_memory = false; @@ -416,9 +414,6 @@ fn rewrite( (0, CmpOp::Ne) => JumpIfNonZero64(target), (imm, op) => JumpCmpStackConst64 { target_ip: target, imm, op }, }); - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); if let JumpIfZero(current) = &mut instrs[i] { *current = target; @@ -478,9 +473,6 @@ fn rewrite( (0, CmpOp::Ne) => JumpIfNonZero64(target), (imm, op) => JumpCmpStackConst64 { target_ip: target, imm, op }, }); - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); if let JumpIfNonZero(current) = &mut instrs[i] { *current = target; @@ -489,9 +481,6 @@ fn rewrite( JumpIfZero32(ip) => { let target = resolve_jump_target(instrs, ip); rewrite!(instrs, i, [LocalGet32(local)] => JumpIfLocalZero32 { target_ip: target, local }); - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); if let JumpIfZero32(current) = &mut instrs[i] { *current = target; @@ -500,9 +489,6 @@ fn rewrite( JumpIfNonZero32(ip) => { let target = resolve_jump_target(instrs, ip); rewrite!(instrs, i, [LocalGet32(local)] => JumpIfLocalNonZero32 { target_ip: target, local }); - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); if let JumpIfNonZero32(current) = &mut instrs[i] { *current = target; @@ -511,9 +497,6 @@ fn rewrite( JumpIfZero64(ip) => { let target = resolve_jump_target(instrs, ip); rewrite!(instrs, i, [LocalGet64(local)] => JumpIfLocalZero64 { target_ip: target, local }); - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); if let JumpIfZero64(current) = &mut instrs[i] { *current = target; @@ -522,9 +505,6 @@ fn rewrite( JumpIfNonZero64(ip) => { let target = resolve_jump_target(instrs, ip); rewrite!(instrs, i, [LocalGet64(local)] => JumpIfLocalNonZero64 { target_ip: target, local }); - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); if let JumpIfNonZero64(current) = &mut instrs[i] { *current = target; @@ -536,9 +516,6 @@ fn rewrite( CmpOp::Ne => instrs[i] = JumpIfNonZero32(target_ip), _ => {} } - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); } JumpCmpStackConst64 { target_ip, imm: 0, op } => { @@ -547,9 +524,6 @@ fn rewrite( CmpOp::Ne => instrs[i] = JumpIfNonZero64(target_ip), _ => {} } - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); } JumpCmpLocalConst32 { target_ip, local, imm: 0, op } => { @@ -558,9 +532,6 @@ fn rewrite( CmpOp::Ne => instrs[i] = JumpIfLocalNonZero32 { target_ip, local }, _ => {} } - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); } JumpCmpLocalConst64 { target_ip, local, imm: 0, op } => { @@ -569,9 +540,6 @@ fn rewrite( CmpOp::Ne => instrs[i] = JumpIfLocalNonZero64 { target_ip, local }, _ => {} } - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); } JumpCmpStackConst32 { .. } @@ -584,9 +552,6 @@ fn rewrite( | JumpIfLocalNonZero32 { .. } | JumpIfLocalZero64 { .. } | JumpIfLocalNonZero64 { .. } => { - if optimize_branch_inversion { - invert_conditional_over_jump(instrs, i); - } canonicalize_jump_like(instrs, i); } _ => {} @@ -710,7 +675,6 @@ define_local_source_resolver!( get = LocalGet32, tee = LocalTee32, set = LocalSet32, - copy = LocalCopy32, binop_local_local_tee = BinOpLocalLocalTee32, binop_local_local_set = BinOpLocalLocalSet32, binop_local_const_tee = BinOpLocalConstTee32, @@ -724,7 +688,6 @@ define_local_source_resolver!( get = LocalGet64, tee = LocalTee64, set = LocalSet64, - copy = LocalCopy64, binop_local_local_tee = BinOpLocalLocalTee64, binop_local_local_set = BinOpLocalLocalSet64, binop_local_const_tee = BinOpLocalConstTee64, @@ -736,7 +699,6 @@ define_local_source_resolver!( get = LocalGet128, tee = LocalTee128, set = LocalSet128, - copy = LocalCopy128, binop_local_local_tee = BinOpLocalLocalTee128, binop_local_local_set = BinOpLocalLocalSet128, binop_local_const_tee = BinOpLocalConstTee128, @@ -881,30 +843,6 @@ fn set_jump_target(instr: &mut Instruction, target: u32) { } } -fn invert_jump(instr: Instruction, target: u32) -> Option { - Some(match instr { - Instruction::JumpCmpStackConst32 { imm, op, .. } => { - Instruction::JumpCmpStackConst32 { target_ip: target, imm, op: inverse_cmp_op(op) } - } - Instruction::JumpCmpStackConst64 { imm, op, .. } => { - Instruction::JumpCmpStackConst64 { target_ip: target, imm, op: inverse_cmp_op(op) } - } - Instruction::JumpCmpLocalConst32 { local, imm, op, .. } => { - Instruction::JumpCmpLocalConst32 { target_ip: target, local, imm, op: inverse_cmp_op(op) } - } - Instruction::JumpCmpLocalConst64 { local, imm, op, .. } => { - Instruction::JumpCmpLocalConst64 { target_ip: target, local, imm, op: inverse_cmp_op(op) } - } - Instruction::JumpCmpLocalLocal32 { left, right, op, .. } => { - Instruction::JumpCmpLocalLocal32 { target_ip: target, left, right, op: inverse_cmp_op(op) } - } - Instruction::JumpCmpLocalLocal64 { left, right, op, .. } => { - Instruction::JumpCmpLocalLocal64 { target_ip: target, left, right, op: inverse_cmp_op(op) } - } - _ => return None, - }) -} - fn canonicalize_jump_like(instrs: &mut [Instruction], idx: usize) { let Some(target) = jump_target(instrs[idx]) else { return; @@ -918,42 +856,6 @@ fn canonicalize_jump_like(instrs: &mut [Instruction], idx: usize) { } } -fn invert_conditional_over_jump(instrs: &mut [Instruction], idx: usize) { - let Some(target) = jump_target(instrs[idx]) else { - return; - }; - if matches!(instrs[idx], Instruction::Jump(_)) { - return; - } - - let target_idx = next_non_nop(instrs, target as usize); - if target_idx >= instrs.len() || target_idx <= idx + 1 { - return; - } - - let Some(jump_idx) = ((idx + 1)..target_idx) - .rev() - .find(|&candidate| !matches!(instrs[candidate], Instruction::Nop | Instruction::MergeBarrier)) - else { - return; - }; - - let Instruction::Jump(exit_target) = instrs[jump_idx] else { - return; - }; - if next_non_nop(instrs, jump_idx + 1) != target_idx { - return; - } - - let exit_target = resolve_jump_target(instrs, exit_target); - let Some(inverted) = invert_jump(instrs[idx], exit_target) else { - return; - }; - - instrs[idx] = inverted; - instrs[jump_idx] = Instruction::Nop; -} - fn remove_nop(instructions: &mut Vec, function_data: &mut WasmFunctionData) { let old_len = instructions.len(); if old_len == 0 { diff --git a/crates/parser/src/visit.rs b/crates/parser/src/visit.rs index 9a094ed..de0c8fb 100644 --- a/crates/parser/src/visit.rs +++ b/crates/parser/src/visit.rs @@ -218,14 +218,52 @@ impl<'a, R: WasmModuleResources> wasmparser::VisitOperator<'a> for FunctionBuild fn visit_local_tee(&mut self, idx: u32) -> Self::Output { let resolved_idx = self.local_addr_map[idx as usize]; if let Some(Some(t)) = self.validator.get_operand_type(0) { - self.instructions.push(match t { - wasmparser::ValType::I32 => Instruction::LocalTee32(resolved_idx), - wasmparser::ValType::F32 => Instruction::LocalTee32(resolved_idx), - wasmparser::ValType::I64 => Instruction::LocalTee64(resolved_idx), - wasmparser::ValType::F64 => Instruction::LocalTee64(resolved_idx), - wasmparser::ValType::V128 => Instruction::LocalTee128(resolved_idx), - wasmparser::ValType::Ref(_) => Instruction::LocalTee32(resolved_idx), - }) + let last = self.instructions.last(); + let src = match t { + wasmparser::ValType::I32 | wasmparser::ValType::F32 => { + if let Some(Instruction::LocalGet32(src)) = last { Some(*src) } else { None } + } + wasmparser::ValType::I64 | wasmparser::ValType::F64 => { + if let Some(Instruction::LocalGet64(src)) = last { Some(*src) } else { None } + } + wasmparser::ValType::V128 => { + if let Some(Instruction::LocalGet128(src)) = last { Some(*src) } else { None } + } + wasmparser::ValType::Ref(_) => { + if let Some(Instruction::LocalGet32(src)) = last { Some(*src) } else { None } + } + }; + + if let Some(src) = src { + self.instructions.pop(); + match t { + wasmparser::ValType::I32 | wasmparser::ValType::F32 => { + self.instructions.push(Instruction::LocalCopy32(src, resolved_idx)); + self.instructions.push(Instruction::LocalGet32(resolved_idx)); + } + wasmparser::ValType::I64 | wasmparser::ValType::F64 => { + self.instructions.push(Instruction::LocalCopy64(src, resolved_idx)); + self.instructions.push(Instruction::LocalGet64(resolved_idx)); + } + wasmparser::ValType::V128 => { + self.instructions.push(Instruction::LocalCopy128(src, resolved_idx)); + self.instructions.push(Instruction::LocalGet128(resolved_idx)); + } + wasmparser::ValType::Ref(_) => { + self.instructions.push(Instruction::LocalCopy32(src, resolved_idx)); + self.instructions.push(Instruction::LocalGet32(resolved_idx)); + } + } + } else { + self.instructions.push(match t { + wasmparser::ValType::I32 => Instruction::LocalTee32(resolved_idx), + wasmparser::ValType::F32 => Instruction::LocalTee32(resolved_idx), + wasmparser::ValType::I64 => Instruction::LocalTee64(resolved_idx), + wasmparser::ValType::F64 => Instruction::LocalTee64(resolved_idx), + wasmparser::ValType::V128 => Instruction::LocalTee128(resolved_idx), + wasmparser::ValType::Ref(_) => Instruction::LocalTee32(resolved_idx), + }) + } } } -- cgit v1.3.1