use std::collections::{HashMap, HashSet};
use crate::ast;
use crate::capture::SourceRange;
use crate::ir::{self, Function, LabelId, Region, RegionKind, Stmt};
struct Item<K> {
labels: Vec<K>,
gotos: Vec<K>,
declares: bool,
}
impl<K> Item<K> {
fn neutral() -> Self {
Self {
labels: Vec::new(),
gotos: Vec::new(),
declares: false,
}
}
}
#[derive(Clone, Copy)]
struct Need<K> {
id: usize,
label: K,
kind: RegionKind,
anchor: usize,
first: usize,
last: usize,
}
impl<K> Need<K> {
fn span(&self) -> (usize, usize) {
match self.kind {
RegionKind::Block => (self.first, self.anchor),
RegionKind::Loop => (self.anchor, self.last + 1),
}
}
}
enum Node<K> {
Item(usize),
Region {
label: K,
kind: RegionKind,
body: Vec<Node<K>>,
},
}
fn plan<K: Copy + Eq>(items: &[Item<K>]) -> Option<Vec<Node<K>>> {
let mut needs: Vec<Need<K>> = Vec::new();
for (anchor, item) in items.iter().enumerate() {
for label in &item.labels {
let mut backward: Option<(usize, usize)> = None;
let mut forward: Option<(usize, usize)> = None;
for (at, other) in items.iter().enumerate() {
if !other.gotos.contains(label) {
continue;
}
let side = if at < anchor {
&mut forward
} else {
&mut backward
};
match side {
Some((first, last)) => {
*first = (*first).min(at);
*last = (*last).max(at);
}
None => *side = Some((at, at)),
}
}
for (kind, found) in [(RegionKind::Block, forward), (RegionKind::Loop, backward)] {
if let Some((first, last)) = found {
needs.push(Need {
id: needs.len(),
label: *label,
kind,
anchor,
first,
last,
});
}
}
}
}
build(items, 0, items.len(), &needs)
}
fn build<K: Copy + Eq>(
items: &[Item<K>],
lo: usize,
hi: usize,
needs: &[Need<K>],
) -> Option<Vec<Node<K>>> {
let mut active: Vec<Need<K>> = Vec::new();
for need in needs {
let (start, end) = need.span();
if end <= lo || start >= hi {
continue;
}
if start < lo || end > hi {
return None;
}
active.push(*need);
}
if active.is_empty() {
return Some((lo..hi).map(Node::Item).collect());
}
let blocks = || active.iter().filter(|need| need.kind == RegionKind::Block);
let loops = || active.iter().filter(|need| need.kind == RegionKind::Loop);
if let Some(chosen) = blocks().max_by_key(|need| need.anchor).copied() {
let start = blocks()
.map(|need| need.first)
.min()
.expect("chosen is one");
if splits(items, lo, hi, start, chosen.anchor, chosen, &active) {
return assemble(items, lo, hi, start, chosen.anchor, chosen, needs);
}
}
if let Some(chosen) = loops().min_by_key(|need| need.anchor).copied() {
let mut end = loops()
.map(|need| need.last + 1)
.max()
.expect("chosen is one");
for need in blocks().filter(|need| need.first >= chosen.anchor) {
end = end.max(need.anchor);
}
if splits(items, lo, hi, chosen.anchor, end, chosen, &active) {
return assemble(items, lo, hi, chosen.anchor, end, chosen, needs);
}
}
None
}
#[allow(clippy::too_many_arguments)]
fn splits<K: Copy + Eq>(
items: &[Item<K>],
lo: usize,
hi: usize,
start: usize,
end: usize,
chosen: Need<K>,
active: &[Need<K>],
) -> bool {
if items[start..end].iter().any(|item| item.declares) {
return false;
}
let part = |need: &Need<K>| {
let (from, to) = need.span();
(from >= lo && to <= start) || (from >= start && to <= end) || (from >= end && to <= hi)
};
active.iter().all(|need| need.id == chosen.id || part(need))
}
#[allow(clippy::too_many_arguments)]
fn assemble<K: Copy + Eq>(
items: &[Item<K>],
lo: usize,
hi: usize,
start: usize,
end: usize,
chosen: Need<K>,
needs: &[Need<K>],
) -> Option<Vec<Node<K>>> {
let rest: Vec<Need<K>> = needs
.iter()
.filter(|need| need.id != chosen.id)
.copied()
.collect();
let mut out = build(items, lo, start, &rest)?;
let body = build(items, start, end, &rest)?;
out.push(Node::Region {
label: chosen.label,
kind: chosen.kind,
body,
});
out.extend(build(items, end, hi, &rest)?);
Some(out)
}
fn planned<K: Copy + Eq + std::hash::Hash>(nodes: &[Node<K>], out: &mut HashSet<K>) {
for node in nodes {
if let Node::Region { label, body, .. } = node {
out.insert(*label);
planned(body, out);
}
}
}
pub fn analyze(body: &ast::Block) -> Option<HashSet<String>> {
let mut scan = Scan {
hosts: Vec::new(),
targets: HashSet::new(),
};
scan.list(&entries(&body.items))?;
Some(scan.targets)
}
enum Entry<'a> {
Stmt(&'a ast::Stmt),
Decl,
Nothing,
}
fn entries(items: &[ast::BlockItem]) -> Vec<Entry<'_>> {
items
.iter()
.map(|item| match item {
ast::BlockItem::Stmt(stmt) => Entry::Stmt(stmt),
ast::BlockItem::Decl(_) => Entry::Decl,
ast::BlockItem::StaticAssert(_) | ast::BlockItem::NestedFunction(_) => Entry::Nothing,
})
.collect()
}
struct Scan<'a> {
hosts: Vec<HashSet<&'a str>>,
targets: HashSet<String>,
}
impl<'a> Scan<'a> {
fn list(&mut self, entries: &[Entry<'a>]) -> Option<()> {
let mut hosted: HashSet<&'a str> = HashSet::new();
for entry in entries {
if let Entry::Stmt(stmt) = entry {
for label in labels_of(stmt) {
hosted.insert(label);
}
}
}
let items: Vec<Item<&'a str>> = entries
.iter()
.map(|entry| match entry {
Entry::Stmt(stmt) => {
let stmt: &'a ast::Stmt = stmt;
let mut gotos = Vec::new();
gotos_of(stmt, &hosted, &mut gotos);
Item {
labels: labels_of(stmt),
gotos,
declares: false,
}
}
Entry::Decl => Item {
declares: true,
..Item::neutral()
},
Entry::Nothing => Item::neutral(),
})
.collect();
let plan = plan(&items)?;
let mut built: HashSet<&'a str> = HashSet::new();
planned(&plan, &mut built);
self.targets
.extend(built.into_iter().map(ToOwned::to_owned));
self.hosts.push(hosted);
for entry in entries {
if let Entry::Stmt(stmt) = entry {
self.stmt(stmt)?;
}
}
self.hosts.pop();
Some(())
}
fn knows(&self, name: &str) -> bool {
self.hosts.iter().any(|hosted| hosted.contains(name))
}
fn stmt(&mut self, stmt: &'a ast::Stmt) -> Option<()> {
match &stmt.kind {
ast::StmtKind::Goto(label) => self.knows(&label.name).then_some(()),
ast::StmtKind::GotoPtr(_) => None,
ast::StmtKind::Compound(block) => self.list(&entries(&block.items)),
ast::StmtKind::Labeled { body, .. }
| ast::StmtKind::Case { body, .. }
| ast::StmtKind::Default { body } => self.stmt(body),
ast::StmtKind::If {
then_branch,
else_branch,
..
} => {
self.stmt(then_branch)?;
match else_branch {
Some(branch) => self.stmt(branch),
None => Some(()),
}
}
ast::StmtKind::While { body, .. }
| ast::StmtKind::DoWhile { body, .. }
| ast::StmtKind::For { body, .. } => self.stmt(body),
ast::StmtKind::Switch { body, .. } => match &body.kind {
ast::StmtKind::Compound(block) => {
for item in &block.items {
if let ast::BlockItem::Stmt(stmt) = item {
self.stmt(stmt)?;
}
}
Some(())
}
_ => self.stmt(body),
},
_ => Some(()),
}
}
}
fn labels_of(stmt: &ast::Stmt) -> Vec<&str> {
let mut out = Vec::new();
let mut current = stmt;
loop {
match ¤t.kind {
ast::StmtKind::Labeled { label, body } => {
out.push(label.name.as_str());
current = body;
}
ast::StmtKind::Case { body, .. } | ast::StmtKind::Default { body } => current = body,
_ => return out,
}
}
}
fn gotos_of<'a>(stmt: &'a ast::Stmt, hosted: &HashSet<&'a str>, out: &mut Vec<&'a str>) {
match &stmt.kind {
ast::StmtKind::Goto(label) => {
if let Some(name) = hosted.get(label.name.as_str())
&& !out.contains(name)
{
out.push(*name);
}
}
ast::StmtKind::Labeled { body, .. }
| ast::StmtKind::Case { body, .. }
| ast::StmtKind::Default { body }
| ast::StmtKind::While { body, .. }
| ast::StmtKind::DoWhile { body, .. }
| ast::StmtKind::For { body, .. }
| ast::StmtKind::Switch { body, .. } => gotos_of(body, hosted, out),
ast::StmtKind::If {
then_branch,
else_branch,
..
} => {
gotos_of(then_branch, hosted, out);
if let Some(branch) = else_branch {
gotos_of(branch, hosted, out);
}
}
ast::StmtKind::Compound(block) => {
for item in &block.items {
if let ast::BlockItem::Stmt(stmt) = item {
gotos_of(stmt, hosted, out);
}
}
}
_ => {}
}
}
pub fn restructure(
body: &mut Vec<Stmt>,
names: &HashMap<LabelId, String>,
functions: &[Function],
) -> bool {
let mut rebuild = Rebuild {
names,
functions,
planned: true,
};
rebuild.list(body);
rebuild.planned
}
struct Rebuild<'a> {
names: &'a HashMap<LabelId, String>,
functions: &'a [Function],
planned: bool,
}
impl Rebuild<'_> {
fn list(&mut self, stmts: &mut Vec<Stmt>) {
for stmt in stmts.iter_mut() {
self.stmt(stmt);
}
if !stmts.iter().any(|stmt| !ir_labels_of(stmt).is_empty()) {
return;
}
let items: Vec<Item<LabelId>> = stmts
.iter()
.map(|stmt| {
let mut gotos = Vec::new();
ir_gotos_of(stmt, &mut gotos);
Item {
labels: ir_labels_of(stmt),
gotos,
declares: matches!(stmt, Stmt::Let { .. } | Stmt::Vla(_) | Stmt::Cleanup(_)),
}
})
.collect();
let Some(nodes) = plan(&items) else {
self.planned = false;
return;
};
let mut ranges: HashMap<LabelId, SourceRange> = HashMap::new();
for stmt in stmts.iter() {
ir_label_ranges(stmt, &mut ranges);
}
let mut slots: Vec<Option<Stmt>> = std::mem::take(stmts).into_iter().map(Some).collect();
*stmts = self.apply(nodes, &mut slots, &ranges);
}
fn apply(
&self,
nodes: Vec<Node<LabelId>>,
slots: &mut Vec<Option<Stmt>>,
ranges: &HashMap<LabelId, SourceRange>,
) -> Vec<Stmt> {
let mut out = Vec::with_capacity(nodes.len());
for node in nodes {
match node {
Node::Item(at) => out.push(slots[at].take().expect("one node per statement")),
Node::Region { label, kind, body } => {
let body = self.apply(body, slots, ranges);
let falls_out =
kind == RegionKind::Loop && !ir::always_terminates(&body, self.functions);
out.push(Stmt::Region(Box::new(Region {
label,
name: self.names.get(&label).cloned().unwrap_or_default(),
kind,
body,
falls_out,
range: ranges.get(&label).copied().unwrap_or(SourceRange::at(0)),
})));
}
}
}
out
}
fn stmt(&mut self, stmt: &mut Stmt) {
match stmt {
Stmt::Block(items) => self.list(items),
Stmt::Label { body, .. } | Stmt::Case { body, .. } => self.stmt(body),
Stmt::If {
then_branch,
else_branch,
..
} => {
self.stmt(then_branch);
if let Some(branch) = else_branch {
self.stmt(branch);
}
}
Stmt::While { body, .. } | Stmt::DoWhile { body, .. } => self.stmt(body),
Stmt::For { init, body, .. } => {
for stmt in init {
self.stmt(stmt);
}
self.stmt(body);
}
Stmt::Switch(switch) => {
for stmt in &mut switch.prelude {
self.stmt(stmt);
}
for group in &mut switch.groups {
for stmt in &mut group.body {
self.stmt(stmt);
}
}
}
_ => {}
}
}
}
fn ir_labels_of(stmt: &Stmt) -> Vec<LabelId> {
let mut out = Vec::new();
let mut current = stmt;
loop {
match current {
Stmt::Label { id, body, .. } => {
out.push(*id);
current = body;
}
Stmt::Case { body, .. } => current = body,
_ => return out,
}
}
}
fn ir_label_ranges(stmt: &Stmt, out: &mut HashMap<LabelId, SourceRange>) {
let mut current = stmt;
loop {
match current {
Stmt::Label { id, body, range } => {
out.insert(*id, *range);
current = body;
}
Stmt::Case { body, .. } => current = body,
_ => return,
}
}
}
fn ir_gotos_of(stmt: &Stmt, out: &mut Vec<LabelId>) {
match stmt {
Stmt::Goto { id, .. } => {
if !out.contains(id) {
out.push(*id);
}
}
Stmt::Block(items) => {
for stmt in items {
ir_gotos_of(stmt, out);
}
}
Stmt::Region(region) => {
for stmt in ®ion.body {
ir_gotos_of(stmt, out);
}
}
Stmt::Label { body, .. }
| Stmt::Case { body, .. }
| Stmt::While { body, .. }
| Stmt::DoWhile { body, .. } => ir_gotos_of(body, out),
Stmt::If {
then_branch,
else_branch,
..
} => {
ir_gotos_of(then_branch, out);
if let Some(branch) = else_branch {
ir_gotos_of(branch, out);
}
}
Stmt::For { init, body, .. } => {
for stmt in init {
ir_gotos_of(stmt, out);
}
ir_gotos_of(body, out);
}
Stmt::Switch(switch) => {
for stmt in switch
.prelude
.iter()
.chain(switch.groups.iter().flat_map(|group| group.body.iter()))
{
ir_gotos_of(stmt, out);
}
}
Stmt::SwitchTree(switch) => ir_gotos_of(&switch.body, out),
_ => {}
}
}