Skip to main content

charon_lib/transform/add_missing_info/
detect_drop_flags.rs

1//! Detect the boolean locals that rustc uses to track whether a place needs to be dropped.
2//!
3//! This may miss some drop flags but the ones it detects are guaranteed to be correct, as we do
4//! a control-flow analysis to be sure.
5use std::collections::HashMap;
6use std::ops::ControlFlow;
7
8use derive_generic_visitor::Visitor;
9use itertools::Itertools;
10use macros::EnumAsGetters;
11use petgraph::graphmap::DiGraphMap;
12use petgraph::visit::{Dfs, Walker};
13
14use crate::options::TranslateOptions;
15use crate::transform::TransformCtx;
16use crate::transform::ctx::UllbcPass;
17use crate::ullbc_ast::*;
18
19impl Rvalue {
20    fn as_const_bool(&self) -> Option<bool> {
21        if let Rvalue::Use(Operand::Const(value), _) = self
22            && let ConstantExprKind::Bool(value) = value.kind()
23        {
24            Some(*value)
25        } else {
26            None
27        }
28    }
29}
30
31#[derive(Debug, Default)]
32enum FlagStatus {
33    /// We have tot encountered a switch on this boolean yet.
34    #[default]
35    Unknown,
36    /// This boolean has only been assigned constants and switched on, hence is a candidate drop
37    /// flag. The places are these that were unconditionally dropped in every true branch of
38    /// a switch on this flag.
39    Candidate { dropped_places: Vec<Place> },
40    /// This boolean was used in a way that drop flags aren't used.
41    Discarded,
42}
43
44#[derive(Visitor)]
45struct GatherCandidates<'a> {
46    body: &'a ExprBody,
47    flag_statuses: IndexVec<LocalId, FlagStatus>,
48}
49
50impl GatherCandidates<'_> {
51    /// Follow the true branch until it reaches the false target, and accumulate any place that
52    /// gets unconditionally dropped along the way. We ignore unwind paths and other switches.
53    fn dropped_places_in_true_branch(
54        &self,
55        true_target: BlockId,
56        false_target: BlockId,
57    ) -> Vec<Place> {
58        let mut dropped_places = Vec::new();
59        let mut visited = self.body.body.map_ref(|_| false);
60        let mut pending = vec![true_target];
61
62        while let Some(block_id) = pending.pop() {
63            if block_id == false_target || std::mem::replace(&mut visited[block_id], true) {
64                continue;
65            }
66
67            let terminator = &self.body.body[block_id].terminator;
68            if let TerminatorKind::Drop {
69                kind: DropKind::Precise,
70                place,
71                ..
72            } = &terminator.kind
73                && !dropped_places.contains(place)
74            {
75                dropped_places.push(place.clone());
76            }
77            pending.extend(terminator.targets_ignoring_unwind());
78        }
79
80        dropped_places
81    }
82}
83
84impl VisitBody for GatherCandidates<'_> {
85    fn enter_local_id(&mut self, local: &LocalId) {
86        // If we get here, the local is being used in a way that drop flags aren't used.
87        self.flag_statuses[*local] = FlagStatus::Discarded;
88    }
89
90    fn visit_ullbc_statement(&mut self, statement: &Statement) -> ControlFlow<Self::Break> {
91        match &statement.kind {
92            // A drop flag may only be assigned a constant value.
93            StatementKind::Assign(place, rval)
94                if let Some(_local) = place.as_local()
95                    && rval.as_const_bool().is_some() =>
96            {
97                ControlFlow::Continue(())
98            }
99            _ => self.visit_inner(statement),
100        }
101    }
102
103    fn visit_ullbc_terminator(&mut self, terminator: &Terminator) -> ControlFlow<Self::Break> {
104        // Find branches over booleans with drops in one of the branches.
105        if let TerminatorKind::Switch { data, branches } = &terminator.kind
106            && let SwitchScrutinee::Value(Operand::Copy(flag_place) | Operand::Move(flag_place)) =
107                &data.scrutinee
108            && let Some(flag) = flag_place.as_local()
109            && !matches!(&self.flag_statuses[flag], FlagStatus::Discarded)
110            && let Some((true_branch, false_branch)) = data.as_if()
111        {
112            let true_target = branches[true_branch];
113            let false_target = branches[false_branch];
114            let dropped_places = self.dropped_places_in_true_branch(true_target, false_target);
115            let status = &mut self.flag_statuses[flag];
116            match status {
117                FlagStatus::Unknown => {
118                    *status = FlagStatus::Candidate { dropped_places };
119                }
120                FlagStatus::Candidate {
121                    dropped_places: candidates,
122                } => {
123                    candidates.retain(|place| dropped_places.contains(place));
124                    if dropped_places.is_empty() {
125                        *status = FlagStatus::Discarded;
126                    }
127                }
128                FlagStatus::Discarded => unreachable!(),
129            }
130            ControlFlow::Continue(())
131        } else {
132            self.visit_inner(terminator)
133        }
134    }
135}
136
137#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
138struct AnalysisState {
139    flag_value: bool,
140    place_is_initialized: bool,
141}
142
143impl AnalysisState {
144    /// All possible states.
145    fn all() -> impl Iterator<Item = Self> {
146        [false, true]
147            .into_iter()
148            .cartesian_product([false, true])
149            .map(|(flag_value, place_is_initialized)| Self {
150                flag_value,
151                place_is_initialized,
152            })
153    }
154
155    fn is_consistent(self) -> bool {
156        self.flag_value == self.place_is_initialized
157    }
158}
159
160#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
161enum AnalysisNode {
162    Root,
163    State {
164        block: BlockId,
165        flag: LocalId,
166        state: AnalysisState,
167    },
168    Invalid(LocalId),
169}
170
171#[derive(Debug, Default, Copy, Clone, EnumAsGetters)]
172enum Update<T> {
173    #[default]
174    Unchanged,
175    Set(T),
176}
177
178impl<T: Copy> Update<T> {
179    fn apply(self, value: T) -> T {
180        match self {
181            Self::Unchanged => value,
182            Self::Set(value) => value,
183        }
184    }
185}
186
187#[derive(Debug, Default, Copy, Clone)]
188struct FlagUpdate {
189    flag_value: Update<bool>,
190    place_is_initialized: Update<bool>,
191}
192
193impl FlagUpdate {
194    fn apply(self, state: AnalysisState) -> AnalysisState {
195        AnalysisState {
196            flag_value: self.flag_value.apply(state.flag_value),
197            place_is_initialized: self.place_is_initialized.apply(state.place_is_initialized),
198        }
199    }
200}
201
202/// For each CFG node and each flag, we record the possible states of both the flag value and
203/// initialization status of the corresponding place. We then add an edge to the graph when one can
204/// transition between two such states. We then check if an invalid node is reachable; if not, then
205/// the boolean flag faithfully represents the initialization status of the place!
206struct CheckFlagCorrectness<'a, 'b> {
207    /// List of candidate drop flags, with the place they correspond to.
208    flags: &'b mut SeqHashMap<LocalId, &'a Place>,
209    graph: DiGraphMap<AnalysisNode, ()>,
210    /// For each block, the updates to candidate flags affected by the block.
211    updates: IndexVec<BlockId, HashMap<LocalId, FlagUpdate>>,
212}
213
214#[derive(Visitor)]
215struct ComputeUpdates<'a, 'b, 'c> {
216    analysis: &'a mut CheckFlagCorrectness<'b, 'c>,
217    block_id: BlockId,
218}
219
220impl VisitBody for ComputeUpdates<'_, '_, '_> {
221    fn enter_operand(&mut self, operand: &Operand) {
222        if let Operand::Move(moved) = operand {
223            let updates = &mut self.analysis.updates[self.block_id];
224            for (&flag, place) in self.analysis.flags.iter() {
225                if place.is_subplace(moved) {
226                    updates.entry(flag).or_default().place_is_initialized = Update::Set(false);
227                }
228            }
229        }
230    }
231
232    fn exit_ullbc_statement(&mut self, statement: &Statement) {
233        match &statement.kind {
234            StatementKind::Assign(assigned, rval) => {
235                let updates = &mut self.analysis.updates[self.block_id];
236                for (&flag, place) in self.analysis.flags.iter() {
237                    if assigned.as_local() == Some(flag) {
238                        // Unwrap is ok because we made sure this is only assigned constants.
239                        let value = rval.as_const_bool().unwrap();
240                        updates.entry(flag).or_default().flag_value = Update::Set(value);
241                    }
242                    if place.is_subplace(assigned) {
243                        updates.entry(flag).or_default().place_is_initialized = Update::Set(true);
244                    }
245                }
246            }
247            StatementKind::StorageLive(local) | StatementKind::StorageDead(local) => {
248                let updates = &mut self.analysis.updates[self.block_id];
249                for (&flag, place) in self.analysis.flags.iter() {
250                    if place.local_id() == Some(*local) {
251                        updates.entry(flag).or_default().place_is_initialized = Update::Set(false);
252                    }
253                }
254            }
255            _ => {}
256        }
257    }
258}
259
260impl<'a, 'b> CheckFlagCorrectness<'a, 'b> {
261    /// Check if these drop flags correctly track the corresponding place. Removes the ones that
262    /// don't.
263    fn filter_invalid_flags(body: &ExprBody, flags: &'b mut SeqHashMap<LocalId, &'a Place>) {
264        let updates = body.body.map_ref(|_| HashMap::new());
265        let mut graph = DiGraphMap::new();
266        graph.add_node(AnalysisNode::Root);
267        let mut analysis = Self {
268            flags,
269            updates,
270            graph,
271        };
272        analysis.compute_updates(body);
273        analysis.build_graph(body);
274        for flag in Dfs::new(&analysis.graph, AnalysisNode::Root)
275            .iter(&analysis.graph)
276            .filter_map(|node| match node {
277                AnalysisNode::Invalid(flag) => Some(flag),
278                AnalysisNode::Root | AnalysisNode::State { .. } => None,
279            })
280        {
281            analysis.flags.swap_remove(&flag);
282        }
283    }
284
285    /// Summarize how each block changes each candidate flag and its associated place.
286    fn compute_updates(&mut self, body: &ExprBody) {
287        for (block_id, block) in body.body.iter_enumerated() {
288            ComputeUpdates {
289                analysis: self,
290                block_id,
291            }
292            .visit(block);
293        }
294    }
295
296    /// Build a graph that has an edge for every valid state transition.
297    fn build_graph(&mut self, body: &ExprBody) {
298        for (&flag, place) in self.flags.iter() {
299            // Drop flags get initialized in the first block.
300            let node = if let Some(update) = self.updates[START_BLOCK_ID].get(&flag)
301                && let Some(flag_value) = update.flag_value.as_set().copied()
302            {
303                let initially_initialized = place.local_id().is_none_or(|local| {
304                    local.index() > 0 && local.index() <= body.locals.arg_count
305                });
306                AnalysisNode::State {
307                    block: START_BLOCK_ID,
308                    flag,
309                    state: AnalysisState {
310                        flag_value,
311                        place_is_initialized: initially_initialized,
312                    },
313                }
314            } else {
315                AnalysisNode::Invalid(flag)
316            };
317            self.graph.add_edge(AnalysisNode::Root, node, ());
318        }
319
320        for (block_id, block) in body.body.iter_enumerated() {
321            for (&flag, place) in self.flags.iter() {
322                let update = self.updates[block_id]
323                    .get(&flag)
324                    .copied()
325                    .unwrap_or_default();
326                for initial_state in AnalysisState::all() {
327                    let source = AnalysisNode::State {
328                        block: block_id,
329                        flag,
330                        state: initial_state,
331                    };
332                    let state = update.apply(initial_state);
333                    if let Some(successors) = Self::block_successors(block, flag, place, state) {
334                        for (target, state) in successors {
335                            self.graph.add_edge(
336                                source,
337                                AnalysisNode::State {
338                                    block: target,
339                                    flag,
340                                    state,
341                                },
342                                (),
343                            );
344                        }
345                    } else {
346                        self.graph.add_edge(source, AnalysisNode::Invalid(flag), ());
347                    }
348                }
349            }
350        }
351    }
352
353    /// Compute the reachable states from this one at the end of a block. `None` means that the
354    /// state is inconsistent and the drop flag is invalid.
355    fn block_successors(
356        block: &BlockData,
357        flag: LocalId,
358        place: &Place,
359        mut state: AnalysisState,
360    ) -> Option<Vec<(BlockId, AnalysisState)>> {
361        if let TerminatorKind::Switch { data, branches } = &block.terminator.kind
362            && let SwitchScrutinee::Value(Operand::Copy(flag_place) | Operand::Move(flag_place)) =
363                &data.scrutinee
364            && flag_place.as_local() == Some(flag)
365        {
366            // The only place that state consistency matters is when switching on the boolean flag.
367            if !state.is_consistent() {
368                return None;
369            }
370            let (true_branch, false_branch) = data.as_if()?;
371            let branch = if state.flag_value {
372                true_branch
373            } else {
374                false_branch
375            };
376            Some(vec![(branches[branch], state)])
377        } else {
378            Some(match &block.terminator.kind {
379                TerminatorKind::Call {
380                    call,
381                    target,
382                    on_unwind,
383                } => {
384                    let mut after_return = state;
385                    if place.is_subplace(&call.dest) {
386                        after_return.place_is_initialized = true;
387                    }
388                    vec![(*target, after_return), (*on_unwind, state)]
389                }
390                TerminatorKind::Drop {
391                    place: dropped,
392                    target,
393                    on_unwind,
394                    ..
395                } => {
396                    if place.is_subplace(dropped) {
397                        state.place_is_initialized = false;
398                    }
399                    vec![(*target, state), (*on_unwind, state)]
400                }
401                _ => block
402                    .terminator
403                    .targets()
404                    .into_iter()
405                    .map(|target| (target, state))
406                    .collect(),
407            })
408        }
409    }
410}
411
412pub struct Transform;
413impl UllbcPass for Transform {
414    fn should_run(&self, options: &TranslateOptions) -> bool {
415        options.detect_drop_flags || options.resugar_drops
416    }
417
418    fn transform_body(&self, _ctx: &mut TransformCtx, body: &mut ExprBody) {
419        if !body.body.iter().any(|block| {
420            matches!(
421                block.terminator.kind,
422                TerminatorKind::Drop {
423                    kind: DropKind::Precise,
424                    ..
425                }
426            )
427        }) {
428            return;
429        }
430
431        // Start with anonymous boolean locals.
432        let flag_statuses: IndexVec<LocalId, FlagStatus> = body.locals.locals.map_ref(|local| {
433            if local.name.is_none()
434                && !body.locals.is_return_or_arg(local.index)
435                && local.ty.is_bool()
436            {
437                FlagStatus::default()
438            } else {
439                FlagStatus::Discarded
440            }
441        });
442
443        // Identify booleans that are used like drop flags, and find out what place they track.
444        let flag_statuses = {
445            let mut visitor = GatherCandidates {
446                body,
447                flag_statuses,
448            };
449            for block in &body.body {
450                visitor.visit(block);
451            }
452            visitor.flag_statuses
453        };
454
455        // Keep flags for which we know the tracked place, then check that they do in fact track
456        // the initializedness of that place.
457        let mut flag_candidates: SeqHashMap<LocalId, &Place> = flag_statuses
458            .iter_enumerated()
459            .filter_map(|(flag, status)| match status {
460                FlagStatus::Candidate { dropped_places } => {
461                    let place = dropped_places.iter().exactly_one().ok()?;
462                    Some((flag, place))
463                }
464                FlagStatus::Unknown | FlagStatus::Discarded => None,
465            })
466            .collect();
467        CheckFlagCorrectness::filter_invalid_flags(body, &mut flag_candidates);
468        for (flag, place) in flag_candidates {
469            body.locals.locals[flag].drop_flag_for = Some(place.clone());
470        }
471    }
472}