use std::collections::{HashMap, HashSet};
use shape_vm::mir::types::{
BasicBlockId, BinOp, MirFunction, Operand, Place, Rvalue, SlotId, StatementKind,
TerminatorKind,
};
#[derive(Debug, Clone, Default)]
pub struct BoundsElisionPlan {
pub trusted_pairs: HashSet<(SlotId, SlotId)>,
}
impl BoundsElisionPlan {
pub fn is_trusted(&self, arr: SlotId, iv: SlotId) -> bool {
self.trusted_pairs.contains(&(arr, iv))
}
}
pub fn analyze(mir: &MirFunction) -> BoundsElisionPlan {
let mut plan = BoundsElisionPlan::default();
let mut preds: HashMap<BasicBlockId, Vec<BasicBlockId>> = HashMap::new();
for block in &mir.blocks {
for succ in successors(&block.terminator.kind) {
preds.entry(succ).or_default().push(block.id);
}
}
let length_field_idxs: HashSet<u16> = mir
.field_name_table
.iter()
.filter_map(|(idx, name)| if name == "length" { Some(idx.0) } else { None })
.collect();
if length_field_idxs.is_empty() {
return plan;
}
for header in &mir.blocks {
let TerminatorKind::SwitchBool {
operand: pred_op,
true_bb,
false_bb: _false_bb,
} = &header.terminator.kind
else {
continue;
};
let Some(cond_slot) = operand_local(pred_op) else {
continue;
};
let Some((iv, bnd)) = find_lt_definition(header, cond_slot) else {
continue;
};
let header_preds = preds.get(&header.id).cloned().unwrap_or_default();
let has_back_edge = header_preds.iter().any(|p| p.0 >= header.id.0);
if !has_back_edge {
continue;
}
let Some(body) = mir.blocks.iter().find(|b| b.id == *true_bb) else {
continue;
};
let bnd_array_sources: Vec<SlotId> = mir
.blocks
.iter()
.flat_map(|b| b.statements.iter())
.filter_map(|stmt| {
let StatementKind::Assign(Place::Local(lhs), rvalue) = &stmt.kind else {
return None;
};
if *lhs != bnd {
return None;
}
let arr_slot = rvalue_field_length_source(rvalue, &length_field_idxs)?;
Some(arr_slot)
})
.collect();
if bnd_array_sources.is_empty() {
continue;
}
let arr = bnd_array_sources[0];
if bnd_array_sources.iter().any(|s| *s != arr) {
continue;
}
let is_param = mir.param_slots.contains(&arr);
let max_assigns = if is_param { 0 } else { 1 };
if slot_assignment_count(mir, arr) > max_assigns {
continue;
}
if block_assigns(body, bnd) {
continue;
}
if !iv_starts_non_negative(mir, iv, header.id) {
continue;
}
if !iv_only_monotonic_in_body(body, iv) {
continue;
}
plan.trusted_pairs.insert((arr, iv));
}
plan
}
fn successors(term: &TerminatorKind) -> Vec<BasicBlockId> {
match term {
TerminatorKind::Goto(b) => vec![*b],
TerminatorKind::SwitchBool { true_bb, false_bb, .. } => vec![*true_bb, *false_bb],
TerminatorKind::Call { next, .. } => vec![*next],
TerminatorKind::Return | TerminatorKind::Unreachable => vec![],
}
}
fn operand_local(op: &Operand) -> Option<SlotId> {
match op {
Operand::Copy(Place::Local(s))
| Operand::Move(Place::Local(s))
| Operand::MoveExplicit(Place::Local(s)) => Some(*s),
_ => None,
}
}
fn rvalue_field_length_source(
rvalue: &Rvalue,
length_field_idxs: &HashSet<u16>,
) -> Option<SlotId> {
let inner_op = match rvalue {
Rvalue::Use(op) | Rvalue::Clone(op) => op,
_ => return None,
};
let place = match inner_op {
Operand::Copy(p) | Operand::Move(p) | Operand::MoveExplicit(p) => p,
_ => return None,
};
let Place::Field(base, field_idx) = place else {
return None;
};
if !length_field_idxs.contains(&field_idx.0) {
return None;
}
if let Place::Local(arr_slot) = base.as_ref() {
Some(*arr_slot)
} else {
None
}
}
fn find_lt_definition(
block: &shape_vm::mir::types::BasicBlock,
cond_slot: SlotId,
) -> Option<(SlotId, SlotId)> {
for stmt in &block.statements {
let StatementKind::Assign(Place::Local(lhs), Rvalue::BinaryOp(BinOp::Lt, l, r)) =
&stmt.kind
else {
continue;
};
if *lhs != cond_slot {
continue;
}
let iv = operand_local(l)?;
let bnd = operand_local(r)?;
return Some((iv, bnd));
}
None
}
fn slot_assignment_count(mir: &MirFunction, slot: SlotId) -> usize {
let mut count = 0usize;
for block in &mir.blocks {
for stmt in &block.statements {
if let StatementKind::Assign(Place::Local(s), _) = &stmt.kind {
if *s == slot {
count += 1;
}
}
}
}
count
}
fn block_assigns(block: &shape_vm::mir::types::BasicBlock, slot: SlotId) -> bool {
for stmt in &block.statements {
if let StatementKind::Assign(Place::Local(s), _) = &stmt.kind {
if *s == slot {
return true;
}
}
}
false
}
fn iv_starts_non_negative(mir: &MirFunction, iv: SlotId, header: BasicBlockId) -> bool {
let mut last_const_init: Option<i64> = None;
for block in &mir.blocks {
if block.id.0 >= header.0 {
continue;
}
for stmt in &block.statements {
let StatementKind::Assign(Place::Local(lhs), rv) = &stmt.kind else {
continue;
};
if *lhs != iv {
continue;
}
let value = match rv {
Rvalue::Use(Operand::Constant(shape_vm::mir::types::MirConstant::Int(v))) => Some(*v),
_ => None,
};
last_const_init = value;
}
}
matches!(last_const_init, Some(v) if v >= 0)
}
fn iv_only_monotonic_in_body(body: &shape_vm::mir::types::BasicBlock, iv: SlotId) -> bool {
for stmt in &body.statements {
let StatementKind::Assign(Place::Local(lhs), rv) = &stmt.kind else {
continue;
};
if *lhs != iv {
continue;
}
let Rvalue::BinaryOp(BinOp::Add, l, r) = rv else {
return false;
};
let l_is_iv = operand_local(l) == Some(iv);
let r_is_iv = operand_local(r) == Some(iv);
let const_step = match (l, r) {
(_, Operand::Constant(shape_vm::mir::types::MirConstant::Int(v))) => Some(*v),
(Operand::Constant(shape_vm::mir::types::MirConstant::Int(v)), _) => Some(*v),
_ => None,
};
let Some(step) = const_step else {
return false;
};
if !(l_is_iv || r_is_iv) {
return false;
}
if step < 0 {
return false;
}
}
true
}
#[cfg(test)]
mod tests {
use super::*;
use shape_vm::mir::types::{
BasicBlock, BasicBlockId, BinOp, FieldIdx, LocalTypeInfo, MirConstant, MirFunction,
MirStatement, Place, Point, Rvalue, SlotId, StatementKind, Terminator, TerminatorKind,
};
use shape_ast::ast::Span;
fn s(kind: StatementKind) -> MirStatement {
MirStatement {
kind,
span: Span { start: 0, end: 0 },
point: Point(0),
}
}
fn term(kind: TerminatorKind) -> Terminator {
Terminator {
kind,
span: Span { start: 0, end: 0 },
}
}
fn mir_for_loop_with_arr_index(arr: SlotId, iv: SlotId, bnd: SlotId, cond: SlotId) -> MirFunction {
let length_idx = FieldIdx(7);
let mut field_name_table = std::collections::HashMap::new();
field_name_table.insert(length_idx, "length".to_string());
let bb0 = BasicBlock {
id: BasicBlockId(0),
statements: vec![
s(StatementKind::Assign(
Place::Local(bnd),
Rvalue::Use(Operand::Copy(Place::Field(
Box::new(Place::Local(arr)),
length_idx,
))),
)),
s(StatementKind::Assign(
Place::Local(iv),
Rvalue::Use(Operand::Constant(MirConstant::Int(0))),
)),
],
terminator: term(TerminatorKind::Goto(BasicBlockId(1))),
};
let bb1 = BasicBlock {
id: BasicBlockId(1),
statements: vec![s(StatementKind::Assign(
Place::Local(cond),
Rvalue::BinaryOp(
BinOp::Lt,
Operand::Copy(Place::Local(iv)),
Operand::Copy(Place::Local(bnd)),
),
))],
terminator: term(TerminatorKind::SwitchBool {
operand: Operand::Copy(Place::Local(cond)),
true_bb: BasicBlockId(2),
false_bb: BasicBlockId(3),
}),
};
let sink = SlotId(99);
let bb2 = BasicBlock {
id: BasicBlockId(2),
statements: vec![
s(StatementKind::Assign(
Place::Local(sink),
Rvalue::Use(Operand::Copy(Place::Index(
Box::new(Place::Local(arr)),
Box::new(Operand::Copy(Place::Local(iv))),
))),
)),
s(StatementKind::Assign(
Place::Local(iv),
Rvalue::BinaryOp(
BinOp::Add,
Operand::Copy(Place::Local(iv)),
Operand::Constant(MirConstant::Int(1)),
),
)),
],
terminator: term(TerminatorKind::Goto(BasicBlockId(1))),
};
let bb3 = BasicBlock {
id: BasicBlockId(3),
statements: vec![],
terminator: term(TerminatorKind::Return),
};
MirFunction {
name: "test_fn".to_string(),
blocks: vec![bb0, bb1, bb2, bb3],
num_locals: 100,
param_slots: vec![arr],
param_reference_kinds: vec![None],
local_types: (0..100).map(|_| LocalTypeInfo::Unknown).collect(),
span: Span { start: 0, end: 0 },
field_name_table,
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(),
}
}
#[test]
fn detects_simple_for_loop_index_pattern() {
let arr = SlotId(1);
let iv = SlotId(2);
let bnd = SlotId(3);
let cond = SlotId(4);
let mir = mir_for_loop_with_arr_index(arr, iv, bnd, cond);
let plan = analyze(&mir);
assert!(
plan.is_trusted(arr, iv),
"expected (arr={:?}, iv={:?}) to be trusted; got {:?}",
arr,
iv,
plan.trusted_pairs,
);
}
#[test]
fn rejects_when_arr_is_reassigned() {
let arr = SlotId(1);
let iv = SlotId(2);
let bnd = SlotId(3);
let cond = SlotId(4);
let mut mir = mir_for_loop_with_arr_index(arr, iv, bnd, cond);
mir.blocks[0].statements.push(s(StatementKind::Assign(
Place::Local(arr),
Rvalue::Use(Operand::Constant(MirConstant::None)),
)));
let plan = analyze(&mir);
assert!(
!plan.is_trusted(arr, iv),
"expected access not to be trusted when arr is reassigned",
);
}
#[test]
fn rejects_when_iv_is_negative_initialized() {
let arr = SlotId(1);
let iv = SlotId(2);
let bnd = SlotId(3);
let cond = SlotId(4);
let mut mir = mir_for_loop_with_arr_index(arr, iv, bnd, cond);
mir.blocks[0].statements[1] = s(StatementKind::Assign(
Place::Local(iv),
Rvalue::Use(Operand::Constant(MirConstant::Int(-1))),
));
let plan = analyze(&mir);
assert!(
!plan.is_trusted(arr, iv),
"expected access not to be trusted when iv starts negative",
);
}
#[test]
fn rejects_when_no_back_edge() {
let arr = SlotId(1);
let iv = SlotId(2);
let bnd = SlotId(3);
let cond = SlotId(4);
let mut mir = mir_for_loop_with_arr_index(arr, iv, bnd, cond);
let bb2 = mir.blocks.iter_mut().find(|b| b.id == BasicBlockId(2)).unwrap();
bb2.terminator = term(TerminatorKind::Return);
let plan = analyze(&mir);
assert!(
!plan.is_trusted(arr, iv),
"expected access not to be trusted without a back edge",
);
}
}