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
