1use itertools::Itertools as _;
55use rustc_const_eval::const_eval::DummyMachine;
56use rustc_const_eval::interpret::{ImmTy, Immediate, InterpCx, OpTy, Projectable};
57use rustc_data_structures::fx::{FxHashMap, FxHashSet, FxIndexSet};
58use rustc_index::IndexVec;
59use rustc_index::bit_set::{DenseBitSet, GrowableBitSet};
60use rustc_middle::mir::interpret::Scalar;
61use rustc_middle::mir::visit::Visitor;
62use rustc_middle::mir::*;
63use rustc_middle::ty::{self, ScalarInt, TyCtxt};
64use rustc_mir_dataflow::value_analysis::{
65 Map, PlaceCollectionMode, PlaceIndex, TrackElem, ValueIndex,
66};
67use rustc_span::{DUMMY_SP, bug};
68use tracing::{debug, instrument, trace};
69
70use crate::PassPolicy;
71use crate::cost_checker::CostChecker;
72
73pub(super) struct JumpThreading;
74
75const MAX_COST: u8 = 100;
76
77impl<'tcx> crate::MirPass<'tcx> for JumpThreading {
78 fn policy(&self, ctx: &crate::PassCtx<'_>) -> PassPolicy {
79 PassPolicy::optional(ctx.mir_opt_level() >= 2 && !ctx.target.is_like_gpu)
85 }
86
87 #[instrument(skip_all level = "debug")]
88 fn run_pass(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) {
89 let def_id = body.source.def_id();
90 debug!(?def_id);
91
92 if tcx.is_coroutine(def_id) {
94 trace!("Skipped for coroutine {:?}", def_id);
95 return;
96 }
97
98 let typing_env = body.typing_env(tcx);
99 let mut finder = TOFinder {
100 tcx,
101 typing_env,
102 ecx: InterpCx::new(tcx, DUMMY_SP, typing_env, DummyMachine),
103 body,
104 map: Map::new(tcx, body, PlaceCollectionMode::OnDemand),
105 maybe_loop_headers: maybe_loop_headers(body),
106 entry_states: IndexVec::from_elem(ConditionSet::default(), &body.basic_blocks),
107 };
108
109 for (bb, bbdata) in traversal::postorder(body) {
110 if bbdata.is_cleanup {
111 continue;
112 }
113
114 let mut state = finder.populate_from_outgoing_edges(bb);
115 trace!("output_states[{bb:?}] = {state:?}");
116
117 finder.process_terminator(bb, &mut state);
118 trace!("pre_terminator_states[{bb:?}] = {state:?}");
119
120 for stmt in bbdata.statements.iter().rev() {
121 if state.is_empty() {
122 break;
123 }
124
125 finder.process_statement(stmt, &mut state);
126
127 if let Some((lhs, tail)) = finder.mutated_statement(stmt) {
132 finder.flood_state(lhs, tail, &mut state);
133 }
134 }
135
136 trace!("entry_states[{bb:?}] = {state:?}");
137 finder.entry_states[bb] = state;
138 }
139
140 let mut entry_states = finder.entry_states;
141 simplify_conditions(body, &mut entry_states);
142 remove_costly_conditions(tcx, typing_env, body, &mut entry_states);
143
144 if let Some(opportunities) = OpportunitySet::new(body, entry_states) {
145 opportunities.apply();
146 }
147 }
148}
149
150struct TOFinder<'a, 'tcx> {
151 tcx: TyCtxt<'tcx>,
152 typing_env: ty::TypingEnv<'tcx>,
153 ecx: InterpCx<'tcx, DummyMachine>,
154 body: &'a Body<'tcx>,
155 map: Map<'tcx>,
156 maybe_loop_headers: DenseBitSet<BasicBlock>,
157 entry_states: IndexVec<BasicBlock, ConditionSet>,
162}
163
164rustc_index::newtype_index! {
165 #[orderable]
166 #[debug_format = "_c{}"]
167 struct ConditionIndex {}
168}
169
170#[derive(Copy, Clone, Debug, Hash, Eq, PartialEq)]
173struct Condition {
174 place: ValueIndex,
175 value: ScalarInt,
176 polarity: Polarity,
177}
178
179#[derive(Copy, Clone, Debug, Hash, Eq, PartialEq)]
180enum Polarity {
181 Ne,
182 Eq,
183}
184
185impl Condition {
186 fn matches(&self, place: ValueIndex, value: ScalarInt) -> bool {
187 self.place == place && (self.value == value) == (self.polarity == Polarity::Eq)
188 }
189}
190
191#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
193enum EdgeEffect {
194 Goto { target: BasicBlock },
196 Chain { succ_block: BasicBlock, succ_condition: ConditionIndex },
198}
199
200impl EdgeEffect {
201 fn block(self) -> BasicBlock {
202 match self {
203 EdgeEffect::Goto { target: bb } | EdgeEffect::Chain { succ_block: bb, .. } => bb,
204 }
205 }
206
207 fn replace_block(&mut self, target: BasicBlock, new_target: BasicBlock) {
208 match self {
209 EdgeEffect::Goto { target: bb } | EdgeEffect::Chain { succ_block: bb, .. } => {
210 if *bb == target {
211 *bb = new_target
212 }
213 }
214 }
215 }
216}
217
218#[derive(Clone, Debug, Default)]
219struct ConditionSet {
220 active: Vec<(ConditionIndex, Condition)>,
221 fulfilled: Vec<ConditionIndex>,
222 targets: IndexVec<ConditionIndex, Vec<EdgeEffect>>,
223}
224
225impl ConditionSet {
226 fn is_empty(&self) -> bool {
227 self.active.is_empty()
228 }
229
230 #[tracing::instrument(level = "trace", skip(self))]
231 fn push_condition(&mut self, c: Condition, target: BasicBlock) {
232 let index = self.targets.push(vec![EdgeEffect::Goto { target }]);
233 self.active.push((index, c));
234 }
235
236 fn fulfill_if(&mut self, f: impl Fn(Condition, &Vec<EdgeEffect>) -> bool) {
238 self.active.retain(|&(index, condition)| {
239 let targets = &self.targets[index];
240 if f(condition, targets) {
241 trace!(?index, ?condition, "fulfill");
242 self.fulfilled.push(index);
243 false
244 } else {
245 true
246 }
247 })
248 }
249
250 fn fulfill_matches(&mut self, place: ValueIndex, value: ScalarInt) {
252 self.fulfill_if(|c, _| c.matches(place, value))
253 }
254
255 fn retain(&mut self, mut f: impl FnMut(Condition) -> bool) {
256 self.active.retain(|&(_, c)| f(c))
257 }
258
259 fn retain_mut(&mut self, mut f: impl FnMut(Condition) -> Option<Condition>) {
260 self.active.retain_mut(|(_, c)| {
261 if let Some(new) = f(*c) {
262 *c = new;
263 true
264 } else {
265 false
266 }
267 })
268 }
269
270 fn for_each_mut(&mut self, f: impl Fn(&mut Condition)) {
271 for (_, c) in &mut self.active {
272 f(c)
273 }
274 }
275}
276
277impl<'a, 'tcx> TOFinder<'a, 'tcx> {
278 fn place(&mut self, place: Place<'tcx>, tail: Option<TrackElem>) -> Option<PlaceIndex> {
279 self.map.register_place(self.tcx, self.body, place, tail)
280 }
281
282 fn value(&mut self, place: PlaceIndex) -> Option<ValueIndex> {
283 self.map.register_value(self.tcx, self.typing_env, place)
284 }
285
286 fn place_value(&mut self, place: Place<'tcx>, tail: Option<TrackElem>) -> Option<ValueIndex> {
287 let place = self.place(place, tail)?;
288 self.value(place)
289 }
290
291 #[instrument(level = "trace", skip(self))]
293 fn populate_from_outgoing_edges(&mut self, bb: BasicBlock) -> ConditionSet {
294 let bbdata = &self.body[bb];
295
296 debug_assert!(self.entry_states[bb].is_empty());
298
299 let state_len =
300 bbdata.terminator().successors().map(|succ| self.entry_states[succ].active.len()).sum();
301 let mut state = ConditionSet {
302 active: Vec::with_capacity(state_len),
303 targets: IndexVec::with_capacity(state_len),
304 fulfilled: Vec::new(),
305 };
306
307 let mut known_conditions =
309 FxIndexSet::with_capacity_and_hasher(state_len, Default::default());
310 let mut insert = |condition, succ_block, succ_condition| {
311 let (index, new) = known_conditions.insert_full(condition);
312 let index = ConditionIndex::from_usize(index);
313 if new {
314 state.active.push((index, condition));
315 let _index = state.targets.push(Vec::new());
316 debug_assert_eq!(_index, index);
317 }
318 let target = EdgeEffect::Chain { succ_block, succ_condition };
319 debug_assert!(
320 !state.targets[index].contains(&target),
321 "duplicate targets for index={index:?} as {target:?} targets={:#?}",
322 &state.targets[index],
323 );
324 state.targets[index].push(target);
325 };
326
327 let mut seen = FxHashSet::default();
329 for succ in bbdata.terminator().successors() {
330 if !seen.insert(succ) {
331 continue;
332 }
333
334 if self.maybe_loop_headers.contains(succ) {
336 continue;
337 }
338
339 for &(succ_index, cond) in self.entry_states[succ].active.iter() {
340 insert(cond, succ, succ_index);
341 }
342 }
343
344 let num_conditions = known_conditions.len();
345 debug_assert_eq!(num_conditions, state.active.len());
346 debug_assert_eq!(num_conditions, state.targets.len());
347 state.fulfilled.reserve(num_conditions);
348
349 state
350 }
351
352 fn flood_state(
354 &self,
355 place: Place<'tcx>,
356 extra_elem: Option<TrackElem>,
357 state: &mut ConditionSet,
358 ) {
359 if state.is_empty() {
360 return;
361 }
362 let mut places_to_exclude = FxHashSet::default();
363 self.map.for_each_aliasing_place(place.as_ref(), extra_elem, &mut |vi| {
364 places_to_exclude.insert(vi);
365 });
366 trace!(?places_to_exclude, "flood_state");
367 if places_to_exclude.is_empty() {
368 return;
369 }
370 state.retain(|c| !places_to_exclude.contains(&c.place));
371 }
372
373 #[instrument(level = "trace", skip(self), ret)]
387 fn mutated_statement(
388 &self,
389 stmt: &Statement<'tcx>,
390 ) -> Option<(Place<'tcx>, Option<TrackElem>)> {
391 match stmt.kind {
392 StatementKind::Assign((place, _)) => Some((place, None)),
393 StatementKind::SetDiscriminant { ref place, variant_index: _ } => {
394 Some((**place, Some(TrackElem::Discriminant)))
395 }
396 StatementKind::StorageLive(local) | StatementKind::StorageDead(local) => {
397 Some((Place::from(local), None))
398 }
399 | StatementKind::Intrinsic(NonDivergingIntrinsic::Assume(..))
400 | StatementKind::Intrinsic(NonDivergingIntrinsic::CopyNonOverlapping(..))
402 | StatementKind::AscribeUserType(..)
403 | StatementKind::Coverage(..)
404 | StatementKind::FakeRead(..)
405 | StatementKind::ConstEvalCounter
406 | StatementKind::PlaceMention(..)
407 | StatementKind::BackwardIncompatibleDropHint { .. }
408 | StatementKind::Nop => None,
409 }
410 }
411
412 #[instrument(level = "trace", skip(self, state))]
413 fn process_immediate(&mut self, lhs: PlaceIndex, rhs: ImmTy<'tcx>, state: &mut ConditionSet) {
414 if let Some(lhs) = self.value(lhs)
415 && let Immediate::Scalar(Scalar::Int(int)) = *rhs
416 {
417 state.fulfill_matches(lhs, int)
418 }
419 }
420
421 #[instrument(level = "trace", skip(self, state))]
423 fn process_constant(
424 &mut self,
425 lhs: PlaceIndex,
426 constant: OpTy<'tcx>,
427 state: &mut ConditionSet,
428 ) {
429 self.map.for_each_projection_value(
430 lhs,
431 constant,
432 &mut |elem, op| match elem {
433 TrackElem::Field(idx) => self.ecx.project_field(op, idx).discard_err(),
434 TrackElem::Variant(idx) => self.ecx.project_downcast(op, idx).discard_err(),
435 TrackElem::Discriminant => {
436 let variant = self.ecx.read_discriminant(op).discard_err()?;
437 let discr_value =
438 self.ecx.discriminant_for_variant(op.layout.ty, variant).discard_err()?;
439 Some(discr_value.into())
440 }
441 TrackElem::DerefLen => {
442 let op: OpTy<'_> = self.ecx.deref_pointer(op).discard_err()?.into();
443 let len_usize = op.len(&self.ecx).discard_err()?;
444 let layout = self.ecx.layout_of(self.tcx.types.usize).unwrap();
445 Some(ImmTy::from_uint(len_usize, layout).into())
446 }
447 },
448 &mut |place, op| {
449 if let Some(place) = self.map.value(place)
450 && let Some(imm) = self.ecx.read_immediate_raw(op).discard_err()
451 && let Some(imm) = imm.right()
452 && let Immediate::Scalar(Scalar::Int(int)) = *imm
453 {
454 state.fulfill_matches(place, int)
455 }
456 },
457 );
458 }
459
460 #[instrument(level = "trace", skip(self, state))]
461 fn process_copy(&mut self, lhs: PlaceIndex, rhs: PlaceIndex, state: &mut ConditionSet) {
462 let mut renames = FxHashMap::default();
463 self.map.register_copy_tree(
464 lhs, rhs, &mut |lhs, rhs| {
467 renames.insert(lhs, rhs);
468 },
469 );
470 state.for_each_mut(|c| {
471 if let Some(rhs) = renames.get(&c.place) {
472 c.place = *rhs
473 }
474 });
475 }
476
477 #[instrument(level = "trace", skip(self, state))]
478 fn process_operand(&mut self, lhs: PlaceIndex, rhs: &Operand<'tcx>, state: &mut ConditionSet) {
479 match rhs {
480 Operand::Constant(constant) => {
482 let Some(constant) =
483 self.ecx.eval_mir_constant(&constant.const_, constant.span, None).discard_err()
484 else {
485 return;
486 };
487 self.process_constant(lhs, constant, state);
488 }
489 Operand::Move(rhs) | Operand::Copy(rhs) => {
491 let Some(rhs) = self.place(*rhs, None) else { return };
492 self.process_copy(lhs, rhs, state)
493 }
494 Operand::RuntimeChecks(_) => {}
495 }
496 }
497
498 #[instrument(level = "trace", skip(self, state))]
499 fn process_assign(
500 &mut self,
501 lhs_place: &Place<'tcx>,
502 rvalue: &Rvalue<'tcx>,
503 state: &mut ConditionSet,
504 ) {
505 let Some(lhs) = self.place(*lhs_place, None) else { return };
506 match rvalue {
507 Rvalue::Use(operand, _) => self.process_operand(lhs, operand, state),
508 Rvalue::Discriminant(rhs) => {
510 let Some(rhs) = self.place(*rhs, Some(TrackElem::Discriminant)) else { return };
511 self.process_copy(lhs, rhs, state)
512 }
513 Rvalue::Aggregate(kind, operands) => {
515 let agg_ty = lhs_place.ty(self.body, self.tcx).ty;
516 let lhs = match kind {
517 AggregateKind::Adt(.., Some(_)) => return,
519 AggregateKind::Adt(_, variant_index, ..) if agg_ty.is_enum() => {
520 let discr_ty = agg_ty.discriminant_ty(self.tcx);
521 let discr_target =
522 self.map.register_place_index(discr_ty, lhs, TrackElem::Discriminant);
523 if let Some(discr_value) =
524 self.ecx.discriminant_for_variant(agg_ty, *variant_index).discard_err()
525 {
526 self.process_immediate(discr_target, discr_value, state);
527 }
528 self.map.register_place_index(
529 agg_ty,
530 lhs,
531 TrackElem::Variant(*variant_index),
532 )
533 }
534 _ => lhs,
535 };
536 for (field_index, operand) in operands.iter_enumerated() {
537 let operand_ty = operand.ty(self.body, self.tcx);
538 let field = self.map.register_place_index(
539 operand_ty,
540 lhs,
541 TrackElem::Field(field_index),
542 );
543 self.process_operand(field, operand, state);
544 }
545 }
546 Rvalue::UnaryOp(UnOp::Not, Operand::Move(operand) | Operand::Copy(operand)) => {
548 let layout = self.ecx.layout_of(operand.ty(self.body, self.tcx).ty).unwrap();
549 let Some(lhs) = self.value(lhs) else { return };
550 let Some(operand) = self.place_value(*operand, None) else { return };
551 state.retain_mut(|mut c| {
552 if c.place == lhs {
553 let value = self
554 .ecx
555 .unary_op(UnOp::Not, &ImmTy::from_scalar_int(c.value, layout))
556 .discard_err()?
557 .to_scalar_int()
558 .discard_err()?;
559 c.place = operand;
560 c.value = value;
561 }
562 Some(c)
563 });
564 }
565 Rvalue::BinaryOp(
568 op,
569 (Operand::Move(operand) | Operand::Copy(operand), Operand::Constant(value))
570 | (Operand::Constant(value), Operand::Move(operand) | Operand::Copy(operand)),
571 ) => {
572 let equals = match op {
573 BinOp::Eq => ScalarInt::TRUE,
574 BinOp::Ne => ScalarInt::FALSE,
575 _ => return,
576 };
577 if value.const_.ty().is_floating_point() {
578 return;
583 }
584 let Some(lhs) = self.value(lhs) else { return };
585 let Some(operand) = self.place_value(*operand, None) else { return };
586 let Some(value) = value.const_.try_eval_scalar_int(self.tcx, self.typing_env)
587 else {
588 return;
589 };
590 state.for_each_mut(|c| {
591 if c.place == lhs {
592 let polarity =
593 if c.matches(lhs, equals) { Polarity::Eq } else { Polarity::Ne };
594 c.place = operand;
595 c.value = value;
596 c.polarity = polarity;
597 }
598 });
599 }
600
601 _ => {}
602 }
603 }
604
605 #[instrument(level = "trace", skip(self, state))]
606 fn process_statement(&mut self, stmt: &Statement<'tcx>, state: &mut ConditionSet) {
607 match &stmt.kind {
611 StatementKind::SetDiscriminant { place, variant_index } => {
614 let Some(discr_target) = self.place(**place, Some(TrackElem::Discriminant)) else {
615 return;
616 };
617 let enum_ty = place.ty(self.body, self.tcx).ty;
618 let Some(discr) =
622 self.ecx.discriminant_for_variant(enum_ty, *variant_index).discard_err()
623 else {
624 return;
625 };
626 self.process_immediate(discr_target, discr, state)
627 }
628 StatementKind::Intrinsic(NonDivergingIntrinsic::Assume(
630 Operand::Copy(place) | Operand::Move(place),
631 )) => {
632 let Some(place) = self.place_value(*place, None) else { return };
633 state.fulfill_matches(place, ScalarInt::TRUE);
634 }
635 StatementKind::Assign((lhs_place, rhs)) => self.process_assign(lhs_place, rhs, state),
636 _ => {}
637 }
638 }
639
640 #[instrument(level = "trace", skip(self, state))]
642 fn process_terminator(&mut self, bb: BasicBlock, state: &mut ConditionSet) {
643 let term = self.body.basic_blocks[bb].terminator();
644 let place_to_flood = match term.kind {
645 TerminatorKind::FalseEdge { .. }
647 | TerminatorKind::FalseUnwind { .. }
648 | TerminatorKind::Yield { .. } => bug!("{term:?} invalid"),
649 TerminatorKind::InlineAsm { .. } => {
651 state.active.clear();
652 return;
653 }
654 TerminatorKind::SwitchInt { ref discr, ref targets } => {
656 return self.process_switch_int(discr, targets, state);
657 }
658 TerminatorKind::UnwindResume
660 | TerminatorKind::UnwindTerminate(_)
661 | TerminatorKind::Return
662 | TerminatorKind::Unreachable
663 | TerminatorKind::CoroutineDrop
664 | TerminatorKind::Assert { .. }
666 | TerminatorKind::Goto { .. } => None,
667 TerminatorKind::Drop { place: destination, .. }
669 | TerminatorKind::Call { destination, .. } => Some(destination),
670 TerminatorKind::TailCall { .. } => Some(RETURN_PLACE.into()),
671 };
672
673 if let Some(place_to_flood) = place_to_flood {
675 self.flood_state(place_to_flood, None, state);
676 }
677 }
678
679 #[instrument(level = "trace", skip(self))]
680 fn process_switch_int(
681 &mut self,
682 discr: &Operand<'tcx>,
683 targets: &SwitchTargets,
684 state: &mut ConditionSet,
685 ) {
686 let Some(discr) = discr.place() else { return };
687 let Some(discr_idx) = self.place_value(discr, None) else { return };
688
689 let discr_ty = discr.ty(self.body, self.tcx).ty;
690 let Ok(discr_layout) = self.ecx.layout_of(discr_ty) else { return };
691
692 if targets.is_distinct() {
695 for &(index, c) in state.active.iter() {
696 if c.place != discr_idx {
697 continue;
698 }
699
700 let mut edges_fulfilling_condition = FxHashSet::default();
702
703 for (branch, tgt) in targets.iter() {
705 if let Some(branch) = ScalarInt::try_from_uint(branch, discr_layout.size)
706 && c.matches(discr_idx, branch)
707 {
708 edges_fulfilling_condition.insert(tgt);
709 }
710 }
711
712 if c.polarity == Polarity::Ne
717 && let value = c.value.to_bits(discr_layout.size)
718 && targets.all_values().contains(&value.into())
719 {
720 edges_fulfilling_condition.insert(targets.otherwise());
721 }
722
723 let condition_targets = &state.targets[index];
727
728 let new_edges: Vec<_> = condition_targets
729 .iter()
730 .copied()
731 .filter(|&target| match target {
732 EdgeEffect::Goto { .. } => false,
733 EdgeEffect::Chain { succ_block, .. } => {
734 edges_fulfilling_condition.contains(&succ_block)
735 }
736 })
737 .collect();
738
739 if new_edges.len() == condition_targets.len() {
740 state.fulfilled.push(index);
743 } else {
744 let index = state.targets.push(new_edges);
747 state.fulfilled.push(index);
748 }
749 }
750 }
751
752 let mut mk_condition = |value, polarity, target| {
754 let c = Condition { place: discr_idx, value, polarity };
755 state.push_condition(c, target);
756 };
757 if let Some((value, then_, else_)) = targets.as_static_if() {
758 let Some(value) = ScalarInt::try_from_uint(value, discr_layout.size) else { return };
760 mk_condition(value, Polarity::Eq, then_);
761 mk_condition(value, Polarity::Ne, else_);
762 } else {
763 for (value, target) in targets.iter() {
766 if let Some(value) = ScalarInt::try_from_uint(value, discr_layout.size) {
767 mk_condition(value, Polarity::Eq, target);
768 }
769 }
770 }
771 }
772}
773
774#[instrument(level = "debug", skip(body, entry_states))]
776fn simplify_conditions(body: &Body<'_>, entry_states: &mut IndexVec<BasicBlock, ConditionSet>) {
777 let basic_blocks = &body.basic_blocks;
778 let reverse_postorder = basic_blocks.reverse_postorder();
779
780 let mut predecessors = IndexVec::from_elem(0, &entry_states);
783 predecessors[START_BLOCK] = 1; for &bb in reverse_postorder {
785 let term = basic_blocks[bb].terminator();
786 for s in term.successors() {
787 predecessors[s] += 1;
788 }
789 }
790
791 let mut fulfill_in_pred_count = IndexVec::from_fn_n(
793 |bb: BasicBlock| IndexVec::from_elem_n(0, entry_states[bb].targets.len()),
794 entry_states.len(),
795 );
796
797 for &bb in reverse_postorder {
799 let preds = predecessors[bb];
800 trace!(?bb, ?preds);
801
802 if preds == 0 {
804 continue;
805 }
806
807 let state = &mut entry_states[bb];
808 trace!(?state);
809
810 trace!(fulfilled_count = ?fulfill_in_pred_count[bb]);
812 for (condition, &cond_preds) in fulfill_in_pred_count[bb].iter_enumerated() {
813 if cond_preds == preds {
814 trace!(?condition);
815 state.fulfilled.push(condition);
816 }
817 }
818
819 let mut targets: Vec<_> = state
822 .fulfilled
823 .iter()
824 .flat_map(|&index| state.targets[index].iter().copied())
825 .collect();
826 targets.sort();
827 targets.dedup();
828 trace!(?targets);
829
830 let mut successors = basic_blocks[bb].terminator().successors().collect::<Vec<_>>();
832
833 targets.reverse();
834 while let Some(target) = targets.pop() {
835 match target {
836 EdgeEffect::Goto { target } => {
837 predecessors[target] += 1;
840 for &s in successors.iter() {
841 predecessors[s] -= 1;
842 }
843 targets.retain(|t| t.block() == target);
845 successors.clear();
846 successors.push(target);
847 }
848 EdgeEffect::Chain { succ_block, succ_condition } => {
849 let count = successors.iter().filter(|&&s| s == succ_block).count();
852 fulfill_in_pred_count[succ_block][succ_condition] += count;
853 }
854 }
855 }
856 }
857}
858
859#[instrument(level = "debug", skip(tcx, typing_env, body, entry_states))]
860fn remove_costly_conditions<'tcx>(
861 tcx: TyCtxt<'tcx>,
862 typing_env: ty::TypingEnv<'tcx>,
863 body: &Body<'tcx>,
864 entry_states: &mut IndexVec<BasicBlock, ConditionSet>,
865) {
866 let basic_blocks = &body.basic_blocks;
867
868 let mut costs = IndexVec::from_elem(None, basic_blocks);
869 let mut cost = |bb: BasicBlock| -> u8 {
870 let c = *costs[bb].get_or_insert_with(|| {
871 let bbdata = &basic_blocks[bb];
872 let mut cost = CostChecker::new(tcx, typing_env, None, body);
873 cost.visit_basic_block_data(bb, bbdata);
874 cost.cost().try_into().unwrap_or(MAX_COST)
875 });
876 trace!("cost[{bb:?}] = {c}");
877 c
878 };
879
880 let mut condition_cost = IndexVec::from_fn_n(
882 |bb: BasicBlock| IndexVec::from_elem_n(MAX_COST, entry_states[bb].targets.len()),
883 entry_states.len(),
884 );
885
886 let reverse_postorder = basic_blocks.reverse_postorder();
887
888 for &bb in reverse_postorder.iter().rev() {
889 let state = &entry_states[bb];
890 trace!(?bb, ?state);
891
892 let mut current_costs = IndexVec::from_elem(0u8, &state.targets);
893
894 for (condition, targets) in state.targets.iter_enumerated() {
895 for &target in targets {
896 match target {
897 EdgeEffect::Goto { .. } => {}
899 EdgeEffect::Chain { succ_block, succ_condition }
901 if entry_states[succ_block].fulfilled.contains(&succ_condition) => {}
902 EdgeEffect::Chain { succ_block, succ_condition } => {
904 let duplication_cost = cost(succ_block);
906 let target_cost =
908 *condition_cost[succ_block].get(succ_condition).unwrap_or(&MAX_COST);
909 let cost = current_costs[condition]
910 .saturating_add(duplication_cost)
911 .saturating_add(target_cost);
912 trace!(?condition, ?succ_block, ?duplication_cost, ?target_cost);
913 current_costs[condition] = cost;
914 }
915 }
916 }
917 }
918
919 trace!("condition_cost[{bb:?}] = {:?}", current_costs);
920 condition_cost[bb] = current_costs;
921 }
922
923 trace!(?condition_cost);
924
925 for &bb in reverse_postorder {
926 for (index, targets) in entry_states[bb].targets.iter_enumerated_mut() {
927 if condition_cost[bb][index] >= MAX_COST {
928 trace!(?bb, ?index, ?targets, c = ?condition_cost[bb][index], "remove");
929 targets.clear()
930 }
931 }
932 }
933}
934
935struct OpportunitySet<'a, 'tcx> {
936 basic_blocks: &'a mut IndexVec<BasicBlock, BasicBlockData<'tcx>>,
937 entry_states: IndexVec<BasicBlock, ConditionSet>,
938 duplicates: FxHashMap<(BasicBlock, ConditionIndex), BasicBlock>,
941}
942
943impl<'a, 'tcx> OpportunitySet<'a, 'tcx> {
944 fn new(
945 body: &'a mut Body<'tcx>,
946 mut entry_states: IndexVec<BasicBlock, ConditionSet>,
947 ) -> Option<OpportunitySet<'a, 'tcx>> {
948 trace!(def_id = ?body.source.def_id(), "apply");
949
950 if entry_states.iter().all(|state| state.fulfilled.is_empty()) {
951 return None;
952 }
953
954 for state in entry_states.iter_mut() {
956 state.active = Default::default();
957 }
958 let duplicates = Default::default();
959 let basic_blocks = body.basic_blocks.as_mut();
960 Some(OpportunitySet { basic_blocks, entry_states, duplicates })
961 }
962
963 #[instrument(level = "debug", skip(self))]
965 fn apply(mut self) {
966 let mut worklist = Vec::with_capacity(self.basic_blocks.len());
967 worklist.push(START_BLOCK);
968
969 let mut visited = GrowableBitSet::with_capacity(self.basic_blocks.len());
971
972 while let Some(bb) = worklist.pop() {
973 if !visited.insert(bb) {
974 continue;
975 }
976
977 self.apply_once(bb);
978
979 worklist.extend(self.basic_blocks[bb].terminator().successors());
982 }
983 }
984
985 #[instrument(level = "debug", skip(self))]
987 fn apply_once(&mut self, bb: BasicBlock) {
988 let state = &mut self.entry_states[bb];
989 trace!(?state);
990
991 let mut targets: Vec<_> = state
994 .fulfilled
995 .iter()
996 .flat_map(|&index| std::mem::take(&mut state.targets[index]))
997 .collect();
998 targets.sort();
999 targets.dedup();
1000 trace!(?targets);
1001
1002 targets.reverse();
1004 while let Some(target) = targets.pop() {
1005 debug!(?target);
1006 trace!(term = ?self.basic_blocks[bb].terminator().kind);
1007
1008 debug_assert!(
1012 self.basic_blocks[bb].terminator().successors().contains(&target.block()),
1013 "missing {target:?} in successors for {bb:?}, term={:?}",
1014 self.basic_blocks[bb].terminator(),
1015 );
1016
1017 match target {
1018 EdgeEffect::Goto { target } => {
1019 self.apply_goto(bb, target);
1020
1021 targets.retain(|t| t.block() == target);
1023 for ts in self.entry_states[bb].targets.iter_mut() {
1025 ts.retain(|t| t.block() == target);
1026 }
1027 }
1028 EdgeEffect::Chain { succ_block, succ_condition } => {
1029 let new_succ_block = self.apply_chain(bb, succ_block, succ_condition);
1030
1031 if let Some(new_succ_block) = new_succ_block {
1033 for t in targets.iter_mut() {
1034 t.replace_block(succ_block, new_succ_block)
1035 }
1036 for t in
1038 self.entry_states[bb].targets.iter_mut().flat_map(|ts| ts.iter_mut())
1039 {
1040 t.replace_block(succ_block, new_succ_block)
1041 }
1042 }
1043 }
1044 }
1045
1046 trace!(post_term = ?self.basic_blocks[bb].terminator().kind);
1047 }
1048 }
1049
1050 #[instrument(level = "debug", skip(self))]
1051 fn apply_goto(&mut self, bb: BasicBlock, target: BasicBlock) {
1052 self.basic_blocks[bb].terminator_mut().kind = TerminatorKind::Goto { target };
1053 }
1054
1055 #[instrument(level = "debug", skip(self), ret)]
1056 fn apply_chain(
1057 &mut self,
1058 bb: BasicBlock,
1059 target: BasicBlock,
1060 condition: ConditionIndex,
1061 ) -> Option<BasicBlock> {
1062 if self.entry_states[target].fulfilled.contains(&condition) {
1063 trace!("fulfilled");
1065 return None;
1066 }
1067
1068 let new_target = *self.duplicates.entry((target, condition)).or_insert_with(|| {
1074 let new_target = self.basic_blocks.push(self.basic_blocks[target].clone());
1077 trace!(?target, ?new_target, ?condition, "clone");
1078
1079 let mut condition_set = self.entry_states[target].clone();
1082 condition_set.fulfilled.push(condition);
1083 let _new_target = self.entry_states.push(condition_set);
1084 debug_assert_eq!(new_target, _new_target);
1085
1086 new_target
1087 });
1088 trace!(?target, ?new_target, ?condition, "reuse");
1089
1090 self.basic_blocks[bb].terminator_mut().successors_mut(|s| {
1093 if *s == target {
1094 *s = new_target;
1095 }
1096 });
1097
1098 Some(new_target)
1099 }
1100}
1101
1102fn maybe_loop_headers(body: &Body<'_>) -> DenseBitSet<BasicBlock> {
1108 let mut maybe_loop_headers = DenseBitSet::new_empty(body.basic_blocks.len());
1109 let mut visited = DenseBitSet::new_empty(body.basic_blocks.len());
1110 for (bb, bbdata) in traversal::postorder(body) {
1111 for succ in bbdata.terminator().successors() {
1114 if !visited.contains(succ) {
1115 maybe_loop_headers.insert(succ);
1116 }
1117 }
1118
1119 let _new = visited.insert(bb);
1122 debug_assert!(_new);
1123 }
1124
1125 maybe_loop_headers
1126}