1use 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 #[default]
35 Unknown,
36 Candidate { dropped_places: Vec<Place> },
40 Discarded,
42}
43
44#[derive(Visitor)]
45struct GatherCandidates<'a> {
46 body: &'a ExprBody,
47 flag_statuses: IndexVec<LocalId, FlagStatus>,
48}
49
50impl GatherCandidates<'_> {
51 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 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 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 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 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
202struct CheckFlagCorrectness<'a, 'b> {
207 flags: &'b mut SeqHashMap<LocalId, &'a Place>,
209 graph: DiGraphMap<AnalysisNode, ()>,
210 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 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 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 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 fn build_graph(&mut self, body: &ExprBody) {
298 for (&flag, place) in self.flags.iter() {
299 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 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 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 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 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 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}