https://github.com/erichkeane created 
https://github.com/llvm/llvm-project/pull/225969

It isn't clear how we missed this, but classic codegen does checks this, so we 
should too.

Claude Helped diagnose/debug, but I did the copy/pasting :D

>From 6f82f899b8847d3d8c86cfcaba2c07de0c46e62f Mon Sep 17 00:00:00 2001
From: erichkeane <[email protected]>
Date: Wed, 23 Sep 2026 16:00:55 -0700
Subject: [PATCH] [CIR] Make sure ptr-cast-to-vbase is guarded.

It isn't clear how we missed this, but classic codegen does checks this,
so we should too.

Claude Helped diagnose/debug, but I did the copy/pasting :D
---
 clang/lib/CIR/CodeGen/CIRGenClass.cpp         | 38 +++++++++++++++--
 clang/test/CIR/CodeGen/vbase.cpp              | 28 +++++++++++++
 .../test/CIR/CodeGenCXX/virtual-base-cast.cpp | 42 +++++++++++++++----
 3 files changed, 96 insertions(+), 12 deletions(-)

diff --git a/clang/lib/CIR/CodeGen/CIRGenClass.cpp 
b/clang/lib/CIR/CodeGen/CIRGenClass.cpp
index c8edc1ccd02c6..28db1763fd3cb 100644
--- a/clang/lib/CIR/CodeGen/CIRGenClass.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenClass.cpp
@@ -1267,17 +1267,49 @@ Address CIRGenFunction::getAddressOfBaseClass(
 
   assert(!cir::MissingFeatures::sanitizers());
 
+  mlir::Location mlirLoc = getLoc(loc);
+
+  // Computing the virtual offset requires reading the vtable, which is only
+  // safe to do once we know the pointer isn't null. Guard the whole
+  // computation, mirroring classic CodeGen's cast.notnull/cast.end split.
+  if (vBase && nullCheckValue) {
+    CharUnits alignment =
+        cgm.getVBaseAlignment(value.getAlignment(), derived, vBase)
+            .alignmentAtOffset(nonVirtualOffset);
+    mlir::Type basePtrTy = builder.getPointerTo(baseValueTy);
+    mlir::Value ptrIsNull = builder.createPtrIsNull(value.getPointer());
+    mlir::Value result =
+        cir::TernaryOp::create(
+            builder, mlirLoc, ptrIsNull,
+            [&](mlir::OpBuilder &, mlir::Location) {
+              builder.createYield(
+                  mlirLoc, builder.getNullPtr(basePtrTy, mlirLoc).getResult());
+            },
+            [&](mlir::OpBuilder &, mlir::Location) {
+              mlir::Value virtualOffset =
+                  cgm.getCXXABI().getVirtualBaseClassOffset(
+                      mlirLoc, *this, value, derived, vBase);
+              Address adjusted = applyNonVirtualAndVirtualOffset(
+                  mlirLoc, *this, value, nonVirtualOffset, virtualOffset,
+                  derived, vBase, baseValueTy, /*assumeNotNull=*/true);
+              adjusted = adjusted.withElementType(builder, baseValueTy);
+              builder.createYield(mlirLoc, adjusted.getPointer());
+            })
+            .getResult();
+    return Address(result, baseValueTy, alignment);
+  }
+
   // Compute the virtual offset.
   mlir::Value virtualOffset = nullptr;
   if (vBase) {
     virtualOffset = cgm.getCXXABI().getVirtualBaseClassOffset(
-        getLoc(loc), *this, value, derived, vBase);
+        mlirLoc, *this, value, derived, vBase);
   }
 
   // Apply both offsets.
   value = applyNonVirtualAndVirtualOffset(
-      getLoc(loc), *this, value, nonVirtualOffset, virtualOffset, derived,
-      vBase, baseValueTy, not nullCheckValue);
+      mlirLoc, *this, value, nonVirtualOffset, virtualOffset, derived, vBase,
+      baseValueTy, not nullCheckValue);
 
   // Cast to the destination type.
   value = value.withElementType(builder, baseValueTy);
diff --git a/clang/test/CIR/CodeGen/vbase.cpp b/clang/test/CIR/CodeGen/vbase.cpp
index b480609af620d..262bc9fdc0110 100644
--- a/clang/test/CIR/CodeGen/vbase.cpp
+++ b/clang/test/CIR/CodeGen/vbase.cpp
@@ -138,3 +138,31 @@ void ppp() { B b; }
 // OGCG:   %[[BASE_A_ADDR:.*]] = getelementptr inbounds i8, ptr %[[THIS]], i64 
12
 // OGCG:   store ptr getelementptr inbounds inrange(-24, 0) (i8, ptr @_ZTV1B, 
i64 24), ptr %[[THIS]]
 // OGCG:   ret void
+
+// Pointer to virtual base must null-check.
+A *conv(B *p) { return p; }
+
+// CIR-LABEL: cir.func {{.*}} @_Z4convP1B(
+// CIR:   %[[P:.*]] = cir.load {{.*}} : !cir.ptr<!cir.ptr<!rec_B>>, 
!cir.ptr<!rec_B>
+// CIR:   %[[IS_NULL:.*]] = cir.cmp eq %[[P]], {{.*}} : !cir.ptr<!rec_B>
+// CIR:   cir.ternary(%[[IS_NULL]], true {
+// CIR:     %[[NULLPTR:.*]] = cir.const #cir.ptr<null> : !cir.ptr<!rec_A>
+// CIR:     cir.yield %[[NULLPTR]] : !cir.ptr<!rec_A>
+// CIR:   }, false {
+// CIR:     cir.vtable.get_vptr %[[P]]
+// CIR:     cir.yield {{.*}} : !cir.ptr<!rec_A>
+// CIR:   }) : (!cir.bool) -> !cir.ptr<!rec_A>
+
+// LLVM: define {{.*}} ptr @_Z4convP1B(
+// LLVM:   %[[P:.*]] = load ptr, ptr {{.*}}
+// LLVM:   %[[IS_NULL:.*]] = icmp eq ptr %[[P]], null
+// LLVM:   br i1 %[[IS_NULL]], label %{{.*}}, label %{{.*}}
+// LLVM:   load ptr, ptr %[[P]]
+// LLVM: phi ptr
+
+// OGCG: define {{.*}} ptr @_Z4convP1B(
+// OGCG:   %[[P:.*]] = load ptr, ptr {{.*}}
+// OGCG:   %[[IS_NULL:.*]] = icmp eq ptr %[[P]], null
+// OGCG:   br i1 %[[IS_NULL]], label %{{.*}}, label %{{.*}}
+// OGCG:   load ptr, ptr %[[P]]
+// OGCG: phi ptr
diff --git a/clang/test/CIR/CodeGenCXX/virtual-base-cast.cpp 
b/clang/test/CIR/CodeGenCXX/virtual-base-cast.cpp
index dc4bae2031b30..ebe05f2f535fc 100644
--- a/clang/test/CIR/CodeGenCXX/virtual-base-cast.cpp
+++ b/clang/test/CIR/CodeGenCXX/virtual-base-cast.cpp
@@ -15,24 +15,36 @@ D* x;
 
 // This uses the vtable to get the offset to the base object. The offset from
 // the vptr to the base object offset in the vtable is a compile-time constant.
+// Since computing that offset requires dereferencing the vtable pointer, the
+// whole computation is guarded by a null check on the source pointer.
 // CIR: %[[X_ADDR:.*]] = cir.get_global @x : !cir.ptr<!cir.ptr<!rec_D>>
 // CIR: %[[X:.*]] = cir.load{{.*}} %[[X_ADDR]]
-// CIR: %[[X_VPTR_ADDR:.*]] = cir.vtable.get_vptr %[[X]] : !cir.ptr<!rec_D> -> 
!cir.ptr<!cir.vptr>
-// CIR: %[[X_VPTR_BASE:.*]] = cir.load{{.*}} %[[X_VPTR_ADDR]] : 
!cir.ptr<!cir.vptr>, !cir.vptr
-// CIR: %[[X_BASE_I8PTR:.*]] = cir.cast bitcast %[[X_VPTR_BASE]] : !cir.vptr 
-> !cir.ptr<!u8i>
-// CIR:  %[[OFFSET_OFFSET:.*]] = cir.const #cir.int<-32> : !s64i
-// CIR:  %[[OFFSET_PTR:.*]] = cir.ptr_stride %[[X_BASE_I8PTR]], 
%[[OFFSET_OFFSET]] : (!cir.ptr<!u8i>, !s64i) -> !cir.ptr<!u8i>
-// CIR:  %[[OFFSET_PTR_CAST:.*]] = cir.cast bitcast %[[OFFSET_PTR]] : 
!cir.ptr<!u8i> -> !cir.ptr<!s64i>
-// CIR:  %[[OFFSET:.*]] = cir.load{{.*}} %[[OFFSET_PTR_CAST]] : 
!cir.ptr<!s64i>, !s64i
-// CIR:  %[[VBASE_ADDR:.*]] = cir.ptr_stride {{.*}}, %[[OFFSET]] : 
(!cir.ptr<!u8i>, !s64i) -> !cir.ptr<!u8i>
-// CIR:  cir.cast bitcast %[[VBASE_ADDR]] : !cir.ptr<!u8i> -> !cir.ptr<!rec_D>
+// CIR: %[[IS_NULL:.*]] = cir.cmp eq %[[X]], {{.*}} : !cir.ptr<!rec_D>
+// CIR: cir.ternary(%[[IS_NULL]], true {
+// CIR:   cir.const #cir.ptr<null> : !cir.ptr<!rec_A>
+// CIR:   cir.yield {{.*}} : !cir.ptr<!rec_A>
+// CIR: }, false {
+// CIR:   %[[X_VPTR_ADDR:.*]] = cir.vtable.get_vptr %[[X]] : !cir.ptr<!rec_D> 
-> !cir.ptr<!cir.vptr>
+// CIR:   %[[X_VPTR_BASE:.*]] = cir.load{{.*}} %[[X_VPTR_ADDR]] : 
!cir.ptr<!cir.vptr>, !cir.vptr
+// CIR:   %[[X_BASE_I8PTR:.*]] = cir.cast bitcast %[[X_VPTR_BASE]] : !cir.vptr 
-> !cir.ptr<!u8i>
+// CIR:   %[[OFFSET_OFFSET:.*]] = cir.const #cir.int<-32> : !s64i
+// CIR:   %[[OFFSET_PTR:.*]] = cir.ptr_stride %[[X_BASE_I8PTR]], 
%[[OFFSET_OFFSET]] : (!cir.ptr<!u8i>, !s64i) -> !cir.ptr<!u8i>
+// CIR:   %[[OFFSET_PTR_CAST:.*]] = cir.cast bitcast %[[OFFSET_PTR]] : 
!cir.ptr<!u8i> -> !cir.ptr<!s64i>
+// CIR:   %[[OFFSET:.*]] = cir.load{{.*}} %[[OFFSET_PTR_CAST]] : 
!cir.ptr<!s64i>, !s64i
+// CIR:   %[[VBASE_ADDR:.*]] = cir.ptr_stride {{.*}}, %[[OFFSET]] : 
(!cir.ptr<!u8i>, !s64i) -> !cir.ptr<!u8i>
+// CIR:   cir.cast bitcast %[[VBASE_ADDR]] : !cir.ptr<!u8i> -> !cir.ptr<!rec_D>
+// CIR: })
 
 // LLVM-LABEL: @_Z1av(
 // LLVM:       [[OBJ:%.*]] = load ptr, ptr @x
+// LLVM-NEXT:  [[IS_NULL:%.*]] = icmp eq ptr [[OBJ]], null
+// LLVM-NEXT:  br i1 [[IS_NULL]], label %[[NULL_BB:.*]], label 
%[[NOTNULL_BB:.*]]
+// LLVM:       [[NOTNULL_BB]]:
 // LLVM-NEXT:  [[VTABLE:%.*]] = load ptr, ptr [[OBJ]]
 // LLVM-NEXT:  [[VBASE_OFFSET_PTR:%.*]] = getelementptr i8, ptr [[VTABLE]], 
i64 -32
 // LLVM-NEXT:  [[VBASE_OFFSET:%.*]] = load i64, ptr [[VBASE_OFFSET_PTR]]
 // LLVM-NEXT:  [[ADD_PTR:%.*]] = getelementptr i8, ptr [[OBJ]], i64 
[[VBASE_OFFSET]]
+// LLVM:       phi ptr
 // LLVM:       ret ptr
 
 // OGCG-LABEL:  @_Z1av(
@@ -53,10 +65,14 @@ A* a() { return x; }
 
 // LLVM-LABEL: @_Z1bv(
 // LLVM:       [[OBJ:%.*]] = load ptr, ptr @x
+// LLVM-NEXT:  [[IS_NULL:%.*]] = icmp eq ptr [[OBJ]], null
+// LLVM-NEXT:  br i1 [[IS_NULL]], label %[[NULL_BB:.*]], label 
%[[NOTNULL_BB:.*]]
+// LLVM:       [[NOTNULL_BB]]:
 // LLVM-NEXT:  [[VTABLE:%.*]] = load ptr, ptr [[OBJ]]
 // LLVM-NEXT:  [[VBASE_OFFSET_PTR:%.*]] = getelementptr i8, ptr [[VTABLE]], 
i64 -40
 // LLVM-NEXT:  [[VBASE_OFFSET:%.*]] = load i64, ptr [[VBASE_OFFSET_PTR]]
 // LLVM-NEXT:  [[ADD_PTR:%.*]] = getelementptr i8, ptr [[OBJ]], i64 
[[VBASE_OFFSET]]
+// LLVM:       phi ptr
 // LLVM:       ret ptr
 
 // OGCG-LABEL:  @_Z1bv(
@@ -78,11 +94,15 @@ B* b() { return x; }
 
 // LLVM-LABEL: @_Z1cv(
 // LLVM:       [[OBJ:%.*]] = load ptr, ptr @x
+// LLVM-NEXT:  [[IS_NULL:%.*]] = icmp eq ptr [[OBJ]], null
+// LLVM-NEXT:  br i1 [[IS_NULL]], label %[[NULL_BB:.*]], label 
%[[NOTNULL_BB:.*]]
+// LLVM:       [[NOTNULL_BB]]:
 // LLVM-NEXT:  [[VTABLE:%.*]] = load ptr, ptr [[OBJ]]
 // LLVM-NEXT:  [[VBASE_OFFSET_PTR:%.*]] = getelementptr i8, ptr [[VTABLE]], 
i64 -48
 // LLVM-NEXT:  [[VBASE_OFFSET:%.*]] = load i64, ptr [[VBASE_OFFSET_PTR]]
 // LLVM-NEXT:  [[OFFSET:%.*]] = add i64 [[VBASE_OFFSET]], 16
 // LLVM-NEXT:  [[ADD_PTR:%.*]] = getelementptr i8, ptr [[OBJ]], i64 [[OFFSET]]
+// LLVM:       phi ptr
 // LLVM:       ret ptr
 
 // OGCG-LABEL:  @_Z1cv(
@@ -116,11 +136,15 @@ F* y;
 
 // LLVM-LABEL: @_Z1dv(
 // LLVM:       [[OBJ:%.*]] = load ptr, ptr @y
+// LLVM-NEXT:  [[IS_NULL:%.*]] = icmp eq ptr [[OBJ]], null
+// LLVM-NEXT:  br i1 [[IS_NULL]], label %[[NULL_BB:.*]], label 
%[[NOTNULL_BB:.*]]
+// LLVM:       [[NOTNULL_BB]]:
 // LLVM-NEXT:  [[VTABLE:%.*]] = load ptr, ptr [[OBJ]]
 // LLVM-NEXT:  [[VBASE_OFFSET_PTR:%.*]] = getelementptr i8, ptr [[VTABLE]], 
i64 -48
 // LLVM-NEXT:  [[VBASE_OFFSET:%.*]] = load i64, ptr [[VBASE_OFFSET_PTR]]
 // LLVM-NEXT:  [[OFFSET:%.*]] = add i64 [[VBASE_OFFSET]], 16
 // LLVM-NEXT:  [[ADD_PTR:%.*]] = getelementptr i8, ptr [[OBJ]], i64 [[OFFSET]]
+// LLVM:       phi ptr
 // LLVM:       ret ptr
 
 // OGCG-LABEL:  @_Z1dv(

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

Reply via email to