Support for initialization of classes with abstract base classes. (#6802)

When initializing `.base` in class initialization, use `partial Base` as
the destination type rather than `Base`. Treat `partial Base` as not
being abstract even when `Base` is.

Allow conversion from a `partial T` initializer to a `T` initializer.
Store the vptr while performing the conversion. Do not store the vptr
when performing a `partial T` initialization, only when performing a
non-partial `T` initialization.
This commit is contained in:
Richard Smith
2026-02-27 01:27:25 +00:00
committed by GitHub
parent be88dfd744
commit 41dd256d56
22 changed files with 717 additions and 233 deletions
+234 -48
View File
@@ -21,6 +21,7 @@
#include "toolchain/check/impl_lookup.h"
#include "toolchain/check/import_ref.h"
#include "toolchain/check/inst.h"
#include "toolchain/check/member_access.h"
#include "toolchain/check/operator.h"
#include "toolchain/check/pattern_match.h"
#include "toolchain/check/type.h"
@@ -33,6 +34,7 @@
#include "toolchain/sem_ir/generic.h"
#include "toolchain/sem_ir/ids.h"
#include "toolchain/sem_ir/inst.h"
#include "toolchain/sem_ir/inst_kind.h"
#include "toolchain/sem_ir/type.h"
#include "toolchain/sem_ir/type_info.h"
#include "toolchain/sem_ir/typed_insts.h"
@@ -206,8 +208,7 @@ static auto ConvertAggregateElement(
llvm::ArrayRef<SemIR::InstId> src_literal_elems,
ConversionTarget::Kind kind, SemIR::InstId target_id,
SemIR::TypeInstId target_elem_type_inst, PendingBlock* target_block,
size_t src_field_index, size_t target_field_index,
SemIR::ClassType* vtable_class_type = nullptr) -> SemIR::InstId {
size_t src_field_index, size_t target_field_index) -> SemIR::InstId {
auto src_elem_type =
context.types().GetTypeIdForTypeInstId(src_elem_type_inst);
auto target_elem_type =
@@ -235,7 +236,7 @@ static auto ConvertAggregateElement(
target.storage_id = MakeElementAccessInst<TargetAccessInstT>(
context, loc_id, target_id, target_elem_type, *target_block,
target_field_index);
return Convert(context, loc_id, src_elem_id, target, vtable_class_type);
return Convert(context, loc_id, src_elem_id, target);
}
// Performs a conversion from a tuple to an array type. This function only
@@ -473,6 +474,136 @@ static auto ConvertTupleToType(Context& context, SemIR::LocId loc_id,
return context.types().GetTypeInstId(tuple_type_id);
}
// Create a reference to the vtable pointer for a class. Returns None if the
// class has no vptr.
static auto CreateVtablePtrRef(Context& context, SemIR::LocId loc_id,
SemIR::ClassType vtable_class_type)
-> SemIR::InstId {
auto vtable_decl_id =
context.classes().Get(vtable_class_type.class_id).vtable_decl_id;
if (!vtable_decl_id.has_value()) {
return SemIR::InstId::None;
}
LoadImportRef(context, vtable_decl_id);
auto canonical_vtable_decl_id =
context.constant_values().GetConstantInstId(vtable_decl_id);
return AddInst<SemIR::VtablePtr>(
context, loc_id,
{.type_id = GetPointerType(context, SemIR::VtableType::TypeInstId),
.vtable_id = context.insts()
.GetAs<SemIR::VtableDecl>(canonical_vtable_decl_id)
.vtable_id,
.specific_id = vtable_class_type.specific_id});
}
// Returns whether the given expression performs in-place initialization (or is
// invalid).
static auto IsInPlaceInitializing(Context& context, SemIR::InstId result_id) {
auto category = SemIR::GetExprCategory(context.sem_ir(), result_id);
return category == SemIR::ExprCategory::InPlaceInitializing ||
(category == SemIR::ExprCategory::ReprInitializing &&
SemIR::InitRepr::ForType(context.sem_ir(),
context.insts().Get(result_id).type_id())
.kind == SemIR::InitRepr::InPlace) ||
category == SemIR::ExprCategory::Error;
}
// Returns the index of the vptr field in the given struct type fields, or
// None if there is no vptr field.
static auto GetVptrFieldIndex(llvm::ArrayRef<SemIR::StructTypeField> fields)
-> SemIR::ElementIndex {
// If the type introduces a vptr, it will always be the first field.
bool has_vptr =
!fields.empty() && fields.front().name_id == SemIR::NameId::Vptr;
return has_vptr ? SemIR::ElementIndex(0) : SemIR::ElementIndex::None;
}
// Builds a member access expression naming the vptr field of the given class
// object. This is analogous to what `PerformMemberAccess` for `NameId::Vptr`
// would return if the vptr could be found by name lookup.
static auto PerformVptrAccess(Context& context, SemIR::LocId loc_id,
SemIR::InstId class_ref_id) -> SemIR::InstId {
auto class_type_id = context.insts().Get(class_ref_id).type_id();
while (class_ref_id.has_value()) {
// The type of `ref_id` must be a class type.
if (class_type_id == SemIR::ErrorInst::TypeId) {
return SemIR::ErrorInst::InstId;
}
auto class_type = context.types().GetAs<SemIR::ClassType>(class_type_id);
auto& class_info = context.classes().Get(class_type.class_id);
// Get the object representation.
auto object_repr_id =
class_info.GetObjectRepr(context.sem_ir(), class_type.specific_id);
if (object_repr_id == SemIR::ErrorInst::TypeId) {
return SemIR::ErrorInst::InstId;
}
if (context.types().Is<SemIR::CustomLayoutType>(object_repr_id)) {
context.TODO(loc_id, "accessing vptr of custom layout class");
return SemIR::ErrorInst::InstId;
}
// Check to see if this class introduces the vptr.
auto repr_struct_type =
context.types().GetAs<SemIR::StructType>(object_repr_id);
auto repr_fields =
context.struct_type_fields().Get(repr_struct_type.fields_id);
if (auto vptr_field_index = GetVptrFieldIndex(repr_fields);
vptr_field_index.has_value()) {
return AddInst<SemIR::ClassElementAccess>(
context, loc_id,
{.type_id = context.types().GetTypeIdForTypeInstId(
repr_fields[vptr_field_index.index].type_inst_id),
.base_id = class_ref_id,
.index = vptr_field_index});
}
// Otherwise, step through to the base class and try again.
CARBON_CHECK(class_info.base_id.has_value(),
"Could not find vptr for dynamic class");
auto base_decl = context.insts().GetAs<SemIR::BaseDecl>(class_info.base_id);
class_type_id = context.types().GetTypeIdForTypeInstId(
repr_fields[base_decl.index.index].type_inst_id);
class_ref_id =
AddInst<SemIR::ClassElementAccess>(context, loc_id,
{.type_id = class_type_id,
.base_id = class_ref_id,
.index = base_decl.index});
}
return class_ref_id;
}
// Converts an initializer for a type `partial T` to an initializer for `T` by
// initializing the vptr if necessary.
static auto ConvertPartialInitializerToNonPartial(
Context& context, ConversionTarget target,
SemIR::ClassType vtable_class_type, SemIR::InstId result_id)
-> SemIR::InstId {
auto loc_id = SemIR::LocId(result_id);
auto vptr_id = CreateVtablePtrRef(context, loc_id, vtable_class_type);
if (!vptr_id.has_value()) {
// No vtable pointer in this class, nothing to do.
return result_id;
}
CARBON_CHECK(
IsInPlaceInitializing(context, result_id),
"Type with vptr should have in-place initializing representation");
target.storage_access_block->InsertHere();
auto dest_id = PerformVptrAccess(context, loc_id, target.storage_id);
auto vptr_init_id = AddInst<SemIR::InPlaceInit>(
context, loc_id,
{.type_id = context.insts().Get(dest_id).type_id(),
.src_id = vptr_id,
.dest_id = dest_id});
return AddInst<SemIR::UpdateInit>(context, loc_id,
{.type_id = target.type_id,
.base_init_id = result_id,
.update_init_id = vptr_init_id});
}
// Common implementation for ConvertStructToStruct and ConvertStructToClass.
template <typename TargetAccessInstT>
static auto ConvertStructToStructOrClass(
@@ -487,10 +618,9 @@ static auto ConvertStructToStructOrClass(
auto& sem_ir = context.sem_ir();
auto src_elem_fields = sem_ir.struct_type_fields().Get(src_type.fields_id);
auto dest_elem_fields = sem_ir.struct_type_fields().Get(dest_type.fields_id);
bool dest_has_vptr = !dest_elem_fields.empty() &&
dest_elem_fields.front().name_id == SemIR::NameId::Vptr;
int dest_vptr_offset = (dest_has_vptr ? 1 : 0);
auto dest_elem_fields_size = dest_elem_fields.size() - dest_vptr_offset;
auto dest_vptr_index = GetVptrFieldIndex(dest_elem_fields);
auto dest_elem_fields_size =
dest_elem_fields.size() - (dest_vptr_index.has_value() ? 1 : 0);
auto value = sem_ir.insts().Get(value_id);
SemIR::LocId value_loc_id(value_id);
@@ -541,7 +671,7 @@ static auto ConvertStructToStructOrClass(
// of the source.
// TODO: Annotate diagnostics coming from here with the element index.
auto new_block =
literal_elems_id.has_value() && !dest_has_vptr
literal_elems_id.has_value() && !dest_vptr_index.has_value()
? SemIR::CopyOnWriteInstBlock(&sem_ir, literal_elems_id)
: SemIR::CopyOnWriteInstBlock(
&sem_ir, SemIR::CopyOnWriteInstBlock::UninitializedBlock{
@@ -559,18 +689,20 @@ static auto ConvertStructToStructOrClass(
{.type_id = vptr_type_id,
.base_id = target.storage_id,
.index = SemIR::ElementIndex(i)});
auto vtable_decl_id =
context.classes().Get(vtable_class_type->class_id).vtable_decl_id;
LoadImportRef(context, vtable_decl_id);
auto canonical_vtable_decl_id =
context.constant_values().GetConstantInstId(vtable_decl_id);
auto vtable_ptr_id = AddInst<SemIR::VtablePtr>(
context, value_loc_id,
{.type_id = GetPointerType(context, SemIR::VtableType::TypeInstId),
.vtable_id = context.insts()
.GetAs<SemIR::VtableDecl>(canonical_vtable_decl_id)
.vtable_id,
.specific_id = vtable_class_type->specific_id});
auto vtable_ptr_id = SemIR::InstId::None;
if (vtable_class_type) {
vtable_ptr_id =
CreateVtablePtrRef(context, value_loc_id, *vtable_class_type);
// Track that we initialized the vptr so we don't do it again.
vtable_class_type = nullptr;
} else {
// For a partial class type, we leave the vtable pointer uninitialized.
// TODO: Consider storing a specified value such as null for hardening.
vtable_ptr_id = AddInst<SemIR::UninitializedValue>(
context, value_loc_id,
{.type_id =
GetPointerType(context, SemIR::VtableType::TypeInstId)});
}
auto init_id = AddInst<SemIR::InPlaceInit>(context, value_loc_id,
{.type_id = vptr_type_id,
.src_id = vtable_ptr_id,
@@ -609,18 +741,48 @@ static auto ConvertStructToStructOrClass(
}
auto src_field = src_elem_fields[src_field_index];
// When initializing the `.base` field of a class, the destination type is
// `partial Base`, not `Base`.
// TODO: Skip this if the source field is an initializing expression of the
// non-partial type in order to produce smaller IR.
auto dest_field_type_inst_id = dest_field.type_inst_id;
if (dest_field.name_id == SemIR::NameId::Base) {
auto partial_type_id = GetQualifiedType(
context,
context.types().GetTypeIdForTypeInstId(dest_field.type_inst_id),
SemIR::TypeQualifiers::Partial);
dest_field_type_inst_id = context.types().GetTypeInstId(partial_type_id);
}
// TODO: This call recurses back into conversion. Switch to an iterative
// approach.
auto dest_field_index = src_field_index;
if (dest_vptr_index.has_value() &&
static_cast<int32_t>(src_field_index) >= dest_vptr_index.index) {
dest_field_index += 1;
}
auto init_id =
ConvertAggregateElement<SemIR::StructAccess, TargetAccessInstT>(
context, value_loc_id, value_id, src_field.type_inst_id,
literal_elems, inner_kind, target.storage_id,
dest_field.type_inst_id, target.storage_access_block,
src_field_index, src_field_index + dest_vptr_offset,
vtable_class_type);
dest_field_type_inst_id, target.storage_access_block,
src_field_index, dest_field_index);
if (init_id == SemIR::ErrorInst::InstId) {
return SemIR::ErrorInst::InstId;
}
// When initializing the base, adjust the type of the initializer from
// `partial Base` to `Base`. This isn't strictly correct, since we haven't
// finished initializing a `Base` until we store to the vptr, but is better
// than having an inconsistent type for the struct field initializer.
if (dest_field_type_inst_id != dest_field.type_inst_id) {
init_id = AddInst<SemIR::AsCompatible>(
context, value_loc_id,
{.type_id =
context.types().GetTypeIdForTypeInstId(dest_field.type_inst_id),
.source_id = init_id});
}
new_block.Set(i, init_id);
}
@@ -629,10 +791,15 @@ static auto ConvertStructToStructOrClass(
target.storage_access_block->InsertHere();
CARBON_CHECK(is_init,
"Converting directly to a class value is not supported");
return AddInst<SemIR::ClassInit>(context, value_loc_id,
{.type_id = target.type_id,
.elements_id = new_block.id(),
.dest_id = target.storage_id});
auto result_id = AddInst<SemIR::ClassInit>(context, value_loc_id,
{.type_id = target.type_id,
.elements_id = new_block.id(),
.dest_id = target.storage_id});
if (vtable_class_type) {
result_id = ConvertPartialInitializerToNonPartial(
context, target, *vtable_class_type, result_id);
}
return result_id;
} else if (is_init) {
target.storage_access_block->InsertHere();
return AddInst<SemIR::StructInit>(context, value_loc_id,
@@ -664,11 +831,11 @@ static auto ConvertStructToClass(Context& context, SemIR::StructType src_type,
SemIR::ClassType dest_type,
SemIR::InstId value_id,
ConversionTarget target,
SemIR::ClassType* vtable_class_type)
-> SemIR::InstId {
bool is_partial = false) -> SemIR::InstId {
PendingBlock target_block(&context);
auto& dest_class_info = context.classes().Get(dest_type.class_id);
CARBON_CHECK(dest_class_info.inheritance_kind != SemIR::Class::Abstract);
CARBON_CHECK(is_partial ||
dest_class_info.inheritance_kind != SemIR::Class::Abstract);
auto object_repr_id =
dest_class_info.GetObjectRepr(context.sem_ir(), dest_type.specific_id);
if (object_repr_id == SemIR::ErrorInst::TypeId) {
@@ -693,7 +860,7 @@ static auto ConvertStructToClass(Context& context, SemIR::StructType src_type,
auto result_id = ConvertStructToStructOrClass<SemIR::ClassElementAccess>(
context, src_type, dest_struct_type, value_id, target,
vtable_class_type ? vtable_class_type : &dest_type);
is_partial ? nullptr : &dest_type);
if (need_temporary) {
target_block.InsertHere();
@@ -894,11 +1061,11 @@ static auto CanRemoveQualifiers(SemIR::TypeQualifiers quals,
return false;
}
if (quals.HasAnyOf(SemIR::TypeQualifiers::Partial) &&
(!allow_unsafe || SemIR::IsInitializerCategory(cat))) {
// TODO: Allow removing `partial` for initializing expressions as a safe
// conversion. `PerformBuiltinConversion` will need to initialize the vptr
// as part of the conversion.
if (quals.HasAnyOf(SemIR::TypeQualifiers::Partial) && !allow_unsafe &&
!SemIR::IsInitializerCategory(cat)) {
// Removing `partial` is an unsafe conversion for a non-initializing
// expression. But it's OK for an initializing expression because we will
// initialize the vptr as part of the conversion.
return false;
}
@@ -939,10 +1106,9 @@ static auto DiagnoseConversionFailureToConstraintValue(
}
}
static auto PerformBuiltinConversion(
Context& context, SemIR::LocId loc_id, SemIR::InstId value_id,
ConversionTarget target, SemIR::ClassType* vtable_class_type = nullptr)
-> SemIR::InstId {
static auto PerformBuiltinConversion(Context& context, SemIR::LocId loc_id,
SemIR::InstId value_id,
ConversionTarget target) -> SemIR::InstId {
auto& sem_ir = context.sem_ir();
auto value = sem_ir.insts().Get(value_id);
auto value_type_id = value.type_id();
@@ -1112,6 +1278,19 @@ static auto PerformBuiltinConversion(
}
}
if ((removed_quals & SemIR::TypeQualifiers::Partial) !=
SemIR::TypeQualifiers::None &&
SemIR::IsInitializerCategory(category)) {
auto unqual_target_type_id =
context.types().GetUnqualifiedType(target.type_id);
if (auto target_class_type =
context.types().TryGetAs<SemIR::ClassType>(
unqual_target_type_id)) {
value_id = ConvertPartialInitializerToNonPartial(
context, target, *target_class_type, value_id);
}
}
value_id = AddInst<SemIR::AsCompatible>(
context, loc_id,
{.type_id = target.type_id, .source_id = value_id});
@@ -1167,19 +1346,28 @@ static auto PerformBuiltinConversion(
}
}
// Split the qualifiers off the target type.
// TODO: Most conversions should probably be looking at the unqualified target
// type.
auto [target_unqual_type_id, target_quals] =
context.types().GetUnqualifiedTypeAndQualifiers(target.type_id);
auto target_unqual_type_inst =
sem_ir.types().GetAsInst(target_unqual_type_id);
// A struct {.f_1: T_1, .f_2: T_2, ..., .f_n: T_n} converts to a class type
// if it converts to the struct type that is the class's representation type
// (a struct with the same fields as the class, plus a base field where
// relevant).
if (auto target_class_type = target_type_inst.TryAs<SemIR::ClassType>()) {
if (auto target_class_type =
target_unqual_type_inst.TryAs<SemIR::ClassType>()) {
if (auto src_struct_type =
sem_ir.types().TryGetAs<SemIR::StructType>(value_type_id)) {
if (!context.classes()
.Get(target_class_type->class_id)
.adapt_id.has_value()) {
return ConvertStructToClass(context, *src_struct_type,
*target_class_type, value_id, target,
vtable_class_type);
return ConvertStructToClass(
context, *src_struct_type, *target_class_type, value_id, target,
target_quals.HasAnyOf(SemIR::TypeQualifiers::Partial));
}
}
@@ -1697,8 +1885,7 @@ auto CategoryConverter::DoStep(const SemIR::InstId expr_id,
}
auto Convert(Context& context, SemIR::LocId loc_id, SemIR::InstId expr_id,
ConversionTarget target, SemIR::ClassType* vtable_class_type)
-> SemIR::InstId {
ConversionTarget target) -> SemIR::InstId {
auto& sem_ir = context.sem_ir();
auto orig_expr_id = expr_id;
@@ -1798,8 +1985,7 @@ auto Convert(Context& context, SemIR::LocId loc_id, SemIR::InstId expr_id,
TryToCompleteType(context, context.insts().Get(expr_id).type_id(), loc_id);
// Check whether any builtin conversion applies.
expr_id = PerformBuiltinConversion(context, loc_id, expr_id, target,
vtable_class_type);
expr_id = PerformBuiltinConversion(context, loc_id, expr_id, target);
if (expr_id == SemIR::ErrorInst::InstId) {
return expr_id;
}