use std::collections::{HashMap, HashSet};
use crate::capture::SourceRange;
use crate::ir::{
BinOp, BreakTarget, CaseRange, Expr, ExprKind, LabelId, LoopId, Object, ObjectId, Place,
PlaceKind, Stmt, Storage, SwitchId, is_always_true,
};
#[derive(Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Debug)]
pub struct BlockId(pub u32);
impl BlockId {
pub(crate) fn index(self) -> usize {
self.0 as usize
}
}
#[derive(Clone, Debug)]
pub struct Local {
pub object: ObjectId,
pub rust_name: String,
}
#[derive(Clone, Debug)]
pub struct BasicBlock {
pub stmts: Vec<Stmt>,
pub term: Terminator,
}
#[derive(Clone, Debug)]
pub enum Terminator {
Jump {
target: BlockId,
range: SourceRange,
},
Branch {
cond: Expr,
then_blk: BlockId,
else_blk: BlockId,
},
Switch {
value: Expr,
cases: Vec<(CaseRange, BlockId)>,
default: BlockId,
range: SourceRange,
},
Return {
value: Option<Expr>,
range: SourceRange,
},
Unreachable,
InvalidTarget,
}
impl Terminator {
pub(crate) fn successors(&self) -> Vec<BlockId> {
match self {
Terminator::Jump { target, .. } => vec![*target],
Terminator::Branch {
then_blk, else_blk, ..
} => vec![*then_blk, *else_blk],
Terminator::Switch { cases, default, .. } => {
let mut out: Vec<BlockId> = cases.iter().map(|(_, blk)| *blk).collect();
out.push(*default);
out
}
Terminator::Return { .. } | Terminator::Unreachable | Terminator::InvalidTarget => {
Vec::new()
}
}
}
fn map_targets(&mut self, mut f: impl FnMut(BlockId) -> BlockId) {
match self {
Terminator::Jump { target, .. } => *target = f(*target),
Terminator::Branch {
then_blk, else_blk, ..
} => {
*then_blk = f(*then_blk);
*else_blk = f(*else_blk);
}
Terminator::Switch { cases, default, .. } => {
for (_, blk) in cases.iter_mut() {
*blk = f(*blk);
}
*default = f(*default);
}
Terminator::Return { .. } | Terminator::Unreachable | Terminator::InvalidTarget => {}
}
}
}
#[derive(Clone, Debug)]
pub struct Cfg {
pub locals: Vec<Local>,
pub blocks: Vec<BasicBlock>,
pub labels: HashMap<LabelId, u32>,
pub shape: Option<crate::reloop::Plan>,
}
pub fn lower(
body: Vec<Stmt>,
params: &[ObjectId],
objects: &[Object],
taken: &[LabelId],
goto_value: Option<ObjectId>,
table: Option<ObjectId>,
names: &HashMap<LabelId, String>,
) -> Cfg {
let mut used: HashSet<String> = params
.iter()
.map(|id| objects[id.0 as usize].name.clone())
.collect();
used.insert("__cinrs_state".to_owned());
let mut lowerer = Lowerer {
objects,
blocks: Vec::new(),
current: None,
labels: HashMap::new(),
loops: HashMap::new(),
switches: HashMap::new(),
locals: Vec::new(),
used,
cleanups: Vec::new(),
label_cleanups: HashMap::new(),
taken: Vec::new(),
goto_value,
table,
dispatch: None,
};
lowerer.collect_label_cleanups(&body, 0);
for id in taken {
let block = lowerer.label_block(*id);
lowerer.taken.push((*id, block));
}
if let Some(object) = goto_value {
let name = objects[object.0 as usize].name.clone();
let rust_name = lowerer.unique_name(&name);
lowerer.locals.push(Local { object, rust_name });
}
let entry = lowerer.new_block();
lowerer.current = Some(entry);
lowerer.stmts(body);
if let Some(open) = lowerer.current.take() {
lowerer.blocks[open.index()].term = Terminator::Unreachable;
}
let labels = taken
.iter()
.enumerate()
.map(|(index, id)| (*id, index as u32 + 1))
.collect();
lowerer.finish(entry, labels, names)
}
#[derive(Clone, Copy)]
struct LoopBlocks {
brk: BlockId,
brk_cleanups: usize,
cont: BlockId,
cont_cleanups: usize,
}
struct SwitchFrame {
brk: BlockId,
brk_cleanups: usize,
cases: Vec<(CaseRange, BlockId)>,
default: Option<BlockId>,
}
struct Lowerer<'a> {
objects: &'a [Object],
blocks: Vec<BasicBlock>,
current: Option<BlockId>,
labels: HashMap<LabelId, BlockId>,
loops: HashMap<LoopId, LoopBlocks>,
switches: HashMap<SwitchId, SwitchFrame>,
locals: Vec<Local>,
used: HashSet<String>,
cleanups: Vec<Expr>,
label_cleanups: HashMap<LabelId, usize>,
taken: Vec<(LabelId, BlockId)>,
goto_value: Option<ObjectId>,
table: Option<ObjectId>,
dispatch: Option<BlockId>,
}
impl Lowerer<'_> {
fn new_block(&mut self) -> BlockId {
let id = BlockId(self.blocks.len() as u32);
self.blocks.push(BasicBlock {
stmts: Vec::new(),
term: Terminator::Unreachable,
});
id
}
fn current(&mut self) -> BlockId {
match self.current {
Some(id) => id,
None => {
let id = self.new_block();
self.current = Some(id);
id
}
}
}
fn push(&mut self, stmt: Stmt) {
let id = self.current();
self.blocks[id.index()].stmts.push(stmt);
}
fn terminate(&mut self, term: Terminator) {
let id = self.current();
self.blocks[id.index()].term = term;
self.current = None;
}
fn jump(&mut self, target: BlockId, range: SourceRange) {
self.terminate(Terminator::Jump { target, range });
}
fn continue_at(&mut self, block: BlockId) {
self.current = Some(block);
}
fn label_block(&mut self, id: LabelId) -> BlockId {
if let Some(block) = self.labels.get(&id) {
return *block;
}
let block = self.new_block();
self.labels.insert(id, block);
block
}
fn collect_label_cleanups(&mut self, stmts: &[Stmt], depth: usize) {
let mut depth = depth;
for stmt in stmts {
match stmt {
Stmt::Cleanup(_) => depth += 1,
Stmt::Block(items) => self.collect_label_cleanups(items, depth),
Stmt::Label { id, body, .. } => {
self.label_cleanups.insert(*id, depth);
self.collect_label_cleanups(std::slice::from_ref(body), depth);
}
Stmt::If {
then_branch,
else_branch,
..
} => {
self.collect_label_cleanups(std::slice::from_ref(then_branch), depth);
if let Some(branch) = else_branch {
self.collect_label_cleanups(std::slice::from_ref(branch), depth);
}
}
Stmt::While { body, .. } | Stmt::DoWhile { body, .. } => {
self.collect_label_cleanups(std::slice::from_ref(body), depth);
}
Stmt::For { init, body, .. } => {
let inner = depth + init.iter().filter(|s| s.is_cleanup()).count();
self.collect_label_cleanups(std::slice::from_ref(body), inner);
}
Stmt::SwitchTree(switch) => {
self.collect_label_cleanups(std::slice::from_ref(&switch.body), depth);
}
Stmt::Case { body, .. } => {
self.collect_label_cleanups(std::slice::from_ref(body), depth);
}
_ => {}
}
}
}
fn leave_cleanups(&mut self, depth: usize) {
if self.cleanups.len() <= depth || self.current.is_none() {
return;
}
let owed: Vec<Expr> = self.cleanups[depth..].iter().rev().cloned().collect();
for call in owed {
self.push(Stmt::Expr(call));
}
}
fn stmts(&mut self, stmts: Vec<Stmt>) {
for stmt in stmts {
self.stmt(stmt);
}
}
fn scope(&mut self, stmts: Vec<Stmt>) {
let depth = self.cleanups.len();
self.stmts(stmts);
self.leave_cleanups(depth);
self.cleanups.truncate(depth);
}
fn stmt(&mut self, stmt: Stmt) {
match stmt {
Stmt::Nop => {}
Stmt::Expr(expr) => self.push(Stmt::Expr(expr)),
asm @ Stmt::Asm(_) => self.push(asm),
Stmt::Let {
object,
init,
explicit,
} => self.local(object, init, explicit),
Stmt::Vla(def) => self.vla(*def),
Stmt::Cleanup(def) => self.cleanups.push(def.call),
Stmt::Block(items) => self.scope(items),
Stmt::If {
cond,
then_branch,
else_branch,
} => self.if_stmt(cond, *then_branch, else_branch.map(|b| *b)),
Stmt::While {
id,
cond,
body,
range,
} => self.while_stmt(id, cond, *body, range),
Stmt::DoWhile {
id,
body,
cond,
range,
} => self.do_while(id, *body, cond, range),
Stmt::For {
id,
init,
cond,
step,
body,
range,
} => self.for_stmt(id, init, cond, step, *body, range),
Stmt::SwitchTree(switch) => self.switch(*switch),
Stmt::Case {
switch,
value,
body,
range,
} => self.case(switch, value, *body, range),
Stmt::Label { id, body, range } => {
let block = self.label_block(id);
self.jump(block, range);
self.continue_at(block);
self.stmt(*body);
}
Stmt::Goto { id, range } => {
let block = self.label_block(id);
let depth = self.label_cleanups.get(&id).copied().unwrap_or(0);
self.leave_cleanups(depth.min(self.cleanups.len()));
self.jump(block, range);
}
Stmt::GotoPtr { target, range } => self.goto_ptr(target, range),
Stmt::Break { target, range } => {
let block = match target {
BreakTarget::Loop(id) => self.loops.get(&id).map(|l| (l.brk, l.brk_cleanups)),
BreakTarget::Switch(id) => {
self.switches.get(&id).map(|s| (s.brk, s.brk_cleanups))
}
};
match block {
Some((block, depth)) => {
self.leave_cleanups(depth);
self.jump(block, range);
}
None => self.terminate(Terminator::Unreachable),
}
}
Stmt::Continue { id, range } => {
match self.loops.get(&id).map(|l| (l.cont, l.cont_cleanups)) {
Some((block, depth)) => {
self.leave_cleanups(depth);
self.jump(block, range);
}
None => self.terminate(Terminator::Unreachable),
}
}
Stmt::Return { value, range } => {
self.leave_cleanups(0);
self.terminate(Terminator::Return { value, range });
}
Stmt::Switch(_) => {
unreachable!("sema lowers every switch into a SwitchTree in CFG mode")
}
Stmt::Region(_) => {
unreachable!("a region is only built for a body that stays structured")
}
}
}
fn goto_ptr(&mut self, target: Expr, range: SourceRange) {
let Some(object) = self.goto_value else {
self.terminate(Terminator::Unreachable);
return;
};
let depth = self
.taken
.iter()
.map(|(id, _)| self.label_cleanups.get(id).copied().unwrap_or(0))
.min()
.unwrap_or(0);
self.leave_cleanups(depth.min(self.cleanups.len()));
let place = self.goto_value_place(object, range);
let ty = place.ty;
let folded = self
.table
.and_then(|table| crate::ir::table_read(&target, table))
.cloned();
let value = match folded {
Some(index) => Expr::new(
ExprKind::Binary {
op: BinOp::Add,
lhs: Box::new(Expr::new(ExprKind::Cast(Box::new(index)), ty, range)),
rhs: Box::new(Expr::new(ExprKind::Int(1), ty, range)),
},
ty,
range,
),
None => Expr::new(ExprKind::Cast(Box::new(target)), ty, range),
};
self.push(Stmt::Expr(Expr::new(
ExprKind::Assign {
place,
value: Box::new(value),
},
ty,
range,
)));
let dispatch = self.dispatch(object, range);
self.jump(dispatch, range);
}
fn goto_value_place(&self, object: ObjectId, range: SourceRange) -> Place {
Place {
kind: PlaceKind::Object(object),
ty: self.objects[object.0 as usize].ty,
is_const: false,
range,
}
}
fn dispatch(&mut self, object: ObjectId, range: SourceRange) -> BlockId {
if let Some(block) = self.dispatch {
return block;
}
let block = self.new_block();
let invalid = self.new_block();
self.blocks[invalid.index()].term = Terminator::InvalidTarget;
let place = self.goto_value_place(object, range);
let ty = place.ty;
let cases = self
.taken
.iter()
.enumerate()
.map(|(index, (_, target))| (CaseRange::single(index as i128 + 1), *target))
.collect();
self.blocks[block.index()].term = Terminator::Switch {
value: Expr::new(ExprKind::Load(place), ty, range),
cases,
default: invalid,
range,
};
self.dispatch = Some(block);
block
}
fn local(&mut self, object: ObjectId, init: Expr, explicit: bool) {
let info = &self.objects[object.0 as usize];
if !matches!(info.storage, Storage::Automatic) {
return;
}
let (ty, is_const, range, name) = (info.ty, info.is_const, info.range, info.name.clone());
let rust_name = self.unique_name(&name);
self.locals.push(Local { object, rust_name });
if !explicit {
return;
}
let place = Place {
kind: PlaceKind::Object(object),
ty,
is_const,
range,
};
self.push(Stmt::Expr(Expr::new(
ExprKind::Assign {
place,
value: Box::new(init),
},
ty,
range,
)));
}
fn vla(&mut self, def: crate::ir::VlaDef) {
for object in [def.storage, def.object] {
let name = self.objects[object.0 as usize].name.clone();
let rust_name = self.unique_name(&name);
self.locals.push(Local { object, rust_name });
}
self.push(Stmt::Vla(Box::new(def)));
}
fn unique_name(&mut self, name: &str) -> String {
if self.used.insert(name.to_owned()) {
return name.to_owned();
}
for n in 1u32.. {
let candidate = format!("{name}_{n}");
if self.used.insert(candidate.clone()) {
return candidate;
}
}
unreachable!("the loop above always terminates")
}
fn if_stmt(&mut self, cond: Expr, then_branch: Stmt, else_branch: Option<Stmt>) {
let range = cond.range;
let then_blk = self.new_block();
let else_blk = self.new_block();
let join = if else_branch.is_some() {
self.new_block()
} else {
else_blk
};
self.terminate(Terminator::Branch {
cond,
then_blk,
else_blk,
});
self.continue_at(then_blk);
self.stmt(then_branch);
self.jump(join, range);
if let Some(else_branch) = else_branch {
self.continue_at(else_blk);
self.stmt(else_branch);
self.jump(join, range);
}
self.continue_at(join);
}
fn while_stmt(&mut self, id: LoopId, cond: Expr, body: Stmt, range: SourceRange) {
let head = self.new_block();
let body_blk = self.new_block();
let exit = self.new_block();
self.jump(head, range);
self.continue_at(head);
self.test(cond, body_blk, exit, range);
let depth = self.cleanups.len();
self.loops.insert(
id,
LoopBlocks {
brk: exit,
brk_cleanups: depth,
cont: head,
cont_cleanups: depth,
},
);
self.continue_at(body_blk);
self.stmt(body);
self.jump(head, range);
self.continue_at(exit);
}
fn do_while(&mut self, id: LoopId, body: Stmt, cond: Expr, range: SourceRange) {
let body_blk = self.new_block();
let test = self.new_block();
let exit = self.new_block();
self.jump(body_blk, range);
let depth = self.cleanups.len();
self.loops.insert(
id,
LoopBlocks {
brk: exit,
brk_cleanups: depth,
cont: test,
cont_cleanups: depth,
},
);
self.continue_at(body_blk);
self.stmt(body);
self.jump(test, range);
self.continue_at(test);
self.test(cond, body_blk, exit, range);
self.continue_at(exit);
}
fn for_stmt(
&mut self,
id: LoopId,
init: Vec<Stmt>,
cond: Option<Expr>,
step: Option<Expr>,
body: Stmt,
range: SourceRange,
) {
let outer = self.cleanups.len();
self.stmts(init);
let inner = self.cleanups.len();
let head = self.new_block();
let body_blk = self.new_block();
let step_blk = self.new_block();
let exit = self.new_block();
self.jump(head, range);
self.continue_at(head);
match cond {
Some(cond) => self.test(cond, body_blk, exit, range),
None => self.jump(body_blk, range),
}
self.loops.insert(
id,
LoopBlocks {
brk: exit,
brk_cleanups: inner,
cont: step_blk,
cont_cleanups: inner,
},
);
self.continue_at(body_blk);
self.stmt(body);
self.jump(step_blk, range);
self.continue_at(step_blk);
if let Some(step) = step {
self.push(Stmt::Expr(step));
}
self.jump(head, range);
self.continue_at(exit);
self.leave_cleanups(outer);
self.cleanups.truncate(outer);
}
fn test(&mut self, cond: Expr, then_blk: BlockId, else_blk: BlockId, range: SourceRange) {
if is_always_true(&cond) {
self.jump(then_blk, range);
return;
}
self.terminate(Terminator::Branch {
cond,
then_blk,
else_blk,
});
}
fn switch(&mut self, switch: crate::ir::SwitchTree) {
let crate::ir::SwitchTree {
id,
scrutinee,
body,
range,
} = switch;
let dispatch = self.current();
let exit = self.new_block();
self.switches.insert(
id,
SwitchFrame {
brk: exit,
brk_cleanups: self.cleanups.len(),
cases: Vec::new(),
default: None,
},
);
self.current = None;
self.stmt(*body);
self.jump(exit, range);
let frame = self
.switches
.remove(&id)
.expect("the frame was just inserted");
self.blocks[dispatch.index()].term = Terminator::Switch {
value: scrutinee,
cases: frame.cases,
default: frame.default.unwrap_or(exit),
range,
};
self.continue_at(exit);
}
fn case(&mut self, switch: SwitchId, value: Option<CaseRange>, body: Stmt, range: SourceRange) {
let block = self.new_block();
self.jump(block, range);
self.continue_at(block);
if let Some(frame) = self.switches.get_mut(&switch) {
match value {
Some(value) => frame.cases.push((value, block)),
None => frame.default = Some(block),
}
}
self.stmt(body);
}
fn finish(
mut self,
entry: BlockId,
labels: HashMap<LabelId, u32>,
names: &HashMap<LabelId, String>,
) -> Cfg {
let entry = self.thread_jumps(entry);
self.merge_chains(entry);
self.renumber(entry, labels, names)
}
fn thread_jumps(&mut self, entry: BlockId) -> BlockId {
let resolved: Vec<BlockId> = (0..self.blocks.len())
.map(|index| self.resolve(BlockId(index as u32)))
.collect();
for block in &mut self.blocks {
block.term.map_targets(|target| resolved[target.index()]);
}
for block in self.labels.values_mut() {
*block = resolved[block.index()];
}
resolved[entry.index()]
}
fn resolve(&self, mut block: BlockId) -> BlockId {
let mut seen = HashSet::new();
while seen.insert(block) {
let candidate = &self.blocks[block.index()];
if !candidate.stmts.is_empty() {
break;
}
match candidate.term {
Terminator::Jump { target, .. } if target != block => block = target,
_ => break,
}
}
block
}
fn merge_chains(&mut self, entry: BlockId) {
let reachable = self.reachable(entry);
let mut predecessors = vec![0usize; self.blocks.len()];
for id in &reachable {
for successor in self.blocks[id.index()].term.successors() {
predecessors[successor.index()] += 1;
}
}
for id in &reachable {
while let Terminator::Jump { target, .. } = self.blocks[id.index()].term {
if target == *id || target == entry || predecessors[target.index()] != 1 {
break;
}
let mut stmts = std::mem::take(&mut self.blocks[target.index()].stmts);
let term = std::mem::replace(
&mut self.blocks[target.index()].term,
Terminator::Unreachable,
);
self.blocks[id.index()].stmts.append(&mut stmts);
self.blocks[id.index()].term = term;
predecessors[target.index()] = 0;
}
}
}
fn reachable(&self, entry: BlockId) -> Vec<BlockId> {
let mut seen = HashSet::new();
let mut stack = vec![entry];
let mut out = Vec::new();
while let Some(id) = stack.pop() {
if !seen.insert(id) {
continue;
}
out.push(id);
stack.extend(self.blocks[id.index()].term.successors());
}
out.sort_unstable();
out
}
fn renumber(
mut self,
entry: BlockId,
labels: HashMap<LabelId, u32>,
names: &HashMap<LabelId, String>,
) -> Cfg {
let mut order = Vec::with_capacity(self.blocks.len());
let mut visited = vec![false; self.blocks.len()];
let mut stack = vec![(entry, 0usize)];
visited[entry.index()] = true;
while let Some((id, next)) = stack.pop() {
let successors = self.blocks[id.index()].term.successors();
if next < successors.len() {
stack.push((id, next + 1));
let successor = successors[successors.len() - 1 - next];
if !visited[successor.index()] {
visited[successor.index()] = true;
stack.push((successor, 0));
}
continue;
}
order.push(id);
}
order.reverse();
let mut index_of = vec![None; self.blocks.len()];
for (index, id) in order.iter().enumerate() {
index_of[id.index()] = Some(BlockId(index as u32));
}
let mut blocks = Vec::with_capacity(order.len());
for id in order {
let mut block = std::mem::replace(
&mut self.blocks[id.index()],
BasicBlock {
stmts: Vec::new(),
term: Terminator::Unreachable,
},
);
block.term.map_targets(|target| {
index_of[target.index()].expect("a reachable block only names reachable blocks")
});
blocks.push(block);
}
let mut named: Vec<(LabelId, BlockId)> =
self.labels.iter().map(|(a, b)| (*a, *b)).collect();
named.sort_unstable();
let mut block_labels: HashMap<BlockId, String> = HashMap::new();
for (id, block) in named {
let (Some(name), Some(numbered)) = (names.get(&id), index_of[block.index()]) else {
continue;
};
block_labels.entry(numbered).or_insert_with(|| name.clone());
}
if let Some(numbered) = self.dispatch.and_then(|block| index_of[block.index()]) {
block_labels
.entry(numbered)
.or_insert_with(|| "dispatch".to_owned());
}
let shape = crate::reloop::plan(&blocks, &block_labels);
Cfg {
locals: self.locals,
blocks,
labels,
shape,
}
}
}