use super::pass_manager::ExecutionUnitPass;
use super::shared::{collect_all_used_registers, def_reg};
use crate::HashMap;
use crate::PassOptions;
use crate::ir::*;
use num_bigint::BigUint;
pub(in crate::optimizer) struct XorChainFoldingPass;
impl ExecutionUnitPass for XorChainFoldingPass {
fn name(&self) -> &'static str {
"xor_chain_folding"
}
fn run(&self, eu: &mut ExecutionUnit<RegionedAbsoluteAddr>, _options: &PassOptions) {
let mut any_changed = false;
let mut max_reg = eu.register_map.keys().map(|r| r.0).max().unwrap_or(0);
for block in eu.blocks.values_mut() {
if fold_xor_chains(&mut block.instructions, &mut eu.register_map, &mut max_reg) {
any_changed = true;
}
}
if !any_changed {
return;
}
let used = collect_all_used_registers(eu);
for block in eu.blocks.values_mut() {
block.instructions.retain(|inst| {
if let Some(d) = def_reg(inst) {
used.contains(&d)
|| matches!(inst, SIRInstruction::Store(..) | SIRInstruction::Commit(..))
} else {
true
}
});
}
}
}
fn fold_xor_chains(
instructions: &mut Vec<SIRInstruction<RegionedAbsoluteAddr>>,
register_map: &mut HashMap<RegisterId, RegisterType>,
next_reg: &mut usize,
) -> bool {
let mut defs: HashMap<RegisterId, SIRInstruction<RegionedAbsoluteAddr>> = HashMap::default();
for inst in instructions.iter() {
if let Some(d) = def_reg(inst) {
defs.insert(d, inst.clone());
}
}
let mut replacements: Vec<(usize, RegisterId, RegisterId, u64, usize)> = Vec::new();
for (idx, inst) in instructions.iter().enumerate() {
let SIRInstruction::Binary(dst, lhs, BinaryOp::Xor, rhs) = inst else {
continue;
};
let mut bits: Vec<usize> = Vec::new();
let mut source: Option<RegisterId> = None;
if collect_xor_bits(*lhs, &defs, &mut bits, &mut source)
&& collect_xor_bits(*rhs, &defs, &mut bits, &mut source)
&& bits.len() >= 3
{
if let Some(src) = source {
let src_width = register_map.get(&src).map(|t| t.width()).unwrap_or(64);
let mut mask: u64 = 0;
for &pos in &bits {
if pos < 64 {
mask |= 1u64 << pos;
}
}
if mask != 0 {
replacements.push((idx, *dst, src, mask, src_width));
}
}
}
}
if replacements.is_empty() {
return false;
}
for (idx, dst, src, mask, src_width) in replacements.into_iter().rev() {
let mask_width = src_width.min(64);
*next_reg += 1;
let mask_reg = RegisterId(*next_reg);
register_map.insert(
mask_reg,
RegisterType::Bit {
width: mask_width,
signed: false,
},
);
*next_reg += 1;
let masked_reg = RegisterId(*next_reg);
register_map.insert(
masked_reg,
RegisterType::Bit {
width: mask_width,
signed: false,
},
);
let mask_value = SIRValue {
payload: BigUint::from(mask),
mask: BigUint::ZERO,
};
instructions.insert(idx, SIRInstruction::Imm(mask_reg, mask_value));
instructions.insert(
idx + 1,
SIRInstruction::Binary(masked_reg, src, BinaryOp::And, mask_reg),
);
instructions[idx + 2] = SIRInstruction::Unary(dst, UnaryOp::Xor, masked_reg);
}
true
}
fn collect_xor_bits(
reg: RegisterId,
defs: &HashMap<RegisterId, SIRInstruction<RegionedAbsoluteAddr>>,
bits: &mut Vec<usize>,
source: &mut Option<RegisterId>,
) -> bool {
let Some(def) = defs.get(®) else {
return false;
};
match def {
SIRInstruction::Slice(_, src, offset, 1) => {
match source {
Some(s) if *s != *src => return false,
None => *source = Some(*src),
_ => {}
}
bits.push(*offset);
true
}
SIRInstruction::Load(_, addr, SIROffset::Static(offset), 1) => {
let _ = (addr, offset);
false
}
SIRInstruction::Binary(_, lhs, BinaryOp::Xor, rhs) => {
collect_xor_bits(*lhs, defs, bits, source) && collect_xor_bits(*rhs, defs, bits, source)
}
SIRInstruction::Unary(_, UnaryOp::Ident, src) => collect_xor_bits(*src, defs, bits, source),
_ => false,
}
}