diff --git a/furc/include/furc/ast/declaration.hpp b/furc/include/furc/ast/declaration.hpp index 2f501c0..557b618 100644 --- a/furc/include/furc/ast/declaration.hpp +++ b/furc/include/furc/ast/declaration.hpp @@ -36,6 +36,8 @@ public: front::token name() const { return m_name; } public: + void accept(visitor& visitor) const override; + std::ostream& print(std::ostream& os) const override; protected: bool equal(const node& rhs) const override; @@ -65,6 +67,8 @@ public: const function_body_h& body() const { return m_body; } public: + void accept(visitor& visitor) const override; + std::ostream& print(std::ostream& os) const override; protected: bool equal(const node& rhs) const override; diff --git a/furc/include/furc/ast/expression.hpp b/furc/include/furc/ast/expression.hpp index a911651..57f2417 100644 --- a/furc/include/furc/ast/expression.hpp +++ b/furc/include/furc/ast/expression.hpp @@ -36,6 +36,8 @@ public: handle&& move_name() { return std::move(m_name); } public: expression_node_t expression_type() const override { return expression_node_t::VarRead; } +public: + void accept(visitor& visitor) const override; std::ostream& print(std::ostream& os) const override; protected: @@ -66,6 +68,8 @@ public: expression_node_h&& move_node() { return std::move(m_node); } public: expression_node_t expression_type() const override { return expression_node_t::Unaryop; } +public: + void accept(visitor& visitor) const override; std::ostream& print(std::ostream& os) const override; protected: @@ -105,6 +109,8 @@ public: expression_node_h&& move_rhs() { return std::move(m_rhs); }; public: expression_node_t expression_type() const override { return expression_node_t::Binop; } +public: + void accept(visitor& visitor) const override; std::ostream& print(std::ostream& os) const override; protected: @@ -128,6 +134,8 @@ public: const expression_node_h& rhs() const { return m_rhs; } public: expression_node_t expression_type() const override { return expression_node_t::VarAssign; } +public: + void accept(visitor& visitor) const override; std::ostream& print(std::ostream& os) const override; protected: diff --git a/furc/include/furc/ast/fwd.hpp b/furc/include/furc/ast/fwd.hpp index 6c72906..81be929 100644 --- a/furc/include/furc/ast/fwd.hpp +++ b/furc/include/furc/ast/fwd.hpp @@ -10,8 +10,8 @@ namespace ast { class node; -template -using node_handle = handle; +template +using node_handle = handle; class literal_node; using literal_node_h = node_handle; diff --git a/furc/include/furc/ast/literal.hpp b/furc/include/furc/ast/literal.hpp index 1938ee7..1338580 100644 --- a/furc/include/furc/ast/literal.hpp +++ b/furc/include/furc/ast/literal.hpp @@ -33,6 +33,8 @@ public: const handle& value() const { return m_value; } public: + void accept(visitor& visitor) const override; + std::ostream& print(std::ostream& os) const override; protected: bool equal(const node& rhs) const override; @@ -51,6 +53,8 @@ public: bool operator==(front::integer_token integer) const { return m_value == integer; } public: + void accept(visitor& visitor) const override; + std::ostream& print(std::ostream& os) const override; protected: bool equal(const node& rhs) const override; diff --git a/furc/include/furc/ast/node.hpp b/furc/include/furc/ast/node.hpp index 26c3ba0..c3568a4 100644 --- a/furc/include/furc/ast/node.hpp +++ b/furc/include/furc/ast/node.hpp @@ -2,6 +2,7 @@ #define FURC_AST_NODE_HPP #include "furc/ast/fwd.hpp" +#include "furc/ast/visitor.hpp" namespace furc { namespace ast { @@ -39,6 +40,8 @@ public: bool operator==(const node& rhs) const { return category() == rhs.category() && equal(rhs); } bool operator!=(const node& rhs) const { return !this->operator==(rhs); } public: + virtual void accept(visitor& visitor) const = 0; + virtual std::ostream& print(std::ostream& os) const = 0; friend std::ostream& operator<<(std::ostream& os, const node& node) { return node.print(os); } diff --git a/furc/include/furc/ast/program.hpp b/furc/include/furc/ast/program.hpp index 1f5b137..62a45f3 100644 --- a/furc/include/furc/ast/program.hpp +++ b/furc/include/furc/ast/program.hpp @@ -19,6 +19,8 @@ public: const std::vector>& declarations() const { return m_declarations; } public: + void accept(visitor& visitor) const override; + std::ostream& print(std::ostream& os) const override; protected: bool equal(const node& rhs) const override; diff --git a/furc/include/furc/ast/statement.hpp b/furc/include/furc/ast/statement.hpp index 2dfb3cc..f9d1092 100644 --- a/furc/include/furc/ast/statement.hpp +++ b/furc/include/furc/ast/statement.hpp @@ -31,6 +31,8 @@ public: node_handle value() const { return m_value; } public: statement_node_t statement_type() const override { return statement_node_t::Return; } +public: + void accept(visitor& visitor) const override; std::ostream& print(std::ostream& os) const override; protected: diff --git a/furc/include/furc/ast/visitor.hpp b/furc/include/furc/ast/visitor.hpp new file mode 100644 index 0000000..32cc2a8 --- /dev/null +++ b/furc/include/furc/ast/visitor.hpp @@ -0,0 +1,34 @@ +#ifndef FURC_AST_VISITOR_HPP +#define FURC_AST_VISITOR_HPP + +#include "furc/ast/fwd.hpp" + +namespace furc { +namespace ast { + +class visitor { +public: + virtual ~visitor() = default; + + visitor(visitor&&) = default; + visitor& operator=(visitor&&) = default; + visitor(const visitor&) = default; + visitor& operator=(const visitor&) = default; +public: + virtual void visit_string_literal_node(const string_literal_node&) {} + virtual void visit_integer_literal_node(const integer_literal_node&) {} + virtual void visit_var_read_expression_node(const var_read_expression_node&) {} + virtual void visit_unaryop_expression_node(const unaryop_expression_node&) {} + virtual void visit_binop_expression_node(const binop_expression_node&) {} + virtual void visit_var_assign_expression_node(const var_assign_expression_node&) {} + virtual void visit_function_declaration_node(const function_declaration_node&) {} + virtual void visit_function_definition_node(const function_definition_node&) {} + virtual void visit_return_statement_node(const return_statement_node&) {} + + virtual void visit_error(const node_handle& handle) {} +}; + +} // namespace ast +} // namespace furc + +#endif // FURC_AST_VISITOR_HPP \ No newline at end of file diff --git a/furc/src/ast.cpp b/furc/src/ast.cpp index 600ccbd..d6a656f 100644 --- a/furc/src/ast.cpp +++ b/furc/src/ast.cpp @@ -12,13 +12,8 @@ bool literal_node::equal(const node& rhs) const { return literal_type() == reinterpret_cast(rhs).literal_type(); } -std::ostream& integer_literal_node::print(std::ostream& os) const { - if (m_value.has_error()) return os << m_value.error(); - return os << *m_value; -} - -bool integer_literal_node::equal(const node& rhs) const { - return literal_node::equal(rhs) && m_value == reinterpret_cast(rhs).m_value; +void string_literal_node::accept(visitor& visitor) const { + visitor.visit_string_literal_node(*this); } std::ostream& string_literal_node::print(std::ostream& os) const { @@ -30,10 +25,27 @@ bool string_literal_node::equal(const node& rhs) const { return literal_node::equal(rhs) && m_value == reinterpret_cast(rhs).m_value; } +void integer_literal_node::accept(visitor& visitor) const { + visitor.visit_integer_literal_node(*this); +} + +std::ostream& integer_literal_node::print(std::ostream& os) const { + if (m_value.has_error()) return os << m_value.error(); + return os << *m_value; +} + +bool integer_literal_node::equal(const node& rhs) const { + return literal_node::equal(rhs) && m_value == reinterpret_cast(rhs).m_value; +} + bool expression_node::equal(const node& rhs) const { return expression_type() == reinterpret_cast(rhs).expression_type(); } +void var_read_expression_node::accept(visitor& visitor) const { + visitor.visit_var_read_expression_node(*this); +} + std::ostream& var_read_expression_node::print(std::ostream& os) const { if (m_name.present()) return os << *m_name; return os << m_name.error(); @@ -56,6 +68,10 @@ std::ostream& operator<<(std::ostream& os, unaryop_expression_node_t type) { return os; } +void unaryop_expression_node::accept(visitor& visitor) const { + visitor.visit_unaryop_expression_node(*this); +} + std::ostream& unaryop_expression_node::print(std::ostream& os) const { switch (m_type) { case unaryop_expression_node_t::Positive: @@ -91,6 +107,10 @@ std::ostream& operator<<(std::ostream& os, binop_expression_node_t type) { } } +void binop_expression_node::accept(visitor& visitor) const { + visitor.visit_binop_expression_node(*this); +} + std::ostream& binop_expression_node::print(std::ostream& os) const { if (m_type == binop_expression_node_t::None) return os; return os << '(' << *m_lhs << ' ' << m_type << ' ' << *m_rhs << ')'; @@ -101,6 +121,10 @@ bool binop_expression_node::equal(const node& rhsNode) const { return expression_node::equal(rhsNode) && m_type == rhs.m_type && m_lhs == rhs.m_lhs && m_rhs == rhs.m_rhs; } +void var_assign_expression_node::accept(visitor& visitor) const { + visitor.visit_var_assign_expression_node(*this); +} + std::ostream& var_assign_expression_node::print(std::ostream& os) const { return os << '(' << *m_lhs << ' ' << m_compound << "= " << *m_rhs << ')'; } @@ -114,6 +138,10 @@ bool declaration_node::equal(const node& rhs) const { return declaration_type() == reinterpret_cast(rhs).declaration_type(); } +void function_declaration_node::accept(visitor& visitor) const { + visitor.visit_function_declaration_node(*this); +} + std::ostream& function_declaration_node::print(std::ostream& os) const { return os << "function " << m_name->string << " declaration"; } @@ -122,6 +150,10 @@ bool function_declaration_node::equal(const node& rhs) const { return declaration_node::equal(rhs) && m_name == reinterpret_cast(rhs).m_name; } +void function_definition_node::accept(visitor& visitor) const { + visitor.visit_function_definition_node(*this); +} + std::ostream& function_definition_node::print(std::ostream& os) const { function_declaration_node::print(os); os << ':'; @@ -142,6 +174,10 @@ bool statement_node::equal(const node& rhs) const { return statement_type() == reinterpret_cast(rhs).statement_type(); } +void return_statement_node::accept(visitor& visitor) const { + visitor.visit_return_statement_node(*this); +} + std::ostream& return_statement_node::print(std::ostream& os) const { os << "return statement"; if (m_value.present()) return os << ' ' << *m_value; @@ -152,8 +188,14 @@ bool return_statement_node::equal(const node& rhs) const { return statement_node::equal(rhs) && m_value == reinterpret_cast(rhs).m_value; } -bool program_node::equal(const node& rhs) const { - return m_declarations == reinterpret_cast(rhs).m_declarations; +void program_node::accept(visitor& visitor) const { + for (const auto& decl : m_declarations) { + if (decl.has_error()) { + visitor.visit_error(decl); + } else { + decl->accept(visitor); + } + } } std::ostream& program_node::print(std::ostream& os) const { @@ -164,4 +206,8 @@ std::ostream& program_node::print(std::ostream& os) const { return os; } +bool program_node::equal(const node& rhs) const { + return m_declarations == reinterpret_cast(rhs).m_declarations; +} + } // namespace furc::ast \ No newline at end of file