diff --git a/toolchain/check/check_unit.cpp b/toolchain/check/check_unit.cpp index d1cdcbcafbd9..187567fb0b6b 100644 --- a/toolchain/check/check_unit.cpp +++ b/toolchain/check/check_unit.cpp @@ -378,10 +378,11 @@ auto CheckUnit::ProcessNodeIds() -> bool { bool result; auto parse_kind = context_.parse_tree().node_kind(node_id); switch (parse_kind) { -#define CARBON_PARSE_NODE_KIND(Name) \ - case Parse::NodeKind::Name: { \ - result = HandleParseNode(context_, Parse::Name##Id(node_id)); \ - break; \ +#define CARBON_PARSE_NODE_KIND(Name) \ + case Parse::NodeKind::Name: { \ + result = HandleParseNode( \ + context_, context_.parse_tree().As(node_id)); \ + break; \ } #include "toolchain/parse/node_kind.def" } diff --git a/toolchain/check/import.cpp b/toolchain/check/import.cpp index e0cd3d0612ef..af20bb8a7627 100644 --- a/toolchain/check/import.cpp +++ b/toolchain/check/import.cpp @@ -133,8 +133,9 @@ auto AddImportNamespace(Context& context, SemIR::TypeId namespace_type_id, ? MakeImportedLocIdAndInst(context, import_loc_id.import_ir_inst_id(), namespace_inst) // TODO: Check that this actually is an `AnyNamespaceId`. - : SemIR::LocIdAndInst(Parse::AnyNamespaceId(import_loc_id.node_id()), - namespace_inst); + : SemIR::LocIdAndInst( + Parse::AnyNamespaceId::UnsafeMake(import_loc_id.node_id()), + namespace_inst); auto namespace_id = AddPlaceholderInstInNoBlock(context, namespace_inst_and_loc); context.import_ref_ids().push_back(namespace_id); diff --git a/toolchain/check/node_stack.h b/toolchain/check/node_stack.h index 025b5cc4623d..d4666aa889d7 100644 --- a/toolchain/check/node_stack.h +++ b/toolchain/check/node_stack.h @@ -140,8 +140,8 @@ class NodeStack { auto PopForSoloNodeId() -> Parse::NodeIdForKind { Entry back = PopEntry(); RequireIdKind(RequiredParseKind, Id::Kind::None); - RequireParseKind(back.node_id); - return Parse::NodeIdForKind(back.node_id); + return parse_tree_->As>( + back.node_id); } // Pops the top of the stack if it is the given kind, and returns the @@ -192,7 +192,7 @@ class NodeStack { template auto PopWithNodeId() -> auto { auto id = Peek(); - Parse::NodeIdForKind node_id( + auto node_id = parse_tree_->As>( stack_.pop_back_val().node_id); return std::make_pair(node_id, id); } @@ -201,8 +201,9 @@ class NodeStack { template auto PopWithNodeId() -> auto { auto id = Peek(); - Parse::NodeIdInCategory node_id( - stack_.pop_back_val().node_id); + auto node_id = + parse_tree_->As>( + stack_.pop_back_val().node_id); return std::make_pair(node_id, id); } @@ -302,7 +303,9 @@ class NodeStack { template auto Peek() const -> auto { Entry back = stack_.back(); - RequireParseKind(back.node_id); + CARBON_CHECK(RequiredParseKind == parse_tree_->node_kind(back.node_id), + "Expected {0}, found {1}", RequiredParseKind, + parse_tree_->node_kind(back.node_id)); constexpr Id::Kind RequiredIdKind = NodeKindToIdKind(RequiredParseKind); return Peek(); } @@ -589,14 +592,6 @@ class NodeStack { SemIR::IdKind(NodeKindToIdKind(parse_kind))); } - // Require an entry to have the given Parse::NodeKind. - template - auto RequireParseKind(Parse::NodeId node_id) const -> void { - auto actual_kind = parse_tree_->node_kind(node_id); - CARBON_CHECK(RequiredParseKind == actual_kind, "Expected {0}, found {1}", - RequiredParseKind, actual_kind); - } - // Require an entry to have the given Parse::NodeCategory. template auto RequireParseCategory(Parse::NodeId node_id) const -> void { diff --git a/toolchain/parse/context.cpp b/toolchain/parse/context.cpp index 55ffda244524..9c39257df92f 100644 --- a/toolchain/parse/context.cpp +++ b/toolchain/parse/context.cpp @@ -449,8 +449,8 @@ auto Context::AddFunctionDefinitionStart(Lex::TokenIndex token, bool has_error) -> void { if (ParsingInDeferredDefinitionScope(*this)) { deferred_definition_stack_.push_back(tree_->deferred_definitions_.Add( - {.start_id = - FunctionDefinitionStartId(NodeId(tree_->node_impls_.size()))})); + {.start_id = FunctionDefinitionStartId::UnsafeMake( + NodeId(tree_->node_impls_.size()))})); } AddNode(NodeKind::FunctionDefinitionStart, token, has_error); @@ -462,7 +462,7 @@ auto Context::AddFunctionDefinition(Lex::TokenIndex token, bool has_error) auto definition_index = deferred_definition_stack_.pop_back_val(); auto& definition = tree_->deferred_definitions_.Get(definition_index); definition.definition_id = - FunctionDefinitionId(NodeId(tree_->node_impls_.size())); + FunctionDefinitionId::UnsafeMake(NodeId(tree_->node_impls_.size())); definition.next_definition_index = DeferredDefinitionIndex(tree_->deferred_definitions().size()); } diff --git a/toolchain/parse/extract.cpp b/toolchain/parse/extract.cpp index adebeb9f3019..bface9a8a8c4 100644 --- a/toolchain/parse/extract.cpp +++ b/toolchain/parse/extract.cpp @@ -155,7 +155,7 @@ struct Extractable> { static auto Extract(NodeExtractor& extractor) -> std::optional> { if (extractor.MatchesNodeIdForKind(Kind)) { - return NodeIdForKind(extractor.ExtractNode()); + return NodeIdForKind::UnsafeMake(extractor.ExtractNode()); } else { return std::nullopt; } @@ -182,7 +182,7 @@ struct Extractable> { static auto Extract(NodeExtractor& extractor) -> std::optional> { if (extractor.MatchesNodeIdInCategory(Category)) { - return NodeIdInCategory(extractor.ExtractNode()); + return NodeIdInCategory::UnsafeMake(extractor.ExtractNode()); } else { return std::nullopt; } @@ -227,7 +227,7 @@ struct Extractable> { static auto Extract(NodeExtractor& extractor) -> std::optional> { if (extractor.MatchesNodeIdOneOf({T::Kind...})) { - return NodeIdOneOf(extractor.ExtractNode()); + return NodeIdOneOf::UnsafeMake(extractor.ExtractNode()); } else { return std::nullopt; } diff --git a/toolchain/parse/handle_import_and_package.cpp b/toolchain/parse/handle_import_and_package.cpp index 108ab422777c..108579f2c81f 100644 --- a/toolchain/parse/handle_import_and_package.cpp +++ b/toolchain/parse/handle_import_and_package.cpp @@ -40,7 +40,7 @@ static auto HandleDeclContent(Context& context, Context::StateStackEntry state, llvm::function_refvoid> on_parse_error) -> void { Tree::PackagingNames names{ - .node_id = ImportDeclId(NodeId(state.subtree_start)), + .node_id = ImportDeclId::UnsafeMake(NodeId(state.subtree_start)), .is_export = is_export}; // Parse the package name. diff --git a/toolchain/parse/node_ids.h b/toolchain/parse/node_ids.h index b66349ae3ef5..a7c695dc0c02 100644 --- a/toolchain/parse/node_ids.h +++ b/toolchain/parse/node_ids.h @@ -40,9 +40,20 @@ template struct NodeIdForKind : public NodeId { // NOLINTNEXTLINE(readability-identifier-naming) static const NodeKind& Kind; - constexpr explicit NodeIdForKind(NodeId node_id) : NodeId(node_id) {} + + // Provide a factory function for construction from `NodeId`. This doesn't + // validate the type, so it's unsafe. + static constexpr auto UnsafeMake(NodeId node_id) -> NodeIdForKind { + return NodeIdForKind(node_id); + } + // NOLINTNEXTLINE(google-explicit-constructor) constexpr NodeIdForKind(NoneNodeId /*none*/) : NodeId(NoneIndex) {} + + private: + // Private to prevent accidental explicit construction from an untyped + // NodeId. + explicit constexpr NodeIdForKind(NodeId node_id) : NodeId(node_id) {} }; template const NodeKind& NodeIdForKind::Kind = K; @@ -54,6 +65,12 @@ const NodeKind& NodeIdForKind::Kind = K; // NodeId that matches any NodeKind whose `category()` overlaps with `Category`. template struct NodeIdInCategory : public NodeId { + // Provide a factory function for construction from `NodeId`. This doesn't + // validate the type, so it's unsafe. + static constexpr auto UnsafeMake(NodeId node_id) -> NodeIdInCategory { + return NodeIdInCategory(node_id); + } + // Support conversion from `NodeIdForKind` if Kind's category // overlaps with `Category`. template @@ -62,9 +79,13 @@ struct NodeIdInCategory : public NodeId { CARBON_CHECK(Kind.category().HasAnyOf(Category)); } - constexpr explicit NodeIdInCategory(NodeId node_id) : NodeId(node_id) {} // NOLINTNEXTLINE(google-explicit-constructor) constexpr NodeIdInCategory(NoneNodeId /*none*/) : NodeId(NoneIndex) {} + + private: + // Private to prevent accidental explicit construction from an untyped + // NodeId. + explicit constexpr NodeIdInCategory(NodeId node_id) : NodeId(node_id) {} }; // Aliases for `NodeIdInCategory` to describe particular categories of nodes. @@ -84,16 +105,26 @@ using AnyPackageNameId = NodeIdInCategory; // NodeId with kind that matches one of the `T::Kind`s. template + requires(sizeof...(T) >= 2) struct NodeIdOneOf : public NodeId { - static_assert(sizeof...(T) >= 2, "Expected at least two types."); - constexpr explicit NodeIdOneOf(NodeId node_id) : NodeId(node_id) {} - template - // NOLINTNEXTLINE(google-explicit-constructor) - NodeIdOneOf(NodeIdForKind node_id) : NodeId(node_id) { - static_assert(((T::Kind == Kind) || ...)); + // Provide a factory function for construction from `NodeId`. This doesn't + // validate the type, so it's unsafe. + static constexpr auto UnsafeMake(NodeId node_id) -> NodeIdOneOf { + return NodeIdOneOf(node_id); } + + template + requires((T::Kind == Kind) || ...) + // NOLINTNEXTLINE(google-explicit-constructor) + NodeIdOneOf(NodeIdForKind node_id) : NodeId(node_id) {} + // NOLINTNEXTLINE(google-explicit-constructor) constexpr NodeIdOneOf(NoneNodeId /*none*/) : NodeId(NoneIndex) {} + + private: + // Private to prevent accidental explicit construction from an untyped + // NodeId. + explicit constexpr NodeIdOneOf(NodeId node_id) : NodeId(node_id) {} }; using AnyClassDeclId = diff --git a/toolchain/parse/tree.h b/toolchain/parse/tree.h index 6221561686fb..4fcf487a7c14 100644 --- a/toolchain/parse/tree.h +++ b/toolchain/parse/tree.h @@ -146,7 +146,7 @@ class Tree : public Printable { auto TryAs(NodeId n) const -> std::optional { CARBON_DCHECK(n.has_value()); if (ConvertTo::AllowedFor(node_kind(n))) { - return T(n); + return T::UnsafeMake(n); } else { return std::nullopt; } @@ -157,8 +157,8 @@ class Tree : public Printable { template auto As(NodeId n) const -> T { CARBON_DCHECK(n.has_value()); - CARBON_CHECK(ConvertTo::AllowedFor(node_kind(n))); - return T(n); + CARBON_DCHECK(ConvertTo::AllowedFor(node_kind(n))); + return T::UnsafeMake(n); } auto packaging_decl() const -> const std::optional& {