use std::collections::{HashMap, VecDeque};
use rucc_ir::{Block, Def, Func, Inst, Value};
#[must_use]
pub fn on_entry(func: &Func) -> Vec<(u32, Block, Value)> {
let mut assigned: HashMap<u32, Vec<(Block, Option<Inst>, Value)>> = HashMap::new();
for value in func.values() {
let place = match func[value].def {
Def::Result { inst, .. } => func.block_of(inst).map(|block| (block, Some(inst))),
Def::Param { block, index } => {
let param = usize::try_from(index)
.ok()
.and_then(|index| func[block].params.get(index).copied());
(func.is_placed(block) && param == Some(value)).then_some((block, None))
}
};
if let Some((block, after)) = place {
for decl in func.value_decls(value) {
assigned.entry(decl).or_default().push((block, after, value));
}
}
for start in func.value_starts(value) {
if let Some((block, after)) = func.start_place(start) {
assigned.entry(start.decl).or_default().push((block, after, value));
}
}
}
assigned.retain(|_, all| all.iter().any(|&(_, _, value)| value != all[0].2));
if assigned.is_empty() {
return Vec::new();
}
let blocks: Vec<Block> = func.blocks().collect();
let count = func.counts().blocks;
let mut position = vec![0usize; func.counts().insts];
let mut succs: Vec<Vec<Block>> = vec![Vec::new(); count];
for &block in &blocks {
for (at, inst) in func.insts(block).enumerate() {
position[inst.index()] = at + 1;
}
if let Some(terminator) = func.terminator(block) {
succs[block.index()].extend(func.successors(terminator).map(|call| call.block));
}
}
let entry = func.entry();
let mut decls: Vec<u32> = assigned.keys().copied().collect();
decls.sort_unstable();
let mut out = Vec::new();
for decl in decls {
let mut last: Vec<Option<(usize, Held)>> = vec![None; count];
let mut top: Vec<Option<Held>> = vec![None; count];
for &(block, after, value) in &assigned[&decl] {
let at = after.map_or(0, |inst| position[inst.index()]);
let slot = &mut last[block.index()];
*slot = match *slot {
Some((have, _)) if have > at => *slot,
Some((have, held)) if have == at => Some((at, held.meet(Held::Known(value)))),
_ => Some((at, Held::Known(value))),
};
if at == 0 {
let slot = &mut top[block.index()];
*slot = Some(slot.map_or(Held::Known(value), |held| held.meet(Held::Known(value))));
}
}
let mut into: Vec<Held> = vec![Held::Unvisited; count];
let mut queued = vec![false; count];
let mut waiting = VecDeque::new();
if let Some(entry) = entry {
into[entry.index()] = Held::Unknown;
queued[entry.index()] = true;
waiting.push_back(entry);
}
for &(block, _, _) in &assigned[&decl] {
if !queued[block.index()] {
queued[block.index()] = true;
waiting.push_back(block);
}
}
while let Some(block) = waiting.pop_front() {
queued[block.index()] = false;
let out = last[block.index()].map_or(into[block.index()], |(_, held)| held);
for &succ in &succs[block.index()] {
if Some(succ) == entry {
continue;
}
let now = into[succ.index()].meet(out);
if now != into[succ.index()] {
into[succ.index()] = now;
if !queued[succ.index()] {
queued[succ.index()] = true;
waiting.push_back(succ);
}
}
}
}
for &block in &blocks {
let held = top[block.index()].unwrap_or(into[block.index()]);
if let Held::Known(value) = held {
out.push((decl, block, value));
}
}
}
out
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Held {
Unvisited,
Known(Value),
Unknown,
}
impl Held {
fn meet(self, other: Held) -> Held {
match (self, other) {
(Held::Unvisited, held) | (held, Held::Unvisited) => held,
(Held::Known(one), Held::Known(two)) if one == two => self,
_ => Held::Unknown,
}
}
}
#[cfg(test)]
mod tests {
use rucc_base::Symbol;
use rucc_ir::{Builder, Signature, Start, Type};
use super::*;
#[test]
fn a_value_computed_on_this_trip_is_what_the_blocks_after_it_hold() {
let mut func = Func::new(Symbol::from_raw(0), Signature::new());
let entry = func.create_block();
let head = func.create_block();
let body = func.create_block();
let after = func.create_block();
let exit = func.create_block();
let i = func.append_param(head, Type::int(32));
let mut build = Builder::new(&mut func, entry);
let zero = build.iconst(Type::int(32), 0);
build.jump(head, &[zero]);
Builder::new(&mut func, head).jump(body, &[]);
let mut build = Builder::new(&mut func, body);
let one = build.iconst(Type::int(32), 1);
let next = build.binary(rucc_ir::Opcode::Add, i, one, rucc_ir::Flags::NONE);
build.jump(after, &[]);
let mut build = Builder::new(&mut func, after);
let done = build.icmp(rucc_ir::IntPred::Eq, i, next);
build.br_if(done, exit, &[], head, &[next]);
Builder::new(&mut func, exit).ret(&[]);
for value in [zero, i, next] {
func.declare_value(value, 3);
}
let held = on_entry(&func);
assert!(held.contains(&(3, head, i)), "{held:?}");
assert!(held.contains(&(3, body, i)), "{held:?}");
assert!(held.contains(&(3, after, next)), "{held:?}");
assert!(held.contains(&(3, exit, next)), "{held:?}");
assert!(!held.iter().any(|&(_, block, _)| block == entry), "{held:?}");
}
#[test]
fn arms_that_disagree_say_nothing_and_a_start_says_what_it_gave() {
let mut func = Func::new(Symbol::from_raw(0), Signature::new());
let entry = func.create_block();
let left = func.create_block();
let right = func.create_block();
let join = func.create_block();
let flag = func.append_param(entry, Type::I1);
let mut build = Builder::new(&mut func, entry);
let first = build.iconst(Type::int(32), 1);
build.br_if(flag, left, &[], right, &[]);
let mut build = Builder::new(&mut func, left);
let second = build.iconst(Type::int(32), 2);
let saved = build.iconst(Type::int(32), 3);
build.jump(join, &[]);
Builder::new(&mut func, right).jump(join, &[]);
Builder::new(&mut func, join).ret(&[]);
func.declare_value(first, 5);
func.declare_value(second, 5);
let held = on_entry(&func);
assert!(held.contains(&(5, left, first)), "{held:?}");
assert!(held.contains(&(5, right, first)), "{held:?}");
assert!(!held.iter().any(|&(_, block, _)| block == join), "{held:?}");
let Def::Result { inst: loaded, .. } = func[saved].def else { unreachable!() };
func.declare_value_from(first, Start { decl: 5, block: left, after: Some(loaded) });
func.declare_value_from(second, Start { decl: 6, block: left, after: None });
let held = on_entry(&func);
assert!(held.contains(&(5, join, first)), "{held:?}");
assert!(held.contains(&(5, left, first)), "{held:?}");
assert!(!held.iter().any(|&(decl, _, _)| decl == 6), "{held:?}");
}
}