diff --git a/common/fuzzing/carbon.proto b/common/fuzzing/carbon.proto index d5f96bcd2729..8676e759c1d8 100644 --- a/common/fuzzing/carbon.proto +++ b/common/fuzzing/carbon.proto @@ -121,6 +121,28 @@ message ArrayTypeLiteral { optional Expression size = 2; } +message IsWhereClause { + optional Expression type = 1; + optional Expression constraint = 2; +} + +message EqualsWhereClause { + optional Expression lhs = 1; + optional Expression rhs = 2; +} + +message WhereClause { + oneof kind { + IsWhereClause is = 1; + EqualsWhereClause equals = 2; + } +} + +message WhereExpression { + optional Expression base = 1; + repeated WhereClause clauses = 2; +} + message Expression { oneof kind { CallExpression call = 1; @@ -145,6 +167,7 @@ message Expression { UnimplementedExpression unimplemented_expression = 20; ArrayTypeLiteral array_type_literal = 21; CompoundMemberAccessExpression compound_member_access = 22; + WhereExpression where = 23; } } diff --git a/common/fuzzing/proto_to_carbon.cpp b/common/fuzzing/proto_to_carbon.cpp index 94187374a87a..cd76deb8da23 100644 --- a/common/fuzzing/proto_to_carbon.cpp +++ b/common/fuzzing/proto_to_carbon.cpp @@ -352,6 +352,33 @@ static auto ExpressionToCarbon(const Fuzzing::Expression& expression, out << "]"; break; } + + case Fuzzing::Expression::kWhere: { + const Fuzzing::WhereExpression& where = expression.where(); + ExpressionToCarbon(where.base(), out); + out << " where "; + llvm::ListSeparator sep(" and "); + for (const auto& clause : where.clauses()) { + out << sep; + switch (clause.kind_case()) { + case Fuzzing::WhereClause::kIs: + ExpressionToCarbon(clause.is().type(), out); + out << " is "; + ExpressionToCarbon(clause.is().constraint(), out); + break; + case Fuzzing::WhereClause::kEquals: + ExpressionToCarbon(clause.equals().lhs(), out); + out << " == "; + ExpressionToCarbon(clause.equals().rhs(), out); + break; + case Fuzzing::WhereClause::KIND_NOT_SET: + // Arbitrary default to avoid invalid syntax. + out << ".Self == .Self"; + break; + } + } + break; + } } } diff --git a/explorer/ast/ast_node.h b/explorer/ast/ast_node.h index 6e156f8f23e1..4d3865b14b80 100644 --- a/explorer/ast/ast_node.h +++ b/explorer/ast/ast_node.h @@ -45,7 +45,7 @@ class AstNode { // 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. + // Print identifying information about the node, such as its name. virtual void PrintID(llvm::raw_ostream& out) const = 0; LLVM_DUMP_METHOD void Dump() const { Print(llvm::errs()); } diff --git a/explorer/ast/ast_rtti.txt b/explorer/ast/ast_rtti.txt index e3f367fcabff..39f6b7fc9b87 100644 --- a/explorer/ast/ast_rtti.txt +++ b/explorer/ast/ast_rtti.txt @@ -59,6 +59,10 @@ abstract class Expression : AstNode; class IdentifierExpression : Expression; class IntrinsicExpression : Expression; class IfExpression : Expression; + class WhereExpression : Expression; class UnimplementedExpression : Expression; class ArrayTypeLiteral : Expression; class InstantiateImpl : Expression; +abstract class WhereClause : AstNode; + class IsWhereClause : WhereClause; + class EqualsWhereClause : WhereClause; diff --git a/explorer/ast/expression.cpp b/explorer/ast/expression.cpp index 72a8e6996874..ba860d382e91 100644 --- a/explorer/ast/expression.cpp +++ b/explorer/ast/expression.cpp @@ -171,6 +171,15 @@ void Expression::Print(llvm::raw_ostream& out) const { << if_expr.then_expression() << " else " << if_expr.else_expression(); break; } + case ExpressionKind::WhereExpression: { + const auto& where = cast(*this); + out << where.base() << " where "; + llvm::ListSeparator sep(" and "); + for (const WhereClause* clause : where.clauses()) { + out << sep << *clause; + } + break; + } case ExpressionKind::InstantiateImpl: { const auto& inst_impl = cast(*this); out << "instantiate " << *inst_impl.generic_impl(); @@ -246,6 +255,7 @@ void Expression::PrintID(llvm::raw_ostream& out) const { case ExpressionKind::SimpleMemberAccessExpression: case ExpressionKind::CompoundMemberAccessExpression: case ExpressionKind::IfExpression: + case ExpressionKind::WhereExpression: case ExpressionKind::TupleLiteral: case ExpressionKind::StructLiteral: case ExpressionKind::StructTypeLiteral: @@ -261,4 +271,23 @@ void Expression::PrintID(llvm::raw_ostream& out) const { } } +WhereClause::~WhereClause() = default; + +void WhereClause::Print(llvm::raw_ostream& out) const { + switch (kind()) { + case WhereClauseKind::IsWhereClause: { + auto& clause = cast(*this); + out << clause.type() << " is " << clause.constraint(); + break; + } + case WhereClauseKind::EqualsWhereClause: { + auto& clause = cast(*this); + out << clause.lhs() << " == " << clause.rhs(); + break; + } + } +} + +void WhereClause::PrintID(llvm::raw_ostream& out) const { out << "..."; } + } // namespace Carbon diff --git a/explorer/ast/expression.h b/explorer/ast/expression.h index 72bbfdb9f0f0..6436e1f036d0 100644 --- a/explorer/ast/expression.h +++ b/explorer/ast/expression.h @@ -666,6 +666,108 @@ class IfExpression : public Expression { Nonnull else_expression_; }; +// A clause appearing on the right-hand side of a `where` operator that forms a +// more precise constraint from a more general one. +class WhereClause : public AstNode { + public: + ~WhereClause() 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 InheritsFromWhereClause(node->kind()); + } + + auto kind() const -> WhereClauseKind { + return static_cast(root_kind()); + } + + protected: + WhereClause(WhereClauseKind kind, SourceLocation source_loc) + : AstNode(static_cast(kind), source_loc) {} +}; + +// An `is` where clause. +// +// For example, `ConstraintA where .Type is ConstraintB` requires that the +// associated type `.Type` implements the constraint `ConstraintB`. +class IsWhereClause : public WhereClause { + public: + explicit IsWhereClause(SourceLocation source_loc, Nonnull type, + Nonnull constraint) + : WhereClause(WhereClauseKind::IsWhereClause, source_loc), + type_(type), + constraint_(constraint) {} + + static auto classof(const AstNode* node) { + return InheritsFromIsWhereClause(node->kind()); + } + + auto type() const -> const Expression& { return *type_; } + auto type() -> Expression& { return *type_; } + + auto constraint() const -> const Expression& { return *constraint_; } + auto constraint() -> Expression& { return *constraint_; } + + private: + Nonnull type_; + Nonnull constraint_; +}; + +// An `==` where clause. +// +// For example, `Constraint where .Type == i32` requires that the associated +// type `.Type` is `i32`. +class EqualsWhereClause : public WhereClause { + public: + explicit EqualsWhereClause(SourceLocation source_loc, + Nonnull lhs, Nonnull rhs) + : WhereClause(WhereClauseKind::EqualsWhereClause, source_loc), + lhs_(lhs), + rhs_(rhs) {} + + static auto classof(const AstNode* node) { + return InheritsFromEqualsWhereClause(node->kind()); + } + + auto lhs() const -> const Expression& { return *lhs_; } + auto lhs() -> Expression& { return *lhs_; } + + auto rhs() const -> const Expression& { return *rhs_; } + auto rhs() -> Expression& { return *rhs_; } + + private: + Nonnull lhs_; + Nonnull rhs_; +}; + +// A `where` expression: `AddableWith(i32) where .Result == i32`. +class WhereExpression : public Expression { + public: + explicit WhereExpression(SourceLocation source_loc, Nonnull base, + std::vector> clauses) + : Expression(AstNodeKind::WhereExpression, source_loc), + base_(base), + clauses_(std::move(clauses)) {} + + static auto classof(const AstNode* node) -> bool { + return InheritsFromWhereExpression(node->kind()); + } + + auto base() const -> const Expression& { return *base_; } + auto base() -> Expression& { return *base_; } + + auto clauses() const -> llvm::ArrayRef> { + return clauses_; + } + auto clauses() -> llvm::ArrayRef> { return clauses_; } + + private: + Nonnull base_; + std::vector> clauses_; +}; + // Instantiate a generic impl. class InstantiateImpl : public Expression { public: diff --git a/explorer/fuzzing/ast_to_proto.cpp b/explorer/fuzzing/ast_to_proto.cpp index 28e5d3d19d02..45a5a62f699c 100644 --- a/explorer/fuzzing/ast_to_proto.cpp +++ b/explorer/fuzzing/ast_to_proto.cpp @@ -179,6 +179,35 @@ static auto ExpressionToProto(const Expression& expression) break; } + case ExpressionKind::WhereExpression: { + const auto& where = cast(expression); + auto* where_proto = expression_proto.mutable_where(); + *where_proto->mutable_base() = ExpressionToProto(where.base()); + for (const WhereClause* where : where.clauses()) { + Fuzzing::WhereClause clause_proto; + switch (where->kind()) { + case WhereClauseKind::IsWhereClause: { + auto* is_proto = clause_proto.mutable_is(); + *is_proto->mutable_type() = + ExpressionToProto(cast(where)->type()); + *is_proto->mutable_constraint() = + ExpressionToProto(cast(where)->constraint()); + break; + } + case WhereClauseKind::EqualsWhereClause: { + auto* equals_proto = clause_proto.mutable_equals(); + *equals_proto->mutable_lhs() = + ExpressionToProto(cast(where)->lhs()); + *equals_proto->mutable_rhs() = + ExpressionToProto(cast(where)->rhs()); + break; + } + } + *where_proto->add_clauses() = clause_proto; + } + break; + } + case ExpressionKind::IntrinsicExpression: { const auto& intrinsic = cast(expression); auto* intrinsic_proto = expression_proto.mutable_intrinsic(); diff --git a/explorer/interpreter/interpreter.cpp b/explorer/interpreter/interpreter.cpp index e6c340cea19f..08fb63bc29fd 100644 --- a/explorer/interpreter/interpreter.cpp +++ b/explorer/interpreter/interpreter.cpp @@ -412,6 +412,7 @@ auto Interpreter::StepLvalue() -> ErrorOr { case ExpressionKind::ValueLiteral: case ExpressionKind::IntrinsicExpression: case ExpressionKind::IfExpression: + case ExpressionKind::WhereExpression: case ExpressionKind::ArrayTypeLiteral: case ExpressionKind::InstantiateImpl: CARBON_FATAL() << "Can't treat expression as lvalue: " << exp; @@ -1082,6 +1083,10 @@ auto Interpreter::StepExp() -> ErrorOr { } break; } + case ExpressionKind::WhereExpression: { + return todo_.FinishAction( + &cast(exp.static_type()).constraint_type()); + } case ExpressionKind::UnimplementedExpression: CARBON_FATAL() << "Unimplemented: " << exp; case ExpressionKind::ArrayTypeLiteral: { diff --git a/explorer/interpreter/resolve_names.cpp b/explorer/interpreter/resolve_names.cpp index ba78f9792021..6173cad34024 100644 --- a/explorer/interpreter/resolve_names.cpp +++ b/explorer/interpreter/resolve_names.cpp @@ -88,6 +88,9 @@ static auto AddExposedNames(const Declaration& declaration, static auto ResolveNames(Expression& expression, const StaticScope& enclosing_scope) -> ErrorOr; +static auto ResolveNames(WhereClause& clause, + const StaticScope& enclosing_scope) + -> ErrorOr; static auto ResolveNames(Pattern& pattern, StaticScope& enclosing_scope) -> ErrorOr; static auto ResolveNames(Statement& statement, StaticScope& enclosing_scope) @@ -177,6 +180,18 @@ static auto ResolveNames(Expression& expression, ResolveNames(if_expr.else_expression(), enclosing_scope)); break; } + case ExpressionKind::WhereExpression: { + auto& where = cast(expression); + // TODO: Introduce `.Self` into scope? + // StaticScope where_scope; + // where_scope.AddParent(&enclosing_scope); + // where_scope.Add(".Self", ???); + CARBON_RETURN_IF_ERROR(ResolveNames(where.base(), enclosing_scope)); + for (Nonnull clause : where.clauses()) { + CARBON_RETURN_IF_ERROR(ResolveNames(*clause, enclosing_scope)); + } + break; + } case ExpressionKind::ArrayTypeLiteral: { auto& array_literal = cast(expression); CARBON_RETURN_IF_ERROR(ResolveNames( @@ -202,6 +217,29 @@ static auto ResolveNames(Expression& expression, return Success(); } +static auto ResolveNames(WhereClause& clause, + const StaticScope& enclosing_scope) + -> ErrorOr { + switch (clause.kind()) { + case WhereClauseKind::IsWhereClause: { + auto& is_clause = cast(clause); + CARBON_RETURN_IF_ERROR(ResolveNames(is_clause.type(), enclosing_scope)); + CARBON_RETURN_IF_ERROR( + ResolveNames(is_clause.constraint(), enclosing_scope)); + break; + } + case WhereClauseKind::EqualsWhereClause: { + auto& equals_clause = cast(clause); + CARBON_RETURN_IF_ERROR( + ResolveNames(equals_clause.lhs(), enclosing_scope)); + CARBON_RETURN_IF_ERROR( + ResolveNames(equals_clause.rhs(), enclosing_scope)); + break; + } + } + return Success(); +} + static auto ResolveNames(Pattern& pattern, StaticScope& enclosing_scope) -> ErrorOr { switch (pattern.kind()) { diff --git a/explorer/interpreter/type_checker.cpp b/explorer/interpreter/type_checker.cpp index 0394e8b0ab9e..7c3a586dd9fe 100644 --- a/explorer/interpreter/type_checker.cpp +++ b/explorer/interpreter/type_checker.cpp @@ -18,6 +18,7 @@ #include "explorer/interpreter/impl_scope.h" #include "explorer/interpreter/interpreter.h" #include "explorer/interpreter/value.h" +#include "llvm/ADT/DenseSet.h" #include "llvm/ADT/StringExtras.h" #include "llvm/Support/Casting.h" #include "llvm/Support/Error.h" @@ -733,6 +734,140 @@ auto TypeChecker::ArgumentDeduction( } } +// Builder for constraint types. +// +// This type supports incrementally building a constraint type by adding +// constraints one at a time, and will deduplicate the constraints as it goes. +// +// TODO: The deduplication here is very inefficient. We should use value +// canonicalization or hashing or similar to speed this up. +class ConstraintTypeBuilder { + public: + ConstraintTypeBuilder(Nonnull arena, SourceLocation source_loc) + : self_binding_(arena->New( + source_loc, ".Self", arena->New(source_loc))) {} + ConstraintTypeBuilder(Nonnull self_binding) + : self_binding_(self_binding) {} + ConstraintTypeBuilder(Nonnull constraint) + : self_binding_(constraint->self_binding()), + impl_constraints_(constraint->impl_constraints()), + equality_constraints_(constraint->equality_constraints()), + lookup_contexts_(constraint->lookup_contexts()) {} + + // Produce a type that refers to the `.Self` type of the constraint. + auto MakeSelfType(Nonnull arena) const + -> Nonnull { + return arena->New(self_binding_); + } + + // Add an `impl` constraint -- `T is C` if not already present. + void AddImplConstraint(ConstraintType::ImplConstraint impl) { + for (ConstraintType::ImplConstraint existing : impl_constraints_) { + if (TypeEqual(existing.type, impl.type) && + TypeEqual(existing.interface, impl.interface)) { + return; + } + } + impl_constraints_.push_back(std::move(impl)); + } + + // Add an `impl` constraint -- `A == B`, merging as necessary. + void AddEqualityConstraint(ConstraintType::EqualityConstraint equal) { + // Check to see if any of the given values are already part of an equality + // constraint. If so, discard the value and note that we'll merge that + // constraint into the new one. + llvm::SmallDenseSet merged_constraints; + { + size_t kept = 0; + for (size_t i = 0; i != equal.values.size(); ++i) { + if (std::optional found_in = + FindInEqualityConstraints(equal.values[i])) { + merged_constraints.insert(*found_in); + } else { + equal.values[kept++] = equal.values[i]; + } + } + equal.values.resize(kept); + } + + // Merge and discard any constraints that overlapped `equal`. + if (!merged_constraints.empty()) { + size_t kept = 0; + for (size_t i = 0; i != equality_constraints_.size(); ++i) { + if (merged_constraints.contains(i)) { + equal.values.insert(equal.values.end(), + equality_constraints_[i].values.begin(), + equality_constraints_[i].values.end()); + } else { + if (kept != i) { + equality_constraints_[kept] = std::move(equality_constraints_[i]); + } + ++kept; + } + } + equality_constraints_[kept++] = std::move(equal); + equality_constraints_.resize(kept); + } else { + equality_constraints_.push_back(std::move(equal)); + } + } + + // Add a context for qualified name lookup, if not already present. + void AddLookupContext(ConstraintType::LookupContext context) { + for (ConstraintType::LookupContext existing : lookup_contexts_) { + if (ValueEqual(existing.context, context.context)) { + return; + } + } + lookup_contexts_.push_back(std::move(context)); + } + + // Add all the constraints from another constraint type. The constraints must + // not refer to that other constraint type's self binding, because it will no + // longer be in scope. + void Add(Nonnull constraint) { + for (const auto& impl_constraint : constraint->impl_constraints()) { + AddImplConstraint(impl_constraint); + } + + for (const auto& equality_constraint : constraint->equality_constraints()) { + AddEqualityConstraint(equality_constraint); + } + + for (const auto& lookup_context : constraint->lookup_contexts()) { + AddLookupContext(lookup_context); + } + } + + // Convert the builder into a ConstraintType. Note that this consumes the + // builder. + auto Build(Nonnull arena_) && -> Nonnull { + return arena_->New( + self_binding_, std::move(impl_constraints_), + std::move(equality_constraints_), std::move(lookup_contexts_)); + } + + private: + // Find the given value in the equality constraints, returning the index of + // the constraint that contains it, if any. + auto FindInEqualityConstraints(const Value* value) -> std::optional { + for (size_t i = 0; i != equality_constraints_.size(); ++i) { + for (const Value* v : equality_constraints_[i].values) { + if (ValueEqual(value, v)) { + return i; + } + } + } + return std::nullopt; + } + + private: + Nonnull self_binding_; + std::vector impl_constraints_; + std::vector equality_constraints_; + std::vector lookup_contexts_; +}; + auto TypeChecker::Substitute( const std::map, Nonnull>& dict, Nonnull type) const -> Nonnull { @@ -806,39 +941,34 @@ auto TypeChecker::Substitute( } case Value::Kind::ConstraintType: { const auto& constraint = cast(*type); - std::vector impl_constraints; - impl_constraints.reserve(constraint.impl_constraints().size()); + ConstraintTypeBuilder builder(constraint.self_binding()); for (const auto& impl_constraint : constraint.impl_constraints()) { - impl_constraints.push_back( + builder.AddImplConstraint( {.type = Substitute(dict, impl_constraint.type), .interface = cast( Substitute(dict, impl_constraint.interface))}); } - std::vector equality_constraints; - equality_constraints.reserve(constraint.equality_constraints().size()); for (const auto& equality_constraint : constraint.equality_constraints()) { std::vector> values; for (const Value* value : equality_constraint.values) { - values.push_back(Substitute(dict, value)); + // Ensure we don't create any duplicates through substitution. + if (std::find_if(values.begin(), values.end(), [&](const Value* v) { + return ValueEqual(v, value); + }) == values.end()) { + values.push_back(Substitute(dict, value)); + } } - equality_constraints.push_back({.values = values}); + builder.AddEqualityConstraint({.values = std::move(values)}); } - // TODO: Coalesce same-type constraints that are now overlapping. - std::vector lookup_contexts; - lookup_contexts.reserve(constraint.lookup_contexts().size()); for (const auto& lookup_context : constraint.lookup_contexts()) { - lookup_contexts.push_back( + builder.AddLookupContext( {.context = Substitute(dict, lookup_context.context)}); } - // TODO: If the self_binding is substituted, should we track that - // somehow? Nonnull new_constraint = - arena_->New( - constraint.self_binding(), std::move(impl_constraints), - std::move(equality_constraints), std::move(lookup_contexts)); + std::move(builder).Build(arena_); if (trace_stream_) { **trace_stream_ << "substitution: " << constraint << " => " << *new_constraint << "\n"; @@ -991,75 +1121,25 @@ auto TypeChecker::SatisfyImpls( auto TypeChecker::MakeConstraintForInterface( SourceLocation source_loc, Nonnull iface_type) -> Nonnull { - auto* self_binding = arena_->New( - source_loc, ".Self", arena_->New(source_loc)); - auto* self = arena_->New(self_binding); - std::vector impl_constraints = { - ConstraintType::ImplConstraint{.type = self, .interface = iface_type}}; - std::vector equality_constraints = {}; - std::vector lookup_contexts = { - {.context = iface_type}}; - return arena_->New(self_binding, std::move(impl_constraints), - std::move(equality_constraints), - std::move(lookup_contexts)); + ConstraintTypeBuilder builder(arena_, source_loc); + builder.AddImplConstraint( + {.type = builder.MakeSelfType(arena_), .interface = iface_type}); + builder.AddLookupContext({.context = iface_type}); + return std::move(builder).Build(arena_); } auto TypeChecker::CombineConstraints( SourceLocation source_loc, llvm::ArrayRef> constraints) -> Nonnull { - auto* self_binding = arena_->New( - source_loc, ".Self", arena_->New(source_loc)); - auto* self = arena_->New(self_binding); - std::vector impl_constraints; - std::vector equality_constraints; - std::vector lookup_contexts; + ConstraintTypeBuilder builder(arena_, source_loc); + auto* self = builder.MakeSelfType(arena_); for (Nonnull constraint : constraints) { BindingMap map; map[constraint->self_binding()] = self; - // TODO: Remove duplicates - for (ConstraintType::ImplConstraint impl : constraint->impl_constraints()) { - impl_constraints.push_back( - {.type = Substitute(map, impl.type), - .interface = cast(Substitute(map, impl.interface))}); - } - for (ConstraintType::EqualityConstraint same : - constraint->equality_constraints()) { - std::vector> values; - for (const Value* value : same.values) { - values.push_back(Substitute(map, value)); - } - auto AddEqualityConstraint = - [&](std::vector> values) { - // TODO: This is really inefficient. Use value canonicalization or - // hashing or similar to avoid the quadratic scan here. - for (const Value* value : values) { - for (ConstraintType::EqualityConstraint& existing : - equality_constraints) { - for (const Value* existing_value : existing.values) { - if (ValueEqual(value, existing_value)) { - // There is overlap between two equality constraints. - // Combine them into a single constraint. - // TODO: Remove duplicates - existing.values.insert(existing.values.end(), - values.begin(), values.end()); - return; - } - } - } - } - equality_constraints.push_back({.values = std::move(values)}); - }; - AddEqualityConstraint(std::move(values)); - } - // TODO: Remove duplicates - for (ConstraintType::LookupContext lookup : constraint->lookup_contexts()) { - lookup_contexts.push_back({.context = Substitute(map, lookup.context)}); - } + builder.Add(cast(Substitute(map, constraint))); } - return arena_->New(self_binding, std::move(impl_constraints), - std::move(equality_constraints), - std::move(lookup_contexts)); + return std::move(builder).Build(arena_); } auto TypeChecker::DeduceCallBindings( @@ -1915,6 +1995,81 @@ auto TypeChecker::TypeCheckExp(Nonnull e, e->set_value_category(ValueCategory::Let); return Success(); } + case ExpressionKind::WhereExpression: { + auto& where = cast(*e); + CARBON_RETURN_IF_ERROR(TypeCheckExp(&where.base(), impl_scope)); + for (Nonnull clause : where.clauses()) { + CARBON_RETURN_IF_ERROR(TypeCheckWhereClause(clause, impl_scope)); + } + + const ConstraintType* base; + const Value& base_type = where.base().static_type(); + if (auto* constraint_type_type = + dyn_cast(&base_type)) { + base = &constraint_type_type->constraint_type(); + } else if (auto* interface_type_type = + dyn_cast(&base_type)) { + base = MakeConstraintForInterface( + e->source_loc(), &interface_type_type->interface_type()); + } else { + return CompilationError(e->source_loc()) + << "expected constraint as first operand of `where` expression, " + << "found " << base_type; + } + + // Apply the `where` clauses. + ConstraintTypeBuilder builder(base); + for (Nonnull clause : where.clauses()) { + switch (clause->kind()) { + case WhereClauseKind::IsWhereClause: { + const auto& is_clause = cast(*clause); + CARBON_ASSIGN_OR_RETURN( + Nonnull type, + InterpExp(&is_clause.type(), arena_, trace_stream_)); + CARBON_ASSIGN_OR_RETURN( + Nonnull constraint, + InterpExp(&is_clause.constraint(), arena_, trace_stream_)); + if (auto* interface = dyn_cast(constraint)) { + // `where X is Y` produces an `impl` constraint. + builder.AddImplConstraint({.type = type, .interface = interface}); + } else if (auto* constraint_type = + dyn_cast(constraint)) { + // Transform `where .B is (C where .D is E)` into + // `where .B is C and .B.D is E` then add all the resulting + // constraints. + BindingMap map; + map[constraint_type->self_binding()] = type; + builder.Add(cast(Substitute(map, constraint))); + } else { + return CompilationError(is_clause.constraint().source_loc()) + << "expression after `is` does not resolve to a " + "constraint, found value " + << *constraint << " of type " + << is_clause.constraint().static_type(); + } + break; + } + case WhereClauseKind::EqualsWhereClause: { + const auto& equals_clause = cast(*clause); + CARBON_ASSIGN_OR_RETURN( + Nonnull lhs, + InterpExp(&equals_clause.lhs(), arena_, trace_stream_)); + CARBON_ASSIGN_OR_RETURN( + Nonnull rhs, + InterpExp(&equals_clause.rhs(), arena_, trace_stream_)); + if (!ValueEqual(lhs, rhs)) { + builder.AddEqualityConstraint({.values = {lhs, rhs}}); + } + break; + } + } + } + + where.set_static_type( + arena_->New(std::move(builder).Build(arena_))); + where.set_value_category(ValueCategory::Let); + return Success(); + } case ExpressionKind::UnimplementedExpression: CARBON_FATAL() << "Unimplemented: " << *e; case ExpressionKind::ArrayTypeLiteral: { @@ -2009,6 +2164,35 @@ auto TypeChecker::TypeCheckTypeExp(Nonnull type_expression, return type; } +auto TypeChecker::TypeCheckWhereClause(Nonnull clause, + const ImplScope& impl_scope) + -> ErrorOr { + switch (clause->kind()) { + case WhereClauseKind::IsWhereClause: { + auto& is_clause = cast(*clause); + CARBON_RETURN_IF_ERROR(TypeCheckTypeExp(&is_clause.type(), impl_scope)); + CARBON_RETURN_IF_ERROR(TypeCheckExp(&is_clause.constraint(), impl_scope)); + if (!isa( + is_clause.constraint().static_type())) { + return CompilationError(is_clause.constraint().source_loc()) + << "expression after `is` does not resolve to a constraint, " + << "found " << is_clause.constraint().static_type(); + } + return Success(); + } + case WhereClauseKind::EqualsWhereClause: { + auto& equals_clause = cast(*clause); + CARBON_RETURN_IF_ERROR(TypeCheckExp(&equals_clause.lhs(), impl_scope)); + CARBON_RETURN_IF_ERROR(TypeCheckExp(&equals_clause.rhs(), impl_scope)); + CARBON_RETURN_IF_ERROR(ExpectExactType( + clause->source_loc(), "values in `where ==` constraint", + &equals_clause.lhs().static_type(), + &equals_clause.rhs().static_type())); + return Success(); + } + } +} + auto TypeChecker::TypeCheckPattern( Nonnull p, std::optional> expected, ImplScope& impl_scope, ValueCategory enclosing_value_category) diff --git a/explorer/interpreter/type_checker.h b/explorer/interpreter/type_checker.h index 4d4fe131ef07..cb79e935ee4b 100644 --- a/explorer/interpreter/type_checker.h +++ b/explorer/interpreter/type_checker.h @@ -119,6 +119,11 @@ class TypeChecker { const ImplScope& impl_scope, bool concrete = true) -> ErrorOr>; + // Type checks and interprets `clause`, and validates it represents a valid + // `where` clause. + auto TypeCheckWhereClause(Nonnull clause, + const ImplScope& impl_scope) -> ErrorOr; + // Equivalent to TypeCheckExp, but operates on the AST rooted at `p`. // // `expected` is the type that this pattern is expected to have, if the diff --git a/explorer/syntax/lexer.lpp b/explorer/syntax/lexer.lpp index 06ab8e883f24..f4e93ae024f9 100644 --- a/explorer/syntax/lexer.lpp +++ b/explorer/syntax/lexer.lpp @@ -67,6 +67,7 @@ IF "if" IMPL "impl" IMPORT "import" INTERFACE "interface" +IS "is" LEFT_CURLY_BRACE "{" LEFT_PARENTHESIS "(" LEFT_SQUARE_BRACKET "[" @@ -94,6 +95,7 @@ TYPE "Type" UNDERSCORE "_" UNIMPL_EXAMPLE "__unimplemented_example_infix" VAR "var" +WHERE "where" WHILE "while" /* table-end */ @@ -173,6 +175,7 @@ string_literal \"([^\\\"\n\v\f\r]|\\.)*\" {IMPL} { return SIMPLE_TOKEN(IMPL); } {IMPORT} { return SIMPLE_TOKEN(IMPORT); } {INTERFACE} { return SIMPLE_TOKEN(INTERFACE); } +{IS} { return SIMPLE_TOKEN(IS); } {LEFT_CURLY_BRACE} { return SIMPLE_TOKEN(LEFT_CURLY_BRACE); } {LEFT_PARENTHESIS} { return SIMPLE_TOKEN(LEFT_PARENTHESIS); } {LEFT_SQUARE_BRACKET} { return SIMPLE_TOKEN(LEFT_SQUARE_BRACKET); } @@ -197,6 +200,7 @@ string_literal \"([^\\\"\n\v\f\r]|\\.)*\" {UNDERSCORE} { return SIMPLE_TOKEN(UNDERSCORE); } {UNIMPL_EXAMPLE} { return SIMPLE_TOKEN(UNIMPL_EXAMPLE); } {VAR} { return SIMPLE_TOKEN(VAR); } +{WHERE} { return SIMPLE_TOKEN(WHERE); } {WHILE} { return SIMPLE_TOKEN(WHILE); } /* table-end */ diff --git a/explorer/syntax/parser.ypp b/explorer/syntax/parser.ypp index 24c537a09949..83c49779b729 100644 --- a/explorer/syntax/parser.ypp +++ b/explorer/syntax/parser.ypp @@ -141,6 +141,9 @@ %type > and_expression %type > or_lhs %type > or_expression +%type > where_clause +%type >> where_clause_list +%type > where_expression %type > statement_expression %type > if_expression %type > expression @@ -212,6 +215,7 @@ IMPL IMPORT INTERFACE + IS LEFT_CURLY_BRACE LEFT_PARENTHESIS LEFT_SQUARE_BRACKET @@ -239,6 +243,7 @@ UNDERSCORE UNIMPL_EXAMPLE VAR + WHERE WHILE // table-end // Used to track EOF. @@ -522,11 +527,31 @@ or_expression: std::vector>({$1, $3})); } ; +where_clause: + comparison_operand IS comparison_operand + { $$ = arena->New(context.source_loc(), $1, $3); } +| comparison_operand EQUAL_EQUAL comparison_operand + { $$ = arena->New(context.source_loc(), $1, $3); } +; +where_clause_list: + where_clause + { $$ = {$1}; } +| where_clause_list AND where_clause + { + $$ = std::move($1); + $$.push_back($3); + } +; +where_expression: + type_expression WHERE where_clause_list + { $$ = arena->New(context.source_loc(), $1, $3); } +; statement_expression: ref_deref_expression | predicate_expression | and_expression | or_expression +| where_expression ; if_expression: statement_expression diff --git a/explorer/testdata/constraint/combine_equality.carbon b/explorer/testdata/constraint/combine_equality.carbon new file mode 100644 index 000000000000..ea2ee7ce8d52 --- /dev/null +++ b/explorer/testdata/constraint/combine_equality.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} %{explorer} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{explorer} --parser_debug --trace_file=- %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{explorer} %s + +package ExplorerTest api; + +interface I {} +impl i32 as I {} + +fn F(A:! i32, B:! i32, C:! i32, D:! i32, E:! i32, + T:! I where A == B and C == D and C == E and B == D) { + // CHECK: COMPILATION ERROR: {{.*}}/explorer/testdata/constraint/combine_equality.carbon:[[@LINE+1]]: member access, F not in constraint interface I where .Self:! Type is interface I, A:! i32 == B:! i32 == E:! i32 == C:! i32 == D:! i32 + T.F(); +} + +fn Main() -> i32 { + F(1, 1, 1, 1, 1, i32); + return 0; +} diff --git a/explorer/testdata/constraint/combined_interfaces.carbon b/explorer/testdata/constraint/combined_interfaces.carbon index c336854b33dc..b044544e769c 100644 --- a/explorer/testdata/constraint/combined_interfaces.carbon +++ b/explorer/testdata/constraint/combined_interfaces.carbon @@ -7,7 +7,7 @@ // RUN: %{explorer} --parser_debug --trace_file=- %s 2>&1 | \ // RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s // AUTOUPDATE: %{explorer} %s -// CHECK: result: 52 +// CHECK: result: 526 package ExplorerTest api; @@ -16,6 +16,7 @@ interface B { fn H() -> i32; } fn Get1[T:! A & B](n: T) -> i32 { return n.F() + n.H(); } fn Get2[T:! B & A](n: T) -> i32 { return n.G(); } +fn Get3[T:! B & A & A & B & A](n: T) -> i32 { return n.G() + n.H(); } impl i32 as A { fn F() -> i32 { return 1; } @@ -27,5 +28,5 @@ impl i32 as B { fn Main() -> i32 { var z: i32 = 0; - return Get1(z) * 10 + Get2(z); + return Get1(z) * 100 + Get2(z) * 10 + Get3(z); } diff --git a/explorer/testdata/constraint/fail_where_equals_different_types.carbon b/explorer/testdata/constraint/fail_where_equals_different_types.carbon new file mode 100644 index 000000000000..ccd0a04d6f83 --- /dev/null +++ b/explorer/testdata/constraint/fail_where_equals_different_types.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} %{explorer} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{explorer} --parser_debug --trace_file=- %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{explorer} %s + +package ExplorerTest api; + +interface A {} + +// CHECK: COMPILATION ERROR: {{.*}}/explorer/testdata/constraint/fail_where_equals_different_types.carbon:[[@LINE+3]]: type error in values in `where ==` constraint +// CHECK: expected: i32 +// CHECK: actual: Type +alias B = A where 4 == i32; + +fn Main() -> i32 { return 0; } diff --git a/explorer/testdata/constraint/fail_where_is_non_constraint.carbon b/explorer/testdata/constraint/fail_where_is_non_constraint.carbon new file mode 100644 index 000000000000..76421268eb85 --- /dev/null +++ b/explorer/testdata/constraint/fail_where_is_non_constraint.carbon @@ -0,0 +1,18 @@ +// Part of the Carbon Language project, under the Apache License v2.0 with LLVM +// Exceptions. See /LICENSE for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +// RUN: %{not} %{explorer} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{explorer} --parser_debug --trace_file=- %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{explorer} %s + +package ExplorerTest api; + +interface A {} + +// CHECK: COMPILATION ERROR: {{.*}}/explorer/testdata/constraint/fail_where_is_non_constraint.carbon:[[@LINE+1]]: expression after `is` does not resolve to a constraint, found value i32 of type Type +alias B = A where i32 is i32; + +fn Main() -> i32 { return 0; } diff --git a/explorer/testdata/constraint/fail_where_is_non_type.carbon b/explorer/testdata/constraint/fail_where_is_non_type.carbon new file mode 100644 index 000000000000..f20429429078 --- /dev/null +++ b/explorer/testdata/constraint/fail_where_is_non_type.carbon @@ -0,0 +1,18 @@ +// Part of the Carbon Language project, under the Apache License v2.0 with LLVM +// Exceptions. See /LICENSE for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +// RUN: %{not} %{explorer} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{explorer} --parser_debug --trace_file=- %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{explorer} %s + +package ExplorerTest api; + +interface A {} + +// CHECK: COMPILATION ERROR: {{.*}}/explorer/testdata/constraint/fail_where_is_non_type.carbon:[[@LINE+1]]: expression after `is` does not resolve to a constraint, found i32 +alias B = A where i32 is 5; + +fn Main() -> i32 { return 0; } diff --git a/explorer/testdata/constraint/fail_where_non_type_is.carbon b/explorer/testdata/constraint/fail_where_non_type_is.carbon new file mode 100644 index 000000000000..a132a7276c60 --- /dev/null +++ b/explorer/testdata/constraint/fail_where_non_type_is.carbon @@ -0,0 +1,18 @@ +// Part of the Carbon Language project, under the Apache License v2.0 with LLVM +// Exceptions. See /LICENSE for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +// RUN: %{not} %{explorer} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{not} %{explorer} --parser_debug --trace_file=- %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{explorer} %s + +package ExplorerTest api; + +interface A {} + +// CHECK: COMPILATION ERROR: {{.*}}/explorer/testdata/constraint/fail_where_non_type_is.carbon:[[@LINE+1]]: Expected a type, but got 4 +alias B = A where 4 is A; + +fn Main() -> i32 { return 0; } diff --git a/explorer/testdata/constraint/nondependent_where.carbon b/explorer/testdata/constraint/nondependent_where.carbon new file mode 100644 index 000000000000..b151286415b0 --- /dev/null +++ b/explorer/testdata/constraint/nondependent_where.carbon @@ -0,0 +1,24 @@ +// 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: %{explorer} %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes=false %s +// RUN: %{explorer} --parser_debug --trace_file=- %s 2>&1 | \ +// RUN: %{FileCheck} --match-full-lines --allow-unused-prefixes %s +// AUTOUPDATE: %{explorer} %s +// CHECK: result: 1 + +package ExplorerTest api; + +interface A { fn Get() -> Self; } +impl i32 as A { fn Get() -> Self { return 1; } } + +alias AlsoA = A where i32 is A and 4 == 4; + +fn F[T:! AlsoA](x: T) -> T { return T.Get(); } + +fn Main() -> i32 { + var z: i32 = 0; + return F(z); +}