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
14pub(super) struct MatchBranchSimplification;
16
17impl<'tcx> crate::MirPass<'tcx> for MatchBranchSimplification {
18 fn policy(&self, ctx: &crate::PassCtx<'_>) -> PassPolicy {
19 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 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 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 Rvalue::Use(Operand::Constant(Box::new(first_const.clone())), WithRetag::No),
77 ))))
78 } else {
79 None
80 }
81 }
82
83 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 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 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 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 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 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 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 return None;
275 }
276
277 for &(case, rvalue) in rvals.iter() {
278 match rvalue {
279 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 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 Some(StatementKind::Assign(Box::new((
311 dest,
312 Rvalue::Use(Operand::Copy(copy_src_place), WithRetag::No),
313 ))))
314 }
315
316 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 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 if index == 0
345 && dest.is_stable_offset()
347 && 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
357fn 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 _ => return false,
366 };
367 if targets.all_targets().contains(&switch_bb) {
369 return false;
370 }
371 if !targets.is_distinct() {
373 return false;
374 }
375 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 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 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 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
464fn 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 _ => 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
520fn 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 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
544fn 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
553fn 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}