use std::collections::{HashMap, HashSet};
use rowan::TextRange;
use rowan::ast::AstNode as _;
use smol_str::SmolStr;
use crate::ast::{AstToken as _, CallExpr};
use crate::syntax::{NodePtr, SyntaxElement, SyntaxKind, SyntaxNode};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct BlockId(u32);
impl BlockId {
pub fn index(self) -> usize {
self.0 as usize
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BasicBlock {
pub stmts: Vec<TextRange>,
pub terminator: Terminator,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Terminator {
Goto(BlockId),
Branch {
then_blk: BlockId,
else_blk: BlockId,
},
Return,
Diverge,
Unreachable,
}
impl Terminator {
fn successors(self) -> impl Iterator<Item = BlockId> {
let (a, b) = match self {
Terminator::Goto(target) => (Some(target), None),
Terminator::Branch { then_blk, else_blk } => (Some(then_blk), Some(else_blk)),
Terminator::Return | Terminator::Diverge | Terminator::Unreachable => (None, None),
};
a.into_iter().chain(b)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ControlFlowGraph {
blocks: Vec<BasicBlock>,
entry: BlockId,
reachable: Vec<bool>,
}
impl Default for ControlFlowGraph {
fn default() -> Self {
Self {
blocks: vec![BasicBlock {
stmts: Vec::new(),
terminator: Terminator::Return,
}],
entry: BlockId(0),
reachable: vec![true],
}
}
}
impl ControlFlowGraph {
pub fn blocks(&self) -> &[BasicBlock] {
&self.blocks
}
pub fn iter(&self) -> impl Iterator<Item = (BlockId, &BasicBlock)> {
self.blocks
.iter()
.enumerate()
.map(|(i, block)| (BlockId(i as u32), block))
}
pub fn entry(&self) -> BlockId {
self.entry
}
pub fn block(&self, id: BlockId) -> &BasicBlock {
&self.blocks[id.index()]
}
pub fn is_reachable(&self, id: BlockId) -> bool {
self.reachable[id.index()]
}
fn build_region(
stmts: &[SyntaxElement],
labels: HashSet<SmolStr>,
opaque_macros: bool,
) -> Self {
let mut builder = Builder {
blocks: Vec::new(),
labels,
label_blocks: HashMap::new(),
};
let entry = builder.new_block();
if let Some(exit) = builder.lower_seq(stmts, entry, None) {
builder.set_term(exit, Terminator::Return);
}
let mut roots = vec![entry];
if opaque_macros {
roots.extend(builder.label_blocks.values().copied());
}
let reachable = reachable_from(&builder.blocks, &roots);
ControlFlowGraph {
blocks: builder.blocks,
entry,
reachable,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct FileControlFlow {
toplevel: ControlFlowGraph,
regions: Vec<(NodePtr, ControlFlowGraph)>,
unreachable: HashSet<TextRange>,
}
impl FileControlFlow {
pub fn build(root: &SyntaxNode) -> Self {
let toplevel = build_region_of(root);
let regions: Vec<_> = root
.descendants()
.filter(|node| is_region_owner(node.kind()))
.filter(|node| body_block(node).is_some())
.map(|node| (NodePtr::new(&node), build_region_of(&node)))
.collect();
let mut this = Self {
toplevel,
regions,
unreachable: HashSet::new(),
};
this.unreachable = this
.graphs()
.flat_map(|cfg| {
cfg.iter()
.filter(|(id, _)| !cfg.is_reachable(*id))
.flat_map(|(_, block)| block.stmts.iter().copied())
})
.collect();
this
}
pub fn toplevel(&self) -> &ControlFlowGraph {
&self.toplevel
}
pub fn regions(&self) -> &[(NodePtr, ControlFlowGraph)] {
&self.regions
}
pub fn region(&self, ptr: NodePtr) -> Option<&ControlFlowGraph> {
self.regions
.iter()
.find(|(p, _)| *p == ptr)
.map(|(_, cfg)| cfg)
}
fn graphs(&self) -> impl Iterator<Item = &ControlFlowGraph> {
std::iter::once(&self.toplevel).chain(self.regions.iter().map(|(_, cfg)| cfg))
}
pub fn is_unreachable(&self, range: TextRange) -> bool {
self.unreachable.contains(&range)
}
pub fn render(&self, src: &str) -> String {
let mut out = String::new();
out.push_str("region: <toplevel>\n");
render_region(&self.toplevel, src, &mut out);
for (ptr, cfg) in &self.regions {
let head = snippet(src, ptr.text_range());
out.push_str(&format!("region: {head}\n"));
render_region(cfg, src, &mut out);
}
out
}
}
fn build_region_of(owner: &SyntaxNode) -> ControlFlowGraph {
let body = if owner.kind() == SyntaxKind::ROOT {
owner.clone()
} else {
body_block(owner).expect("region owner has a body block")
};
let mut labels = HashSet::new();
collect_labels(&body, &mut labels);
let opaque_macros = contains_opaque_macro(&body);
ControlFlowGraph::build_region(®ion_statements(&body), labels, opaque_macros)
}
fn is_region_owner(kind: SyntaxKind) -> bool {
matches!(
kind,
SyntaxKind::FUNCTION_DEF
| SyntaxKind::MACRO_DEF
| SyntaxKind::DO_EXPR
| SyntaxKind::MODULE_DEF
)
}
fn body_block(node: &SyntaxNode) -> Option<SyntaxNode> {
node.children().find(|c| c.kind() == SyntaxKind::BLOCK)
}
struct Builder {
blocks: Vec<BasicBlock>,
labels: HashSet<SmolStr>,
label_blocks: HashMap<SmolStr, BlockId>,
}
#[derive(Clone, Copy)]
struct LoopCtx {
header: BlockId,
after: BlockId,
}
enum Jump {
Diverge,
Break,
Continue,
Goto(SmolStr),
}
impl Builder {
fn new_block(&mut self) -> BlockId {
let id = BlockId(u32::try_from(self.blocks.len()).expect("block count fits in u32"));
self.blocks.push(BasicBlock {
stmts: Vec::new(),
terminator: Terminator::Return,
});
id
}
fn set_term(&mut self, block: BlockId, terminator: Terminator) {
self.blocks[block.index()].terminator = terminator;
}
fn push_stmt(&mut self, block: BlockId, range: TextRange) {
self.blocks[block.index()].stmts.push(range);
}
fn has_predecessor(&self, target: BlockId) -> bool {
self.blocks
.iter()
.any(|block| block.terminator.successors().any(|s| s == target))
}
fn label_block(&mut self, name: &SmolStr) -> BlockId {
if let Some(id) = self.label_blocks.get(name) {
return *id;
}
let id = self.new_block();
self.label_blocks.insert(name.clone(), id);
id
}
fn jump_terminator(&mut self, jump: &Jump, loop_ctx: Option<LoopCtx>) -> Option<Terminator> {
match jump {
Jump::Diverge => Some(Terminator::Diverge),
Jump::Break => loop_ctx.map(|lc| Terminator::Goto(lc.after)),
Jump::Continue => loop_ctx.map(|lc| Terminator::Goto(lc.header)),
Jump::Goto(name) => self
.labels
.contains(name)
.then(|| Terminator::Goto(self.label_block(name))),
}
}
fn lower_seq(
&mut self,
stmts: &[SyntaxElement],
mut cur: BlockId,
loop_ctx: Option<LoopCtx>,
) -> Option<BlockId> {
let mut i = 0;
while i < stmts.len() {
match self.lower_stmt(&stmts[i], cur, loop_ctx) {
Some(next) => {
cur = next;
i += 1;
}
None => {
let rest = &stmts[i + 1..];
let resume = rest
.iter()
.position(|stmt| label_definition(stmt).is_some());
let dead = &rest[..resume.unwrap_or(rest.len())];
if !dead.is_empty() {
let dead_blk = self.new_block();
self.set_term(dead_blk, Terminator::Unreachable);
for stmt in dead {
self.push_stmt(dead_blk, stmt.text_range());
}
}
let offset = resume?;
let label = &rest[offset];
let name = label_definition(label).expect("resume lands on a label");
cur = self.label_block(&name);
self.push_stmt(cur, label.text_range());
i += offset + 2;
}
}
}
Some(cur)
}
fn lower_stmt(
&mut self,
stmt: &SyntaxElement,
cur: BlockId,
loop_ctx: Option<LoopCtx>,
) -> Option<BlockId> {
let Some(node) = stmt.as_node() else {
self.push_stmt(cur, stmt.text_range());
return Some(cur);
};
if let Some(name) = label_definition(stmt) {
let target = self.label_block(&name);
self.set_term(cur, Terminator::Goto(target));
self.push_stmt(target, node.text_range());
return Some(target);
}
if let Some(jump) = jump_of(node) {
self.push_stmt(cur, node.text_range());
return match self.jump_terminator(&jump, loop_ctx) {
Some(terminator) => {
self.set_term(cur, terminator);
None
}
None => Some(cur),
};
}
match node.kind() {
SyntaxKind::BLOCK => self.lower_seq(®ion_statements(node), cur, loop_ctx),
SyntaxKind::BEGIN_EXPR | SyntaxKind::LET_EXPR => match body_block(node) {
Some(body) => self.lower_seq(®ion_statements(&body), cur, loop_ctx),
None => Some(cur),
},
SyntaxKind::IF_EXPR => {
self.push_stmt(cur, node.text_range());
let (guarded, unguarded) = if_arms(node);
self.lower_if(&guarded, unguarded.as_ref(), cur, loop_ctx)
}
SyntaxKind::FOR_EXPR | SyntaxKind::WHILE_EXPR => {
self.push_stmt(cur, node.text_range());
let body = body_block(node);
let never_exits = node.kind() == SyntaxKind::WHILE_EXPR
&& has_literal_true_test(node)
&& !body.as_ref().is_some_and(contains_opaque_macro);
let stmts = body.map(|b| region_statements(&b));
self.lower_loop(cur, stmts.as_deref().unwrap_or_default(), never_exits)
}
SyntaxKind::TRY_EXPR => {
self.push_stmt(cur, node.text_range());
self.lower_try(node, cur, loop_ctx)
}
SyntaxKind::BINARY_EXPR => {
self.push_stmt(cur, node.text_range());
match short_circuit_jump(node) {
Some((jump_node, jump)) => {
self.lower_conditional_jump(cur, jump_node.text_range(), &jump, loop_ctx)
}
None => Some(cur),
}
}
_ => {
self.push_stmt(cur, node.text_range());
Some(cur)
}
}
}
fn lower_conditional_jump(
&mut self,
cur: BlockId,
range: TextRange,
jump: &Jump,
loop_ctx: Option<LoopCtx>,
) -> Option<BlockId> {
let Some(terminator) = self.jump_terminator(jump, loop_ctx) else {
return Some(cur);
};
let taken = self.new_block();
let fallthrough = self.new_block();
self.set_term(
cur,
Terminator::Branch {
then_blk: taken,
else_blk: fallthrough,
},
);
self.push_stmt(taken, range);
self.set_term(taken, terminator);
Some(fallthrough)
}
fn lower_if(
&mut self,
guarded: &[Vec<SyntaxElement>],
unguarded: Option<&Vec<SyntaxElement>>,
cur: BlockId,
loop_ctx: Option<LoopCtx>,
) -> Option<BlockId> {
let Some((first, rest)) = guarded.split_first() else {
return Some(cur);
};
let then_blk = self.new_block();
let then_exit = self.lower_seq(first, then_blk, loop_ctx);
if rest.is_empty() && unguarded.is_none() {
let join = self.new_block();
self.set_term(
cur,
Terminator::Branch {
then_blk,
else_blk: join,
},
);
if let Some(exit) = then_exit {
self.set_term(exit, Terminator::Goto(join));
}
return Some(join);
}
let else_blk = self.new_block();
let else_exit = if rest.is_empty() {
self.lower_seq(unguarded.expect("checked above"), else_blk, loop_ctx)
} else {
self.lower_if(rest, unguarded, else_blk, loop_ctx)
};
self.set_term(cur, Terminator::Branch { then_blk, else_blk });
self.join(&[then_exit, else_exit])
}
fn lower_loop(
&mut self,
cur: BlockId,
body: &[SyntaxElement],
never_exits: bool,
) -> Option<BlockId> {
let header = self.new_block();
self.set_term(cur, Terminator::Goto(header));
let after = self.new_block();
let body_blk = self.new_block();
if never_exits {
self.set_term(header, Terminator::Goto(body_blk));
} else {
self.set_term(
header,
Terminator::Branch {
then_blk: body_blk,
else_blk: after,
},
);
}
let loop_ctx = LoopCtx { header, after };
if let Some(exit) = self.lower_seq(body, body_blk, Some(loop_ctx)) {
self.set_term(exit, Terminator::Goto(header)); }
if never_exits && !self.has_predecessor(after) {
self.set_term(after, Terminator::Unreachable);
None
} else {
Some(after)
}
}
fn lower_try(
&mut self,
node: &SyntaxNode,
cur: BlockId,
loop_ctx: Option<LoopCtx>,
) -> Option<BlockId> {
let catch = clause_body(node, SyntaxKind::CATCH_CLAUSE);
let finally = clause_body(node, SyntaxKind::FINALLY_CLAUSE);
let try_blk = self.new_block();
let catch_blk = catch.as_ref().map(|_| self.new_block());
let finally_blk = finally.as_ref().map(|_| self.new_block());
match (catch_blk, finally_blk) {
(Some(catch_blk), Some(finally_blk)) => {
let handler = self.new_block();
self.set_term(
cur,
Terminator::Branch {
then_blk: try_blk,
else_blk: handler,
},
);
self.set_term(
handler,
Terminator::Branch {
then_blk: catch_blk,
else_blk: finally_blk,
},
);
}
(Some(handler), None) | (None, Some(handler)) => self.set_term(
cur,
Terminator::Branch {
then_blk: try_blk,
else_blk: handler,
},
),
(None, None) => self.set_term(cur, Terminator::Goto(try_blk)),
}
let try_body = body_block(node)
.map(|b| region_statements(&b))
.unwrap_or_default();
let try_exit = self.lower_seq(&try_body, try_blk, loop_ctx);
let catch_exit = catch
.zip(catch_blk)
.and_then(|(stmts, blk)| self.lower_seq(®ion_statements(&stmts), blk, loop_ctx));
match (finally, finally_blk) {
(Some(stmts), Some(blk)) => {
for exit in [try_exit, catch_exit].into_iter().flatten() {
self.set_term(exit, Terminator::Goto(blk));
}
let finally_exit = self.lower_seq(®ion_statements(&stmts), blk, loop_ctx);
self.join(&[finally_exit])
}
_ => self.join(&[try_exit, catch_exit]),
}
}
fn join(&mut self, exits: &[Option<BlockId>]) -> Option<BlockId> {
if exits.iter().all(Option::is_none) {
return None;
}
let join = self.new_block();
for exit in exits.iter().flatten() {
self.set_term(*exit, Terminator::Goto(join));
}
Some(join)
}
}
fn reachable_from(blocks: &[BasicBlock], roots: &[BlockId]) -> Vec<bool> {
let mut seen = vec![false; blocks.len()];
let mut stack = Vec::from(roots);
for root in roots {
seen[root.index()] = true;
}
while let Some(id) = stack.pop() {
for next in blocks[id.index()].terminator.successors() {
if !seen[next.index()] {
seen[next.index()] = true;
stack.push(next);
}
}
}
seen
}
fn region_statements(container: &SyntaxNode) -> Vec<SyntaxElement> {
container
.children_with_tokens()
.filter(|element| !is_ignorable(element.kind()))
.collect()
}
fn is_ignorable(kind: SyntaxKind) -> bool {
matches!(
kind,
SyntaxKind::WHITESPACE
| SyntaxKind::NEWLINE
| SyntaxKind::COMMENT
| SyntaxKind::BLOCK_COMMENT
| SyntaxKind::SEMICOLON
| SyntaxKind::TOPLEVEL_SEMICOLON
)
}
fn if_arms(node: &SyntaxNode) -> (Vec<Vec<SyntaxElement>>, Option<Vec<SyntaxElement>>) {
let mut guarded = Vec::new();
let mut unguarded = None;
if let Some(body) = body_block(node) {
guarded.push(region_statements(&body));
}
for child in node.children() {
match child.kind() {
SyntaxKind::ELSEIF_CLAUSE => {
if let Some(body) = body_block(&child) {
guarded.push(region_statements(&body));
}
}
SyntaxKind::ELSE_CLAUSE => {
unguarded = body_block(&child).map(|body| region_statements(&body));
}
_ => {}
}
}
(guarded, unguarded)
}
fn clause_body(node: &SyntaxNode, kind: SyntaxKind) -> Option<SyntaxNode> {
node.children()
.find(|c| c.kind() == kind)
.and_then(|clause| body_block(&clause))
}
fn has_literal_true_test(node: &SyntaxNode) -> bool {
node.children()
.find(|c| c.kind() == SyntaxKind::CONDITION)
.and_then(|cond| cond.children().next())
.is_some_and(|test| {
test.kind() == SyntaxKind::LITERAL
&& test
.children_with_tokens()
.filter_map(|e| e.into_token())
.any(|t| t.kind() == SyntaxKind::TRUE_KW)
})
}
fn jump_of(node: &SyntaxNode) -> Option<Jump> {
match node.kind() {
SyntaxKind::RETURN_EXPR => Some(Jump::Diverge),
SyntaxKind::BREAK_EXPR => Some(Jump::Break),
SyntaxKind::CONTINUE_EXPR => Some(Jump::Continue),
SyntaxKind::CALL_EXPR => {
let callee = CallExpr::cast(node.clone())?.callee_ident()?;
matches!(callee.text(), "throw" | "error" | "rethrow").then_some(Jump::Diverge)
}
SyntaxKind::MACRO_CALL => macro_with_label(node, "goto").map(Jump::Goto),
_ => None,
}
}
fn label_definition(stmt: &SyntaxElement) -> Option<SmolStr> {
let node = stmt.as_node()?;
(node.kind() == SyntaxKind::MACRO_CALL)
.then(|| macro_with_label(node, "label"))
.flatten()
}
fn macro_with_label(node: &SyntaxNode, macro_name: &str) -> Option<SmolStr> {
let name = node
.children()
.find(|c| c.kind() == SyntaxKind::MACRO_NAME)?;
let simple = name
.children_with_tokens()
.filter_map(|e| e.into_token())
.filter(|t| t.kind() == SyntaxKind::IDENT)
.last()?;
if simple.text() != macro_name {
return None;
}
let arg = node
.children()
.skip_while(|c| c.kind() != SyntaxKind::MACRO_NAME)
.skip(1)
.find(|c| !is_ignorable(c.kind()))?;
let label = match arg.kind() {
SyntaxKind::NAME => arg,
SyntaxKind::ARG_LIST => arg
.children()
.find(|c| c.kind() == SyntaxKind::ARG)?
.children()
.find(|c| c.kind() == SyntaxKind::NAME)?,
_ => return None,
};
label
.children_with_tokens()
.filter_map(|e| e.into_token())
.find(|t| t.kind() == SyntaxKind::IDENT)
.map(|t| SmolStr::new(t.text()))
}
fn contains_opaque_macro(body: &SyntaxNode) -> bool {
body.children().any(|child| {
if is_region_owner(child.kind()) {
return false;
}
let opaque = child.kind() == SyntaxKind::MACRO_CALL
&& macro_with_label(&child, "goto").is_none()
&& macro_with_label(&child, "label").is_none();
opaque || contains_opaque_macro(&child)
})
}
fn collect_labels(body: &SyntaxNode, out: &mut HashSet<SmolStr>) {
for child in body.children() {
if is_region_owner(child.kind()) {
continue;
}
if let Some(name) = macro_with_label(&child, "label") {
out.insert(name);
}
collect_labels(&child, out);
}
}
fn short_circuit_jump(node: &SyntaxNode) -> Option<(SyntaxNode, Jump)> {
if !node
.children_with_tokens()
.filter_map(|e| e.into_token())
.any(|t| matches!(t.kind(), SyntaxKind::AND_AND | SyntaxKind::OR_OR))
{
return None;
}
let rhs = node.children().nth(1)?;
if let Some(jump) = jump_of(&rhs) {
return Some((rhs, jump));
}
(rhs.kind() == SyntaxKind::BINARY_EXPR)
.then(|| short_circuit_jump(&rhs))
.flatten()
}
fn render_region(cfg: &ControlFlowGraph, src: &str, out: &mut String) {
for (id, block) in cfg.iter() {
let stmts = block
.stmts
.iter()
.map(|range| snippet(src, *range))
.collect::<Vec<_>>()
.join("; ");
let dead = if cfg.is_reachable(id) { "" } else { " (dead)" };
let term = match block.terminator {
Terminator::Goto(t) => format!("-> bb{}", t.index()),
Terminator::Branch { then_blk, else_blk } => format!(
"-> then bb{}, else bb{}",
then_blk.index(),
else_blk.index()
),
Terminator::Return => "-> return".to_string(),
Terminator::Diverge => "-> diverge".to_string(),
Terminator::Unreachable => "-> unreachable".to_string(),
};
let i = id.index();
out.push_str(&format!(" bb{i}{dead}: [{stmts}] {term}\n"));
}
}
fn snippet(src: &str, range: TextRange) -> String {
let flat = src[range].split_whitespace().collect::<Vec<_>>().join(" ");
match flat.char_indices().nth(40) {
Some((cut, _)) => format!("{}…", &flat[..cut]),
None => flat,
}
}