From 981427588b38bbc4269cfe4d58a7b5e906c9e1a2 Mon Sep 17 00:00:00 2001 From: Calvin Date: Sat, 28 Jan 2023 13:00:49 -0700 Subject: [PATCH] Refactor `name` property into `FunctionDeclaration` (#2555) As per the `TODO` comment in the file, the `CallableDeclaration::name` property was refactored into `FunctionDeclaration`, which is currently the only namable callable declaration kind. I was able to rely on the existing `GetName` function to replace call sites of the accessor function with nearly identical swaps. Fixes #2531 --- explorer/ast/declaration.cpp | 8 +++++--- explorer/ast/declaration.h | 21 +++++++++++---------- explorer/interpreter/resolve_names.cpp | 6 ++++-- explorer/interpreter/type_checker.cpp | 16 ++++++++++------ explorer/interpreter/value.cpp | 4 ++-- 5 files changed, 32 insertions(+), 23 deletions(-) diff --git a/explorer/ast/declaration.cpp b/explorer/ast/declaration.cpp index 85d0c0d3d75c..7bf82dfb86ac 100644 --- a/explorer/ast/declaration.cpp +++ b/explorer/ast/declaration.cpp @@ -156,7 +156,7 @@ void Declaration::PrintID(llvm::raw_ostream& out) const { out << "fn " << cast(*this).name(); break; case DeclarationKind::DestructorDeclaration: - out << cast(*this).name(); + out << *GetName(*this); break; case DeclarationKind::ClassDeclaration: { const auto& class_decl = cast(*this); @@ -224,7 +224,7 @@ auto GetName(const Declaration& declaration) case DeclarationKind::FunctionDeclaration: return cast(declaration).name(); case DeclarationKind::DestructorDeclaration: - return cast(declaration).name(); + return "destructor"; case DeclarationKind::ClassDeclaration: return cast(declaration).name(); case DeclarationKind::MixinDeclaration: { @@ -368,7 +368,9 @@ auto FunctionDeclaration::Create(Nonnull arena, } void CallableDeclaration::PrintDepth(int depth, llvm::raw_ostream& out) const { - out << "fn " << name_ << " "; + auto name = GetName(*this); + CARBON_CHECK(name) << "Unexpected missing name for `" << *this << "`."; + out << "fn " << *name << " "; if (!deduced_parameters_.empty()) { out << "["; llvm::ListSeparator sep; diff --git a/explorer/ast/declaration.h b/explorer/ast/declaration.h index 164313e9504f..3725f58b0df7 100644 --- a/explorer/ast/declaration.h +++ b/explorer/ast/declaration.h @@ -131,7 +131,7 @@ enum class VirtualOverride { None, Abstract, Virtual, Impl }; class CallableDeclaration : public Declaration { public: - CallableDeclaration(AstNodeKind kind, SourceLocation loc, std::string name, + CallableDeclaration(AstNodeKind kind, SourceLocation loc, std::vector> deduced_params, std::optional> self_pattern, Nonnull param_pattern, @@ -139,7 +139,6 @@ class CallableDeclaration : public Declaration { std::optional> body, VirtualOverride virt_override) : Declaration(kind, loc), - name_(std::move(name)), deduced_parameters_(std::move(deduced_params)), self_pattern_(self_pattern), param_pattern_(param_pattern), @@ -149,8 +148,6 @@ class CallableDeclaration : public Declaration { void PrintDepth(int depth, llvm::raw_ostream& out) const; - // TODO: Move name() and name_ to FunctionDeclaration - auto name() const -> const std::string& { return name_; } auto deduced_parameters() const -> llvm::ArrayRef> { return deduced_parameters_; @@ -173,7 +170,6 @@ class CallableDeclaration : public Declaration { auto is_method() const -> bool { return self_pattern_.has_value(); } private: - std::string name_; std::vector> deduced_parameters_; std::optional> self_pattern_; Nonnull param_pattern_; @@ -204,13 +200,18 @@ class FunctionDeclaration : public CallableDeclaration { std::optional> body, VirtualOverride virt_override) : CallableDeclaration(AstNodeKind::FunctionDeclaration, source_loc, - std::move(name), std::move(deduced_params), - self_pattern, param_pattern, return_term, body, - virt_override) {} + std::move(deduced_params), self_pattern, + param_pattern, return_term, body, virt_override), + name_(std::move(name)) {} static auto classof(const AstNode* node) -> bool { return InheritsFromFunctionDeclaration(node->kind()); } + + auto name() const -> const std::string& { return name_; } + + private: + std::string name_; }; class DestructorDeclaration : public CallableDeclaration { @@ -232,8 +233,8 @@ class DestructorDeclaration : public CallableDeclaration { ReturnTerm return_term, std::optional> body) : CallableDeclaration(AstNodeKind::DestructorDeclaration, source_loc, - "destructor", std::move(deduced_params), - self_pattern, param_pattern, return_term, body, + std::move(deduced_params), self_pattern, + param_pattern, return_term, body, // TODO: Add virtual destructors VirtualOverride::None) {} diff --git a/explorer/interpreter/resolve_names.cpp b/explorer/interpreter/resolve_names.cpp index 7be76f0b0019..2d9fcc6d4490 100644 --- a/explorer/interpreter/resolve_names.cpp +++ b/explorer/interpreter/resolve_names.cpp @@ -576,7 +576,9 @@ static auto ResolveNames(Declaration& declaration, StaticScope& enclosing_scope, auto& function = cast(declaration); StaticScope function_scope; function_scope.AddParent(&enclosing_scope); - enclosing_scope.MarkDeclared(function.name()); + const auto name = GetName(function); + CARBON_CHECK(name) << "Unexpected missing name for `" << function << "`."; + enclosing_scope.MarkDeclared(std::string(*name)); for (Nonnull binding : function.deduced_parameters()) { CARBON_RETURN_IF_ERROR(ResolveNames(*binding, function_scope)); } @@ -590,7 +592,7 @@ static auto ResolveNames(Declaration& declaration, StaticScope& enclosing_scope, CARBON_RETURN_IF_ERROR(ResolveNames( **function.return_term().type_expression(), function_scope)); } - enclosing_scope.MarkUsable(function.name()); + enclosing_scope.MarkUsable(std::string(*name)); if (function.body().has_value() && bodies != ResolveFunctionBodies::Skip) { CARBON_RETURN_IF_ERROR(ResolveNames(**function.body(), function_scope)); diff --git a/explorer/interpreter/type_checker.cpp b/explorer/interpreter/type_checker.cpp index 295fa74d337b..2a9c09115e3f 100644 --- a/explorer/interpreter/type_checker.cpp +++ b/explorer/interpreter/type_checker.cpp @@ -4455,8 +4455,10 @@ auto TypeChecker::ExpectReturnOnAllPaths( auto TypeChecker::DeclareCallableDeclaration(Nonnull f, const ScopeInfo& scope_info) -> ErrorOr { + const auto name = GetName(*f); + CARBON_CHECK(name) << "Unexpected missing name for `" << *f << "`."; if (trace_stream_) { - **trace_stream_ << "** declaring function " << f->name() << "\n"; + **trace_stream_ << "** declaring function " << *name << "\n"; } ImplScope function_scope; function_scope.AddParent(scope_info.innermost_scope); @@ -4537,7 +4539,7 @@ auto TypeChecker::DeclareCallableDeclaration(Nonnull f, CARBON_FATAL() << "f is not a callable declaration"; } - if (f->name() == "Main") { + if (name == "Main") { if (!f->return_term().type_expression().has_value()) { return ProgramError(f->return_term().source_loc()) << "`Main` must have an explicit return type"; @@ -4553,8 +4555,8 @@ auto TypeChecker::DeclareCallableDeclaration(Nonnull f, } if (trace_stream_) { - **trace_stream_ << "** finished declaring function " << f->name() - << " of type " << f->static_type() << "\n"; + **trace_stream_ << "** finished declaring function " << *name << " of type " + << f->static_type() << "\n"; } return Success(); } @@ -4562,8 +4564,10 @@ auto TypeChecker::DeclareCallableDeclaration(Nonnull f, auto TypeChecker::TypeCheckCallableDeclaration(Nonnull f, const ImplScope& impl_scope) -> ErrorOr { + auto name = GetName(*f); + CARBON_CHECK(name) << "Unexpected missing name for `" << *f << "`."; if (trace_stream_) { - **trace_stream_ << "** checking function " << f->name() << "\n"; + **trace_stream_ << "** checking function " << *name << "\n"; } // If f->return_term().is_auto(), the function body was already // type checked in DeclareFunctionDeclaration. @@ -4583,7 +4587,7 @@ auto TypeChecker::TypeCheckCallableDeclaration(Nonnull f, } } if (trace_stream_) { - **trace_stream_ << "** finished checking function " << f->name() << "\n"; + **trace_stream_ << "** finished checking function " << *name << "\n"; } return Success(); } diff --git a/explorer/interpreter/value.cpp b/explorer/interpreter/value.cpp index fcf31350eb57..2276addbc4c4 100644 --- a/explorer/interpreter/value.cpp +++ b/explorer/interpreter/value.cpp @@ -1166,7 +1166,7 @@ auto FindFunction(std::string_view name, break; } case DeclarationKind::FunctionDeclaration: { - const auto& fun = cast(*member); + const auto& fun = cast(*member); if (fun.name() == name) { return &cast(**fun.constant_value()); } @@ -1194,7 +1194,7 @@ auto MixinPseudoType::FindFunction(const std::string_view& name) const break; } case DeclarationKind::FunctionDeclaration: { - const auto& fun = cast(*member); + const auto& fun = cast(*member); if (fun.name() == name) { return &cast(**fun.constant_value()); }