Skip to main content

rustc_mir_transform/
match_branches.rs

1use rustc_abi::Integer;
2use rustc_const_eval::const_eval::mk_eval_cx_for_const_val;
3use rustc_index::bit_set::DenseBitSet;
4use rustc_middle::mir::*;
5use rustc_middle::ty::layout::{IntegerExt, TyAndLayout};
6use rustc_middle::ty::util::Discr;
7use rustc_middle::ty::{self, ScalarInt, Ty, TyCtxt};
8
9use super::simplify::simplify_cfg;
10use crate::PassPolicy;
11use crate::patch::MirPatch;
12use crate::unreachable_prop::remove_successors_from_switch;
13
14/// Unifies all targets into one basic block if each statement can have the same statement.
15pub(super) struct MatchBranchSimplification;
16
17impl<'tcx> crate::MirPass<'tcx> for MatchBranchSimplification {
18    fn policy(&self, ctx: &crate::PassCtx<'_>) -> PassPolicy {
19        // Enable only under -Zmir-opt-level=2 as this can make programs less debuggable.
20        PassPolicy::optional(ctx.mir_opt_level() >= 2)
21    }
22
23    fn run_pass(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) {
24        let typing_env = body.typing_env(tcx);
25        let mut changed = false;
26        for bb in body.basic_blocks.indices() {
27            if !candidate_match(body, bb) {
28                continue;
29            };
30            changed |= simplify_match(tcx, typing_env, body, bb)
31        }
32
33        if changed {
34            simplify_cfg(tcx, body);
35        }
36    }
37}
38
39struct SimplifyMatch<'tcx, 'a> {
40    tcx: TyCtxt<'tcx>,
41    typing_env: ty::TypingEnv<'tcx>,
42    patch: MirPatch<'tcx>,
43    body: &'a Body<'tcx>,
44    switch_bb: BasicBlock,
45    discr: &'a Operand<'tcx>,
46    discr_local: Option<Local>,
47    discr_ty: Ty<'tcx>,
48    borrowed_locals: Option<DenseBitSet<Local>>,
49}
50
51impl<'tcx, 'a> SimplifyMatch<'tcx, 'a> {
52    fn discr_local(&mut self) -> Local {
53        *self.discr_local.get_or_insert_with(|| {
54            // Introduce a temporary for the discriminant value.
55            let source_info = self.body.basic_blocks[self.switch_bb].terminator().source_info;
56            self.patch.new_temp(self.discr_ty, source_info.span)
57        })
58    }
59
60    /// Unifies the assignments if all rvalues are constants and equal.
61    fn unify_if_equal_const(
62        &self,
63        dest: Place<'tcx>,
64        consts: &[(u128, &ConstOperand<'tcx>)],
65        otherwise: Option<&ConstOperand<'tcx>>,
66    ) -> Option<StatementKind<'tcx>> {
67        let (_, first_const, mut others) = split_first_case(consts, otherwise);
68        let first_scalar_int = first_const.const_.try_eval_scalar_int(self.tcx, self.typing_env)?;
69        if others.all(|const_| {
70            const_.const_.try_eval_scalar_int(self.tcx, self.typing_env) == Some(first_scalar_int)
71        }) {
72            Some(StatementKind::Assign(Box::new((
73                dest,
74                // We didn't remember the `WithRetag` of the original assignments, so in case
75                // one of them had "no", we also have to use "no" here.
76                Rvalue::Use(Operand::Constant(Box::new(first_const.clone())), WithRetag::No),
77            ))))
78        } else {
79            None
80        }
81    }
82
83    /// If a source block is found that switches between two blocks that are exactly
84    /// the same modulo const bool assignments (e.g., one assigns true another false
85    /// to the same place), unify a target block statements into the source block,
86    /// using Eq / Ne comparison with switch value where const bools value differ.
87    ///
88    /// For example:
89    ///
90    /// ```ignore (MIR)
91    /// bb0: {
92    ///     switchInt(move _3) -> [42_isize: bb1, otherwise: bb2];
93    /// }
94    ///
95    /// bb1: {
96    ///     _2 = const true;
97    ///     goto -> bb3;
98    /// }
99    ///
100    /// bb2: {
101    ///     _2 = const false;
102    ///     goto -> bb3;
103    /// }
104    /// ```
105    ///
106    /// into:
107    ///
108    /// ```ignore (MIR)
109    /// bb0: {
110    ///    _2 = Eq(move _3, const 42_isize);
111    ///    goto -> bb3;
112    /// }
113    /// ```
114    fn unify_by_eq_op(
115        &mut self,
116        dest: Place<'tcx>,
117        consts: &[(u128, &ConstOperand<'tcx>)],
118        otherwise: Option<&ConstOperand<'tcx>>,
119    ) -> Option<StatementKind<'tcx>> {
120        // FIXME: extend to any case.
121        let (first_case, first_const, mut others) = split_first_case(consts, otherwise);
122        if !first_const.ty().is_bool() {
123            return None;
124        }
125        let first_bool = first_const.const_.try_eval_bool(self.tcx, self.typing_env)?;
126        if others.all(|const_| {
127            const_.const_.try_eval_bool(self.tcx, self.typing_env) == Some(!first_bool)
128        }) {
129            // Make value conditional on switch condition.
130            let size =
131                self.tcx.layout_of(self.typing_env.as_query_input(self.discr_ty)).unwrap().size;
132            let const_cmp = Operand::const_from_scalar(
133                self.tcx,
134                self.discr_ty,
135                rustc_const_eval::interpret::Scalar::from_uint(first_case, size),
136                rustc_span::DUMMY_SP,
137            );
138            let op = if first_bool { BinOp::Eq } else { BinOp::Ne };
139            let rval = Rvalue::BinaryOp(
140                op,
141                Box::new((Operand::Copy(Place::from(self.discr_local())), const_cmp)),
142            );
143            Some(StatementKind::Assign(Box::new((dest, rval))))
144        } else {
145            None
146        }
147    }
148
149    /// Unifies the assignments if all rvalues can be cast from the discriminant value by IntToInt.
150    ///
151    /// For example:
152    ///
153    /// ```ignore (MIR)
154    /// bb0: {
155    ///     switchInt(_1) -> [1: bb2, 2: bb3, 3: bb4, otherwise: bb1];
156    /// }
157    ///
158    /// bb1: {
159    ///     unreachable;
160    /// }
161    ///
162    /// bb2: {
163    ///     _0 = const 1_i16;
164    ///     goto -> bb5;
165    /// }
166    ///
167    /// bb3: {
168    ///     _0 = const 2_i16;
169    ///     goto -> bb5;
170    /// }
171    ///
172    /// bb4: {
173    ///     _0 = const 3_i16;
174    ///     goto -> bb5;
175    /// }
176    /// ```
177    ///
178    /// into:
179    ///
180    /// ```ignore (MIR)
181    /// bb0: {
182    ///    _0 = _1 as i16 (IntToInt);
183    ///    goto -> bb5;
184    /// }
185    /// ```
186    fn unify_by_int_to_int(
187        &mut self,
188        dest: Place<'tcx>,
189        consts: &[(u128, &ConstOperand<'tcx>)],
190    ) -> Option<StatementKind<'tcx>> {
191        let (_, first_const) = consts[0];
192        if !first_const.ty().is_integral() {
193            return None;
194        }
195        let discr_layout =
196            self.tcx.layout_of(self.typing_env.as_query_input(self.discr_ty)).unwrap();
197        if consts.iter().all(|&(case, const_)| {
198            let Some(scalar_int) = const_.const_.try_eval_scalar_int(self.tcx, self.typing_env)
199            else {
200                return false;
201            };
202            can_cast(self.tcx, case, discr_layout, const_.ty(), scalar_int)
203        }) {
204            let operand = Operand::Copy(Place::from(self.discr_local()));
205            let rval = if first_const.ty() == self.discr_ty {
206                Rvalue::Use(operand, WithRetag::No)
207            } else {
208                Rvalue::Cast(CastKind::IntToInt, operand, first_const.ty())
209            };
210            Some(StatementKind::Assign(Box::new((dest, rval))))
211        } else {
212            None
213        }
214    }
215
216    /// This is primarily used to unify these copy statements that simplified the canonical enum clone method by GVN.
217    /// The GVN simplified
218    /// ```ignore (syntax-highlighting-only)
219    /// match a {
220    ///     Foo::A(x) => Foo::A(*x),
221    ///     Foo::B => Foo::B
222    /// }
223    /// ```
224    /// to
225    /// ```ignore (syntax-highlighting-only)
226    /// match a {
227    ///     Foo::A(_x) => a, // copy a
228    ///     Foo::B => Foo::B
229    /// }
230    /// ```
231    /// This will simplify into a copy statement.
232    fn unify_by_copy(
233        &mut self,
234        dest: Place<'tcx>,
235        rvals: &[(u128, &Rvalue<'tcx>)],
236    ) -> Option<StatementKind<'tcx>> {
237        let bbs = &self.body.basic_blocks;
238        // Check if the copy source matches the following pattern.
239        // _2 = discriminant(*_1); // "*_1" is the expected the copy source.
240        // switchInt(move _2) -> [0: bb3, 1: bb2, otherwise: bb1];
241        let &Statement {
242            kind: StatementKind::Assign((discr_place, Rvalue::Discriminant(copy_src_place))),
243            ..
244        } = bbs[self.switch_bb].statements.last()?
245        else {
246            return None;
247        };
248        if self.discr.place() != Some(discr_place) {
249            return None;
250        }
251        let src_ty = copy_src_place.ty(self.body.local_decls(), self.tcx);
252        if !src_ty.ty.is_enum() || src_ty.variant_index.is_some() {
253            return None;
254        }
255        let dest_ty = dest.ty(self.body.local_decls(), self.tcx);
256        if dest_ty.ty != src_ty.ty || dest_ty.variant_index.is_some() {
257            return None;
258        }
259        let ty::Adt(def, _) = dest_ty.ty.kind() else {
260            return None;
261        };
262
263        if copy_src_place.is_indirect() {
264            // If the src place is indirect, only permit generating the copy when the dest place is
265            // never borrowed.
266            let borrowed_locals = self
267                .borrowed_locals
268                .get_or_insert_with(|| rustc_mir_dataflow::impls::borrowed_locals(self.body));
269            if borrowed_locals.contains(dest.local) {
270                return None;
271            }
272        } else if copy_src_place.local == dest.local {
273            // Also forbid the case where the source and dest are fields of the same local
274            return None;
275        }
276
277        for &(case, rvalue) in rvals.iter() {
278            match rvalue {
279                // Check if `_3 = const Foo::B` can be transformed to `_3 = copy *_1`.
280                Rvalue::Use(Operand::Constant(constant), _)
281                    if let Const::Val(const_, ty) = constant.const_ =>
282                {
283                    let (ecx, op) = mk_eval_cx_for_const_val(
284                        self.tcx.at(constant.span),
285                        self.typing_env,
286                        const_,
287                        ty,
288                    )?;
289                    let variant = ecx.read_discriminant(&op).discard_err()?;
290                    if !def.variants()[variant].fields.is_empty() {
291                        return None;
292                    }
293                    let Discr { val, .. } = ty.discriminant_for_variant(self.tcx, variant)?;
294                    if val != case {
295                        return None;
296                    }
297                }
298                Rvalue::Use(Operand::Copy(src_place), _) if *src_place == copy_src_place => {}
299                // Check if `_3 = Foo::B` can be transformed to `_3 = copy *_1`.
300                Rvalue::Aggregate(AggregateKind::Adt(_, variant_index, _, _, None), fields)
301                    if fields.is_empty()
302                        && let Some(Discr { val, .. }) =
303                            src_ty.ty.discriminant_for_variant(self.tcx, *variant_index)
304                        && val == case => {}
305                _ => return None,
306            }
307        }
308        // We didn't remember the `WithRetag` of the original assignments, so in case
309        // one of them had "no", we also have to use "no" here.
310        Some(StatementKind::Assign(Box::new((
311            dest,
312            Rvalue::Use(Operand::Copy(copy_src_place), WithRetag::No),
313        ))))
314    }
315
316    /// Returns a new statement if we can use the statement replace all statements.
317    fn try_unify_stmts(
318        &mut self,
319        index: usize,
320        stmts: &[(u128, &StatementKind<'tcx>)],
321        otherwise: Option<&StatementKind<'tcx>>,
322    ) -> Option<StatementKind<'tcx>> {
323        if let Some(new_stmt) = identical_stmts(stmts, otherwise) {
324            return Some(new_stmt);
325        }
326
327        let (dest, rvals, otherwise) = candidate_assign(stmts, otherwise)?;
328        if let Some((consts, otherwise)) = candidate_const(&rvals, otherwise) {
329            if let Some(new_stmt) = self.unify_if_equal_const(dest, &consts, otherwise) {
330                return Some(new_stmt);
331            }
332            if let Some(new_stmt) = self.unify_by_eq_op(dest, &consts, otherwise) {
333                return Some(new_stmt);
334            }
335            // Requires the otherwise is unreachable.
336            if otherwise.is_none()
337                && let Some(new_stmt) = self.unify_by_int_to_int(dest, &consts)
338            {
339                return Some(new_stmt);
340            }
341        }
342
343        // We only know the first statement is safe to introduce new dereferences.
344        if index == 0
345            // We cannot create overlapping assignments.
346            && dest.is_stable_offset()
347            // Requires the otherwise is unreachable.
348            && otherwise.is_none()
349            && let Some(new_stmt) = self.unify_by_copy(dest, &rvals)
350        {
351            return Some(new_stmt);
352        }
353        None
354    }
355}
356
357/// Returns the first case target if all targets have an equal number of statements and identical destination.
358fn candidate_match<'tcx>(body: &Body<'tcx>, switch_bb: BasicBlock) -> bool {
359    use itertools::Itertools;
360    let targets = match &body.basic_blocks[switch_bb].terminator().kind {
361        TerminatorKind::SwitchInt {
362            discr: Operand::Copy(_) | Operand::Move(_), targets, ..
363        } => targets,
364        // Only optimize switch int statements
365        _ => return false,
366    };
367    // We require that the possible target blocks don't contain this block.
368    if targets.all_targets().contains(&switch_bb) {
369        return false;
370    }
371    // We require that the possible target blocks all be distinct.
372    if !targets.is_distinct() {
373        return false;
374    }
375    // Check that destinations are identical, and if not, then don't optimize this block
376    targets
377        .all_targets()
378        .iter()
379        .map(|&bb| &body.basic_blocks[bb])
380        .filter(|bb| !bb.is_empty_unreachable())
381        .map(|bb| (bb.statements.len(), &bb.terminator().kind))
382        .all_equal()
383}
384
385fn simplify_match<'tcx>(
386    tcx: TyCtxt<'tcx>,
387    typing_env: ty::TypingEnv<'tcx>,
388    body: &mut Body<'tcx>,
389    switch_bb: BasicBlock,
390) -> bool {
391    let (discr, targets) = match &body.basic_blocks[switch_bb].terminator().kind {
392        TerminatorKind::SwitchInt { discr, targets, .. } => (discr, targets),
393        _ => unreachable!(),
394    };
395    let mut simplify_match = SimplifyMatch {
396        tcx,
397        typing_env,
398        patch: MirPatch::new(body),
399        body,
400        switch_bb,
401        discr,
402        discr_local: None,
403        discr_ty: discr.ty(body.local_decls(), tcx),
404        borrowed_locals: None,
405    };
406    let reachable_cases: Vec<_> =
407        targets.iter().filter(|&(_, bb)| !body.basic_blocks[bb].is_empty_unreachable()).collect();
408    let mut new_stmts = Vec::new();
409    let otherwise = if body.basic_blocks[targets.otherwise()].is_empty_unreachable() {
410        None
411    } else {
412        Some(targets.otherwise())
413    };
414    // We can patch the terminator to goto because there is a single target.
415    match (reachable_cases.len(), otherwise.is_none()) {
416        (1, true) | (0, false) => {
417            let mut patch = simplify_match.patch;
418            remove_successors_from_switch(tcx, switch_bb, body, &mut patch, |bb| {
419                body.basic_blocks[bb].is_empty_unreachable()
420            });
421            patch.apply(body);
422            return true;
423        }
424        _ => {}
425    }
426    let Some(&(_, first_case_bb)) = reachable_cases.first() else {
427        return false;
428    };
429    let stmt_len = body.basic_blocks[first_case_bb].statements.len();
430    let mut cases = Vec::with_capacity(stmt_len);
431    // Check at each position in the basic blocks whether these statements can be unified.
432    for index in 0..stmt_len {
433        cases.clear();
434        let otherwise = otherwise.map(|bb| &body.basic_blocks[bb].statements[index].kind);
435        for &(case, bb) in &reachable_cases {
436            cases.push((case, &body.basic_blocks[bb].statements[index].kind));
437        }
438        let Some(new_stmt) = simplify_match.try_unify_stmts(index, &cases, otherwise) else {
439            return false;
440        };
441        new_stmts.push(new_stmt);
442    }
443    // Take ownership of items now that we know we can optimize.
444    let discr = discr.clone();
445
446    let statement_index = body.basic_blocks[switch_bb].statements.len();
447    let parent_end = Location { block: switch_bb, statement_index };
448    let mut patch = simplify_match.patch;
449    if let Some(discr_local) = simplify_match.discr_local {
450        patch.add_statement(parent_end, StatementKind::StorageLive(discr_local));
451        patch.add_assign(parent_end, Place::from(discr_local), Rvalue::Use(discr, WithRetag::No));
452    }
453    for new_stmt in new_stmts {
454        patch.add_statement(parent_end, new_stmt);
455    }
456    if let Some(discr_local) = simplify_match.discr_local {
457        patch.add_statement(parent_end, StatementKind::StorageDead(discr_local));
458    }
459    patch.patch_terminator(switch_bb, body.basic_blocks[first_case_bb].terminator().kind.clone());
460    patch.apply(body);
461    true
462}
463
464/// Check if the cast constant using `IntToInt` is equal to the target constant.
465fn can_cast(
466    tcx: TyCtxt<'_>,
467    src_val: impl Into<u128>,
468    src_layout: TyAndLayout<'_>,
469    cast_ty: Ty<'_>,
470    target_scalar: ScalarInt,
471) -> bool {
472    let from_scalar = ScalarInt::try_from_uint(src_val.into(), src_layout.size).unwrap();
473    let v = match src_layout.ty.kind() {
474        ty::Uint(_) => from_scalar.to_uint(src_layout.size),
475        ty::Int(_) => from_scalar.to_int(src_layout.size) as u128,
476        // We can also transform the values of other integer representations (such as char),
477        // although this may not be practical in real-world scenarios.
478        _ => return false,
479    };
480    let size = match *cast_ty.kind() {
481        ty::Int(t) => Integer::from_int_ty(&tcx, t).size(),
482        ty::Uint(t) => Integer::from_uint_ty(&tcx, t).size(),
483        _ => return false,
484    };
485    let v = size.truncate(v);
486    let cast_scalar = ScalarInt::try_from_uint(v, size).unwrap();
487    cast_scalar == target_scalar
488}
489
490fn candidate_assign<'tcx, 'a>(
491    stmts: &'a [(u128, &'a StatementKind<'tcx>)],
492    otherwise: Option<&'a StatementKind<'tcx>>,
493) -> Option<(Place<'tcx>, Vec<(u128, &'a Rvalue<'tcx>)>, Option<&'a Rvalue<'tcx>>)> {
494    let (_, first_stmt) = stmts[0];
495    let (dest, _) = first_stmt.as_assign()?;
496    let otherwise = if let Some(otherwise) = otherwise {
497        let Some((otherwise_dest, rval)) = otherwise.as_assign() else {
498            return None;
499        };
500        if otherwise_dest != dest {
501            return None;
502        }
503        Some(rval)
504    } else {
505        None
506    };
507    let rvals = stmts
508        .into_iter()
509        .map(|&(case, stmt)| {
510            let (other_dest, rval) = stmt.as_assign()?;
511            if other_dest != dest {
512                return None;
513            }
514            Some((case, rval))
515        })
516        .try_collect()?;
517    Some((*dest, rvals, otherwise))
518}
519
520// Returns all ConstOperands if all Rvalues are ConstOperands.
521fn candidate_const<'tcx, 'a>(
522    rvals: &'a [(u128, &'a Rvalue<'tcx>)],
523    otherwise: Option<&'a Rvalue<'tcx>>,
524) -> Option<(Vec<(u128, &'a ConstOperand<'tcx>)>, Option<&'a ConstOperand<'tcx>>)> {
525    // We ignore the retag mode here, which means the `Use` we insert later must be without retag.
526    let otherwise = if let Some(otherwise) = otherwise {
527        let Rvalue::Use(Operand::Constant(const_), _) = otherwise else {
528            return None;
529        };
530        Some(&**const_)
531    } else {
532        None
533    };
534    let consts = rvals
535        .into_iter()
536        .map(|&(case, rval)| {
537            let Rvalue::Use(Operand::Constant(const_), _) = rval else { return None };
538            Some((case, &**const_))
539        })
540        .try_collect()?;
541    Some((consts, otherwise))
542}
543
544// Returns the first case and others (including otherwise if present).
545fn split_first_case<'a, T>(
546    stmts: &'a [(u128, &'a T)],
547    otherwise: Option<&'a T>,
548) -> (u128, &'a T, impl Iterator<Item = &'a T>) {
549    let (first_case, first) = stmts[0];
550    (first_case, first, stmts[1..].into_iter().map(|&(_, val)| val).chain(otherwise))
551}
552
553// If all statements are identical, we can optimize.
554fn identical_stmts<'tcx>(
555    stmts: &[(u128, &StatementKind<'tcx>)],
556    otherwise: Option<&StatementKind<'tcx>>,
557) -> Option<StatementKind<'tcx>> {
558    use itertools::Itertools;
559    let (_, first_stmt, others) = split_first_case(stmts, otherwise);
560    if std::iter::once(first_stmt).chain(others).all_equal() {
561        return Some(first_stmt.clone());
562    }
563    None
564}