Skip to main content

flux_refineck/ghost_statements/
fold_unfold.rs

1use std::{collections::hash_map::Entry, fmt, iter};
2
3use flux_common::{tracked_span_assert_eq, tracked_span_bug, tracked_span_dbg_assert_eq};
4use flux_middle::{
5    PlaceExt as _, def_id_to_string, global_env::GlobalEnv, queries::QueryResult, query_bug, rty,
6};
7use flux_rustc_bridge::{
8    mir::{
9        BasicBlock, Body, BorrowKind, FIRST_VARIANT, FieldIdx, Local, Location,
10        NonDivergingIntrinsic, Operand, Place, PlaceElem, PlaceRef, Rvalue, Statement,
11        StatementKind, Terminator, TerminatorKind, UnOp, VariantIdx,
12    },
13    ty::{AdtDef, GenericArgs, GenericArgsExt as _, List, Mutability, Ty, TyKind},
14};
15use itertools::{Itertools, repeat_n};
16use rustc_data_structures::{fx::FxHashMap, unord::UnordMap};
17use rustc_hir::def_id::DefId;
18use rustc_index::{Idx, IndexVec, bit_set::DenseBitSet};
19use rustc_middle::mir::{FakeReadCause, START_BLOCK};
20
21use super::{GhostStatements, StatementsAt};
22use crate::{
23    ghost_statements::{GhostStatement, Point},
24    queue::WorkQueue,
25};
26
27pub(crate) fn add_ghost_statements<'tcx>(
28    stmts: &mut GhostStatements,
29    genv: GlobalEnv<'_, 'tcx>,
30    body: &Body<'tcx>,
31    fn_sig: Option<&rty::EarlyBinder<rty::PolyFnSig>>,
32) -> QueryResult {
33    let mut bb_envs = UnordMap::default();
34    FoldUnfoldAnalysis::new(genv, body, &mut bb_envs, Infer).run(fn_sig)?;
35
36    FoldUnfoldAnalysis::new(genv, body, &mut bb_envs, Elaboration { stmts }).run(fn_sig)
37}
38
39#[derive(Clone)]
40struct Env {
41    map: IndexVec<Local, PlaceNode>,
42}
43
44impl Env {
45    fn new(body: &Body) -> Self {
46        Self {
47            map: body
48                .local_decls
49                .iter()
50                .map(|decl| PlaceNode::Ty(decl.ty.clone()))
51                .collect(),
52        }
53    }
54
55    fn projection<'a>(&mut self, genv: GlobalEnv, place: &'a Place) -> QueryResult<ProjResult<'a>> {
56        let (node, place, modified) = self.ensure_unfolded(genv, place)?;
57        if modified {
58            Ok(ProjResult::Unfold(place))
59        } else if node.ensure_folded() {
60            Ok(ProjResult::Fold(place))
61        } else {
62            Ok(ProjResult::None)
63        }
64    }
65
66    fn downcast(&mut self, genv: GlobalEnv, place: &Place, variant_idx: VariantIdx) -> QueryResult {
67        let (node, ..) = self.ensure_unfolded(genv, place)?;
68        node.downcast(genv, variant_idx)?;
69        Ok(())
70    }
71
72    fn ensure_unfolded<'a>(
73        &mut self,
74        genv: GlobalEnv,
75        place: &'a Place,
76    ) -> QueryResult<(&mut PlaceNode, PlaceRef<'a>, Modified)> {
77        let mut node = &mut self.map[place.local];
78        let mut modified = false;
79        let mut i = 0;
80        while i < place.projection.len() {
81            let elem = place.projection[i];
82            let (n, m) = match elem {
83                PlaceElem::Deref => node.deref(),
84                PlaceElem::Field(f) => node.field(genv, f)?,
85                PlaceElem::Downcast(_, idx) => node.downcast(genv, idx)?,
86                PlaceElem::Index(_) | PlaceElem::ConstantIndex { .. } => break,
87            };
88            node = n;
89            modified |= m;
90            i += 1;
91        }
92        Ok((node, place.as_ref().truncate(i), modified))
93    }
94
95    fn join(&mut self, genv: GlobalEnv, mut other: Env) -> QueryResult<Modified> {
96        let mut modified = false;
97        for (local, node) in self.map.iter_enumerated_mut() {
98            let (m, _) = node.join(genv, &mut other.map[local], false)?;
99            modified |= m;
100        }
101        Ok(modified)
102    }
103
104    fn collect_fold_unfolds_at_goto(&self, target: &Env, stmts: &mut StatementsAt) {
105        for (local, node) in self.map.iter_enumerated() {
106            node.collect_fold_unfolds(&target.map[local], &mut Place::new(local, vec![]), stmts);
107        }
108    }
109
110    fn collect_folds_at_ret(&self, body: &Body, stmts: &mut StatementsAt) {
111        for local in body.args_iter() {
112            self.map[local].collect_folds_at_ret(&mut Place::new(local, vec![]), stmts);
113        }
114    }
115}
116
117type Modified = bool;
118
119struct FoldUnfoldAnalysis<'a, 'genv, 'tcx, M> {
120    genv: GlobalEnv<'genv, 'tcx>,
121    body: &'a Body<'tcx>,
122    bb_envs: &'a mut UnordMap<BasicBlock, Env>,
123    visited: DenseBitSet<BasicBlock>,
124    queue: WorkQueue<'a>,
125    discriminants: UnordMap<Place, Place>,
126    point: Point,
127    mode: M,
128}
129
130trait Mode: Sized {
131    const _NAME: &'static str;
132
133    fn projection(
134        analysis: &mut FoldUnfoldAnalysis<Self>,
135        env: &mut Env,
136        place: &Place,
137    ) -> QueryResult;
138
139    fn goto_join_point(
140        analysis: &mut FoldUnfoldAnalysis<Self>,
141        target: BasicBlock,
142        env: Env,
143    ) -> QueryResult<bool>;
144
145    fn ret(analysis: &mut FoldUnfoldAnalysis<Self>, env: &Env);
146}
147
148struct Infer;
149
150struct Elaboration<'a> {
151    stmts: &'a mut GhostStatements,
152}
153
154impl Elaboration<'_> {
155    fn insert_at(&mut self, point: Point, stmt: GhostStatement) {
156        self.stmts.insert_at(point, stmt);
157    }
158}
159
160#[derive(Debug)]
161enum ProjResult<'a> {
162    None,
163    Fold(PlaceRef<'a>),
164    Unfold(PlaceRef<'a>),
165}
166
167impl Mode for Infer {
168    const _NAME: &'static str = "infer";
169
170    fn projection(
171        analysis: &mut FoldUnfoldAnalysis<Self>,
172        env: &mut Env,
173        place: &Place,
174    ) -> QueryResult {
175        env.projection(analysis.genv, place)?;
176        Ok(())
177    }
178
179    fn goto_join_point(
180        analysis: &mut FoldUnfoldAnalysis<Self>,
181        target: BasicBlock,
182        env: Env,
183    ) -> QueryResult<bool> {
184        let modified = match analysis.bb_envs.entry(target) {
185            Entry::Occupied(mut entry) => entry.get_mut().join(analysis.genv, env)?,
186            Entry::Vacant(entry) => {
187                entry.insert(env);
188                true
189            }
190        };
191        Ok(modified)
192    }
193
194    fn ret(_: &mut FoldUnfoldAnalysis<Self>, _: &Env) {}
195}
196
197impl Mode for Elaboration<'_> {
198    const _NAME: &'static str = "elaboration";
199
200    fn projection(
201        analysis: &mut FoldUnfoldAnalysis<Self>,
202        env: &mut Env,
203        place: &Place,
204    ) -> QueryResult {
205        match env.projection(analysis.genv, place)? {
206            ProjResult::None => {}
207            ProjResult::Fold(place_ref) => {
208                tracked_span_assert_eq!(place_ref, place.as_ref());
209                let place = place.clone();
210                analysis
211                    .mode
212                    .insert_at(analysis.point, GhostStatement::Fold(place));
213            }
214            // `place_ref` is the longest prefix of `place` that `ensure_unfolded` could walk
215            // through. It is a strict prefix when `place` indexes into an array or slice, in
216            // which case we unfold up to the array/slice itself.
217            ProjResult::Unfold(place_ref) => {
218                match place_ref.last_projection() {
219                    Some((base, PlaceElem::Deref | PlaceElem::Field(..))) => {
220                        analysis
221                            .mode
222                            .insert_at(analysis.point, GhostStatement::Unfold(base.to_place()));
223                    }
224                    _ => Err(query_bug!("invalid projection for unfolding {place_ref:?}"))?,
225                }
226            }
227        }
228        Ok(())
229    }
230
231    fn goto_join_point(
232        analysis: &mut FoldUnfoldAnalysis<Self>,
233        target: BasicBlock,
234        env: Env,
235    ) -> QueryResult<bool> {
236        env.collect_fold_unfolds_at_goto(
237            &analysis.bb_envs[&target],
238            &mut analysis.mode.stmts.at(analysis.point),
239        );
240        Ok(!analysis.visited.contains(target))
241    }
242
243    fn ret(analysis: &mut FoldUnfoldAnalysis<Self>, env: &Env) {
244        env.collect_folds_at_ret(analysis.body, &mut analysis.mode.stmts.at(analysis.point));
245    }
246}
247
248#[derive(Clone)]
249enum PlaceNode {
250    Deref(Ty, Box<PlaceNode>),
251    Downcast(AdtDef, GenericArgs, VariantIdx, Vec<PlaceNode>),
252    Closure(DefId, GenericArgs, Vec<PlaceNode>),
253    Generator(DefId, GenericArgs, Vec<PlaceNode>),
254    Tuple(List<Ty>, Vec<PlaceNode>),
255    Ty(Ty),
256}
257
258impl<M: Mode> FoldUnfoldAnalysis<'_, '_, '_, M> {
259    fn run(mut self, fn_sig: Option<&rty::EarlyBinder<rty::PolyFnSig>>) -> QueryResult {
260        let mut env = Env::new(self.body);
261
262        if let Some(fn_sig) = fn_sig {
263            let fn_sig = fn_sig.as_ref().skip_binder().as_ref().skip_binder();
264            for (local, ty) in iter::zip(self.body.args_iter(), fn_sig.inputs()) {
265                if let rty::TyKind::StrgRef(..) | rty::Ref!(.., Mutability::Mut) = ty.kind() {
266                    M::projection(&mut self, &mut env, &Place::new(local, vec![PlaceElem::Deref]))?;
267                }
268            }
269        }
270        self.goto(START_BLOCK, env)?;
271        while let Some(bb) = self.queue.pop() {
272            self.basic_block(bb, self.bb_envs[&bb].clone())?;
273        }
274        Ok(())
275    }
276
277    fn basic_block(&mut self, bb: BasicBlock, mut env: Env) -> QueryResult {
278        self.visited.insert(bb);
279        let data = &self.body.basic_blocks[bb];
280        for (statement_index, stmt) in data.statements.iter().enumerate() {
281            self.point = Point::BeforeLocation(Location { block: bb, statement_index });
282            self.statement(stmt, &mut env)?;
283        }
284        if let Some(terminator) = &data.terminator {
285            self.point = Point::BeforeLocation(self.body.terminator_loc(bb));
286            let successors = self.terminator(terminator, env)?;
287            for (env, target) in successors {
288                self.point = Point::Edge(bb, target);
289                self.goto(target, env)?;
290            }
291        }
292        Ok(())
293    }
294
295    fn statement(&mut self, stmt: &Statement, env: &mut Env) -> QueryResult {
296        match &stmt.kind {
297            StatementKind::FakeRead(deref!((FakeReadCause::ForIndex, place))) => {
298                M::projection(self, env, place)?;
299            }
300            StatementKind::Assign(place, rvalue) => {
301                match rvalue {
302                    Rvalue::UnaryOp(UnOp::PtrMetadata, Operand::Copy(place))
303                    | Rvalue::UnaryOp(UnOp::PtrMetadata, Operand::Move(place)) => {
304                        let deref_place = place.deref();
305                        M::projection(self, env, &deref_place)?;
306                    }
307                    Rvalue::Use(op, _) | Rvalue::Cast(_, op, _) | Rvalue::UnaryOp(_, op) => {
308                        self.operand(op, env)?;
309                    }
310                    Rvalue::Ref(.., bk, place) => {
311                        // Fake borrows should not cause the place to fold
312                        if !matches!(bk, BorrowKind::Fake(_)) {
313                            M::projection(self, env, place)?;
314                        }
315                    }
316                    Rvalue::RawPtr(_, place) => {
317                        M::projection(self, env, place)?;
318                    }
319                    Rvalue::BinaryOp(_, op1, op2) => {
320                        self.operand(op1, env)?;
321                        self.operand(op2, env)?;
322                    }
323                    Rvalue::Aggregate(_, args) => {
324                        for arg in args {
325                            self.operand(arg, env)?;
326                        }
327                    }
328
329                    Rvalue::Discriminant(discr) => {
330                        M::projection(self, env, discr)?;
331                        self.discriminants.insert(place.clone(), discr.clone());
332                    }
333                    Rvalue::Repeat(op, _) => {
334                        self.operand(op, env)?;
335                    }
336                }
337                M::projection(self, env, place)?;
338            }
339            StatementKind::Intrinsic(NonDivergingIntrinsic::Assume(op)) => {
340                self.operand(op, env)?;
341            }
342            StatementKind::SetDiscriminant(_, _)
343            | StatementKind::FakeRead(_)
344            | StatementKind::AscribeUserType(_, _)
345            | StatementKind::PlaceMention(_)
346            | StatementKind::Nop => {}
347        }
348        Ok(())
349    }
350
351    fn operand(&mut self, op: &Operand, env: &mut Env) -> QueryResult {
352        match op {
353            Operand::Copy(place) | Operand::Move(place) => {
354                M::projection(self, env, place)?;
355            }
356            Operand::Constant(_) => {}
357        }
358        Ok(())
359    }
360
361    fn terminator(
362        &mut self,
363        terminator: &Terminator,
364        mut env: Env,
365    ) -> QueryResult<Vec<(Env, BasicBlock)>> {
366        let mut successors = vec![];
367        match &terminator.kind {
368            TerminatorKind::Return => {
369                M::ret(self, &env);
370            }
371            TerminatorKind::Call { args, destination, target, .. } => {
372                for arg in args {
373                    self.operand(arg, &mut env)?;
374                }
375                M::projection(self, &mut env, destination)?;
376                if let Some(target) = target {
377                    successors.push((env, *target));
378                }
379            }
380            TerminatorKind::SwitchInt { discr, targets } => {
381                let is_match = match discr {
382                    Operand::Copy(place) | Operand::Move(place) => {
383                        M::projection(self, &mut env, place)?;
384                        self.discriminants.remove(place)
385                    }
386                    Operand::Constant(_) => None,
387                };
388                if let Some(place) = is_match {
389                    let discr_ty = place.ty(self.genv, &self.body.local_decls)?.ty;
390                    let (adt, _) = discr_ty.expect_adt();
391
392                    let mut remaining: FxHashMap<u128, VariantIdx> = adt
393                        .discriminants()
394                        .map(|(idx, discr)| (discr, idx))
395                        .collect();
396                    for (bits, target) in targets.iter() {
397                        let variant_idx = remaining
398                            .remove(&bits)
399                            .expect("value doesn't correspond to any variant");
400
401                        // We do not insert unfolds in match arms because they are explicit
402                        // unfold points.
403                        let mut env = env.clone();
404                        env.downcast(self.genv, &place, variant_idx)?;
405                        successors.push((env, target));
406                    }
407                    if remaining.len() == 1 {
408                        let (_, variant_idx) = remaining
409                            .into_iter()
410                            .next()
411                            .unwrap_or_else(|| tracked_span_bug!());
412                        env.downcast(self.genv, &place, variant_idx)?;
413                    }
414                    successors.push((env, targets.otherwise()));
415                } else {
416                    let n = targets.all_targets().len();
417                    for (env, target) in iter::zip(repeat_n(env, n), targets.all_targets()) {
418                        successors.push((env, *target));
419                    }
420                }
421            }
422            TerminatorKind::Goto { target } => {
423                successors.push((env, *target));
424            }
425            TerminatorKind::Yield { resume, resume_arg, .. } => {
426                M::projection(self, &mut env, resume_arg)?;
427                successors.push((env, *resume));
428            }
429            TerminatorKind::Drop { place, target, .. } => {
430                M::projection(self, &mut env, place)?;
431                successors.push((env, *target));
432            }
433            TerminatorKind::Assert { cond, target, .. } => {
434                self.operand(cond, &mut env)?;
435                successors.push((env, *target));
436            }
437            TerminatorKind::FalseEdge { real_target, .. } => {
438                successors.push((env, *real_target));
439            }
440            TerminatorKind::FalseUnwind { real_target, .. } => {
441                successors.push((env, *real_target));
442            }
443            TerminatorKind::Unreachable
444            | TerminatorKind::UnwindResume
445            | TerminatorKind::CoroutineDrop => {}
446        }
447        Ok(successors)
448    }
449
450    fn goto(&mut self, target: BasicBlock, env: Env) -> QueryResult {
451        if self.body.is_join_point(target) {
452            if M::goto_join_point(self, target, env)? {
453                self.queue.insert(target);
454            }
455            Ok(())
456        } else {
457            self.basic_block(target, env)
458        }
459    }
460}
461
462impl<'a, 'genv, 'tcx, M> FoldUnfoldAnalysis<'a, 'genv, 'tcx, M> {
463    pub(crate) fn new(
464        genv: GlobalEnv<'genv, 'tcx>,
465        body: &'a Body<'tcx>,
466        bb_envs: &'a mut UnordMap<BasicBlock, Env>,
467        mode: M,
468    ) -> Self {
469        Self {
470            genv,
471            body,
472            bb_envs,
473            discriminants: Default::default(),
474            point: Point::FunEntry,
475            visited: DenseBitSet::new_empty(body.basic_blocks.len()),
476            queue: WorkQueue::empty(body.basic_blocks.len(), &body.dominator_order_rank),
477            mode,
478        }
479    }
480}
481
482impl PlaceNode {
483    fn deref(&mut self) -> (&mut PlaceNode, Modified) {
484        match self {
485            PlaceNode::Deref(_, node) => (node, false),
486            PlaceNode::Ty(ty) => {
487                *self = PlaceNode::Deref(ty.clone(), Box::new(PlaceNode::Ty(ty.deref())));
488                let PlaceNode::Deref(_, node) = self else { unreachable!() };
489                (node, true)
490            }
491            _ => tracked_span_bug!("deref of non-deref place: `{:?}`", self),
492        }
493    }
494
495    fn downcast(
496        &mut self,
497        genv: GlobalEnv,
498        idx: VariantIdx,
499    ) -> QueryResult<(&mut PlaceNode, Modified)> {
500        match self {
501            PlaceNode::Downcast(.., idx2, _) => {
502                debug_assert_eq!(idx, *idx2);
503                Ok((self, false))
504            }
505            PlaceNode::Ty(ty) => {
506                if let TyKind::Adt(adt_def, args) = ty.kind() {
507                    let fields = downcast(genv, adt_def, args, idx)?;
508                    *self = PlaceNode::Downcast(adt_def.clone(), args.clone(), idx, fields);
509                    Ok((self, true))
510                } else {
511                    tracked_span_bug!("invalid downcast `{self:?}`");
512                }
513            }
514            _ => tracked_span_bug!("invalid downcast `{self:?}`"),
515        }
516    }
517
518    fn field(&mut self, genv: GlobalEnv, f: FieldIdx) -> QueryResult<(&mut PlaceNode, Modified)> {
519        let (fields, unfolded) = self.fields(genv)?;
520        Ok((&mut fields[f.as_usize()], unfolded))
521    }
522
523    fn fields(&mut self, genv: GlobalEnv) -> QueryResult<(&mut Vec<PlaceNode>, bool)> {
524        match self {
525            PlaceNode::Ty(ty) => {
526                let fields = match ty.kind() {
527                    TyKind::Adt(adt_def, args) => {
528                        let fields = downcast_struct(genv, adt_def, args)?;
529                        *self = PlaceNode::Downcast(
530                            adt_def.clone(),
531                            args.clone(),
532                            FIRST_VARIANT,
533                            fields,
534                        );
535                        let PlaceNode::Downcast(.., fields) = self else { unreachable!() };
536                        fields
537                    }
538                    TyKind::Closure(def_id, args) => {
539                        let fields = args
540                            .as_closure()
541                            .upvar_tys()
542                            .iter()
543                            .cloned()
544                            .map(PlaceNode::Ty)
545                            .collect_vec();
546                        *self = PlaceNode::Closure(*def_id, args.clone(), fields);
547                        let PlaceNode::Closure(.., fields) = self else { unreachable!() };
548                        fields
549                    }
550                    TyKind::Tuple(fields) => {
551                        let node_fields = fields.iter().cloned().map(PlaceNode::Ty).collect();
552                        *self = PlaceNode::Tuple(fields.clone(), node_fields);
553                        let PlaceNode::Tuple(.., fields) = self else { unreachable!() };
554                        fields
555                    }
556                    TyKind::Coroutine(def_id, args) => {
557                        let fields = args
558                            .as_coroutine()
559                            .upvar_tys()
560                            .cloned()
561                            .map(PlaceNode::Ty)
562                            .collect_vec();
563                        *self = PlaceNode::Generator(*def_id, args.clone(), fields);
564                        let PlaceNode::Generator(.., fields) = self else { unreachable!() };
565                        fields
566                    }
567                    _ => tracked_span_bug!("implicit downcast of non-struct: `{ty:?}`"),
568                };
569                Ok((fields, true))
570            }
571            PlaceNode::Downcast(.., fields)
572            | PlaceNode::Tuple(.., fields)
573            | PlaceNode::Closure(.., fields)
574            | PlaceNode::Generator(.., fields) => Ok((fields, false)),
575            PlaceNode::Deref(..) => {
576                tracked_span_bug!("projection field of non-adt non-tuple place: `{self:?}`")
577            }
578        }
579    }
580
581    fn ensure_folded(&mut self) -> Modified {
582        match self {
583            PlaceNode::Deref(ty, _) => {
584                *self = PlaceNode::Ty(ty.clone());
585                true
586            }
587            PlaceNode::Downcast(adt, args, ..) => {
588                *self = PlaceNode::Ty(Ty::mk_adt(adt.clone(), args.clone()));
589                true
590            }
591            PlaceNode::Closure(did, args, _) => {
592                *self = PlaceNode::Ty(Ty::mk_closure(*did, args.clone()));
593                true
594            }
595            PlaceNode::Generator(did, args, _) => {
596                *self = PlaceNode::Ty(Ty::mk_coroutine(*did, args.clone()));
597                true
598            }
599            PlaceNode::Tuple(fields, ..) => {
600                *self = PlaceNode::Ty(Ty::mk_tuple(fields.clone()));
601                true
602            }
603            PlaceNode::Ty(_) => false,
604        }
605    }
606
607    fn join(
608        &mut self,
609        genv: GlobalEnv,
610        other: &mut PlaceNode,
611        in_mut_ref: bool,
612    ) -> QueryResult<(bool, bool)> {
613        let mut modified1 = false;
614        let mut modified2 = false;
615
616        let (fields1, fields2) = match (&mut *self, &mut *other) {
617            (PlaceNode::Deref(ty1, node1), PlaceNode::Deref(ty2, node2)) => {
618                debug_assert_eq!(ty1, ty2);
619                return node1.join(genv, node2, in_mut_ref || ty1.is_mut_ref());
620            }
621            (PlaceNode::Tuple(_, fields1), PlaceNode::Tuple(_, fields2)) => (fields1, fields2),
622            (PlaceNode::Closure(.., fields1), PlaceNode::Closure(.., fields2)) => {
623                (fields1, fields2)
624            }
625            (PlaceNode::Generator(.., fields1), PlaceNode::Generator(.., fields2)) => {
626                (fields1, fields2)
627            }
628            (
629                PlaceNode::Downcast(adt1, args1, variant1, fields1),
630                PlaceNode::Downcast(adt2, args2, variant2, fields2),
631            ) => {
632                debug_assert_eq!(adt1, adt2);
633                if variant1 == variant2 {
634                    (fields1, fields2)
635                } else {
636                    *self = PlaceNode::Ty(Ty::mk_adt(adt1.clone(), args1.clone()));
637                    *other = PlaceNode::Ty(Ty::mk_adt(adt2.clone(), args2.clone()));
638                    return Ok((true, true));
639                }
640            }
641            (PlaceNode::Ty(_), PlaceNode::Ty(_)) => return Ok((false, false)),
642            (PlaceNode::Ty(_), _) => {
643                let (m1, m2) = other.join(genv, self, in_mut_ref)?;
644                return Ok((m2, m1));
645            }
646            (PlaceNode::Deref(ty, _), _) => {
647                *self = PlaceNode::Ty(ty.clone());
648                return Ok((true, false));
649            }
650            (PlaceNode::Tuple(_, fields1), _) => {
651                let (fields2, m) = other.fields(genv)?;
652                modified2 |= m;
653                (fields1, fields2)
654            }
655            (PlaceNode::Closure(.., fields1), _) | (PlaceNode::Generator(.., fields1), _) => {
656                let (fields2, m) = other.fields(genv)?;
657                modified2 |= m;
658                (fields1, fields2)
659            }
660
661            (PlaceNode::Downcast(adt, args, .., fields1), _) => {
662                if adt.is_struct() && !in_mut_ref {
663                    let (fields2, m) = other.fields(genv)?;
664                    modified2 |= m;
665                    (fields1, fields2)
666                } else {
667                    *self = PlaceNode::Ty(Ty::mk_adt(adt.clone(), args.clone()));
668                    return Ok((true, false));
669                }
670            }
671        };
672        for (node1, node2) in iter::zip(fields1, fields2) {
673            let (m1, m2) = node1.join(genv, node2, in_mut_ref)?;
674            modified1 |= m1;
675            modified2 |= m2;
676        }
677        Ok((modified1, modified2))
678    }
679
680    /// Collect necessary fold/unfold operations such that `self` is unfolded at the same level than `target`
681    fn collect_fold_unfolds(
682        &self,
683        target: &PlaceNode,
684        place: &mut Place,
685        stmts: &mut StatementsAt,
686    ) {
687        let (fields1, fields2) = match (self, target) {
688            (PlaceNode::Deref(_, node1), PlaceNode::Deref(_, node2)) => {
689                place.projection.push(PlaceElem::Deref);
690                node1.collect_fold_unfolds(node2, place, stmts);
691                place.projection.pop();
692                return;
693            }
694            (PlaceNode::Tuple(_, fields1), PlaceNode::Tuple(_, fields2)) => (fields1, fields2),
695            (PlaceNode::Closure(.., fields1), PlaceNode::Closure(.., fields2))
696            | (PlaceNode::Generator(.., fields1), PlaceNode::Generator(.., fields2)) => {
697                (fields1, fields2)
698            }
699            (
700                PlaceNode::Downcast(adt1, .., idx1, fields1),
701                PlaceNode::Downcast(adt2, .., idx2, fields2),
702            ) => {
703                tracked_span_dbg_assert_eq!(adt1.did(), adt2.did());
704                tracked_span_dbg_assert_eq!(idx1, idx2);
705                (fields1, fields2)
706            }
707            (PlaceNode::Ty(_), PlaceNode::Ty(_)) => return,
708            (PlaceNode::Ty(_), _) => {
709                target.collect_unfolds(place, stmts);
710                return;
711            }
712            (_, PlaceNode::Ty(_)) => {
713                stmts.insert(GhostStatement::Fold(place.clone()));
714                return;
715            }
716            _ => tracked_span_bug!("{self:?} {target:?}"),
717        };
718        for (i, (node1, node2)) in iter::zip(fields1, fields2).enumerate() {
719            place.projection.push(PlaceElem::Field(FieldIdx::new(i)));
720            node1.collect_fold_unfolds(node2, place, stmts);
721            place.projection.pop();
722        }
723    }
724
725    fn collect_unfolds(&self, place: &mut Place, stmts: &mut StatementsAt) {
726        match self {
727            PlaceNode::Ty(_) => {}
728            PlaceNode::Deref(_, node) => {
729                if node.is_ty() {
730                    stmts.insert(GhostStatement::Unfold(place.clone()));
731                } else {
732                    place.projection.push(PlaceElem::Deref);
733                    node.collect_unfolds(place, stmts);
734                    place.projection.pop();
735                }
736            }
737            PlaceNode::Downcast(.., fields)
738            | PlaceNode::Closure(.., fields)
739            | PlaceNode::Generator(.., fields)
740            | PlaceNode::Tuple(.., fields) => {
741                let all_leaves = fields.iter().all(PlaceNode::is_ty);
742                if all_leaves {
743                    stmts.insert(GhostStatement::Unfold(place.clone()));
744                } else {
745                    if let Some(idx) = self.enum_variant() {
746                        place.projection.push(PlaceElem::Downcast(None, idx));
747                    }
748                    for (i, node) in fields.iter().enumerate() {
749                        place.projection.push(PlaceElem::Field(FieldIdx::new(i)));
750                        node.collect_unfolds(place, stmts);
751                        place.projection.pop();
752                    }
753                    if self.enum_variant().is_some() {
754                        place.projection.pop();
755                    }
756                }
757            }
758        }
759    }
760
761    fn collect_folds_at_ret(&self, place: &mut Place, stmts: &mut StatementsAt) {
762        let fields = match self {
763            PlaceNode::Deref(ty, deref_ty) => {
764                place.projection.push(PlaceElem::Deref);
765                if ty.is_mut_ref() {
766                    stmts.insert(GhostStatement::Fold(place.clone()));
767                } else if ty.is_box() {
768                    deref_ty.collect_folds_at_ret(place, stmts);
769                }
770                place.projection.pop();
771                return;
772            }
773            PlaceNode::Downcast(adt, _, idx, fields) => {
774                if adt.is_enum() {
775                    place.projection.push(PlaceElem::Downcast(None, *idx));
776                }
777                fields
778            }
779            PlaceNode::Closure(_, _, fields)
780            | PlaceNode::Generator(_, _, fields)
781            | PlaceNode::Tuple(_, fields) => fields,
782            PlaceNode::Ty(_) => return,
783        };
784        for (i, node) in fields.iter().enumerate() {
785            place.projection.push(PlaceElem::Field(FieldIdx::new(i)));
786            node.collect_folds_at_ret(place, stmts);
787            place.projection.pop();
788        }
789        if let PlaceNode::Downcast(adt, ..) = self
790            && adt.is_enum()
791        {
792            place.projection.pop();
793        }
794    }
795
796    fn enum_variant(&self) -> Option<VariantIdx> {
797        if let PlaceNode::Downcast(adt, _, idx, _) = self
798            && adt.is_enum()
799        {
800            Some(*idx)
801        } else {
802            None
803        }
804    }
805
806    /// Returns `true` if the place node is [`Ty`].
807    ///
808    /// [`Ty`]: PlaceNode::Ty
809    #[must_use]
810    fn is_ty(&self) -> bool {
811        matches!(self, Self::Ty(..))
812    }
813}
814
815fn downcast(
816    genv: GlobalEnv,
817    adt_def: &AdtDef,
818    args: &GenericArgs,
819    variant: VariantIdx,
820) -> QueryResult<Vec<PlaceNode>> {
821    adt_def
822        .variant(variant)
823        .fields
824        .iter()
825        .map(|field| {
826            let ty = genv.lower_type_of(field.did)?.subst(args);
827            QueryResult::Ok(PlaceNode::Ty(ty))
828        })
829        .try_collect()
830}
831
832fn downcast_struct(
833    genv: GlobalEnv,
834    adt_def: &AdtDef,
835    args: &GenericArgs,
836) -> QueryResult<Vec<PlaceNode>> {
837    adt_def
838        .non_enum_variant()
839        .fields
840        .iter()
841        .map(|field| {
842            let ty = genv.lower_type_of(field.did)?.subst(args);
843            QueryResult::Ok(PlaceNode::Ty(ty))
844        })
845        .try_collect()
846}
847
848impl fmt::Debug for Env {
849    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
850        write!(
851            f,
852            "{}",
853            self.map
854                .iter_enumerated()
855                .format_with(", ", |(local, node), f| f(&format_args!("{local:?}: {node:?}")))
856        )
857    }
858}
859
860impl fmt::Debug for PlaceNode {
861    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
862        match self {
863            PlaceNode::Deref(_, node) => write!(f, "*({node:?})"),
864            PlaceNode::Downcast(adt, args, variant, fields) => {
865                write!(f, "{}", def_id_to_string(adt.did()))?;
866                if !args.is_empty() {
867                    write!(f, "<{:?}>", args.iter().format(", "),)?;
868                }
869                write!(f, "::{}", adt.variant(*variant).name)?;
870                if !fields.is_empty() {
871                    write!(f, "({:?})", fields.iter().format(", "),)?;
872                }
873                Ok(())
874            }
875            PlaceNode::Closure(did, args, fields) => {
876                write!(f, "Closure {}", def_id_to_string(*did))?;
877                if !args.is_empty() {
878                    write!(f, "<{:?}>", args.iter().format(", "),)?;
879                }
880                if !fields.is_empty() {
881                    write!(f, "({:?})", fields.iter().format(", "),)?;
882                }
883                Ok(())
884            }
885            PlaceNode::Generator(did, args, fields) => {
886                write!(f, "Generator {}", def_id_to_string(*did))?;
887                if !args.is_empty() {
888                    write!(f, "<{:?}>", args.iter().format(", "),)?;
889                }
890                if !fields.is_empty() {
891                    write!(f, "({:?})", fields.iter().format(", "),)?;
892                }
893                Ok(())
894            }
895            PlaceNode::Tuple(_, fields) => write!(f, "({:?})", fields.iter().format(", ")),
896            PlaceNode::Ty(ty) => write!(f, "•{ty:?}"),
897        }
898    }
899}