use crate::mir::{BlockId, Function, InstId, InstKind, MirType, Value, ValueId};
use solar_data_structures::map::FxHashMap;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum CopySource {
Value(ValueId),
Temp(u32),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum CopyDest {
Value(ValueId),
Temp(u32),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ParallelCopy {
pub src: CopySource,
pub dst: CopyDest,
pub ty: MirType,
}
#[derive(Clone, Debug, Default)]
pub struct BlockCopies {
pub copies: Vec<ParallelCopy>,
}
#[derive(Debug)]
pub struct PhiEliminationResult {
pub block_copies: FxHashMap<BlockId, BlockCopies>,
pub phis_to_remove: Vec<(BlockId, usize)>,
}
pub struct PhiEliminator;
impl PhiEliminator {
#[must_use]
pub fn analyze(func: &Function) -> PhiEliminationResult {
let mut block_copies: FxHashMap<BlockId, BlockCopies> = FxHashMap::default();
let mut phis_to_remove = Vec::new();
for (block_id, block) in func.blocks.iter_enumerated() {
for (inst_idx, &inst_id) in block.instructions.iter().enumerate() {
let inst = func.instruction(inst_id);
if let InstKind::Phi(incoming) = &inst.kind {
let phi_dst = find_phi_dst(func, inst_id);
if let Some(dst) = phi_dst {
let ty = func.value(dst).ty();
for &(pred_block, src_val) in incoming {
block_copies.entry(pred_block).or_default().copies.push(ParallelCopy {
src: CopySource::Value(src_val),
dst: CopyDest::Value(dst),
ty,
});
}
phis_to_remove.push((block_id, inst_idx));
}
}
}
}
let mut temp_counter = 0u32;
for copies in block_copies.values_mut() {
sequentialize_copies(&mut copies.copies, &mut temp_counter);
}
PhiEliminationResult { block_copies, phis_to_remove }
}
}
#[must_use]
pub fn eliminate_phis(func: &Function) -> PhiEliminationResult {
PhiEliminator::analyze(func)
}
fn find_phi_dst(func: &Function, inst_id: InstId) -> Option<ValueId> {
for (val_id, val) in func.values.iter_enumerated() {
if let Value::Inst(def_inst) = val
&& *def_inst == inst_id
{
return Some(val_id);
}
}
None
}
fn src_value(src: &CopySource) -> Option<ValueId> {
match src {
CopySource::Value(v) => Some(*v),
CopySource::Temp(_) => None,
}
}
fn dst_value(dst: &CopyDest) -> Option<ValueId> {
match dst {
CopyDest::Value(v) => Some(*v),
CopyDest::Temp(_) => None,
}
}
fn sequentialize_copies(copies: &mut Vec<ParallelCopy>, temp_counter: &mut u32) {
if copies.len() <= 1 {
return;
}
let pending: Vec<ParallelCopy> = std::mem::take(copies);
let mut result: Vec<ParallelCopy> = Vec::with_capacity(pending.len() + 2);
let mut writes_to: FxHashMap<ValueId, usize> = FxHashMap::default();
for (i, copy) in pending.iter().enumerate() {
if let Some(dst) = dst_value(©.dst) {
writes_to.insert(dst, i);
}
}
let mut blocked_by: Vec<usize> = vec![0; pending.len()];
for (i, copy) in pending.iter().enumerate() {
if let Some(src) = src_value(©.src)
&& let Some(&writer_idx) = writes_to.get(&src)
&& writer_idx != i
{
blocked_by[writer_idx] += 1;
}
}
let mut emitted = vec![false; pending.len()];
loop {
let mut made_progress = false;
for i in 0..pending.len() {
if emitted[i] {
continue;
}
if blocked_by[i] == 0 {
result.push(pending[i].clone());
emitted[i] = true;
made_progress = true;
if let Some(src) = src_value(&pending[i].src)
&& let Some(&blocked_writer) = writes_to.get(&src)
&& blocked_writer != i
&& !emitted[blocked_writer]
{
blocked_by[blocked_writer] = blocked_by[blocked_writer].saturating_sub(1);
}
}
}
if !made_progress {
break_cycles(
&pending,
&mut emitted,
&mut blocked_by,
&writes_to,
&mut result,
temp_counter,
);
if !emitted.iter().all(|&e| e) {
continue;
}
break;
}
if emitted.iter().all(|&e| e) {
break;
}
}
*copies = result;
}
fn break_cycles(
pending: &[ParallelCopy],
emitted: &mut [bool],
blocked_by: &mut [usize],
writes_to: &FxHashMap<ValueId, usize>,
result: &mut Vec<ParallelCopy>,
temp_counter: &mut u32,
) {
let cycle_start = pending
.iter()
.enumerate()
.find(|(i, _)| !emitted[*i] && blocked_by[*i] > 0)
.map(|(i, _)| i);
let Some(start_idx) = cycle_start else {
for (i, copy) in pending.iter().enumerate() {
if !emitted[i] {
result.push(copy.clone());
emitted[i] = true;
}
}
return;
};
let mut cycle_indices = vec![start_idx];
let mut current = start_idx;
while let Some(src) = src_value(&pending[current].src) {
if let Some(&pred_idx) = writes_to.get(&src) {
if emitted[pred_idx] {
break;
}
if pred_idx == start_idx {
break;
}
if cycle_indices.contains(&pred_idx) {
break;
}
cycle_indices.push(pred_idx);
current = pred_idx;
} else {
break;
}
}
let break_idx = cycle_indices[0];
let break_copy = &pending[break_idx];
let temp_id = *temp_counter;
*temp_counter += 1;
result.push(ParallelCopy {
src: break_copy.src.clone(),
dst: CopyDest::Temp(temp_id),
ty: break_copy.ty,
});
if let Some(src) = src_value(&break_copy.src)
&& let Some(&blocked_writer) = writes_to.get(&src)
&& blocked_writer != break_idx
{
blocked_by[blocked_writer] = blocked_by[blocked_writer].saturating_sub(1);
}
for &idx in &cycle_indices[1..] {
if !emitted[idx] && blocked_by[idx] == 0 {
result.push(pending[idx].clone());
emitted[idx] = true;
if let Some(src) = src_value(&pending[idx].src)
&& let Some(&blocked_writer) = writes_to.get(&src)
&& !emitted[blocked_writer]
{
blocked_by[blocked_writer] = blocked_by[blocked_writer].saturating_sub(1);
}
}
}
result.push(ParallelCopy {
src: CopySource::Temp(temp_id),
dst: break_copy.dst.clone(),
ty: break_copy.ty,
});
emitted[break_idx] = true;
}
#[cfg(test)]
mod tests {
use super::*;
fn copy(src: usize, dst: usize) -> ParallelCopy {
ParallelCopy {
src: CopySource::Value(ValueId::from_usize(src)),
dst: CopyDest::Value(ValueId::from_usize(dst)),
ty: MirType::uint256(),
}
}
fn has_temp(copies: &[ParallelCopy]) -> bool {
copies.iter().any(|c| matches!(c.src, CopySource::Temp(_)))
|| copies.iter().any(|c| matches!(c.dst, CopyDest::Temp(_)))
}
#[test]
fn test_no_cycle() {
let mut copies = vec![copy(0, 1), copy(2, 3)];
let mut temp_counter = 0;
sequentialize_copies(&mut copies, &mut temp_counter);
assert_eq!(copies.len(), 2);
assert!(!has_temp(&copies));
}
#[test]
fn test_chain() {
let mut copies = vec![copy(1, 0), copy(0, 2)];
let mut temp_counter = 0;
sequentialize_copies(&mut copies, &mut temp_counter);
let write_to_a_idx =
copies.iter().position(|c| matches!(c.dst, CopyDest::Value(v) if v.index() == 0));
let read_from_a_idx =
copies.iter().position(|c| matches!(c.src, CopySource::Value(v) if v.index() == 0));
assert!(read_from_a_idx.unwrap() < write_to_a_idx.unwrap());
}
#[test]
fn test_simple_cycle() {
let mut copies = vec![copy(1, 0), copy(0, 1)];
let mut temp_counter = 0;
sequentialize_copies(&mut copies, &mut temp_counter);
assert!(copies.len() >= 3, "Cycle should introduce temporary copies");
assert!(has_temp(&copies), "Should use temporaries for cycles");
assert!(temp_counter >= 1, "Should allocate at least one temp");
}
#[test]
fn test_three_way_cycle() {
let mut copies = vec![copy(1, 0), copy(2, 1), copy(0, 2)];
let mut temp_counter = 0;
sequentialize_copies(&mut copies, &mut temp_counter);
assert!(copies.len() >= 4, "3-way cycle should introduce temporary copies");
assert!(has_temp(&copies), "Should use temporaries for cycles");
}
#[test]
fn test_independent_copies() {
let mut copies = vec![copy(10, 0), copy(11, 1), copy(12, 2)];
let mut temp_counter = 0;
sequentialize_copies(&mut copies, &mut temp_counter);
assert_eq!(copies.len(), 3);
assert!(!has_temp(&copies));
}
}