use rucc_ir::{Block, BlockCall, Def, Extra, Func, Imm, Inst, IntPred, Opcode, SwitchInfo, Value};
use crate::fold::constant;
use crate::range::ops::Truth;
use crate::range::query::Ranges;
use crate::simplify_cfg;
use crate::{Analyses, Fuel, Pass, Preserved, Stats};
const BRANCH_DECIDED: &str =
"branch the value ranges settle whichever way control reached it replaced by a jump";
const CASE_REMOVED: &str =
"case whose value the switched value cannot hold taken out of the switch";
const SWITCH_REMOVED: &str =
"switch none of whose cases the switched value can reach replaced by a jump to its default";
const BRANCH_UNDECIDED: &str = "branch kept, the value ranges do not settle which way it goes";
const NO_CASE_REMOVED: &str = "switch kept whole, the value ranges rule none of its cases out";
const NO_FUEL: &str = "branch or switch kept, the pass ran out of fuel";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Prune;
impl Pass for Prune {
fn name(&self) -> &'static str {
"prune"
}
fn describe(&self) -> &'static str {
"a branch the ranges settle becomes a jump, and a case they rule out leaves its switch"
}
fn preserves(&self) -> Preserved {
Preserved::NONE
}
fn run(&self, func: &mut Func, an: &mut Analyses, fuel: &mut Fuel) -> Stats {
let mut stats = Stats::new();
if func.entry().is_none() {
return stats;
}
let plan = answers(func, an, &mut stats);
if plan.branches.is_empty() && plan.switches.is_empty() {
return stats;
}
'apply: {
for (term, call) in plan.branches {
if !fuel.take() {
stats.missed(NO_FUEL);
break 'apply;
}
simplify_cfg::jump_to(func, term, call);
stats.optimized(BRANCH_DECIDED);
}
for (term, keeping) in plan.switches {
if !fuel.take() {
stats.missed(NO_FUEL);
break 'apply;
}
let removed = shrink(func, term, &keeping);
for _ in 0..removed {
stats.optimized(CASE_REMOVED);
}
if keeping.is_empty() {
stats.optimized(SWITCH_REMOVED);
}
}
}
an.clear();
simplify_cfg::sweep(func, an, &mut stats);
stats
}
}
#[derive(Debug, Default)]
struct Plan {
branches: Vec<(Inst, BlockCall)>,
switches: Vec<(Inst, Vec<usize>)>,
}
fn answers(func: &Func, an: &mut Analyses, stats: &mut Stats) -> Plan {
let mut plan = Plan::default();
let cfg = an.cfg(func).clone();
let dom = an.dominators(func).clone();
let mut ranges = Ranges::new(func, &cfg, &dom);
for block in func.blocks() {
if !cfg.reaches(block) {
continue;
}
let Some(term) = func.terminator(block) else { continue };
match func[term].opcode {
Opcode::BrIf => match decided(func, &mut ranges, block, term) {
Answer::Jump(call) => plan.branches.push((term, call)),
Answer::Unsettled => stats.missed(BRANCH_UNDECIDED),
Answer::NotAsked => (),
},
Opcode::Switch => match reachable(func, &mut ranges, block, term) {
Some(keeping) => plan.switches.push((term, keeping)),
None => stats.missed(NO_CASE_REMOVED),
},
_ => (),
}
}
plan
}
enum Answer {
Jump(BlockCall),
Unsettled,
NotAsked,
}
fn decided(func: &Func, ranges: &mut Ranges<'_>, block: Block, term: Inst) -> Answer {
let data = &func[term];
let Extra::Targets(targets) = data.extra else { return Answer::NotAsked };
let Some(&cond) = func[data.args].first() else { return Answer::NotAsked };
if constant(func, cond).is_some() {
return Answer::NotAsked;
}
let calls = &func[targets];
let together =
|two: &[BlockCall]| two[0].block == two[1].block && func[two[0].args] == func[two[1].args];
if calls.len() == 2 && together(calls) {
return Answer::NotAsked;
}
let arm = match settled(func, ranges, block, cond) {
Some(true) => 0,
Some(false) => 1,
None => return Answer::Unsettled,
};
func[targets].get(arm).copied().map_or(Answer::NotAsked, Answer::Jump)
}
fn settled(func: &Func, ranges: &mut Ranges<'_>, block: Block, cond: Value) -> Option<bool> {
if let Some((pred, lhs, rhs)) = comparison(func, cond) {
return match ranges.compare(pred, lhs, rhs, block) {
Truth::Always => Some(true),
Truth::Never => Some(false),
Truth::Either => None,
};
}
let range = ranges.at(cond, block);
if range.nonzero() {
return Some(true);
}
(range.singleton() == Some(0)).then_some(false)
}
fn comparison(func: &Func, value: Value) -> Option<(IntPred, Value, Value)> {
let Def::Result { inst, .. } = func[value].def else { return None };
if func[inst].opcode != Opcode::ICmp {
return None;
}
let Extra::IntPred(pred) = func[inst].extra else { return None };
let &[lhs, rhs] = func[func[inst].args].first_chunk::<2>()?;
Some((pred, lhs, rhs))
}
fn reachable(func: &Func, ranges: &mut Ranges<'_>, block: Block, term: Inst) -> Option<Vec<usize>> {
let Extra::Switch(at) = func[term].extra else { return None };
let info = func[at];
let arg = *func[func[term].args].first()?;
let range = ranges.at(arg, block);
if range.is_full() {
return None;
}
let cases = &func[info.cases];
let keeping: Vec<usize> =
(0..cases.len()).filter(|&at| range.contains(cases[at].unsigned())).collect();
(keeping.len() < cases.len()).then_some(keeping)
}
fn shrink(func: &mut Func, term: Inst, keeping: &[usize]) -> usize {
let Extra::Switch(at) = func[term].extra else { return 0 };
let info = func[at];
let all = func[info.targets].to_vec();
let values = func[info.cases].to_vec();
let removed = values.len() - keeping.len();
let default = all[0];
if keeping.is_empty() {
simplify_cfg::jump_to(func, term, default);
return removed;
}
let targets: Vec<BlockCall> =
std::iter::once(default).chain(keeping.iter().map(|&at| all[at + 1])).collect();
let cases: Vec<Imm> = keeping.iter().map(|&at| values[at]).collect();
let targets = func.push_block_calls(&targets);
let cases = func.push_imms(&cases);
let fresh = func.add_switch(SwitchInfo { targets, cases });
func[term].extra = Extra::Switch(fresh);
removed
}
#[cfg(test)]
mod tests {
use rucc_base::Interner;
use rucc_ir::{Block, Builder, Flags, Func, IntPred, Opcode, Signature, Type};
use super::Prune;
use crate::stats::Kind;
use crate::{Analyses, Fuel, Pass, Stats};
fn prune(func: &mut Func) -> Stats {
Prune.run(func, &mut Analyses::new(), &mut Fuel::unlimited())
}
fn terminator(func: &Func, block: usize) -> Opcode {
let block = Block::from_usize(block);
func[func.terminator(block).expect("every block here has one")].opcode
}
fn goes_to(func: &Func, block: usize) -> Vec<usize> {
let block = Block::from_usize(block);
let term = func.terminator(block).expect("every block here has one");
func.successors(term).map(|call| call.block.index()).collect()
}
fn nested(outer: IntPred, bound: i128, inner: IntPred) -> Func {
let mut names = Interner::new();
let ty = Type::int(32);
let mut func = Func::new(names.intern("f"), Signature::new().with_params(&[ty]));
let entry = func.create_block();
let middle = func.create_block();
let arms = [func.create_block(), func.create_block()];
let away = func.create_block();
let x = func.append_param(entry, ty);
let mut build = Builder::new(&mut func, entry);
let edge = build.iconst(ty, bound);
let first = build.icmp(outer, x, edge);
build.br_if(first, middle, &[], away, &[]);
let mut build = Builder::new(&mut func, middle);
let five = build.iconst(ty, 5);
let second = build.icmp(inner, x, five);
build.br_if(second, arms[0], &[], arms[1], &[]);
for block in [arms[0], arms[1], away] {
let mut build = Builder::new(&mut func, block);
build.ret(&[]);
}
func
}
#[test]
fn a_branch_the_ranges_settle_becomes_a_jump_to_the_arm_they_settle_on() {
let mut func = nested(IntPred::Sgt, 10, IntPred::Sgt);
let stats = prune(&mut func);
assert!(stats.changed());
assert_eq!(terminator(&func, 1), Opcode::Jump);
assert_eq!(goes_to(&func, 1), [2]);
assert_eq!(stats.count(Kind::Optimized, super::BRANCH_DECIDED), 1);
}
#[test]
fn a_branch_the_ranges_settle_the_other_way_jumps_to_the_other_arm() {
let mut func = nested(IntPred::Sgt, 10, IntPred::Slt);
let stats = prune(&mut func);
assert!(stats.changed());
assert_eq!(terminator(&func, 1), Opcode::Jump);
assert_eq!(goes_to(&func, 1), [3]);
}
#[test]
fn a_branch_the_ranges_do_not_settle_keeps_its_two_arms() {
let mut func = nested(IntPred::Sgt, 3, IntPred::Sgt);
let stats = prune(&mut func);
assert!(!stats.changed());
assert_eq!(terminator(&func, 1), Opcode::BrIf);
assert_eq!(stats.count(Kind::Missed, super::BRANCH_UNDECIDED), 2);
}
#[test]
fn a_branch_on_a_constant_is_left_for_the_control_flow_pass() {
let mut names = Interner::new();
let mut func = Func::new(names.intern("f"), Signature::new());
let entry = func.create_block();
let arms = [func.create_block(), func.create_block()];
let mut build = Builder::new(&mut func, entry);
let always = build.iconst(Type::int(1), 1);
build.br_if(always, arms[0], &[], arms[1], &[]);
for block in arms {
let mut build = Builder::new(&mut func, block);
build.ret(&[]);
}
let stats = prune(&mut func);
assert!(!stats.changed());
assert_eq!(terminator(&func, 0), Opcode::BrIf);
assert_eq!(stats.count(Kind::Missed, super::BRANCH_UNDECIDED), 0);
}
fn masked(mask: i128, cases: &[i128]) -> Func {
let mut names = Interner::new();
let ty = Type::int(32);
let mut func = Func::new(names.intern("f"), Signature::new().with_params(&[ty]));
let entry = func.create_block();
let default = func.create_block();
let arms: Vec<Block> = cases.iter().map(|_| func.create_block()).collect();
let x = func.append_param(entry, ty);
let mut build = Builder::new(&mut func, entry);
let bits = build.iconst(ty, mask);
let narrowed = build.binary(Opcode::And, x, bits, Flags::NONE);
let pairs: Vec<(i128, Block)> =
cases.iter().copied().zip(arms.iter().copied()).collect::<Vec<_>>();
build.switch(narrowed, default, &pairs);
for block in std::iter::once(default).chain(arms) {
let mut build = Builder::new(&mut func, block);
build.ret(&[]);
}
func
}
#[test]
fn a_case_the_switched_value_cannot_hold_leaves_the_switch() {
let mut func = masked(3, &[0, 2, 7]);
let stats = prune(&mut func);
assert!(stats.changed());
assert_eq!(terminator(&func, 0), Opcode::Switch);
assert_eq!(goes_to(&func, 0), [1, 2, 3]);
assert_eq!(stats.count(Kind::Optimized, super::CASE_REMOVED), 1);
}
#[test]
fn a_switch_no_case_of_which_can_be_reached_jumps_to_its_default() {
let mut func = masked(1, &[5, 9]);
let stats = prune(&mut func);
assert!(stats.changed());
assert_eq!(terminator(&func, 0), Opcode::Jump);
assert_eq!(goes_to(&func, 0), [1]);
assert_eq!(stats.count(Kind::Optimized, super::SWITCH_REMOVED), 1);
}
#[test]
fn a_switch_whose_cases_the_operand_can_all_hold_is_kept_whole() {
let mut func = masked(7, &[0, 2, 7]);
let stats = prune(&mut func);
assert!(!stats.changed());
assert_eq!(terminator(&func, 0), Opcode::Switch);
assert_eq!(stats.count(Kind::Missed, super::NO_CASE_REMOVED), 1);
}
#[test]
fn a_branch_two_values_are_related_on_settles_without_either_being_pinned_down() {
let mut names = Interner::new();
let ty = Type::int(32);
let mut func = Func::new(names.intern("f"), Signature::new().with_params(&[ty, ty]));
let entry = func.create_block();
let middle = func.create_block();
let arms = [func.create_block(), func.create_block()];
let away = func.create_block();
let a = func.append_param(entry, ty);
let b = func.append_param(entry, ty);
let mut build = Builder::new(&mut func, entry);
let first = build.icmp(IntPred::Slt, a, b);
build.br_if(first, middle, &[], away, &[]);
let mut build = Builder::new(&mut func, middle);
let second = build.icmp(IntPred::Sle, a, b);
build.br_if(second, arms[0], &[], arms[1], &[]);
for block in [arms[0], arms[1], away] {
let mut build = Builder::new(&mut func, block);
build.ret(&[]);
}
let stats = prune(&mut func);
assert!(stats.changed());
assert_eq!(terminator(&func, 1), Opcode::Jump);
assert_eq!(goes_to(&func, 1), [2]);
}
#[test]
fn a_value_that_cannot_be_zero_is_a_branch_that_always_holds() {
let mut names = Interner::new();
let wide = Type::int(32);
let mut func = Func::new(names.intern("f"), Signature::new().with_params(&[wide]));
let entry = func.create_block();
let middle = func.create_block();
let arms = [func.create_block(), func.create_block()];
let away = func.create_block();
let x = func.append_param(entry, wide);
let mut build = Builder::new(&mut func, entry);
let five = build.iconst(wide, 5);
let is_five = build.icmp(IntPred::Eq, x, five);
build.br_if(is_five, middle, &[], away, &[]);
let mut build = Builder::new(&mut func, middle);
let bit = build.unary(Opcode::Trunc, x, Type::int(1));
build.br_if(bit, arms[0], &[], arms[1], &[]);
for block in [arms[0], arms[1], away] {
let mut build = Builder::new(&mut func, block);
build.ret(&[]);
}
let stats = prune(&mut func);
assert!(stats.changed());
assert_eq!(terminator(&func, 1), Opcode::Jump);
assert_eq!(goes_to(&func, 1), [2]);
}
#[test]
fn a_function_with_no_body_is_left_alone() {
let mut names = Interner::new();
let mut func = Func::new(names.intern("f"), Signature::new());
let stats = prune(&mut func);
assert!(!stats.changed());
}
#[test]
fn no_fuel_leaves_the_branch_where_it_is() {
let mut func = nested(IntPred::Sgt, 10, IntPred::Sgt);
let stats = Prune.run(&mut func, &mut Analyses::new(), &mut Fuel::of(0));
assert!(!stats.changed());
assert_eq!(terminator(&func, 1), Opcode::BrIf);
assert_eq!(stats.count(Kind::Missed, super::NO_FUEL), 1);
}
}