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)