From 85d9d2ec30add54f3bbf1210b7d09a464a8bf39b Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Mon, 17 Aug 2026 16:35:40 +0200 Subject: [PATCH 01/12] Create a new pass --- .../src/collapse_yields.rs | 67 +++++++++++++++++++ compiler/rustc_mir_transform/src/lib.rs | 2 + 2 files changed, 69 insertions(+) create mode 100644 compiler/rustc_mir_transform/src/collapse_yields.rs diff --git a/compiler/rustc_mir_transform/src/collapse_yields.rs b/compiler/rustc_mir_transform/src/collapse_yields.rs new file mode 100644 index 0000000000000..690b2f0e48abc --- /dev/null +++ b/compiler/rustc_mir_transform/src/collapse_yields.rs @@ -0,0 +1,67 @@ +use rustc_middle::mir::{BasicBlock, Body, Operand, Place, TerminatorKind}; +use rustc_middle::ty::TyCtxt; +use tracing::instrument; + +use crate::MirPass; +use crate::pass_manager::PassPolicy; + +pub(super) struct CollapseIdenticalYields; + +impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { + #[instrument(level = "debug", skip(self, tcx, body), ret)] + fn run_pass(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) { + if body.coroutine_kind().is_none() { + return; + } + + tracing::debug!("running pass for {}", tcx.def_path_str(body.source.def_id())); + + let yields = body + .basic_blocks + .iter_enumerated() + .filter_map(|(bb, bb_data)| match &bb_data.terminator.as_ref()?.kind { + TerminatorKind::Yield { value, resume, resume_arg, drop } => Some(Yield { + basic_block: bb, + value: value.clone(), + resume: *resume, + resume_arg: *resume_arg, + drop: *drop, + }), + _ => None, + }) + .collect::>(); + + let resume_loops = + yields.iter().map(|yield_val| ResumeLoop::search(yield_val, body)).collect::>(); + + tracing::trace!("resume loops: {resume_loops:?}"); + } + + fn policy(&self, _sess: &rustc_session::Session) -> PassPolicy { + PassPolicy::optimization(true) + } +} + +struct Yield<'tcx> { + basic_block: BasicBlock, + /// The value to return. + value: Operand<'tcx>, + /// Where to resume to. + resume: BasicBlock, + /// The place to store the resume argument in. + resume_arg: Place<'tcx>, + /// Cleanup to be done if the coroutine is dropped at this suspend point. + drop: Option, +} + +#[derive(Debug)] +struct ResumeLoop { + inner: Vec, +} + +impl ResumeLoop { + fn search(yield_val: &Yield, body: &Body<'_>) -> Option { + // FIXME + None + } +} diff --git a/compiler/rustc_mir_transform/src/lib.rs b/compiler/rustc_mir_transform/src/lib.rs index 95a708f125660..bc20cfb172752 100644 --- a/compiler/rustc_mir_transform/src/lib.rs +++ b/compiler/rustc_mir_transform/src/lib.rs @@ -136,6 +136,7 @@ declare_passes! { pub mod cleanup_post_borrowck : CleanupPostBorrowck; mod copy_prop : CopyProp; + mod collapse_yields : CollapseIdenticalYields; mod coroutine : StateTransform; mod coverage : InstrumentCoverage; mod ctfe_limit : CtfeLimit; @@ -661,6 +662,7 @@ fn run_runtime_lowering_passes<'tcx>(tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) { &add_moves_for_packed_drops::AddMovesForPackedDrops, &erase_deref_temps::EraseDerefTemps, &elaborate_box_derefs::ElaborateBoxDerefs, + &collapse_yields::CollapseIdenticalYields, &coroutine::StateTransform, &Lint(known_panics_lint::KnownPanicsLint), ]; From 1f7efdfdeffcef2b99507cb1be7258197cd87280 Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Tue, 18 Aug 2026 16:29:02 +0200 Subject: [PATCH 02/12] Build out more of the detection and translation --- .../src/collapse_yields.rs | 174 ++++++++++++++++-- 1 file changed, 163 insertions(+), 11 deletions(-) diff --git a/compiler/rustc_mir_transform/src/collapse_yields.rs b/compiler/rustc_mir_transform/src/collapse_yields.rs index 690b2f0e48abc..278b1322b2190 100644 --- a/compiler/rustc_mir_transform/src/collapse_yields.rs +++ b/compiler/rustc_mir_transform/src/collapse_yields.rs @@ -1,4 +1,9 @@ -use rustc_middle::mir::{BasicBlock, Body, Operand, Place, TerminatorKind}; +use rustc_data_structures::fx::FxHashMap; +use rustc_data_structures::graph::Successors; +use rustc_data_structures::indexmap::IndexSet; +use rustc_middle::mir::{ + BasicBlock, Body, Local, Operand, Place, Statement, StatementKind, TerminatorKind, +}; use rustc_middle::ty::TyCtxt; use tracing::instrument; @@ -16,7 +21,7 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { tracing::debug!("running pass for {}", tcx.def_path_str(body.source.def_id())); - let yields = body + let mut yields = body .basic_blocks .iter_enumerated() .filter_map(|(bb, bb_data)| match &bb_data.terminator.as_ref()?.kind { @@ -31,10 +36,24 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { }) .collect::>(); - let resume_loops = - yields.iter().map(|yield_val| ResumeLoop::search(yield_val, body)).collect::>(); + while let Some(base_yield) = yields.pop() { + for i in (0..yields.len()).rev() { + let compare_yield = &yields[i]; - tracing::trace!("resume loops: {resume_loops:?}"); + let Some(translation) = base_yield.try_find_translation(compare_yield, body) else { + // No translation, so these yields aren't equivalent + continue; + }; + + // FIXME: impl doing the translation with a visitor + tracing::trace!( + "Translate {:?} to {:?}: {:?}", + compare_yield.basic_block, + base_yield.basic_block, + translation + ); + } + } } fn policy(&self, _sess: &rustc_session::Session) -> PassPolicy { @@ -42,6 +61,7 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { } } +#[derive(Debug)] struct Yield<'tcx> { basic_block: BasicBlock, /// The value to return. @@ -54,14 +74,146 @@ struct Yield<'tcx> { drop: Option, } +impl Yield<'_> { + fn try_find_translation(&self, other: &Yield<'_>, body: &Body<'_>) -> Option { + let mut self_successors = self.all_successors(body); + let mut other_successors = other.all_successors(body); + + let mut map = TranslationMap::new(); + + for (self_successor, other_successor) in (&mut self_successors).zip(&mut other_successors) { + if self_successor == other_successor { + continue; + } + map = map.try_add_translation(self_successor, other_successor, body)?; + } + + if self_successors.next().is_some() || other_successors.next().is_some() { + // Can't be the same if they're not the same length + return None; + } + + Some(map) + } + + fn all_successors(&self, body: &Body<'_>) -> impl Iterator { + let mut seen: IndexSet = IndexSet::default(); + let mut todo = vec![self.basic_block]; + + std::iter::from_fn(move || { + let next = todo.pop()?; + for successor in body.basic_blocks.successors(next) { + if !seen.insert(successor) { + continue; + } + + todo.push(successor); + } + Some(next) + }) + } +} + #[derive(Debug)] -struct ResumeLoop { - inner: Vec, +struct TranslationMap { + blocks: FxHashMap, + locals: FxHashMap, } -impl ResumeLoop { - fn search(yield_val: &Yield, body: &Body<'_>) -> Option { - // FIXME - None +impl TranslationMap { + fn new() -> Self { + Self { blocks: Default::default(), locals: Default::default() } + } + + fn try_add_translation( + mut self, + self_bb: BasicBlock, + other_bb: BasicBlock, + body: &Body<'_>, + ) -> Option { + let self_data = &body.basic_blocks[self_bb]; + let other_data = &body.basic_blocks[other_bb]; + + if self_data.is_cleanup != other_data.is_cleanup { + return None; + } + + if self_data.statements.len() != other_data.statements.len() { + return None; + } + + for (self_statement, other_statement) in + self_data.statements.iter().zip(other_data.statements.iter()) + { + let locals_map = + Self::statements_functionally_equivalent(self_statement, other_statement)?; + for (self_local, other_local) in locals_map { + if let Some(old_other_local) = self.locals.insert(self_local, other_local) { + if old_other_local != other_local { + // We're not equivalent since a local has changed + return None; + } + } + } + } + + // FIXME: Consider terminators + + self.blocks.insert(self_bb, other_bb); + + Some(self) + } + + fn statements_functionally_equivalent( + l: &Statement<'_>, + r: &Statement<'_>, + ) -> Option> { + let mut map = FxHashMap::default(); + + match (&l.kind, &r.kind) { + (StatementKind::Assign(l_assign), StatementKind::Assign(r_assign)) => { + map.insert(l_assign.0.local, r_assign.0.local); + unimplemented!(); + } + (StatementKind::FakeRead(l_fake_read), StatementKind::FakeRead(r_fake_read)) => { + unimplemented!() + } + ( + StatementKind::SetDiscriminant { place: l_place, variant_index: l_variant_index }, + StatementKind::SetDiscriminant { place: r_place, variant_index: r_variant_index }, + ) => unimplemented!(), + (StatementKind::StorageLive(l_local), StatementKind::StorageLive(r_local)) => { + unimplemented!() + } + (StatementKind::StorageDead(l_local), StatementKind::StorageDead(r_local)) => { + unimplemented!() + } + (StatementKind::PlaceMention(l_place), StatementKind::PlaceMention(r_place)) => { + unimplemented!() + } + ( + StatementKind::AscribeUserType(l_ascribe_user_type, l_variance), + StatementKind::AscribeUserType(r_ascribe_user_type, r_variance), + ) => unimplemented!(), + ( + StatementKind::Coverage(l_coverage_kind), + StatementKind::Coverage(r_coverage_kind), + ) => { + unimplemented!() + } + ( + StatementKind::Intrinsic(l_non_diverging_intrinsic), + StatementKind::Intrinsic(r_non_diverging_intrinsic), + ) => unimplemented!(), + (StatementKind::ConstEvalCounter, StatementKind::ConstEvalCounter) => {} + (StatementKind::Nop, StatementKind::Nop) => {} + ( + StatementKind::BackwardIncompatibleDropHint { place: l_place, reason: l_reason }, + StatementKind::BackwardIncompatibleDropHint { place: r_place, reason: r_reason }, + ) => unimplemented!(), + _ => {} + } + + Some(map) } } From 83fd47c619f3555d7d864d60a5e1eb1058a9396a Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Mon, 24 Aug 2026 15:34:42 +0200 Subject: [PATCH 03/12] Impl full comparison --- .../src/collapse_yields.rs | 589 ++++++++++++++++-- 1 file changed, 538 insertions(+), 51 deletions(-) diff --git a/compiler/rustc_mir_transform/src/collapse_yields.rs b/compiler/rustc_mir_transform/src/collapse_yields.rs index 278b1322b2190..f900bd2ba5830 100644 --- a/compiler/rustc_mir_transform/src/collapse_yields.rs +++ b/compiler/rustc_mir_transform/src/collapse_yields.rs @@ -1,8 +1,11 @@ +use std::mem::discriminant; + use rustc_data_structures::fx::FxHashMap; use rustc_data_structures::graph::Successors; -use rustc_data_structures::indexmap::IndexSet; +use rustc_data_structures::indexmap::{IndexMap, IndexSet}; use rustc_middle::mir::{ - BasicBlock, Body, Local, Operand, Place, Statement, StatementKind, TerminatorKind, + AssertKind, BasicBlock, Body, Local, NonDivergingIntrinsic, Operand, Rvalue, Statement, + StatementKind, TerminatorKind, }; use rustc_middle::ty::TyCtxt; use tracing::instrument; @@ -25,13 +28,7 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { .basic_blocks .iter_enumerated() .filter_map(|(bb, bb_data)| match &bb_data.terminator.as_ref()?.kind { - TerminatorKind::Yield { value, resume, resume_arg, drop } => Some(Yield { - basic_block: bb, - value: value.clone(), - resume: *resume, - resume_arg: *resume_arg, - drop: *drop, - }), + TerminatorKind::Yield { .. } => Some(Yield { basic_block: bb }), _ => None, }) .collect::>(); @@ -44,10 +41,14 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { // No translation, so these yields aren't equivalent continue; }; + let Some(translation) = translation.check_local_types(body) else { + // The translated locals don't have the same types, so yields are not equivalent + continue; + }; // FIXME: impl doing the translation with a visitor tracing::trace!( - "Translate {:?} to {:?}: {:?}", + "Successfully translated {:?} to {:?}:\n{:?}", compare_yield.basic_block, base_yield.basic_block, translation @@ -62,20 +63,12 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { } #[derive(Debug)] -struct Yield<'tcx> { +struct Yield { basic_block: BasicBlock, - /// The value to return. - value: Operand<'tcx>, - /// Where to resume to. - resume: BasicBlock, - /// The place to store the resume argument in. - resume_arg: Place<'tcx>, - /// Cleanup to be done if the coroutine is dropped at this suspend point. - drop: Option, } -impl Yield<'_> { - fn try_find_translation(&self, other: &Yield<'_>, body: &Body<'_>) -> Option { +impl Yield { + fn try_find_translation(&self, other: &Yield, body: &Body<'_>) -> Option { let mut self_successors = self.all_successors(body); let mut other_successors = other.all_successors(body); @@ -114,15 +107,35 @@ impl Yield<'_> { } } +#[derive(Debug)] +struct LocalTranslationMap { + locals: IndexMap, +} + +impl LocalTranslationMap { + fn new() -> Self { + Self { locals: Default::default() } + } + + fn insert(mut self, l: Local, r: Local) -> Option { + if let Some(old_r) = self.locals.insert(l, r) { + if old_r != r { + return None; + } + } + Some(self) + } +} + #[derive(Debug)] struct TranslationMap { blocks: FxHashMap, - locals: FxHashMap, + locals: LocalTranslationMap, } impl TranslationMap { fn new() -> Self { - Self { blocks: Default::default(), locals: Default::default() } + Self { blocks: Default::default(), locals: LocalTranslationMap::new() } } fn try_add_translation( @@ -142,76 +155,550 @@ impl TranslationMap { return None; } + if self_data.terminator.is_some() != other_data.terminator.is_some() { + return None; + } + for (self_statement, other_statement) in self_data.statements.iter().zip(other_data.statements.iter()) { - let locals_map = - Self::statements_functionally_equivalent(self_statement, other_statement)?; - for (self_local, other_local) in locals_map { - if let Some(old_other_local) = self.locals.insert(self_local, other_local) { - if old_other_local != other_local { - // We're not equivalent since a local has changed - return None; - } - } - } + self.locals = Self::statements_functionally_equivalent( + self.locals, + self_statement, + other_statement, + )?; } - // FIXME: Consider terminators + if let (Some(l_tk), Some(r_tk)) = (&self_data.terminator, &other_data.terminator) { + self.locals = Self::terminator_kinds_functionally_equivalent( + self.locals, + &l_tk.kind, + &r_tk.kind, + )?; + } self.blocks.insert(self_bb, other_bb); Some(self) } - fn statements_functionally_equivalent( - l: &Statement<'_>, - r: &Statement<'_>, - ) -> Option> { - let mut map = FxHashMap::default(); + fn check_local_types(self, body: &Body<'_>) -> Option { + for (l, r) in &self.locals.locals { + if body.local_decls[*l].ty != body.local_decls[*r].ty { + return None; + } + } + + Some(self) + } + fn statements_functionally_equivalent<'tcx>( + mut map: LocalTranslationMap, + l: &Statement<'tcx>, + r: &Statement<'tcx>, + ) -> Option { match (&l.kind, &r.kind) { (StatementKind::Assign(l_assign), StatementKind::Assign(r_assign)) => { - map.insert(l_assign.0.local, r_assign.0.local); - unimplemented!(); + map = map.insert(l_assign.0.local, r_assign.0.local)?; + map = Self::rvalues_functionally_equivalent(map, &l_assign.1, &r_assign.1)?; } (StatementKind::FakeRead(l_fake_read), StatementKind::FakeRead(r_fake_read)) => { - unimplemented!() + if l_fake_read.0 != r_fake_read.0 { + return None; + } + map = map.insert(l_fake_read.1.local, r_fake_read.1.local)?; } ( StatementKind::SetDiscriminant { place: l_place, variant_index: l_variant_index }, StatementKind::SetDiscriminant { place: r_place, variant_index: r_variant_index }, - ) => unimplemented!(), + ) => { + if l_variant_index != r_variant_index { + return None; + } + map = map.insert(l_place.local, r_place.local)?; + } (StatementKind::StorageLive(l_local), StatementKind::StorageLive(r_local)) => { - unimplemented!() + map = map.insert(*l_local, *r_local)?; } (StatementKind::StorageDead(l_local), StatementKind::StorageDead(r_local)) => { - unimplemented!() + map = map.insert(*l_local, *r_local)?; } (StatementKind::PlaceMention(l_place), StatementKind::PlaceMention(r_place)) => { - unimplemented!() + map = map.insert(l_place.local, r_place.local)?; } ( StatementKind::AscribeUserType(l_ascribe_user_type, l_variance), StatementKind::AscribeUserType(r_ascribe_user_type, r_variance), - ) => unimplemented!(), + ) => { + if l_ascribe_user_type.1 != r_ascribe_user_type.1 { + return None; + } + if l_variance != r_variance { + return None; + } + map = map.insert(l_ascribe_user_type.0.local, r_ascribe_user_type.0.local)?; + } ( StatementKind::Coverage(l_coverage_kind), StatementKind::Coverage(r_coverage_kind), ) => { - unimplemented!() + if discriminant(l_coverage_kind) != discriminant(r_coverage_kind) { + return None; + } } ( StatementKind::Intrinsic(l_non_diverging_intrinsic), StatementKind::Intrinsic(r_non_diverging_intrinsic), - ) => unimplemented!(), + ) => { + map = Self::non_diverging_intrinsics_functionally_equivalent( + map, + l_non_diverging_intrinsic, + r_non_diverging_intrinsic, + )?; + } (StatementKind::ConstEvalCounter, StatementKind::ConstEvalCounter) => {} (StatementKind::Nop, StatementKind::Nop) => {} ( StatementKind::BackwardIncompatibleDropHint { place: l_place, reason: l_reason }, StatementKind::BackwardIncompatibleDropHint { place: r_place, reason: r_reason }, - ) => unimplemented!(), - _ => {} + ) => { + if l_reason != r_reason { + return None; + } + map = map.insert(l_place.local, r_place.local)?; + } + _ => { + // By definition not equal + return None; + } + } + + Some(map) + } + + fn rvalues_functionally_equivalent<'tcx>( + mut map: LocalTranslationMap, + l: &Rvalue<'tcx>, + r: &Rvalue<'tcx>, + ) -> Option { + match (l, r) { + (Rvalue::Use(l_operand, l_retag), Rvalue::Use(r_operand, r_retag)) => { + if l_retag != r_retag { + return None; + } + map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + } + (Rvalue::Repeat(l_operand, l_const), Rvalue::Repeat(r_operand, r_const)) => { + if l_const != r_const { + return None; + } + map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + } + ( + Rvalue::Ref(l_region, l_borrow_kind, l_place), + Rvalue::Ref(r_region, r_borrow_kind, r_place), + ) => { + if l_region != r_region { + return None; + } + if l_borrow_kind != r_borrow_kind { + return None; + } + map = map.insert(l_place.local, r_place.local)?; + } + (Rvalue::ThreadLocalRef(l_def_id), Rvalue::ThreadLocalRef(r_def_id)) => { + if l_def_id != r_def_id { + return None; + } + } + (Rvalue::RawPtr(l_raw_ptr_kind, l_place), Rvalue::RawPtr(r_raw_ptr_kind, r_place)) => { + if l_raw_ptr_kind != r_raw_ptr_kind { + return None; + } + map = map.insert(l_place.local, r_place.local)?; + } + ( + Rvalue::Cast(l_cast_kind, l_operand, l_ty), + Rvalue::Cast(r_cast_kind, r_operand, r_ty), + ) => { + if l_cast_kind != r_cast_kind { + return None; + } + if l_ty != r_ty { + return None; + } + map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + } + (Rvalue::BinaryOp(l_bin_op, l_operands), Rvalue::BinaryOp(r_bin_op, r_operands)) => { + if l_bin_op != r_bin_op { + return None; + } + map = Self::operands_functionally_equivalent(map, &l_operands.0, &r_operands.0)?; + map = Self::operands_functionally_equivalent(map, &l_operands.1, &r_operands.1)?; + } + (Rvalue::UnaryOp(l_un_op, l_operand), Rvalue::UnaryOp(r_un_op, r_operand)) => { + if l_un_op != r_un_op { + return None; + } + map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + } + (Rvalue::Discriminant(l_place), Rvalue::Discriminant(r_place)) => { + map = map.insert(l_place.local, r_place.local)?; + } + ( + Rvalue::Aggregate(l_aggregate_kind, l_index_vec), + Rvalue::Aggregate(r_aggregate_kind, r_index_vec), + ) => { + if l_aggregate_kind != r_aggregate_kind { + return None; + } + if l_index_vec != r_index_vec { + return None; + } + } + (Rvalue::CopyForDeref(l_place), Rvalue::CopyForDeref(r_place)) => { + map = map.insert(l_place.local, r_place.local)?; + } + ( + Rvalue::WrapUnsafeBinder(l_operand, l_ty), + Rvalue::WrapUnsafeBinder(r_operand, r_ty), + ) => { + if l_ty != r_ty { + return None; + } + map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + } + ( + Rvalue::Reborrow(l_ty, l_mutability, l_place), + Rvalue::Reborrow(r_ty, r_mutability, r_place), + ) => { + if l_ty != r_ty { + return None; + } + if l_mutability != r_mutability { + return None; + } + map = map.insert(l_place.local, r_place.local)?; + } + _ => { + return None; + } + } + + Some(map) + } + + fn operands_functionally_equivalent<'tcx>( + mut map: LocalTranslationMap, + l: &Operand<'tcx>, + r: &Operand<'tcx>, + ) -> Option { + match (l, r) { + (Operand::Copy(l_place), Operand::Copy(r_place)) => { + map = map.insert(l_place.local, r_place.local)?; + } + (Operand::Move(l_place), Operand::Move(r_place)) => { + map = map.insert(l_place.local, r_place.local)?; + } + (Operand::Constant(l_const_operand), Operand::Constant(r_const_operand)) => { + if l_const_operand.user_ty != r_const_operand.user_ty { + return None; + } + if l_const_operand.const_ != r_const_operand.const_ { + return None; + } + } + ( + Operand::RuntimeChecks(l_runtime_checks), + Operand::RuntimeChecks(r_runtime_checks), + ) => { + if l_runtime_checks != r_runtime_checks { + return None; + } + } + _ => { + return None; + } + } + + Some(map) + } + + fn non_diverging_intrinsics_functionally_equivalent<'tcx>( + mut map: LocalTranslationMap, + l: &NonDivergingIntrinsic<'tcx>, + r: &NonDivergingIntrinsic<'tcx>, + ) -> Option { + match (l, r) { + ( + NonDivergingIntrinsic::Assume(l_operand), + NonDivergingIntrinsic::Assume(r_operand), + ) => { + map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + } + ( + NonDivergingIntrinsic::CopyNonOverlapping(l_copy_non_overlapping), + NonDivergingIntrinsic::CopyNonOverlapping(r_copy_non_overlapping), + ) => { + map = Self::operands_functionally_equivalent( + map, + &l_copy_non_overlapping.src, + &r_copy_non_overlapping.src, + )?; + map = Self::operands_functionally_equivalent( + map, + &l_copy_non_overlapping.dst, + &r_copy_non_overlapping.dst, + )?; + map = Self::operands_functionally_equivalent( + map, + &l_copy_non_overlapping.count, + &r_copy_non_overlapping.count, + )?; + } + _ => { + return None; + } + } + + Some(map) + } + + fn terminator_kinds_functionally_equivalent<'tcx>( + mut map: LocalTranslationMap, + l: &TerminatorKind<'tcx>, + r: &TerminatorKind<'tcx>, + ) -> Option { + match (l, r) { + (TerminatorKind::Goto { target: _ }, TerminatorKind::Goto { target: _ }) => {} + ( + TerminatorKind::SwitchInt { discr: l_discr, targets: _ }, + TerminatorKind::SwitchInt { discr: r_discr, targets: _ }, + ) => { + map = Self::operands_functionally_equivalent(map, l_discr, r_discr)?; + } + (TerminatorKind::UnwindResume, TerminatorKind::UnwindResume) => {} + ( + TerminatorKind::UnwindTerminate(l_unwind_terminate_reason), + TerminatorKind::UnwindTerminate(r_unwind_terminate_reason), + ) => { + if l_unwind_terminate_reason != r_unwind_terminate_reason { + return None; + } + } + (TerminatorKind::Return, TerminatorKind::Return) => {} + (TerminatorKind::Unreachable, TerminatorKind::Unreachable) => {} + ( + TerminatorKind::Drop { + place: l_place, + target: _, + unwind: l_unwind, + replace: l_replace, + drop: _, + }, + TerminatorKind::Drop { + place: r_place, + target: _, + unwind: r_unwind, + replace: r_replace, + drop: _, + }, + ) => { + if discriminant(l_unwind) != discriminant(r_unwind) { + return None; + } + if l_replace != r_replace { + return None; + } + map = map.insert(l_place.local, r_place.local)?; + } + ( + TerminatorKind::Call { + func: l_func, + args: l_args, + destination: l_destination, + target: _, + unwind: l_unwind, + call_source: l_call_source, + fn_span: _, + }, + TerminatorKind::Call { + func: r_func, + args: r_args, + destination: r_destination, + target: _, + unwind: r_unwind, + call_source: r_call_source, + fn_span: _, + }, + ) => { + if discriminant(l_unwind) != discriminant(r_unwind) { + return None; + } + if l_call_source != r_call_source { + return None; + } + if l_args.len() != r_args.len() { + return None; + } + map = Self::operands_functionally_equivalent(map, l_func, r_func)?; + map = map.insert(l_destination.local, r_destination.local)?; + for (l_arg, r_arg) in l_args.iter().zip(r_args.iter()) { + map = Self::operands_functionally_equivalent(map, &l_arg.node, &r_arg.node)?; + } + } + ( + TerminatorKind::TailCall { func: l_func, args: l_args, fn_span: _ }, + TerminatorKind::TailCall { func: r_func, args: r_args, fn_span: _ }, + ) => { + if l_args.len() != r_args.len() { + return None; + } + map = Self::operands_functionally_equivalent(map, l_func, r_func)?; + for (l_arg, r_arg) in l_args.iter().zip(r_args.iter()) { + map = Self::operands_functionally_equivalent(map, &l_arg.node, &r_arg.node)?; + } + } + ( + TerminatorKind::Assert { + cond: l_cond, + expected: l_expected, + msg: l_msg, + target: _, + unwind: l_unwind, + }, + TerminatorKind::Assert { + cond: r_cond, + expected: r_expected, + msg: r_msg, + target: _, + unwind: r_unwind, + }, + ) => { + if l_expected != r_expected { + return None; + } + if discriminant(l_unwind) != discriminant(r_unwind) { + return None; + } + map = Self::operand_assert_kinds_functionally_equivalent(map, l_msg, r_msg)?; + map = Self::operands_functionally_equivalent(map, l_cond, r_cond)?; + } + ( + TerminatorKind::Yield { + value: l_value, + resume: _, + resume_arg: l_resume_arg, + drop: _, + }, + TerminatorKind::Yield { + value: r_value, + resume: _, + resume_arg: r_resume_arg, + drop: _, + }, + ) => { + map = Self::operands_functionally_equivalent(map, l_value, r_value)?; + map = map.insert(l_resume_arg.local, r_resume_arg.local)?; + } + (TerminatorKind::CoroutineDrop, TerminatorKind::CoroutineDrop) => {} + ( + TerminatorKind::FalseEdge { real_target: _, imaginary_target: _ }, + TerminatorKind::FalseEdge { real_target: _, imaginary_target: _ }, + ) => {} + ( + TerminatorKind::FalseUnwind { real_target: _, unwind: l_unwind }, + TerminatorKind::FalseUnwind { real_target: _, unwind: r_unwind }, + ) => { + if discriminant(l_unwind) != discriminant(r_unwind) { + return None; + } + } + (TerminatorKind::InlineAsm { .. }, TerminatorKind::InlineAsm { .. }) => { + // Let's not risk messing with asm... + return None; + } + _ => { + return None; + } + } + Some(map) + } + + fn operand_assert_kinds_functionally_equivalent<'tcx>( + mut map: LocalTranslationMap, + l: &AssertKind>, + r: &AssertKind>, + ) -> Option { + match (l, r) { + ( + AssertKind::BoundsCheck { len: l_len, index: l_index }, + AssertKind::BoundsCheck { len: r_len, index: r_index }, + ) => { + map = Self::operands_functionally_equivalent(map, l_len, r_len)?; + map = Self::operands_functionally_equivalent(map, l_index, r_index)?; + } + ( + AssertKind::Overflow(l_bin_op, l_op_0, l_op_1), + AssertKind::Overflow(r_bin_op, r_op_0, r_op_1), + ) => { + if l_bin_op != r_bin_op { + return None; + } + map = Self::operands_functionally_equivalent(map, l_op_0, r_op_0)?; + map = Self::operands_functionally_equivalent(map, l_op_1, r_op_1)?; + } + (AssertKind::OverflowNeg(l_op), AssertKind::OverflowNeg(r_op)) => { + map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + } + (AssertKind::DivisionByZero(l_op), AssertKind::DivisionByZero(r_op)) => { + map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + } + (AssertKind::RemainderByZero(l_op), AssertKind::RemainderByZero(r_op)) => { + map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + } + ( + AssertKind::ResumedAfterReturn(l_coroutine_kind), + AssertKind::ResumedAfterReturn(r_coroutine_kind), + ) => { + if l_coroutine_kind != r_coroutine_kind { + return None; + } + } + ( + AssertKind::ResumedAfterPanic(l_coroutine_kind), + AssertKind::ResumedAfterPanic(r_coroutine_kind), + ) => { + if l_coroutine_kind != r_coroutine_kind { + return None; + } + } + ( + AssertKind::ResumedAfterDrop(l_coroutine_kind), + AssertKind::ResumedAfterDrop(r_coroutine_kind), + ) => { + if l_coroutine_kind != r_coroutine_kind { + return None; + } + } + ( + AssertKind::MisalignedPointerDereference { required: l_required, found: l_found }, + AssertKind::MisalignedPointerDereference { required: r_required, found: r_found }, + ) => { + map = Self::operands_functionally_equivalent(map, l_required, r_required)?; + map = Self::operands_functionally_equivalent(map, l_found, r_found)?; + } + (AssertKind::NullPointerDereference, AssertKind::NullPointerDereference) => {} + (AssertKind::NullReferenceConstructed, AssertKind::NullReferenceConstructed) => {} + ( + AssertKind::InvalidEnumConstruction(l_op), + AssertKind::InvalidEnumConstruction(r_op), + ) => { + map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + } + _ => { + return None; + } } Some(map) From 22942775ad739a36e608833205167c8ef0b5cfad Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Tue, 25 Aug 2026 13:40:58 +0200 Subject: [PATCH 04/12] Redirect entry points --- .../src/collapse_yields.rs | 102 ++++++++++++++++-- 1 file changed, 93 insertions(+), 9 deletions(-) diff --git a/compiler/rustc_mir_transform/src/collapse_yields.rs b/compiler/rustc_mir_transform/src/collapse_yields.rs index f900bd2ba5830..269e72270f0ed 100644 --- a/compiler/rustc_mir_transform/src/collapse_yields.rs +++ b/compiler/rustc_mir_transform/src/collapse_yields.rs @@ -1,17 +1,22 @@ +use std::borrow::Cow; use std::mem::discriminant; use rustc_data_structures::fx::FxHashMap; use rustc_data_structures::graph::Successors; use rustc_data_structures::indexmap::{IndexMap, IndexSet}; use rustc_middle::mir::{ - AssertKind, BasicBlock, Body, Local, NonDivergingIntrinsic, Operand, Rvalue, Statement, - StatementKind, TerminatorKind, + AssertKind, BasicBlock, Body, Local, MirDumper, NonDivergingIntrinsic, OUTERMOST_SOURCE_SCOPE, + Operand, Place, Rvalue, SourceInfo, Statement, StatementKind, TerminatorKind, WithRetag, }; use rustc_middle::ty::TyCtxt; +use rustc_mir_dataflow::Analysis; +use rustc_mir_dataflow::impls::{MaybeStorageLive, always_storage_live_locals}; +use rustc_span::DUMMY_SP; use tracing::instrument; use crate::MirPass; use crate::pass_manager::PassPolicy; +use crate::simplify::remove_dead_blocks; pub(super) struct CollapseIdenticalYields; @@ -22,22 +27,33 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { return; } + if let Some(dumper) = MirDumper::new(tcx, "collapse_yields_before", body) { + dumper.dump_mir(body); + } + tracing::debug!("running pass for {}", tcx.def_path_str(body.source.def_id())); let mut yields = body .basic_blocks .iter_enumerated() - .filter_map(|(bb, bb_data)| match &bb_data.terminator.as_ref()?.kind { - TerminatorKind::Yield { .. } => Some(Yield { basic_block: bb }), - _ => None, + .filter_map(|(bb, bb_data)| { + if let TerminatorKind::Yield { .. } = &bb_data.terminator.as_ref()?.kind { + Some(Yield { basic_block: bb }) + } else { + None + } }) .collect::>(); + // Sort so we always translate from high bbs to low bbs + yields.sort_unstable_by(|y1, y2| y1.basic_block.cmp(&y2.basic_block).reverse()); while let Some(base_yield) = yields.pop() { + tracing::trace!("Comparing base yield {:?} to others:", base_yield.basic_block); for i in (0..yields.len()).rev() { let compare_yield = &yields[i]; - let Some(translation) = base_yield.try_find_translation(compare_yield, body) else { + let Some(translation) = compare_yield.try_find_translation(&base_yield, body) + else { // No translation, so these yields aren't equivalent continue; }; @@ -46,13 +62,17 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { continue; }; - // FIXME: impl doing the translation with a visitor tracing::trace!( - "Successfully translated {:?} to {:?}:\n{:?}", + "Successfully translated yield {:?} to yield {:?}:\n{:?}", compare_yield.basic_block, base_yield.basic_block, translation ); + + yields.remove(i); + + translation.redirect_entry_points(tcx, body); + remove_dead_blocks(body); } } } @@ -129,7 +149,7 @@ impl LocalTranslationMap { #[derive(Debug)] struct TranslationMap { - blocks: FxHashMap, + blocks: IndexMap, locals: LocalTranslationMap, } @@ -192,6 +212,70 @@ impl TranslationMap { Some(self) } + fn redirect_entry_points<'tcx>(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) { + let body_predecessors = body.basic_blocks.predecessors().clone(); + let always_live_locals = always_storage_live_locals(body); + let mut results = MaybeStorageLive::new(Cow::Borrowed(&always_live_locals)) + .iterate_to_fixpoint(tcx, body, Some("callapse_yields")) + .into_results_cursor(body); + + let from_live_locals = self + .blocks + .keys() + .map(|from| { + results.seek_to_block_start(*from); + (*from, results.get().clone()) + }) + .collect::>(); + + for (from, to) in self.blocks.iter() { + let predecessors = &body_predecessors[*from]; + // These predecessors go to the 'from' block, but need to go to the 'to' block + // We also need to translate the live locals + // So we insert the translations into the predecessor and change the successor to 'to' + + for predecessor in predecessors { + let predecessor_data = &mut body.basic_blocks_mut()[*predecessor]; + + let live_locals = &from_live_locals[from]; + for from_local in live_locals.iter() { + if let Some(to_local) = self.locals.locals.get(&from_local) { + if from_local == *to_local { + continue; + } + + if !always_live_locals.contains(*to_local) { + predecessor_data.statements.push(Statement::new( + SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, + StatementKind::StorageLive(*to_local), + )); + } + predecessor_data.statements.push(Statement::new( + SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, + StatementKind::Assign(Box::new(( + Place::from(*to_local), + Rvalue::Use(Operand::Move(Place::from(from_local)), WithRetag::Yes), + ))), + )); + if !always_live_locals.contains(from_local) { + predecessor_data.statements.push(Statement::new( + SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, + StatementKind::StorageDead(from_local), + )); + } + } + } + + // Change the terminator to go to the 'to' block + predecessor_data.terminator_mut().successors_mut(|successor| { + if *successor == *from { + *successor = *to + } + }); + } + } + } + fn statements_functionally_equivalent<'tcx>( mut map: LocalTranslationMap, l: &Statement<'tcx>, From 5ed1363f3f3e14b572057dae05d25de90d28c422 Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Fri, 28 Aug 2026 13:52:39 +0200 Subject: [PATCH 05/12] Improve algorithm so it's more obvious and move dead block elimation to last --- .../src/collapse_yields.rs | 62 +++++++++++-------- 1 file changed, 36 insertions(+), 26 deletions(-) diff --git a/compiler/rustc_mir_transform/src/collapse_yields.rs b/compiler/rustc_mir_transform/src/collapse_yields.rs index 269e72270f0ed..b35915fb61ca8 100644 --- a/compiler/rustc_mir_transform/src/collapse_yields.rs +++ b/compiler/rustc_mir_transform/src/collapse_yields.rs @@ -1,7 +1,8 @@ use std::borrow::Cow; use std::mem::discriminant; -use rustc_data_structures::fx::FxHashMap; +use itertools::Itertools; +use rustc_data_structures::fx::{FxHashMap, FxHashSet}; use rustc_data_structures::graph::Successors; use rustc_data_structures::indexmap::{IndexMap, IndexSet}; use rustc_middle::mir::{ @@ -31,7 +32,7 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { dumper.dump_mir(body); } - tracing::debug!("running pass for {}", tcx.def_path_str(body.source.def_id())); + tracing::debug!("running pass for {}", tcx.def_path_debug_str(body.source.def_id())); let mut yields = body .basic_blocks @@ -47,34 +48,43 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { // Sort so we always translate from high bbs to low bbs yields.sort_unstable_by(|y1, y2| y1.basic_block.cmp(&y2.basic_block).reverse()); - while let Some(base_yield) = yields.pop() { - tracing::trace!("Comparing base yield {:?} to others:", base_yield.basic_block); - for i in (0..yields.len()).rev() { - let compare_yield = &yields[i]; + let mut collapsed_yields = FxHashSet::default(); - let Some(translation) = compare_yield.try_find_translation(&base_yield, body) - else { - // No translation, so these yields aren't equivalent - continue; - }; - let Some(translation) = translation.check_local_types(body) else { - // The translated locals don't have the same types, so yields are not equivalent - continue; - }; + for compare_yields in yields.iter().combinations(2) { + let base_yield = compare_yields[0]; + let compare_yield = compare_yields[1]; - tracing::trace!( - "Successfully translated yield {:?} to yield {:?}:\n{:?}", - compare_yield.basic_block, - base_yield.basic_block, - translation - ); + if collapsed_yields.contains(&compare_yield) { + continue; + } - yields.remove(i); + tracing::trace!( + "Comparing yield {:?} to yield {:?}", + base_yield.basic_block, + compare_yield.basic_block + ); - translation.redirect_entry_points(tcx, body); - remove_dead_blocks(body); - } + let Some(translation) = compare_yield.try_find_translation(&base_yield, body) else { + // No translation, so these yields aren't equivalent + continue; + }; + let Some(translation) = translation.check_local_types(body) else { + // The translated locals don't have the same types, so yields are not equivalent + continue; + }; + + tracing::trace!( + "Successfully translated yield {:?} to yield {:?}:\n{:?}", + compare_yield.basic_block, + base_yield.basic_block, + translation + ); + + collapsed_yields.insert(compare_yield); + + translation.redirect_entry_points(tcx, body); } + remove_dead_blocks(body); } fn policy(&self, _sess: &rustc_session::Session) -> PassPolicy { @@ -82,7 +92,7 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { } } -#[derive(Debug)] +#[derive(Debug, Hash, PartialEq, Eq)] struct Yield { basic_block: BasicBlock, } From cbaa9eae4eac326e82af21c5f5ad88dbfee4bca4 Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Fri, 28 Aug 2026 16:12:10 +0200 Subject: [PATCH 06/12] Create separate blocks for locals translation so it doesn't influence paths that shouldn't be translated --- .../src/collapse_yields.rs | 38 +++++++++++++------ 1 file changed, 26 insertions(+), 12 deletions(-) diff --git a/compiler/rustc_mir_transform/src/collapse_yields.rs b/compiler/rustc_mir_transform/src/collapse_yields.rs index b35915fb61ca8..465c3d10ad679 100644 --- a/compiler/rustc_mir_transform/src/collapse_yields.rs +++ b/compiler/rustc_mir_transform/src/collapse_yields.rs @@ -6,8 +6,9 @@ use rustc_data_structures::fx::{FxHashMap, FxHashSet}; use rustc_data_structures::graph::Successors; use rustc_data_structures::indexmap::{IndexMap, IndexSet}; use rustc_middle::mir::{ - AssertKind, BasicBlock, Body, Local, MirDumper, NonDivergingIntrinsic, OUTERMOST_SOURCE_SCOPE, - Operand, Place, Rvalue, SourceInfo, Statement, StatementKind, TerminatorKind, WithRetag, + AssertKind, BasicBlock, BasicBlockData, Body, Local, MirDumper, NonDivergingIntrinsic, + OUTERMOST_SOURCE_SCOPE, Operand, Place, Rvalue, SourceInfo, Statement, StatementKind, + Terminator, TerminatorKind, WithRetag, }; use rustc_middle::ty::TyCtxt; use rustc_mir_dataflow::Analysis; @@ -245,8 +246,28 @@ impl TranslationMap { // So we insert the translations into the predecessor and change the successor to 'to' for predecessor in predecessors { + let inbetween = body.basic_blocks_mut().push(BasicBlockData::new( + Some(Terminator { + source_info: SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, + kind: TerminatorKind::Goto { target: *to }, + attributes: Default::default(), + }), + false, + )); + let predecessor_data = &mut body.basic_blocks_mut()[*predecessor]; + // Change the terminator to go to the 'to' block + predecessor_data.terminator_mut().successors_mut(|successor| { + if *successor == *from { + *successor = inbetween; + } + }); + let is_cleanup = body.basic_blocks_mut()[*from].is_cleanup; + + let inbetween_data = &mut body.basic_blocks_mut()[inbetween]; + inbetween_data.is_cleanup = is_cleanup; + let live_locals = &from_live_locals[from]; for from_local in live_locals.iter() { if let Some(to_local) = self.locals.locals.get(&from_local) { @@ -255,12 +276,12 @@ impl TranslationMap { } if !always_live_locals.contains(*to_local) { - predecessor_data.statements.push(Statement::new( + inbetween_data.statements.push(Statement::new( SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, StatementKind::StorageLive(*to_local), )); } - predecessor_data.statements.push(Statement::new( + inbetween_data.statements.push(Statement::new( SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, StatementKind::Assign(Box::new(( Place::from(*to_local), @@ -268,20 +289,13 @@ impl TranslationMap { ))), )); if !always_live_locals.contains(from_local) { - predecessor_data.statements.push(Statement::new( + inbetween_data.statements.push(Statement::new( SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, StatementKind::StorageDead(from_local), )); } } } - - // Change the terminator to go to the 'to' block - predecessor_data.terminator_mut().successors_mut(|successor| { - if *successor == *from { - *successor = *to - } - }); } } } From 15562404a91db9c0ff209a95641a9d63bd3e6d13 Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Wed, 2 Sep 2026 11:30:53 +0200 Subject: [PATCH 07/12] Add a test --- tests/mir-opt/coroutine/async_collapse.rs | 98 +++++++++++++++++++++++ 1 file changed, 98 insertions(+) create mode 100644 tests/mir-opt/coroutine/async_collapse.rs diff --git a/tests/mir-opt/coroutine/async_collapse.rs b/tests/mir-opt/coroutine/async_collapse.rs new file mode 100644 index 0000000000000..b61259186ad99 --- /dev/null +++ b/tests/mir-opt/coroutine/async_collapse.rs @@ -0,0 +1,98 @@ +// This test makes sure that the collapse_yields MIR pass eliminates +// identical yields so the statemachine will be smaller + +//@ edition:2018 + +#![crate_type = "lib"] + +async fn a(arg: i32) -> i32 { + arg +} + +pub async fn b(val: bool) { + // Check there's only one suspension point, not two + + // CHECK-LABEL: fn b::{closure#0}( + // CHECK-COUNT-1: Suspend + // CHECK-NOT: Suspend + + if val { + a(1).await; + } else { + a(-1).await; + } +} + +pub async fn c(val: bool) { + // Check there's only three suspension points, not four + + // CHECK-LABEL: fn c::{closure#0}( + // CHECK-COUNT-3: Suspend + // CHECK-NOT: Suspend + + b(val).await; + + if val { + a(1).await; + } else { + a(-1).await; + } + + b(val).await; +} + +enum ManyOptions { + A, + B, + C, + D, +} + +pub async fn d(val: ManyOptions) { + // Check there's only one suspension points, not four + + // CHECK-LABEL: fn d::{closure#0}( + // CHECK-COUNT-1: Suspend + // CHECK-NOT: Suspend + + match val { + ManyOptions::A => { + a(0).await; + } + ManyOptions::B => { + println!("It's B!"); // Doing diverging work before the await shouldn't matter + a(1).await; + } + ManyOptions::C => { + a(2).await; + } + ManyOptions::D => { + a(3).await; + } + } +} +pub async fn e(val: ManyOptions) { + // Check there's only two suspension points, not four + + // CHECK-LABEL: fn e::{closure#0}( + // CHECK-COUNT-2: Suspend + // CHECK-NOT: Suspend + + match val { + ManyOptions::A => { + a(0).await; + } + ManyOptions::B => { + a(1).await; + println!("It's B!"); + // Doing diverging work after the await *does* matter. + // Nonetheless, the others should still be optimized. + } + ManyOptions::C => { + a(2).await; + } + ManyOptions::D => { + a(3).await; + } + } +} From 77b2ab576a26f78c5197d14f32d7ea8c69db8098 Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Wed, 2 Sep 2026 13:56:19 +0200 Subject: [PATCH 08/12] Rename from collapse to merge --- compiler/rustc_mir_transform/src/lib.rs | 4 ++-- .../src/{collapse_yields.rs => merge_yields.rs} | 14 +++++--------- .../{async_collapse.rs => async_merge.rs} | 2 +- 3 files changed, 8 insertions(+), 12 deletions(-) rename compiler/rustc_mir_transform/src/{collapse_yields.rs => merge_yields.rs} (98%) rename tests/mir-opt/coroutine/{async_collapse.rs => async_merge.rs} (96%) diff --git a/compiler/rustc_mir_transform/src/lib.rs b/compiler/rustc_mir_transform/src/lib.rs index bc20cfb172752..3553bd7cec31f 100644 --- a/compiler/rustc_mir_transform/src/lib.rs +++ b/compiler/rustc_mir_transform/src/lib.rs @@ -136,7 +136,7 @@ declare_passes! { pub mod cleanup_post_borrowck : CleanupPostBorrowck; mod copy_prop : CopyProp; - mod collapse_yields : CollapseIdenticalYields; + mod merge_yields : MergeYields; mod coroutine : StateTransform; mod coverage : InstrumentCoverage; mod ctfe_limit : CtfeLimit; @@ -662,7 +662,7 @@ fn run_runtime_lowering_passes<'tcx>(tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) { &add_moves_for_packed_drops::AddMovesForPackedDrops, &erase_deref_temps::EraseDerefTemps, &elaborate_box_derefs::ElaborateBoxDerefs, - &collapse_yields::CollapseIdenticalYields, + &merge_yields::MergeYields, &coroutine::StateTransform, &Lint(known_panics_lint::KnownPanicsLint), ]; diff --git a/compiler/rustc_mir_transform/src/collapse_yields.rs b/compiler/rustc_mir_transform/src/merge_yields.rs similarity index 98% rename from compiler/rustc_mir_transform/src/collapse_yields.rs rename to compiler/rustc_mir_transform/src/merge_yields.rs index 465c3d10ad679..7769e4874ef73 100644 --- a/compiler/rustc_mir_transform/src/collapse_yields.rs +++ b/compiler/rustc_mir_transform/src/merge_yields.rs @@ -20,19 +20,15 @@ use crate::MirPass; use crate::pass_manager::PassPolicy; use crate::simplify::remove_dead_blocks; -pub(super) struct CollapseIdenticalYields; +pub(super) struct MergeYields; -impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { +impl<'tcx> MirPass<'tcx> for MergeYields { #[instrument(level = "debug", skip(self, tcx, body), ret)] fn run_pass(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) { if body.coroutine_kind().is_none() { return; } - if let Some(dumper) = MirDumper::new(tcx, "collapse_yields_before", body) { - dumper.dump_mir(body); - } - tracing::debug!("running pass for {}", tcx.def_path_debug_str(body.source.def_id())); let mut yields = body @@ -49,13 +45,13 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { // Sort so we always translate from high bbs to low bbs yields.sort_unstable_by(|y1, y2| y1.basic_block.cmp(&y2.basic_block).reverse()); - let mut collapsed_yields = FxHashSet::default(); + let mut merged_yields = FxHashSet::default(); for compare_yields in yields.iter().combinations(2) { let base_yield = compare_yields[0]; let compare_yield = compare_yields[1]; - if collapsed_yields.contains(&compare_yield) { + if merged_yields.contains(&compare_yield) { continue; } @@ -81,7 +77,7 @@ impl<'tcx> MirPass<'tcx> for CollapseIdenticalYields { translation ); - collapsed_yields.insert(compare_yield); + merged_yields.insert(compare_yield); translation.redirect_entry_points(tcx, body); } diff --git a/tests/mir-opt/coroutine/async_collapse.rs b/tests/mir-opt/coroutine/async_merge.rs similarity index 96% rename from tests/mir-opt/coroutine/async_collapse.rs rename to tests/mir-opt/coroutine/async_merge.rs index b61259186ad99..85282ba63f09f 100644 --- a/tests/mir-opt/coroutine/async_collapse.rs +++ b/tests/mir-opt/coroutine/async_merge.rs @@ -1,4 +1,4 @@ -// This test makes sure that the collapse_yields MIR pass eliminates +// This test makes sure that the MergeYields MIR pass eliminates // identical yields so the statemachine will be smaller //@ edition:2018 From b932817aca4320bc7f30524637e14602a09c9d69 Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Wed, 2 Sep 2026 15:35:28 +0200 Subject: [PATCH 09/12] Add docs, do a pass over names and move to results rather than using options weirdly --- .../rustc_mir_transform/src/merge_yields.rs | 390 +++++++++++------- 1 file changed, 240 insertions(+), 150 deletions(-) diff --git a/compiler/rustc_mir_transform/src/merge_yields.rs b/compiler/rustc_mir_transform/src/merge_yields.rs index 7769e4874ef73..a23cb13600f1e 100644 --- a/compiler/rustc_mir_transform/src/merge_yields.rs +++ b/compiler/rustc_mir_transform/src/merge_yields.rs @@ -1,3 +1,86 @@ +//! Implementation of the [MergeYields] pass. +//! +//! The [super::coroutine::StateTransform] will take all yields in a body and make them into a +//! suspension point in the generated statemachine. This pass needs to run before that happens. +//! +//! The idea of this pass it to find all yield points and compare them. +//! If they are *functionally* identical, we can merge them. When that happens, the resulting +//! state machine will have fewer states, which is good for binary size and (probably) +//! performance. +//! +//! Take this code as example: +//! ```ignore (demonstration only) +//! if _0 { +//! a(_1).await; +//! } else { +//! a(_2).await; +//! } +//! ``` +//! +//! This gets turned into this (simplified) MIR shape before the state transform: +//! ```txt +//! ┌─────────┐ +//! ┌────────┼switch _0┼────────┐ +//! │ └─────────┘ │ +//! │ │ +//! ┌──────────▼──────────┐ ┌──────────▼───────────┐ +//! │create future A as _3│ │create future A as _4 │ +//! │with local _1 │ │with local _2 │ +//! └──────────┬──────────┘ └──────────┬───────────┘ +//! ┌───▼───┐ ┌───▼───┐ +//! ┌───────►poll _3│ ┌───────►poll _4│ +//! │ └──┬─┬──┘ │ └──┬─┬──┘ +//! │ ┌─────┐ │ │ ┌─────┐ │ ┌─────┐ │ │ ┌─────┐ +//! └──┼yield◄─┘ └─►ready│ └──┼yield◄─┘ └─►ready│ +//! └─────┘ └──┬──┘ └─────┘ └──┬──┘ +//! └───────┬───────────────────┘ +//! │ +//! ┌───────▼────────┐ +//! │next thing to do│ +//! └────────────────┘ +//! ``` +//! +//! To compare the yields, we take all the successors of each yield and walk them. +//! For each block we check if: +//! - The terminators are the same +//! - The statements are the same +//! +//! The only thing they're allowed to differ in are the indices of the locals and successor blocks. +//! Everything else must be the same and the locals must also consistently map onto each other. +//! So if we figured out during the walk that _40 maps to _50, +//! but then later we find out _40 now maps to _60, the walk is stopped and the yields are not +//! deemed identical. +//! +//! If this all went ok, we only need to check the translations of the locals and make sure +//! they're all of the same type. +//! +//! When we find identical yields, we need to merge them. This is done by taking all entry points +//! into the second yield and rewriting them to point into the first yield. +//! +//! Ultimately our example will look like this: +//! ```txt +//! ┌─────────┐ +//! ┌────────┼switch _0┼────────┐ +//! │ └─────────┘ │ +//! │ │ +//! ┌──────────▼──────────┐ ┌──────────▼───────────┐ +//! │create future A as _3│ │create future A as _4 │ +//! │with local _1 │ │with local _2 │ +//! └──────────┬──────────┘ └──────────┬───────────┘ +//! ┌───▼───┐ ┌─────▼──────┐ +//! ┌───────►poll _3◄─────────────────┼_3 = move _4│ +//! │ └──┬─┬──┘ └────────────┘ +//! │ ┌─────┐ │ │ ┌─────┐ +//! └──┼yield◄─┘ └─►ready│ +//! └─────┘ └──┬──┘ +//! └───────┐ +//! │ +//! ┌───────▼────────┐ +//! │next thing to do│ +//! └────────────────┘ +//! ``` +//! + use std::borrow::Cow; use std::mem::discriminant; @@ -6,7 +89,7 @@ use rustc_data_structures::fx::{FxHashMap, FxHashSet}; use rustc_data_structures::graph::Successors; use rustc_data_structures::indexmap::{IndexMap, IndexSet}; use rustc_middle::mir::{ - AssertKind, BasicBlock, BasicBlockData, Body, Local, MirDumper, NonDivergingIntrinsic, + AssertKind, BasicBlock, BasicBlockData, Body, Local, NonDivergingIntrinsic, OUTERMOST_SOURCE_SCOPE, Operand, Place, Rvalue, SourceInfo, Statement, StatementKind, Terminator, TerminatorKind, WithRetag, }; @@ -51,7 +134,8 @@ impl<'tcx> MirPass<'tcx> for MergeYields { let base_yield = compare_yields[0]; let compare_yield = compare_yields[1]; - if merged_yields.contains(&compare_yield) { + if merged_yields.contains(base_yield) || merged_yields.contains(compare_yield) { + // Skip comparison if we've already merged this yield continue; } @@ -61,11 +145,11 @@ impl<'tcx> MirPass<'tcx> for MergeYields { compare_yield.basic_block ); - let Some(translation) = compare_yield.try_find_translation(&base_yield, body) else { + let Ok(translation) = compare_yield.try_find_translation(base_yield, body) else { // No translation, so these yields aren't equivalent continue; }; - let Some(translation) = translation.check_local_types(body) else { + if translation.check_local_types(body).is_err() { // The translated locals don't have the same types, so yields are not equivalent continue; }; @@ -77,10 +161,14 @@ impl<'tcx> MirPass<'tcx> for MergeYields { translation ); - merged_yields.insert(compare_yield); - + // The compare yield must be removed and every entry point is redirected to the equivalent entry into base + // The removal itself is done at the end of the pass translation.redirect_entry_points(tcx, body); + + // Avoid comparing this yield again later since it has been removed + merged_yields.insert(compare_yield); } + remove_dead_blocks(body); } @@ -95,7 +183,11 @@ struct Yield { } impl Yield { - fn try_find_translation(&self, other: &Yield, body: &Body<'_>) -> Option { + fn try_find_translation( + &self, + other: &Yield, + body: &Body<'_>, + ) -> Result { let mut self_successors = self.all_successors(body); let mut other_successors = other.all_successors(body); @@ -105,15 +197,15 @@ impl Yield { if self_successor == other_successor { continue; } - map = map.try_add_translation(self_successor, other_successor, body)?; + map.try_add_translation(self_successor, other_successor, body)?; } if self_successors.next().is_some() || other_successors.next().is_some() { // Can't be the same if they're not the same length - return None; + return Err(TranslationError); } - Some(map) + Ok(map) } fn all_successors(&self, body: &Body<'_>) -> impl Iterator { @@ -144,13 +236,17 @@ impl LocalTranslationMap { Self { locals: Default::default() } } - fn insert(mut self, l: Local, r: Local) -> Option { + /// Insert locals for translation. + /// + /// If `l` already exists but has a different `r` as value already, + /// then None is returned. This signifies the translation has failed. + fn insert(&mut self, l: Local, r: Local) -> Result<(), TranslationError> { if let Some(old_r) = self.locals.insert(l, r) { if old_r != r { - return None; + return Err(TranslationError); } } - Some(self) + Ok(()) } } @@ -166,57 +262,49 @@ impl TranslationMap { } fn try_add_translation( - mut self, + &mut self, self_bb: BasicBlock, other_bb: BasicBlock, body: &Body<'_>, - ) -> Option { + ) -> Result<(), TranslationError> { let self_data = &body.basic_blocks[self_bb]; let other_data = &body.basic_blocks[other_bb]; if self_data.is_cleanup != other_data.is_cleanup { - return None; + return Err(TranslationError); } if self_data.statements.len() != other_data.statements.len() { - return None; + return Err(TranslationError); } if self_data.terminator.is_some() != other_data.terminator.is_some() { - return None; + return Err(TranslationError); } for (self_statement, other_statement) in self_data.statements.iter().zip(other_data.statements.iter()) { - self.locals = Self::statements_functionally_equivalent( - self.locals, - self_statement, - other_statement, - )?; + Self::check_statements(&mut self.locals, self_statement, other_statement)?; } if let (Some(l_tk), Some(r_tk)) = (&self_data.terminator, &other_data.terminator) { - self.locals = Self::terminator_kinds_functionally_equivalent( - self.locals, - &l_tk.kind, - &r_tk.kind, - )?; + Self::check_terminator_kinds(&mut self.locals, &l_tk.kind, &r_tk.kind)?; } self.blocks.insert(self_bb, other_bb); - Some(self) + Ok(()) } - fn check_local_types(self, body: &Body<'_>) -> Option { + fn check_local_types(&self, body: &Body<'_>) -> Result<(), TranslationError> { for (l, r) in &self.locals.locals { if body.local_decls[*l].ty != body.local_decls[*r].ty { - return None; + return Err(TranslationError); } } - Some(self) + Ok(()) } fn redirect_entry_points<'tcx>(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) { @@ -296,65 +384,65 @@ impl TranslationMap { } } - fn statements_functionally_equivalent<'tcx>( - mut map: LocalTranslationMap, + fn check_statements<'tcx>( + map: &mut LocalTranslationMap, l: &Statement<'tcx>, r: &Statement<'tcx>, - ) -> Option { + ) -> Result<(), TranslationError> { match (&l.kind, &r.kind) { (StatementKind::Assign(l_assign), StatementKind::Assign(r_assign)) => { - map = map.insert(l_assign.0.local, r_assign.0.local)?; - map = Self::rvalues_functionally_equivalent(map, &l_assign.1, &r_assign.1)?; + map.insert(l_assign.0.local, r_assign.0.local)?; + Self::check_rvalues(map, &l_assign.1, &r_assign.1)?; } (StatementKind::FakeRead(l_fake_read), StatementKind::FakeRead(r_fake_read)) => { if l_fake_read.0 != r_fake_read.0 { - return None; + return Err(TranslationError); } - map = map.insert(l_fake_read.1.local, r_fake_read.1.local)?; + map.insert(l_fake_read.1.local, r_fake_read.1.local)?; } ( StatementKind::SetDiscriminant { place: l_place, variant_index: l_variant_index }, StatementKind::SetDiscriminant { place: r_place, variant_index: r_variant_index }, ) => { if l_variant_index != r_variant_index { - return None; + return Err(TranslationError); } - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } (StatementKind::StorageLive(l_local), StatementKind::StorageLive(r_local)) => { - map = map.insert(*l_local, *r_local)?; + map.insert(*l_local, *r_local)?; } (StatementKind::StorageDead(l_local), StatementKind::StorageDead(r_local)) => { - map = map.insert(*l_local, *r_local)?; + map.insert(*l_local, *r_local)?; } (StatementKind::PlaceMention(l_place), StatementKind::PlaceMention(r_place)) => { - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } ( StatementKind::AscribeUserType(l_ascribe_user_type, l_variance), StatementKind::AscribeUserType(r_ascribe_user_type, r_variance), ) => { if l_ascribe_user_type.1 != r_ascribe_user_type.1 { - return None; + return Err(TranslationError); } if l_variance != r_variance { - return None; + return Err(TranslationError); } - map = map.insert(l_ascribe_user_type.0.local, r_ascribe_user_type.0.local)?; + map.insert(l_ascribe_user_type.0.local, r_ascribe_user_type.0.local)?; } ( StatementKind::Coverage(l_coverage_kind), StatementKind::Coverage(r_coverage_kind), ) => { if discriminant(l_coverage_kind) != discriminant(r_coverage_kind) { - return None; + return Err(TranslationError); } } ( StatementKind::Intrinsic(l_non_diverging_intrinsic), StatementKind::Intrinsic(r_non_diverging_intrinsic), ) => { - map = Self::non_diverging_intrinsics_functionally_equivalent( + Self::check_non_diverging_intrinsics( map, l_non_diverging_intrinsic, r_non_diverging_intrinsic, @@ -367,149 +455,149 @@ impl TranslationMap { StatementKind::BackwardIncompatibleDropHint { place: r_place, reason: r_reason }, ) => { if l_reason != r_reason { - return None; + return Err(TranslationError); } - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } _ => { // By definition not equal - return None; + return Err(TranslationError); } } - Some(map) + Ok(()) } - fn rvalues_functionally_equivalent<'tcx>( - mut map: LocalTranslationMap, + fn check_rvalues<'tcx>( + map: &mut LocalTranslationMap, l: &Rvalue<'tcx>, r: &Rvalue<'tcx>, - ) -> Option { + ) -> Result<(), TranslationError> { match (l, r) { (Rvalue::Use(l_operand, l_retag), Rvalue::Use(r_operand, r_retag)) => { if l_retag != r_retag { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + Self::check_operands(map, l_operand, r_operand)?; } (Rvalue::Repeat(l_operand, l_const), Rvalue::Repeat(r_operand, r_const)) => { if l_const != r_const { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + Self::check_operands(map, l_operand, r_operand)?; } ( Rvalue::Ref(l_region, l_borrow_kind, l_place), Rvalue::Ref(r_region, r_borrow_kind, r_place), ) => { if l_region != r_region { - return None; + return Err(TranslationError); } if l_borrow_kind != r_borrow_kind { - return None; + return Err(TranslationError); } - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } (Rvalue::ThreadLocalRef(l_def_id), Rvalue::ThreadLocalRef(r_def_id)) => { if l_def_id != r_def_id { - return None; + return Err(TranslationError); } } (Rvalue::RawPtr(l_raw_ptr_kind, l_place), Rvalue::RawPtr(r_raw_ptr_kind, r_place)) => { if l_raw_ptr_kind != r_raw_ptr_kind { - return None; + return Err(TranslationError); } - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } ( Rvalue::Cast(l_cast_kind, l_operand, l_ty), Rvalue::Cast(r_cast_kind, r_operand, r_ty), ) => { if l_cast_kind != r_cast_kind { - return None; + return Err(TranslationError); } if l_ty != r_ty { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + Self::check_operands(map, l_operand, r_operand)?; } (Rvalue::BinaryOp(l_bin_op, l_operands), Rvalue::BinaryOp(r_bin_op, r_operands)) => { if l_bin_op != r_bin_op { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, &l_operands.0, &r_operands.0)?; - map = Self::operands_functionally_equivalent(map, &l_operands.1, &r_operands.1)?; + Self::check_operands(map, &l_operands.0, &r_operands.0)?; + Self::check_operands(map, &l_operands.1, &r_operands.1)?; } (Rvalue::UnaryOp(l_un_op, l_operand), Rvalue::UnaryOp(r_un_op, r_operand)) => { if l_un_op != r_un_op { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + Self::check_operands(map, l_operand, r_operand)?; } (Rvalue::Discriminant(l_place), Rvalue::Discriminant(r_place)) => { - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } ( Rvalue::Aggregate(l_aggregate_kind, l_index_vec), Rvalue::Aggregate(r_aggregate_kind, r_index_vec), ) => { if l_aggregate_kind != r_aggregate_kind { - return None; + return Err(TranslationError); } if l_index_vec != r_index_vec { - return None; + return Err(TranslationError); } } (Rvalue::CopyForDeref(l_place), Rvalue::CopyForDeref(r_place)) => { - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } ( Rvalue::WrapUnsafeBinder(l_operand, l_ty), Rvalue::WrapUnsafeBinder(r_operand, r_ty), ) => { if l_ty != r_ty { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + Self::check_operands(map, l_operand, r_operand)?; } ( Rvalue::Reborrow(l_ty, l_mutability, l_place), Rvalue::Reborrow(r_ty, r_mutability, r_place), ) => { if l_ty != r_ty { - return None; + return Err(TranslationError); } if l_mutability != r_mutability { - return None; + return Err(TranslationError); } - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } _ => { - return None; + return Err(TranslationError); } } - Some(map) + Ok(()) } - fn operands_functionally_equivalent<'tcx>( - mut map: LocalTranslationMap, + fn check_operands<'tcx>( + map: &mut LocalTranslationMap, l: &Operand<'tcx>, r: &Operand<'tcx>, - ) -> Option { + ) -> Result<(), TranslationError> { match (l, r) { (Operand::Copy(l_place), Operand::Copy(r_place)) => { - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } (Operand::Move(l_place), Operand::Move(r_place)) => { - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } (Operand::Constant(l_const_operand), Operand::Constant(r_const_operand)) => { if l_const_operand.user_ty != r_const_operand.user_ty { - return None; + return Err(TranslationError); } if l_const_operand.const_ != r_const_operand.const_ { - return None; + return Err(TranslationError); } } ( @@ -517,69 +605,69 @@ impl TranslationMap { Operand::RuntimeChecks(r_runtime_checks), ) => { if l_runtime_checks != r_runtime_checks { - return None; + return Err(TranslationError); } } _ => { - return None; + return Err(TranslationError); } } - Some(map) + Ok(()) } - fn non_diverging_intrinsics_functionally_equivalent<'tcx>( - mut map: LocalTranslationMap, + fn check_non_diverging_intrinsics<'tcx>( + map: &mut LocalTranslationMap, l: &NonDivergingIntrinsic<'tcx>, r: &NonDivergingIntrinsic<'tcx>, - ) -> Option { + ) -> Result<(), TranslationError> { match (l, r) { ( NonDivergingIntrinsic::Assume(l_operand), NonDivergingIntrinsic::Assume(r_operand), ) => { - map = Self::operands_functionally_equivalent(map, l_operand, r_operand)?; + Self::check_operands(map, l_operand, r_operand)?; } ( NonDivergingIntrinsic::CopyNonOverlapping(l_copy_non_overlapping), NonDivergingIntrinsic::CopyNonOverlapping(r_copy_non_overlapping), ) => { - map = Self::operands_functionally_equivalent( + Self::check_operands( map, &l_copy_non_overlapping.src, &r_copy_non_overlapping.src, )?; - map = Self::operands_functionally_equivalent( + Self::check_operands( map, &l_copy_non_overlapping.dst, &r_copy_non_overlapping.dst, )?; - map = Self::operands_functionally_equivalent( + Self::check_operands( map, &l_copy_non_overlapping.count, &r_copy_non_overlapping.count, )?; } _ => { - return None; + return Err(TranslationError); } } - Some(map) + Ok(()) } - fn terminator_kinds_functionally_equivalent<'tcx>( - mut map: LocalTranslationMap, + fn check_terminator_kinds<'tcx>( + map: &mut LocalTranslationMap, l: &TerminatorKind<'tcx>, r: &TerminatorKind<'tcx>, - ) -> Option { + ) -> Result<(), TranslationError> { match (l, r) { (TerminatorKind::Goto { target: _ }, TerminatorKind::Goto { target: _ }) => {} ( TerminatorKind::SwitchInt { discr: l_discr, targets: _ }, TerminatorKind::SwitchInt { discr: r_discr, targets: _ }, ) => { - map = Self::operands_functionally_equivalent(map, l_discr, r_discr)?; + Self::check_operands(map, l_discr, r_discr)?; } (TerminatorKind::UnwindResume, TerminatorKind::UnwindResume) => {} ( @@ -587,7 +675,7 @@ impl TranslationMap { TerminatorKind::UnwindTerminate(r_unwind_terminate_reason), ) => { if l_unwind_terminate_reason != r_unwind_terminate_reason { - return None; + return Err(TranslationError); } } (TerminatorKind::Return, TerminatorKind::Return) => {} @@ -609,12 +697,12 @@ impl TranslationMap { }, ) => { if discriminant(l_unwind) != discriminant(r_unwind) { - return None; + return Err(TranslationError); } if l_replace != r_replace { - return None; + return Err(TranslationError); } - map = map.insert(l_place.local, r_place.local)?; + map.insert(l_place.local, r_place.local)?; } ( TerminatorKind::Call { @@ -637,18 +725,18 @@ impl TranslationMap { }, ) => { if discriminant(l_unwind) != discriminant(r_unwind) { - return None; + return Err(TranslationError); } if l_call_source != r_call_source { - return None; + return Err(TranslationError); } if l_args.len() != r_args.len() { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_func, r_func)?; - map = map.insert(l_destination.local, r_destination.local)?; + Self::check_operands(map, l_func, r_func)?; + map.insert(l_destination.local, r_destination.local)?; for (l_arg, r_arg) in l_args.iter().zip(r_args.iter()) { - map = Self::operands_functionally_equivalent(map, &l_arg.node, &r_arg.node)?; + Self::check_operands(map, &l_arg.node, &r_arg.node)?; } } ( @@ -656,11 +744,11 @@ impl TranslationMap { TerminatorKind::TailCall { func: r_func, args: r_args, fn_span: _ }, ) => { if l_args.len() != r_args.len() { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_func, r_func)?; + Self::check_operands(map, l_func, r_func)?; for (l_arg, r_arg) in l_args.iter().zip(r_args.iter()) { - map = Self::operands_functionally_equivalent(map, &l_arg.node, &r_arg.node)?; + Self::check_operands(map, &l_arg.node, &r_arg.node)?; } } ( @@ -680,13 +768,13 @@ impl TranslationMap { }, ) => { if l_expected != r_expected { - return None; + return Err(TranslationError); } if discriminant(l_unwind) != discriminant(r_unwind) { - return None; + return Err(TranslationError); } - map = Self::operand_assert_kinds_functionally_equivalent(map, l_msg, r_msg)?; - map = Self::operands_functionally_equivalent(map, l_cond, r_cond)?; + Self::check_operand_assert_kinds(map, l_msg, r_msg)?; + Self::check_operands(map, l_cond, r_cond)?; } ( TerminatorKind::Yield { @@ -702,8 +790,8 @@ impl TranslationMap { drop: _, }, ) => { - map = Self::operands_functionally_equivalent(map, l_value, r_value)?; - map = map.insert(l_resume_arg.local, r_resume_arg.local)?; + Self::check_operands(map, l_value, r_value)?; + map.insert(l_resume_arg.local, r_resume_arg.local)?; } (TerminatorKind::CoroutineDrop, TerminatorKind::CoroutineDrop) => {} ( @@ -715,58 +803,58 @@ impl TranslationMap { TerminatorKind::FalseUnwind { real_target: _, unwind: r_unwind }, ) => { if discriminant(l_unwind) != discriminant(r_unwind) { - return None; + return Err(TranslationError); } } (TerminatorKind::InlineAsm { .. }, TerminatorKind::InlineAsm { .. }) => { // Let's not risk messing with asm... - return None; + return Err(TranslationError); } _ => { - return None; + return Err(TranslationError); } } - Some(map) + Ok(()) } - fn operand_assert_kinds_functionally_equivalent<'tcx>( - mut map: LocalTranslationMap, + fn check_operand_assert_kinds<'tcx>( + map: &mut LocalTranslationMap, l: &AssertKind>, r: &AssertKind>, - ) -> Option { + ) -> Result<(), TranslationError> { match (l, r) { ( AssertKind::BoundsCheck { len: l_len, index: l_index }, AssertKind::BoundsCheck { len: r_len, index: r_index }, ) => { - map = Self::operands_functionally_equivalent(map, l_len, r_len)?; - map = Self::operands_functionally_equivalent(map, l_index, r_index)?; + Self::check_operands(map, l_len, r_len)?; + Self::check_operands(map, l_index, r_index)?; } ( AssertKind::Overflow(l_bin_op, l_op_0, l_op_1), AssertKind::Overflow(r_bin_op, r_op_0, r_op_1), ) => { if l_bin_op != r_bin_op { - return None; + return Err(TranslationError); } - map = Self::operands_functionally_equivalent(map, l_op_0, r_op_0)?; - map = Self::operands_functionally_equivalent(map, l_op_1, r_op_1)?; + Self::check_operands(map, l_op_0, r_op_0)?; + Self::check_operands(map, l_op_1, r_op_1)?; } (AssertKind::OverflowNeg(l_op), AssertKind::OverflowNeg(r_op)) => { - map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + Self::check_operands(map, l_op, r_op)?; } (AssertKind::DivisionByZero(l_op), AssertKind::DivisionByZero(r_op)) => { - map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + Self::check_operands(map, l_op, r_op)?; } (AssertKind::RemainderByZero(l_op), AssertKind::RemainderByZero(r_op)) => { - map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + Self::check_operands(map, l_op, r_op)?; } ( AssertKind::ResumedAfterReturn(l_coroutine_kind), AssertKind::ResumedAfterReturn(r_coroutine_kind), ) => { if l_coroutine_kind != r_coroutine_kind { - return None; + return Err(TranslationError); } } ( @@ -774,7 +862,7 @@ impl TranslationMap { AssertKind::ResumedAfterPanic(r_coroutine_kind), ) => { if l_coroutine_kind != r_coroutine_kind { - return None; + return Err(TranslationError); } } ( @@ -782,15 +870,15 @@ impl TranslationMap { AssertKind::ResumedAfterDrop(r_coroutine_kind), ) => { if l_coroutine_kind != r_coroutine_kind { - return None; + return Err(TranslationError); } } ( AssertKind::MisalignedPointerDereference { required: l_required, found: l_found }, AssertKind::MisalignedPointerDereference { required: r_required, found: r_found }, ) => { - map = Self::operands_functionally_equivalent(map, l_required, r_required)?; - map = Self::operands_functionally_equivalent(map, l_found, r_found)?; + Self::check_operands(map, l_required, r_required)?; + Self::check_operands(map, l_found, r_found)?; } (AssertKind::NullPointerDereference, AssertKind::NullPointerDereference) => {} (AssertKind::NullReferenceConstructed, AssertKind::NullReferenceConstructed) => {} @@ -798,13 +886,15 @@ impl TranslationMap { AssertKind::InvalidEnumConstruction(l_op), AssertKind::InvalidEnumConstruction(r_op), ) => { - map = Self::operands_functionally_equivalent(map, l_op, r_op)?; + Self::check_operands(map, l_op, r_op)?; } _ => { - return None; + return Err(TranslationError); } } - Some(map) + Ok(()) } } + +struct TranslationError; From d1fa25a3fd8af4d5c8a393f433d1bb7a9814da63 Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Mon, 7 Sep 2026 15:41:22 +0200 Subject: [PATCH 10/12] Update policy --- compiler/rustc_mir_transform/src/merge_yields.rs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/compiler/rustc_mir_transform/src/merge_yields.rs b/compiler/rustc_mir_transform/src/merge_yields.rs index a23cb13600f1e..f9211d83de409 100644 --- a/compiler/rustc_mir_transform/src/merge_yields.rs +++ b/compiler/rustc_mir_transform/src/merge_yields.rs @@ -99,9 +99,9 @@ use rustc_mir_dataflow::impls::{MaybeStorageLive, always_storage_live_locals}; use rustc_span::DUMMY_SP; use tracing::instrument; -use crate::MirPass; use crate::pass_manager::PassPolicy; use crate::simplify::remove_dead_blocks; +use crate::{MirPass, PassCtx}; pub(super) struct MergeYields; @@ -172,8 +172,8 @@ impl<'tcx> MirPass<'tcx> for MergeYields { remove_dead_blocks(body); } - fn policy(&self, _sess: &rustc_session::Session) -> PassPolicy { - PassPolicy::optimization(true) + fn policy(&self, _ctx: &PassCtx<'_>) -> PassPolicy { + PassPolicy::optional(true) } } From 3a50fa07d5152e49f0d94506bfd0b388c958cb4b Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Wed, 30 Sep 2026 09:46:08 +0200 Subject: [PATCH 11/12] Update after rebase --- compiler/rustc_mir_transform/src/merge_yields.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/compiler/rustc_mir_transform/src/merge_yields.rs b/compiler/rustc_mir_transform/src/merge_yields.rs index f9211d83de409..420a61569cba8 100644 --- a/compiler/rustc_mir_transform/src/merge_yields.rs +++ b/compiler/rustc_mir_transform/src/merge_yields.rs @@ -334,7 +334,7 @@ impl TranslationMap { Some(Terminator { source_info: SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, kind: TerminatorKind::Goto { target: *to }, - attributes: Default::default(), + loop_hint_attrs: Default::default(), }), false, )); From b857943555793073f95f277000039522eaa063ea Mon Sep 17 00:00:00 2001 From: Dion Dokter Date: Wed, 30 Sep 2026 14:35:45 +0200 Subject: [PATCH 12/12] Check if locals are initialize before moving them --- .../rustc_mir_transform/src/merge_yields.rs | 94 +++++++++++++------ 1 file changed, 66 insertions(+), 28 deletions(-) diff --git a/compiler/rustc_mir_transform/src/merge_yields.rs b/compiler/rustc_mir_transform/src/merge_yields.rs index 420a61569cba8..41a2800ced213 100644 --- a/compiler/rustc_mir_transform/src/merge_yields.rs +++ b/compiler/rustc_mir_transform/src/merge_yields.rs @@ -89,13 +89,16 @@ use rustc_data_structures::fx::{FxHashMap, FxHashSet}; use rustc_data_structures::graph::Successors; use rustc_data_structures::indexmap::{IndexMap, IndexSet}; use rustc_middle::mir::{ - AssertKind, BasicBlock, BasicBlockData, Body, Local, NonDivergingIntrinsic, + AssertKind, BasicBlock, BasicBlockData, Body, Local, MirDumper, NonDivergingIntrinsic, OUTERMOST_SOURCE_SCOPE, Operand, Place, Rvalue, SourceInfo, Statement, StatementKind, Terminator, TerminatorKind, WithRetag, }; use rustc_middle::ty::TyCtxt; use rustc_mir_dataflow::Analysis; -use rustc_mir_dataflow::impls::{MaybeStorageLive, always_storage_live_locals}; +use rustc_mir_dataflow::impls::{ + MaybeInitializedPlaces, MaybeStorageLive, always_storage_live_locals, +}; +use rustc_mir_dataflow::move_paths::MoveData; use rustc_span::DUMMY_SP; use tracing::instrument; @@ -112,6 +115,10 @@ impl<'tcx> MirPass<'tcx> for MergeYields { return; } + if let Some(dumper) = MirDumper::new(tcx, "merge_yields_before", body) { + dumper.dump_mir(body); + } + tracing::debug!("running pass for {}", tcx.def_path_debug_str(body.source.def_id())); let mut yields = body @@ -170,6 +177,10 @@ impl<'tcx> MirPass<'tcx> for MergeYields { } remove_dead_blocks(body); + + if let Some(dumper) = MirDumper::new(tcx, "merge_yields_after", body) { + dumper.dump_mir(body); + } } fn policy(&self, _ctx: &PassCtx<'_>) -> PassPolicy { @@ -310,16 +321,30 @@ impl TranslationMap { fn redirect_entry_points<'tcx>(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) { let body_predecessors = body.basic_blocks.predecessors().clone(); let always_live_locals = always_storage_live_locals(body); - let mut results = MaybeStorageLive::new(Cow::Borrowed(&always_live_locals)) - .iterate_to_fixpoint(tcx, body, Some("callapse_yields")) + let mut maybe_live_cursor = MaybeStorageLive::new(Cow::Borrowed(&always_live_locals)) + .iterate_to_fixpoint(tcx, body, Some("merge_yields")) + .into_results_cursor(body); + + let bb_live_locals = self + .blocks + .keys() + .map(|bb| { + maybe_live_cursor.seek_to_block_start(*bb); + (*bb, maybe_live_cursor.get().clone()) + }) + .collect::>(); + + let move_data = MoveData::gather_moves(body, tcx, |_| true); + let mut maybe_initialized_cursor = MaybeInitializedPlaces::new(tcx, body, &move_data) + .iterate_to_fixpoint(tcx, body, Some("merge_yields")) .into_results_cursor(body); - let from_live_locals = self + let bb_initialized_locals = self .blocks .keys() - .map(|from| { - results.seek_to_block_start(*from); - (*from, results.get().clone()) + .map(|bb| { + maybe_initialized_cursor.seek_to_block_start(*bb); + (*bb, maybe_initialized_cursor.get().clone()) }) .collect::>(); @@ -352,32 +377,45 @@ impl TranslationMap { let inbetween_data = &mut body.basic_blocks_mut()[inbetween]; inbetween_data.is_cleanup = is_cleanup; - let live_locals = &from_live_locals[from]; - for from_local in live_locals.iter() { - if let Some(to_local) = self.locals.locals.get(&from_local) { - if from_local == *to_local { - continue; - } - - if !always_live_locals.contains(*to_local) { - inbetween_data.statements.push(Statement::new( - SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, - StatementKind::StorageLive(*to_local), - )); - } + for (from_local, to_local) in self.locals.locals.iter() { + if from_local == to_local { + // This local doesn't need translation + continue; + } + + if !bb_live_locals[from].contains(*from_local) { + // This local isn't live at this point + continue; + } + + if !always_live_locals.contains(*to_local) { + // Mark the new (to) local live + inbetween_data.statements.push(Statement::new( + SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, + StatementKind::StorageLive(*to_local), + )); + } + if let Some(mpi) = move_data.rev_lookup.find_local(*from_local) + && bb_initialized_locals[&from].contains(mpi) + { + // Move the initialized local `from` -> `to` inbetween_data.statements.push(Statement::new( SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, StatementKind::Assign(Box::new(( Place::from(*to_local), - Rvalue::Use(Operand::Move(Place::from(from_local)), WithRetag::Yes), + Rvalue::Use( + Operand::Move(Place::from(*from_local)), + WithRetag::Yes, + ), ))), )); - if !always_live_locals.contains(from_local) { - inbetween_data.statements.push(Statement::new( - SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, - StatementKind::StorageDead(from_local), - )); - } + } + if !always_live_locals.contains(*from_local) { + // Mark the old (from) local dead + inbetween_data.statements.push(Statement::new( + SourceInfo { span: DUMMY_SP, scope: OUTERMOST_SOURCE_SCOPE }, + StatementKind::StorageDead(*from_local), + )); } } }