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 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 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 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 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 #[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}