use std::collections::{BTreeMap, BTreeSet};
use crate::analysis::{SsaFunction, SsaVarId, VariableOrigin};
fn var_key(var: SsaVarId) -> u32 {
u32::try_from(var.index()).unwrap_or(u32::MAX)
}
struct ValueClasses {
parent: Vec<u32>,
}
impl ValueClasses {
fn new(capacity: usize) -> Self {
let parent = (0..u32::try_from(capacity).unwrap_or(u32::MAX)).collect();
Self { parent }
}
fn find(&mut self, var: u32) -> u32 {
let mut root = var;
while let Some(&parent) = self.parent.get(root as usize) {
if parent == root {
break;
}
root = parent;
}
let mut current = var;
while let Some(&parent) = self.parent.get(current as usize) {
if parent == root {
break;
}
if let Some(slot) = self.parent.get_mut(current as usize) {
*slot = root;
}
current = parent;
}
root
}
fn union(&mut self, a: u32, b: u32) {
let (ra, rb) = (self.find(a), self.find(b));
if ra != rb {
if let Some(slot) = self.parent.get_mut(rb as usize) {
*slot = ra;
}
}
}
}
pub fn promote_cross_block_values(ssa: &mut SsaFunction) -> usize {
let capacity = ssa.var_id_capacity();
if capacity == 0 {
return 0;
}
let mut def_block: BTreeMap<SsaVarId, usize> = BTreeMap::new();
for block in ssa.blocks() {
let id = block.id();
for phi in block.phi_nodes() {
def_block.insert(phi.result(), id);
}
for instr in block.instructions() {
if let Some(dest) = instr.def() {
def_block.insert(dest, id);
}
}
}
let mut crosses: BTreeSet<SsaVarId> = BTreeSet::new();
let mut classes = ValueClasses::new(capacity);
for block in ssa.blocks() {
let id = block.id();
for phi in block.phi_nodes() {
let result = phi.result();
crosses.insert(result);
for operand in phi.operands() {
let value = operand.value();
crosses.insert(value);
classes.union(var_key(result), var_key(value));
}
}
for instr in block.instructions() {
for used in instr.uses() {
if def_block.get(&used).is_some_and(|&def| def != id) {
crosses.insert(used);
}
}
}
}
if crosses.is_empty() {
return 0;
}
let mut members: BTreeMap<u32, Vec<SsaVarId>> = BTreeMap::new();
let mut anchored: BTreeSet<u32> = BTreeSet::new();
for &var in &crosses {
let root = classes.find(var_key(var));
members.entry(root).or_default().push(var);
if let Some(variable) = ssa.variable(var) {
if matches!(
variable.origin(),
VariableOrigin::Argument(_) | VariableOrigin::Local(_)
) {
anchored.insert(root);
}
}
}
let mut next_local = u16::try_from(ssa.num_locals()).unwrap_or(u16::MAX);
let mut promotions: Vec<(SsaVarId, u16)> = Vec::new();
for (root, vars) in members {
if anchored.contains(&root) {
continue;
}
if next_local == u16::MAX {
break;
}
let slot = next_local;
next_local = next_local.saturating_add(1);
for var in vars {
promotions.push((var, slot));
}
}
if promotions.is_empty() {
return 0;
}
let num_args = u32::try_from(ssa.num_args()).unwrap_or(0);
for (var, slot) in &promotions {
if let Some(variable) = ssa.variable_mut(*var) {
variable.set_origin(VariableOrigin::Local(*slot));
}
ssa.set_rename_group(*var, num_args.saturating_add(u32::from(*slot)));
}
let added = usize::from(next_local).saturating_sub(ssa.num_locals());
ssa.set_num_locals(usize::from(next_local), ssa.original_num_locals());
added
}