From 7e9d644e1f6647845ca8e317fcc5f747c33245d9 Mon Sep 17 00:00:00 2001 From: Jon Ross-Perkins Date: Fri, 20 Oct 2023 15:40:18 -0700 Subject: [PATCH] Switch File functions, classes, and types to ValueStores (#3316) Building on #3313, start using ValueStore on File. Functions and classes are straightforward. Types here I present as a borderline case where maybe we want a more bespoke API, but maybe this is okay? Most other things probably need a slightly different API, which although I might do that for a consistent interface, felt more out-of-scope for this change. --- toolchain/base/value_store.h | 19 +++-- toolchain/check/context.cpp | 5 +- toolchain/check/declaration_name_stack.cpp | 2 +- toolchain/check/handle_call_expression.cpp | 2 +- toolchain/check/handle_class.cpp | 4 +- toolchain/check/handle_function.cpp | 9 +-- toolchain/check/handle_name.cpp | 3 +- toolchain/check/handle_statement.cpp | 2 +- toolchain/lower/file_context.cpp | 12 ++-- toolchain/lower/handle.cpp | 2 +- toolchain/sem_ir/entry_point.cpp | 2 +- toolchain/sem_ir/file.cpp | 22 ++---- toolchain/sem_ir/file.h | 84 +++++----------------- toolchain/sem_ir/formatter.cpp | 34 +++++---- 14 files changed, 76 insertions(+), 126 deletions(-) diff --git a/toolchain/base/value_store.h b/toolchain/base/value_store.h index 6476b05a66ce..20c4a133bc05 100644 --- a/toolchain/base/value_store.h +++ b/toolchain/base/value_store.h @@ -80,19 +80,25 @@ constexpr StringId StringId::Invalid(StringId::InvalidIndex); // A simple wrapper for accumulating values, providing IDs to later retrieve the // value. This does not do deduplication. -template -class ValueStore : public Printable> { +template +class ValueStore : public Printable> { public: // Stores the value and returns an ID to reference it. - auto Add(typename IdT::IndexedType value) -> IdT { + auto Add(ValueT value) -> IdT { IdT id = IdT(values_.size()); CARBON_CHECK(id.index >= 0) << "Id overflow"; values_.push_back(std::move(value)); return id; } + // Returns a mutable value for an ID. + auto Get(IdT id) -> ValueT& { + CARBON_CHECK(id.index >= 0) << id.index; + return values_[id.index]; + } + // Returns the value for an ID. - auto Get(IdT id) const -> const typename IdT::IndexedType& { + auto Get(IdT id) const -> const ValueT& { CARBON_CHECK(id.index >= 0) << id.index; return values_[id.index]; } @@ -105,8 +111,11 @@ class ValueStore : public Printable> { } } + auto array_ref() const -> llvm::ArrayRef { return values_; } + auto size() const -> int { return values_.size(); } + private: - llvm::SmallVector values_; + llvm::SmallVector values_; }; // Storage for StringRefs. The caller is responsible for ensuring storage is diff --git a/toolchain/check/context.cpp b/toolchain/check/context.cpp index b9b6e859aa50..179c09d61e8b 100644 --- a/toolchain/check/context.cpp +++ b/toolchain/check/context.cpp @@ -246,7 +246,8 @@ auto Context::AddCurrentCodeBlockToFunction() -> void { .GetNodeAs(return_scope_stack().back()) .function_id; semantics_ir() - .GetFunction(function_id) + .functions() + .Get(function_id) .body_block_ids.push_back(node_block_stack().PeekOrAdd()); } @@ -709,7 +710,7 @@ auto Context::CanonicalizeTypeImpl( } auto node_id = make_node(); - auto type_id = semantics_ir_->AddType(node_id); + auto type_id = semantics_ir_->types().Add({.node_id = node_id}); CARBON_CHECK(canonical_types_.insert({node_id, type_id}).second); type_node_storage_.push_back( std::make_unique(canonical_id, type_id)); diff --git a/toolchain/check/declaration_name_stack.cpp b/toolchain/check/declaration_name_stack.cpp index fd1d7731f7ad..b44eced9626e 100644 --- a/toolchain/check/declaration_name_stack.cpp +++ b/toolchain/check/declaration_name_stack.cpp @@ -145,7 +145,7 @@ auto DeclarationNameStack::UpdateScopeIfNeeded(NameContext& name_context) context_->semantics_ir().GetNode(name_context.resolved_node_id); switch (resolved_node.kind()) { case SemIR::ClassDeclaration::Kind: { - auto& class_info = context_->semantics_ir().GetClass( + const auto& class_info = context_->semantics_ir().classes().Get( resolved_node.As().class_id); // TODO: Check that the class is complete rather than that it has a scope. if (class_info.scope_id.is_valid()) { diff --git a/toolchain/check/handle_call_expression.cpp b/toolchain/check/handle_call_expression.cpp index 196d3b378041..ed7674759bdf 100644 --- a/toolchain/check/handle_call_expression.cpp +++ b/toolchain/check/handle_call_expression.cpp @@ -29,7 +29,7 @@ auto HandleCallExpression(Context& context, Parse::Node parse_node) -> bool { } auto function_id = function_name->function_id; - const auto& callable = context.semantics_ir().GetFunction(function_id); + const auto& callable = context.semantics_ir().functions().Get(function_id); // For functions with an implicit return type, the return type is the empty // tuple type. diff --git a/toolchain/check/handle_class.cpp b/toolchain/check/handle_class.cpp index 34db8e4d439f..4cd1cf6cda88 100644 --- a/toolchain/check/handle_class.cpp +++ b/toolchain/check/handle_class.cpp @@ -51,7 +51,7 @@ static auto BuildClassDeclaration(Context& context) // TODO: If this is an invalid redeclaration of a non-class entity or there // was an error in the qualifier, we will have lost track of the class name // here. We should keep track of it even if the name is invalid. - class_decl.class_id = context.semantics_ir().AddClass( + class_decl.class_id = context.semantics_ir().classes().Add( {.name_id = name_context.state == DeclarationNameStack::NameContext::State::Unresolved ? name_context.unresolved_name_id @@ -73,7 +73,7 @@ auto HandleClassDeclaration(Context& context, Parse::Node /*parse_node*/) auto HandleClassDefinitionStart(Context& context, Parse::Node parse_node) -> bool { auto [class_id, class_decl_id] = BuildClassDeclaration(context); - auto& class_info = context.semantics_ir().GetClass(class_id); + auto& class_info = context.semantics_ir().classes().Get(class_id); // Track that this declaration is the definition. if (class_info.definition_id.is_valid()) { diff --git a/toolchain/check/handle_function.cpp b/toolchain/check/handle_function.cpp index 0ef92e12c5de..9cf3fae6362b 100644 --- a/toolchain/check/handle_function.cpp +++ b/toolchain/check/handle_function.cpp @@ -80,7 +80,7 @@ static auto BuildFunctionDeclaration(Context& context, bool is_definition) // IDs in the signature. if (is_definition) { auto& function_info = - context.semantics_ir().GetFunction(function_decl.function_id); + context.semantics_ir().functions().Get(function_decl.function_id); function_info.param_refs_id = param_refs_id; function_info.return_type_id = return_type_id; function_info.return_slot_id = return_slot_id; @@ -93,7 +93,7 @@ static auto BuildFunctionDeclaration(Context& context, bool is_definition) // Create a new function if this isn't a valid redeclaration. if (!function_decl.function_id.is_valid()) { - function_decl.function_id = context.semantics_ir().AddFunction( + function_decl.function_id = context.semantics_ir().functions().Add( {.name_id = name_context.state == DeclarationNameStack::NameContext::State::Unresolved ? name_context.unresolved_name_id @@ -138,7 +138,8 @@ auto HandleFunctionDefinition(Context& context, Parse::Node parse_node) // and otherwise add an implicit `return;`. if (context.is_current_position_reachable()) { if (context.semantics_ir() - .GetFunction(function_id) + .functions() + .Get(function_id) .return_type_id.is_valid()) { CARBON_DIAGNOSTIC( MissingReturnStatement, Error, @@ -160,7 +161,7 @@ auto HandleFunctionDefinitionStart(Context& context, Parse::Node parse_node) // Process the declaration portion of the function. auto [function_id, decl_id] = BuildFunctionDeclaration(context, /*is_definition=*/true); - auto& function = context.semantics_ir().GetFunction(function_id); + auto& function = context.semantics_ir().functions().Get(function_id); // Track that this declaration is the definition. if (function.definition_id.is_valid()) { diff --git a/toolchain/check/handle_name.cpp b/toolchain/check/handle_name.cpp index 9484c84b6ed7..48508ad6fb50 100644 --- a/toolchain/check/handle_name.cpp +++ b/toolchain/check/handle_name.cpp @@ -19,7 +19,8 @@ static auto GetAsNameScope(Context& context, SemIR::NodeId base_id) return base_as_namespace->name_scope_id; } if (auto base_as_class = base.TryAs()) { - auto& class_info = context.semantics_ir().GetClass(base_as_class->class_id); + auto& class_info = + context.semantics_ir().classes().Get(base_as_class->class_id); if (!class_info.scope_id.is_valid()) { CARBON_DIAGNOSTIC(QualifiedExpressionInIncompleteClassScope, Error, "Member access into incomplete class `{0}`.", diff --git a/toolchain/check/handle_statement.cpp b/toolchain/check/handle_statement.cpp index 26479ca067eb..195d456397e7 100644 --- a/toolchain/check/handle_statement.cpp +++ b/toolchain/check/handle_statement.cpp @@ -32,7 +32,7 @@ auto HandleReturnStatement(Context& context, Parse::Node parse_node) -> bool { auto fn_node = context.semantics_ir().GetNodeAs( context.return_scope_stack().back()); const auto& callable = - context.semantics_ir().GetFunction(fn_node.function_id); + context.semantics_ir().functions().Get(fn_node.function_id); if (context.parse_tree().node_kind(context.node_stack().PeekParseNode()) == Parse::NodeKind::ReturnStatementStart) { diff --git a/toolchain/lower/file_context.cpp b/toolchain/lower/file_context.cpp index cb09bb536753..590739d94621 100644 --- a/toolchain/lower/file_context.cpp +++ b/toolchain/lower/file_context.cpp @@ -32,22 +32,22 @@ auto FileContext::Run() -> std::unique_ptr { CARBON_CHECK(llvm_module_) << "Run can only be called once."; // Lower types. - auto types = semantics_ir_->types(); + auto types = semantics_ir_->types().array_ref(); types_.resize_for_overwrite(types.size()); for (auto [i, type] : llvm::enumerate(types)) { types_[i] = BuildType(type.node_id); } // Lower function declarations. - functions_.resize_for_overwrite(semantics_ir_->functions_size()); - for (auto i : llvm::seq(semantics_ir_->functions_size())) { + functions_.resize_for_overwrite(semantics_ir_->functions().size()); + for (auto i : llvm::seq(semantics_ir_->functions().size())) { functions_[i] = BuildFunctionDeclaration(SemIR::FunctionId(i)); } // TODO: Lower global variable declarations. // Lower function definitions. - for (auto i : llvm::seq(semantics_ir_->functions_size())) { + for (auto i : llvm::seq(semantics_ir_->functions().size())) { BuildFunctionDefinition(SemIR::FunctionId(i)); } @@ -58,7 +58,7 @@ auto FileContext::Run() -> std::unique_ptr { auto FileContext::BuildFunctionDeclaration(SemIR::FunctionId function_id) -> llvm::Function* { - const auto& function = semantics_ir().GetFunction(function_id); + const auto& function = semantics_ir().functions().Get(function_id); const bool has_return_slot = function.return_slot_id.is_valid(); auto param_refs = semantics_ir().GetNodeBlock(function.param_refs_id); @@ -141,7 +141,7 @@ auto FileContext::BuildFunctionDeclaration(SemIR::FunctionId function_id) auto FileContext::BuildFunctionDefinition(SemIR::FunctionId function_id) -> void { - const auto& function = semantics_ir().GetFunction(function_id); + const auto& function = semantics_ir().functions().Get(function_id); const auto& body_block_ids = function.body_block_ids; if (body_block_ids.empty()) { // Function is probably defined in another file; not an error. diff --git a/toolchain/lower/handle.cpp b/toolchain/lower/handle.cpp index 23f21b4e8208..62c2553f6bea 100644 --- a/toolchain/lower/handle.cpp +++ b/toolchain/lower/handle.cpp @@ -344,7 +344,7 @@ auto HandleStructAccess(FunctionContext& context, SemIR::NodeId node_id, auto fields = context.semantics_ir().GetNodeBlock( context.semantics_ir() .GetNodeAs( - context.semantics_ir().GetType(struct_type_id)) + context.semantics_ir().types().Get(struct_type_id).node_id) .fields_id); auto field = context.semantics_ir().GetNodeAs( fields[node.index.index]); diff --git a/toolchain/sem_ir/entry_point.cpp b/toolchain/sem_ir/entry_point.cpp index e2833669d3d6..b8aa92e75298 100644 --- a/toolchain/sem_ir/entry_point.cpp +++ b/toolchain/sem_ir/entry_point.cpp @@ -13,7 +13,7 @@ static constexpr llvm::StringLiteral EntryPointFunction = "Run"; auto IsEntryPoint(const SemIR::File& file, SemIR::FunctionId function_id) -> bool { // TODO: Check if `file` is in the `Main` package. - const auto& function = file.GetFunction(function_id); + const auto& function = file.functions().Get(function_id); // TODO: Check if `function` is in a namespace. return function.name_id.is_valid() && file.strings().Get(function.name_id) == EntryPointFunction; diff --git a/toolchain/sem_ir/file.cpp b/toolchain/sem_ir/file.cpp index 370128d10802..2d1c5be5e5c1 100644 --- a/toolchain/sem_ir/file.cpp +++ b/toolchain/sem_ir/file.cpp @@ -94,7 +94,7 @@ auto File::Verify() const -> ErrorOr { // Check that every code block has a terminator sequence that appears at the // end of the block. - for (const Function& function : functions_) { + for (const Function& function : functions_.array_ref()) { for (NodeBlockId block_id : function.body_block_ids) { TerminatorKind prior_kind = TerminatorKind::NotTerminator; for (NodeId node_id : GetNodeBlock(block_id)) { @@ -125,7 +125,7 @@ auto File::Verify() const -> ErrorOr { static constexpr int BaseIndent = 4; static constexpr int IndentStep = 2; -// Define PrintList for ArrayRef. +// Prints a list of elements. template > static auto PrintList( @@ -142,16 +142,6 @@ static auto PrintList( out << "]\n"; } -// Adapt PrintList for a vector. -template > -static auto PrintList( - llvm::raw_ostream& out, llvm::StringLiteral name, - const llvm::SmallVector& list, - PrintT print = [](llvm::raw_ostream& out, const T& val) { out << val; }) { - PrintList(out, name, llvm::ArrayRef(list), print); -} - // PrintBlock is only used for vectors. template static auto PrintBlock(llvm::raw_ostream& out, llvm::StringLiteral block_name, @@ -179,9 +169,9 @@ auto File::Print(llvm::raw_ostream& out, bool include_builtins) const -> void { << " - cross_reference_irs_size: " << cross_reference_irs_.size() << "\n"; - PrintList(out, "functions", functions_); - PrintList(out, "classes", classes_); - PrintList(out, "types", types_); + PrintList(out, "functions", functions_.array_ref()); + PrintList(out, "classes", classes_.array_ref()); + PrintList(out, "types", types_.array_ref()); PrintBlock(out, "type_blocks", type_blocks_); llvm::ArrayRef nodes = nodes_; @@ -316,7 +306,7 @@ auto File::StringifyTypeExpression(NodeId outer_node_id, } case ClassDeclaration::Kind: { auto class_name_id = - GetClass(node.As().class_id).name_id; + classes().Get(node.As().class_id).name_id; out << strings().Get(class_name_id); break; } diff --git a/toolchain/sem_ir/file.h b/toolchain/sem_ir/file.h index 347996538eb3..41104fe2ea04 100644 --- a/toolchain/sem_ir/file.h +++ b/toolchain/sem_ir/file.h @@ -153,42 +153,6 @@ class File : public Printable { return *cross_reference_irs_[xref_id.index]; } - // Adds a callable, returning an ID to reference it. - auto AddFunction(Function function) -> FunctionId { - FunctionId id(functions_.size()); - // TODO: Return failure on overflow instead of crashing. - CARBON_CHECK(id.index >= 0); - functions_.push_back(function); - return id; - } - - // Returns the requested callable. - auto GetFunction(FunctionId function_id) const -> const Function& { - return functions_[function_id.index]; - } - - // Returns the requested callable. - auto GetFunction(FunctionId function_id) -> Function& { - return functions_[function_id.index]; - } - - // Adds a class, returning an ID to reference it. - auto AddClass(Class class_info) -> ClassId { - ClassId id(classes_.size()); - // TODO: Return failure on overflow instead of crashing. - CARBON_CHECK(id.index >= 0); - classes_.push_back(class_info); - return id; - } - - // Returns the requested class. - auto GetClass(ClassId class_id) const -> const Class& { - return classes_[class_id.index]; - } - - // Returns the requested class. - auto GetClass(ClassId class_id) -> Class& { return classes_[class_id.index]; } - // Adds a name scope, returning an ID to reference it. auto AddNameScope() -> NameScopeId { NameScopeId name_scopes_id(name_scopes_.size()); @@ -287,15 +251,6 @@ class File : public Printable { return node_blocks_[block_id.index]; } - // Adds a type, returning an ID to reference it. - auto AddType(NodeId node_id) -> TypeId { - TypeId type_id(types_.size()); - // Should never happen, will always overflow node_ids first. - CARBON_DCHECK(type_id.index >= 0); - types_.push_back({.node_id = node_id}); - return type_id; - } - // Marks a type as complete, and sets its value representation. auto CompleteType(TypeId object_type_id, ValueRepresentation value_representation) -> void { @@ -303,20 +258,10 @@ class File : public Printable { // We already know our builtin types are complete. return; } - CARBON_CHECK(types_[object_type_id.index].value_representation.kind == + CARBON_CHECK(types().Get(object_type_id).value_representation.kind == ValueRepresentation::Unknown) << "Type " << object_type_id << " completed more than once"; - types_[object_type_id.index].value_representation = value_representation; - } - - // Gets the node ID for a type. This doesn't handle TypeType or InvalidType in - // order to avoid a check; callers that need that should use - // GetTypeAllowBuiltinTypes. - auto GetType(TypeId type_id) const -> NodeId { - // Double-check it's not called with TypeType or InvalidType. - CARBON_CHECK(type_id.index >= 0) - << "Invalid argument for GetType: " << type_id; - return types_[type_id.index].node_id; + types().Get(object_type_id).value_representation = value_representation; } auto GetTypeAllowBuiltinTypes(TypeId type_id) const -> NodeId { @@ -327,7 +272,7 @@ class File : public Printable { } else if (type_id == TypeId::Invalid) { return NodeId::Invalid; } else { - return GetType(type_id); + return types().Get(type_id).node_id; } } @@ -338,7 +283,7 @@ class File : public Printable { // TypeType and InvalidType are their own value representation. return {.kind = ValueRepresentation::Copy, .type_id = type_id}; } - return types_[type_id.index].value_representation; + return types().Get(type_id).value_representation; } // Determines whether the given type is known to be complete. This does not @@ -349,7 +294,7 @@ class File : public Printable { // Gets the pointee type of the given type, which must be a pointer type. auto GetPointeeType(TypeId pointer_id) const -> TypeId { - return GetNodeAs(GetType(pointer_id)).pointee_id; + return GetNodeAs(types().Get(pointer_id).node_id).pointee_id; } // Adds a type block with the given content, returning an ID to reference it. @@ -399,13 +344,18 @@ class File : public Printable { return value_stores_->strings(); } - auto functions_size() const -> int { return functions_.size(); } - auto classes_size() const -> int { return classes_.size(); } + auto functions() -> ValueStore& { return functions_; } + auto functions() const -> const ValueStore& { + return functions_; + } + auto classes() -> ValueStore& { return classes_; } + auto classes() const -> const ValueStore& { return classes_; } + auto types() -> ValueStore& { return types_; } + auto types() const -> const ValueStore& { return types_; } + auto nodes_size() const -> int { return nodes_.size(); } auto node_blocks_size() const -> int { return node_blocks_.size(); } - auto types() const -> llvm::ArrayRef { return types_; } - auto top_node_block_id() const -> NodeBlockId { return top_node_block_id_; } auto set_top_node_block_id(NodeBlockId block_id) -> void { top_node_block_id_ = block_id; @@ -450,10 +400,10 @@ class File : public Printable { std::string filename_; // Storage for callable objects. - llvm::SmallVector functions_; + ValueStore functions_; // Storage for classes. - llvm::SmallVector classes_; + ValueStore classes_; // Related IRs. There will always be at least 2 entries, the builtin IR (used // for references of builtins) followed by the current IR (used for references @@ -464,7 +414,7 @@ class File : public Printable { llvm::SmallVector> name_scopes_; // Descriptions of types used in this file. - llvm::SmallVector types_; + ValueStore types_; // Type blocks within the IR. These reference entries in types_. Storage for // the data is provided by allocator_. diff --git a/toolchain/sem_ir/formatter.cpp b/toolchain/sem_ir/formatter.cpp index e1745cf043f0..d49f44b59b64 100644 --- a/toolchain/sem_ir/formatter.cpp +++ b/toolchain/sem_ir/formatter.cpp @@ -38,8 +38,8 @@ class NodeNamer { semantics_ir_(semantics_ir) { nodes.resize(semantics_ir.nodes_size()); labels.resize(semantics_ir.node_blocks_size()); - scopes.resize(1 + semantics_ir.functions_size() + - semantics_ir.classes_size()); + scopes.resize(1 + semantics_ir.functions().size() + + semantics_ir.classes().size()); // Build the package scope. GetScopeInfo(ScopeIndex::Package).name = @@ -47,10 +47,9 @@ class NodeNamer { CollectNamesInBlock(ScopeIndex::Package, semantics_ir.top_node_block_id()); // Build each function scope. - for (int i : llvm::seq(semantics_ir.functions_size())) { + for (auto [i, fn] : llvm::enumerate(semantics_ir.functions().array_ref())) { auto fn_id = FunctionId(i); auto fn_scope = GetScopeFor(fn_id); - const auto& fn = semantics_ir.GetFunction(fn_id); // TODO: Provide a location for the function for use as a // disambiguator. auto fn_loc = Parse::Node::Invalid; @@ -78,10 +77,10 @@ class NodeNamer { } // Build each class scope. - for (int i : llvm::seq(semantics_ir.classes_size())) { + for (auto [i, class_info] : + llvm::enumerate(semantics_ir.classes().array_ref())) { auto class_id = ClassId(i); auto class_scope = GetScopeFor(class_id); - const auto& class_info = semantics_ir.GetClass(class_id); // TODO: Provide a location for the class for use as a // disambiguator. auto class_loc = Parse::Node::Invalid; @@ -102,7 +101,7 @@ class NodeNamer { // Returns the scope index corresponding to a class. auto GetScopeFor(ClassId class_id) -> ScopeIndex { - return static_cast(1 + semantics_ir_.functions_size() + + return static_cast(1 + semantics_ir_.functions().size() + class_id.index); } @@ -420,16 +419,15 @@ class NodeNamer { continue; } case FunctionDeclaration::Kind: { - add_node_name_id( - semantics_ir_ - .GetFunction(node.As().function_id) - .name_id); + add_node_name_id(semantics_ir_.functions() + .Get(node.As().function_id) + .name_id); continue; } case ClassDeclaration::Kind: { - add_node_name_id( - semantics_ir_.GetClass(node.As().class_id) - .name_id); + add_node_name_id(semantics_ir_.classes() + .Get(node.As().class_id) + .name_id); continue; } case NameReference::Kind: { @@ -498,17 +496,17 @@ class Formatter { } out_ << "}\n"; - for (int i : llvm::seq(semantics_ir_.classes_size())) { + for (int i : llvm::seq(semantics_ir_.classes().size())) { FormatClass(ClassId(i)); } - for (int i : llvm::seq(semantics_ir_.functions_size())) { + for (int i : llvm::seq(semantics_ir_.functions().size())) { FormatFunction(FunctionId(i)); } } auto FormatClass(ClassId id) -> void { - const Class& class_info = semantics_ir_.GetClass(id); + const Class& class_info = semantics_ir_.classes().Get(id); out_ << "\nclass "; FormatClassName(id); @@ -527,7 +525,7 @@ class Formatter { } auto FormatFunction(FunctionId id) -> void { - const Function& fn = semantics_ir_.GetFunction(id); + const Function& fn = semantics_ir_.functions().Get(id); out_ << "\nfn "; FormatFunctionName(id);