From 782bd873162982550fe5e7142d9d3de9da95aa34 Mon Sep 17 00:00:00 2001 From: Richard Smith Date: Tue, 21 Mar 2023 12:36:05 -0700 Subject: [PATCH] AST cloning mechanism (#2699) This is intended to be used for template instantiation, but for now takes no stance as to what it's cloning. This is tested by parsing and cloning all of our test files, and checking that the result of converting each AST to proto is the same as the result of cloning and then converting each AST to proto. --- common/fuzzing/carbon.proto | 7 +- common/fuzzing/proto_to_carbon.cpp | 9 ++ explorer/ast/BUILD | 63 +++++----- explorer/ast/ast_node.h | 14 +++ explorer/ast/bindings.cpp | 10 ++ explorer/ast/bindings.h | 7 +- explorer/ast/clone_context.cpp | 92 ++++++++++++++ explorer/ast/clone_context.h | 166 ++++++++++++++++++++++++++ explorer/ast/declaration.cpp | 59 +++++++-- explorer/ast/declaration.h | 119 ++++++++++++++++-- explorer/ast/expression.cpp | 36 ++++++ explorer/ast/expression.h | 164 ++++++++++++++++++++++++- explorer/ast/impl_binding.cpp | 18 +++ explorer/ast/impl_binding.h | 5 +- explorer/ast/pattern.cpp | 12 ++ explorer/ast/pattern.h | 39 +++++- explorer/ast/return_term.h | 7 ++ explorer/ast/statement.cpp | 4 + explorer/ast/statement.h | 94 ++++++++++++++- explorer/ast/value.h | 31 +++++ explorer/ast/value_node.h | 10 ++ explorer/ast/value_transform.h | 136 +++++++++++++-------- explorer/common/arena.h | 32 +++++ explorer/fuzzing/BUILD | 22 ++++ explorer/fuzzing/ast_to_proto.cpp | 3 + explorer/fuzzing/clone_test.cpp | 73 +++++++++++ explorer/gen_rtti.py | 28 ++++- explorer/interpreter/type_checker.cpp | 2 +- 28 files changed, 1144 insertions(+), 118 deletions(-) create mode 100644 explorer/ast/clone_context.cpp create mode 100644 explorer/ast/clone_context.h create mode 100644 explorer/ast/impl_binding.cpp create mode 100644 explorer/fuzzing/clone_test.cpp diff --git a/common/fuzzing/carbon.proto b/common/fuzzing/carbon.proto index d691e5703763..6025b7990b53 100644 --- a/common/fuzzing/carbon.proto +++ b/common/fuzzing/carbon.proto @@ -430,9 +430,10 @@ message ImplDeclaration { } optional ImplKind kind = 1; - optional Expression impl_type = 2; - optional Expression interface = 3; - repeated Declaration members = 4; + repeated GenericBinding deduced_parameters = 2; + optional Expression impl_type = 3; + optional Expression interface = 4; + repeated Declaration members = 5; } message MatchFirstDeclaration { diff --git a/common/fuzzing/proto_to_carbon.cpp b/common/fuzzing/proto_to_carbon.cpp index b8b84aebf87c..da63c6475dc9 100644 --- a/common/fuzzing/proto_to_carbon.cpp +++ b/common/fuzzing/proto_to_carbon.cpp @@ -939,6 +939,15 @@ static auto DeclarationToCarbon(const Fuzzing::Declaration& declaration, out << "external "; } out << "impl "; + if (!impl.deduced_parameters().empty()) { + out << "forall ["; + llvm::ListSeparator sep; + for (const Fuzzing::GenericBinding& p : impl.deduced_parameters()) { + out << sep; + GenericBindingToCarbon(p, out); + } + out << "]"; + } ExpressionToCarbon(impl.impl_type(), out); out << " as "; ExpressionToCarbon(impl.interface(), out); diff --git a/explorer/ast/BUILD b/explorer/ast/BUILD index 1ce82ed6dfd8..5c27faf06e41 100644 --- a/explorer/ast/BUILD +++ b/explorer/ast/BUILD @@ -4,9 +4,28 @@ package(default_visibility = ["//explorer:__subpackages__"]) +AST_HDRS = [ + "address.h", + "ast.h", + "ast_node.h", + "bindings.h", + "clone_context.h", + "declaration.h", + "element.h", + "element_path.h", + "expression.h", + "impl_binding.h", + "pattern.h", + "return_term.h", + "statement.h", + "value.h", + "value_node.h", + "value_transform.h", +] + genrule( name = "ast_rtti", - srcs = ["ast_rtti.txt"], + srcs = ["ast_rtti.txt"] + AST_HDRS, outs = [ "ast_rtti.h", "ast_rtti.cpp", @@ -14,58 +33,31 @@ genrule( cmd = "./$(location //explorer:gen_rtti)" + " $(location ast_rtti.txt)" + " $(location ast_rtti.h) $(location ast_rtti.cpp)" + - " $(rootpath ast_rtti.h)", + " $(rootpath ast_rtti.h)" + + "".join([" $(rootpath " + f + ")" for f in AST_HDRS]), tools = ["//explorer:gen_rtti"], ) -cc_library( - name = "ast_node", - srcs = [ - "ast_node.cpp", - "ast_rtti.cpp", - ], - hdrs = [ - "ast_node.h", - "ast_rtti.h", - ], - deps = [ - "//explorer/common:source_location", - "@llvm-project//llvm:Support", - ], -) - cc_library( name = "ast", srcs = [ + "ast_node.cpp", + "ast_rtti.cpp", "bindings.cpp", + "clone_context.cpp", "declaration.cpp", "element.cpp", "expression.cpp", + "impl_binding.cpp", "pattern.cpp", "statement.cpp", "value.cpp", ], - hdrs = [ - "address.h", - "ast.h", - "bindings.h", - "declaration.h", - "element.h", - "element_path.h", - "expression.h", - "impl_binding.h", - "pattern.h", - "return_term.h", - "statement.h", - "value.h", - "value_node.h", - "value_transform.h", - ], + hdrs = AST_HDRS + ["ast_rtti.h"], textual_hdrs = [ "value_kinds.def", ], deps = [ - ":ast_node", ":library_name", ":paren_contents", ":value_category", @@ -91,7 +83,6 @@ cc_library( hdrs = ["ast_test_matchers.h"], deps = [ ":ast", - ":ast_node", "@com_google_googletest//:gtest", "@llvm-project//llvm:Support", ], diff --git a/explorer/ast/ast_node.h b/explorer/ast/ast_node.h index 4d3865b14b80..b36040539611 100644 --- a/explorer/ast/ast_node.h +++ b/explorer/ast/ast_node.h @@ -11,6 +11,8 @@ namespace Carbon { +class CloneContext; + // Base class for all nodes in the AST. // // Every class derived from this class must be listed in ast_rtti.txt. See @@ -35,6 +37,14 @@ namespace Carbon { // The definitions of `InheritsFromFoo` and `FooKind` are generated from // ast_rtti.txt, and are implicitly provided by this header. // +// Every AST node is expected to provide a cloning constructor: +// +// explicit MyAstNode(CloneContext& context, const MyAstNode& other); +// +// The cloning constructor should behave like a copy constructor, but pointers +// to other AST nodes should be passed through context.Clone to clone the +// referenced object. +// // TODO: To support generic traversal, add children() method, and ensure that // all AstNodes are reachable from a root AstNode. class AstNode { @@ -65,6 +75,10 @@ class AstNode { explicit AstNode(AstNodeKind kind, SourceLocation source_loc) : kind_(kind), source_loc_(source_loc) {} + // Clone this AstNode. + explicit AstNode(CloneContext& /*context*/, const AstNode& other) + : kind_(other.kind_), source_loc_(other.source_loc_) {} + // Equivalent to kind(), but will not be hidden by `kind()` methods of // derived classes. auto root_kind() const -> AstNodeKind { return kind_; } diff --git a/explorer/ast/bindings.cpp b/explorer/ast/bindings.cpp index ade2a9c9487f..44082301d77a 100644 --- a/explorer/ast/bindings.cpp +++ b/explorer/ast/bindings.cpp @@ -7,9 +7,19 @@ #include "common/error.h" #include "explorer/ast/impl_binding.h" #include "explorer/ast/pattern.h" +#include "explorer/ast/value.h" namespace Carbon { +Bindings::Bindings(CloneContext& context, const Bindings& other) { + for (auto [binding, value] : other.args_) { + args_.insert({context.Remap(binding), context.Clone(value)}); + } + for (auto [binding, value] : other.witnesses_) { + witnesses_.insert({context.Remap(binding), context.Clone(value)}); + } +} + void Bindings::Add(Nonnull binding, Nonnull value, std::optional> witness) { diff --git a/explorer/ast/bindings.h b/explorer/ast/bindings.h index 026e13d0c6f6..d4d5e62d98a7 100644 --- a/explorer/ast/bindings.h +++ b/explorer/ast/bindings.h @@ -8,6 +8,7 @@ #include #include +#include "explorer/ast/clone_context.h" #include "explorer/common/nonnull.h" #include "llvm/ADT/ArrayRef.h" @@ -45,16 +46,18 @@ class Bindings { // Create an instantiated set of bindings for use during evaluation, // containing both arguments and witnesses. - Bindings(BindingMap args, ImplWitnessMap witnesses) + explicit Bindings(BindingMap args, ImplWitnessMap witnesses) : args_(std::move(args)), witnesses_(std::move(witnesses)) {} enum NoWitnessesTag { NoWitnesses }; // Create a set of bindings for use during type-checking, containing only the // arguments but not the corresponding witnesses. - Bindings(BindingMap args, NoWitnessesTag /*unused*/) + explicit Bindings(BindingMap args, NoWitnessesTag /*unused*/) : args_(std::move(args)) {} + explicit Bindings(CloneContext& context, const Bindings& other); + template auto Decompose(F f) const { return f(args_, witnesses_); diff --git a/explorer/ast/clone_context.cpp b/explorer/ast/clone_context.cpp new file mode 100644 index 000000000000..5f2a33398c98 --- /dev/null +++ b/explorer/ast/clone_context.cpp @@ -0,0 +1,92 @@ +// 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 + +#include "explorer/ast/clone_context.h" + +#include "explorer/ast/ast_node.h" +#include "explorer/ast/value_transform.h" + +namespace Carbon { + +auto CloneContext::CloneBase(Nonnull node) + -> Nonnull { + auto [it, added] = nodes_.insert({node, nullptr}); + CARBON_CHECK(added) << (it->second + ? "node was cloned multiple times: " + : "node was remapped before it was cloned: ") + << *node; + + // The implementation is generated in ast_rtti.cpp. + CloneImpl(*arena_, *this, *node, &it->second); + + // Cloning may have invalidated our iterator; redo lookup. + auto* result = nodes_[node]; + CARBON_CHECK(result) << "CloneImpl didn't set the result pointer"; + return result; +} + +class CloneContext::CloneValueTransform + : public ValueTransform { + public: + CloneValueTransform(Nonnull context, Nonnull arena) + : ValueTransform(arena), context_(context) {} + + using ValueTransform::operator(); + + // Transforming a pointer to an AstNode should remap the node. Values do not + // own the nodes they point to, apart from the exceptions handled below. + template + auto operator()(Nonnull node, int /*unused*/ = 0) + -> std::enable_if_t, + Nonnull> { + return context_->Remap(node); + } + + // Transforming a value node view should clone it. The value node view does + // not itself own the node it points to, so this is a shallow clone. + auto operator()(ValueNodeView value_node) -> ValueNodeView { + return context_->Clone(value_node); + } + + // A FunctionType may or may not own its bindings. + auto operator()(Nonnull fn_type) + -> Nonnull { + for (auto* binding : fn_type->deduced_bindings()) { + context_->MaybeCloneBase(binding); + } + for (auto [index, binding] : fn_type->generic_parameters()) { + context_->MaybeCloneBase(binding); + } + return ValueTransform::operator()(fn_type); + } + + // A ConstraintType owns its self binding, so we need to clone it. + auto operator()(Nonnull constraint) + -> Nonnull { + context_->Clone(constraint->self_binding()); + return ValueTransform::operator()(constraint); + } + + private: + Nonnull context_; +}; + +auto CloneContext::CloneBase(Nonnull value) -> Nonnull { + return const_cast(CloneValueTransform(this, arena_).Transform(value)); +} + +auto CloneContext::CloneBase(Nonnull elem) + -> Nonnull { + return const_cast( + CloneValueTransform(this, arena_).Transform(elem)); +} + +void CloneContext::MaybeCloneBase(Nonnull node) { + auto it = nodes_.find(node); + if (it == nodes_.end()) { + Clone(node); + } +} + +} // namespace Carbon diff --git a/explorer/ast/clone_context.h b/explorer/ast/clone_context.h new file mode 100644 index 000000000000..8a6eb0fa1bff --- /dev/null +++ b/explorer/ast/clone_context.h @@ -0,0 +1,166 @@ +// 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 + +#ifndef CARBON_EXPLORER_AST_CLONE_CONTEXT_H_ +#define CARBON_EXPLORER_AST_CLONE_CONTEXT_H_ + +#include +#include +#include + +#include "common/check.h" +#include "explorer/ast/ast_rtti.h" +#include "explorer/common/arena.h" +#include "llvm/ADT/DenseMap.h" +#include "llvm/Support/Casting.h" + +namespace Carbon { + +class Element; +class Value; + +// A context for performing a deep copy of some fragment of the AST. +// +// This class carries the state necessary to make the copy, including ensuring +// that each node is cloned only once and mapping from old nodes to new ones. +class CloneContext { + public: + explicit CloneContext(Nonnull arena) : arena_(arena) {} + + CloneContext(const CloneContext&) = delete; + auto operator=(const CloneContext&) -> const CloneContext& = delete; + + // Clone an AST element. + template + auto Clone(Nonnull node) -> Nonnull { + if constexpr (std::is_convertible_v) { + const AstNode* base_node = node; + // Note, we can't use `llvm::cast` here because we might not have + // finished cloning `base_node` yet and its kind might not be set. This + // happens when there is a pointer cycle in the AST. + return static_cast(CloneBase(base_node)); + } else if constexpr (std::is_convertible_v) { + const Value* base_value = node; + return static_cast(CloneBase(base_value)); + } else { + static_assert(std::is_convertible_v, + "unknown pointer type to clone"); + const Element* base_elem = node; + return static_cast(CloneBase(base_elem)); + } + } + + // Clone anything with a clone constructor, that is, a constructor of the + // form: + // + // explicit MyType(CloneContext&, const MyType&) + // + // Clone constructors should call Clone on their owned elements to form a new + // value. Pointers returned by Clone should not be inspected by the clone + // constructor, as the pointee is not necessarily fully initialized until the + // overall cloning process completes. + // + // Clone constructors should call Remap on values that they do not own, such + // as for the declaration named by an IdentifierExpression. + template + auto Clone(const T& other) + -> std::enable_if_t, + T> { + return T(*this, other); + } + + template + auto Clone(std::optional node) -> std::optional { + if (node) { + return Clone(*node); + } + return std::nullopt; + } + + template + auto Clone(const std::vector& nodes) -> std::vector { + std::vector result; + result.reserve(nodes.size()); + for (const auto& node : nodes) { + result.push_back(Clone(node)); + } + return result; + } + + // Find the new or existing node corresponding to the given node. This should + // be used when a cloned node has a non-owning reference to another node, + // that might refer to something being cloned or might refer to the original + // object. The returned node might not be fully constructed and should not be + // inspected. + template + auto Remap(Nonnull node) -> Nonnull { + // Note, we can't use `llvm::cast` here because we might not have + // finished cloning `base_node` yet and its kind might not be set. This + // happens when there is a pointer cycle in the AST. + T* cloned = static_cast(nodes_[node]); + return cloned ? cloned : node; + } + + // It's safe to remap a `const` object by remapping the non-const version and + // adding back the `const`. + template + auto Remap(Nonnull node) -> Nonnull { + return Remap(const_cast(node)); + } + + template + auto Remap(std::optional node) -> std::optional { + if (node) { + return Remap(*node); + } + return std::nullopt; + } + + template + auto Remap(const std::vector& nodes) -> std::vector { + std::vector result; + result.reserve(nodes.size()); + for (const auto& node : nodes) { + result.push_back(Remap(node)); + } + return result; + } + + template + auto GetExistingClone(Nonnull node) -> Nonnull { + AstNode* cloned = nodes_.lookup(node); + CARBON_CHECK(cloned) << "expected node to be cloned"; + return llvm::cast(cloned); + } + + private: + // A value transform that remaps or clones AST elements referred to by the + // value being transformed. + class CloneValueTransform; + + // Clone the given node, and remember the mapping from the original to the + // new node for remapping. + auto CloneBase(Nonnull node) -> Nonnull; + + // Clone the given value, replacing references to cloned local declarations + // with references to the copies. + auto CloneBase(Nonnull value) -> Nonnull; + + // Clone the given element reference. + auto CloneBase(Nonnull elem) -> Nonnull; + + // Clone the given node if it's not already been cloned. This should be used + // very sparingly, in cases where ownership is unclear. + void MaybeCloneBase(Nonnull node); + + // Arena to allocate new nodes within. + Nonnull arena_; + + // Mapping from old nodes to new nodes. + llvm::DenseMap nodes_; +}; + +} // namespace Carbon + +#endif // CARBON_EXPLORER_AST_CLONE_CONTEXT_H_ diff --git a/explorer/ast/declaration.cpp b/explorer/ast/declaration.cpp index 1caea54f192c..2295677d98aa 100644 --- a/explorer/ast/declaration.cpp +++ b/explorer/ast/declaration.cpp @@ -4,6 +4,7 @@ #include "explorer/ast/declaration.h" +#include "explorer/ast/value.h" #include "llvm/ADT/StringExtras.h" #include "llvm/Support/Casting.h" @@ -152,8 +153,16 @@ void Declaration::PrintID(llvm::raw_ostream& out) const { out << "external "; break; } - out << "impl " << *impl_decl.impl_type() << " as " - << impl_decl.interface(); + out << "impl "; + if (!impl_decl.deduced_parameters().empty()) { + out << "forall ["; + llvm::ListSeparator sep; + for (auto* param : impl_decl.deduced_parameters()) { + out << sep << *param; + } + out << "] "; + } + out << *impl_decl.impl_type() << " as " << impl_decl.interface(); break; } case DeclarationKind::MatchFirstDeclaration: @@ -271,12 +280,6 @@ auto GetName(const Declaration& declaration) } } -void GenericBinding::Print(llvm::raw_ostream& out) const { - out << name() << ":! " << type(); -} - -void GenericBinding::PrintID(llvm::raw_ostream& out) const { out << name(); } - void ReturnTerm::Print(llvm::raw_ostream& out) const { switch (kind_) { case ReturnKind::Omitted: @@ -405,6 +408,27 @@ void CallableDeclaration::PrintDepth(int depth, llvm::raw_ostream& out) const { } } +ClassDeclaration::ClassDeclaration(CloneContext& context, + const ClassDeclaration& other) + : Declaration(context, other), + name_(other.name_), + extensibility_(other.extensibility_), + self_decl_(context.Clone(other.self_decl_)), + type_params_(context.Clone(other.type_params_)), + base_expr_(context.Clone(other.base_expr_)), + members_(context.Clone(other.members_)), + base_type_(context.Clone(other.base_type_)) {} + +ConstraintTypeDeclaration::ConstraintTypeDeclaration( + CloneContext& context, const ConstraintTypeDeclaration& other) + : Declaration(context, other), + name_(other.name_), + params_(context.Clone(other.params_)), + self_type_(context.Clone(other.self_type_)), + self_(context.Clone(other.self_)), + members_(context.Clone(other.members_)), + constraint_type_(context.Clone(other.constraint_type_)) {} + auto ImplDeclaration::Create(Nonnull arena, SourceLocation source_loc, ImplKind kind, Nonnull impl_type, Nonnull interface, @@ -428,6 +452,19 @@ auto ImplDeclaration::Create(Nonnull arena, SourceLocation source_loc, interface, resolved_params, members); } +ImplDeclaration::ImplDeclaration(CloneContext& context, + const ImplDeclaration& other) + : Declaration(context, other), + kind_(other.kind_), + deduced_parameters_(context.Clone(other.deduced_parameters_)), + impl_type_(context.Clone(other.impl_type_)), + self_decl_(context.Clone(other.self_decl_)), + interface_(context.Clone(other.interface_)), + constraint_type_(context.Clone(other.constraint_type_)), + members_(context.Clone(other.members_)), + impl_bindings_(context.Remap(other.impl_bindings_)), + match_first_(context.Remap(other.match_first_)) {} + void AlternativeSignature::Print(llvm::raw_ostream& out) const { out << "alt " << name(); if (auto params = parameters()) { @@ -449,4 +486,10 @@ auto ChoiceDeclaration::FindAlternative(std::string_view name) const return std::nullopt; } +MixDeclaration::MixDeclaration(CloneContext& context, + const MixDeclaration& other) + : Declaration(context, other), + mixin_(context.Clone(other.mixin_)), + mixin_value_(context.Clone(other.mixin_value_)) {} + } // namespace Carbon diff --git a/explorer/ast/declaration.h b/explorer/ast/declaration.h index b35f93752bff..0463942cd76c 100644 --- a/explorer/ast/declaration.h +++ b/explorer/ast/declaration.h @@ -13,6 +13,7 @@ #include "common/check.h" #include "common/ostream.h" #include "explorer/ast/ast_node.h" +#include "explorer/ast/clone_context.h" #include "explorer/ast/impl_binding.h" #include "explorer/ast/pattern.h" #include "explorer/ast/return_term.h" @@ -116,9 +117,16 @@ class Declaration : public AstNode { // Constructs a Declaration representing syntax at the given line number. // `kind` must be the enumerator corresponding to the most-derived type being // constructed. - Declaration(AstNodeKind kind, SourceLocation source_loc) + explicit Declaration(AstNodeKind kind, SourceLocation source_loc) : AstNode(kind, source_loc) {} + explicit Declaration(CloneContext& context, const Declaration& other) + : AstNode(context, other), + static_type_(context.Clone(other.static_type_)), + constant_value_(context.Clone(other.constant_value_)), + is_declared_(other.is_declared_), + is_type_checked_(other.is_type_checked_) {} + private: std::optional> static_type_; std::optional> constant_value_; @@ -183,6 +191,10 @@ class NamespaceDeclaration : public Declaration { : Declaration(AstNodeKind::NamespaceDeclaration, source_loc), name_(std::move(name)) {} + explicit NamespaceDeclaration(CloneContext& context, + const NamespaceDeclaration& other) + : Declaration(context, other), name_(other.name_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromNamespaceDeclaration(node->kind()); } @@ -214,6 +226,16 @@ class CallableDeclaration : public Declaration { body_(body), virt_override_(virt_override) {} + explicit CallableDeclaration(CloneContext& context, + const CallableDeclaration& other) + : Declaration(context, other), + deduced_parameters_(context.Clone(other.deduced_parameters_)), + self_pattern_(context.Clone(other.self_pattern_)), + param_pattern_(context.Clone(other.param_pattern_)), + return_term_(context.Clone(other.return_term_)), + body_(context.Clone(other.body_)), + virt_override_(other.virt_override_) {} + void PrintDepth(int depth, llvm::raw_ostream& out) const; auto deduced_parameters() const @@ -272,6 +294,10 @@ class FunctionDeclaration : public CallableDeclaration { param_pattern, return_term, body, virt_override), name_(std::move(name)) {} + explicit FunctionDeclaration(CloneContext& context, + const FunctionDeclaration& other) + : CallableDeclaration(context, other), name_(other.name_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromFunctionDeclaration(node->kind()); } @@ -306,6 +332,10 @@ class DestructorDeclaration : public CallableDeclaration { // TODO: Add virtual destructors VirtualOverride::None) {} + explicit DestructorDeclaration(CloneContext& context, + const DestructorDeclaration& other) + : CallableDeclaration(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromDestructorDeclaration(node->kind()); } @@ -318,6 +348,9 @@ class SelfDeclaration : public Declaration { explicit SelfDeclaration(SourceLocation source_loc) : Declaration(AstNodeKind::SelfDeclaration, source_loc) {} + explicit SelfDeclaration(CloneContext& context, const SelfDeclaration& other) + : Declaration(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromSelfDeclaration(node->kind()); } @@ -346,6 +379,9 @@ class ClassDeclaration : public Declaration { base_expr_(base), members_(std::move(members)) {} + explicit ClassDeclaration(CloneContext& context, + const ClassDeclaration& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromClassDeclaration(node->kind()); } @@ -364,7 +400,8 @@ class ClassDeclaration : public Declaration { auto members() const -> llvm::ArrayRef> { return members_; } - auto destructor() const -> std::optional> { + auto destructor() const + -> std::optional> { for (const auto& x : members_) { if (x->kind() == DeclarationKind::DestructorDeclaration) { return llvm::cast(x); @@ -396,8 +433,6 @@ class ClassDeclaration : public Declaration { std::optional> type_params_; std::optional> base_expr_; std::vector> members_; - std::optional> destructor_; - std::optional> base_; std::optional> base_type_; }; @@ -416,6 +451,14 @@ class MixinDeclaration : public Declaration { self_(self), members_(std::move(members)) {} + explicit MixinDeclaration(CloneContext& context, + const MixinDeclaration& other) + : Declaration(context, other), + name_(other.name_), + params_(context.Clone(other.params_)), + self_(context.Clone(other.self_)), + members_(context.Clone(other.members_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromMixinDeclaration(node->kind()); } @@ -448,6 +491,8 @@ class MixDeclaration : public Declaration { : Declaration(AstNodeKind::MixDeclaration, source_loc), mixin_(mixin_type) {} + explicit MixDeclaration(CloneContext& context, const MixDeclaration& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromMixDeclaration(node->kind()); } @@ -455,14 +500,14 @@ class MixDeclaration : public Declaration { auto mixin() const -> const Expression& { return **mixin_; } auto mixin() -> Expression& { return **mixin_; } - auto mixin_value() const -> const MixinPseudoType& { return *mixin_value_; } + auto mixin_value() const -> const MixinPseudoType& { return **mixin_value_; } void set_mixin_value(Nonnull mixin_value) { mixin_value_ = mixin_value; } private: std::optional> mixin_; - Nonnull mixin_value_; + std::optional> mixin_value_; }; class AlternativeSignature : public AstNode { @@ -473,6 +518,13 @@ class AlternativeSignature : public AstNode { name_(std::move(name)), parameters_(parameters) {} + explicit AlternativeSignature(CloneContext& context, + const AlternativeSignature& other) + : AstNode(context, other), + name_(other.name_), + parameters_(context.Clone(other.parameters_)), + parameters_static_type_(context.Clone(other.parameters_static_type_)) {} + void Print(llvm::raw_ostream& out) const override; void PrintID(llvm::raw_ostream& out) const override; @@ -520,6 +572,13 @@ class ChoiceDeclaration : public Declaration { type_params_(type_params), alternatives_(std::move(alternatives)) {} + explicit ChoiceDeclaration(CloneContext& context, + const ChoiceDeclaration& other) + : Declaration(context, other), + name_(other.name_), + type_params_(context.Clone(other.type_params_)), + alternatives_(context.Clone(other.alternatives_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromChoiceDeclaration(node->kind()); } @@ -564,6 +623,13 @@ class VariableDeclaration : public Declaration { initializer_(initializer), value_category_(value_category) {} + explicit VariableDeclaration(CloneContext& context, + const VariableDeclaration& other) + : Declaration(context, other), + binding_(context.Clone(other.binding_)), + initializer_(context.Clone(other.initializer_)), + value_category_(other.value_category_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromVariableDeclaration(node->kind()); } @@ -611,6 +677,9 @@ class ConstraintTypeDeclaration : public Declaration { self_ = arena->New(source_loc, "Self", self_type_ref); } + explicit ConstraintTypeDeclaration(CloneContext& context, + const ConstraintTypeDeclaration& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromConstraintTypeDeclaration(node->kind()); } @@ -671,6 +740,10 @@ class InterfaceDeclaration : public ConstraintTypeDeclaration { source_loc, std::move(name), params, std::move(members)) {} + explicit InterfaceDeclaration(CloneContext& context, + const InterfaceDeclaration& other) + : ConstraintTypeDeclaration(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromInterfaceDeclaration(node->kind()); } @@ -689,6 +762,10 @@ class ConstraintDeclaration : public ConstraintTypeDeclaration { source_loc, std::move(name), params, std::move(members)) {} + explicit ConstraintDeclaration(CloneContext& context, + const ConstraintDeclaration& other) + : ConstraintTypeDeclaration(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromConstraintDeclaration(node->kind()); } @@ -702,6 +779,10 @@ class InterfaceExtendsDeclaration : public Declaration { : Declaration(AstNodeKind::InterfaceExtendsDeclaration, source_loc), base_(base) {} + explicit InterfaceExtendsDeclaration(CloneContext& context, + const InterfaceExtendsDeclaration& other) + : Declaration(context, other), base_(context.Clone(other.base_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromInterfaceExtendsDeclaration(node->kind()); } @@ -723,6 +804,12 @@ class InterfaceImplDeclaration : public Declaration { impl_type_(impl_type), constraint_(constraint) {} + explicit InterfaceImplDeclaration(CloneContext& context, + const InterfaceImplDeclaration& other) + : Declaration(context, other), + impl_type_(context.Clone(other.impl_type_)), + constraint_(context.Clone(other.constraint_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromInterfaceImplDeclaration(node->kind()); } @@ -747,6 +834,10 @@ class AssociatedConstantDeclaration : public Declaration { : Declaration(AstNodeKind::AssociatedConstantDeclaration, source_loc), binding_(binding) {} + explicit AssociatedConstantDeclaration( + CloneContext& context, const AssociatedConstantDeclaration& other) + : Declaration(context, other), binding_(context.Clone(other.binding_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromAssociatedConstantDeclaration(node->kind()); } @@ -780,12 +871,14 @@ class ImplDeclaration : public Declaration { std::vector> members) : Declaration(AstNodeKind::ImplDeclaration, source_loc), kind_(kind), + deduced_parameters_(std::move(deduced_params)), impl_type_(impl_type), self_decl_(self_decl), interface_(interface), - deduced_parameters_(std::move(deduced_params)), members_(std::move(members)) {} + explicit ImplDeclaration(CloneContext& context, const ImplDeclaration& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromImplDeclaration(node->kind()); } @@ -837,11 +930,11 @@ class ImplDeclaration : public Declaration { private: ImplKind kind_; + std::vector> deduced_parameters_; Nonnull impl_type_; Nonnull self_decl_; Nonnull interface_; std::optional> constraint_type_; - std::vector> deduced_parameters_; std::vector> members_; std::vector> impl_bindings_; std::optional> match_first_; @@ -854,6 +947,10 @@ class MatchFirstDeclaration : public Declaration { : Declaration(AstNodeKind::MatchFirstDeclaration, source_loc), impls_(std::move(impls)) {} + explicit MatchFirstDeclaration(CloneContext& context, + const MatchFirstDeclaration& other) + : Declaration(context, other), impls_(context.Clone(other.impls_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromMatchFirstDeclaration(node->kind()); } @@ -877,6 +974,12 @@ class AliasDeclaration : public Declaration { name_(std::move(name)), target_(target) {} + explicit AliasDeclaration(CloneContext& context, + const AliasDeclaration& other) + : Declaration(context, other), + name_(other.name_), + target_(context.Clone(other.target_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromAliasDeclaration(node->kind()); } diff --git a/explorer/ast/expression.cpp b/explorer/ast/expression.cpp index 7ad9c811ed1f..52d297775d46 100644 --- a/explorer/ast/expression.cpp +++ b/explorer/ast/expression.cpp @@ -365,6 +365,35 @@ void Expression::PrintID(llvm::raw_ostream& out) const { } } +DotSelfExpression::DotSelfExpression(CloneContext& context, + const DotSelfExpression& other) + : Expression(context, other), + name_(other.name_), + self_binding_(context.Remap(other.self_binding_)) {} + +MemberAccessExpression::MemberAccessExpression( + CloneContext& context, const MemberAccessExpression& other) + : Expression(context, other), + object_(context.Clone(other.object_)), + is_type_access_(other.is_type_access_), + is_addr_me_method_(other.is_addr_me_method_), + impl_(context.Clone(other.impl_)), + constant_value_(context.Clone(other.constant_value_)) {} + +SimpleMemberAccessExpression::SimpleMemberAccessExpression( + CloneContext& context, const SimpleMemberAccessExpression& other) + : RewritableMixin(context, other), + member_name_(other.member_name_), + member_(context.Clone(other.member_)), + found_in_interface_(context.Clone(other.found_in_interface_)), + value_node_(context.Clone(other.value_node_)) {} + +CompoundMemberAccessExpression::CompoundMemberAccessExpression( + CloneContext& context, const CompoundMemberAccessExpression& other) + : MemberAccessExpression(context, other), + path_(context.Clone(other.path_)), + member_(context.Clone(other.member_)) {} + WhereClause::~WhereClause() = default; void WhereClause::Print(llvm::raw_ostream& out) const { @@ -389,4 +418,11 @@ void WhereClause::Print(llvm::raw_ostream& out) const { void WhereClause::PrintID(llvm::raw_ostream& out) const { out << "..."; } +WhereExpression::WhereExpression(CloneContext& context, + const WhereExpression& other) + : RewritableMixin(context, other), + self_binding_(context.Clone(other.self_binding_)), + clauses_(context.Clone(other.clauses_)), + enclosing_dot_self_(context.Remap(other.enclosing_dot_self_)) {} + } // namespace Carbon diff --git a/explorer/ast/expression.h b/explorer/ast/expression.h index 6c7901e11198..e0983a2ba9a2 100644 --- a/explorer/ast/expression.h +++ b/explorer/ast/expression.h @@ -78,7 +78,7 @@ class Expression : public AstNode { // Determines whether the expression has already been type-checked. Should // only be used by type-checking. - auto is_type_checked() -> bool { + auto is_type_checked() const -> bool { return static_type_.has_value() && value_category_.has_value(); } @@ -89,6 +89,9 @@ class Expression : public AstNode { Expression(AstNodeKind kind, SourceLocation source_loc) : AstNode(kind, source_loc) {} + explicit Expression(CloneContext& context, const Expression& other) + : AstNode(context, other) {} + private: std::optional> static_type_; std::optional value_category_; @@ -101,6 +104,10 @@ class RewritableMixin : public Base { public: using Base::Base; + explicit RewritableMixin(CloneContext& context, const RewritableMixin& other) + : Base(context, other), + rewritten_form_(context.Clone(other.rewritten_form_)) {} + // Set the rewritten form of this expression. Can only be called during type // checking. auto set_rewritten_form(Nonnull rewritten_form) -> void { @@ -127,6 +134,10 @@ class FieldInitializer { FieldInitializer(std::string name, Nonnull expression) : name_(std::move(name)), expression_(expression) {} + explicit FieldInitializer(CloneContext& context, + const FieldInitializer& other) + : name_(other.name_), expression_(context.Clone(other.expression_)) {} + auto name() const -> const std::string& { return name_; } auto expression() const -> const Expression& { return *expression_; } @@ -177,6 +188,12 @@ class IdentifierExpression : public Expression { : Expression(AstNodeKind::IdentifierExpression, source_loc), name_(std::move(name)) {} + explicit IdentifierExpression(CloneContext& context, + const IdentifierExpression& other) + : Expression(context, other), + name_(other.name_), + value_node_(context.Clone(other.value_node_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIdentifierExpression(node->kind()); } @@ -214,6 +231,9 @@ class DotSelfExpression : public Expression { explicit DotSelfExpression(SourceLocation source_loc) : Expression(AstNodeKind::DotSelfExpression, source_loc) {} + explicit DotSelfExpression(CloneContext& context, + const DotSelfExpression& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromDotSelfExpression(node->kind()); } @@ -239,6 +259,9 @@ class MemberAccessExpression : public Expression { Nonnull object) : Expression(kind, source_loc), object_(object) {} + explicit MemberAccessExpression(CloneContext& context, + const MemberAccessExpression& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromMemberAccessExpression(node->kind()); } @@ -314,6 +337,9 @@ class SimpleMemberAccessExpression object), member_name_(std::move(member_name)) {} + explicit SimpleMemberAccessExpression( + CloneContext& context, const SimpleMemberAccessExpression& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromSimpleMemberAccessExpression(node->kind()); } @@ -386,6 +412,9 @@ class CompoundMemberAccessExpression : public MemberAccessExpression { source_loc, object), path_(path) {} + explicit CompoundMemberAccessExpression( + CloneContext& context, const CompoundMemberAccessExpression& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromCompoundMemberAccessExpression(node->kind()); } @@ -420,6 +449,11 @@ class IndexExpression : public Expression { object_(object), offset_(offset) {} + explicit IndexExpression(CloneContext& context, const IndexExpression& other) + : Expression(context, other), + object_(context.Clone(other.object_)), + offset_(context.Clone(other.offset_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIndexExpression(node->kind()); } @@ -446,6 +480,11 @@ class BaseAccessExpression : public MemberAccessExpression { set_value_category(ValueCategory::Let); } + explicit BaseAccessExpression(CloneContext& context, + const BaseAccessExpression& other) + : MemberAccessExpression(context, other), + base_(context.Clone(other.base_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromBaseAccessExpression(node->kind()); } @@ -461,6 +500,9 @@ class IntLiteral : public Expression { explicit IntLiteral(SourceLocation source_loc, int value) : Expression(AstNodeKind::IntLiteral, source_loc), value_(value) {} + explicit IntLiteral(CloneContext& context, const IntLiteral& other) + : Expression(context, other), value_(other.value_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIntLiteral(node->kind()); } @@ -476,6 +518,9 @@ class BoolLiteral : public Expression { explicit BoolLiteral(SourceLocation source_loc, bool value) : Expression(AstNodeKind::BoolLiteral, source_loc), value_(value) {} + explicit BoolLiteral(CloneContext& context, const BoolLiteral& other) + : Expression(context, other), value_(other.value_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromBoolLiteral(node->kind()); } @@ -492,6 +537,9 @@ class StringLiteral : public Expression { : Expression(AstNodeKind::StringLiteral, source_loc), value_(std::move(value)) {} + explicit StringLiteral(CloneContext& context, const StringLiteral& other) + : Expression(context, other), value_(other.value_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromStringLiteral(node->kind()); } @@ -507,6 +555,10 @@ class StringTypeLiteral : public Expression { explicit StringTypeLiteral(SourceLocation source_loc) : Expression(AstNodeKind::StringTypeLiteral, source_loc) {} + explicit StringTypeLiteral(CloneContext& context, + const StringTypeLiteral& other) + : Expression(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromStringTypeLiteral(node->kind()); } @@ -522,6 +574,9 @@ class TupleLiteral : public Expression { : Expression(AstNodeKind::TupleLiteral, source_loc), fields_(std::move(fields)) {} + explicit TupleLiteral(CloneContext& context, const TupleLiteral& other) + : Expression(context, other), fields_(context.Clone(other.fields_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromTupleLiteral(node->kind()); } @@ -545,6 +600,9 @@ class StructLiteral : public Expression { : Expression(AstNodeKind::StructLiteral, loc), fields_(std::move(fields)) {} + explicit StructLiteral(CloneContext& context, const StructLiteral& other) + : Expression(context, other), fields_(context.Clone(other.fields_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromStructLiteral(node->kind()); } @@ -564,6 +622,11 @@ class ConstantValueLiteral : public Expression { std::optional> constant_value = std::nullopt) : Expression(kind, source_loc), constant_value_(constant_value) {} + explicit ConstantValueLiteral(CloneContext& context, + const ConstantValueLiteral& other) + : Expression(context, other), + constant_value_(context.Clone(other.constant_value_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromConstantValueLiteral(node->kind()); } @@ -599,6 +662,11 @@ class StructTypeLiteral : public ConstantValueLiteral { << "`{}` is represented as a StructLiteral, not a StructTypeLiteral."; } + explicit StructTypeLiteral(CloneContext& context, + const StructTypeLiteral& other) + : ConstantValueLiteral(context, other), + fields_(context.Clone(other.fields_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromStructTypeLiteral(node->kind()); } @@ -618,6 +686,12 @@ class OperatorExpression : public RewritableMixin { op_(op), arguments_(std::move(arguments)) {} + explicit OperatorExpression(CloneContext& context, + const OperatorExpression& other) + : RewritableMixin(context, other), + op_(other.op_), + arguments_(context.Clone(other.arguments_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromOperatorExpression(node->kind()); } @@ -645,6 +719,12 @@ class CallExpression : public Expression { argument_(argument), bindings_({}, {}) {} + explicit CallExpression(CloneContext& context, const CallExpression& other) + : Expression(context, other), + function_(context.Clone(other.function_)), + argument_(context.Clone(other.argument_)), + bindings_(context.Clone(other.bindings_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromCallExpression(node->kind()); } @@ -687,6 +767,12 @@ class FunctionTypeLiteral : public ConstantValueLiteral { parameter_(parameter), return_type_(return_type) {} + explicit FunctionTypeLiteral(CloneContext& context, + const FunctionTypeLiteral& other) + : ConstantValueLiteral(context, other), + parameter_(context.Clone(other.parameter_)), + return_type_(context.Clone(other.return_type_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromFunctionTypeLiteral(node->kind()); } @@ -706,6 +792,9 @@ class BoolTypeLiteral : public Expression { explicit BoolTypeLiteral(SourceLocation source_loc) : Expression(AstNodeKind::BoolTypeLiteral, source_loc) {} + explicit BoolTypeLiteral(CloneContext& context, const BoolTypeLiteral& other) + : Expression(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromBoolTypeLiteral(node->kind()); } @@ -716,6 +805,9 @@ class IntTypeLiteral : public Expression { explicit IntTypeLiteral(SourceLocation source_loc) : Expression(AstNodeKind::IntTypeLiteral, source_loc) {} + explicit IntTypeLiteral(CloneContext& context, const IntTypeLiteral& other) + : Expression(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIntTypeLiteral(node->kind()); } @@ -726,6 +818,10 @@ class ContinuationTypeLiteral : public Expression { explicit ContinuationTypeLiteral(SourceLocation source_loc) : Expression(AstNodeKind::ContinuationTypeLiteral, source_loc) {} + explicit ContinuationTypeLiteral(CloneContext& context, + const ContinuationTypeLiteral& other) + : Expression(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromContinuationTypeLiteral(node->kind()); } @@ -736,6 +832,9 @@ class TypeTypeLiteral : public Expression { explicit TypeTypeLiteral(SourceLocation source_loc) : Expression(AstNodeKind::TypeTypeLiteral, source_loc) {} + explicit TypeTypeLiteral(CloneContext& context, const TypeTypeLiteral& other) + : Expression(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromTypeTypeLiteral(node->kind()); } @@ -754,6 +853,9 @@ class ValueLiteral : public ConstantValueLiteral { set_value_category(value_category); } + explicit ValueLiteral(CloneContext& context, const ValueLiteral& other) + : ConstantValueLiteral(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromValueLiteral(node->kind()); } @@ -792,6 +894,12 @@ class IntrinsicExpression : public Expression { intrinsic_(intrinsic), args_(args) {} + explicit IntrinsicExpression(CloneContext& context, + const IntrinsicExpression& other) + : Expression(context, other), + intrinsic_(other.intrinsic_), + args_(context.Clone(other.args_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIntrinsicExpression(node->kind()); } @@ -817,6 +925,12 @@ class IfExpression : public Expression { then_expression_(then_expression), else_expression_(else_expression) {} + explicit IfExpression(CloneContext& context, const IfExpression& other) + : Expression(context, other), + condition_(context.Clone(other.condition_)), + then_expression_(context.Clone(other.then_expression_)), + else_expression_(context.Clone(other.else_expression_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIfExpression(node->kind()); } @@ -861,8 +975,11 @@ class WhereClause : public AstNode { } protected: - WhereClause(WhereClauseKind kind, SourceLocation source_loc) + explicit WhereClause(WhereClauseKind kind, SourceLocation source_loc) : AstNode(static_cast(kind), source_loc) {} + + explicit WhereClause(CloneContext& context, const WhereClause& other) + : AstNode(context, other) {} }; // An `is` where clause. @@ -877,6 +994,11 @@ class IsWhereClause : public WhereClause { type_(type), constraint_(constraint) {} + explicit IsWhereClause(CloneContext& context, const IsWhereClause& other) + : WhereClause(context, other), + type_(context.Clone(other.type_)), + constraint_(context.Clone(other.constraint_)) {} + static auto classof(const AstNode* node) { return InheritsFromIsWhereClause(node->kind()); } @@ -904,6 +1026,12 @@ class EqualsWhereClause : public WhereClause { lhs_(lhs), rhs_(rhs) {} + explicit EqualsWhereClause(CloneContext& context, + const EqualsWhereClause& other) + : WhereClause(context, other), + lhs_(context.Clone(other.lhs_)), + rhs_(context.Clone(other.rhs_)) {} + static auto classof(const AstNode* node) { return InheritsFromEqualsWhereClause(node->kind()); } @@ -932,6 +1060,12 @@ class RewriteWhereClause : public WhereClause { member_name_(std::move(member_name)), replacement_(replacement) {} + explicit RewriteWhereClause(CloneContext& context, + const RewriteWhereClause& other) + : WhereClause(context, other), + member_name_(other.member_name_), + replacement_(context.Clone(other.replacement_)) {} + static auto classof(const AstNode* node) { return InheritsFromRewriteWhereClause(node->kind()); } @@ -959,6 +1093,8 @@ class WhereExpression : public RewritableMixin { self_binding_(self_binding), clauses_(std::move(clauses)) {} + explicit WhereExpression(CloneContext& context, const WhereExpression& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromWhereExpression(node->kind()); } @@ -1003,6 +1139,12 @@ class BuiltinConvertExpression : public Expression { set_value_category(ValueCategory::Let); } + explicit BuiltinConvertExpression(CloneContext& context, + const BuiltinConvertExpression& other) + : Expression(context, other), + source_expression_(context.Clone(other.source_expression_)), + rewritten_form_(context.Clone(other.rewritten_form_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromBuiltinConvertExpression(node->kind()); } @@ -1049,6 +1191,12 @@ class UnimplementedExpression : public Expression { AddChildren(children...); } + explicit UnimplementedExpression(CloneContext& context, + const UnimplementedExpression& other) + : Expression(context, other), + label_(other.label_), + children_(context.Clone(other.children_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromUnimplementedExpression(node->kind()); } @@ -1076,13 +1224,19 @@ class ArrayTypeLiteral : public ConstantValueLiteral { public: // Constructs an array type literal which uses the given expressions to // represent the element type and size. - ArrayTypeLiteral(SourceLocation source_loc, - Nonnull element_type_expression, - Nonnull size_expression) + explicit ArrayTypeLiteral(SourceLocation source_loc, + Nonnull element_type_expression, + Nonnull size_expression) : ConstantValueLiteral(AstNodeKind::ArrayTypeLiteral, source_loc), element_type_expression_(element_type_expression), size_expression_(size_expression) {} + explicit ArrayTypeLiteral(CloneContext& context, + const ArrayTypeLiteral& other) + : ConstantValueLiteral(context, other), + element_type_expression_(context.Clone(other.element_type_expression_)), + size_expression_(context.Clone(other.size_expression_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromArrayTypeLiteral(node->kind()); } diff --git a/explorer/ast/impl_binding.cpp b/explorer/ast/impl_binding.cpp new file mode 100644 index 000000000000..008ae3578f91 --- /dev/null +++ b/explorer/ast/impl_binding.cpp @@ -0,0 +1,18 @@ +// 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 + +#include "explorer/ast/impl_binding.h" + +#include "explorer/ast/pattern.h" + +namespace Carbon { + +ImplBinding::ImplBinding(CloneContext& context, const ImplBinding& other) + : AstNode(context, other), + type_var_(context.Remap(other.type_var_)), + iface_(context.Clone(other.iface_)), + symbolic_identity_(context.Clone(other.symbolic_identity_)), + original_(context.Remap(other.original_)) {} + +} // namespace Carbon diff --git a/explorer/ast/impl_binding.h b/explorer/ast/impl_binding.h index 59cca394a3e3..1ff9b0a6aee1 100644 --- a/explorer/ast/impl_binding.h +++ b/explorer/ast/impl_binding.h @@ -10,14 +10,13 @@ #include "common/check.h" #include "common/ostream.h" #include "explorer/ast/ast_node.h" -#include "explorer/ast/pattern.h" #include "explorer/ast/value_category.h" namespace Carbon { class Value; class Expression; -class ImplBinding; +class GenericBinding; // `ImplBinding` plays the role of the parameter for passing witness // tables to a generic. However, unlike regular parameters @@ -37,6 +36,8 @@ class ImplBinding : public AstNode { type_var_(type_var), iface_(iface) {} + explicit ImplBinding(CloneContext& context, const ImplBinding& other); + static auto classof(const AstNode* node) -> bool { return InheritsFromImplBinding(node->kind()); } diff --git a/explorer/ast/pattern.cpp b/explorer/ast/pattern.cpp index 4be25458b1b8..f5780adc7036 100644 --- a/explorer/ast/pattern.cpp +++ b/explorer/ast/pattern.cpp @@ -8,6 +8,8 @@ #include "common/ostream.h" #include "explorer/ast/expression.h" +#include "explorer/ast/impl_binding.h" +#include "explorer/ast/value.h" #include "explorer/common/arena.h" #include "explorer/common/error_builders.h" #include "llvm/ADT/StringExtras.h" @@ -170,4 +172,14 @@ auto ParenExpressionToParenPattern(Nonnull arena, return result; } +GenericBinding::GenericBinding(CloneContext& context, + const GenericBinding& other) + : Pattern(context, other), + name_(other.name_), + type_(context.Clone(other.type_)), + symbolic_identity_(context.Clone(other.symbolic_identity_)), + impl_binding_(context.Clone(other.impl_binding_)), + original_(context.Remap(other.original_)), + named_as_type_via_dot_self_(other.named_as_type_via_dot_self_) {} + } // namespace Carbon diff --git a/explorer/ast/pattern.h b/explorer/ast/pattern.h index a3ecdd4dff92..60b75bd7030f 100644 --- a/explorer/ast/pattern.h +++ b/explorer/ast/pattern.h @@ -12,6 +12,7 @@ #include "common/ostream.h" #include "explorer/ast/ast_node.h" #include "explorer/ast/ast_rtti.h" +#include "explorer/ast/clone_context.h" #include "explorer/ast/expression.h" #include "explorer/ast/value_category.h" #include "explorer/ast/value_node.h" @@ -33,6 +34,11 @@ class Value; // details. class Pattern : public AstNode { public: + explicit Pattern(CloneContext& context, const Pattern& other) + : AstNode(context, other), + static_type_(context.Clone(other.static_type_)), + value_(context.Clone(other.value_)) {} + Pattern(const Pattern&) = delete; auto operator=(const Pattern&) -> Pattern& = delete; @@ -123,6 +129,9 @@ class AutoPattern : public Pattern { explicit AutoPattern(SourceLocation source_loc) : Pattern(AstNodeKind::AutoPattern, source_loc) {} + explicit AutoPattern(CloneContext& context, const AutoPattern& other) + : Pattern(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromAutoPattern(node->kind()); } @@ -133,6 +142,9 @@ class VarPattern : public Pattern { explicit VarPattern(SourceLocation source_loc, Nonnull pattern) : Pattern(AstNodeKind::VarPattern, source_loc), pattern_(pattern) {} + explicit VarPattern(CloneContext& context, const VarPattern& other) + : Pattern(context, other), pattern_(context.Clone(other.pattern_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromVarPattern(node->kind()); } @@ -160,6 +172,12 @@ class BindingPattern : public Pattern { type_(type), value_category_(value_category) {} + explicit BindingPattern(CloneContext& context, const BindingPattern& other) + : Pattern(context, other), + name_(other.name_), + type_(context.Clone(other.type_)), + value_category_(other.value_category_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromBindingPattern(node->kind()); } @@ -211,6 +229,9 @@ class AddrPattern : public Pattern { Nonnull binding) : Pattern(AstNodeKind::AddrPattern, source_loc), binding_(binding) {} + explicit AddrPattern(CloneContext& context, const AddrPattern& other) + : Pattern(context, other), binding_(context.Clone(other.binding_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromAddrPattern(node->kind()); } @@ -229,6 +250,9 @@ class TuplePattern : public Pattern { : Pattern(AstNodeKind::TuplePattern, source_loc), fields_(std::move(fields)) {} + explicit TuplePattern(CloneContext& context, const TuplePattern& other) + : Pattern(context, other), fields_(context.Clone(other.fields_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromTuplePattern(node->kind()); } @@ -252,8 +276,7 @@ class GenericBinding : public Pattern { name_(std::move(name)), type_(type) {} - void Print(llvm::raw_ostream& out) const override; - void PrintID(llvm::raw_ostream& out) const override; + explicit GenericBinding(CloneContext& context, const GenericBinding& other); static auto classof(const AstNode* node) -> bool { return InheritsFromGenericBinding(node->kind()); @@ -376,6 +399,13 @@ class AlternativePattern : public Pattern { alternative_name_(std::move(alternative_name)), arguments_(arguments) {} + explicit AlternativePattern(CloneContext& context, + const AlternativePattern& other) + : Pattern(context, other), + choice_type_(context.Clone(other.choice_type_)), + alternative_name_(other.alternative_name_), + arguments_(context.Clone(other.arguments_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromAlternativePattern(node->kind()); } @@ -405,6 +435,11 @@ class ExpressionPattern : public Pattern { : Pattern(AstNodeKind::ExpressionPattern, expression->source_loc()), expression_(expression) {} + explicit ExpressionPattern(CloneContext& context, + const ExpressionPattern& other) + : Pattern(context, other), + expression_(context.Clone(other.expression_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromExpressionPattern(node->kind()); } diff --git a/explorer/ast/return_term.h b/explorer/ast/return_term.h index bbde4c6dbd76..b3c7ff2eec64 100644 --- a/explorer/ast/return_term.h +++ b/explorer/ast/return_term.h @@ -10,6 +10,7 @@ #include "common/check.h" #include "common/ostream.h" +#include "explorer/ast/clone_context.h" #include "explorer/ast/expression.h" #include "explorer/common/nonnull.h" #include "explorer/common/source_location.h" @@ -26,6 +27,12 @@ class Value; // Each of these forms has a corresponding factory function. class ReturnTerm { public: + explicit ReturnTerm(CloneContext& context, const ReturnTerm& other) + : kind_(other.kind_), + type_expression_(context.Clone(other.type_expression_)), + static_type_(context.Clone(other.static_type_)), + source_loc_(other.source_loc_) {} + ReturnTerm(const ReturnTerm&) = default; auto operator=(const ReturnTerm&) -> ReturnTerm& = default; diff --git a/explorer/ast/statement.cpp b/explorer/ast/statement.cpp index d634f74686e4..c9d195f63d62 100644 --- a/explorer/ast/statement.cpp +++ b/explorer/ast/statement.cpp @@ -5,6 +5,7 @@ #include "explorer/ast/statement.h" #include "common/check.h" +#include "explorer/ast/declaration.h" #include "explorer/common/arena.h" #include "llvm/Support/Casting.h" @@ -170,4 +171,7 @@ auto AssignOperatorToString(AssignOperator op) -> std::string_view { } } +Return::Return(CloneContext& context, const Return& other) + : Statement(context, other), function_(context.Remap(other.function_)) {} + } // namespace Carbon diff --git a/explorer/ast/statement.h b/explorer/ast/statement.h index df653fe75442..110cecb72bd2 100644 --- a/explorer/ast/statement.h +++ b/explorer/ast/statement.h @@ -10,6 +10,7 @@ #include "common/ostream.h" #include "explorer/ast/ast_node.h" +#include "explorer/ast/clone_context.h" #include "explorer/ast/expression.h" #include "explorer/ast/pattern.h" #include "explorer/ast/return_term.h" @@ -43,8 +44,11 @@ class Statement : public AstNode { } protected: - Statement(AstNodeKind kind, SourceLocation source_loc) + explicit Statement(AstNodeKind kind, SourceLocation source_loc) : AstNode(kind, source_loc) {} + + explicit Statement(CloneContext& context, const Statement& other) + : AstNode(context, other) {} }; class Block : public Statement { @@ -53,6 +57,10 @@ class Block : public Statement { : Statement(AstNodeKind::Block, source_loc), statements_(std::move(statements)) {} + explicit Block(CloneContext& context, const Block& other) + : Statement(context, other), + statements_(context.Clone(other.statements_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromBlock(node->kind()); } @@ -75,6 +83,11 @@ class ExpressionStatement : public Statement { : Statement(AstNodeKind::ExpressionStatement, source_loc), expression_(expression) {} + explicit ExpressionStatement(CloneContext& context, + const ExpressionStatement& other) + : Statement(context, other), + expression_(context.Clone(other.expression_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromExpressionStatement(node->kind()); } @@ -112,6 +125,13 @@ class Assign : public Statement { rhs_(rhs), op_(op) {} + explicit Assign(CloneContext& context, const Assign& other) + : Statement(context, other), + lhs_(context.Clone(other.lhs_)), + rhs_(context.Clone(other.rhs_)), + op_(other.op_), + rewritten_form_(context.Clone(other.rewritten_form_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromAssign(node->kind()); } @@ -156,6 +176,13 @@ class IncrementDecrement : public Statement { argument_(argument), is_increment_(is_increment) {} + explicit IncrementDecrement(CloneContext& context, + const IncrementDecrement& other) + : Statement(context, other), + argument_(context.Clone(other.argument_)), + is_increment_(other.is_increment_), + rewritten_form_(context.Clone(other.rewritten_form_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIncrementDecrement(node->kind()); } @@ -199,6 +226,14 @@ class VariableDefinition : public Statement { value_category_(value_category), def_type_(def_type) {} + explicit VariableDefinition(CloneContext& context, + const VariableDefinition& other) + : Statement(context, other), + pattern_(context.Clone(other.pattern_)), + init_(context.Clone(other.init_)), + value_category_(other.value_category_), + def_type_(other.def_type_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromVariableDefinition(node->kind()); } @@ -243,6 +278,12 @@ class If : public Statement { then_block_(then_block), else_block_(else_block) {} + explicit If(CloneContext& context, const If& other) + : Statement(context, other), + condition_(context.Clone(other.condition_)), + then_block_(context.Clone(other.then_block_)), + else_block_(context.Clone(other.else_block_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromIf(node->kind()); } @@ -290,6 +331,8 @@ class Return : public Statement { Return(AstNodeKind node_kind, SourceLocation source_loc) : Statement(node_kind, source_loc) {} + explicit Return(CloneContext& context, const Return& other); + private: std::optional> function_; }; @@ -299,6 +342,9 @@ class ReturnVar : public Return { explicit ReturnVar(SourceLocation source_loc) : Return(AstNodeKind::ReturnVar, source_loc) {} + explicit ReturnVar(CloneContext& context, const ReturnVar& other) + : Return(context, other), value_node_(context.Clone(other.value_node_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromReturnVar(node->kind()); } @@ -329,6 +375,12 @@ class ReturnExpression : public Return { expression_(expression), is_omitted_expression_(is_omitted_expression) {} + explicit ReturnExpression(CloneContext& context, + const ReturnExpression& other) + : Return(context, other), + expression_(context.Clone(other.expression_)), + is_omitted_expression_(other.is_omitted_expression_) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromReturnExpression(node->kind()); } @@ -355,6 +407,11 @@ class While : public Statement { condition_(condition), body_(body) {} + explicit While(CloneContext& context, const While& other) + : Statement(context, other), + condition_(context.Clone(other.condition_)), + body_(context.Clone(other.body_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromWhile(node->kind()); } @@ -381,6 +438,12 @@ class For : public Statement { loop_target_(loop_target), body_(body) {} + explicit For(CloneContext& context, const For& other) + : Statement(context, other), + variable_declaration_(context.Clone(other.variable_declaration_)), + loop_target_(context.Clone(other.loop_target_)), + body_(context.Clone(other.body_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromFor(node->kind()); } @@ -409,6 +472,9 @@ class Break : public Statement { explicit Break(SourceLocation source_loc) : Statement(AstNodeKind::Break, source_loc) {} + explicit Break(CloneContext& context, const Break& other) + : Statement(context, other), loop_(context.Clone(other.loop_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromBreak(node->kind()); } @@ -436,6 +502,9 @@ class Continue : public Statement { explicit Continue(SourceLocation source_loc) : Statement(AstNodeKind::Continue, source_loc) {} + explicit Continue(CloneContext& context, const Continue& other) + : Statement(context, other), loop_(context.Clone(other.loop_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromContinue(node->kind()); } @@ -462,9 +531,13 @@ class Match : public Statement { public: class Clause { public: - Clause(Nonnull pattern, Nonnull statement) + explicit Clause(Nonnull pattern, Nonnull statement) : pattern_(pattern), statement_(statement) {} + explicit Clause(CloneContext& context, const Clause& other) + : pattern_(context.Clone(other.pattern_)), + statement_(context.Clone(other.statement_)) {} + auto pattern() const -> const Pattern& { return *pattern_; } auto pattern() -> Pattern& { return *pattern_; } auto statement() const -> const Statement& { return *statement_; } @@ -481,6 +554,11 @@ class Match : public Statement { expression_(expression), clauses_(std::move(clauses)) {} + explicit Match(CloneContext& context, const Match& other) + : Statement(context, other), + expression_(context.Clone(other.expression_)), + clauses_(context.Clone(other.clauses_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromMatch(node->kind()); } @@ -515,6 +593,12 @@ class Continuation : public Statement { name_(std::move(name)), body_(body) {} + explicit Continuation(CloneContext& context, const Continuation& other) + : Statement(context, other), + name_(other.name_), + body_(context.Clone(other.body_)), + static_type_(context.Clone(other.static_type_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromContinuation(node->kind()); } @@ -558,6 +642,9 @@ class Run : public Statement { Run(SourceLocation source_loc, Nonnull argument) : Statement(AstNodeKind::Run, source_loc), argument_(argument) {} + explicit Run(CloneContext& context, const Run& other) + : Statement(context, other), argument_(context.Clone(other.argument_)) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromRun(node->kind()); } @@ -580,6 +667,9 @@ class Await : public Statement { explicit Await(SourceLocation source_loc) : Statement(AstNodeKind::Await, source_loc) {} + explicit Await(CloneContext& context, const Await& other) + : Statement(context, other) {} + static auto classof(const AstNode* node) -> bool { return InheritsFromAwait(node->kind()); } diff --git a/explorer/ast/value.h b/explorer/ast/value.h index a709f96429a4..9b1a31e6d448 100644 --- a/explorer/ast/value.h +++ b/explorer/ast/value.h @@ -614,6 +614,11 @@ class FunctionType : public Value { // // fn MakeEmptyVector(T:! type) -> Vector(T); struct GenericParameter { + template + auto Decompose(F f) const { + return f(index, binding); + } + size_t index; Nonnull binding; }; @@ -923,6 +928,11 @@ class NamedConstraintType : public Value { // A constraint that requires implementation of an interface. struct ImplConstraint { + template + auto Decompose(F f) const { + return f(type, interface); + } + // The type that is required to implement the interface. Nonnull type; // The interface that is required to be implemented. @@ -931,6 +941,11 @@ struct ImplConstraint { // A constraint that requires an intrinsic property of a type. struct IntrinsicConstraint { + template + auto Decompose(F f) const { + return f(type, kind, arguments); + } + // Print the intrinsic constraint. void Print(llvm::raw_ostream& out) const; @@ -951,6 +966,11 @@ struct IntrinsicConstraint { // A constraint that a collection of values are known to be the same. struct EqualityConstraint { + template + auto Decompose(F f) const { + return f(values); + } + // Visit the values in this equality constraint that are a single step away // from the given value according to this equality constraint. That is: if // `value` is identical to a value in `values`, then call the visitor on all @@ -969,6 +989,12 @@ struct EqualityConstraint { // A constraint indicating that access to an associated constant should be // replaced by another value. struct RewriteConstraint { + template + auto Decompose(F f) const { + return f(constant, unconverted_replacement, unconverted_replacement_type, + converted_replacement); + } + // The associated constant value that is rewritten. Nonnull constant; // The replacement in its original type. @@ -981,6 +1007,11 @@ struct RewriteConstraint { // A context in which we might look up a name. struct LookupContext { + template + auto Decompose(F f) const { + return f(context); + } + Nonnull context; }; diff --git a/explorer/ast/value_node.h b/explorer/ast/value_node.h index dd7ad58057a5..114c108773f3 100644 --- a/explorer/ast/value_node.h +++ b/explorer/ast/value_node.h @@ -10,6 +10,7 @@ #include #include "explorer/ast/ast_node.h" +#include "explorer/ast/clone_context.h" #include "explorer/ast/value_category.h" #include "explorer/common/nonnull.h" @@ -84,6 +85,15 @@ class ValueNodeView { return llvm::cast(base).value_category(); }) {} + explicit ValueNodeView(CloneContext& context, const ValueNodeView& other) + : base_(context.Remap(other.base_)), + // We assume the clone is the same kind of node as the original. + constant_value_(other.constant_value_), + symbolic_identity_(other.symbolic_identity_), + print_(other.print_), + static_type_(other.static_type_), + value_category_(other.value_category_) {} + ValueNodeView(const ValueNodeView&) = default; ValueNodeView(ValueNodeView&&) = default; auto operator=(const ValueNodeView&) -> ValueNodeView& = default; diff --git a/explorer/ast/value_transform.h b/explorer/ast/value_transform.h index 58f1078d5721..e01735b49c65 100644 --- a/explorer/ast/value_transform.h +++ b/explorer/ast/value_transform.h @@ -10,14 +10,22 @@ namespace Carbon { +template +constexpr bool is_list_constructible_impl = false; + +template +constexpr bool is_list_constructible_impl< + T, decltype(T{std::declval()...}), Args...> = true; + // A no-op visitor used to implement `IsRecursivelyTransformable`. The // `operator()` function returns `true_type` if it's called with arguments that -// can be used to construct `T`, and `false_type` otherwise. +// can be used to direct-list-initialize `T`, and `false_type` otherwise. template struct IsRecursivelyTransformableVisitor { template auto operator()(Args&&... args) - -> std::integral_constant>; + -> std::integral_constant>; }; // A type trait that indicates whether `T` is transformable. A transformable @@ -35,8 +43,61 @@ constexpr bool IsRecursivelyTransformable< T, decltype(std::declval().Decompose( IsRecursivelyTransformableVisitor{}))> = true; +// Unwrapper for the case where there's nothing to unwrap. +class NoOpUnwrapper { + public: + template + auto UnwrapOr(T&& value, const U&) -> T { + return std::forward(value); + } + + template + auto Wrap(T&& value) -> T&& { + return std::forward(value); + } + + constexpr bool failed() const { return false; } +}; + +// Helper to temporarily unwrap the ErrorOr around a value, and then put it +// back when we're done with the overall computation. +class ErrorUnwrapper { + public: + // Unwrap the `ErrorOr` from the given value, or collect the error and return + // the given fallback value on failure. + template + auto UnwrapOr(ErrorOr value, const U& fallback) -> T { + if (!value.ok()) { + status_ = std::move(value).error(); + return fallback; + } + return std::move(*value); + } + template + auto UnwrapOr(T&& value, const U&) -> T { + return std::forward(value); + } + + // Wrap the given value into `ErrorOr`, returning our collected error if any, + // or the given value if we succeeded. + template + auto Wrap(T&& value) -> ErrorOr { + if (!status_.ok()) { + Error error = std::move(status_).error(); + status_ = Success(); + return error; + } + return std::forward(value); + } + + bool failed() const { return !status_.ok(); } + + private: + ErrorOr status_ = Success(); +}; + // Base class for transforms of visitable data types. -template +template class TransformBase { public: explicit TransformBase(Nonnull arena) : arena_(arena) {} @@ -44,45 +105,21 @@ class TransformBase { // Transform the given value, and produce either the transformed value or an // error. template - auto Transform(const T& v) -> ErrorOr { - auto result = TransformOrOriginal(v); - if (!status_.ok()) { - Error error = std::move(status_).error(); - status_ = Success(); - return error; - } - return result; + auto Transform(const T& v) -> decltype(auto) { + return unwrapper_.Wrap(TransformOrOriginal(v)); } protected: - // Given an original value and the result of calling `operator()`, find the - // transformed value we should use. - // - // If `operator()` returns `ErrorOr`, then on failure, collect the error - // and return the untransformed value; otherwise, return the transformed - // value. - template - auto CollectError(const T& /*original*/, const U& transformed) -> U { - return transformed; - } - template - auto CollectError(const T& original, ErrorOr transformed) -> U { - if (!transformed.ok()) { - status_ = std::move(transformed).error(); - return original; - } - return std::move(*transformed); - } - // Transform the given value, or return the original if transformation fails. template auto TransformOrOriginal(const T& v) - -> decltype(CollectError(v, std::declval()(v))) { + -> decltype(std::declval().UnwrapOr( + std::declval()(v), v)) { // If we've already failed, don't do any more transformations. - if (!status_.ok()) { + if (unwrapper_.failed()) { return v; } - return CollectError(v, static_cast(*this)(v)); + return unwrapper_.UnwrapOr(static_cast(*this)(v), v); } // Transformable values are recursively transformed by default. @@ -91,10 +128,10 @@ class TransformBase { auto operator()(const T& value) -> T { return value.Decompose([&](const auto&... elements) { return [&](auto&&... transformed_elements) { - if (status_.ok()) { - return T{decltype(transformed_elements)(transformed_elements)...}; + if (unwrapper_.failed()) { + return value; } - return value; + return T{decltype(transformed_elements)(transformed_elements)...}; }(TransformOrOriginal(elements)...); }); } @@ -109,11 +146,11 @@ class TransformBase { -> decltype(AllocateTrait::New( arena_, decltype(transformed_elements)(transformed_elements)...)) { - if (status_.ok()) { - return AllocateTrait::New( - arena_, decltype(transformed_elements)(transformed_elements)...); + if (unwrapper_.failed()) { + return value; } - return value; + return AllocateTrait::New( + arena_, decltype(transformed_elements)(transformed_elements)...); }(TransformOrOriginal(elements)...); }); } @@ -176,17 +213,17 @@ class TransformBase { private: Nonnull arena_; - // Temporary storage for an error that was produced during transformation - // that has not yet been handed back to the caller. - ErrorOr status_ = Success(); + // Unwrapper for results. Used to remove an ErrorOr<...> wrapper temporarily + // during recursive transformations and re-apply it when we're done. + ResultUnwrapper unwrapper_; }; // Base class for transforms of `Value`s. -template -class ValueTransform : public TransformBase { +template +class ValueTransform : public TransformBase { public: - using TransformBase::TransformBase; - using TransformBase::operator(); + using TransformBase::TransformBase; + using TransformBase::operator(); // Leave references to AST nodes alone by default. // The 'int = 0' parameter avoids this function hiding the `operator()(const @@ -247,6 +284,11 @@ class ValueTransform : public TransformBase { -> Nonnull { return value_ptr; } + + // Preserve constraint kind for intrinsic constraints. + auto operator()(IntrinsicConstraint::Kind kind) -> IntrinsicConstraint::Kind { + return kind; + } }; } // namespace Carbon diff --git a/explorer/common/arena.h b/explorer/common/arena.h index 3e5d1f4b6843..ee7950c0893d 100644 --- a/explorer/common/arena.h +++ b/explorer/common/arena.h @@ -15,6 +15,17 @@ namespace Carbon { class Arena { public: + // Values of this type can be passed as the first argument to New in order to + // have the address of the created object written to the given pointer before + // the constructor is run. This is used during cloning to support pointer + // cycles within the AST. + template + struct WriteAddressTo { + Nonnull target; + }; + template + WriteAddressTo(T** target) -> WriteAddressTo; + // Allocates an object in the arena, returning a pointer to it. template < typename T, typename... Args, @@ -27,6 +38,15 @@ class Arena { return ptr; } + // Allocates an object in the arena, writing its address to the given pointer. + template < + typename T, typename U, typename... Args, + typename std::enable_if_t>* = nullptr> + void New(WriteAddressTo addr, Args&&... args) { + arena_.push_back(std::make_unique>( + addr, std::forward(args)...)); + } + private: // Virtualizes arena entries so that a single vector can contain many types, // avoiding templated statics. @@ -39,10 +59,22 @@ class Arena { template class ArenaEntryTyped : public ArenaEntry { public: + struct WriteAddressHelper {}; + template explicit ArenaEntryTyped(Args&&... args) : instance_(std::forward(args)...) {} + template + explicit ArenaEntryTyped(WriteAddressHelper, Args&&... args) + : ArenaEntryTyped(std::forward(args)...) {} + + template + explicit ArenaEntryTyped(WriteAddressTo write_address, Args&&... args) + : ArenaEntryTyped( + (*write_address.target = &instance_, WriteAddressHelper{}), + std::forward(args)...) {} + auto Instance() -> Nonnull { return Nonnull(&instance_); } private: diff --git a/explorer/fuzzing/BUILD b/explorer/fuzzing/BUILD index 39f9642a599e..8feef0cfe675 100644 --- a/explorer/fuzzing/BUILD +++ b/explorer/fuzzing/BUILD @@ -116,6 +116,28 @@ cc_test( ], ) +cc_test( + name = "clone_test", + srcs = ["clone_test.cpp"], + args = [ + "$(locations //explorer:standard_libraries)", + "$(locations //explorer/testdata:carbon_files)", + ], + data = [ + "//explorer:standard_libraries", + "//explorer/testdata:carbon_files", + ], + deps = [ + ":ast_to_proto_lib", + "//common/fuzzing:carbon_cc_proto", + "//explorer/ast", + "//explorer/syntax", + "@com_google_googletest//:gtest", + "@com_google_protobuf//:protobuf_headers", + "@llvm-project//llvm:Support", + ], +) + cc_fuzz_test( name = "explorer_fuzzer", testonly = 1, diff --git a/explorer/fuzzing/ast_to_proto.cpp b/explorer/fuzzing/ast_to_proto.cpp index 0c3e2074981e..93e596336229 100644 --- a/explorer/fuzzing/ast_to_proto.cpp +++ b/explorer/fuzzing/ast_to_proto.cpp @@ -823,6 +823,9 @@ static auto DeclarationToProto(const Declaration& declaration) impl_proto->set_kind(Fuzzing::ImplDeclaration::ExternalImpl); break; } + for (Nonnull binding : impl.deduced_parameters()) { + *impl_proto->add_deduced_parameters() = GenericBindingToProto(*binding); + } *impl_proto->mutable_impl_type() = ExpressionToProto(*impl.impl_type()); *impl_proto->mutable_interface() = ExpressionToProto(impl.interface()); for (const auto& member : impl.members()) { diff --git a/explorer/fuzzing/clone_test.cpp b/explorer/fuzzing/clone_test.cpp new file mode 100644 index 000000000000..c4da163dfb73 --- /dev/null +++ b/explorer/fuzzing/clone_test.cpp @@ -0,0 +1,73 @@ +// 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 + +#include +#include +#include + +#include "explorer/ast/clone_context.h" +#include "explorer/fuzzing/ast_to_proto.h" +#include "explorer/syntax/parse.h" + +namespace Carbon::Testing { +namespace { + +static std::vector* carbon_files = nullptr; + +auto CloneAST(Arena& arena, const AST& ast) -> AST { + CloneContext context(&arena); + return { + .package = ast.package, + .is_api = ast.is_api, + .imports = ast.imports, + .declarations = context.Clone(ast.declarations), + .main_call = context.Clone(ast.main_call), + .num_prelude_declarations = ast.num_prelude_declarations, + }; +} + +// Returns a string representation of `ast`. +auto AstToString(const AST& ast) -> std::string { + std::string s; + llvm::raw_string_ostream out(s); + out << "package " << ast.package.package << (ast.is_api ? "api" : "impl") + << ";\n"; + for (auto* declaration : ast.declarations) { + out << *declaration << "\n"; + } + return s; +} + +TEST(CloneTest, SameProtoAfterClone) { + int parsed_ok_count = 0; + for (const llvm::StringRef f : *carbon_files) { + Carbon::Arena arena; + const ErrorOr ast = Carbon::Parse(&arena, f, /*parser_debug=*/false); + if (ast.ok()) { + ++parsed_ok_count; + const AST clone = CloneAST(arena, *ast); + const Fuzzing::CompilationUnit orig_proto = AstToProto(*ast); + const Fuzzing::CompilationUnit clone_proto = AstToProto(clone); + // TODO: Use EqualsProto once it's available. + EXPECT_TRUE(google::protobuf::util::MessageDifferencer::Equals( + orig_proto, clone_proto)) + << "clone produced a different AST. original:\n" + << AstToString(*ast) << "clone:\n" + << AstToString(clone); + } + } + // Makes sure files were actually processed. + EXPECT_GT(parsed_ok_count, 0); +} + +} // namespace +} // namespace Carbon::Testing + +auto main(int argc, char** argv) -> int { + ::testing::InitGoogleTest(&argc, argv); + // gtest should remove flags, leaving just input files. + std::vector carbon_files(&argv[1], &argv[argc]); + Carbon::Testing::carbon_files = &carbon_files; + return RUN_ALL_TESTS(); +} diff --git a/explorer/gen_rtti.py b/explorer/gen_rtti.py index 374c6870d86f..0574b1155516 100755 --- a/explorer/gen_rtti.py +++ b/explorer/gen_rtti.py @@ -48,6 +48,9 @@ numeric value, so you can use `static_cast` to convert between the enum types for different classes that have a common root, so long as the enumerator value is present in both types. As a result, `InheritsFromFoo` can be used to determine whether casting to `FooKind` is safe. + +This also generates the dispatch function for the clone constructor, which +behaves like a virtual function. """ __copyright__ = """ @@ -185,7 +188,7 @@ def main() -> None: input_filename = sys.argv[1] header_filename = sys.argv[2] cpp_filename = sys.argv[3] - relative_header = sys.argv[4] + ast_headers = sys.argv[4:] with open(input_filename) as file: lines = file.readlines() @@ -282,6 +285,12 @@ def main() -> None: ) print("}\n") + print("class Arena;") + print("class AstNode;") + print("class CloneContext;") + print("void CloneImpl(Arena& arena, CloneContext& context,") + print(" const AstNode& node, AstNode** result);\n") + print("} // namespace Carbon\n") print(f"#endif // {guard_macro}") @@ -291,7 +300,8 @@ def main() -> None: sys.stdout = cpp_file print(f"// Generated from {input_filename} by explorer/gen_rtti.py\n") - print(f'#include "{relative_header}"') + for h in ast_headers: + print(f'#include "{h}"') print("\nnamespace Carbon {\n") for node in classes.values(): if node.kind != Class.Kind.CONCRETE: @@ -308,6 +318,20 @@ def main() -> None: print(" }") print("}\n") + print("void CloneImpl(Arena& arena, CloneContext& context,") + print(" const AstNode& node, AstNode** result) {") + print(" switch(node.kind()) {") + for node in classes.values(): + if node.kind == Class.Kind.CONCRETE and node.Root().name == "AstNode": + print(f" case AstNodeKind::{node.name}:") + print( + f" return arena.New<{node.name}>(" + + "Arena::WriteAddressTo{result}, context, " + + f"static_cast(node));" + ) + print(" }") + print("}\n") + print("} // namespace Carbon\n") cpp_file.close() diff --git a/explorer/interpreter/type_checker.cpp b/explorer/interpreter/type_checker.cpp index c72966558a6b..fdff55b889ac 100644 --- a/explorer/interpreter/type_checker.cpp +++ b/explorer/interpreter/type_checker.cpp @@ -2000,7 +2000,7 @@ auto TypeChecker::RebuildValue(Nonnull value) const } class TypeChecker::SubstituteTransform - : public ValueTransform { + : public ValueTransform { public: SubstituteTransform(Nonnull type_checker, const Bindings& bindings)