use super::cfg::ControlFlowGraph;
use super::types::*;
use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone)]
pub struct LivenessResult {
pub live_in: HashMap<BasicBlockId, HashSet<SlotId>>,
pub live_out: HashMap<BasicBlockId, HashSet<SlotId>>,
}
impl LivenessResult {
pub fn is_live_after(
&self,
block: BasicBlockId,
stmt_idx: usize,
slot: SlotId,
mir: &MirFunction,
) -> bool {
let bb = mir.block(block);
let mut live = self.live_out.get(&block).cloned().unwrap_or_default();
for i in (stmt_idx + 1..bb.statements.len()).rev() {
let stmt = &bb.statements[i];
update_liveness_for_statement(&mut live, &stmt.kind);
}
add_terminator_uses(&mut live, &bb.terminator.kind);
live.contains(&slot)
}
pub fn is_live_at_entry(&self, block: BasicBlockId, slot: SlotId) -> bool {
self.live_in
.get(&block)
.map_or(false, |set| set.contains(&slot))
}
}
pub fn compute_liveness(mir: &MirFunction, cfg: &ControlFlowGraph) -> LivenessResult {
let mut live_in: HashMap<BasicBlockId, HashSet<SlotId>> = HashMap::new();
let mut live_out: HashMap<BasicBlockId, HashSet<SlotId>> = HashMap::new();
for block in &mir.blocks {
live_in.insert(block.id, HashSet::new());
live_out.insert(block.id, HashSet::new());
}
let mut changed = true;
while changed {
changed = false;
let rpo = cfg.reverse_postorder();
for &block_id in rpo.iter().rev() {
let block = mir.block(block_id);
let mut new_live_out = HashSet::new();
for &succ in cfg.successors(block_id) {
if let Some(succ_in) = live_in.get(&succ) {
new_live_out.extend(succ_in);
}
}
let mut new_live_in = new_live_out.clone();
add_terminator_uses(&mut new_live_in, &block.terminator.kind);
for stmt in block.statements.iter().rev() {
update_liveness_for_statement(&mut new_live_in, &stmt.kind);
}
if new_live_in != *live_in.get(&block_id).unwrap_or(&HashSet::new()) {
changed = true;
live_in.insert(block_id, new_live_in);
}
if new_live_out != *live_out.get(&block_id).unwrap_or(&HashSet::new()) {
changed = true;
live_out.insert(block_id, new_live_out);
}
}
}
LivenessResult { live_in, live_out }
}
fn update_liveness_for_statement(live: &mut HashSet<SlotId>, kind: &StatementKind) {
match kind {
StatementKind::Assign(place, rvalue) => {
if let Place::Local(slot) = place {
live.remove(slot);
}
add_rvalue_uses(live, rvalue);
}
StatementKind::Drop(place) => {
live.insert(place.root_local());
}
StatementKind::TaskBoundary(operands, _kind) => {
for operand in operands {
add_operand_uses(live, operand);
}
}
StatementKind::ClosureCapture { operands, .. } => {
for operand in operands {
add_operand_uses(live, operand);
}
}
StatementKind::ArrayStore { operands, .. } => {
for operand in operands {
add_operand_uses(live, operand);
}
}
StatementKind::ObjectStore { operands, .. } => {
for operand in operands {
add_operand_uses(live, operand);
}
}
StatementKind::EnumStore { operands, .. } => {
for operand in operands {
add_operand_uses(live, operand);
}
}
StatementKind::Nop => {}
}
}
fn add_rvalue_uses(live: &mut HashSet<SlotId>, rvalue: &Rvalue) {
match rvalue {
Rvalue::Use(op) | Rvalue::Clone(op) | Rvalue::UnaryOp(_, op) => {
add_operand_uses(live, op);
}
Rvalue::Borrow(_, place) => {
live.insert(place.root_local());
}
Rvalue::BinaryOp(_, lhs, rhs) => {
add_operand_uses(live, lhs);
add_operand_uses(live, rhs);
}
Rvalue::Aggregate(ops) => {
for op in ops {
add_operand_uses(live, op);
}
}
Rvalue::EnumTest { operand, .. }
| Rvalue::EnumPayload { operand, .. }
| Rvalue::TypePatternTest { operand, .. }
| Rvalue::EnumDiscriminantTest { operand, .. } => {
add_operand_uses(live, operand);
}
}
}
fn add_operand_uses(live: &mut HashSet<SlotId>, op: &Operand) {
match op {
Operand::Copy(place) | Operand::Move(place) | Operand::MoveExplicit(place) => {
live.insert(place.root_local());
}
Operand::Constant(_) => {}
}
}
fn add_terminator_uses(live: &mut HashSet<SlotId>, kind: &TerminatorKind) {
match kind {
TerminatorKind::SwitchBool { operand, .. } => {
add_operand_uses(live, operand);
}
TerminatorKind::Call { func, args, .. } => {
add_operand_uses(live, func);
for arg in args {
add_operand_uses(live, arg);
}
}
TerminatorKind::Goto(_) | TerminatorKind::Return | TerminatorKind::Unreachable => {}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn span() -> shape_ast::ast::Span {
shape_ast::ast::Span { start: 0, end: 1 }
}
fn make_stmt(kind: StatementKind, point: u32) -> MirStatement {
MirStatement {
kind,
span: span(),
point: Point(point),
}
}
fn make_terminator(kind: TerminatorKind) -> Terminator {
Terminator { kind, span: span() }
}
#[test]
fn test_simple_liveness() {
let mir = MirFunction {
name: "test".to_string(),
blocks: vec![BasicBlock {
id: BasicBlockId(0),
statements: vec![
make_stmt(
StatementKind::Assign(
Place::Local(SlotId(0)),
Rvalue::Use(Operand::Constant(MirConstant::Int(1))),
),
0,
),
make_stmt(
StatementKind::Assign(
Place::Local(SlotId(1)),
Rvalue::Use(Operand::Copy(Place::Local(SlotId(0)))),
),
1,
),
],
terminator: make_terminator(TerminatorKind::Return),
}],
num_locals: 2,
param_slots: vec![],
param_reference_kinds: vec![],
local_types: vec![LocalTypeInfo::Copy, LocalTypeInfo::Copy],
span: span(),
field_name_table: std::collections::HashMap::new(),
local_struct_type_names: std::collections::HashMap::new(),
local_typed_array_element_types: std::collections::HashMap::new(),
local_declared_scalar_types: std::collections::HashMap::new(),
};
let cfg = ControlFlowGraph::build(&mir);
let liveness = compute_liveness(&mir, &cfg);
assert!(liveness.is_live_after(BasicBlockId(0), 0, SlotId(0), &mir));
assert!(!liveness.is_live_after(BasicBlockId(0), 1, SlotId(0), &mir));
}
#[test]
fn test_branch_liveness() {
let mir = MirFunction {
name: "test".to_string(),
blocks: vec![
BasicBlock {
id: BasicBlockId(0),
statements: vec![make_stmt(
StatementKind::Assign(
Place::Local(SlotId(0)),
Rvalue::Use(Operand::Constant(MirConstant::Int(1))),
),
0,
)],
terminator: make_terminator(TerminatorKind::SwitchBool {
operand: Operand::Copy(Place::Local(SlotId(2))),
true_bb: BasicBlockId(1),
false_bb: BasicBlockId(2),
}),
},
BasicBlock {
id: BasicBlockId(1),
statements: vec![make_stmt(
StatementKind::Assign(
Place::Local(SlotId(1)),
Rvalue::Use(Operand::Copy(Place::Local(SlotId(0)))),
),
1,
)],
terminator: make_terminator(TerminatorKind::Goto(BasicBlockId(3))),
},
BasicBlock {
id: BasicBlockId(2),
statements: vec![],
terminator: make_terminator(TerminatorKind::Goto(BasicBlockId(3))),
},
BasicBlock {
id: BasicBlockId(3),
statements: vec![],
terminator: make_terminator(TerminatorKind::Return),
},
],
num_locals: 3,
param_slots: vec![],
param_reference_kinds: vec![],
local_types: vec![
LocalTypeInfo::Copy,
LocalTypeInfo::Copy,
LocalTypeInfo::Copy,
],
span: span(),
field_name_table: std::collections::HashMap::new(),
local_struct_type_names: std::collections::HashMap::new(),
local_typed_array_element_types: std::collections::HashMap::new(),
local_declared_scalar_types: std::collections::HashMap::new(),
};
let cfg = ControlFlowGraph::build(&mir);
let liveness = compute_liveness(&mir, &cfg);
assert!(
liveness
.live_out
.get(&BasicBlockId(0))
.map_or(false, |s| s.contains(&SlotId(0)))
);
}
}