[clang] [CIR] Make sure ptr-cast-to-vbase is guarded. (PR #225969)
Erich Keane via cfe-commits
cfe-commits at lists.llvm.org
Wed Sep 23 16:02:47 PDT 2026
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 <ekeane at nvidia.com>
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(
More information about the cfe-commits
mailing list