From 4e64b1948d9b34d1c1d66a6553f4510209b988d2 Mon Sep 17 00:00:00 2001 From: Richard Smith Date: Fri, 13 Oct 2023 14:29:50 -0700 Subject: [PATCH] Bare-bones support for forward-declared classes. (#3294) This is mostly scaffolding, but is just about enough for pointers to classes to work properly as types. --- toolchain/check/context.cpp | 6 ++ toolchain/check/handle_class.cpp | 53 ++++++++++++---- toolchain/check/node_stack.h | 13 ++++ .../testdata/basics/builtin_nodes.carbon | 2 + .../multifile_raw_and_textual_ir.carbon | 4 ++ .../testdata/basics/multifile_raw_ir.carbon | 4 ++ .../testdata/basics/raw_and_textual_ir.carbon | 2 + toolchain/check/testdata/basics/raw_ir.carbon | 2 + .../testdata/class/forward_declared.carbon | 22 +++++++ toolchain/lower/handle.cpp | 6 ++ toolchain/sem_ir/file.cpp | 15 +++++ toolchain/sem_ir/file.h | 32 ++++++++++ toolchain/sem_ir/file_test.cpp | 1 + toolchain/sem_ir/formatter.cpp | 61 ++++++++++++++++++- toolchain/sem_ir/node.h | 13 ++++ toolchain/sem_ir/node_kind.def | 2 + 16 files changed, 225 insertions(+), 13 deletions(-) create mode 100644 toolchain/check/testdata/class/forward_declared.carbon diff --git a/toolchain/check/context.cpp b/toolchain/check/context.cpp index 4ea1d4f218e3..58269a92a731 100644 --- a/toolchain/check/context.cpp +++ b/toolchain/check/context.cpp @@ -336,6 +336,10 @@ static auto ProfileType(Context& semantics_context, SemIR::Node node, case SemIR::Builtin::Kind: canonical_id.AddInteger(node.As().builtin_kind.AsInt()); break; + case SemIR::ClassDeclaration::Kind: + canonical_id.AddInteger( + node.As().class_id.index); + break; case SemIR::CrossReference::Kind: { // TODO: Cross-references should be canonicalized by looking at their // target rather than treating them as new unique types. @@ -385,6 +389,8 @@ auto Context::CanonicalizeTypeAndAddNodeIfNew(SemIR::Node node) } auto Context::CanonicalizeType(SemIR::NodeId node_id) -> SemIR::TypeId { + node_id = FollowNameReferences(node_id); + auto it = canonical_types_.find(node_id); if (it != canonical_types_.end()) { return it->second; diff --git a/toolchain/check/handle_class.cpp b/toolchain/check/handle_class.cpp index abd10bf24a86..62db93afacf3 100644 --- a/toolchain/check/handle_class.cpp +++ b/toolchain/check/handle_class.cpp @@ -6,21 +6,52 @@ namespace Carbon::Check { -auto HandleClassDeclaration(Context& context, Parse::Node parse_node) -> bool { - return context.TODO(parse_node, "HandleClassDeclaration"); +auto HandleClassIntroducer(Context& context, Parse::Node parse_node) -> bool { + // Create a node block to hold the nodes created as part of the class + // signature, such as generic parameters. + context.node_block_stack().Push(); + // Push the bracketing node. + context.node_stack().Push(parse_node); + // A name should always follow. + context.declaration_name_stack().Push(); + return true; +} + +static auto BuildClassDeclaration(Context& context) -> void { + auto name_context = context.declaration_name_stack().Pop(); + + auto class_keyword = + context.node_stack() + .PopForSoloParseNode(); + + // TODO: Track this somewhere. + context.node_block_stack().Pop(); + + auto class_id = context.semantics_ir().AddClass( + {.name_id = name_context.state == + DeclarationNameStack::NameContext::State::Unresolved + ? name_context.unresolved_name_id + : SemIR::StringId(SemIR::StringId::InvalidIndex)}); + auto class_decl_id = context.AddNode(SemIR::ClassDeclaration( + class_keyword, SemIR::TypeId::TypeType, class_id)); + context.declaration_name_stack().AddNameToLookup(name_context, class_decl_id); +} + +auto HandleClassDeclaration(Context& context, Parse::Node /*parse_node*/) + -> bool { + BuildClassDeclaration(context); + return true; +} + +auto HandleClassDefinitionStart(Context& context, Parse::Node parse_node) + -> bool { + BuildClassDeclaration(context); + // TODO: Introduce `Self`. + return context.TODO(parse_node, "HandleClassDefinitionStart"); } auto HandleClassDefinition(Context& context, Parse::Node parse_node) -> bool { return context.TODO(parse_node, "HandleClassDefinition"); } -auto HandleClassDefinitionStart(Context& context, Parse::Node parse_node) - -> bool { - return context.TODO(parse_node, "HandleClassDefinitionStart"); -} - -auto HandleClassIntroducer(Context& context, Parse::Node parse_node) -> bool { - return context.TODO(parse_node, "HandleClassIntroducer"); -} - } // namespace Carbon::Check diff --git a/toolchain/check/node_stack.h b/toolchain/check/node_stack.h index 99a354dd3b46..7c3e6f076844 100644 --- a/toolchain/check/node_stack.h +++ b/toolchain/check/node_stack.h @@ -108,6 +108,11 @@ class NodeStack { RequireParseKind(back.first); return back; } + if constexpr (RequiredIdKind == IdKind::ClassId) { + auto back = PopWithParseNode(); + RequireParseKind(back.first); + return back; + } if constexpr (RequiredIdKind == IdKind::StringId) { auto back = PopWithParseNode(); RequireParseKind(back.first); @@ -154,6 +159,9 @@ class NodeStack { if constexpr (RequiredIdKind == IdKind::FunctionId) { return back.id(); } + if constexpr (RequiredIdKind == IdKind::ClassId) { + return back.id(); + } if constexpr (RequiredIdKind == IdKind::StringId) { return back.id(); } @@ -177,6 +185,7 @@ class NodeStack { NodeId, NodeBlockId, FunctionId, + ClassId, StringId, TypeId, // No associated ID type. @@ -272,6 +281,7 @@ class NodeStack { case Parse::NodeKind::Name: return IdKind::StringId; case Parse::NodeKind::ArrayExpressionSemi: + case Parse::NodeKind::ClassIntroducer: case Parse::NodeKind::CodeBlockStart: case Parse::NodeKind::FunctionIntroducer: case Parse::NodeKind::IfStatementElse: @@ -302,6 +312,9 @@ class NodeStack { if constexpr (std::is_same_v) { return IdKind::FunctionId; } + if constexpr (std::is_same_v) { + return IdKind::ClassId; + } if constexpr (std::is_same_v) { return IdKind::StringId; } diff --git a/toolchain/check/testdata/basics/builtin_nodes.carbon b/toolchain/check/testdata/basics/builtin_nodes.carbon index e8d5febacdfa..672b526af9cc 100644 --- a/toolchain/check/testdata/basics/builtin_nodes.carbon +++ b/toolchain/check/testdata/basics/builtin_nodes.carbon @@ -11,6 +11,8 @@ // CHECK:STDOUT: - cross_reference_irs_size: 1 // CHECK:STDOUT: functions: [ // CHECK:STDOUT: ] +// CHECK:STDOUT: classes: [ +// CHECK:STDOUT: ] // CHECK:STDOUT: integers: [ // CHECK:STDOUT: ] // CHECK:STDOUT: reals: [ diff --git a/toolchain/check/testdata/basics/multifile_raw_and_textual_ir.carbon b/toolchain/check/testdata/basics/multifile_raw_and_textual_ir.carbon index f0dc38ac9001..ae2acb188164 100644 --- a/toolchain/check/testdata/basics/multifile_raw_and_textual_ir.carbon +++ b/toolchain/check/testdata/basics/multifile_raw_and_textual_ir.carbon @@ -20,6 +20,8 @@ fn B() {} // CHECK:STDOUT: functions: [ // CHECK:STDOUT: {name: str0, param_refs: block0, body: [block1]}, // CHECK:STDOUT: ] +// CHECK:STDOUT: classes: [ +// CHECK:STDOUT: ] // CHECK:STDOUT: integers: [ // CHECK:STDOUT: ] // CHECK:STDOUT: reals: [ @@ -61,6 +63,8 @@ fn B() {} // CHECK:STDOUT: functions: [ // CHECK:STDOUT: {name: str0, param_refs: block0, body: [block1]}, // CHECK:STDOUT: ] +// CHECK:STDOUT: classes: [ +// CHECK:STDOUT: ] // CHECK:STDOUT: integers: [ // CHECK:STDOUT: ] // CHECK:STDOUT: reals: [ diff --git a/toolchain/check/testdata/basics/multifile_raw_ir.carbon b/toolchain/check/testdata/basics/multifile_raw_ir.carbon index 2b9f8aefa7cf..afaaf0788130 100644 --- a/toolchain/check/testdata/basics/multifile_raw_ir.carbon +++ b/toolchain/check/testdata/basics/multifile_raw_ir.carbon @@ -20,6 +20,8 @@ fn B() {} // CHECK:STDOUT: functions: [ // CHECK:STDOUT: {name: str0, param_refs: block0, body: [block1]}, // CHECK:STDOUT: ] +// CHECK:STDOUT: classes: [ +// CHECK:STDOUT: ] // CHECK:STDOUT: integers: [ // CHECK:STDOUT: ] // CHECK:STDOUT: reals: [ @@ -52,6 +54,8 @@ fn B() {} // CHECK:STDOUT: functions: [ // CHECK:STDOUT: {name: str0, param_refs: block0, body: [block1]}, // CHECK:STDOUT: ] +// CHECK:STDOUT: classes: [ +// CHECK:STDOUT: ] // CHECK:STDOUT: integers: [ // CHECK:STDOUT: ] // CHECK:STDOUT: reals: [ diff --git a/toolchain/check/testdata/basics/raw_and_textual_ir.carbon b/toolchain/check/testdata/basics/raw_and_textual_ir.carbon index 75131165858c..1ca55903e988 100644 --- a/toolchain/check/testdata/basics/raw_and_textual_ir.carbon +++ b/toolchain/check/testdata/basics/raw_and_textual_ir.carbon @@ -18,6 +18,8 @@ fn Foo(n: i32) -> (i32, f64) { // CHECK:STDOUT: functions: [ // CHECK:STDOUT: {name: str0, param_refs: block1, return_type: type3, return_slot: node+4, body: [block4]}, // CHECK:STDOUT: ] +// CHECK:STDOUT: classes: [ +// CHECK:STDOUT: ] // CHECK:STDOUT: integers: [ // CHECK:STDOUT: 2, // CHECK:STDOUT: ] diff --git a/toolchain/check/testdata/basics/raw_ir.carbon b/toolchain/check/testdata/basics/raw_ir.carbon index 4d5e33886e96..aecf8ec13a0c 100644 --- a/toolchain/check/testdata/basics/raw_ir.carbon +++ b/toolchain/check/testdata/basics/raw_ir.carbon @@ -18,6 +18,8 @@ fn Foo(n: i32) -> (i32, f64) { // CHECK:STDOUT: functions: [ // CHECK:STDOUT: {name: str0, param_refs: block1, return_type: type3, return_slot: node+4, body: [block4]}, // CHECK:STDOUT: ] +// CHECK:STDOUT: classes: [ +// CHECK:STDOUT: ] // CHECK:STDOUT: integers: [ // CHECK:STDOUT: 2, // CHECK:STDOUT: ] diff --git a/toolchain/check/testdata/class/forward_declared.carbon b/toolchain/check/testdata/class/forward_declared.carbon new file mode 100644 index 000000000000..18fe8d9f1cff --- /dev/null +++ b/toolchain/check/testdata/class/forward_declared.carbon @@ -0,0 +1,22 @@ +// Part of the Carbon Language project, under the Apache License v2.0 with LLVM +// Exceptions. See /LICENSE for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +// AUTOUPDATE + +class Class; + +fn F(p: Class*) -> Class* { return p; } + +// CHECK:STDOUT: file "forward_declared.carbon" { +// CHECK:STDOUT: %Class: type = class_declaration @Class +// CHECK:STDOUT: %F: = fn_decl @F +// CHECK:STDOUT: } +// CHECK:STDOUT: +// CHECK:STDOUT: class @Class; +// CHECK:STDOUT: +// CHECK:STDOUT: fn @F(%p: Class*) -> Class* { +// CHECK:STDOUT: !entry: +// CHECK:STDOUT: %p.ref: Class* = name_reference "p", %p +// CHECK:STDOUT: return %p.ref +// CHECK:STDOUT: } diff --git a/toolchain/lower/handle.cpp b/toolchain/lower/handle.cpp index 6dcb6167e4ee..84fc819e28c0 100644 --- a/toolchain/lower/handle.cpp +++ b/toolchain/lower/handle.cpp @@ -163,6 +163,12 @@ auto HandleCall(FunctionContext& context, SemIR::NodeId node_id, } } +auto HandleClassDeclaration(FunctionContext& /*context*/, + SemIR::NodeId /*node_id*/, + SemIR::ClassDeclaration /*node*/) -> void { + // No action to perform. +} + auto HandleDereference(FunctionContext& context, SemIR::NodeId node_id, SemIR::Dereference node) -> void { context.SetLocal(node_id, context.GetLocal(node.pointer_id)); diff --git a/toolchain/sem_ir/file.cpp b/toolchain/sem_ir/file.cpp index 14fc1b30d4d9..f05336cb3cc0 100644 --- a/toolchain/sem_ir/file.cpp +++ b/toolchain/sem_ir/file.cpp @@ -150,6 +150,7 @@ auto File::Print(llvm::raw_ostream& out, bool include_builtins) const -> void { << "\n"; PrintList(out, "functions", functions_); + PrintList(out, "classes", classes_); // Integer values are APInts, and default to a signed print, but we currently // treat them as unsigned. PrintList(out, "integers", integers_, @@ -179,6 +180,7 @@ static auto GetTypePrecedence(NodeKind kind) -> int { switch (kind) { case ArrayType::Kind: case Builtin::Kind: + case ClassDeclaration::Kind: case StructType::Kind: case TupleType::Kind: return 0; @@ -285,6 +287,12 @@ auto File::StringifyType(TypeId type_id, bool in_type_context) const } break; } + case ClassDeclaration::Kind: { + auto class_name_id = + GetClass(node.As().class_id).name_id; + out << GetString(class_name_id); + break; + } case ConstType::Kind: { if (step.index == 0) { out << "const "; @@ -463,6 +471,7 @@ auto GetExpressionCategory(const File& file, NodeId node_id) case BindValue::Kind: case BlockArg::Kind: case BoolLiteral::Kind: + case ClassDeclaration::Kind: case ConstType::Kind: case IntegerLiteral::Kind: case Parameter::Kind: @@ -637,6 +646,12 @@ auto GetValueRepresentation(const File& file, TypeId type_id) return {.kind = ValueRepresentation::Pointer, .type = type_id}; } + case ClassDeclaration::Kind: { + // TODO: Pick the default value representation in a smarter way. + // TODO: Allow the value representation for a class to be customized. + return {.kind = ValueRepresentation::Pointer, .type = type_id}; + } + case Builtin::Kind: // clang warns on unhandled enum values; clang-tidy is incorrect here. // NOLINTNEXTLINE(bugprone-switch-missing-default-case) diff --git a/toolchain/sem_ir/file.h b/toolchain/sem_ir/file.h index 9a5f833be98d..896b8ae69c54 100644 --- a/toolchain/sem_ir/file.h +++ b/toolchain/sem_ir/file.h @@ -51,6 +51,17 @@ struct Function : public Printable { llvm::SmallVector body_block_ids; }; +// A class. +struct Class : public Printable { + auto Print(llvm::raw_ostream& out) const -> void { + out << "{name: " << name_id; + out << "}"; + } + + // The class name. + StringId name_id; +}; + // TODO: Replace this with a Rational type, per the design: // docs/design/expressions/literals.md struct Real : public Printable { @@ -121,6 +132,23 @@ class File : public Printable { return functions_[function_id.index]; } + // Adds a class, returning an ID to reference it. + auto AddClass(Class class_info) -> ClassId { + ClassId id(classes_.size()); + // TODO: Return failure on overflow instead of crashing. + CARBON_CHECK(id.index >= 0); + classes_.push_back(class_info); + return id; + } + + // Returns the requested class. + auto GetClass(ClassId class_id) const -> const Class& { + return classes_[class_id.index]; + } + + // Returns the requested class. + auto GetClass(ClassId class_id) -> Class& { return classes_[class_id.index]; } + // Adds an integer value, returning an ID to reference it. auto AddInteger(llvm::APInt integer) -> IntegerId { IntegerId id(integers_.size()); @@ -327,6 +355,7 @@ class File : public Printable { -> std::string; auto functions_size() const -> int { return functions_.size(); } + auto classes_size() const -> int { return classes_.size(); } auto nodes_size() const -> int { return nodes_.size(); } auto node_blocks_size() const -> int { return node_blocks_.size(); } @@ -375,6 +404,9 @@ class File : public Printable { // Storage for callable objects. llvm::SmallVector functions_; + // Storage for classes. + llvm::SmallVector classes_; + // Related IRs. There will always be at least 2 entries, the builtin IR (used // for references of builtins) followed by the current IR (used for references // crossing node blocks). diff --git a/toolchain/sem_ir/file_test.cpp b/toolchain/sem_ir/file_test.cpp index d411c3e15363..592e1dcd14a7 100644 --- a/toolchain/sem_ir/file_test.cpp +++ b/toolchain/sem_ir/file_test.cpp @@ -47,6 +47,7 @@ TEST(SemIRTest, YAML) { auto file = Yaml::Sequence(ElementsAre(Yaml::Mapping(ElementsAre( Pair("cross_reference_irs_size", "1"), Pair("functions", Yaml::Sequence(SizeIs(1))), + Pair("classes", Yaml::Sequence(SizeIs(0))), Pair("integers", Yaml::Sequence(ElementsAre("0"))), Pair("reals", Yaml::Sequence(IsEmpty())), Pair("strings", Yaml::Sequence(ElementsAre("F", "x"))), diff --git a/toolchain/sem_ir/formatter.cpp b/toolchain/sem_ir/formatter.cpp index 268cb5a309e2..bdd2cc0ec8e6 100644 --- a/toolchain/sem_ir/formatter.cpp +++ b/toolchain/sem_ir/formatter.cpp @@ -37,7 +37,8 @@ class NodeNamer { semantics_ir_(semantics_ir) { nodes.resize(semantics_ir.nodes_size()); labels.resize(semantics_ir.node_blocks_size()); - scopes.resize(1 + semantics_ir.functions_size()); + scopes.resize(1 + semantics_ir.functions_size() + + semantics_ir.classes_size()); // Build the package scope. GetScopeInfo(ScopeIndex::Package).name = @@ -74,11 +75,33 @@ class NodeNamer { AddBlockLabel(fn_scope, block_id); } } + + // Build each class scope. + for (int i : llvm::seq(semantics_ir.classes_size())) { + auto class_id = ClassId(i); + auto class_scope = GetScopeFor(class_id); + const auto& class_info = semantics_ir.GetClass(class_id); + // TODO: Provide a location for the class for use as a + // disambiguator. + auto class_loc = Parse::Node::Invalid; + GetScopeInfo(class_scope).name = globals.AllocateName( + *this, class_loc, + class_info.name_id.is_valid() + ? semantics_ir.GetString(class_info.name_id).str() + : ""); + // TODO: Handle names declared in the class scope. + } } // Returns the scope index corresponding to a function. auto GetScopeFor(FunctionId fn_id) -> ScopeIndex { - return static_cast(fn_id.index + 1); + return static_cast(1 + fn_id.index); + } + + // Returns the scope index corresponding to a class. + auto GetScopeFor(ClassId class_id) -> ScopeIndex { + return static_cast(1 + semantics_ir_.functions_size() + + class_id.index); } // Returns the IR name to use for a function. @@ -89,6 +112,14 @@ class NodeNamer { return GetScopeInfo(GetScopeFor(fn_id)).name.str(); } + // Returns the IR name to use for a class. + auto GetNameFor(ClassId class_id) -> llvm::StringRef { + if (!class_id.is_valid()) { + return "invalid"; + } + return GetScopeInfo(GetScopeFor(class_id)).name.str(); + } + // Returns the IR name to use for a node, when referenced from a given scope. auto GetNameFor(ScopeIndex scope_idx, NodeId node_id) -> std::string { if (!node_id.is_valid()) { @@ -393,6 +424,12 @@ class NodeNamer { .name_id); continue; } + case ClassDeclaration::Kind: { + add_node_name_id( + semantics_ir_.GetClass(node.As().class_id) + .name_id); + continue; + } case NameReference::Kind: { add_node_name( semantics_ir_.GetString(node.As().name_id).str() + @@ -458,11 +495,25 @@ class Formatter { } out_ << "}\n"; + for (int i : llvm::seq(semantics_ir_.classes_size())) { + FormatClass(ClassId(i)); + } + for (int i : llvm::seq(semantics_ir_.functions_size())) { FormatFunction(FunctionId(i)); } } + auto FormatClass(ClassId id) -> void { + const Class& class_info = semantics_ir_.GetClass(id); + + out_ << "\nclass "; + FormatClassName(id); + // TODO: Format class definitions. + (void)class_info; + out_ << ";\n"; + } + auto FormatFunction(FunctionId id) -> void { const Function& fn = semantics_ir_.GetFunction(id); @@ -728,6 +779,8 @@ class Formatter { auto FormatArg(FunctionId id) -> void { FormatFunctionName(id); } + auto FormatArg(ClassId id) -> void { FormatClassName(id); } + auto FormatArg(IntegerId id) -> void { semantics_ir_.GetInteger(id).print(out_, /*isSigned=*/false); } @@ -815,6 +868,10 @@ class Formatter { out_ << node_namer_.GetNameFor(id); } + auto FormatClassName(ClassId id) -> void { + out_ << node_namer_.GetNameFor(id); + } + auto FormatType(TypeId id) -> void { if (!id.is_valid()) { out_ << "invalid"; diff --git a/toolchain/sem_ir/node.h b/toolchain/sem_ir/node.h index 38c551661ef7..530dd9accd72 100644 --- a/toolchain/sem_ir/node.h +++ b/toolchain/sem_ir/node.h @@ -59,6 +59,15 @@ struct FunctionId : public IndexBase, public Printable { } }; +// The ID of a class. +struct ClassId : public IndexBase, public Printable { + using IndexBase::IndexBase; + auto Print(llvm::raw_ostream& out) const -> void { + out << "class"; + IndexBase::Print(out); + } +}; + // The ID of a cross-referenced IR. struct CrossReferenceIRId : public IndexBase, public Printable { @@ -308,6 +317,10 @@ struct Call { FunctionId function_id; }; +struct ClassDeclaration { + ClassId class_id; +}; + struct ConstType { TypeId inner_id; }; diff --git a/toolchain/sem_ir/node_kind.def b/toolchain/sem_ir/node_kind.def index a11756a1990e..a1c9deca9589 100644 --- a/toolchain/sem_ir/node_kind.def +++ b/toolchain/sem_ir/node_kind.def @@ -59,6 +59,8 @@ CARBON_SEMANTICS_NODE_KIND_IMPL(BranchIf, "br", None, TerminatorSequence) CARBON_SEMANTICS_NODE_KIND_IMPL(BranchWithArg, "br", None, Terminator) CARBON_SEMANTICS_NODE_KIND_IMPL(Builtin, "builtin", Typed, NotTerminator) CARBON_SEMANTICS_NODE_KIND_IMPL(Call, "call", Typed, NotTerminator) +CARBON_SEMANTICS_NODE_KIND_IMPL(ClassDeclaration, "class_declaration", Typed, + NotTerminator) CARBON_SEMANTICS_NODE_KIND_IMPL(ConstType, "const_type", Typed, NotTerminator) CARBON_SEMANTICS_NODE_KIND_IMPL(Dereference, "dereference", Typed, NotTerminator)