diff --git a/explorer/ast/declaration.cpp b/explorer/ast/declaration.cpp index be4ebfe60511..a0c27c9af8ba 100644 --- a/explorer/ast/declaration.cpp +++ b/explorer/ast/declaration.cpp @@ -237,35 +237,56 @@ void ReturnTerm::Print(llvm::raw_ostream& out) const { } } -// Look for the `me` parameter in the `deduced_parameters` -// and put it in the `me_pattern`. -static auto MoveMeParameterToPattern( - SourceLocation source_loc, std::optional>& me_pattern, - const std::vector>& deduced_params) - -> ErrorOr>> { +namespace { + +// The deduced parameters of a function declaration. +struct DeducedParameters { + // The `me` parameter, if any. + std::optional> me_pattern; + + // All other deduced parameters. std::vector> resolved_params; +}; + +// Split the `me` pattern (if any) out of `deduced_params`. +auto SplitDeducedParameters( + SourceLocation source_loc, + const std::vector>& deduced_params) + -> ErrorOr { + DeducedParameters result; for (Nonnull param : deduced_params) { switch (param->kind()) { case AstNodeKind::GenericBinding: - resolved_params.push_back(&cast(*param)); + result.resolved_params.push_back(&cast(*param)); break; case AstNodeKind::BindingPattern: { - Nonnull bp = &cast(*param); - if (me_pattern.has_value() || bp->name() != "me") { + Nonnull binding = &cast(*param); + if (binding->name() != "me") { return CompilationError(source_loc) << "illegal binding pattern in implicit parameter list"; } - me_pattern = bp; + if (result.me_pattern.has_value()) { + return CompilationError(source_loc) + << "parameter list cannot contain more than one `me` " + "parameter"; + } + result.me_pattern = binding; break; } case AstNodeKind::AddrPattern: { - Nonnull abp = &cast(*param); - Nonnull bp = &cast(abp->binding()); - if (me_pattern.has_value() || bp->name() != "me") { + Nonnull addr_pattern = &cast(*param); + Nonnull binding = + &cast(addr_pattern->binding()); + if (binding->name() != "me") { return CompilationError(source_loc) << "illegal binding pattern in implicit parameter list"; } - me_pattern = abp; + if (result.me_pattern.has_value()) { + return CompilationError(source_loc) + << "parameter list cannot contain more than one `me` " + "parameter"; + } + result.me_pattern = addr_pattern; break; } default: @@ -273,8 +294,9 @@ static auto MoveMeParameterToPattern( << "illegal AST node in implicit parameter list"; } } - return resolved_params; + return result; } +} // namespace auto DestructorDeclaration::CreateDestructor( Nonnull arena, SourceLocation source_loc, @@ -282,31 +304,27 @@ auto DestructorDeclaration::CreateDestructor( Nonnull param_pattern, ReturnTerm return_term, std::optional> body) -> ErrorOr> { - std::vector> resolved_params; - std::optional> me_pattern; - CARBON_ASSIGN_OR_RETURN( - resolved_params, - MoveMeParameterToPattern(source_loc, me_pattern, deduced_params)); + DeducedParameters split_params; + CARBON_ASSIGN_OR_RETURN(split_params, + SplitDeducedParameters(source_loc, deduced_params)); return arena->New( - source_loc, std::move(resolved_params), me_pattern, param_pattern, - return_term, body); + source_loc, std::move(split_params.resolved_params), + split_params.me_pattern, param_pattern, return_term, body); } auto FunctionDeclaration::Create(Nonnull arena, SourceLocation source_loc, std::string name, std::vector> deduced_params, - std::optional> me_pattern, Nonnull param_pattern, ReturnTerm return_term, std::optional> body) -> ErrorOr> { - std::vector> resolved_params; - CARBON_ASSIGN_OR_RETURN( - resolved_params, - MoveMeParameterToPattern(source_loc, me_pattern, deduced_params)); - return arena->New(source_loc, name, - std::move(resolved_params), me_pattern, - param_pattern, return_term, body); + DeducedParameters split_params; + CARBON_ASSIGN_OR_RETURN(split_params, + SplitDeducedParameters(source_loc, deduced_params)); + return arena->New( + source_loc, name, std::move(split_params.resolved_params), + split_params.me_pattern, param_pattern, return_term, body); } void CallableDeclaration::PrintDepth(int depth, llvm::raw_ostream& out) const { diff --git a/explorer/ast/declaration.h b/explorer/ast/declaration.h index 74fab7a6bd10..b359c8790634 100644 --- a/explorer/ast/declaration.h +++ b/explorer/ast/declaration.h @@ -158,7 +158,6 @@ class FunctionDeclaration : public CallableDeclaration { static auto Create(Nonnull arena, SourceLocation source_loc, std::string name, std::vector> deduced_params, - std::optional> me_pattern, Nonnull param_pattern, ReturnTerm return_term, std::optional> body) diff --git a/explorer/syntax/parser.ypp b/explorer/syntax/parser.ypp index 5d7abf98f58b..7821d5bb6a79 100644 --- a/explorer/syntax/parser.ypp +++ b/explorer/syntax/parser.ypp @@ -180,7 +180,6 @@ %type > struct_type_literal_contents %type > tuple %type binding_lhs -%type >> receiver %type > variable_declaration %type > paren_expression_base %type > paren_expression_contents @@ -1041,19 +1040,11 @@ impl_deduced_params: | FORALL LEFT_SQUARE_BRACKET deduced_param_list RIGHT_SQUARE_BRACKET { $$ = $3; } ; -receiver: - // Empty - { $$ = std::nullopt; } -| LEFT_CURLY_BRACE variable_declaration RIGHT_CURLY_BRACE - { $$ = $2; } -| LEFT_CURLY_BRACE ADDR variable_declaration RIGHT_CURLY_BRACE - { $$ = arena->New(context.source_loc(), $3); } -; function_declaration: - FN identifier deduced_params receiver maybe_empty_tuple_pattern return_term block + FN identifier deduced_params maybe_empty_tuple_pattern return_term block { ErrorOr fn = FunctionDeclaration::Create( - arena, context.source_loc(), $2, $3, $4, $5, $6, $7); + arena, context.source_loc(), $2, $3, $4, $5, $6); if (fn.ok()) { $$ = *fn; } else { @@ -1061,10 +1052,10 @@ function_declaration: YYERROR; } } -| FN identifier deduced_params receiver maybe_empty_tuple_pattern return_term SEMICOLON +| FN identifier deduced_params maybe_empty_tuple_pattern return_term SEMICOLON { ErrorOr fn = FunctionDeclaration::Create( - arena, context.source_loc(), $2, $3, $4, $5, $6, std::nullopt); + arena, context.source_loc(), $2, $3, $4, $5, std::nullopt); if (fn.ok()) { $$ = *fn; } else {