https://github.com/erichkeane updated 
https://github.com/llvm/llvm-project/pull/228267

>From 487ee56d9b1b83fb2e7ff2e4a4ddc53b782b8209 Mon Sep 17 00:00:00 2001
From: erichkeane <[email protected]>
Date: Thu, 1 Oct 2026 15:15:34 -0700
Subject: [PATCH] [CIR][CIRSimplify] Fix non-adjacent switch-case-merging

We previously assumed that all switch cases in a switch were adjacent/in
the same scope while merging switch cases. This wasn't true, since a
case can be inside of a scope.

Skip out on merging them not only when they have something between them,
but also when they aren't exactly 'neighbors'.

Note: Claude helped diagnose and write my test case.
---
 .../CIR/Dialect/Transforms/CIRSimplify.cpp    | 13 +++-
 clang/test/CIR/Transforms/switch-fold.cir     | 64 +++++++++++++++++++
 2 files changed, 74 insertions(+), 3 deletions(-)

diff --git a/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp 
b/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp
index 50bba6688c882..ad1e58e6ab5cd 100644
--- a/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp
+++ b/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp
@@ -311,9 +311,16 @@ struct SimplifySwitch : public OpRewritePattern<SwitchOp> {
     };
 
     for (CaseOp c : cases) {
-      if (!cascadingCases.empty() &&
-          !isa_and_nonnull<CaseOp>(c->getPrevNode())) {
-
+      // Cascading cases must be textually adjacent to the previously
+      // collected cascading case. This is false when something is in the
+      // way (e.g. a goto) or when the previous cascading case was found
+      // nested inside a sibling case's body (e.g. a case label that falls
+      // through into a compound statement) rather than next to `c`.
+      bool isAdjacentToLastCascadingCase =
+          !cascadingCases.empty() &&
+          c->getPrevNode() == cascadingCases.back().getOperation();
+
+      if (!cascadingCases.empty() && !isAdjacentToLastCascadingCase) {
         if (cascadingCases.size() > 1)
           mergeLastCascadingAndFlush();
         else
diff --git a/clang/test/CIR/Transforms/switch-fold.cir 
b/clang/test/CIR/Transforms/switch-fold.cir
index b76d982540eba..5101a303e12d8 100644
--- a/clang/test/CIR/Transforms/switch-fold.cir
+++ b/clang/test/CIR/Transforms/switch-fold.cir
@@ -302,4 +302,68 @@ module {
   //CHECK:     }
   //CHECK:     cir.yield
   //CHECK:   }
+
+  // A case label nested inside a sibling case's body (e.g. falling through
+  // into a compound statement, as with Duff's device) must not be folded
+  // into an unrelated later case just because that later case's *sibling*
+  // happens to be a cir.case.
+  //
+  //   switch (op) {
+  //   case 1: {
+  //   case 2:;
+  //   }
+  //     return 2;
+  //   case 3:
+  //     return 3;
+  //   }
+  cir.func @noFoldCascadeNestedInSiblingCase(%arg0: !s32i) -> !s32i {
+    %0 = cir.alloca "__retval" align(4) : !cir.ptr<!s32i>
+    cir.scope {
+      cir.switch (%arg0 : !s32i) {
+        cir.case(equal, [#cir.int<1> : !s32i]) {
+          cir.scope {
+            cir.case(equal, [#cir.int<2> : !s32i]) {
+              cir.yield
+            }
+          }
+          %1 = cir.const #cir.int<2> : !s32i
+          cir.store %1, %0 : !s32i, !cir.ptr<!s32i>
+          %2 = cir.load %0 : !cir.ptr<!s32i>, !s32i
+          cir.return %2 : !s32i
+        }
+        cir.case(equal, [#cir.int<3> : !s32i]) {
+          %1 = cir.const #cir.int<3> : !s32i
+          cir.store %1, %0 : !s32i, !cir.ptr<!s32i>
+          %2 = cir.load %0 : !cir.ptr<!s32i>, !s32i
+          cir.return %2 : !s32i
+        }
+        cir.yield
+      }
+    }
+    %1 = cir.const #cir.int<0> : !s32i
+    cir.store %1, %0 : !s32i, !cir.ptr<!s32i>
+    %2 = cir.load %0 : !cir.ptr<!s32i>, !s32i
+    cir.return %2 : !s32i
+  }
+  //CHECK: cir.func @noFoldCascadeNestedInSiblingCase
+  //CHECK:   cir.switch(%[[COND:.*]] : !s32i) {
+  //CHECK:     cir.case(equal, [#cir.int<1> : !s32i]) {
+  //CHECK:       cir.scope {
+  //CHECK:         cir.case(equal, [#cir.int<2> : !s32i]) {
+  //CHECK:           cir.yield
+  //CHECK:         }
+  //CHECK:       }
+  //CHECK:       %[[TWO:.*]] = cir.const #cir.int<2> : !s32i
+  //CHECK:       cir.store %[[TWO]], %{{.*}} : !s32i, !cir.ptr<!s32i>
+  //CHECK:       %[[RET2:.*]] = cir.load %{{.*}} : !cir.ptr<!s32i>, !s32i
+  //CHECK:       cir.return %[[RET2]] : !s32i
+  //CHECK:     }
+  //CHECK:     cir.case(equal, [#cir.int<3> : !s32i]) {
+  //CHECK:       %[[THREE:.*]] = cir.const #cir.int<3> : !s32i
+  //CHECK:       cir.store %[[THREE]], %{{.*}} : !s32i, !cir.ptr<!s32i>
+  //CHECK:       %[[RET3:.*]] = cir.load %{{.*}} : !cir.ptr<!s32i>, !s32i
+  //CHECK:       cir.return %[[RET3]] : !s32i
+  //CHECK:     }
+  //CHECK:     cir.yield
+  //CHECK:   }
 }

_______________________________________________
cfe-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits

Reply via email to