use crate::{
analysis::CfgInfo,
mir::{
BlockId, Function, Immediate, InstId, InstKind, MirType, Value, ValueId, utils as mir_utils,
},
pass::FunctionPass,
};
use solar_data_structures::map::{FxHashMap, FxHashSet};
const MAX_VN_SWEEPS: usize = 10;
const MAX_ROUNDS: usize = 4;
type ClassId = ValueId;
#[derive(Debug, Default)]
pub struct GlobalValueNumberer {
pub eliminated_count: usize,
}
pub struct GvnPass;
impl FunctionPass for GvnPass {
fn name(&self) -> &str {
"gvn"
}
fn run_on_function(&mut self, func: &mut Function) -> bool {
GlobalValueNumberer::new().run(func) != 0
}
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
struct ExprKey {
kind: ExprKind,
ty: MirType,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
enum ExprKind {
Add(ClassId, ClassId),
Sub(ClassId, ClassId),
Mul(ClassId, ClassId),
Div(ClassId, ClassId),
SDiv(ClassId, ClassId),
Mod(ClassId, ClassId),
SMod(ClassId, ClassId),
Exp(ClassId, ClassId),
AddMod(ClassId, ClassId, ClassId),
MulMod(ClassId, ClassId, ClassId),
And(ClassId, ClassId),
Or(ClassId, ClassId),
Xor(ClassId, ClassId),
Not(ClassId),
Shl(ClassId, ClassId),
Shr(ClassId, ClassId),
Sar(ClassId, ClassId),
Byte(ClassId, ClassId),
Lt(ClassId, ClassId),
SLt(ClassId, ClassId),
Eq(ClassId, ClassId),
IsZero(ClassId),
Select(ClassId, ClassId, ClassId),
SignExtend(ClassId, ClassId),
CalldataLoad(ClassId),
BlockHash(ClassId),
BlobHash(ClassId),
LoadImmutable(u32),
Phi(BlockId, Vec<(BlockId, ClassId)>),
}
struct ReplaceCtx<'a> {
vn: &'a [ClassId],
cfg: &'a CfgInfo,
inst_results: &'a FxHashMap<InstId, ValueId>,
replacements: &'a mut FxHashMap<ValueId, ValueId>,
dead: &'a mut FxHashSet<InstId>,
}
impl GlobalValueNumberer {
pub fn new() -> Self {
Self::default()
}
pub fn run(&mut self, func: &mut Function) -> usize {
self.eliminated_count = 0;
for _ in 0..MAX_ROUNDS {
if !self.run_round(func) {
break;
}
}
self.eliminated_count
}
fn run_round(&mut self, func: &mut Function) -> bool {
let cfg = CfgInfo::new(func);
let inst_results = func.inst_results();
let Some(vn) = Self::compute_value_numbers(func, cfg.rpo(), &inst_results) else {
return false;
};
let mut leaders: FxHashMap<ClassId, ValueId> = FxHashMap::default();
for (value_id, value) in func.values.iter_enumerated() {
if !matches!(value, Value::Inst(_)) && vn[value_id.index()] == value_id {
leaders.insert(value_id, value_id);
}
}
let mut replacements = FxHashMap::default();
let mut dead = FxHashSet::default();
let mut ctx = ReplaceCtx {
vn: &vn,
cfg: &cfg,
inst_results: &inst_results,
replacements: &mut replacements,
dead: &mut dead,
};
self.replace_in_block(func, func.entry_block, &mut leaders, &mut ctx);
if replacements.is_empty() {
return false;
}
Self::apply_replacements_to_all_blocks(func, &replacements);
for block in func.blocks.iter_mut() {
block.instructions.retain(|id| !dead.contains(id));
}
true
}
fn compute_value_numbers(
func: &Function,
rpo: &[BlockId],
inst_results: &FxHashMap<InstId, ValueId>,
) -> Option<Vec<ClassId>> {
let mut vn: Vec<ClassId> = func.values.indices().collect();
let mut immediate_reps: FxHashMap<Immediate, ValueId> = FxHashMap::default();
let mut arg_reps: FxHashMap<u32, ValueId> = FxHashMap::default();
for (value_id, value) in func.values.iter_enumerated() {
match value {
Value::Immediate(imm) => {
vn[value_id.index()] = *immediate_reps.entry(imm.clone()).or_insert(value_id);
}
Value::Arg { index, .. } => {
vn[value_id.index()] = *arg_reps.entry(*index).or_insert(value_id);
}
Value::Inst(_) | Value::Undef(_) | Value::Error(_) => {}
}
}
for _ in 0..MAX_VN_SWEEPS {
let mut table: FxHashMap<ExprKey, ClassId> = FxHashMap::default();
let mut changed = false;
for &block_id in rpo {
for &inst_id in &func.blocks[block_id].instructions {
let Some(&result) = inst_results.get(&inst_id) else { continue };
let inst = &func.instructions[inst_id];
let Some(ty) = inst.result_ty else { continue };
let Some(class) =
Self::instruction_class(block_id, &inst.kind, ty, result, &vn, &mut table)
else {
continue;
};
if vn[result.index()] != class {
vn[result.index()] = class;
changed = true;
}
}
}
if !changed {
return Some(vn);
}
}
None
}
fn instruction_class(
block_id: BlockId,
kind: &InstKind,
ty: MirType,
result: ValueId,
vn: &[ClassId],
table: &mut FxHashMap<ExprKey, ClassId>,
) -> Option<ClassId> {
if let InstKind::Phi(incoming) = kind {
let Some((&(_, first), rest)) = incoming.split_first() else { return Some(result) };
if rest.iter().all(|&(_, value)| vn[value.index()] == vn[first.index()]) {
return Some(vn[first.index()]);
}
let Some(incoming) = Self::phi_key_incoming(incoming, vn) else { return Some(result) };
let key = ExprKey { kind: ExprKind::Phi(block_id, incoming), ty };
return Some(*table.entry(key).or_insert(result));
}
let kind = Self::expr_kind(kind, vn)?;
Some(*table.entry(ExprKey { kind, ty }).or_insert(result))
}
fn phi_key_incoming(
incoming: &[(BlockId, ValueId)],
vn: &[ClassId],
) -> Option<Vec<(BlockId, ClassId)>> {
let mut entries: Vec<(BlockId, ClassId)> =
incoming.iter().map(|&(pred, value)| (pred, vn[value.index()])).collect();
entries.sort_by_key(|&(pred, class)| (pred.index(), class.index()));
entries.dedup();
if entries.windows(2).any(|pair| pair[0].0 == pair[1].0) {
return None;
}
Some(entries)
}
fn expr_kind(kind: &InstKind, vn: &[ClassId]) -> Option<ExprKind> {
let class = |value: ValueId| vn[value.index()];
let sorted = |a: ValueId, b: ValueId| {
let (a, b) = (class(a), class(b));
if b.index() < a.index() { (b, a) } else { (a, b) }
};
Some(match *kind {
InstKind::Add(a, b) => {
let (a, b) = sorted(a, b);
ExprKind::Add(a, b)
}
InstKind::Mul(a, b) => {
let (a, b) = sorted(a, b);
ExprKind::Mul(a, b)
}
InstKind::And(a, b) => {
let (a, b) = sorted(a, b);
ExprKind::And(a, b)
}
InstKind::Or(a, b) => {
let (a, b) = sorted(a, b);
ExprKind::Or(a, b)
}
InstKind::Xor(a, b) => {
let (a, b) = sorted(a, b);
ExprKind::Xor(a, b)
}
InstKind::Eq(a, b) => {
let (a, b) = sorted(a, b);
ExprKind::Eq(a, b)
}
InstKind::AddMod(a, b, n) => {
let (a, b) = sorted(a, b);
ExprKind::AddMod(a, b, class(n))
}
InstKind::MulMod(a, b, n) => {
let (a, b) = sorted(a, b);
ExprKind::MulMod(a, b, class(n))
}
InstKind::Sub(a, b) => ExprKind::Sub(class(a), class(b)),
InstKind::Div(a, b) => ExprKind::Div(class(a), class(b)),
InstKind::SDiv(a, b) => ExprKind::SDiv(class(a), class(b)),
InstKind::Mod(a, b) => ExprKind::Mod(class(a), class(b)),
InstKind::SMod(a, b) => ExprKind::SMod(class(a), class(b)),
InstKind::Exp(a, b) => ExprKind::Exp(class(a), class(b)),
InstKind::Shl(a, b) => ExprKind::Shl(class(a), class(b)),
InstKind::Shr(a, b) => ExprKind::Shr(class(a), class(b)),
InstKind::Sar(a, b) => ExprKind::Sar(class(a), class(b)),
InstKind::Byte(a, b) => ExprKind::Byte(class(a), class(b)),
InstKind::SignExtend(a, b) => ExprKind::SignExtend(class(a), class(b)),
InstKind::Lt(a, b) => ExprKind::Lt(class(a), class(b)),
InstKind::Gt(a, b) => ExprKind::Lt(class(b), class(a)),
InstKind::SLt(a, b) => ExprKind::SLt(class(a), class(b)),
InstKind::SGt(a, b) => ExprKind::SLt(class(b), class(a)),
InstKind::IsZero(a) => ExprKind::IsZero(class(a)),
InstKind::Not(a) => ExprKind::Not(class(a)),
InstKind::Select(condition, then_value, else_value) => {
ExprKind::Select(class(condition), class(then_value), class(else_value))
}
InstKind::CalldataLoad(a) => ExprKind::CalldataLoad(class(a)),
InstKind::BlockHash(a) => ExprKind::BlockHash(class(a)),
InstKind::BlobHash(a) => ExprKind::BlobHash(class(a)),
InstKind::LoadImmutable(offset) => ExprKind::LoadImmutable(offset),
_ => return None,
})
}
fn replace_in_block(
&mut self,
func: &Function,
block_id: BlockId,
leaders: &mut FxHashMap<ClassId, ValueId>,
ctx: &mut ReplaceCtx<'_>,
) {
for &inst_id in &func.blocks[block_id].instructions {
let Some(&result) = ctx.inst_results.get(&inst_id) else { continue };
let kind = &func.instructions[inst_id].kind;
if !matches!(kind, InstKind::Phi(_)) && Self::expr_kind(kind, ctx.vn).is_none() {
continue;
}
let class = ctx.vn[result.index()];
if let Some(&leader) = leaders.get(&class) {
if leader != result {
ctx.replacements.insert(result, leader);
ctx.dead.insert(inst_id);
self.eliminated_count += 1;
}
} else {
leaders.insert(class, result);
}
}
for &child in ctx.cfg.dominators().children(block_id) {
let mut child_leaders = leaders.clone();
self.replace_in_block(func, child, &mut child_leaders, ctx);
}
}
fn apply_replacements_to_all_blocks(
func: &mut Function,
replacements: &FxHashMap<ValueId, ValueId>,
) {
let block_ids: Vec<_> = func.blocks.indices().collect();
for block_id in block_ids {
Self::apply_replacements(func, block_id, replacements);
}
}
fn apply_replacements(
func: &mut Function,
block_id: BlockId,
replacements: &FxHashMap<ValueId, ValueId>,
) {
let inst_ids: Vec<InstId> = func.blocks[block_id].instructions.clone();
for inst_id in inst_ids {
let inst = &mut func.instructions[inst_id];
if mir_utils::replace_inst_uses_canonicalized(&mut inst.kind, replacements) != 0 {
if mir_utils::is_memory_inst(&inst.kind) {
inst.metadata.set_memory_region(None);
}
if matches!(
inst.kind,
InstKind::SLoad(_)
| InstKind::SStore(_, _)
| InstKind::TLoad(_)
| InstKind::TStore(_, _)
) {
inst.metadata.set_storage_alias(None);
}
}
}
if let Some(term) = &mut func.blocks[block_id].terminator {
mir_utils::replace_terminator_uses_canonicalized(term, replacements);
}
}
}