diff --git a/common/fuzzing/carbon.proto b/common/fuzzing/carbon.proto index 76437d8062d6..998ad9c37d89 100644 --- a/common/fuzzing/carbon.proto +++ b/common/fuzzing/carbon.proto @@ -142,6 +142,11 @@ message BindingPattern { optional Pattern type = 2; } +message GenericBinding { + optional string name = 1; + optional Expression type = 2; +} + message TuplePattern { repeated Pattern fields = 1; } @@ -170,6 +175,7 @@ message Pattern { ExpressionPattern expression_pattern = 4; AutoPattern auto_pattern = 5; VarPattern var_pattern = 6; + GenericBinding generic_binding = 7; } } @@ -264,11 +270,6 @@ message ReturnTerm { optional Expression type = 2; } -message GenericBinding { - optional string name = 1; - optional Expression type = 2; -} - message FunctionDeclaration { optional string name = 1; repeated GenericBinding deduced_parameters = 2; diff --git a/executable_semantics/ast/BUILD b/executable_semantics/ast/BUILD index 24572db010db..ea254ae4ac6a 100644 --- a/executable_semantics/ast/BUILD +++ b/executable_semantics/ast/BUILD @@ -70,12 +70,13 @@ cc_test( ) cc_library( - name = "generic_binding", + name = "impl_binding", hdrs = [ - "generic_binding.h", + "impl_binding.h", ], deps = [ ":ast_node", + ":pattern", ":source_location", ":value_category", "//common:check", @@ -93,7 +94,7 @@ cc_library( ], deps = [ ":ast_node", - ":generic_binding", + ":impl_binding", ":pattern", ":return_term", ":source_location", @@ -125,7 +126,6 @@ cc_library( hdrs = ["expression.h"], deps = [ ":ast_node", - ":generic_binding", ":paren_contents", ":source_location", ":static_scope", diff --git a/executable_semantics/ast/ast_node.h b/executable_semantics/ast/ast_node.h index 1757b8b341a6..ed5fdfc4b2b3 100644 --- a/executable_semantics/ast/ast_node.h +++ b/executable_semantics/ast/ast_node.h @@ -43,7 +43,10 @@ class AstNode { auto operator=(AstNode&&) -> AstNode& = delete; virtual ~AstNode() = 0; + // Print the AST rooted at the node. virtual void Print(llvm::raw_ostream& out) const = 0; + // Print identifying information about the node, such as it's name. + virtual void PrintID(llvm::raw_ostream& out) const = 0; LLVM_DUMP_METHOD void Dump() const { Print(llvm::errs()); } // Returns an enumerator specifying the concrete type of this node. diff --git a/executable_semantics/ast/ast_rtti.txt b/executable_semantics/ast/ast_rtti.txt index 7ff4206b2080..54bce5fd1ebe 100644 --- a/executable_semantics/ast/ast_rtti.txt +++ b/executable_semantics/ast/ast_rtti.txt @@ -7,6 +7,7 @@ abstract class Pattern : AstNode; class AutoPattern : Pattern; class VarPattern : Pattern; class BindingPattern : Pattern; + class GenericBinding : Pattern; class TuplePattern : Pattern; class AlternativePattern : Pattern; class ExpressionPattern : Pattern; @@ -17,7 +18,6 @@ abstract class Declaration : AstNode; class VariableDeclaration : Declaration; class InterfaceDeclaration : Declaration; class ImplDeclaration : Declaration; -class GenericBinding : AstNode; class ImplBinding : AstNode; class AlternativeSignature : AstNode; abstract class Statement : AstNode; diff --git a/executable_semantics/ast/declaration.cpp b/executable_semantics/ast/declaration.cpp index 1a316058e44e..7b4074fc5adf 100644 --- a/executable_semantics/ast/declaration.cpp +++ b/executable_semantics/ast/declaration.cpp @@ -17,7 +17,8 @@ void Declaration::Print(llvm::raw_ostream& out) const { switch (kind()) { case DeclarationKind::InterfaceDeclaration: { const auto& iface_decl = cast(*this); - out << "interface " << iface_decl.name() << " {\n"; + PrintID(out); + out << " {\n"; for (Nonnull m : iface_decl.members()) { out << *m; } @@ -26,15 +27,8 @@ void Declaration::Print(llvm::raw_ostream& out) const { } case DeclarationKind::ImplDeclaration: { const auto& impl_decl = cast(*this); - switch (impl_decl.kind()) { - case ImplKind::InternalImpl: - break; - case ImplKind::ExternalImpl: - out << "external "; - break; - } - out << "impl " << *impl_decl.impl_type() << " as " - << impl_decl.interface() << " {\n"; + PrintID(out); + out << " {\n"; for (Nonnull m : impl_decl.members()) { out << *m; } @@ -47,7 +41,11 @@ void Declaration::Print(llvm::raw_ostream& out) const { case DeclarationKind::ClassDeclaration: { const auto& class_decl = cast(*this); - out << "class " << class_decl.name() << " {\n"; + PrintID(out); + if (class_decl.type_params().has_value()) { + out << **class_decl.type_params(); + } + out << " {\n"; for (Nonnull m : class_decl.members()) { out << *m; } @@ -57,7 +55,8 @@ void Declaration::Print(llvm::raw_ostream& out) const { case DeclarationKind::ChoiceDeclaration: { const auto& choice = cast(*this); - out << "choice " << choice.name() << " {\n"; + PrintID(out); + out << " {\n"; for (Nonnull alt : choice.alternatives()) { out << *alt << ";\n"; } @@ -67,7 +66,7 @@ void Declaration::Print(llvm::raw_ostream& out) const { case DeclarationKind::VariableDeclaration: { const auto& var = cast(*this); - out << "var " << var.binding(); + PrintID(out); if (var.has_initializer()) { out << " = " << var.initializer(); } @@ -77,6 +76,50 @@ void Declaration::Print(llvm::raw_ostream& out) const { } } +void Declaration::PrintID(llvm::raw_ostream& out) const { + switch (kind()) { + case DeclarationKind::InterfaceDeclaration: { + const auto& iface_decl = cast(*this); + out << "interface" << iface_decl.name(); + break; + } + case DeclarationKind::ImplDeclaration: { + const auto& impl_decl = cast(*this); + switch (impl_decl.kind()) { + case ImplKind::InternalImpl: + break; + case ImplKind::ExternalImpl: + out << "external "; + break; + } + out << "impl " << *impl_decl.impl_type() << " as " + << impl_decl.interface(); + break; + } + case DeclarationKind::FunctionDeclaration: + out << "fn " << cast(*this).name(); + break; + + case DeclarationKind::ClassDeclaration: { + const auto& class_decl = cast(*this); + out << "class " << class_decl.name(); + break; + } + + case DeclarationKind::ChoiceDeclaration: { + const auto& choice = cast(*this); + out << "choice " << choice.name(); + break; + } + + case DeclarationKind::VariableDeclaration: { + const auto& var = cast(*this); + out << "var " << var.binding(); + break; + } + } +} + auto GetName(const Declaration& declaration) -> std::optional { switch (declaration.kind()) { case DeclarationKind::FunctionDeclaration: @@ -98,6 +141,8 @@ 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: @@ -170,4 +215,8 @@ void AlternativeSignature::Print(llvm::raw_ostream& out) const { out << "alt " << name() << " " << signature(); } +void AlternativeSignature::PrintID(llvm::raw_ostream& out) const { + out << name(); +} + } // namespace Carbon diff --git a/executable_semantics/ast/declaration.h b/executable_semantics/ast/declaration.h index b83b1d036c53..34c1b66c8184 100644 --- a/executable_semantics/ast/declaration.h +++ b/executable_semantics/ast/declaration.h @@ -11,7 +11,7 @@ #include "common/ostream.h" #include "executable_semantics/ast/ast_node.h" -#include "executable_semantics/ast/generic_binding.h" +#include "executable_semantics/ast/impl_binding.h" #include "executable_semantics/ast/pattern.h" #include "executable_semantics/ast/return_term.h" #include "executable_semantics/ast/source_location.h" @@ -40,6 +40,7 @@ class Declaration : public AstNode { auto operator=(const Declaration&) -> Declaration& = delete; void Print(llvm::raw_ostream& out) const override; + void PrintID(llvm::raw_ostream& out) const override; static auto classof(const AstNode* node) -> bool { return InheritsFromDeclaration(node->kind()); @@ -67,6 +68,23 @@ class Declaration : public AstNode { // and after typechecking it's guaranteed to be true. auto has_static_type() const -> bool { return static_type_.has_value(); } + // Sets the value returned by constant_value(). Can only be called once, + // during typechecking. + void set_constant_value(Nonnull value) { + CHECK(!constant_value_.has_value()); + constant_value_ = value; + } + + // See static_scope.h for API. + auto constant_value() const -> std::optional> { + return constant_value_; + } + + // See static_scope.h for API. + auto symbolic_identity() const -> std::optional> { + return constant_value_; + } + protected: // Constructs a Declaration representing syntax at the given line number. // `kind` must be the enumerator corresponding to the most-derived type being @@ -76,6 +94,7 @@ class Declaration : public AstNode { private: std::optional> static_type_; + std::optional> constant_value_; }; class FunctionDeclaration : public Declaration { @@ -130,16 +149,6 @@ class FunctionDeclaration : public Declaration { auto body() -> std::optional> { return body_; } auto value_category() const -> ValueCategory { return ValueCategory::Let; } - auto constant_value() const -> std::optional> { - return constant_value_; - } - - // Sets the value returned by constant_value(). Can only be called once, - // during typechecking. - void set_constant_value(Nonnull value) { - CHECK(!constant_value_.has_value()); - constant_value_ = value; - } auto is_method() const -> bool { return me_pattern_.has_value(); } @@ -150,7 +159,6 @@ class FunctionDeclaration : public Declaration { Nonnull param_pattern_; ReturnTerm return_term_; std::optional> body_; - std::optional> constant_value_; }; class ClassDeclaration : public Declaration { @@ -158,9 +166,11 @@ class ClassDeclaration : public Declaration { using ImplementsCarbonValueNode = void; ClassDeclaration(SourceLocation source_loc, std::string name, + std::optional> type_params, std::vector> members) : Declaration(AstNodeKind::ClassDeclaration, source_loc), name_(std::move(name)), + type_params_(type_params), members_(std::move(members)) {} static auto classof(const AstNode* node) -> bool { @@ -168,26 +178,23 @@ class ClassDeclaration : public Declaration { } auto name() const -> const std::string& { return name_; } + auto type_params() const -> std::optional> { + return type_params_; + } + auto type_params() -> std::optional> { + return type_params_; + } + auto members() const -> llvm::ArrayRef> { return members_; } auto value_category() const -> ValueCategory { return ValueCategory::Let; } - auto constant_value() const -> std::optional> { - return constant_value_; - } - - // Sets the value returned by constant_value(). Can only be called once, - // during typechecking. - void set_constant_value(Nonnull value) { - CHECK(!constant_value_.has_value()); - constant_value_ = value; - } private: std::string name_; + std::optional> type_params_; std::vector> members_; - std::optional> constant_value_; }; class AlternativeSignature : public AstNode { @@ -199,6 +206,7 @@ class AlternativeSignature : public AstNode { signature_(signature) {} void Print(llvm::raw_ostream& out) const override; + void PrintID(llvm::raw_ostream& out) const override; static auto classof(const AstNode* node) -> bool { return InheritsFromAlternativeSignature(node->kind()); @@ -237,21 +245,10 @@ class ChoiceDeclaration : public Declaration { } auto value_category() const -> ValueCategory { return ValueCategory::Let; } - auto constant_value() const -> std::optional> { - return constant_value_; - } - - // Sets the value returned by constant_value(). Can only be called once, - // during typechecking. - void set_constant_value(Nonnull value) { - CHECK(!constant_value_.has_value()); - constant_value_ = value; - } private: std::string name_; std::vector> alternatives_; - std::optional> constant_value_; }; // Global variable definition implements the Declaration concept. @@ -311,21 +308,10 @@ class InterfaceDeclaration : public Declaration { auto self() -> Nonnull { return self_; } auto value_category() const -> ValueCategory { return ValueCategory::Let; } - auto constant_value() const -> std::optional> { - return constant_value_; - } - - // Sets the value returned by constant_value(). Can only be called once, - // during typechecking. - void set_constant_value(Nonnull value) { - CHECK(!constant_value_.has_value()); - constant_value_ = value; - } private: std::string name_; std::vector> members_; - std::optional> constant_value_; Nonnull self_; }; @@ -364,14 +350,6 @@ class ImplDeclaration : public Declaration { auto members() const -> llvm::ArrayRef> { return members_; } - // Return the witness table for this impl. - auto constant_value() const -> std::optional> { - return constant_value_; - } - void set_constant_value(Nonnull value) { - CHECK(!constant_value_.has_value()); - constant_value_ = value; - } auto value_category() const -> ValueCategory { return ValueCategory::Let; } private: @@ -380,7 +358,6 @@ class ImplDeclaration : public Declaration { Nonnull interface_; std::optional> interface_type_; std::vector> members_; - std::optional> constant_value_; }; // Return the name of a declaration, if it has one. diff --git a/executable_semantics/ast/expression.cpp b/executable_semantics/ast/expression.cpp index 64f6bf969907..8c2ae1c613e8 100644 --- a/executable_semantics/ast/expression.cpp +++ b/executable_semantics/ast/expression.cpp @@ -116,12 +116,6 @@ void Expression::Print(llvm::raw_ostream& out) const { PrintFields(out, cast(*this).fields(), ": "); out << "}"; break; - case ExpressionKind::IntLiteral: - out << cast(*this).value(); - break; - case ExpressionKind::BoolLiteral: - out << (cast(*this).value() ? "true" : "false"); - break; case ExpressionKind::PrimitiveOperatorExpression: { out << "("; const auto& op = cast(*this); @@ -142,9 +136,6 @@ void Expression::Print(llvm::raw_ostream& out) const { out << ")"; break; } - case ExpressionKind::IdentifierExpression: - out << cast(*this).name(); - break; case ExpressionKind::CallExpression: { const auto& call = cast(*this); out << call.function(); @@ -155,26 +146,6 @@ void Expression::Print(llvm::raw_ostream& out) const { } break; } - case ExpressionKind::BoolTypeLiteral: - out << "Bool"; - break; - case ExpressionKind::IntTypeLiteral: - out << "i32"; - break; - case ExpressionKind::StringLiteral: - out << "\""; - out.write_escaped(cast(*this).value()); - out << "\""; - break; - case ExpressionKind::StringTypeLiteral: - out << "String"; - break; - case ExpressionKind::TypeTypeLiteral: - out << "Type"; - break; - case ExpressionKind::ContinuationTypeLiteral: - out << "Continuation"; - break; case ExpressionKind::FunctionTypeLiteral: { const auto& fn = cast(*this); out << "fn " << fn.parameter() << " -> " << fn.return_type(); @@ -205,6 +176,64 @@ void Expression::Print(llvm::raw_ostream& out) const { out << ")"; break; } + case ExpressionKind::IdentifierExpression: + case ExpressionKind::IntLiteral: + case ExpressionKind::BoolLiteral: + case ExpressionKind::BoolTypeLiteral: + case ExpressionKind::IntTypeLiteral: + case ExpressionKind::StringLiteral: + case ExpressionKind::StringTypeLiteral: + case ExpressionKind::TypeTypeLiteral: + case ExpressionKind::ContinuationTypeLiteral: + PrintID(out); + break; + } +} + +void Expression::PrintID(llvm::raw_ostream& out) const { + switch (kind()) { + case ExpressionKind::IdentifierExpression: + out << cast(*this).name(); + break; + case ExpressionKind::IntLiteral: + out << cast(*this).value(); + break; + case ExpressionKind::BoolLiteral: + out << (cast(*this).value() ? "true" : "false"); + break; + case ExpressionKind::BoolTypeLiteral: + out << "Bool"; + break; + case ExpressionKind::IntTypeLiteral: + out << "i32"; + break; + case ExpressionKind::StringLiteral: + out << "\""; + out.write_escaped(cast(*this).value()); + out << "\""; + break; + case ExpressionKind::StringTypeLiteral: + out << "String"; + break; + case ExpressionKind::TypeTypeLiteral: + out << "Type"; + break; + case ExpressionKind::ContinuationTypeLiteral: + out << "Continuation"; + break; + case ExpressionKind::IndexExpression: + case ExpressionKind::FieldAccessExpression: + case ExpressionKind::IfExpression: + case ExpressionKind::TupleLiteral: + case ExpressionKind::StructLiteral: + case ExpressionKind::StructTypeLiteral: + case ExpressionKind::CallExpression: + case ExpressionKind::PrimitiveOperatorExpression: + case ExpressionKind::IntrinsicExpression: + case ExpressionKind::UnimplementedExpression: + case ExpressionKind::FunctionTypeLiteral: + out << "..."; + break; } } diff --git a/executable_semantics/ast/expression.h b/executable_semantics/ast/expression.h index 80e362aed781..c47e68da4fa1 100644 --- a/executable_semantics/ast/expression.h +++ b/executable_semantics/ast/expression.h @@ -5,6 +5,7 @@ #ifndef EXECUTABLE_SEMANTICS_AST_EXPRESSION_H_ #define EXECUTABLE_SEMANTICS_AST_EXPRESSION_H_ +#include #include #include #include @@ -12,7 +13,6 @@ #include "common/ostream.h" #include "executable_semantics/ast/ast_node.h" -#include "executable_semantics/ast/generic_binding.h" #include "executable_semantics/ast/paren_contents.h" #include "executable_semantics/ast/source_location.h" #include "executable_semantics/ast/static_scope.h" @@ -25,12 +25,14 @@ namespace Carbon { class Value; class VariableType; +class ImplBinding; class Expression : public AstNode { public: ~Expression() override = 0; void Print(llvm::raw_ostream& out) const override; + void PrintID(llvm::raw_ostream& out) const override; static auto classof(const AstNode* node) { return InheritsFromExpression(node->kind()); @@ -43,7 +45,10 @@ class Expression : public AstNode { } // The static type of this expression. Cannot be called before typechecking. - auto static_type() const -> const Value& { return **static_type_; } + auto static_type() const -> const Value& { + CHECK(static_type_.has_value()); + return **static_type_; + } // Sets the static type of this expression. Can only be called once, during // typechecking. @@ -354,7 +359,10 @@ class PrimitiveOperatorExpression : public Expression { std::vector> arguments_; }; -class ImplBinding; +class GenericBinding; + +using BindingMap = + std::map, Nonnull>; class CallExpression : public Expression { public: @@ -390,10 +398,17 @@ class CallExpression : public Expression { impls_ = impls; } + auto deduced_args() const -> const BindingMap& { return deduced_args_; } + + void set_deduced_args(const BindingMap& deduced_args) { + deduced_args_ = deduced_args; + } + private: Nonnull function_; Nonnull argument_; std::map, ValueNodeView> impls_; + BindingMap deduced_args_; }; class FunctionTypeLiteral : public Expression { diff --git a/executable_semantics/ast/generic_binding.h b/executable_semantics/ast/impl_binding.h similarity index 52% rename from executable_semantics/ast/generic_binding.h rename to executable_semantics/ast/impl_binding.h index a67a94551fcb..8790c662412e 100644 --- a/executable_semantics/ast/generic_binding.h +++ b/executable_semantics/ast/impl_binding.h @@ -10,6 +10,7 @@ #include "common/check.h" #include "common/ostream.h" #include "executable_semantics/ast/ast_node.h" +#include "executable_semantics/ast/pattern.h" #include "executable_semantics/ast/value_category.h" namespace Carbon { @@ -18,71 +19,6 @@ class Value; class Expression; class ImplBinding; -// TODO: expand the kinds of things that can be deduced parameters. -// For now, only generic parameters are supported. -class GenericBinding : public AstNode { - public: - using ImplementsCarbonValueNode = void; - - GenericBinding(SourceLocation source_loc, std::string name, - Nonnull type) - : AstNode(AstNodeKind::GenericBinding, source_loc), - name_(std::move(name)), - type_(type) {} - - void Print(llvm::raw_ostream& out) const override; - - static auto classof(const AstNode* node) -> bool { - return InheritsFromGenericBinding(node->kind()); - } - - auto name() const -> const std::string& { return name_; } - auto type() const -> const Expression& { return *type_; } - auto type() -> Expression& { return *type_; } - - // The static type of the binding. Cannot be called before typechecking. - auto static_type() const -> const Value& { return **static_type_; } - - // Sets the static type of the binding. Can only be called once, during - // typechecking. - void set_static_type(Nonnull type) { - CHECK(!static_type_.has_value()); - static_type_ = type; - } - - auto value_category() const -> ValueCategory { return ValueCategory::Let; } - auto constant_value() const -> std::optional> { - return constant_value_; - } - - // Sets the value returned by constant_value(). Can only be called once, - // during typechecking. - void set_constant_value(Nonnull value) { - CHECK(!constant_value_.has_value()); - constant_value_ = value; - } - - // The impl binding associated with this type variable. - auto impl_binding() const -> std::optional> { - return impl_binding_; - } - // Set the impl binding. - void set_impl_binding(Nonnull binding) { - CHECK(!impl_binding_.has_value()); - impl_binding_ = binding; - } - - private: - std::string name_; - Nonnull type_; - std::optional> static_type_; - std::optional> constant_value_; - std::optional> impl_binding_; -}; - -using BindingMap = - std::map, Nonnull>; - // The run-time counterpart of a `GenericBinding`. // // Once a generic binding has been declared, it can be used @@ -106,6 +42,7 @@ class ImplBinding : public AstNode { return InheritsFromImplBinding(node->kind()); } void Print(llvm::raw_ostream& out) const override; + void PrintID(llvm::raw_ostream& out) const override; // The binding for the type variable. auto type_var() const -> Nonnull { return type_var_; } @@ -116,6 +53,9 @@ class ImplBinding : public AstNode { auto constant_value() const -> std::optional> { return std::nullopt; } + auto symbolic_identity() const -> std::optional> { + return std::nullopt; + } // The static type of the impl. Cannot be called before typechecking. auto static_type() const -> const Value& { return **static_type_; } diff --git a/executable_semantics/ast/pattern.cpp b/executable_semantics/ast/pattern.cpp index 20cf9cab963e..41efd10a5385 100644 --- a/executable_semantics/ast/pattern.cpp +++ b/executable_semantics/ast/pattern.cpp @@ -29,6 +29,11 @@ void Pattern::Print(llvm::raw_ostream& out) const { out << binding.name() << ": " << binding.type(); break; } + case PatternKind::GenericBinding: { + const auto& binding = cast(*this); + out << binding.name() << ":! " << binding.type(); + break; + } case PatternKind::TuplePattern: { const auto& tuple = cast(*this); out << "("; @@ -54,6 +59,40 @@ void Pattern::Print(llvm::raw_ostream& out) const { } } +void Pattern::PrintID(llvm::raw_ostream& out) const { + switch (kind()) { + case PatternKind::AutoPattern: + out << "auto"; + break; + case PatternKind::BindingPattern: { + const auto& binding = cast(*this); + out << binding.name(); + break; + } + case PatternKind::GenericBinding: { + const auto& binding = cast(*this); + out << binding.name(); + break; + } + case PatternKind::TuplePattern: { + out << "(...)"; + break; + } + case PatternKind::AlternativePattern: { + const auto& alternative = cast(*this); + out << alternative.choice_type() << "." << alternative.alternative_name() + << "(...)"; + break; + } + case PatternKind::VarPattern: + out << "var ..."; + break; + case PatternKind::ExpressionPattern: + out << "..."; + break; + } +} + // Equivalent to `GetBindings`, but stores its output in `bindings` instead of // returning it. static void GetBindingsImpl( @@ -73,6 +112,7 @@ static void GetBindingsImpl( return; case PatternKind::AutoPattern: case PatternKind::ExpressionPattern: + case PatternKind::GenericBinding: return; case PatternKind::VarPattern: GetBindingsImpl(cast(pattern).pattern(), bindings); diff --git a/executable_semantics/ast/pattern.h b/executable_semantics/ast/pattern.h index 5238263335ea..3d909ec29499 100644 --- a/executable_semantics/ast/pattern.h +++ b/executable_semantics/ast/pattern.h @@ -38,6 +38,7 @@ class Pattern : public AstNode { ~Pattern() override = 0; void Print(llvm::raw_ostream& out) const override; + void PrintID(llvm::raw_ostream& out) const override; static auto classof(const AstNode* node) -> bool { return InheritsFromPattern(node->kind()); @@ -50,7 +51,10 @@ class Pattern : public AstNode { } // The static type of this pattern. Cannot be called before typechecking. - auto static_type() const -> const Value& { return **static_type_; } + auto static_type() const -> const Value& { + CHECK(static_type_.has_value()); + return **static_type_; + } // Sets the static type of this expression. Can only be called once, during // typechecking. @@ -168,6 +172,9 @@ class BindingPattern : public Pattern { auto constant_value() const -> std::optional> { return std::nullopt; } + auto symbolic_identity() const -> std::optional> { + return std::nullopt; + } private: std::string name_; @@ -195,6 +202,58 @@ class TuplePattern : public Pattern { std::vector> fields_; }; +class GenericBinding : public Pattern { + public: + using ImplementsCarbonValueNode = void; + + GenericBinding(SourceLocation source_loc, std::string name, + Nonnull type) + : Pattern(AstNodeKind::GenericBinding, source_loc), + name_(std::move(name)), + type_(type) {} + + void Print(llvm::raw_ostream& out) const override; + void PrintID(llvm::raw_ostream& out) const override; + + static auto classof(const AstNode* node) -> bool { + return InheritsFromGenericBinding(node->kind()); + } + + auto name() const -> const std::string& { return name_; } + auto type() const -> const Expression& { return *type_; } + auto type() -> Expression& { return *type_; } + + auto value_category() const -> ValueCategory { return ValueCategory::Let; } + + auto constant_value() const -> std::optional> { + return std::nullopt; + } + + auto symbolic_identity() const -> std::optional> { + return symbolic_identity_; + } + void set_symbolic_identity(Nonnull value) { + CHECK(!symbolic_identity_.has_value()); + symbolic_identity_ = value; + } + + // The impl binding associated with this type variable. + auto impl_binding() const -> std::optional> { + return impl_binding_; + } + // Set the impl binding. + void set_impl_binding(Nonnull binding) { + CHECK(!impl_binding_.has_value()); + impl_binding_ = binding; + } + + private: + std::string name_; + Nonnull type_; + std::optional> symbolic_identity_; + std::optional> impl_binding_; +}; + // Converts paren_contents to a Pattern, interpreting the parentheses as // grouping if their contents permit that interpretation, or as forming a // tuple otherwise. diff --git a/executable_semantics/ast/statement.h b/executable_semantics/ast/statement.h index ca0153f3ab16..1a7c3bfda964 100644 --- a/executable_semantics/ast/statement.h +++ b/executable_semantics/ast/statement.h @@ -29,6 +29,7 @@ class Statement : public AstNode { ~Statement() override = 0; void Print(llvm::raw_ostream& out) const override { PrintDepth(-1, out); } + void PrintID(llvm::raw_ostream& out) const override { PrintDepth(1, out); } void PrintDepth(int depth, llvm::raw_ostream& out) const; static auto classof(const AstNode* node) { @@ -350,6 +351,9 @@ class Continuation : public Statement { auto constant_value() const -> std::optional> { return std::nullopt; } + auto symbolic_identity() const -> std::optional> { + return std::nullopt; + } private: std::string name_; diff --git a/executable_semantics/ast/static_scope.h b/executable_semantics/ast/static_scope.h index a669a29b9490..b2ca222047f2 100644 --- a/executable_semantics/ast/static_scope.h +++ b/executable_semantics/ast/static_scope.h @@ -38,15 +38,22 @@ static constexpr bool ImplementsValueNode = false; with a value, such as declarations and bindings. The interface consists of the following methods: + // Returns the constant associated with the node. + // This is called by the interpreter, not the type checker. + auto constant_value() const -> std::optional>; + + // Returns the symbolic compile-time identity of the node. + // This is called by the type checker, not the interpreter. + auto symbolic_identity() const -> std::optional>; + // Returns the static type of an IdentifierExpression that names *this. auto static_type() const -> const Value&; // Returns the value category of an IdentifierExpression that names *this. auto value_category() const -> ValueCategory; - // Print the node for diagnostic or tracing purposes. - void Print(llvm::raw_ostream& out) const; - + // Print the node's identity (e.g. its name). + void PrintID(llvm::raw_ostream& out) const; */ // TODO: consider turning the above documentation into real code, as sketched @@ -70,9 +77,13 @@ class ValueNodeView { [](const AstNode& base) -> std::optional> { return llvm::cast(base).constant_value(); }), + symbolic_identity_( + [](const AstNode& base) -> std::optional> { + return llvm::cast(base).symbolic_identity(); + }), print_([](const AstNode& base, llvm::raw_ostream& out) -> void { // TODO: change this to print a summary of the node - return llvm::cast(base).Print(out); + return llvm::cast(base).PrintID(out); }), static_type_([](const AstNode& base) -> const Value& { return llvm::cast(base).static_type(); @@ -94,6 +105,11 @@ class ValueNodeView { return constant_value_(*base_); } + // Returns node->symbolic_identity() + auto symbolic_identity() const -> std::optional> { + return symbolic_identity_(*base_); + } + void Print(llvm::raw_ostream& out) const { print_(*base_, out); } // Returns node->static_type() @@ -123,6 +139,8 @@ class ValueNodeView { Nonnull base_; std::function>(const AstNode&)> constant_value_; + std::function>(const AstNode&)> + symbolic_identity_; std::function print_; std::function static_type_; std::function value_category_; diff --git a/executable_semantics/fuzzing/BUILD b/executable_semantics/fuzzing/BUILD index 58b482154805..9603caa0cb2f 100644 --- a/executable_semantics/fuzzing/BUILD +++ b/executable_semantics/fuzzing/BUILD @@ -11,7 +11,6 @@ cc_library( "//executable_semantics/ast", "//executable_semantics/ast:declaration", "//executable_semantics/ast:expression", - "//executable_semantics/ast:generic_binding", "@llvm-project//llvm:Support", ], ) diff --git a/executable_semantics/fuzzing/ast_to_proto.cpp b/executable_semantics/fuzzing/ast_to_proto.cpp index 2eda74df0d99..43757e8c5f4c 100644 --- a/executable_semantics/fuzzing/ast_to_proto.cpp +++ b/executable_semantics/fuzzing/ast_to_proto.cpp @@ -8,7 +8,6 @@ #include "executable_semantics/ast/declaration.h" #include "executable_semantics/ast/expression.h" -#include "executable_semantics/ast/generic_binding.h" #include "llvm/Support/Casting.h" namespace Carbon { @@ -240,6 +239,14 @@ static auto BindingPatternToProto(const BindingPattern& pattern) return pattern_proto; } +static auto GenericBindingToProto(const GenericBinding& binding) + -> Fuzzing::GenericBinding { + Fuzzing::GenericBinding binding_proto; + binding_proto.set_name(binding.name()); + *binding_proto.mutable_type() = ExpressionToProto(binding.type()); + return binding_proto; +} + static auto TuplePatternToProto(const TuplePattern& tuple_pattern) -> Fuzzing::TuplePattern { Fuzzing::TuplePattern tuple_pattern_proto; @@ -252,6 +259,11 @@ static auto TuplePatternToProto(const TuplePattern& tuple_pattern) static auto PatternToProto(const Pattern& pattern) -> Fuzzing::Pattern { Fuzzing::Pattern pattern_proto; switch (pattern.kind()) { + case PatternKind::GenericBinding: { + const auto& binding = cast(pattern); + *pattern_proto.mutable_generic_binding() = GenericBindingToProto(binding); + break; + } case PatternKind::BindingPattern: { const auto& binding = cast(pattern); *pattern_proto.mutable_binding_pattern() = BindingPatternToProto(binding); @@ -422,14 +434,6 @@ static auto ReturnTermToProto(const ReturnTerm& return_term) return return_term_proto; } -static auto GenericBindingToProto(const GenericBinding& binding) - -> Fuzzing::GenericBinding { - Fuzzing::GenericBinding binding_proto; - binding_proto.set_name(binding.name()); - *binding_proto.mutable_type() = ExpressionToProto(binding.type()); - return binding_proto; -} - static auto DeclarationToProto(const Declaration& declaration) -> Fuzzing::Declaration { Fuzzing::Declaration declaration_proto; diff --git a/executable_semantics/interpreter/action_stack.cpp b/executable_semantics/interpreter/action_stack.cpp index a9b9f2fc144d..662bb310be9c 100644 --- a/executable_semantics/interpreter/action_stack.cpp +++ b/executable_semantics/interpreter/action_stack.cpp @@ -51,10 +51,11 @@ void ActionStack::Initialize(ValueNodeView value_node, auto ActionStack::ValueOfNode(ValueNodeView value_node, SourceLocation source_loc) const -> ErrorOr> { - if (std::optional> constant_value = - value_node.constant_value(); - constant_value.has_value()) { - return *constant_value; + std::optional value = (phase_ == Phase::CompileTime) + ? value_node.symbolic_identity() + : value_node.constant_value(); + if (value.has_value()) { + return *value; } for (const std::unique_ptr& action : todo_) { // TODO: have static name resolution identify the scope of value_node diff --git a/executable_semantics/interpreter/action_stack.h b/executable_semantics/interpreter/action_stack.h index 67234e5b5478..f1a16d243811 100644 --- a/executable_semantics/interpreter/action_stack.h +++ b/executable_semantics/interpreter/action_stack.h @@ -15,16 +15,19 @@ namespace Carbon { +// Selects between compile-time and run-time behavior. +enum class Phase { CompileTime, RunTime }; + // The stack of Actions currently being executed by the interpreter. class ActionStack { public: // Constructs an empty compile-time ActionStack. - ActionStack() = default; + ActionStack() : phase_(Phase::CompileTime) {} // Constructs an empty run-time ActionStack that allocates global variables // on `heap`. explicit ActionStack(Nonnull heap) - : globals_(RuntimeScope(heap)) {} + : globals_(RuntimeScope(heap)), phase_(Phase::RunTime) {} void Print(llvm::raw_ostream& out) const; LLVM_DUMP_METHOD void Dump() const { Print(llvm::errs()); } @@ -120,6 +123,7 @@ class ActionStack { Stack> todo_; std::optional> result_; std::optional globals_; + Phase phase_; }; } // namespace Carbon diff --git a/executable_semantics/interpreter/impl_scope.cpp b/executable_semantics/interpreter/impl_scope.cpp index e809c4652135..8b61f109e44a 100644 --- a/executable_semantics/interpreter/impl_scope.cpp +++ b/executable_semantics/interpreter/impl_scope.cpp @@ -6,6 +6,7 @@ #include "executable_semantics/common/error.h" #include "executable_semantics/interpreter/value.h" +#include "llvm/ADT/StringExtras.h" #include "llvm/Support/Casting.h" using llvm::cast; @@ -79,4 +80,17 @@ auto ImplScope::ResolveHere(Nonnull iface_type, } } +// TODO: Add indentation when printing the parents. +void ImplScope::Print(llvm::raw_ostream& out) const { + out << "impls: "; + llvm::ListSeparator sep; + for (Impl impl : impls_) { + out << sep << *(impl.type) << " as " << *(impl.interface); + } + out << "\n"; + for (const Nonnull& parent : parent_scopes_) { + out << *parent; + } +} + } // namespace Carbon diff --git a/executable_semantics/interpreter/impl_scope.h b/executable_semantics/interpreter/impl_scope.h index d60549702151..ebf1c7b8d465 100644 --- a/executable_semantics/interpreter/impl_scope.h +++ b/executable_semantics/interpreter/impl_scope.h @@ -52,6 +52,8 @@ class ImplScope { auto Resolve(Nonnull iface, Nonnull type, SourceLocation source_loc) const -> ErrorOr; + void Print(llvm::raw_ostream& out) const; + private: auto TryResolve(Nonnull iface_type, Nonnull type, SourceLocation source_loc) const diff --git a/executable_semantics/interpreter/interpreter.cpp b/executable_semantics/interpreter/interpreter.cpp index f3a42087da97..d7ac7992dbf2 100644 --- a/executable_semantics/interpreter/interpreter.cpp +++ b/executable_semantics/interpreter/interpreter.cpp @@ -29,9 +29,6 @@ using llvm::isa; namespace Carbon { -// Selects between compile-time and run-time behavior. -enum class Phase { CompileTime, RunTime }; - // Constructs an ActionStack suitable for the specified phase. static auto MakeTodo(Phase phase, Nonnull heap) -> ActionStack { switch (phase) { @@ -54,7 +51,8 @@ class Interpreter { : arena_(arena), heap_(arena), todo_(MakeTodo(phase, &heap_)), - trace_(trace) {} + trace_(trace), + phase_(phase) {} ~Interpreter(); @@ -90,11 +88,25 @@ class Interpreter { // Returns the result of converting `value` to type `destination_type`. auto Convert(Nonnull value, - Nonnull destination_type) const - -> Nonnull; + Nonnull destination_type, + SourceLocation source_loc) const + -> ErrorOr>; + + // Instantiate a type by replacing all type variables that occur inside the + // type by the current values of those variables. + // + // For example, suppose T=i32 and U=Bool. Then + // __Fn (Point(T)) -> Point(U) + // becomes + // __Fn (Point(i32)) -> Point(Bool) + auto InstantiateType(Nonnull type, + SourceLocation source_loc) const + -> ErrorOr>; void PrintState(llvm::raw_ostream& out); + Phase phase() const { return phase_; } + Nonnull arena_; Heap heap_; @@ -106,6 +118,7 @@ class Interpreter { std::vector> stack_fragments_; bool trace_; + Phase phase_; }; Interpreter::~Interpreter() { @@ -178,7 +191,8 @@ auto Interpreter::CreateStruct(const std::vector& fields, auto PatternMatch(Nonnull p, Nonnull v, SourceLocation source_loc, - std::optional> bindings) -> bool { + std::optional> bindings, + BindingMap& generic_args) -> bool { switch (p->kind()) { case Value::Kind::BindingPlaceholderValue: { CHECK(bindings.has_value()); @@ -188,6 +202,11 @@ auto PatternMatch(Nonnull p, Nonnull v, } return true; } + case Value::Kind::VariableType: { + const auto& var_type = cast(*p); + generic_args[&var_type.binding()] = v; + return true; + } case Value::Kind::TupleValue: switch (v->kind()) { case Value::Kind::TupleValue: { @@ -196,7 +215,7 @@ auto PatternMatch(Nonnull p, Nonnull v, CHECK(p_tup.elements().size() == v_tup.elements().size()); for (size_t i = 0; i < p_tup.elements().size(); ++i) { if (!PatternMatch(p_tup.elements()[i], v_tup.elements()[i], - source_loc, bindings)) { + source_loc, bindings, generic_args)) { return false; } } // for @@ -212,7 +231,8 @@ auto PatternMatch(Nonnull p, Nonnull v, for (size_t i = 0; i < p_struct.elements().size(); ++i) { CHECK(p_struct.elements()[i].name == v_struct.elements()[i].name); if (!PatternMatch(p_struct.elements()[i].value, - v_struct.elements()[i].value, source_loc, bindings)) { + v_struct.elements()[i].value, source_loc, bindings, + generic_args)) { return false; } } @@ -228,7 +248,7 @@ auto PatternMatch(Nonnull p, Nonnull v, return false; } return PatternMatch(&p_alt.argument(), &v_alt.argument(), source_loc, - bindings); + bindings, generic_args); } default: FATAL() << "expected a choice alternative in pattern, not " << *v; @@ -239,11 +259,11 @@ auto PatternMatch(Nonnull p, Nonnull v, const auto& p_fn = cast(*p); const auto& v_fn = cast(*v); if (!PatternMatch(&p_fn.parameters(), &v_fn.parameters(), source_loc, - bindings)) { + bindings, generic_args)) { return false; } if (!PatternMatch(&p_fn.return_type(), &v_fn.return_type(), - source_loc, bindings)) { + source_loc, bindings, generic_args)) { return false; } return true; @@ -349,9 +369,79 @@ auto Interpreter::StepLvalue() -> ErrorOr { } } +auto Interpreter::InstantiateType(Nonnull type, + SourceLocation source_loc) const + -> ErrorOr> { + if (trace_) { + llvm::outs() << "instantiating: " << *type << "\n"; + } + switch (type->kind()) { + case Value::Kind::VariableType: { + if (trace_) { + llvm::outs() << "case VariableType\n"; + } + ASSIGN_OR_RETURN( + Nonnull value, + todo_.ValueOfNode(&cast(*type).binding(), source_loc)); + if (const auto* lvalue = dyn_cast(value)) { + ASSIGN_OR_RETURN(value, heap_.Read(lvalue->address(), source_loc)); + } + return value; + } + case Value::Kind::NominalClassType: { + if (trace_) { + llvm::outs() << "case NominalClassType\n"; + } + const auto& class_type = cast(*type); + BindingMap inst_type_args; + for (const auto& [ty_var, ty_arg] : class_type.type_args()) { + ASSIGN_OR_RETURN(inst_type_args[ty_var], + InstantiateType(ty_arg, source_loc)); + } + if (trace_) { + llvm::outs() << "finished instantiating ty_arg\n"; + } + std::map, Nonnull> witnesses; + for (const auto& [bind, impl] : class_type.impls()) { + ASSIGN_OR_RETURN(Nonnull witness_addr, + todo_.ValueOfNode(impl, source_loc)); + if (trace_) { + llvm::outs() << "witness_addr: " << *witness_addr << "\n"; + } + // If the witness came directly from an `impl` declaration (via + // `constant_value`), then it is a `Witness`. If the witness + // came from the runtime scope, then the `Witness` got wrapped + // in an `LValue` because that's what + // `RuntimeScope::Initialize` does. + Nonnull witness; + if (llvm::isa(witness_addr)) { + witness = cast(witness_addr); + } else if (llvm::isa(witness_addr)) { + ASSIGN_OR_RETURN( + Nonnull witness_value, + heap_.Read(llvm::cast(witness_addr)->address(), + source_loc)); + witness = cast(witness_value); + } else { + FATAL() << "expected a witness or LValue of a witness"; + } + witnesses[bind] = witness; + } + if (trace_) { + llvm::outs() << "finished finding witnesses\n"; + } + return arena_->New(&class_type.declaration(), + inst_type_args, witnesses); + } + default: + return type; + } +} + auto Interpreter::Convert(Nonnull value, - Nonnull destination_type) const - -> Nonnull { + Nonnull destination_type, + SourceLocation source_loc) const + -> ErrorOr> { switch (value->kind()) { case Value::Kind::IntValue: case Value::Kind::FunctionValue: @@ -396,13 +486,19 @@ auto Interpreter::Convert(Nonnull value, destination_struct_type.fields()) { std::optional> old_value = struct_val.FindField(field_name); - new_elements.push_back( - {.name = field_name, .value = Convert(*old_value, field_type)}); + ASSIGN_OR_RETURN(Nonnull val, + Convert(*old_value, field_type, source_loc)); + new_elements.push_back({.name = field_name, .value = val}); } return arena_->New(std::move(new_elements)); } - case Value::Kind::NominalClassType: - return arena_->New(destination_type, value); + case Value::Kind::NominalClassType: { + // Instantiate the `destintation_type` to obtain the runtime + // type of the object. + ASSIGN_OR_RETURN(Nonnull inst_dest, + InstantiateType(destination_type, source_loc)); + return arena_->New(inst_dest, value); + } default: FATAL() << "Can't convert value " << *value << " to type " << *destination_type; @@ -415,8 +511,11 @@ auto Interpreter::Convert(Nonnull value, destination_tuple_type->elements().size()); std::vector> new_elements; for (size_t i = 0; i < tuple->elements().size(); ++i) { - new_elements.push_back(Convert(tuple->elements()[i], - destination_tuple_type->elements()[i])); + ASSIGN_OR_RETURN( + Nonnull val, + Convert(tuple->elements()[i], destination_tuple_type->elements()[i], + source_loc)); + new_elements.push_back(val); } return arena_->New(std::move(new_elements)); } @@ -580,11 +679,27 @@ auto Interpreter::StepExp() -> ErrorOr { alt.alt_name(), alt.choice_name(), act.results()[1])); } case Value::Kind::FunctionValue: { - const FunctionDeclaration& function = - cast(*act.results()[0]).declaration(); - Nonnull converted_args = Convert( - act.results()[1], &function.param_pattern().static_type()); + const FunctionValue& fun_val = + cast(*act.results()[0]); + const FunctionDeclaration& function = fun_val.declaration(); + if (trace_) { + llvm::outs() << "*** call function " << function.name() << "\n"; + } + ASSIGN_OR_RETURN(Nonnull converted_args, + Convert(act.results()[1], + &function.param_pattern().static_type(), + exp.source_loc())); RuntimeScope function_scope(&heap_); + // Bring the class type arguments into scope. + for (const auto& [bind, val] : fun_val.type_args()) { + function_scope.Initialize(bind, val); + } + // Bring the deduced type arguments into scope. + for (const auto& [bind, val] : + cast(exp).deduced_args()) { + function_scope.Initialize(bind, val); + } + // Bring the impl witness tables into scope. for (const auto& [impl_bind, impl_node] : cast(exp).impls()) { @@ -597,9 +712,13 @@ auto Interpreter::StepExp() -> ErrorOr { } function_scope.Initialize(impl_bind, witness); } + for (const auto& [impl_bind, witness] : fun_val.witnesses()) { + function_scope.Initialize(impl_bind, witness); + } + BindingMap generic_args; CHECK(PatternMatch(&function.param_pattern().value(), converted_args, exp.source_loc(), - &function_scope)); + &function_scope, generic_args)); CHECK(function.body().has_value()) << "Calling a function that's missing a body"; return todo_.Spawn( @@ -609,19 +728,75 @@ auto Interpreter::StepExp() -> ErrorOr { case Value::Kind::BoundMethodValue: { const auto& m = cast(*act.results()[0]); const FunctionDeclaration& method = m.declaration(); - Nonnull converted_args = Convert( - act.results()[1], &method.param_pattern().static_type()); + CHECK(method.is_method()); + ASSIGN_OR_RETURN( + Nonnull converted_args, + Convert(act.results()[1], &method.param_pattern().static_type(), + exp.source_loc())); RuntimeScope method_scope(&heap_); + BindingMap generic_args; CHECK(PatternMatch(&method.me_pattern().value(), m.receiver(), - exp.source_loc(), &method_scope)); + exp.source_loc(), &method_scope, generic_args)); CHECK(PatternMatch(&method.param_pattern().value(), converted_args, - exp.source_loc(), &method_scope)); + exp.source_loc(), &method_scope, generic_args)); + // Bring the class type arguments into scope. + for (const auto& [bind, val] : m.type_args()) { + method_scope.Initialize(bind, val); + } + + // Bring the impl witness tables into scope. + for (const auto& [impl_bind, witness] : m.witnesses()) { + method_scope.Initialize(impl_bind, witness); + } CHECK(method.body().has_value()) << "Calling a method that's missing a body"; return todo_.Spawn( std::make_unique(*method.body()), std::move(method_scope)); } + case Value::Kind::NominalClassType: { + const NominalClassType& class_type = + cast(*act.results()[0]); + const ClassDeclaration& class_decl = class_type.declaration(); + RuntimeScope type_params_scope(&heap_); + BindingMap generic_args; + if (class_decl.type_params().has_value()) { + CHECK(PatternMatch(&(*class_decl.type_params())->value(), + act.results()[1], exp.source_loc(), + &type_params_scope, generic_args)); + switch (phase()) { + case Phase::RunTime: { + std::map, const Witness*> + witnesses; + for (const auto& [impl_bind, impl_node] : + cast(exp).impls()) { + ASSIGN_OR_RETURN( + Nonnull witness, + todo_.ValueOfNode(impl_node, exp.source_loc())); + if (witness->kind() == Value::Kind::LValue) { + const LValue& lval = cast(*witness); + ASSIGN_OR_RETURN(witness, heap_.Read(lval.address(), + exp.source_loc())); + } + witnesses[impl_bind] = &cast(*witness); + } + Nonnull inst_class = + arena_->New(&class_type.declaration(), + generic_args, witnesses); + return todo_.FinishAction(inst_class); + } + case Phase::CompileTime: { + Nonnull inst_class = + arena_->New( + &class_type.declaration(), generic_args, + cast(exp).impls()); + return todo_.FinishAction(inst_class); + } + } + } else { + FATAL() << "instantiation of non-generic class " << class_type; + } + } default: return FATAL_RUNTIME_ERROR(exp.source_loc()) << "in call, expected a function, not " << *act.results()[0]; @@ -735,6 +910,10 @@ auto Interpreter::StepPattern() -> ErrorOr { return todo_.FinishAction(arena_->New()); } } + case PatternKind::GenericBinding: { + const auto& binding = cast(pattern); + return todo_.FinishAction(arena_->New(&binding)); + } case PatternKind::TuplePattern: { const auto& tuple = cast(pattern); if (act.pos() < static_cast(tuple.fields().size())) { @@ -805,9 +984,12 @@ auto Interpreter::StepStmt() -> ErrorOr { } auto c = match_stmt.clauses()[clause_num]; RuntimeScope matches(&heap_); - if (PatternMatch(&c.pattern().value(), - Convert(act.results()[0], &c.pattern().static_type()), - stmt.source_loc(), &matches)) { + BindingMap generic_args; + ASSIGN_OR_RETURN(Nonnull val, + Convert(act.results()[0], &c.pattern().static_type(), + stmt.source_loc())); + if (PatternMatch(&c.pattern().value(), val, stmt.source_loc(), &matches, + generic_args)) { // Ensure we don't process any more clauses. act.set_pos(match_stmt.clauses().size() + 1); todo_.MergeScope(std::move(matches)); @@ -825,8 +1007,9 @@ auto Interpreter::StepStmt() -> ErrorOr { return todo_.Spawn( std::make_unique(&cast(stmt).condition())); } else { - Nonnull condition = - Convert(act.results().back(), arena_->New()); + ASSIGN_OR_RETURN(Nonnull condition, + Convert(act.results().back(), arena_->New(), + stmt.source_loc())); if (cast(*condition).value()) { // { {true :: (while ([]) s) :: C, E, F} :: S, H} // -> { { s :: (while (e) s) :: C, E, F } :: S, H} @@ -876,13 +1059,16 @@ auto Interpreter::StepStmt() -> ErrorOr { } else { // { { v :: (x = []) :: C, E, F} :: S, H} // -> { { C, E(x := a), F} :: S, H(a := copy(v))} - Nonnull v = - Convert(act.results()[0], &definition.pattern().static_type()); + ASSIGN_OR_RETURN( + Nonnull v, + Convert(act.results()[0], &definition.pattern().static_type(), + stmt.source_loc())); Nonnull p = &cast(stmt).pattern().value(); RuntimeScope matches(&heap_); - CHECK(PatternMatch(p, v, stmt.source_loc(), &matches)) + BindingMap generic_args; + CHECK(PatternMatch(p, v, stmt.source_loc(), &matches, generic_args)) << stmt.source_loc() << ": internal error in variable definition, match failed"; todo_.MergeScope(std::move(matches)); @@ -912,8 +1098,9 @@ auto Interpreter::StepStmt() -> ErrorOr { // { { v :: (a = []) :: C, E, F} :: S, H} // -> { { C, E, F} :: S, H(a := v)} const auto& lval = cast(*act.results()[0]); - Nonnull rval = - Convert(act.results()[1], &assign.lhs().static_type()); + ASSIGN_OR_RETURN(Nonnull rval, + Convert(act.results()[1], &assign.lhs().static_type(), + stmt.source_loc())); RETURN_IF_ERROR(heap_.Write(lval.address(), rval, stmt.source_loc())); return todo_.FinishAction(); } @@ -925,8 +1112,9 @@ auto Interpreter::StepStmt() -> ErrorOr { return todo_.Spawn( std::make_unique(&cast(stmt).condition())); } else if (act.pos() == 1) { - Nonnull condition = - Convert(act.results()[0], arena_->New()); + ASSIGN_OR_RETURN(Nonnull condition, + Convert(act.results()[0], arena_->New(), + stmt.source_loc())); if (cast(*condition).value()) { // { {true :: if ([]) then_stmt else else_stmt :: C, E, F} :: // S, H} @@ -955,9 +1143,11 @@ auto Interpreter::StepStmt() -> ErrorOr { // { {v :: return [] :: C, E, F} :: {C', E', F'} :: S, H} // -> { {v :: C', E', F'} :: S, H} const FunctionDeclaration& function = cast(stmt).function(); - return todo_.UnwindPast( - *function.body(), - Convert(act.results()[0], &function.return_term().static_type())); + ASSIGN_OR_RETURN( + Nonnull return_value, + Convert(act.results()[0], &function.return_term().static_type(), + stmt.source_loc())); + return todo_.UnwindPast(*function.body(), return_value); } case StatementKind::Continuation: { CHECK(act.pos() == 0); diff --git a/executable_semantics/interpreter/interpreter.h b/executable_semantics/interpreter/interpreter.h index 15497595f8fb..91dc8673ab97 100644 --- a/executable_semantics/interpreter/interpreter.h +++ b/executable_semantics/interpreter/interpreter.h @@ -44,12 +44,14 @@ auto InterpPattern(Nonnull p, Nonnull arena, bool trace) // is not permitted to bind variables. **bindings may be modified even if the // match is unsuccessful, so it should typically be created for the // PatternMatch call and then merged into an existing scope on success. +// The matches for generic variables in the pattern are output in +// `generic_args`. // TODO: consider moving this to a separate header. [[nodiscard]] auto PatternMatch(Nonnull p, Nonnull v, SourceLocation source_loc, - std::optional> bindings) - -> bool; + std::optional> bindings, + BindingMap& generic_args) -> bool; } // namespace Carbon diff --git a/executable_semantics/interpreter/resolve_names.cpp b/executable_semantics/interpreter/resolve_names.cpp index 9687848a2f79..95f3402f4773 100644 --- a/executable_semantics/interpreter/resolve_names.cpp +++ b/executable_semantics/interpreter/resolve_names.cpp @@ -178,6 +178,14 @@ static auto ResolveNames(Pattern& pattern, StaticScope& enclosing_scope) } break; } + case PatternKind::GenericBinding: { + auto& binding = cast(pattern); + RETURN_IF_ERROR(ResolveNames(binding.type(), enclosing_scope)); + if (binding.name() != AnonymousName) { + RETURN_IF_ERROR(enclosing_scope.Add(binding.name(), &binding)); + } + break; + } case PatternKind::TuplePattern: for (Nonnull field : cast(pattern).fields()) { RETURN_IF_ERROR(ResolveNames(*field, enclosing_scope)); @@ -315,8 +323,8 @@ static auto ResolveNames(Declaration& declaration, StaticScope& enclosing_scope) StaticScope function_scope; function_scope.AddParent(&enclosing_scope); for (Nonnull binding : function.deduced_parameters()) { - RETURN_IF_ERROR(function_scope.Add(binding->name(), binding)); RETURN_IF_ERROR(ResolveNames(binding->type(), function_scope)); + RETURN_IF_ERROR(function_scope.Add(binding->name(), binding)); } if (function.is_method()) { RETURN_IF_ERROR(ResolveNames(function.me_pattern(), function_scope)); @@ -336,9 +344,17 @@ static auto ResolveNames(Declaration& declaration, StaticScope& enclosing_scope) StaticScope class_scope; class_scope.AddParent(&enclosing_scope); RETURN_IF_ERROR(class_scope.Add(class_decl.name(), &class_decl)); - for (Nonnull member : class_decl.members()) { - RETURN_IF_ERROR(AddExposedNames(*member, class_scope)); + if (class_decl.type_params().has_value()) { + RETURN_IF_ERROR(ResolveNames(**class_decl.type_params(), class_scope)); } + + // TODO: Disable unqualified access of members by other members for now. + // Put it back later, but in a way that turns unqualified accesses + // into qualified ones, so that generic classes and impls + // behave the in the right way. -Jeremy + // for (Nonnull member : class_decl.members()) { + // AddExposedNames(*member, class_scope); + // } for (Nonnull member : class_decl.members()) { RETURN_IF_ERROR(ResolveNames(*member, class_scope)); } diff --git a/executable_semantics/interpreter/type_checker.cpp b/executable_semantics/interpreter/type_checker.cpp index 8144d3fb03a1..96d466517158 100644 --- a/executable_semantics/interpreter/type_checker.cpp +++ b/executable_semantics/interpreter/type_checker.cpp @@ -122,16 +122,7 @@ auto TypeChecker::ExpectIsConcreteType(SourceLocation source_loc, } } -// Returns true if *source is implicitly convertible to *destination. *source -// and *destination must be concrete types. -static auto IsImplicitlyConvertible(Nonnull source, - Nonnull destination) -> bool; - -// Returns true if source_fields and destination_fields contain the same set -// of names, and each value in source_fields is implicitly convertible to -// the corresponding value in destination_fields. All values in both arguments -// must be types. -static auto FieldTypesImplicitlyConvertible( +auto TypeChecker::FieldTypesImplicitlyConvertible( llvm::ArrayRef source_fields, llvm::ArrayRef destination_fields) { if (source_fields.size() != destination_fields.size()) { @@ -150,8 +141,29 @@ static auto FieldTypesImplicitlyConvertible( return true; } -static auto IsImplicitlyConvertible(Nonnull source, - Nonnull destination) -> bool { +auto TypeChecker::FieldTypes(const NominalClassType& class_type) + -> std::vector { + std::vector field_types; + for (Nonnull m : class_type.declaration().members()) { + switch (m->kind()) { + case DeclarationKind::VariableDeclaration: { + const auto& var = cast(*m); + Nonnull field_type = + Substitute(class_type.type_args(), &var.binding().static_type()); + field_types.push_back( + {.name = var.binding().name(), .value = field_type}); + break; + } + default: + break; + } + } + return field_types; +} + +auto TypeChecker::IsImplicitlyConvertible(Nonnull source, + Nonnull destination) + -> bool { CHECK(IsConcreteType(source)); CHECK(IsConcreteType(destination)); if (TypeEqual(source, destination)) { @@ -172,34 +184,35 @@ static auto IsImplicitlyConvertible(Nonnull source, return false; } case Value::Kind::TupleValue: - switch (destination->kind()) { - case Value::Kind::TupleValue: { - const std::vector>& source_elements = - cast(*source).elements(); - const std::vector>& destination_elements = - cast(*destination).elements(); - if (source_elements.size() != destination_elements.size()) { + if (destination->kind() == Value::Kind::TupleValue) { + const std::vector>& source_elements = + cast(*source).elements(); + const std::vector>& destination_elements = + cast(*destination).elements(); + if (source_elements.size() != destination_elements.size()) { + return false; + } + for (size_t i = 0; i < source_elements.size(); ++i) { + if (!IsImplicitlyConvertible(source_elements[i], + destination_elements[i])) { return false; } - for (size_t i = 0; i < source_elements.size(); ++i) { - if (!IsImplicitlyConvertible(source_elements[i], - destination_elements[i])) { - return false; - } - } - return true; } - default: - return false; + return true; + } else { + return false; } + case Value::Kind::TypeType: + return destination->kind() == Value::Kind::InterfaceType; default: return false; } } -static auto ExpectType(SourceLocation source_loc, const std::string& context, - Nonnull expected, - Nonnull actual) -> ErrorOr { +auto TypeChecker::ExpectType(SourceLocation source_loc, + const std::string& context, + Nonnull expected, + Nonnull actual) -> ErrorOr { if (!IsImplicitlyConvertible(actual, expected)) { return FATAL_COMPILATION_ERROR(source_loc) << "type error in " << context << ": " @@ -212,29 +225,29 @@ static auto ExpectType(SourceLocation source_loc, const std::string& context, auto TypeChecker::ArgumentDeduction(SourceLocation source_loc, BindingMap& deduced, - Nonnull param, - Nonnull arg) + Nonnull param_type, + Nonnull arg_type) -> ErrorOr { - switch (param->kind()) { + switch (param_type->kind()) { case Value::Kind::VariableType: { - const auto& var_type = cast(*param); - auto [it, success] = deduced.insert({&var_type.binding(), arg}); + const auto& var_type = cast(*param_type); + auto [it, success] = deduced.insert({&var_type.binding(), arg_type}); if (!success) { // TODO: can we allow implicit conversions here? - RETURN_IF_ERROR( - ExpectExactType(source_loc, "argument deduction", it->second, arg)); + RETURN_IF_ERROR(ExpectExactType(source_loc, "argument deduction", + it->second, arg_type)); } return Success(); } case Value::Kind::TupleValue: { - if (arg->kind() != Value::Kind::TupleValue) { + if (arg_type->kind() != Value::Kind::TupleValue) { return FATAL_COMPILATION_ERROR(source_loc) << "type error in argument deduction\n" - << "expected: " << *param << "\n" - << "actual: " << *arg; + << "expected: " << *param_type << "\n" + << "actual: " << *arg_type; } - const auto& param_tup = cast(*param); - const auto& arg_tup = cast(*arg); + const auto& param_tup = cast(*param_type); + const auto& arg_tup = cast(*arg_type); if (param_tup.elements().size() != arg_tup.elements().size()) { return FATAL_COMPILATION_ERROR(source_loc) << "mismatch in tuple sizes, expected " @@ -249,14 +262,14 @@ auto TypeChecker::ArgumentDeduction(SourceLocation source_loc, return Success(); } case Value::Kind::StructType: { - if (arg->kind() != Value::Kind::StructType) { + if (arg_type->kind() != Value::Kind::StructType) { return FATAL_COMPILATION_ERROR(source_loc) << "type error in argument deduction\n" - << "expected: " << *param << "\n" - << "actual: " << *arg; + << "expected: " << *param_type << "\n" + << "actual: " << *arg_type; } - const auto& param_struct = cast(*param); - const auto& arg_struct = cast(*arg); + const auto& param_struct = cast(*param_type); + const auto& arg_struct = cast(*arg_type); if (param_struct.fields().size() != arg_struct.fields().size()) { return FATAL_COMPILATION_ERROR(source_loc) << "mismatch in struct field counts, expected " @@ -276,14 +289,14 @@ auto TypeChecker::ArgumentDeduction(SourceLocation source_loc, return Success(); } case Value::Kind::FunctionType: { - if (arg->kind() != Value::Kind::FunctionType) { + if (arg_type->kind() != Value::Kind::FunctionType) { return FATAL_COMPILATION_ERROR(source_loc) << "type error in argument deduction\n" - << "expected: " << *param << "\n" - << "actual: " << *arg; + << "expected: " << *param_type << "\n" + << "actual: " << *arg_type; } - const auto& param_fn = cast(*param); - const auto& arg_fn = cast(*arg); + const auto& param_fn = cast(*param_type); + const auto& arg_fn = cast(*arg_type); // TODO: handle situation when arg has deduced parameters. RETURN_IF_ERROR(ArgumentDeduction( source_loc, deduced, ¶m_fn.parameters(), &arg_fn.parameters())); @@ -292,23 +305,41 @@ auto TypeChecker::ArgumentDeduction(SourceLocation source_loc, return Success(); } case Value::Kind::PointerType: { - if (arg->kind() != Value::Kind::PointerType) { + if (arg_type->kind() != Value::Kind::PointerType) { return FATAL_COMPILATION_ERROR(source_loc) << "type error in argument deduction\n" - << "expected: " << *param << "\n" - << "actual: " << *arg; + << "expected: " << *param_type << "\n" + << "actual: " << *arg_type; } return ArgumentDeduction(source_loc, deduced, - &cast(*param).type(), - &cast(*arg).type()); + &cast(*param_type).type(), + &cast(*arg_type).type()); } // Nothing to do in the case for `auto`. case Value::Kind::AutoType: { return Success(); } + case Value::Kind::NominalClassType: { + const auto& param_class_type = cast(*param_type); + if (arg_type->kind() == Value::Kind::NominalClassType) { + const auto& arg_class_type = cast(*arg_type); + if (param_class_type.declaration().name() == + arg_class_type.declaration().name()) { + for (const auto& [ty, param_ty] : param_class_type.type_args()) { + RETURN_IF_ERROR( + ArgumentDeduction(source_loc, deduced, param_ty, + arg_class_type.type_args().at(ty))); + } + return Success(); + } + } + return FATAL_COMPILATION_ERROR(source_loc) + << "type error in argument deduction\n" + << "expected: " << *param_type << "\n" + << "actual: " << *arg_type; + } // For the following cases, we check for type convertability. case Value::Kind::ContinuationType: - case Value::Kind::NominalClassType: case Value::Kind::InterfaceType: case Value::Kind::ChoiceType: case Value::Kind::IntType: @@ -318,7 +349,7 @@ auto TypeChecker::ArgumentDeduction(SourceLocation source_loc, case Value::Kind::TypeOfClassType: case Value::Kind::TypeOfInterfaceType: case Value::Kind::TypeOfChoiceType: - return ExpectType(source_loc, "argument deduction", param, arg); + return ExpectType(source_loc, "argument deduction", param_type, arg_type); // The rest of these cases should never happen. case Value::Kind::Witness: case Value::Kind::IntValue: @@ -334,7 +365,8 @@ auto TypeChecker::ArgumentDeduction(SourceLocation source_loc, case Value::Kind::AlternativeConstructorValue: case Value::Kind::ContinuationValue: case Value::Kind::StringValue: - FATAL() << "In ArgumentDeduction: expected type, not value " << *param; + FATAL() << "In ArgumentDeduction: expected type, not value " + << *param_type; } } @@ -377,11 +409,24 @@ auto TypeChecker::Substitute( return arena_->New( Substitute(dict, &cast(*type).type())); } + case Value::Kind::NominalClassType: { + const auto& class_type = cast(*type); + BindingMap type_args; + for (const auto& [name, value] : class_type.type_args()) { + type_args[name] = Substitute(dict, value); + } + Nonnull new_class_type = + arena_->New(&class_type.declaration(), type_args); + if (trace_) { + llvm::outs() << "substitution: " << class_type << " => " + << *new_class_type << "\n"; + } + return new_class_type; + } case Value::Kind::AutoType: case Value::Kind::IntType: case Value::Kind::BoolType: case Value::Kind::TypeType: - case Value::Kind::NominalClassType: case Value::Kind::InterfaceType: case Value::Kind::ChoiceType: case Value::Kind::ContinuationType: @@ -506,7 +551,9 @@ auto TypeChecker::TypeCheckExp(Nonnull e, if (std::optional> member = FindMember(access.field(), t_class.declaration().members()); member.has_value()) { - access.set_static_type(&(*member)->static_type()); + Nonnull field_type = + Substitute(t_class.type_args(), &(*member)->static_type()); + access.set_static_type(field_type); switch ((*member)->kind()) { case DeclarationKind::VariableDeclaration: access.set_value_category(access.aggregate().value_category()); @@ -554,7 +601,9 @@ auto TypeChecker::TypeCheckExp(Nonnull e, if (func->is_method()) { break; } - access.set_static_type(&(*member)->static_type()); + Nonnull field_type = Substitute( + class_type.type_args(), &(*member)->static_type()); + access.set_static_type(field_type); access.set_value_category(ValueCategory::Let); return Success(); } @@ -570,7 +619,11 @@ auto TypeChecker::TypeCheckExp(Nonnull e, } } case Value::Kind::VariableType: { - const auto& var_type = cast(aggregate_type); + // This case handles access to a method on a receiver whose type + // is a type variable. For example, `x.foo` where the type of + // `x` is `T` and `foo` and `T` implements an interface that + // includes `foo`. + const VariableType& var_type = cast(aggregate_type); const Value& typeof_var = var_type.binding().static_type(); switch (typeof_var.kind()) { case Value::Kind::InterfaceType: { @@ -586,6 +639,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, Nonnull inst_member_type = Substitute(self_map, &member_type); access.set_static_type(inst_member_type); + CHECK(var_type.binding().impl_binding().has_value()); access.set_impl(*var_type.binding().impl_binding()); return Success(); } else { @@ -596,11 +650,37 @@ auto TypeChecker::TypeCheckExp(Nonnull e, break; } default: - break; + return FATAL_COMPILATION_ERROR(e->source_loc()) + << "field access, unexpected " << aggregate_type + << " of non-interface type " << typeof_var << " in " << *e; + } + break; + } + case Value::Kind::InterfaceType: { + // This case handles access to a class function from a type variable. + // If `T` is a type variable and `foo` is a class function in an + // interface implemented by `T`, then `T.foo` accesses the `foo` class + // function of `T`. + ASSIGN_OR_RETURN(Nonnull var_addr, + InterpExp(&access.aggregate(), arena_, trace_)); + const VariableType& var_type = cast(*var_addr); + const InterfaceType& iface_type = cast(aggregate_type); + const InterfaceDeclaration& iface_decl = iface_type.declaration(); + if (std::optional> member = + FindMember(access.field(), iface_decl.members()); + member.has_value()) { + const Value& member_type = (*member)->static_type(); + Nonnull inst_member_type = + Substitute({{iface_decl.self(), &var_type}}, &member_type); + access.set_static_type(inst_member_type); + CHECK(var_type.binding().impl_binding().has_value()); + access.set_impl(*var_type.binding().impl_binding()); + return Success(); + } else { + return FATAL_COMPILATION_ERROR(e->source_loc()) + << "field access, " << access.field() << " not in " + << iface_decl.name(); } - return FATAL_COMPILATION_ERROR(e->source_loc()) - << "field access, unexpected " << aggregate_type << " in " - << *e; break; } default: @@ -724,30 +804,33 @@ auto TypeChecker::TypeCheckExp(Nonnull e, case ExpressionKind::CallExpression: { auto& call = cast(*e); RETURN_IF_ERROR(TypeCheckExp(&call.function(), impl_scope)); + RETURN_IF_ERROR(TypeCheckExp(&call.argument(), impl_scope)); switch (call.function().static_type().kind()) { case Value::Kind::FunctionType: { const auto& fun_t = cast(call.function().static_type()); - RETURN_IF_ERROR(TypeCheckExp(&call.argument(), impl_scope)); Nonnull parameters = &fun_t.parameters(); Nonnull return_type = &fun_t.return_type(); if (!fun_t.deduced().empty()) { - BindingMap deduced_args; - RETURN_IF_ERROR(ArgumentDeduction(e->source_loc(), deduced_args, - parameters, + BindingMap deduced_type_args; + RETURN_IF_ERROR(ArgumentDeduction(e->source_loc(), + deduced_type_args, parameters, &call.argument().static_type())); + call.set_deduced_args(deduced_type_args); for (Nonnull deduced_param : fun_t.deduced()) { // TODO: change the following to a CHECK once the real checking // has been added to the type checking of function signatures. - if (auto it = deduced_args.find(deduced_param); - it == deduced_args.end()) { + if (auto it = deduced_type_args.find(deduced_param); + it == deduced_type_args.end()) { return FATAL_COMPILATION_ERROR(e->source_loc()) << "could not deduce type argument for type parameter " - << deduced_param->name(); + << deduced_param->name() << "\n" + << "in " << call; } } - parameters = Substitute(deduced_args, parameters); - return_type = Substitute(deduced_args, return_type); + parameters = Substitute(deduced_type_args, parameters); + return_type = Substitute(deduced_type_args, return_type); + // Find impls for all the impl bindings of the function std::map, ValueNodeView> impls; for (Nonnull impl_binding : @@ -756,9 +839,10 @@ auto TypeChecker::TypeCheckExp(Nonnull e, case Value::Kind::InterfaceType: { ASSIGN_OR_RETURN( ValueNodeView impl, - impl_scope.Resolve(impl_binding->interface(), - deduced_args[impl_binding->type_var()], - e->source_loc())); + impl_scope.Resolve( + impl_binding->interface(), + deduced_type_args[impl_binding->type_var()], + e->source_loc())); impls.emplace(impl_binding, impl); break; } @@ -772,6 +856,8 @@ auto TypeChecker::TypeCheckExp(Nonnull e, } call.set_impls(impls); } else { + // No deduced parameters. Check that the argument types + // are convertible to the parameter types. RETURN_IF_ERROR(ExpectType(e->source_loc(), "call", parameters, &call.argument().static_type())); } @@ -779,10 +865,62 @@ auto TypeChecker::TypeCheckExp(Nonnull e, call.set_value_category(ValueCategory::Let); return Success(); } + case Value::Kind::TypeOfClassType: { + // This case handles the application of a generic class to + // a type argument, such as Point(i32). + const ClassDeclaration& class_decl = + cast(call.function().static_type()) + .class_type() + .declaration(); + BindingMap generic_args; + if (class_decl.type_params().has_value()) { + if (trace_) { + llvm::outs() << "pattern matching type params and args "; + } + ASSIGN_OR_RETURN(Nonnull arg, + InterpExp(&call.argument(), arena_, trace_)); + CHECK(PatternMatch(&(*class_decl.type_params())->value(), arg, + call.source_loc(), std::nullopt, generic_args)); + } else { + return FATAL_COMPILATION_ERROR(call.source_loc()) + << "attempt to instantiate a non-generic class: " << *e; + } + // Find impls for all the impl bindings of the class. + std::map, ValueNodeView> impls; + for (const auto& [binding, val] : generic_args) { + if (binding->impl_binding().has_value()) { + Nonnull impl_binding = + *binding->impl_binding(); + switch (impl_binding->interface()->kind()) { + case Value::Kind::InterfaceType: { + ASSIGN_OR_RETURN(ValueNodeView impl, + impl_scope.Resolve(impl_binding->interface(), + generic_args[binding], + call.source_loc())); + impls.emplace(impl_binding, impl); + break; + } + case Value::Kind::TypeType: + break; + default: + return FATAL_COMPILATION_ERROR(e->source_loc()) + << "unexpected type of deduced parameter " + << *impl_binding->interface(); + } + } + } + Nonnull class_type = + arena_->New(&class_decl, generic_args, impls); + call.set_impls(impls); + call.set_static_type(class_type); + call.set_value_category(ValueCategory::Let); + return Success(); + } default: { return FATAL_COMPILATION_ERROR(e->source_loc()) << "in call, expected a function\n" - << *e; + << *e << "\nnot an operator of type " + << call.function().static_type() << "\n"; } } break; @@ -854,6 +992,41 @@ auto TypeChecker::TypeCheckExp(Nonnull e, } } +void TypeChecker::AddPatternImpls(Nonnull p, ImplScope& impl_scope) { + switch (p->kind()) { + case PatternKind::GenericBinding: { + auto& binding = cast(*p); + CHECK(binding.impl_binding().has_value()); + Nonnull impl_binding = *binding.impl_binding(); + impl_scope.Add(impl_binding->interface(), + *impl_binding->type_var()->symbolic_identity(), + impl_binding); + return; + } + case PatternKind::TuplePattern: { + auto& tuple = cast(*p); + for (Nonnull field : tuple.fields()) { + AddPatternImpls(field, impl_scope); + } + return; + } + case PatternKind::AlternativePattern: { + auto& alternative = cast(*p); + AddPatternImpls(&alternative.arguments(), impl_scope); + return; + } + case PatternKind::VarPattern: { + auto& var_pattern = cast(*p); + AddPatternImpls(&var_pattern.pattern(), impl_scope); + return; + } + case PatternKind::ExpressionPattern: + case PatternKind::AutoPattern: + case PatternKind::BindingPattern: + return; + } +} + auto TypeChecker::TypeCheckPattern( Nonnull p, std::optional> expected, const ImplScope& impl_scope, ValueCategory enclosing_value_category) @@ -887,8 +1060,9 @@ auto TypeChecker::TypeCheckPattern( RETURN_IF_ERROR( ExpectType(p->source_loc(), "name binding", type, *expected)); } else { + BindingMap generic_args; if (!PatternMatch(type, *expected, binding.type().source_loc(), - std::nullopt)) { + std::nullopt, generic_args)) { return FATAL_COMPILATION_ERROR(binding.type().source_loc()) << "Type pattern '" << *type << "' does not match actual type '" << **expected << "'"; @@ -907,6 +1081,27 @@ auto TypeChecker::TypeCheckPattern( } return Success(); } + case PatternKind::GenericBinding: { + auto& binding = cast(*p); + RETURN_IF_ERROR(TypeCheckExp(&binding.type(), impl_scope)); + ASSIGN_OR_RETURN(Nonnull type, + InterpExp(&binding.type(), arena_, trace_)); + if (expected) { + return FATAL_COMPILATION_ERROR(binding.type().source_loc()) + << "Generic binding may not occur in pattern with expected " + "type: " + << binding; + } + binding.set_static_type(type); + ASSIGN_OR_RETURN(Nonnull val, + InterpPattern(&binding, arena_, trace_)); + binding.set_symbolic_identity(val); + Nonnull impl_binding = arena_->New( + binding.source_loc(), &binding, &binding.static_type()); + binding.set_impl_binding(impl_binding); + SetValue(&binding, val); + return Success(); + } case PatternKind::TuplePattern: { auto& tuple = cast(*p); std::vector> field_types; @@ -927,6 +1122,9 @@ auto TypeChecker::TypeCheckPattern( } RETURN_IF_ERROR(TypeCheckPattern(field, expected_field_type, impl_scope, enclosing_value_category)); + if (trace_) + llvm::outs() << "finished checking tuple pattern field " << *field + << "\n"; field_types.push_back(&field->static_type()); } tuple.set_static_type(arena_->New(std::move(field_types))); @@ -1183,28 +1381,19 @@ auto TypeChecker::ExpectReturnOnAllPaths( // TODO: Add checking to function definitions to ensure that // all deduced type parameters will be deduced. auto TypeChecker::DeclareFunctionDeclaration(Nonnull f, - const ImplScope& enclosing_scope) + const ImplScope& impl_scope) -> ErrorOr { if (trace_) { llvm::outs() << "** declaring function " << f->name() << "\n"; } // Bring the deduced parameters into scope for (Nonnull deduced : f->deduced_parameters()) { - RETURN_IF_ERROR(TypeCheckExp(&deduced->type(), enclosing_scope)); - SetConstantValue(deduced, arena_->New(deduced)); - ASSIGN_OR_RETURN(Nonnull deduced_type, + RETURN_IF_ERROR(TypeCheckExp(&deduced->type(), impl_scope)); + deduced->set_symbolic_identity(arena_->New(deduced)); + ASSIGN_OR_RETURN(Nonnull type_of_type, InterpExp(&deduced->type(), arena_, trace_)); - deduced->set_static_type(deduced_type); + deduced->set_static_type(type_of_type); } - // Type check the receiver pattern - if (f->is_method()) { - RETURN_IF_ERROR(TypeCheckPattern(&f->me_pattern(), std::nullopt, - enclosing_scope, ValueCategory::Let)); - } - // Type check the parameter pattern - RETURN_IF_ERROR(TypeCheckPattern(&f->param_pattern(), std::nullopt, - enclosing_scope, ValueCategory::Let)); - // Create the impl_bindings std::vector> impl_bindings; for (Nonnull deduced : f->deduced_parameters()) { @@ -1214,6 +1403,23 @@ auto TypeChecker::DeclareFunctionDeclaration(Nonnull f, impl_binding->set_static_type(&deduced->static_type()); impl_bindings.push_back(impl_binding); } + // Bring the impl bindings into scope. + ImplScope function_scope; + function_scope.AddParent(&impl_scope); + for (Nonnull impl_binding : impl_bindings) { + CHECK(impl_binding->type_var()->symbolic_identity().has_value()); + function_scope.Add(impl_binding->interface(), + *impl_binding->type_var()->symbolic_identity(), + impl_binding); + } + // Type check the receiver pattern. + if (f->is_method()) { + RETURN_IF_ERROR(TypeCheckPattern(&f->me_pattern(), std::nullopt, + function_scope, ValueCategory::Let)); + } + // Type check the parameter pattern. + RETURN_IF_ERROR(TypeCheckPattern(&f->param_pattern(), std::nullopt, + function_scope, ValueCategory::Let)); // Evaluate the return type, if we can do so without examining the body. if (std::optional> return_expression = @@ -1221,7 +1427,7 @@ auto TypeChecker::DeclareFunctionDeclaration(Nonnull f, return_expression.has_value()) { // We ignore the return value because return type expressions can't bring // new types into scope. - RETURN_IF_ERROR(TypeCheckExp(*return_expression, enclosing_scope)); + RETURN_IF_ERROR(TypeCheckExp(*return_expression, function_scope)); // Should we be doing SetConstantValue instead? -Jeremy // And shouldn't the type of this be Type? ASSIGN_OR_RETURN(Nonnull ret_type, @@ -1235,15 +1441,7 @@ auto TypeChecker::DeclareFunctionDeclaration(Nonnull f, return FATAL_COMPILATION_ERROR(f->return_term().source_loc()) << "Function declaration has deduced return type but no body"; } - // Bring the impl bindings into scope - ImplScope function_scope; - function_scope.AddParent(&enclosing_scope); - for (Nonnull impl_binding : impl_bindings) { - function_scope.Add(impl_binding->interface(), - *impl_binding->type_var()->constant_value(), - impl_binding); - } - RETURN_IF_ERROR(TypeCheckStmt(*f->body(), enclosing_scope)); + RETURN_IF_ERROR(TypeCheckStmt(*f->body(), function_scope)); if (!f->return_term().is_omitted()) { RETURN_IF_ERROR(ExpectReturnOnAllPaths(f->body(), f->source_loc())); } @@ -1268,7 +1466,8 @@ auto TypeChecker::DeclareFunctionDeclaration(Nonnull f, } if (trace_) { - llvm::outs() << "** finished declaring function " << f->name() << "\n"; + llvm::outs() << "** finished declaring function " << f->name() + << " of type " << f->static_type() << "\n"; } return Success(); } @@ -1287,10 +1486,13 @@ auto TypeChecker::TypeCheckFunctionDeclaration(Nonnull f, function_scope.AddParent(&impl_scope); for (Nonnull impl_binding : cast(f->static_type()).impl_bindings()) { + CHECK(impl_binding->type_var()->symbolic_identity().has_value()); function_scope.Add(impl_binding->interface(), - *impl_binding->type_var()->constant_value(), + *impl_binding->type_var()->symbolic_identity(), impl_binding); } + if (trace_) + llvm::outs() << function_scope; RETURN_IF_ERROR(TypeCheckStmt(*f->body(), function_scope)); if (!f->return_term().is_omitted()) { RETURN_IF_ERROR(ExpectReturnOnAllPaths(f->body(), f->source_loc())); @@ -1305,16 +1507,45 @@ auto TypeChecker::TypeCheckFunctionDeclaration(Nonnull f, auto TypeChecker::DeclareClassDeclaration(Nonnull class_decl, ImplScope& enclosing_scope) -> ErrorOr { - // The declarations of the members may refer to the class, so we - // must set the constant value of the class and its static type - // before we start processing the members. - Nonnull class_type = - arena_->New(class_decl); - SetConstantValue(class_decl, class_type); - class_decl->set_static_type(arena_->New(class_type)); + if (trace_) { + llvm::outs() << "** declaring class " << class_decl->name() << "\n"; + } + if (class_decl->type_params().has_value()) { + ImplScope class_scope; + class_scope.AddParent(&enclosing_scope); + RETURN_IF_ERROR(TypeCheckPattern(*class_decl->type_params(), std::nullopt, + class_scope, ValueCategory::Let)); + AddPatternImpls(*class_decl->type_params(), class_scope); + if (trace_) { + llvm::outs() << class_scope; + } - for (Nonnull m : class_decl->members()) { - RETURN_IF_ERROR(DeclareDeclaration(m, enclosing_scope)); + Nonnull class_type = + arena_->New(class_decl); + SetConstantValue(class_decl, class_type); + class_decl->set_static_type(arena_->New(class_type)); + + for (Nonnull m : class_decl->members()) { + RETURN_IF_ERROR(DeclareDeclaration(m, class_scope)); + } + + // TODO: when/how to bring impls in generic class into scope? + } else { + // The declarations of the members may refer to the class, so we + // must set the constant value of the class and its static type + // before we start processing the members. + Nonnull class_type = + arena_->New(class_decl); + SetConstantValue(class_decl, class_type); + class_decl->set_static_type(arena_->New(class_type)); + + for (Nonnull m : class_decl->members()) { + RETURN_IF_ERROR(DeclareDeclaration(m, enclosing_scope)); + } + } + if (trace_) { + llvm::outs() << "** finished declaring class " << class_decl->name() + << "\n"; } return Success(); } @@ -1322,8 +1553,22 @@ auto TypeChecker::DeclareClassDeclaration(Nonnull class_decl, auto TypeChecker::TypeCheckClassDeclaration( Nonnull class_decl, const ImplScope& impl_scope) -> ErrorOr { + if (trace_) { + llvm::outs() << "** checking class " << class_decl->name() << "\n"; + } + ImplScope class_scope; + class_scope.AddParent(&impl_scope); + if (class_decl->type_params().has_value()) { + AddPatternImpls(*class_decl->type_params(), class_scope); + } + if (trace_) { + llvm::outs() << class_scope; + } for (Nonnull m : class_decl->members()) { - RETURN_IF_ERROR(TypeCheckDeclaration(m, impl_scope)); + RETURN_IF_ERROR(TypeCheckDeclaration(m, class_scope)); + } + if (trace_) { + llvm::outs() << "** finished checking class " << class_decl->name() << "\n"; } return Success(); } @@ -1339,7 +1584,7 @@ auto TypeChecker::DeclareInterfaceDeclaration( RETURN_IF_ERROR(TypeCheckExp(&iface_decl->self()->type(), enclosing_scope)); iface_decl->self()->set_static_type( arena_->New(iface_decl->self())); - SetConstantValue(iface_decl->self(), &iface_decl->self()->static_type()); + iface_decl->self()->set_symbolic_identity(&iface_decl->self()->static_type()); for (Nonnull m : iface_decl->members()) { RETURN_IF_ERROR(DeclareDeclaration(m, enclosing_scope)); diff --git a/executable_semantics/interpreter/type_checker.h b/executable_semantics/interpreter/type_checker.h index 06933781557c..8e0e973243e0 100644 --- a/executable_semantics/interpreter/type_checker.h +++ b/executable_semantics/interpreter/type_checker.h @@ -36,9 +36,9 @@ class TypeChecker { // inside the argument type. // The `deduced` parameter is an accumulator, that is, it holds the // results so-far. - static auto ArgumentDeduction(SourceLocation source_loc, BindingMap& deduced, - Nonnull param, - Nonnull arg) -> ErrorOr; + auto ArgumentDeduction(SourceLocation source_loc, BindingMap& deduced, + Nonnull param_type, + Nonnull arg_type) -> ErrorOr; // Traverses the AST rooted at `e`, populating the static_type() of all nodes // and ensuring they follow Carbon's typing rules. @@ -93,6 +93,9 @@ class TypeChecker { const ImplScope& enclosing_scope) -> ErrorOr; + // Add the impls from the pattern into the given `impl_scope`. + void AddPatternImpls(Nonnull p, ImplScope& impl_scope); + // Checks the statements and (runtime) expressions within the // declaration, such as the body of a function. // Dispatches to one of the following functions. @@ -135,6 +138,29 @@ class TypeChecker { auto ExpectIsConcreteType(SourceLocation source_loc, Nonnull value) -> ErrorOr; + // Returns the field names of the class together with their types. + auto FieldTypes(const NominalClassType& class_type) + -> std::vector; + + // Returns true if source_fields and destination_fields contain the same set + // of names, and each value in source_fields is implicitly convertible to + // the corresponding value in destination_fields. All values in both arguments + // must be types. + auto FieldTypesImplicitlyConvertible( + llvm::ArrayRef source_fields, + llvm::ArrayRef destination_fields); + + // Returns true if *source is implicitly convertible to *destination. *source + // and *destination must be concrete types. + auto IsImplicitlyConvertible(Nonnull source, + Nonnull destination) -> bool; + + // Check whether `actual` is implicitly convertible to `expected` + // and halt with a fatal compilation error if it is not. + auto ExpectType(SourceLocation source_loc, const std::string& context, + Nonnull expected, Nonnull actual) + -> ErrorOr; + auto Substitute(const std::map, Nonnull>& dict, Nonnull type) -> Nonnull; diff --git a/executable_semantics/interpreter/value.cpp b/executable_semantics/interpreter/value.cpp index 09c6a60ff26b..45ff9e2d8b23 100644 --- a/executable_semantics/interpreter/value.cpp +++ b/executable_semantics/interpreter/value.cpp @@ -42,7 +42,12 @@ static auto GetMember(Nonnull arena, Nonnull v, FindMember(f, witness->declaration().members()); mem_decl.has_value()) { const auto& fun_decl = cast(**mem_decl); - return arena->New(&fun_decl, v); + if (fun_decl.is_method()) { + return arena->New(&fun_decl, v); + } else { + // Class function. + return *fun_decl.constant_value(); + } } else { return FATAL_COMPILATION_ERROR(source_loc) << "member " << f << " not in " << *witness; @@ -67,7 +72,9 @@ static auto GetMember(Nonnull arena, Nonnull v, // Look for a field std::optional> field = cast(object.inits()).FindField(f); - if (field == std::nullopt) { + if (field.has_value()) { + return *field; + } else { // Look for a method in the object's class const auto& class_type = cast(object.type()); std::optional> func = @@ -78,14 +85,18 @@ static auto GetMember(Nonnull arena, Nonnull v, << class_type; } else if ((*func)->declaration().is_method()) { // Found a method. Turn it into a bound method. - const auto& m = cast(**func); - return arena->New(&m.declaration(), &object); + const FunctionValue& m = cast(**func); + return arena->New(&m.declaration(), &object, + class_type.type_args(), + class_type.witnesses()); } else { // Found a class function - return *func; + Nonnull fun = arena->New( + &(*func)->declaration(), class_type.type_args(), + class_type.witnesses()); + return fun; } } - return *field; } case Value::Kind::ChoiceType: { const auto& choice = cast(*v); @@ -96,14 +107,17 @@ static auto GetMember(Nonnull arena, Nonnull v, return arena->New(f, choice.name()); } case Value::Kind::NominalClassType: { - const auto& class_type = cast(*v); + // Access a class function. + const NominalClassType& class_type = cast(*v); std::optional> fun = class_type.FindFunction(f); if (fun == std::nullopt) { return FATAL_RUNTIME_ERROR(source_loc) << "class function " << f << " not in " << *v; } - return *fun; + return arena->New(&(*fun)->declaration(), + class_type.type_args(), + class_type.witnesses()); } default: FATAL() << "field access not allowed for value " << *v; @@ -293,6 +307,28 @@ void Value::Print(llvm::raw_ostream& out) const { case Value::Kind::NominalClassType: { const auto& class_type = cast(*this); out << "class " << class_type.declaration().name(); + if (!class_type.type_args().empty()) { + out << "("; + llvm::ListSeparator sep; + for (const auto& [bind, val] : class_type.type_args()) { + out << sep << bind->name() << " = " << *val; + } + out << ")"; + } + if (!class_type.impls().empty()) { + out << " impls "; + llvm::ListSeparator sep; + for (const auto& [impl_bind, impl] : class_type.impls()) { + out << sep << impl; + } + } + if (!class_type.witnesses().empty()) { + out << " witnesses "; + llvm::ListSeparator sep; + for (const auto& [impl_bind, witness] : class_type.witnesses()) { + out << sep << *witness; + } + } break; } case Value::Kind::InterfaceType: { @@ -302,7 +338,7 @@ void Value::Print(llvm::raw_ostream& out) const { } case Value::Kind::Witness: { const auto& witness = cast(*this); - out << "impl " << *witness.declaration().impl_type() << " as " + out << "witness " << *witness.declaration().impl_type() << " as " << witness.declaration().interface(); break; } @@ -310,7 +346,7 @@ void Value::Print(llvm::raw_ostream& out) const { out << "choice " << cast(*this).name(); break; case Value::Kind::VariableType: - out << cast(*this).binding().name(); + out << cast(*this).binding(); break; case Value::Kind::ContinuationValue: { out << cast(*this).stack(); @@ -410,8 +446,18 @@ auto TypeEqual(Nonnull t1, Nonnull t2) -> bool { return true; } case Value::Kind::NominalClassType: - return cast(*t1).declaration().name() == - cast(*t2).declaration().name(); + if (cast(*t1).declaration().name() != + cast(*t2).declaration().name()) { + return false; + } + for (const auto& [ty_var1, ty1] : + cast(*t1).type_args()) { + if (!TypeEqual(ty1, + cast(*t2).type_args().at(ty_var1))) { + return false; + } + } + return true; case Value::Kind::InterfaceType: return cast(*t1).declaration().name() == cast(*t2).declaration().name(); @@ -591,23 +637,6 @@ auto NominalClassType::FindFunction(const std::string& name) const return std::nullopt; } -auto FieldTypes(const NominalClassType& class_type) -> std::vector { - std::vector field_types; - for (Nonnull m : class_type.declaration().members()) { - switch (m->kind()) { - case DeclarationKind::VariableDeclaration: { - const auto& var = cast(*m); - field_types.push_back({.name = var.binding().name(), - .value = &var.binding().static_type()}); - break; - } - default: - break; - } - } - return field_types; -} - auto FindMember(const std::string& name, llvm::ArrayRef> members) -> std::optional> { @@ -623,7 +652,11 @@ auto FindMember(const std::string& name, } void ImplBinding::Print(llvm::raw_ostream& out) const { - out << "impl " << *type_var_ << " as " << *iface_; + out << "impl binding " << *type_var_ << " as " << *iface_; +} + +void ImplBinding::PrintID(llvm::raw_ostream& out) const { + out << *type_var_ << " as " << *iface_; } } // namespace Carbon diff --git a/executable_semantics/interpreter/value.h b/executable_semantics/interpreter/value.h index 0190fe33522d..c53f72871389 100644 --- a/executable_semantics/interpreter/value.h +++ b/executable_semantics/interpreter/value.h @@ -129,6 +129,15 @@ class FunctionValue : public Value { explicit FunctionValue(Nonnull declaration) : Value(Kind::FunctionValue), declaration_(declaration) {} + explicit FunctionValue(Nonnull declaration, + const BindingMap& type_args, + const std::map, + Nonnull>& wits) + : Value(Kind::FunctionValue), + declaration_(declaration), + type_args_(type_args), + witnesses_(wits) {} + static auto classof(const Value* value) -> bool { return value->kind() == Kind::FunctionValue; } @@ -137,8 +146,17 @@ class FunctionValue : public Value { return *declaration_; } + auto type_args() const -> const BindingMap& { return type_args_; } + + auto witnesses() const + -> const std::map, const Witness*>& { + return witnesses_; + } + private: Nonnull declaration_; + BindingMap type_args_; + std::map, Nonnull> witnesses_; }; // A bound method value. It includes the receiver object. @@ -150,6 +168,17 @@ class BoundMethodValue : public Value { declaration_(declaration), receiver_(receiver) {} + explicit BoundMethodValue(Nonnull declaration, + Nonnull receiver, + const BindingMap& type_args, + const std::map, + Nonnull>& wits) + : Value(Kind::BoundMethodValue), + declaration_(declaration), + receiver_(receiver), + type_args_(type_args), + witnesses_(wits) {} + static auto classof(const Value* value) -> bool { return value->kind() == Kind::BoundMethodValue; } @@ -160,9 +189,18 @@ class BoundMethodValue : public Value { auto receiver() const -> Nonnull { return receiver_; } + auto type_args() const -> const BindingMap& { return type_args_; } + + auto witnesses() const + -> const std::map, Nonnull>& { + return witnesses_; + } + private: Nonnull declaration_; Nonnull receiver_; + BindingMap type_args_; + std::map, Nonnull> witnesses_; }; // The value of a location in memory. @@ -465,16 +503,68 @@ class StructType : public Value { }; // A class type. +// TODO: Consider splitting this class into several classes. class NominalClassType : public Value { public: + // Construct a non-generic class type or a generic class type that has + // not yet been applied to type arguments. explicit NominalClassType(Nonnull declaration) : Value(Kind::NominalClassType), declaration_(declaration) {} + // Construct a class type that represents the result of applying the + // given generic class to the `type_args`. + explicit NominalClassType(Nonnull declaration, + const BindingMap& type_args) + : Value(Kind::NominalClassType), + declaration_(declaration), + type_args_(type_args) {} + + // Construct a class type that represents the result of applying the + // given generic class to the `type_args` and that records the result of the + // compile-time search for any required impls. + explicit NominalClassType( + Nonnull declaration, const BindingMap& type_args, + const std::map, ValueNodeView>& impls) + : Value(Kind::NominalClassType), + declaration_(declaration), + type_args_(type_args), + impls_(impls) {} + + // Construct a fully instantiated generic class type to represent the + // run-time type of an object. + explicit NominalClassType(Nonnull declaration, + const BindingMap& type_args, + const std::map, + Nonnull>& wits) + : Value(Kind::NominalClassType), + declaration_(declaration), + type_args_(type_args), + witnesses_(wits) {} + static auto classof(const Value* value) -> bool { return value->kind() == Kind::NominalClassType; } auto declaration() const -> const ClassDeclaration& { return *declaration_; } + auto type_args() const -> const BindingMap& { return type_args_; } + + // Maps each of the class's generic parameters to the AST node that + // identifies the witness table for the corresponding argument. + // Should not be called on 1) a non-generic class, 2) a generic-class + // that is not instantiated, or 3) a fully instantiated runtime type + // of a generic class. + auto impls() const + -> const std::map, ValueNodeView>& { + return impls_; + } + + // Maps each of the class's generic parameters to the witness table + // for the corresponding argument. Should only be called on a fully + // instantiated runtime type of a generic class. + auto witnesses() const + -> const std::map, Nonnull>& { + return witnesses_; + } // Returns the value of the function named `name` in this class, or // nullopt if there is no such function. @@ -483,9 +573,11 @@ class NominalClassType : public Value { private: Nonnull declaration_; + BindingMap type_args_; + std::map, ValueNodeView> impls_; + std::map, Nonnull> witnesses_; }; -auto FieldTypes(const NominalClassType&) -> std::vector; // Return the declaration of the member with the given name. auto FindMember(const std::string& name, llvm::ArrayRef> members) diff --git a/executable_semantics/syntax/parser.ypp b/executable_semantics/syntax/parser.ypp index 70a6f324df29..16b7c7efce80 100644 --- a/executable_semantics/syntax/parser.ypp +++ b/executable_semantics/syntax/parser.ypp @@ -159,6 +159,7 @@ %type > paren_pattern %type > tuple_pattern %type > maybe_empty_tuple_pattern +%type >> type_params %type > paren_pattern_base %type > paren_pattern_contents %type > alternative @@ -580,6 +581,8 @@ non_expression_pattern: $$ = arena->New(context.source_loc(), $1, $3, std::nullopt); } +| binding_lhs COLON_BANG expression + { $$ = arena->New(context.source_loc(), $1, $3); } | paren_pattern { $$ = $1; } | postfix_expression tuple_pattern @@ -868,11 +871,17 @@ alternative_list_contents: $$.push_back(std::move($3)); } ; +type_params: + // Empty + { $$ = std::nullopt; } +| tuple_pattern + { $$ = $1; } +; declaration: function_declaration { $$ = $1; } -| CLASS identifier LEFT_CURLY_BRACE declaration_list RIGHT_CURLY_BRACE - { $$ = arena->New(context.source_loc(), $2, $4); } +| CLASS identifier type_params LEFT_CURLY_BRACE declaration_list RIGHT_CURLY_BRACE + { $$ = arena->New(context.source_loc(), $2, $3, $5); } | CHOICE identifier LEFT_CURLY_BRACE alternative_list RIGHT_CURLY_BRACE { $$ = arena->New(context.source_loc(), $2, $4); } | VAR variable_declaration SEMICOLON diff --git a/executable_semantics/testdata/generic_class/class_function.carbon b/executable_semantics/testdata/generic_class/class_function.carbon new file mode 100644 index 000000000000..f58e7a743af7 --- /dev/null +++ b/executable_semantics/testdata/generic_class/class_function.carbon @@ -0,0 +1,41 @@ +// 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 +// +// RUN: %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: result: 0 + +package ExecutableSemanticsTest api; + +interface Number { + fn Zero() -> Self; + fn Add[me: Self](other: Self) -> Self; +} + +class Point(T:! Number) { + fn Origin() -> Point(T) { + return {.x = T.Zero(), .y = T.Zero()}; + } + fn SumXY(p: Point(T)) -> T { + return p.x.Add(p.y); + } + fn SumFn() -> (__Fn(Point(T)) -> T) { + return Point(T).SumXY; + } + var x: T; + var y: T; +} + +external impl i32 as Number { + fn Zero() -> i32 { return 0; } + fn Add[me: i32](other: i32) -> i32 { return me + other; } +} + +fn Main() -> i32 { + var p: Point(i32) = Point(i32).Origin(); + return p.SumFn()(p); +} diff --git a/executable_semantics/testdata/generic_class/fail_argument_deduction.carbon b/executable_semantics/testdata/generic_class/fail_argument_deduction.carbon new file mode 100644 index 000000000000..dffaa0239e89 --- /dev/null +++ b/executable_semantics/testdata/generic_class/fail_argument_deduction.carbon @@ -0,0 +1,27 @@ +// 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 +// +// RUN: %{not} %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/generic_class/fail_argument_deduction.carbon:26: type error in argument deduction + +package ExecutableSemanticsTest api; + +class Point(T:! Type) { + var x: T; + var y: T; +} + +fn FirstOfTwoPoints[T:! Type](a: Point(T), b: Point(T)) -> Point(T) { + return a; +} + +fn Main() -> i32 { + var p: Point(i32) = {.x = 0, .y = 1}; + var q: Point(Bool) = {.x = true, .y = false}; + return FirstOfTwoPoints(p, q).x; +} diff --git a/executable_semantics/testdata/generic_class/fail_bad_parameter_type.carbon b/executable_semantics/testdata/generic_class/fail_bad_parameter_type.carbon new file mode 100644 index 000000000000..efc090447c3d --- /dev/null +++ b/executable_semantics/testdata/generic_class/fail_bad_parameter_type.carbon @@ -0,0 +1,31 @@ +// 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 +// +// RUN: %{not} %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/generic_class/fail_bad_parameter_type.carbon:16: unexpected type of deduced parameter i32 + +package ExecutableSemanticsTest api; + +class Point(T:! i32) { + + fn Origin(zero: T) -> Point(T) { + return {.x = zero, .y = zero}; + } + + fn GetX[me: Point(T)]() -> T { + return me.x; + } + + var x: T; + var y: T; +} + +fn Main() -> i32 { + var p: Point(i32) = Point(i32).Origin(0); + return p.GetX(); +} diff --git a/executable_semantics/testdata/generic_class/fail_field_access_on_generic.carbon b/executable_semantics/testdata/generic_class/fail_field_access_on_generic.carbon new file mode 100644 index 000000000000..8fb6bc506364 --- /dev/null +++ b/executable_semantics/testdata/generic_class/fail_field_access_on_generic.carbon @@ -0,0 +1,20 @@ +// 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 +// +// RUN: %{not} %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/generic_class/fail_field_access_on_generic.carbon:15: field access, unexpected T:! Type of non-interface type Type in a.x + +package ExecutableSemanticsTest api; + +fn BadFieldAccess[T:! Type](a: T) -> T { + return a.x; +} + +fn Main() -> i32 { + return BadFieldAccess(0); +} diff --git a/executable_semantics/testdata/generic_class/fail_generic_in_pattern.carbon b/executable_semantics/testdata/generic_class/fail_generic_in_pattern.carbon new file mode 100644 index 000000000000..0f5acac1fb7e --- /dev/null +++ b/executable_semantics/testdata/generic_class/fail_generic_in_pattern.carbon @@ -0,0 +1,22 @@ +// Part of the Carbon Language project, under the Apache License v2.0 with LLVM +// Exceptions. See /LICENSE for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +// RUN: %{not} %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/generic_class/fail_generic_in_pattern.carbon:17: Generic binding may not occur in pattern with expected type: T:! i32 + +package ExecutableSemanticsTest api; + +fn Main() -> i32 { + var t: auto = 5; + match (t) { + case T:! i32 => + return 0; + default => + return 1; + } +} diff --git a/executable_semantics/testdata/generic_class/fail_instantiate_non_generic.carbon b/executable_semantics/testdata/generic_class/fail_instantiate_non_generic.carbon new file mode 100644 index 000000000000..39009345b309 --- /dev/null +++ b/executable_semantics/testdata/generic_class/fail_instantiate_non_generic.carbon @@ -0,0 +1,25 @@ +// 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 +// +// RUN: %{not} %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/generic_class/fail_instantiate_non_generic.carbon:23: attempt to instantiate a non-generic class: Point(i32) + +package ExecutableSemanticsTest api; + +class Point { + fn Origin() -> Point { + return {.x = 0, .y = 0}; + } + var x: i32; + var y: i32; +} + +fn Main() -> i32 { + var p: Point(i32) = Point.Origin(); + return 0; +} diff --git a/executable_semantics/testdata/generic_class/fail_point_equal.carbon b/executable_semantics/testdata/generic_class/fail_point_equal.carbon new file mode 100644 index 000000000000..20b275963f61 --- /dev/null +++ b/executable_semantics/testdata/generic_class/fail_point_equal.carbon @@ -0,0 +1,23 @@ +// 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 +// +// RUN: %{not} %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/generic_class/fail_point_equal.carbon:21: type error in name binding: 'class Point(T = i32)' is not implicitly convertible to 'class Point(T = Bool)' + +package ExecutableSemanticsTest api; + +class Point(T:! Type) { + var x: T; + var y: T; +} + +fn Main() -> i32 { + var p: Point(i32) = {.x = 0, .y = 0}; + var q: Point(Bool) = p; + return 0; +} diff --git a/executable_semantics/testdata/generic_class/generic_class_substitution.carbon b/executable_semantics/testdata/generic_class/generic_class_substitution.carbon new file mode 100644 index 000000000000..ee204b056995 --- /dev/null +++ b/executable_semantics/testdata/generic_class/generic_class_substitution.carbon @@ -0,0 +1,39 @@ +// 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 +// +// RUN: %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: result: 0 + +package ExecutableSemanticsTest api; + +interface Number { + fn Zero() -> Self; + fn Add[me: Self](other: Self) -> Self; +} + +class Point(T:! Number) { + fn Origin() -> Point(T) { + return {.x = T.Zero(), .y = T.Zero()}; + } + var x: T; + var y: T; +} + +external impl i32 as Number { + fn Zero() -> i32 { return 0; } + fn Add[me: i32](other: i32) -> i32 { return me + other; } +} + +fn SumXY[U:! Number](p: Point(U)) -> U { + return p.Origin().x.Add(p.y); +} + +fn Main() -> i32 { + var p: Point(i32) = {.x = 0, .y = 0}; + return SumXY(p); +} diff --git a/executable_semantics/testdata/generic_class/generic_fun_and_class.carbon b/executable_semantics/testdata/generic_class/generic_fun_and_class.carbon new file mode 100644 index 000000000000..fc7364f8c6dc --- /dev/null +++ b/executable_semantics/testdata/generic_class/generic_fun_and_class.carbon @@ -0,0 +1,44 @@ +// 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 +// +// RUN: %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: result: 0 + +package ExecutableSemanticsTest api; + +interface Number { + fn Zero() -> Self; + fn Add[me: Self](other: Self) -> Self; +} + +class Point(T:! Number) { + var x: T; + var y: T; +} + +fn Origin[U :! Number](other: U) -> Point(U) { + return {.x = U.Zero(), .y = U.Zero()}; +} + +fn Clone[U :! Number](other: Point(U)) -> Point(U) { + return {.x = other.x, .y = other.y}; +} + +fn SumXY[U :! Number](other: Point(U)) -> U { + return other.x.Add(other.y); +} + +external impl i32 as Number { + fn Zero() -> i32 { return 0; } + fn Add[me: i32](other: i32) -> i32 { return me + other; } +} + +fn Main() -> i32 { + var p: Point(i32) = Origin(0); + return SumXY(Clone(p)); +} diff --git a/executable_semantics/testdata/generic_class/generic_point.carbon b/executable_semantics/testdata/generic_class/generic_point.carbon new file mode 100644 index 000000000000..cd139f19d8f1 --- /dev/null +++ b/executable_semantics/testdata/generic_class/generic_point.carbon @@ -0,0 +1,31 @@ +// 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 +// +// RUN: %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: result: 0 + +package ExecutableSemanticsTest api; + +class Point(T:! Type) { + + fn Origin(zero: T) -> Point(T) { + return {.x = zero, .y = zero}; + } + + fn GetX[me: Point(T)]() -> T { + return me.x; + } + + var x: T; + var y: T; +} + +fn Main() -> i32 { + var p: Point(i32) = Point(i32).Origin(0); + return p.GetX(); +} diff --git a/executable_semantics/testdata/generic_class/point_with_interface.carbon b/executable_semantics/testdata/generic_class/point_with_interface.carbon new file mode 100644 index 000000000000..a717da98fdff --- /dev/null +++ b/executable_semantics/testdata/generic_class/point_with_interface.carbon @@ -0,0 +1,41 @@ +// 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 +// +// RUN: %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: result: 0 + +package ExecutableSemanticsTest api; + +interface Number { + fn Zero() -> Self; + fn Add[me: Self](other: Self) -> Self; +} + +class Point(T:! Number) { + fn Origin() -> Point(T) { + return {.x = T.Zero(), .y = T.Zero()}; + } + fn Clone[me: Point(T)]() -> Point(T) { + return {.x = me.x, .y = me.y}; + } + fn SumXY[me: Point(T)]() -> T { + return me.x.Add(me.y); + } + var x: T; + var y: T; +} + +external impl i32 as Number { + fn Zero() -> i32 { return 0; } + fn Add[me: i32](other: i32) -> i32 { return me + other; } +} + +fn Main() -> i32 { + var p: Point(i32) = Point(i32).Origin(); + return p.Clone().SumXY(); +} diff --git a/executable_semantics/testdata/generic_function/fail_not_addable.carbon b/executable_semantics/testdata/generic_function/fail_not_addable.carbon index 1b8755a4aae2..b04c97be7a10 100644 --- a/executable_semantics/testdata/generic_function/fail_not_addable.carbon +++ b/executable_semantics/testdata/generic_function/fail_not_addable.carbon @@ -9,7 +9,7 @@ // AUTOUPDATE: %{executable_semantics} %s // CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/generic_function/fail_not_addable.carbon:17: type error in addition(1) // CHECK: expected: i32 -// CHECK: actual: T +// CHECK: actual: T:! Type package ExecutableSemanticsTest api; diff --git a/executable_semantics/testdata/interface/class_function.carbon b/executable_semantics/testdata/interface/class_function.carbon new file mode 100644 index 000000000000..258962e9cfe7 --- /dev/null +++ b/executable_semantics/testdata/interface/class_function.carbon @@ -0,0 +1,44 @@ +// 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 +// +// RUN: %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: result: 0 + +package ExecutableSemanticsTest api; + +interface Vector { + fn Zero() -> Self; + fn Add[me: Self](b: Self) -> Self; + fn Scale[me: Self](v: i32) -> Self; +} + +class Point { + var x: i32; + var y: i32; + impl Point as Vector { + fn Zero() -> Point { + return {.x = 0, .y = 0}; + } + fn Add[me: Point](b: Point) -> Point { + return {.x = me.x + b.x, .y = me.y + b.y}; + } + fn Scale[me: Point](v: i32) -> Point { + return {.x = me.x * v, .y = me.y * v}; + } + } +} + +fn AddAndScaleGeneric[T:! Vector](a: T, s: i32) -> T { + return a.Add(T.Zero()).Scale(s); +} + +fn Main() -> i32 { + var a: Point = {.x = 2, .y = 1}; + var p: Point = AddAndScaleGeneric(a, 5); + return p.x - 10; +} diff --git a/executable_semantics/testdata/interface/fail_interface_missing_member.carbon b/executable_semantics/testdata/interface/fail_interface_missing_member.carbon new file mode 100644 index 000000000000..636a110ae219 --- /dev/null +++ b/executable_semantics/testdata/interface/fail_interface_missing_member.carbon @@ -0,0 +1,36 @@ +// 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 +// +// RUN: %{not} %{executable_semantics} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{executable_semantics} --trace %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{executable_semantics} %s +// CHECK: COMPILATION ERROR: {{.*}}/executable_semantics/testdata/interface/fail_interface_missing_member.carbon:19: field access, Scale not in Vector + +package ExecutableSemanticsTest api; + +interface Vector { + fn Add[me: Self](b: Self) -> Self; +} + +fn ScaleGeneric[T:! Vector](a: T, s: i32) -> T { + return a.Scale(s); +} + +class Point { + var x: i32; + var y: i32; + impl Point as Vector { + fn Add[me: Point](b: Point) -> Point { + return {.x = me.x + b.x, .y = me.y + b.y}; + } + } +} + +fn Main() -> i32 { + var a: Point = {.x = 3, .y = 1}; + var b: Point = ScaleGeneric(a, 2); + return b.x - 6; +}