Refactor SemanticsNode factory functions into factory templates. (#2711)

I'm trying to reduce the amount of per-SemanticsNodeKind boilerplate, and make mistakes (e.g., SemanticsNodeKind not matching the Make name, misplacing the type, or Get/Make type mismatches) easier to see.

I could've done this with (more) macros, but felt that the template approach was reasonable enough and likely easier to understand/debug. I'm not sure whether there's more that I could be doing with variadics to reduce the amount of factory code, but this feels good right now.
This commit is contained in:
Jon Ross-Perkins
2023-03-27 11:07:49 -07:00
committed by GitHub
parent 4fa11af13f
commit 7fc203c536
7 changed files with 201 additions and 185 deletions
+175 -163
View File
@@ -124,207 +124,219 @@ struct SemanticsStringId : public IndexBase {
}
};
// The standard structure for nodes.
// The standard structure for SemanticsNode. This is trying to provide a minimal
// amount of information for a node:
//
// - parse_node for error placement.
// - kind for run-time logic when the input Kind is unknown.
// - type_id for quick type checking.
// - Up to two Kind-specific members.
//
// For each Kind in SemanticsNodeKind, a typical flow looks like:
//
// - Create a `SemanticsNode` using `SemanticsNode::Kind::Make()`
// - Access cross-Kind members using `node.type_id()` and similar.
// - Access Kind-specific members using `node.GetAsKind()`, which depending on
// the number of members will return one of NoArgs, a single value, or a
// `std::pair` of values.
// - Using the wrong `node.GetAsKind()` is a programming error, and should
// CHECK-fail in debug modes (opt may too, but it's not an API guarantee).
//
// Internally, each Kind uses the `Factory*` types to provide a boilerplate
// `Make` and `Get` methods.
class SemanticsNode {
public:
struct NoArgs {};
auto GetAsInvalid() const -> NoArgs { CARBON_FATAL() << "Invalid access"; }
// Factory base classes are private, then used for public classes. This class
// has two public and two private sections to prevent accidents.
private:
// Factory templates need to use the raw enum instead of the class wrapper.
using KindTemplateEnum = Internal::SemanticsNodeKindRawEnum;
static auto MakeAssign(ParseTree::Node parse_node, SemanticsNodeId type,
SemanticsNodeId lhs, SemanticsNodeId rhs)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::Assign, type, lhs.index,
rhs.index);
}
auto GetAsAssign() const -> std::pair<SemanticsNodeId, SemanticsNodeId> {
CARBON_CHECK(kind_ == SemanticsNodeKind::Assign);
return {SemanticsNodeId(arg0_), SemanticsNodeId(arg1_)};
}
// Provides Make and Get to support 0, 1, or 2 arguments for a SemanticsNode.
// These are protected so that child factories can opt in to what pieces they
// want to use.
template <KindTemplateEnum Kind, typename... ArgTypes>
class FactoryBase {
protected:
static auto Make(ParseTree::Node parse_node, SemanticsNodeId type_id,
ArgTypes... arg_ids) -> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::Create(Kind), type_id,
arg_ids.index...);
}
static auto MakeBinaryOperatorAdd(ParseTree::Node parse_node,
SemanticsNodeId type, SemanticsNodeId lhs,
SemanticsNodeId rhs) -> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::BinaryOperatorAdd, type,
lhs.index, rhs.index);
}
auto GetAsBinaryOperatorAdd() const
-> std::pair<SemanticsNodeId, SemanticsNodeId> {
CARBON_CHECK(kind_ == SemanticsNodeKind::BinaryOperatorAdd);
return {SemanticsNodeId(arg0_), SemanticsNodeId(arg1_)};
}
static auto Get(SemanticsNode node) {
struct Unused {};
return GetImpl<ArgTypes..., Unused>(node);
}
static auto MakeBindName(ParseTree::Node parse_node, SemanticsNodeId type,
SemanticsStringId name, SemanticsNodeId node)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::BindName, type,
name.index, node.index);
}
auto GetAsBindName() const -> std::pair<SemanticsStringId, SemanticsNodeId> {
CARBON_CHECK(kind_ == SemanticsNodeKind::BindName);
return {SemanticsStringId(arg0_), SemanticsNodeId(arg1_)};
}
private:
// GetImpl handles the different return types based on ArgTypes.
template <typename Arg0Type, typename Arg1Type, typename>
static auto GetImpl(SemanticsNode node) -> std::pair<Arg0Type, Arg1Type> {
CARBON_CHECK(node.kind() == Kind);
return {Arg0Type(node.arg0_), Arg1Type(node.arg1_)};
}
template <typename Arg0Type, typename>
static auto GetImpl(SemanticsNode node) -> Arg0Type {
CARBON_CHECK(node.kind() == Kind);
return Arg0Type(node.arg0_);
}
template <typename>
static auto GetImpl(SemanticsNode node) -> NoArgs {
CARBON_CHECK(node.kind() == Kind);
return NoArgs();
}
};
static auto MakeBuiltin(SemanticsBuiltinKind builtin_kind,
SemanticsNodeId type) -> SemanticsNode {
// Builtins won't have a ParseTree node associated, so we provide the
// default invalid one.
return SemanticsNode(ParseTree::Node::Invalid, SemanticsNodeKind::Builtin,
type, builtin_kind.AsInt());
}
auto GetAsBuiltin() const -> SemanticsBuiltinKind {
CARBON_CHECK(kind_ == SemanticsNodeKind::Builtin);
return SemanticsBuiltinKind::FromInt(arg0_);
}
// Provide Get along with a Make that requires a type.
template <KindTemplateEnum Kind, typename... ArgTypes>
class Factory : public FactoryBase<Kind, ArgTypes...> {
public:
using FactoryBase<Kind, ArgTypes...>::Make;
using FactoryBase<Kind, ArgTypes...>::Get;
};
static auto MakeCall(ParseTree::Node parse_node, SemanticsNodeId type,
SemanticsCallId call_id, SemanticsCallableId callable_id)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::Call, type,
call_id.index, callable_id.index);
}
auto GetAsCall() const -> std::pair<SemanticsCallId, SemanticsCallableId> {
CARBON_CHECK(kind_ == SemanticsNodeKind::Call);
return {SemanticsCallId(arg0_), SemanticsCallableId(arg1_)};
}
// Provides Get along with a Make that assumes a non-changing type.
template <KindTemplateEnum Kind, int32_t TypeIndex, typename... ArgTypes>
class FactoryPreTyped : public FactoryBase<Kind, ArgTypes...> {
public:
static auto Make(ParseTree::Node parse_node, ArgTypes... args) {
SemanticsNodeId type_id(TypeIndex);
return FactoryBase<Kind, ArgTypes...>::Make(parse_node, type_id, args...);
}
using FactoryBase<Kind, ArgTypes...>::Get;
};
static auto MakeCodeBlock(ParseTree::Node parse_node,
SemanticsNodeBlockId node_block) -> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::CodeBlock,
SemanticsNodeId::Invalid, node_block.index);
}
auto GetAsCodeBlock() const -> SemanticsNodeBlockId {
CARBON_CHECK(kind_ == SemanticsNodeKind::CodeBlock);
return SemanticsNodeBlockId(arg0_);
}
public:
// Invalid is in the SemanticsNodeKind enum, but should never be used.
class Invalid {
public:
static auto Get(SemanticsNode /*node*/) -> SemanticsNode::NoArgs {
CARBON_FATAL() << "Invalid access";
}
};
static auto MakeCrossReference(SemanticsNodeId type,
SemanticsCrossReferenceIRId ir,
SemanticsNodeId node) -> SemanticsNode {
return SemanticsNode(ParseTree::Node::Invalid,
SemanticsNodeKind::CrossReference, type, ir.index,
node.index);
}
auto GetAsCrossReference() const
-> std::pair<SemanticsCrossReferenceIRId, SemanticsNodeId> {
CARBON_CHECK(kind_ == SemanticsNodeKind::CrossReference);
return {SemanticsCrossReferenceIRId(arg0_), SemanticsNodeId(arg1_)};
}
using Assign = SemanticsNode::Factory<SemanticsNodeKind::Assign,
SemanticsNodeId /*lhs_id*/,
SemanticsNodeId /*rhs_id*/>;
static auto MakeFunctionDeclaration(ParseTree::Node parse_node,
SemanticsStringId name_id,
SemanticsCallableId signature_id)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::FunctionDeclaration,
SemanticsNodeId::Invalid, name_id.index,
signature_id.index);
}
auto GetAsFunctionDeclaration() const
-> std::pair<SemanticsStringId, SemanticsCallableId> {
CARBON_CHECK(kind_ == SemanticsNodeKind::FunctionDeclaration);
return {SemanticsStringId(arg0_), SemanticsCallableId(arg1_)};
}
using BinaryOperatorAdd =
SemanticsNode::Factory<SemanticsNodeKind::BinaryOperatorAdd,
SemanticsNodeId /*lhs_id*/,
SemanticsNodeId /*rhs_id*/>;
static auto MakeFunctionDefinition(ParseTree::Node parse_node,
SemanticsNodeId decl,
SemanticsNodeBlockId node_block)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::FunctionDefinition,
SemanticsNodeId::Invalid, decl.index,
node_block.index);
}
auto GetAsFunctionDefinition() const
-> std::pair<SemanticsNodeId, SemanticsNodeBlockId> {
CARBON_CHECK(kind_ == SemanticsNodeKind::FunctionDefinition);
return {SemanticsNodeId(arg0_), SemanticsNodeBlockId(arg1_)};
}
using BindName = SemanticsNode::Factory<SemanticsNodeKind::BindName,
SemanticsStringId /*name_id*/,
SemanticsNodeId /*node_id*/>;
static auto MakeIntegerLiteral(ParseTree::Node parse_node,
SemanticsIntegerLiteralId integer)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::IntegerLiteral,
SemanticsNodeId::BuiltinIntegerType, integer.index);
}
auto GetAsIntegerLiteral() const -> SemanticsIntegerLiteralId {
CARBON_CHECK(kind_ == SemanticsNodeKind::IntegerLiteral);
return SemanticsIntegerLiteralId(arg0_);
}
class Builtin {
public:
static auto Make(SemanticsBuiltinKind builtin_kind, SemanticsNodeId type_id)
-> SemanticsNode {
// Builtins won't have a ParseTree node associated, so we provide the
// default invalid one.
// This can't use the standard Make function because of the `AsInt()` cast
// instead of `.index`.
return SemanticsNode(ParseTree::Node::Invalid, SemanticsNodeKind::Builtin,
type_id, builtin_kind.AsInt());
}
static auto Get(SemanticsNode node) -> SemanticsBuiltinKind {
return SemanticsBuiltinKind::FromInt(node.arg0_);
}
};
static auto MakeRealLiteral(ParseTree::Node parse_node,
SemanticsRealLiteralId real) -> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::RealLiteral,
SemanticsNodeId::BuiltinFloatingPointType, real.index);
}
auto GetAsRealLiteral() const -> SemanticsRealLiteralId {
CARBON_CHECK(kind_ == SemanticsNodeKind::RealLiteral);
return SemanticsRealLiteralId(arg0_);
}
using Call = Factory<SemanticsNodeKind::Call, SemanticsCallId /*call_id*/,
SemanticsCallableId /*callable_id*/>;
static auto MakeReturn(ParseTree::Node parse_node) -> SemanticsNode {
// The actual type is `()`. However, code dealing with `return;` should
// understand the type without checking, so it's not necessary but could be
// specified if needed.
return SemanticsNode(parse_node, SemanticsNodeKind::Return,
SemanticsNodeId::Invalid);
}
auto GetAsReturn() const -> NoArgs {
CARBON_CHECK(kind_ == SemanticsNodeKind::Return);
return {};
}
using CodeBlock = FactoryPreTyped<SemanticsNodeKind::CodeBlock,
SemanticsNodeId::InvalidIndex,
SemanticsNodeBlockId /*node_block_id*/>;
static auto MakeReturnExpression(ParseTree::Node parse_node,
SemanticsNodeId type, SemanticsNodeId expr)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::ReturnExpression, type,
expr.index);
}
auto GetAsReturnExpression() const -> SemanticsNodeId {
CARBON_CHECK(kind_ == SemanticsNodeKind::ReturnExpression);
return SemanticsNodeId(arg0_);
}
class CrossReference
: public FactoryBase<SemanticsNodeKind::CrossReference,
SemanticsCrossReferenceIRId /*ir_id*/,
SemanticsNodeId /*node_id*/> {
public:
static auto Make(SemanticsNodeId type_id, SemanticsCrossReferenceIRId ir_id,
SemanticsNodeId node_id) -> SemanticsNode {
// A node's parse tree node must refer to a node in the current parse
// tree. This cannot use the cross-referenced node's parse tree node
// because it will be in a different parse tree.
return FactoryBase::Make(ParseTree::Node::Invalid, type_id, ir_id,
node_id);
}
using FactoryBase::Get;
};
static auto MakeStringLiteral(ParseTree::Node parse_node,
SemanticsStringId string_id) -> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::StringLiteral,
SemanticsNodeId::BuiltinStringType, string_id.index);
}
auto GetAsStringLiteral() const -> SemanticsStringId {
CARBON_CHECK(kind_ == SemanticsNodeKind::StringLiteral);
return SemanticsStringId(arg0_);
}
using FunctionDeclaration = FactoryPreTyped<
SemanticsNodeKind::FunctionDeclaration, SemanticsNodeId::InvalidIndex,
SemanticsStringId /*name_id*/, SemanticsCallableId /*signature_id*/>;
static auto MakeVarStorage(ParseTree::Node parse_node, SemanticsNodeId type)
-> SemanticsNode {
return SemanticsNode(parse_node, SemanticsNodeKind::VarStorage, type);
}
auto GetAsVarStorage() const -> NoArgs {
CARBON_CHECK(kind_ == SemanticsNodeKind::VarStorage);
return NoArgs();
}
using FunctionDefinition = FactoryPreTyped<
SemanticsNodeKind::FunctionDefinition, SemanticsNodeId::InvalidIndex,
SemanticsNodeId /*decl_id*/, SemanticsNodeBlockId /*node_block_id*/>;
using IntegerLiteral =
FactoryPreTyped<SemanticsNodeKind::IntegerLiteral,
SemanticsBuiltinKind::IntegerType.AsInt(),
SemanticsIntegerLiteralId /*integer_id*/>;
using RealLiteral =
FactoryPreTyped<SemanticsNodeKind::RealLiteral,
SemanticsBuiltinKind::FloatingPointType.AsInt(),
SemanticsRealLiteralId /*real_id*/>;
using Return =
FactoryPreTyped<SemanticsNodeKind::Return, SemanticsNodeId::InvalidIndex>;
using ReturnExpression =
Factory<SemanticsNodeKind::ReturnExpression, SemanticsNodeId /*expr_id*/>;
using StringLiteral =
FactoryPreTyped<SemanticsNodeKind::StringLiteral,
SemanticsBuiltinKind::StringType.AsInt(),
SemanticsStringId /*string_id*/>;
using VarStorage = Factory<SemanticsNodeKind::VarStorage>;
SemanticsNode()
: SemanticsNode(ParseTree::Node::Invalid, SemanticsNodeKind::Invalid,
SemanticsNodeId::Invalid) {}
// Provide `node.GetAsKind()` as an instance method for all kinds, essentially
// an alias for`SemanticsNode::Kind::Get(node)`.
#define CARBON_SEMANTICS_NODE_KIND(Name) \
auto GetAs##Name() const { return Name::Get(*this); }
#include "toolchain/semantics/semantics_node_kind.def"
auto parse_node() const -> ParseTree::Node { return parse_node_; }
auto kind() const -> SemanticsNodeKind { return kind_; }
auto type() const -> SemanticsNodeId { return type_; }
auto type_id() const -> SemanticsNodeId { return type_id_; }
auto Print(llvm::raw_ostream& out) const -> void;
private:
// Builtins have peculiar construction, so they are a friend rather than using
// a factory base class.
friend struct SemanticsNodeForBuiltin;
explicit SemanticsNode(ParseTree::Node parse_node, SemanticsNodeKind kind,
SemanticsNodeId type, int32_t arg0 = -1,
int32_t arg1 = -1)
SemanticsNodeId type_id,
int32_t arg0 = SemanticsNodeId::InvalidIndex,
int32_t arg1 = SemanticsNodeId::InvalidIndex)
: parse_node_(parse_node),
kind_(kind),
type_(type),
type_id_(type_id),
arg0_(arg0),
arg1_(arg1) {}
ParseTree::Node parse_node_;
SemanticsNodeKind kind_;
SemanticsNodeId type_;
SemanticsNodeId type_id_;
// Use GetAsKind to access arg0 and arg1.
int32_t arg0_;
int32_t arg1_;
};