use rustc_hir::def_id::DefId;
use rustc_middle::mir::Body;
use rustc_middle::mir::{BasicBlock, Local, StatementKind, TerminatorKind};
use rustc_middle::ty::TyCtxt;
use std::collections::{HashMap, HashSet};
use crate::analysis::dataflow::graph::build_dataflow_graph;
use crate::analysis::dataflow::types::DataflowGraph;
use super::super::{
contract,
def_use::{RelevantPlaces, bind_callsite_roots, operand_uses, terminator_use_def},
path_extractor::{Path, PathStep},
};
use crate::helpers::mir_scan::{Checkpoint, CheckpointLocation};
use crate::analysis::path::{PathNode, PathTree};
use super::{
call_visit,
types::{ProofGoal, RelevantItem},
};
pub(crate) struct BackwardSlicer<'tcx> {
tcx: TyCtxt<'tcx>,
}
impl<'tcx> BackwardSlicer<'tcx> {
pub(crate) fn new(tcx: TyCtxt<'tcx>) -> Self {
Self { tcx }
}
pub(crate) fn visit_path_tree(
&self,
tree: &PathTree,
target_block: usize,
checkpoint: &Checkpoint<'tcx>,
property: &contract::Property<'tcx>,
) -> Vec<ProofGoal<'tcx>> {
self.visit_path_tree_impl(
tree,
target_block,
checkpoint.caller,
checkpoint.block,
Some(checkpoint),
property,
)
}
pub(crate) fn visit_path_tree_for_checkpoint(
&self,
tree: &PathTree,
target_block: usize,
caller: DefId,
checkpoint_loc: CheckpointLocation,
property: &contract::Property<'tcx>,
) -> Vec<ProofGoal<'tcx>> {
self.visit_path_tree_impl(
tree,
target_block,
caller,
checkpoint_loc.block,
None,
property,
)
}
fn visit_path_tree_impl(
&self,
tree: &PathTree,
target_block: usize,
caller: DefId,
checkpoint_block: BasicBlock,
bind_checkpoint: Option<&Checkpoint<'tcx>>,
property: &contract::Property<'tcx>,
) -> Vec<ProofGoal<'tcx>> {
let Some(root) = tree.root() else {
return Vec::new();
};
let checkpoint_loc = CheckpointLocation {
caller,
block: checkpoint_block,
};
let mut bodies: HashMap<DefId, &'tcx Body<'tcx>> = HashMap::new();
let mut flows: HashMap<DefId, DataflowGraph> = HashMap::new();
let mut def_ids: HashSet<DefId> = tree.block_fns().iter().map(|(d, _)| *d).collect();
def_ids.insert(caller);
for d in def_ids {
bodies.insert(d, self.tcx.optimized_mir(d));
flows.insert(d, build_dataflow_graph(self.tcx, d));
}
let leaf_results = Self::build_leaf_items(
self,
tree,
root,
target_block,
checkpoint_block,
bind_checkpoint,
property,
caller,
&bodies,
&flows,
);
let mut results = Vec::new();
for (block_path, backward_items, _relevant) in leaf_results {
let mut items = backward_items;
items.reverse();
let steps: Vec<PathStep> = block_path
.iter()
.map(|&b| PathStep::Block(BasicBlock::from(b)))
.chain(std::iter::once(PathStep::Checkpoint(checkpoint_loc)))
.collect();
results.push(ProofGoal {
path: Path {
target: checkpoint_loc,
steps,
},
items,
block_fn: tree.block_fns().to_vec(),
});
}
results
}
fn build_leaf_items(
visitor: &Self,
tree: &PathTree,
node: &PathNode,
target_block: usize,
checkpoint_block: BasicBlock,
bind_checkpoint: Option<&Checkpoint<'tcx>>,
property: &contract::Property<'tcx>,
caller: DefId,
bodies: &HashMap<DefId, &'tcx Body<'tcx>>,
flows: &HashMap<DefId, DataflowGraph>,
) -> Vec<(Vec<usize>, Vec<RelevantItem<'tcx>>, RelevantPlaces)> {
let (def_id, local_index) = tree.block_fn_of(node.block).unwrap_or((caller, node.block));
let body = &bodies[&def_id];
let flow = &flows[&def_id];
let block = BasicBlock::from(local_index);
let keep_inv = property
.kind()
.is_some_and(|k| needs_invalidation_tracking(&k));
let block_data = &body.basic_blocks[block];
let mut results = Vec::new();
let (checkpoint_items, checkpoint_relevant) = if node.block == target_block {
let mut relevant = RelevantPlaces::from_property(property);
if let Some(cs) = bind_checkpoint {
bind_callsite_roots(visitor.tcx, &mut relevant, cs);
}
let mut items = Vec::new();
items.push(RelevantItem::Terminator {
def_id: caller,
block: checkpoint_block,
});
for (si, stmt) in block_data.statements.iter().enumerate().rev() {
visitor.visit_statement(
def_id,
checkpoint_block,
si,
stmt,
flow,
&mut relevant,
&mut items,
keep_inv,
);
}
Self::re_visit_newly_added(
visitor,
def_id,
checkpoint_block,
block_data,
flow,
&mut relevant,
&mut items,
keep_inv,
);
(items, relevant)
} else {
(Vec::new(), RelevantPlaces::new())
};
for child in &node.children {
let child_results = Self::build_leaf_items(
visitor,
tree,
child,
target_block,
checkpoint_block,
bind_checkpoint,
property,
caller,
bodies,
flows,
);
for (mut child_path, child_items, child_relevant) in child_results {
let mut relevant = child_relevant;
let mut items = child_items;
if let Some(binding) = tree.inline_binding(node.block) {
let dest = Local::from_usize(binding.dest_local);
if relevant.locals.contains(&dest) {
relevant.locals.remove(&dest);
relevant.places.retain(|p| p.local() != Some(dest));
relevant.insert_local(Local::from_usize(0));
}
}
if !tree.is_inlined_call(node.block) {
visitor.visit_terminator(
def_id,
block,
block_data.terminator(),
flow,
body,
&mut relevant,
&mut items,
keep_inv,
);
}
let block_stmt_count = block_data.statements.len();
for (si, stmt) in block_data.statements.iter().enumerate().rev() {
visitor.visit_statement(
def_id,
block,
si,
stmt,
flow,
&mut relevant,
&mut items,
keep_inv,
);
}
if let Some(binding) = tree.inline_binding(node.block) {
for (i, arg_local) in binding.arg_locals.iter().enumerate() {
let param = Local::from_usize(i + 1);
if relevant.locals.contains(¶m) {
relevant.locals.remove(¶m);
relevant.places.retain(|p| p.local() != Some(param));
relevant.insert_local(Local::from_usize(*arg_local));
}
}
}
let dist_to_target = child_path.iter().position(|&b| b == target_block);
if block_stmt_count > 0 && dist_to_target.map_or(false, |d| d <= 2) {
Self::re_visit_newly_added(
visitor,
def_id,
block,
block_data,
flow,
&mut relevant,
&mut items,
keep_inv,
);
}
child_path.insert(0, node.block);
results.push((child_path, items, relevant));
}
}
if !checkpoint_items.is_empty() {
results.push((vec![node.block], checkpoint_items, checkpoint_relevant));
}
results
}
fn re_visit_newly_added(
visitor: &Self,
def_id: DefId,
block: BasicBlock,
block_data: &'tcx rustc_middle::mir::BasicBlockData<'tcx>,
flow: &DataflowGraph,
relevant: &mut RelevantPlaces,
items: &mut Vec<RelevantItem<'tcx>>,
keep_inv: bool,
) {
let newly_added = std::mem::take(&mut relevant.just_added);
if newly_added.is_empty() {
return;
}
for (si, stmt) in block_data.statements.iter().enumerate().rev() {
let defs = match &stmt.kind {
rustc_middle::mir::StatementKind::Assign(assign) => {
let mut d = crate::verify::def_use::RelevantPlaces::new();
d.insert_mir_place(&assign.0);
d
}
_ => continue,
};
let any_new = defs
.places
.iter()
.any(|dp| newly_added.iter().any(|np| dp.local() == np.local()));
if any_new {
visitor.visit_statement(def_id, block, si, stmt, flow, relevant, items, keep_inv);
}
}
}
fn visit_statement(
&self,
def_id: DefId,
block: BasicBlock,
statement_index: usize,
statement: &'tcx rustc_middle::mir::Statement<'tcx>,
flow: &DataflowGraph,
relevant: &mut RelevantPlaces,
items: &mut Vec<RelevantItem<'tcx>>,
keep_invalidations: bool,
) {
if keep_invalidations
&& matches!(
statement.kind,
StatementKind::StorageDead(_) | StatementKind::StorageLive(_)
)
{
items.push(RelevantItem::Statement {
def_id,
block,
statement_index,
});
return;
}
let mut defs = RelevantPlaces::new();
match &statement.kind {
StatementKind::Assign(assign) => {
let (place, _) = &**assign;
defs.insert_mir_place(place);
}
StatementKind::StorageDead(local) => {
defs.insert_local(*local);
}
_ => {}
}
if defs.intersects(relevant) {
let mut uses = collect_statement_uses(statement, block, statement_index, flow, &defs);
items.push(RelevantItem::Statement {
def_id,
block,
statement_index,
});
let already_seen: crate::compat::FxHashSet<crate::verify::def_use::PlaceKey> =
relevant.places.clone();
relevant.remove_all(&defs);
uses.places.retain(|p| !already_seen.contains(p));
relevant.extend(uses);
return;
}
if statement_can_refine(statement) {
let uses = collect_flow_uses(flow, block, statement_index, &defs);
if uses.intersects(relevant) {
items.push(RelevantItem::Statement {
def_id,
block,
statement_index,
});
}
}
}
fn visit_terminator(
&self,
def_id: DefId,
block: BasicBlock,
terminator: &rustc_middle::mir::Terminator<'tcx>,
flow: &DataflowGraph,
body: &Body<'tcx>,
relevant: &mut RelevantPlaces,
items: &mut Vec<RelevantItem<'tcx>>,
keep_invalidations: bool,
) {
if keep_invalidations && matches!(terminator.kind, TerminatorKind::Drop { .. }) {
items.push(RelevantItem::Terminator { def_id, block });
return;
}
if let TerminatorKind::Call {
func,
args,
destination,
..
} = &terminator.kind
{
call_visit::visit(
self.tcx,
def_id,
block,
func,
args,
destination,
flow,
body,
relevant,
items,
);
return;
}
let use_def = terminator_use_def(terminator);
if terminator_is_path_condition(terminator) {
items.push(RelevantItem::Terminator { def_id, block });
relevant.extend(use_def.uses.clone());
return;
}
if use_def.defs.intersects(relevant) {
items.push(RelevantItem::Terminator { def_id, block });
relevant.remove_all(&use_def.defs);
relevant.extend(use_def.uses);
return;
}
if use_def.uses.intersects(relevant) {
items.push(RelevantItem::Terminator { def_id, block });
}
}
}
fn needs_invalidation_tracking(kind: &contract::PropertyKind) -> bool {
matches!(
kind,
contract::PropertyKind::Allocated
| contract::PropertyKind::Init
| contract::PropertyKind::Alive
| contract::PropertyKind::ValidString
| contract::PropertyKind::ValidCStr
| contract::PropertyKind::Owning
)
}
fn statement_can_refine(statement: &rustc_middle::mir::Statement<'_>) -> bool {
matches!(&statement.kind, StatementKind::Assign(assign) if matches!(
&**assign,
(
_,
rustc_middle::mir::Rvalue::BinaryOp(_, _)
| rustc_middle::mir::Rvalue::UnaryOp(_, _)
| rustc_middle::mir::Rvalue::Cast(_, _, _),
)
))
}
fn terminator_is_path_condition(terminator: &rustc_middle::mir::Terminator<'_>) -> bool {
matches!(
terminator.kind,
TerminatorKind::SwitchInt { .. } | TerminatorKind::Assert { .. }
)
}
fn collect_statement_uses<'tcx>(
statement: &'tcx rustc_middle::mir::Statement<'tcx>,
block: BasicBlock,
statement_index: usize,
flow: &DataflowGraph,
defs: &RelevantPlaces,
) -> RelevantPlaces {
let mut uses = collect_flow_uses(flow, block, statement_index, defs);
if let StatementKind::Assign(assign) = &statement.kind {
let (_, rvalue) = &**assign;
for operand in super::super::def_use::rvalue_operands(rvalue) {
uses.extend(operand_uses(operand));
}
if let rustc_middle::mir::Rvalue::Ref(_, _, place)
| rustc_middle::mir::Rvalue::RawPtr(_, place) = rvalue
{
uses.insert_local(place.local);
}
}
uses
}
fn collect_flow_uses(
flow: &DataflowGraph,
block: BasicBlock,
statement_index: usize,
defs: &RelevantPlaces,
) -> RelevantPlaces {
let mut uses = RelevantPlaces::new();
for &local in &defs.locals {
for &edge_idx in &flow.node(local).in_edges {
let edge = &flow.edges[edge_idx];
if edge.block == block.as_usize() && edge.statement_index == statement_index {
uses.insert_local(edge.src);
}
}
}
uses
}