use crate::{
analysis::{CallGraphInfo, CfgInfo},
mir::{
BlockId, Function, FunctionId, Immediate, InstKind, InstructionMetadata, MirType, Module,
Terminator, Value, ValueId, utils::repair_reachability_phis,
},
pass::{FunctionPass, ModulePass},
};
use solar_data_structures::map::{FxHashMap, FxHashSet};
#[derive(Debug, PartialEq)]
struct CanonBlock {
insts: Vec<CanonInst>,
term_mnemonic: &'static str,
term_operands: Vec<CanonOperand>,
}
#[derive(Debug, PartialEq)]
struct CanonInst {
mnemonic: &'static str,
payload: CanonPayload,
operands: Vec<CanonOperand>,
result_ty: Option<MirType>,
metadata: InstructionMetadata,
}
#[derive(Debug, PartialEq)]
enum CanonPayload {
None,
FrameAddr(u64),
Call(FunctionId, usize),
}
#[derive(Debug, PartialEq)]
enum CanonOperand {
Local(usize),
Imm(Immediate),
Outside(ValueId),
}
#[derive(Debug, Default, Clone)]
pub struct CfgSimplifyStats {
pub blocks_merged: usize,
pub empty_blocks_eliminated: usize,
pub terminators_simplified: usize,
pub trivial_phis_simplified: usize,
pub terminal_blocks_deduplicated: usize,
pub dead_functions_eliminated: usize,
pub gas_saved: usize,
}
impl CfgSimplifyStats {
#[must_use]
pub fn total(&self) -> usize {
self.blocks_merged
+ self.empty_blocks_eliminated
+ self.terminators_simplified
+ self.trivial_phis_simplified
+ self.terminal_blocks_deduplicated
+ self.dead_functions_eliminated
}
pub fn combine(&mut self, other: &Self) {
self.blocks_merged += other.blocks_merged;
self.empty_blocks_eliminated += other.empty_blocks_eliminated;
self.terminators_simplified += other.terminators_simplified;
self.trivial_phis_simplified += other.trivial_phis_simplified;
self.terminal_blocks_deduplicated += other.terminal_blocks_deduplicated;
self.dead_functions_eliminated += other.dead_functions_eliminated;
self.gas_saved += other.gas_saved;
}
}
#[derive(Debug, Default)]
pub struct CfgSimplifier {
pub stats: CfgSimplifyStats,
}
pub struct CfgSimplifyPass;
impl FunctionPass for CfgSimplifyPass {
fn name(&self) -> &str {
"cfg-simplify"
}
fn run_on_function(&mut self, func: &mut Function) -> bool {
CfgSimplifier::new().run_to_fixpoint(func).total() != 0
}
}
impl CfgSimplifier {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn run(&mut self, func: &mut Function) -> usize {
self.stats = CfgSimplifyStats::default();
self.simplify_degenerate_terminators(func);
self.merge_blocks(func);
self.eliminate_empty_blocks(func);
self.deduplicate_terminal_blocks(func);
self.simplify_trivial_phis(func);
self.stats.total()
}
fn deduplicate_terminal_blocks(&mut self, func: &mut Function) {
let inst_results = func.inst_results();
let mut kept: Vec<(BlockId, CanonBlock)> = Vec::new();
let mut merges: Vec<(BlockId, BlockId)> = Vec::new();
for block_id in func.blocks.indices() {
if block_id == func.entry_block || func.blocks[block_id].predecessors.is_empty() {
continue;
}
let Some(canon) = Self::canonicalize_terminal_block(func, block_id, &inst_results)
else {
continue;
};
if let Some((keep, _)) = kept.iter().find(|(_, existing)| *existing == canon) {
merges.push((block_id, *keep));
} else {
kept.push((block_id, canon));
}
}
for (dup, keep) in merges {
let predecessors: Vec<_> = func.blocks[dup].predecessors.to_vec();
for pred in predecessors {
self.redirect_terminator(func, pred, dup, keep);
if !func.blocks[keep].predecessors.contains(&pred) {
func.blocks[keep].predecessors.push(pred);
}
}
func.blocks[dup].instructions.clear();
func.blocks[dup].terminator = Some(Terminator::Invalid);
func.blocks[dup].predecessors.clear();
self.stats.terminal_blocks_deduplicated += 1;
}
}
fn canonicalize_terminal_block(
func: &Function,
block_id: BlockId,
inst_results: &FxHashMap<crate::mir::InstId, ValueId>,
) -> Option<CanonBlock> {
let block = &func.blocks[block_id];
let term = block.terminator.as_ref()?;
if matches!(term, Terminator::Invalid) || !term.successors().is_empty() {
return None;
}
let mut local_defs: FxHashMap<ValueId, usize> = FxHashMap::default();
for (position, &inst_id) in block.instructions.iter().enumerate() {
if let Some(&result) = inst_results.get(&inst_id) {
local_defs.insert(result, position);
}
}
let canon_operand = |value: ValueId| {
if let Some(&position) = local_defs.get(&value) {
return CanonOperand::Local(position);
}
match &func.values[value] {
Value::Immediate(imm) => CanonOperand::Imm(imm.clone()),
_ => CanonOperand::Outside(value),
}
};
let mut insts = Vec::with_capacity(block.instructions.len());
for &inst_id in &block.instructions {
let inst = &func.instructions[inst_id];
let extra = match &inst.kind {
InstKind::Phi(_) => return None,
InstKind::InternalFrameAddr(offset) => CanonPayload::FrameAddr(*offset),
InstKind::InternalCall { function, returns, .. } => {
CanonPayload::Call(*function, *returns as usize)
}
_ => CanonPayload::None,
};
let mut metadata = inst.metadata.clone();
metadata.set_hir_expr(None);
metadata.set_source_span(None);
metadata.loop_depth = 0;
insts.push(CanonInst {
mnemonic: inst.kind.mnemonic(),
payload: extra,
operands: inst.kind.operands().into_iter().map(canon_operand).collect(),
result_ty: inst.result_ty,
metadata,
});
}
let term_operands = term.operands().into_iter().map(canon_operand).collect();
Some(CanonBlock { insts, term_mnemonic: term.mnemonic(), term_operands })
}
fn simplify_trivial_phis(&mut self, func: &mut Function) {
let mut candidates = Vec::new();
let mut raw = FxHashMap::default();
for block_id in func.blocks.indices() {
let same_block_phi_results = func.block_phi_results(block_id);
for &inst_id in &func.blocks[block_id].instructions {
let InstKind::Phi(incoming) = &func.instructions[inst_id].kind else {
continue;
};
let Some(phi_value) = func.inst_result_value(inst_id) else {
continue;
};
let Some(replacement) =
Self::trivial_phi_replacement(incoming, phi_value, &same_block_phi_results)
else {
continue;
};
candidates.push((inst_id, phi_value));
raw.insert(phi_value, replacement);
}
}
if raw.is_empty() {
return;
}
let mut replacements = FxHashMap::default();
let mut dead = FxHashSet::default();
for &(inst_id, phi_value) in &candidates {
let mut seen = FxHashSet::from_iter([phi_value]);
let mut target = raw[&phi_value];
let mut cyclic = false;
while let Some(&next) = raw.get(&target) {
if !seen.insert(target) {
cyclic = true;
break;
}
target = next;
}
if !cyclic {
replacements.insert(phi_value, target);
dead.insert(inst_id);
}
}
if replacements.is_empty() {
return;
}
func.replace_uses(&replacements);
for block in func.blocks.iter_mut() {
block.instructions.retain(|inst_id| !dead.contains(inst_id));
}
self.stats.trivial_phis_simplified += dead.len();
}
fn trivial_phi_replacement(
incoming: &[(BlockId, ValueId)],
phi_value: ValueId,
same_block_phi_results: &FxHashSet<ValueId>,
) -> Option<ValueId> {
let mut incoming_values = incoming.iter().map(|(_, value)| *value);
let first = incoming_values.find(|value| *value != phi_value)?;
if same_block_phi_results.contains(&first) {
return None;
}
incoming_values.all(|value| value == phi_value || value == first).then_some(first)
}
fn simplify_degenerate_terminators(&mut self, func: &mut Function) {
let block_ids: Vec<_> = func.blocks.indices().collect();
let mut changed = false;
for block_id in block_ids {
let Some(Terminator::Branch { then_block, else_block, .. }) =
func.blocks[block_id].terminator.as_ref()
else {
continue;
};
if then_block != else_block {
continue;
}
let target = *then_block;
func.blocks[block_id].terminator = Some(Terminator::Jump(target));
self.stats.terminators_simplified += 1;
self.stats.gas_saved += 10;
changed = true;
}
if changed {
repair_reachability_phis(func);
}
}
pub fn run_to_fixpoint(&mut self, func: &mut Function) -> CfgSimplifyStats {
let mut total_stats = CfgSimplifyStats::default();
loop {
let changed = self.run(func);
if changed == 0 {
break;
}
total_stats.combine(&self.stats);
}
total_stats
}
fn merge_blocks(&mut self, func: &mut Function) {
let mut merged = true;
while merged {
merged = false;
let block_ids: Vec<_> = func.blocks.indices().collect();
for block_id in block_ids {
if let Some(target) = self.can_merge(func, block_id) {
self.do_merge(func, block_id, target);
merged = true;
self.stats.blocks_merged += 1;
self.stats.gas_saved += 8;
break;
}
}
}
}
fn can_merge(&self, func: &Function, block_id: BlockId) -> Option<BlockId> {
let block = &func.blocks[block_id];
let Terminator::Jump(target) = block.terminator.as_ref()? else {
return None;
};
if *target == block_id {
return None;
}
if *target == func.entry_block {
return None;
}
let target_block = &func.blocks[*target];
if target_block.predecessors.len() != 1 {
return None;
}
if target_block.predecessors[0] != block_id {
return None;
}
for &inst_id in &target_block.instructions {
let InstKind::Phi(incoming) = &func.instructions[inst_id].kind else {
continue;
};
if !incoming.iter().any(|(pred, _)| *pred == block_id) {
return None;
}
}
Some(*target)
}
fn do_merge(&self, func: &mut Function, block_id: BlockId, target: BlockId) {
let phi_replacements = self.fold_target_phis_for_merge(func, block_id, target);
let target_instructions: Vec<_> = func.blocks[target]
.instructions
.iter()
.copied()
.filter(|&inst_id| !matches!(func.instructions[inst_id].kind, InstKind::Phi(_)))
.collect();
let target_terminator = func.blocks[target].terminator.take();
let target_successors =
target_terminator.as_ref().map(Terminator::successors).unwrap_or_default();
func.blocks[block_id].instructions.extend(target_instructions);
func.blocks[block_id].terminator = target_terminator;
for &succ in &target_successors {
self.redirect_target_phi_incoming(func, target, succ, &[block_id]);
let succ_block = &mut func.blocks[succ];
for pred in &mut succ_block.predecessors {
if *pred == target {
*pred = block_id;
}
}
}
func.blocks[target].instructions.clear();
func.blocks[target].terminator = Some(Terminator::Invalid);
func.blocks[target].predecessors.clear();
func.replace_uses(&phi_replacements);
}
fn fold_target_phis_for_merge(
&self,
func: &Function,
pred: BlockId,
target: BlockId,
) -> FxHashMap<ValueId, ValueId> {
let mut replacements = FxHashMap::default();
for &inst_id in &func.blocks[target].instructions {
let InstKind::Phi(incoming) = &func.instructions[inst_id].kind else {
continue;
};
let Some(phi_value) = func.inst_result_value(inst_id) else {
continue;
};
let Some((_, incoming_value)) =
incoming.iter().find(|(incoming_pred, _)| *incoming_pred == pred)
else {
continue;
};
replacements.insert(phi_value, *incoming_value);
}
replacements
}
fn eliminate_empty_blocks(&mut self, func: &mut Function) {
let mut eliminated = true;
while eliminated {
eliminated = false;
let block_ids: Vec<_> = func.blocks.indices().collect();
for block_id in block_ids {
if block_id == func.entry_block {
continue;
}
if self.is_empty_forwarder(func, block_id)
&& !self.is_loop_preheader_forwarder(func, block_id)
&& self.forwarder_elimination_preserves_phis(func, block_id)
{
self.eliminate_forwarder(func, block_id);
eliminated = true;
self.stats.empty_blocks_eliminated += 1;
self.stats.gas_saved += 8;
break;
}
}
}
}
fn is_empty_forwarder(&self, func: &Function, block_id: BlockId) -> bool {
let block = &func.blocks[block_id];
if !block.instructions.is_empty() {
return false;
}
matches!(&block.terminator, Some(Terminator::Jump(target)) if *target != block_id)
}
fn is_loop_preheader_forwarder(&self, func: &Function, block_id: BlockId) -> bool {
let Some(Terminator::Jump(target)) = func.blocks[block_id].terminator else {
return false;
};
if !matches!(
func.blocks[target].instructions.first(),
Some(&inst) if matches!(func.instructions[inst].kind, InstKind::Phi(_))
) {
return false;
}
let cfg = CfgInfo::new(func);
func.blocks[target]
.predecessors
.iter()
.copied()
.any(|pred| pred != block_id && cfg.dominators().dominates(target, pred))
}
fn forwarder_elimination_preserves_phis(&self, func: &Function, block_id: BlockId) -> bool {
let Some(Terminator::Jump(target)) = func.blocks[block_id].terminator else {
return false;
};
let predecessors = &func.blocks[block_id].predecessors;
for &inst_id in &func.blocks[target].instructions {
let InstKind::Phi(incoming) = &func.instructions[inst_id].kind else {
continue;
};
let Some(&(_, forwarded)) = incoming.iter().find(|(pred, _)| *pred == block_id) else {
continue;
};
for &pred in predecessors {
if incoming.iter().any(|&(other, value)| other == pred && value != forwarded) {
return false;
}
}
}
true
}
fn eliminate_forwarder(&self, func: &mut Function, block_id: BlockId) {
let target = match &func.blocks[block_id].terminator {
Some(Terminator::Jump(t)) => *t,
_ => return,
};
let predecessors: Vec<_> = func.blocks[block_id].predecessors.to_vec();
self.redirect_target_phi_incoming(func, block_id, target, &predecessors);
for pred_id in predecessors {
self.redirect_terminator(func, pred_id, block_id, target);
func.blocks[target].predecessors.push(pred_id);
}
func.blocks[target].predecessors.retain(|p| *p != block_id);
func.blocks[block_id].instructions.clear();
func.blocks[block_id].terminator = Some(Terminator::Invalid);
func.blocks[block_id].predecessors.clear();
}
fn redirect_target_phi_incoming(
&self,
func: &mut Function,
old_pred: BlockId,
target: BlockId,
new_preds: &[BlockId],
) {
for &inst_id in &func.blocks[target].instructions {
let InstKind::Phi(incoming) = &mut func.instructions[inst_id].kind else {
continue;
};
let mut rewritten: Vec<(BlockId, ValueId)> =
Vec::with_capacity(incoming.len() + new_preds.len());
for &(pred, value) in incoming.iter() {
if pred == old_pred {
rewritten.extend(new_preds.iter().map(|&new_pred| (new_pred, value)));
} else {
rewritten.push((pred, value));
}
}
let mut seen = Vec::with_capacity(rewritten.len());
rewritten.retain(|&(pred, _)| {
if seen.contains(&pred) {
false
} else {
seen.push(pred);
true
}
});
*incoming = rewritten;
}
}
fn redirect_terminator(
&self,
func: &mut Function,
block_id: BlockId,
old_target: BlockId,
new_target: BlockId,
) {
let block = &mut func.blocks[block_id];
match &mut block.terminator {
Some(Terminator::Jump(t)) if *t == old_target => {
*t = new_target;
}
Some(Terminator::Branch { then_block, else_block, .. }) => {
if *then_block == old_target {
*then_block = new_target;
}
if *else_block == old_target {
*else_block = new_target;
}
}
Some(Terminator::Switch { default, cases, .. }) => {
if *default == old_target {
*default = new_target;
}
for (_, target) in cases.iter_mut() {
if *target == old_target {
*target = new_target;
}
}
}
_ => {}
}
}
}
#[derive(Debug, Default)]
pub struct DeadFunctionEliminator {
pub stats: CfgSimplifyStats,
}
pub struct FunctionDcePass;
impl ModulePass for FunctionDcePass {
fn name(&self) -> &str {
"function-dce"
}
fn run(&mut self, module: &mut Module) -> bool {
DeadFunctionEliminator::new().run(module) != 0
}
}
impl DeadFunctionEliminator {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn run(&mut self, module: &mut Module) -> usize {
self.stats = CfgSimplifyStats::default();
let call_graph = CallGraphInfo::new(module);
let reachable = call_graph.reachable_from_entries();
if reachable.is_empty() {
return 0;
}
let dead_functions: Vec<FunctionId> =
module.functions.indices().filter(|id| !reachable.contains(id)).collect();
self.stats.dead_functions_eliminated = dead_functions.len();
for func_id in &dead_functions {
let func = &mut module.functions[*func_id];
func.blocks.clear();
func.instructions.clear();
func.values.clear();
}
self.stats.dead_functions_eliminated
}
}
pub fn simplify_cfg(func: &mut Function) -> CfgSimplifyStats {
let mut simplifier = CfgSimplifier::new();
simplifier.run_to_fixpoint(func)
}
pub fn simplify_module_cfg(module: &mut Module) -> CfgSimplifyStats {
let mut total_stats = CfgSimplifyStats::default();
for func_id in module.functions.indices() {
let func = &mut module.functions[func_id];
let stats = simplify_cfg(func);
total_stats.combine(&stats);
}
let mut dfe = DeadFunctionEliminator::new();
dfe.run(module);
total_stats.combine(&dfe.stats);
total_stats
}
#[cfg(test)]
mod tests {
use super::*;
use crate::mir::{FunctionBuilder, Instruction, MirType, Value};
use solar_interface::Ident;
use solar_sema::hir::Visibility;
#[test]
fn dead_function_elimination_keeps_internal_call_targets() {
let mut module = Module::new(Ident::DUMMY);
let live_helper = module.add_function(Function::new(Ident::DUMMY));
let dead_helper = module.add_function(Function::new(Ident::DUMMY));
let mut entry = Function::new(Ident::DUMMY);
entry.selector = Some([0, 0, 0, 1]);
entry.attributes.visibility = Visibility::Public;
{
let mut builder = FunctionBuilder::new(&mut entry);
let value = builder.internal_call(live_helper, Vec::new(), MirType::uint256(), 1);
builder.ret([value]);
}
let entry = module.add_function(entry);
{
let mut builder = FunctionBuilder::new(module.function_mut(live_helper));
let value = builder.imm_u64(1);
builder.ret([value]);
}
{
let mut builder = FunctionBuilder::new(module.function_mut(dead_helper));
let value = builder.imm_u64(2);
builder.ret([value]);
}
let mut dfe = DeadFunctionEliminator::new();
assert_eq!(dfe.run(&mut module), 1);
assert!(!module.function(entry).blocks.is_empty());
assert!(!module.function(live_helper).blocks.is_empty());
assert!(module.function(dead_helper).blocks.is_empty());
assert!(module.function(dead_helper).instructions.is_empty());
assert!(module.function(dead_helper).values.is_empty());
}
#[test]
fn empty_forwarder_rewrites_target_phi_incoming() {
let mut func = Function::new(Ident::DUMMY);
let forwarder;
let direct;
let target;
let value;
let other;
{
let mut builder = FunctionBuilder::new(&mut func);
forwarder = builder.create_block();
direct = builder.create_block();
target = builder.create_block();
value = builder.imm_u64(42);
let cond = builder.imm_bool(true);
builder.branch(cond, forwarder, direct);
builder.switch_to_block(direct);
let seven = builder.imm_u64(7);
other = builder.add(seven, value);
builder.jump(target);
builder.switch_to_block(forwarder);
builder.jump(target);
}
let phi_inst = func.alloc_inst(Instruction::new(
InstKind::Phi(vec![(forwarder, value), (direct, other)]),
Some(MirType::uint256()),
));
let phi_value = func.alloc_value(Value::Inst(phi_inst));
func.blocks[target].instructions.push(phi_inst);
func.blocks[target].terminator =
Some(Terminator::Return { values: vec![phi_value].into() });
let mut simplifier = CfgSimplifier::new();
simplifier.run_to_fixpoint(&mut func);
assert!(matches!(func.blocks[forwarder].terminator, Some(Terminator::Invalid)));
let phi_inst = func.blocks[target].instructions[0];
let InstKind::Phi(incoming) = &func.instructions[phi_inst].kind else {
panic!("expected phi");
};
assert_eq!(incoming.as_slice(), &[(func.entry_block, value), (direct, other)]);
}
#[test]
fn block_merge_rewrites_successor_phi_incoming() {
let mut func = Function::new(Ident::DUMMY);
let source;
let middle;
let other;
let exit;
let result;
let other_value;
{
let mut builder = FunctionBuilder::new(&mut func);
source = builder.create_block();
middle = builder.create_block();
other = builder.create_block();
exit = builder.create_block();
let cond = builder.imm_bool(true);
builder.branch(cond, source, other);
builder.switch_to_block(source);
builder.jump(middle);
builder.switch_to_block(middle);
let one = builder.imm_u64(1);
let two = builder.imm_u64(2);
result = builder.add(one, two);
builder.jump(exit);
builder.switch_to_block(other);
let three = builder.imm_u64(3);
let four = builder.imm_u64(4);
other_value = builder.add(three, four);
builder.jump(exit);
}
let phi_inst = func.alloc_inst(Instruction::new(
InstKind::Phi(vec![(middle, result), (other, other_value)]),
Some(MirType::uint256()),
));
let phi_value = func.alloc_value(Value::Inst(phi_inst));
func.blocks[exit].instructions.push(phi_inst);
func.blocks[exit].terminator = Some(Terminator::Return { values: vec![phi_value].into() });
let mut simplifier = CfgSimplifier::new();
simplifier.run_to_fixpoint(&mut func);
assert!(matches!(func.blocks[middle].terminator, Some(Terminator::Invalid)));
let phi_inst = func.blocks[exit].instructions[0];
let InstKind::Phi(incoming) = &func.instructions[phi_inst].kind else {
panic!("expected phi");
};
assert_eq!(incoming.as_slice(), &[(source, result), (other, other_value)]);
}
}