use std::collections::{BTreeSet, HashMap, HashSet};
use crate::cfg::{BasicBlock, BlockId, Terminator};
type Set = BTreeSet<BlockId>;
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct Exit {
pub targets: Vec<BlockId>,
pub state: Option<u32>,
}
impl Exit {
pub fn nowhere() -> Self {
Self::default()
}
pub fn to(target: BlockId) -> Self {
Self {
targets: vec![target],
state: None,
}
}
pub fn index_of(&self, target: BlockId) -> Option<usize> {
self.targets.iter().position(|block| *block == target)
}
}
pub type Seq = Vec<Shape>;
#[derive(Clone, Debug)]
pub struct Arm {
pub entry: BlockId,
pub body: Seq,
}
#[derive(Clone, Debug)]
pub enum Shape {
Simple {
block: BlockId,
arms: Vec<Arm>,
},
Loop {
id: u32,
name: Option<String>,
entries: Vec<BlockId>,
state: Option<u32>,
body: Seq,
},
Dispatch {
state: u32,
arms: Vec<Arm>,
},
}
impl Shape {
pub fn exit(&self) -> Exit {
match self {
Shape::Simple { block, .. } => Exit::to(*block),
Shape::Loop { entries, state, .. } => Exit {
targets: entries.clone(),
state: *state,
},
Shape::Dispatch { state, arms } => Exit {
targets: arms.iter().map(|arm| arm.entry).collect(),
state: Some(*state),
},
}
}
}
#[derive(Clone, Debug)]
pub struct Plan {
pub body: Seq,
pub states: u32,
}
pub fn plan(blocks: &[BasicBlock], names: &HashMap<BlockId, String>) -> Option<Plan> {
if blocks.is_empty() {
return None;
}
let mut relooper = Relooper::new(blocks, names);
let all: Set = (0..blocks.len() as u32).map(BlockId).collect();
let entries: Set = [BlockId(0)].into_iter().collect();
let body = relooper.process(&entries, &all);
let plan = Plan {
body,
states: relooper.states,
};
let mut seen = vec![0usize; blocks.len()];
count_blocks(&plan.body, &mut seen);
if seen.iter().any(|times| *times != 1) {
return None;
}
if nesting(&plan.body, 0) > MAX_NESTING {
return None;
}
let mut scopes = Vec::new();
resolves(blocks, &plan.body, &Exit::nowhere(), &mut scopes).then_some(plan)
}
const MAX_NESTING: usize = 200;
fn nesting(seq: &Seq, depth: usize) -> usize {
if depth > MAX_NESTING {
return depth;
}
let last = seq.len().saturating_sub(1);
let mut out = depth;
for (index, shape) in seq.iter().enumerate() {
let here = depth + (last - index);
out = out.max(match shape {
Shape::Simple { arms, .. } | Shape::Dispatch { arms, .. } => arms
.iter()
.map(|arm| nesting(&arm.body, here + 1))
.max()
.unwrap_or(here),
Shape::Loop { body, .. } => nesting(body, here + 1),
});
if out > MAX_NESTING {
return out;
}
}
out
}
struct Relooper<'a> {
blocks: &'a [BasicBlock],
names: &'a HashMap<BlockId, String>,
succ: Vec<Vec<BlockId>>,
preds: Vec<Vec<BlockId>>,
cut: HashSet<(BlockId, BlockId)>,
states: u32,
loops: u32,
}
impl<'a> Relooper<'a> {
fn new(blocks: &'a [BasicBlock], names: &'a HashMap<BlockId, String>) -> Self {
let succ: Vec<Vec<BlockId>> = blocks
.iter()
.map(|block| {
let mut out = block.term.successors();
out.sort_unstable();
out.dedup();
out
})
.collect();
let mut preds: Vec<Vec<BlockId>> = vec![Vec::new(); blocks.len()];
for (index, targets) in succ.iter().enumerate() {
for target in targets {
preds[target.0 as usize].push(BlockId(index as u32));
}
}
Self {
blocks,
names,
succ,
preds,
cut: HashSet::new(),
states: 0,
loops: 0,
}
}
fn succ_in(&self, block: BlockId, set: &Set) -> Vec<BlockId> {
self.succ[block.0 as usize]
.iter()
.copied()
.filter(|target| set.contains(target) && !self.cut.contains(&(block, *target)))
.collect()
}
fn has_pred_in(&self, block: BlockId, set: &Set) -> bool {
self.preds[block.0 as usize]
.iter()
.any(|from| set.contains(from) && !self.cut.contains(&(*from, block)))
}
fn process(&mut self, entries: &Set, blocks: &Set) -> Seq {
let mut out = Seq::new();
let mut entries = entries.clone();
let mut blocks = blocks.clone();
loop {
entries.retain(|block| blocks.contains(block));
if entries.is_empty() {
return out;
}
entries = if entries.len() == 1 {
let entry = *entries.iter().next().expect("one entry");
if self.has_pred_in(entry, &blocks) {
self.make_loop(&entries, &mut blocks, &mut out)
} else {
self.make_simple(entry, &mut blocks, &mut out)
}
} else {
match self.make_multiple(&entries, &mut blocks, &mut out) {
Some(next) => next,
None => self.make_loop(&entries, &mut blocks, &mut out),
}
};
}
}
fn make_simple(&mut self, entry: BlockId, blocks: &mut Set, out: &mut Seq) -> Set {
blocks.remove(&entry);
let targets: Set = self.succ_in(entry, blocks).into_iter().collect();
if self.fusable(entry, &targets)
&& let Some((arms, next)) = self.groups(&targets, blocks)
{
out.push(Shape::Simple { block: entry, arms });
return next;
}
out.push(Shape::Simple {
block: entry,
arms: Vec::new(),
});
targets
}
fn fusable(&self, entry: BlockId, targets: &Set) -> bool {
if targets.len() < 2 {
return false;
}
let listed = self.terminator_targets(entry);
targets
.iter()
.all(|target| listed.iter().filter(|other| *other == target).count() <= 1)
}
fn terminator_targets(&self, block: BlockId) -> Vec<BlockId> {
match &self.blocks[block.0 as usize].term {
Terminator::Switch { cases, default, .. } => {
let mut out: Vec<BlockId> = Vec::new();
for (_, target) in cases {
if !out.contains(target) {
out.push(*target);
}
}
out.push(*default);
out
}
other => other.successors(),
}
}
fn make_multiple(&mut self, entries: &Set, blocks: &mut Set, out: &mut Seq) -> Option<Set> {
let (arms, next) = self.groups(entries, blocks)?;
for arm in arms {
out.extend(arm.body);
}
Some(next)
}
fn groups(&mut self, entries: &Set, blocks: &mut Set) -> Option<(Vec<Arm>, Set)> {
let owned = self.independent_groups(entries, blocks);
if owned.is_empty() {
return None;
}
let mut consumed = Set::new();
for (_, group) in &owned {
consumed.extend(group.iter().copied());
}
for block in &consumed {
blocks.remove(block);
}
let mut next: Set = entries
.iter()
.copied()
.filter(|entry| !consumed.contains(entry))
.collect();
for block in &consumed {
next.extend(self.succ_in(*block, blocks));
}
let arms: Vec<Arm> = owned
.into_iter()
.map(|(entry, group)| {
let one: Set = [entry].into_iter().collect();
Arm {
entry,
body: self.process(&one, &group),
}
})
.collect();
Some((arms, next))
}
fn independent_groups(&self, entries: &Set, blocks: &Set) -> Vec<(BlockId, Set)> {
let count = self.blocks.len();
let mut owner: Vec<Option<BlockId>> = vec![None; count];
let mut shared = vec![false; count];
for entry in entries {
owner[entry.0 as usize] = Some(*entry);
}
let mut seen = vec![false; count];
let mut stack: Vec<BlockId> = Vec::new();
for entry in entries {
seen.iter_mut().for_each(|flag| *flag = false);
stack.clear();
stack.extend(self.succ_in(*entry, blocks));
for target in &stack {
seen[target.0 as usize] = true;
}
while let Some(block) = stack.pop() {
match owner[block.0 as usize] {
None => owner[block.0 as usize] = Some(*entry),
Some(other) if other != *entry => shared[block.0 as usize] = true,
Some(_) => {}
}
for target in self.succ_in(block, blocks) {
if !seen[target.0 as usize] {
seen[target.0 as usize] = true;
stack.push(target);
}
}
}
}
let mut out = Vec::new();
for entry in entries {
if shared[entry.0 as usize] {
continue;
}
let group: Set = blocks
.iter()
.copied()
.filter(|block| {
owner[block.0 as usize] == Some(*entry) && !shared[block.0 as usize]
})
.collect();
out.push((*entry, group));
}
out
}
fn make_loop(&mut self, entries: &Set, blocks: &mut Set, out: &mut Seq) -> Set {
let mut inner = entries.clone();
inner.extend(self.reaching(entries, blocks));
if self.leaving(blocks, &inner).len() > 1 {
self.widen_loop(blocks, &mut inner);
}
let follow = self.leaving(blocks, &inner);
for block in &inner {
blocks.remove(block);
}
for entry in entries {
for from in self.preds[entry.0 as usize].clone() {
if inner.contains(&from) {
self.cut.insert((from, *entry));
}
}
}
let id = self.loops;
self.loops += 1;
let name = entries
.iter()
.find_map(|entry| self.names.get(entry))
.cloned();
let (state, entry_list, body) = if entries.len() == 1 {
let body = self.process(entries, &inner);
(None, entries.iter().copied().collect(), body)
} else {
let state = self.states;
self.states += 1;
let mut rest = inner.clone();
match self.groups(entries, &mut rest) {
Some((arms, next)) => {
let list: Vec<BlockId> = arms.iter().map(|arm| arm.entry).collect();
let mut body = vec![Shape::Dispatch { state, arms }];
body.extend(self.process(&next, &rest));
(Some(state), list, body)
}
None => (Some(state), entries.iter().copied().collect(), Seq::new()),
}
};
out.push(Shape::Loop {
id,
name,
entries: entry_list,
state,
body,
});
follow
}
fn widen_loop(&self, blocks: &Set, inner: &mut Set) {
let mut queue: Vec<BlockId> = self.leaving(blocks, inner).into_iter().collect();
while let Some(mut block) = queue.pop() {
loop {
if inner.contains(&block) {
break;
}
let entered = self.preds[block.0 as usize]
.iter()
.filter(|from| blocks.contains(from) && !self.cut.contains(&(**from, block)))
.count();
if entered != 1 {
break;
}
inner.insert(block);
let targets = self.succ_in(block, blocks);
let Some((next, rest)) = targets.split_first() else {
break;
};
queue.extend(rest.iter().copied());
block = *next;
}
}
}
fn leaving(&self, blocks: &Set, part: &Set) -> Set {
let mut out = Set::new();
for block in part {
for target in self.succ_in(*block, blocks) {
if !part.contains(&target) {
out.insert(target);
}
}
}
out
}
fn reaching(&self, targets: &Set, set: &Set) -> Set {
let mut seen = vec![false; self.blocks.len()];
let mut out = Set::new();
let mut stack: Vec<BlockId> = targets.iter().copied().collect();
let mut todo: Vec<BlockId> = Vec::new();
while let Some(block) = stack.pop() {
for from in &self.preds[block.0 as usize] {
if !set.contains(from)
|| self.cut.contains(&(*from, block))
|| seen[from.0 as usize]
{
continue;
}
seen[from.0 as usize] = true;
out.insert(*from);
todo.push(*from);
}
stack.append(&mut todo);
}
out
}
}
fn count_blocks(seq: &Seq, out: &mut [usize]) {
for shape in seq {
match shape {
Shape::Simple { block, arms } => {
out[block.0 as usize] += 1;
for arm in arms {
count_blocks(&arm.body, out);
}
}
Shape::Loop { body, .. } => count_blocks(body, out),
Shape::Dispatch { arms, .. } => {
for arm in arms {
count_blocks(&arm.body, out);
}
}
}
}
}
fn resolves(blocks: &[BasicBlock], seq: &Seq, fall: &Exit, scopes: &mut Vec<Exit>) -> bool {
let exits: Vec<Exit> = seq.iter().map(Shape::exit).collect();
let depth = scopes.len();
for exit in exits.iter().skip(1).rev() {
scopes.push(exit.clone());
}
let mut ok = true;
for (index, shape) in seq.iter().enumerate() {
if index > 0 {
scopes.pop();
}
let next = exits.get(index + 1).unwrap_or(fall);
ok &= resolves_shape(blocks, shape, next, scopes);
}
scopes.truncate(depth);
ok
}
fn resolves_shape(
blocks: &[BasicBlock],
shape: &Shape,
fall: &Exit,
scopes: &mut Vec<Exit>,
) -> bool {
match shape {
Shape::Simple { block, arms } => {
let mut ok = true;
for target in blocks[block.0 as usize].term.successors() {
match arms.iter().find(|arm| arm.entry == target) {
Some(arm) => ok &= resolves(blocks, &arm.body, fall, scopes),
None => {
ok &= fall.index_of(target).is_some()
|| scopes.iter().any(|exit| exit.index_of(target).is_some());
}
}
}
ok
}
Shape::Loop {
entries,
state,
body,
..
} => {
let head = Exit {
targets: entries.clone(),
state: *state,
};
scopes.push(head.clone());
scopes.push(fall.clone());
let ok = resolves(blocks, body, &head, scopes);
scopes.pop();
scopes.pop();
ok
}
Shape::Dispatch { arms, .. } => arms
.iter()
.all(|arm| resolves(blocks, &arm.body, fall, scopes)),
}
}