diff --git a/src/ast/AstVisitor.cpp b/src/ast/AstVisitor.cpp new file mode 100644 index 0000000..664796a --- /dev/null +++ b/src/ast/AstVisitor.cpp @@ -0,0 +1,39 @@ +#include "ast/AstVisitor.hpp" + +#include "ast/ClassNode.hpp" +#include "ast/MainClassNode.hpp" +#include "ast/MethodNode.hpp" +#include "ast/MethodParameterNode.hpp" +#include "ast/MethodWithoutParametersNode.hpp" +#include "ast/Node.h" +#include "ast/VariableNode.hpp" + +void AstVisitor::visit(const Node &node) { + for (const auto &child : node.children) { + child->accept(*this); + } +} + +void AstVisitor::visit(const ClassNode &node) { + visit(static_cast(node)); +} + +void AstVisitor::visit(const MainClassNode &node) { + visit(static_cast(node)); +} + +void AstVisitor::visit(const MethodNode &node) { + visit(static_cast(node)); +} + +void AstVisitor::visit(const MethodWithoutParametersNode &node) { + visit(static_cast(node)); +} + +void AstVisitor::visit(const MethodParameterNode &node) { + visit(static_cast(node)); +} + +void AstVisitor::visit(const VariableNode &node) { + visit(static_cast(node)); +} diff --git a/src/ast/AstVisitor.hpp b/src/ast/AstVisitor.hpp new file mode 100644 index 0000000..f9712bb --- /dev/null +++ b/src/ast/AstVisitor.hpp @@ -0,0 +1,25 @@ +#ifndef AST_VISITOR_HPP +#define AST_VISITOR_HPP + +class Node; +class ClassNode; +class MainClassNode; +class MethodNode; +class MethodWithoutParametersNode; +class MethodParameterNode; +class VariableNode; + +class AstVisitor { + public: + virtual ~AstVisitor() = default; + + virtual void visit(const Node &node); + virtual void visit(const ClassNode &node); + virtual void visit(const MainClassNode &node); + virtual void visit(const MethodNode &node); + virtual void visit(const MethodWithoutParametersNode &node); + virtual void visit(const MethodParameterNode &node); + virtual void visit(const VariableNode &node); +}; + +#endif diff --git a/src/ast/ClassNode.cpp b/src/ast/ClassNode.cpp index 91901f0..afd7b49 100644 --- a/src/ast/ClassNode.cpp +++ b/src/ast/ClassNode.cpp @@ -1,23 +1,8 @@ #include "ast/ClassNode.hpp" -bool ClassNode::buildTable(SymbolTable &st) const { - bool valid = true; - if (st.lookupClass(className)) { - std::cerr << "Error: "; - std::cerr << "(line " << lineno << ") "; - std::cerr << "Class " << className << " already declared.\n"; - valid = false; - } - - st.addClass(className); - auto *currentClass = st.lookupClass(className); - st.enterClassScope(currentClass); - st.addVariable(className, "this"); - bool validBody = body->buildTable(st); - st.exitScope(); +#include "ast/AstVisitor.hpp" - return valid && validBody; -} +void ClassNode::accept(AstVisitor &visitor) const { visitor.visit(*this); } std::string ClassNode::checkTypes(SymbolTable &st) const { st.enterClassScope(className); diff --git a/src/ast/ClassNode.hpp b/src/ast/ClassNode.hpp index d16353c..259da4d 100644 --- a/src/ast/ClassNode.hpp +++ b/src/ast/ClassNode.hpp @@ -14,7 +14,11 @@ class ClassNode : public Node { body = append_child(std::move(body_)); className = id->value; } - bool buildTable(SymbolTable &st) const override; + void accept(AstVisitor &visitor) const override; + + [[nodiscard]] const std::string &getClassName() const { return className; } + [[nodiscard]] const Node &getBodyNode() const { return *body; } + std::string checkTypes(SymbolTable &st) const override; Operand generateIR(CFG &graph, SymbolTable &st) override; }; diff --git a/src/ast/MainClassNode.cpp b/src/ast/MainClassNode.cpp index ee14f93..ac10baa 100644 --- a/src/ast/MainClassNode.cpp +++ b/src/ast/MainClassNode.cpp @@ -1,29 +1,8 @@ #include "ast/MainClassNode.hpp" -bool MainClassNode::buildTable(SymbolTable &st) const { - if (st.lookupClass(mainClassName)) { - std::cerr << "Error: (line " << lineno << ") Class '" << mainClassName - << "' already declared.\n"; - return false; - } - st.addClass(mainClassName); - auto *mainClass = st.lookupClass(mainClassName); - st.enterClassScope(mainClass); +#include "ast/AstVisitor.hpp" - st.addVariable(mainClassName, "this"); - Variable *mainClassThis = st.lookupVariableInScope("this"); - mainClass->addVariable(mainClassThis); - - st.addMethod("void", "main"); - Method *mainClassMethod = st.lookupMethod("main"); - mainClass->addMethod(mainClassMethod); - - st.enterMethodScope(mainClassMethod); - st.addVariable("String[]", mainMethodArgumentName); - st.exitScope(); - st.exitScope(); - return true; -} +void MainClassNode::accept(AstVisitor &visitor) const { visitor.visit(*this); } std::string MainClassNode::checkTypes(SymbolTable &st) const { st.enterClassScope(mainClassName); diff --git a/src/ast/MainClassNode.hpp b/src/ast/MainClassNode.hpp index 5992b4b..24e8dcb 100644 --- a/src/ast/MainClassNode.hpp +++ b/src/ast/MainClassNode.hpp @@ -18,7 +18,16 @@ class MainClassNode : public Node { mainMethodArgumentName = arg->value; } - bool buildTable(SymbolTable &st) const override; + void accept(AstVisitor &visitor) const override; + + [[nodiscard]] const std::string &getMainClassName() const { + return mainClassName; + } + [[nodiscard]] const std::string &getMainMethodArgumentName() const { + return mainMethodArgumentName; + } + [[nodiscard]] const Node &getBodyNode() const { return *body; } + std::string checkTypes(SymbolTable &st) const override; Operand generateIR(CFG &graph, SymbolTable &st) override; }; diff --git a/src/ast/MethodNode.cpp b/src/ast/MethodNode.cpp index 4d50fc8..f2825a8 100644 --- a/src/ast/MethodNode.cpp +++ b/src/ast/MethodNode.cpp @@ -1,25 +1,8 @@ #include "ast/MethodNode.hpp" -bool MethodNode::buildTable(SymbolTable &st) const { - bool valid = true; - if (st.lookupMethod(methodName)) { - std::cerr << "Error: (line " << lineno << ") Method '" << methodName - << "' already declared.\n"; - valid = false; - } - - st.addMethod(methodType, methodName); - auto *currentMethod = st.lookupMethod(methodName); - auto *currentClass = dynamic_cast(st.getCurrentRecord()); - currentClass->addMethod(currentMethod); +#include "ast/AstVisitor.hpp" - st.enterMethodScope(currentMethod); - bool validParams = params->buildTable(st); - bool validBody = body->buildTable(st); - st.exitScope(); - - return valid && validParams && validBody; -} +void MethodNode::accept(AstVisitor &visitor) const { visitor.visit(*this); } std::string MethodNode::checkTypes(SymbolTable &st) const { st.enterMethodScope(methodName); diff --git a/src/ast/MethodNode.hpp b/src/ast/MethodNode.hpp index 1f09b82..5c5d008 100644 --- a/src/ast/MethodNode.hpp +++ b/src/ast/MethodNode.hpp @@ -19,7 +19,14 @@ class MethodNode : public Node { methodName = id->value; methodType = type->value; } - bool buildTable(SymbolTable &st) const override; + + void accept(AstVisitor &visitor) const override; + + [[nodiscard]] const std::string &getMethodName() const { return methodName; } + [[nodiscard]] const std::string &getMethodType() const { return methodType; } + [[nodiscard]] const Node &getParametersNode() const { return *params; } + [[nodiscard]] const Node &getBodyNode() const { return *body; } + std::string checkTypes(SymbolTable &st) const override; Operand generateIR(CFG &graph, SymbolTable &st) override; }; diff --git a/src/ast/MethodParameterNode.cpp b/src/ast/MethodParameterNode.cpp index 03a9be7..8196ac9 100644 --- a/src/ast/MethodParameterNode.cpp +++ b/src/ast/MethodParameterNode.cpp @@ -1,15 +1,7 @@ #include "ast/MethodParameterNode.hpp" -bool MethodParameterNode::buildTable(SymbolTable &st) const { - if (st.lookupVariableInScope(id->value)) { - std::cerr << "Error: (line " << lineno << ") Parameter '" << id->value - << "' already declared.\n"; - return false; - } - st.addVariable(type->value, id->value); - auto *parameter = st.lookupVariable(id->value); - auto *currentScope = st.getCurrentScope(); - auto *currentMethod = dynamic_cast(currentScope->getRecord()); - currentMethod->addParameter(parameter); - return true; +#include "ast/AstVisitor.hpp" + +void MethodParameterNode::accept(AstVisitor &visitor) const { + visitor.visit(*this); } diff --git a/src/ast/MethodParameterNode.hpp b/src/ast/MethodParameterNode.hpp index 139872f..6f6ae01 100644 --- a/src/ast/MethodParameterNode.hpp +++ b/src/ast/MethodParameterNode.hpp @@ -13,7 +13,16 @@ class MethodParameterNode : public Node { type = append_child(std::move(type_)); id = append_child(std::move(id_)); } - bool buildTable(SymbolTable &st) const override; + + void accept(AstVisitor &visitor) const override; + + [[nodiscard]] const std::string &getParameterType() const { + return type->value; + } + [[nodiscard]] const std::string &getParameterName() const { + return id->value; + } + }; #endif diff --git a/src/ast/MethodWithoutParametersNode.cpp b/src/ast/MethodWithoutParametersNode.cpp index 785508d..0f9cf71 100644 --- a/src/ast/MethodWithoutParametersNode.cpp +++ b/src/ast/MethodWithoutParametersNode.cpp @@ -1,21 +1,9 @@ #include "ast/MethodWithoutParametersNode.hpp" -bool MethodWithoutParametersNode::buildTable(SymbolTable &st) const { - if (st.lookupMethod(id->value)) { - std::cerr << "Error: (line " << lineno << ") Method '" << id->value - << "' already declared.\n"; - return false; - } - st.addMethod(type->value, id->value); - auto *currentMethod = st.lookupMethod(id->value); - auto *currentClass = dynamic_cast(st.getCurrentRecord()); - currentClass->addMethod(currentMethod); - - st.enterMethodScope(currentMethod); - bool validBody = body->buildTable(st); - st.exitScope(); +#include "ast/AstVisitor.hpp" - return validBody; +void MethodWithoutParametersNode::accept(AstVisitor &visitor) const { + visitor.visit(*this); } std::string MethodWithoutParametersNode::checkTypes(SymbolTable &st) const { diff --git a/src/ast/MethodWithoutParametersNode.hpp b/src/ast/MethodWithoutParametersNode.hpp index 3fac9e4..d90d663 100644 --- a/src/ast/MethodWithoutParametersNode.hpp +++ b/src/ast/MethodWithoutParametersNode.hpp @@ -16,7 +16,12 @@ class MethodWithoutParametersNode : public Node { body = append_child(std::move(body_)); } - bool buildTable(SymbolTable &st) const override; + void accept(AstVisitor &visitor) const override; + + [[nodiscard]] const std::string &getMethodName() const { return id->value; } + [[nodiscard]] const std::string &getMethodType() const { return type->value; } + [[nodiscard]] const Node &getBodyNode() const { return *body; } + std::string checkTypes(SymbolTable &st) const override; Operand generateIR(CFG &graph, SymbolTable &st) override; }; diff --git a/src/ast/Node.cpp b/src/ast/Node.cpp index 59f4b9c..b58d78f 100644 --- a/src/ast/Node.cpp +++ b/src/ast/Node.cpp @@ -1,14 +1,11 @@ #include "ast/Node.h" + +#include "ast/AstVisitor.hpp" #include "semantic/SymbolTable.hpp" +#include "semantic/SymbolTableVisitor.hpp" bool Node::buildTable(SymbolTable &st) const { - bool valid = true; - for (const auto &child : children) { - if (!child->buildTable(st)) { - valid = false; - } - } - return valid; + return build_symbol_table(*this, st).ok(); } std::string Node::checkTypes(SymbolTable &st) const { @@ -32,6 +29,8 @@ Operand Node::generateIR(CFG &graph, SymbolTable &st) { return "foobar"; } +void Node::accept(AstVisitor &visitor) const { visitor.visit(*this); } + void Node::print(int depth = 0) const { for (int i = 0; i < depth; i++) { std::cout << " "; diff --git a/src/ast/Node.h b/src/ast/Node.h index 2f4a0f7..0e2b0e9 100644 --- a/src/ast/Node.h +++ b/src/ast/Node.h @@ -12,6 +12,8 @@ #include "ir/CFG.hpp" +class AstVisitor; + class Node { public: std::string type{}, value{}; @@ -31,6 +33,8 @@ class Node { virtual Operand generateIR(CFG &graph, SymbolTable &st); + virtual void accept(AstVisitor &visitor) const; + void print(int depth) const; void printGraphviz(int &count, std::ostream &outStream); diff --git a/src/ast/VariableNode.cpp b/src/ast/VariableNode.cpp index 427c2b4..ea99317 100644 --- a/src/ast/VariableNode.cpp +++ b/src/ast/VariableNode.cpp @@ -1,30 +1,8 @@ #include "ast/VariableNode.hpp" -bool VariableNode::buildTable(SymbolTable &st) const { +#include "ast/AstVisitor.hpp" - Variable *lookup = st.lookupVariableInScope(name->value); - if (lookup) { - std::cerr << "Error: (line " << lineno << ") " - << "Variable '" << name->value << "' " - << "already declared.\n"; - return false; - } - st.addVariable(type->value, name->value); - Variable *currentVariable = st.lookupVariable(name->value); - - Record *curRecord = st.getCurrentRecord(); - if (curRecord->getType() == curRecord->getID()) { - // Record is a class - auto *curClass = dynamic_cast(curRecord); - curClass->addVariable(currentVariable); - } else { - // Record is a method - auto *curMethod = dynamic_cast(curRecord); - curMethod->addVariable(currentVariable); - } - - return true; -}; +void VariableNode::accept(AstVisitor &visitor) const { visitor.visit(*this); } std::string VariableNode::checkTypes(SymbolTable &st) const { const auto &variableType = type->value; diff --git a/src/ast/VariableNode.hpp b/src/ast/VariableNode.hpp index 88bde80..4f1e9be 100644 --- a/src/ast/VariableNode.hpp +++ b/src/ast/VariableNode.hpp @@ -14,7 +14,15 @@ class VariableNode : public Node { name = append_child(std::move(name_)); } - bool buildTable(SymbolTable &st) const override; + void accept(AstVisitor &visitor) const override; + + [[nodiscard]] const std::string &getVariableType() const { + return type->value; + } + [[nodiscard]] const std::string &getVariableName() const { + return name->value; + } + std::string checkTypes(SymbolTable &st) const override; }; diff --git a/src/main.cpp b/src/main.cpp index fc09234..debeace 100644 --- a/src/main.cpp +++ b/src/main.cpp @@ -15,6 +15,7 @@ namespace fs = std::filesystem; #include "lexing/SourceBuffer.hpp" #include "lexing/StringViewStream.hpp" #include "parsing/Parser.hpp" +#include "semantic/SymbolTableVisitor.hpp" std::unique_ptr root; int lexical_errors = 0; @@ -273,19 +274,24 @@ int main(int argc, char **argv) { } SymbolTable st; - bool symbolTableSuccess = root->buildTable(st); + lexing::LegacyDiagnosticSink semantic_diag; + auto symbol_table_result = build_symbol_table(*root, st, &semantic_diag); + bool symbolTableSuccess = symbol_table_result.ok(); if (!symbolTableSuccess) { std::cout << "Symbol table construction failed.\n"; } - auto rootType = root->checkTypes(st); - bool typeCheckSuccess = !rootType.empty(); - if (!typeCheckSuccess) { - std::cout << "Type checking failed.\n"; + bool typeCheckSuccess = true; + if (symbolTableSuccess) { + auto rootType = root->checkTypes(st); + typeCheckSuccess = !rootType.empty(); + if (!typeCheckSuccess) { + std::cout << "Type checking failed.\n"; + } } if (!symbolTableSuccess || !typeCheckSuccess) { - return errCode; + return errCodes::SEMANTIC_ERROR; } std::ofstream outStream(outputDirectory / "tree.dot"); diff --git a/src/semantic/SymbolTableVisitor.cpp b/src/semantic/SymbolTableVisitor.cpp new file mode 100644 index 0000000..8d522cb --- /dev/null +++ b/src/semantic/SymbolTableVisitor.cpp @@ -0,0 +1,226 @@ +#include "semantic/SymbolTableVisitor.hpp" + +#include +#include +#include +#include + +#include "ast/ClassNode.hpp" +#include "ast/MainClassNode.hpp" +#include "ast/MethodNode.hpp" +#include "ast/MethodParameterNode.hpp" +#include "ast/MethodWithoutParametersNode.hpp" +#include "ast/Node.h" +#include "ast/VariableNode.hpp" +#include "semantic/Class.hpp" +#include "semantic/Method.hpp" +#include "semantic/Record.hpp" +#include "semantic/Variable.hpp" + +namespace { + +class ScopeExit { + public: + explicit ScopeExit(SymbolTable &table) : table_(&table) {} + + ~ScopeExit() { + if (table_ != nullptr) { + table_->exitScope(); + } + } + + ScopeExit(const ScopeExit &) = delete; + ScopeExit &operator=(const ScopeExit &) = delete; + + private: + SymbolTable *table_ = nullptr; +}; + +} // namespace + +void SymbolTableVisitor::emit_error(int line, std::string message) { + error_count_ += 1; + + if (sink_ == nullptr) { + return; + } + + const auto line_no = static_cast(std::max(line, 1)); + const lexing::SourceSpan span{ + .begin = {.offset = 0, .line = line_no, .column = 1}, + .end = {.offset = 0, .line = line_no, .column = 1}, + }; + + sink_->emit({.severity = lexing::Severity::Error, + .message = std::move(message), + .span = span}); +} + +void SymbolTableVisitor::visit(const Node &node) { + for (const auto &child : node.children) { + child->accept(*this); + } +} + +void SymbolTableVisitor::visit(const ClassNode &node) { + const auto &class_name = node.getClassName(); + + if (table_.lookupClass(class_name) != nullptr) { + emit_error(node.lineno, "Error: (line " + std::to_string(node.lineno) + + ") Class '" + class_name + + "' already declared.\n"); + return; + } + + table_.addClass(class_name); + auto *current_class = table_.lookupClass(class_name); + table_.enterClassScope(current_class); + ScopeExit exit_scope(table_); + + table_.addVariable(class_name, "this"); + auto *this_variable = table_.lookupVariableInScope("this"); + current_class->addVariable(this_variable); + + node.getBodyNode().accept(*this); +} + +void SymbolTableVisitor::visit(const MainClassNode &node) { + const auto &main_class_name = node.getMainClassName(); + + if (table_.lookupClass(main_class_name) != nullptr) { + emit_error(node.lineno, + "Error: (line " + std::to_string(node.lineno) + + ") Class '" + main_class_name + + "' already declared.\n"); + return; + } + + table_.addClass(main_class_name); + auto *main_class = table_.lookupClass(main_class_name); + table_.enterClassScope(main_class); + ScopeExit exit_class_scope(table_); + + table_.addVariable(main_class_name, "this"); + auto *main_class_this = table_.lookupVariableInScope("this"); + main_class->addVariable(main_class_this); + + table_.addMethod("void", "main"); + auto *main_class_method = table_.lookupMethod("main"); + main_class->addMethod(main_class_method); + + table_.enterMethodScope(main_class_method); + { + ScopeExit exit_method_scope(table_); + table_.addVariable("String[]", node.getMainMethodArgumentName()); + } + + node.getBodyNode().accept(*this); +} + +void SymbolTableVisitor::visit(const MethodNode &node) { + auto *current_class = dynamic_cast(table_.getCurrentRecord()); + const auto &method_name = node.getMethodName(); + + if (current_class != nullptr && current_class->lookupMethod(method_name)) { + emit_error(node.lineno, + "Error: (line " + std::to_string(node.lineno) + + ") Method '" + method_name + "' already declared.\n"); + return; + } + + table_.addMethod(node.getMethodType(), method_name); + auto *current_method = table_.lookupMethod(method_name); + + if (current_class != nullptr) { + current_class->addMethod(current_method); + } + + table_.enterMethodScope(current_method); + ScopeExit exit_method_scope(table_); + + node.getParametersNode().accept(*this); + node.getBodyNode().accept(*this); +} + +void SymbolTableVisitor::visit(const MethodWithoutParametersNode &node) { + auto *current_class = dynamic_cast(table_.getCurrentRecord()); + const auto &method_name = node.getMethodName(); + + if (current_class != nullptr && current_class->lookupMethod(method_name)) { + emit_error(node.lineno, + "Error: (line " + std::to_string(node.lineno) + + ") Method '" + method_name + "' already declared.\n"); + return; + } + + table_.addMethod(node.getMethodType(), method_name); + auto *current_method = table_.lookupMethod(method_name); + + if (current_class != nullptr) { + current_class->addMethod(current_method); + } + + table_.enterMethodScope(current_method); + ScopeExit exit_method_scope(table_); + + node.getBodyNode().accept(*this); +} + +void SymbolTableVisitor::visit(const MethodParameterNode &node) { + const auto ¶meter_name = node.getParameterName(); + + if (table_.lookupVariableInScope(parameter_name) != nullptr) { + emit_error(node.lineno, + "Error: (line " + std::to_string(node.lineno) + + ") Parameter '" + parameter_name + + "' already declared.\n"); + return; + } + + table_.addVariable(node.getParameterType(), parameter_name); + auto *parameter = table_.lookupVariable(parameter_name); + + auto *current_scope = table_.getCurrentScope(); + auto *current_method = + dynamic_cast(current_scope != nullptr ? current_scope->getRecord() + : nullptr); + if (current_method != nullptr) { + current_method->addParameter(parameter); + } +} + +void SymbolTableVisitor::visit(const VariableNode &node) { + const auto &variable_name = node.getVariableName(); + + if (table_.lookupVariableInScope(variable_name) != nullptr) { + emit_error(node.lineno, + "Error: (line " + std::to_string(node.lineno) + + ") Variable '" + variable_name + + "' already declared.\n"); + return; + } + + table_.addVariable(node.getVariableType(), variable_name); + auto *current_variable = table_.lookupVariable(variable_name); + + auto *current_record = table_.getCurrentRecord(); + if (auto *current_class = dynamic_cast(current_record)) { + current_class->addVariable(current_variable); + return; + } + + if (auto *current_method = dynamic_cast(current_record)) { + current_method->addVariable(current_variable); + } +} + +SemanticPassResult build_symbol_table(const Node &root, SymbolTable &table, + lexing::DiagnosticSink *sink) { + while (table.getParentScope() != nullptr) { + table.exitScope(); + } + + SymbolTableVisitor visitor(table, sink); + root.accept(visitor); + return visitor.result(); +} diff --git a/src/semantic/SymbolTableVisitor.hpp b/src/semantic/SymbolTableVisitor.hpp new file mode 100644 index 0000000..6e805ec --- /dev/null +++ b/src/semantic/SymbolTableVisitor.hpp @@ -0,0 +1,51 @@ +#ifndef SYMBOL_TABLE_VISITOR_HPP +#define SYMBOL_TABLE_VISITOR_HPP + +#include + +#include "ast/AstVisitor.hpp" +#include "lexing/Diagnostics.hpp" +#include "semantic/SymbolTable.hpp" + +class Node; +class ClassNode; +class MainClassNode; +class MethodNode; +class MethodWithoutParametersNode; +class MethodParameterNode; +class VariableNode; + +struct SemanticPassResult { + int error_count = 0; + + [[nodiscard]] bool ok() const { return error_count == 0; } +}; + +class SymbolTableVisitor : public AstVisitor { + public: + explicit SymbolTableVisitor(SymbolTable &table, + lexing::DiagnosticSink *sink = nullptr) + : table_(table), sink_(sink) {} + + SemanticPassResult result() const { return {.error_count = error_count_}; } + + void visit(const Node &node) override; + void visit(const ClassNode &node) override; + void visit(const MainClassNode &node) override; + void visit(const MethodNode &node) override; + void visit(const MethodWithoutParametersNode &node) override; + void visit(const MethodParameterNode &node) override; + void visit(const VariableNode &node) override; + + private: + SymbolTable &table_; + lexing::DiagnosticSink *sink_ = nullptr; + int error_count_ = 0; + + void emit_error(int line, std::string message); +}; + +SemanticPassResult build_symbol_table(const Node &root, SymbolTable &table, + lexing::DiagnosticSink *sink = nullptr); + +#endif diff --git a/tests/symbol_table_test.cpp b/tests/symbol_table_test.cpp index d97ce01..8f6c859 100644 --- a/tests/symbol_table_test.cpp +++ b/tests/symbol_table_test.cpp @@ -16,6 +16,7 @@ #include "semantic/Method.hpp" #include "semantic/Scope.hpp" #include "semantic/SymbolTable.hpp" +#include "semantic/SymbolTableVisitor.hpp" namespace { @@ -37,6 +38,27 @@ void assert_no_errors(const std::vector &diagnostics) { << "Unexpected diagnostic errors: " << error_count; } +int count_error_diagnostics( + const std::vector &diagnostics) { + int count = 0; + for (const auto &d : diagnostics) { + if (d.severity == lexing::Severity::Error) { + count += 1; + } + } + return count; +} + +const lexing::Diagnostic * +find_first_error(const std::vector &diagnostics) { + for (const auto &d : diagnostics) { + if (d.severity == lexing::Severity::Error) { + return &d; + } + } + return nullptr; +} + std::unique_ptr parse_program(std::string_view source) { CollectingDiagnosticSink diag; auto stream = std::make_unique(source); @@ -164,7 +186,10 @@ class Foo { ASSERT_NE(root, nullptr); SymbolTable st; - ASSERT_TRUE(root->buildTable(st)); + CollectingDiagnosticSink semantic_diag; + const auto symbol_result = build_symbol_table(*root, st, &semantic_diag); + ASSERT_TRUE(symbol_result.ok()); + ASSERT_EQ(count_error_diagnostics(semantic_diag.diagnostics), 0); const Scope *program = st.getCurrentScope(); ASSERT_NE(program, nullptr); @@ -247,7 +272,10 @@ TEST(SymbolTable, GoldenProgram2) { ASSERT_NE(root, nullptr); SymbolTable st; - ASSERT_TRUE(root->buildTable(st)); + CollectingDiagnosticSink semantic_diag; + const auto symbol_result = build_symbol_table(*root, st, &semantic_diag); + ASSERT_TRUE(symbol_result.ok()); + ASSERT_EQ(count_error_diagnostics(semantic_diag.diagnostics), 0); const Scope *program = st.getCurrentScope(); ASSERT_NE(program, nullptr); @@ -439,3 +467,172 @@ TEST(SymbolTable, GoldenProgram2) { const auto root_type = root->checkTypes(st); EXPECT_EQ(root_type, "void"); } + +TEST(SymbolTable, DuplicateClassReportsDiagnostic) { + constexpr std::string_view source = R"(public class Main { + public static void main(String[] args) { + System.out.println(0); + } +} + +class Foo { +} + +class Foo { +} +)"; + + auto root = parse_program(source); + ASSERT_NE(root, nullptr); + + SymbolTable st; + CollectingDiagnosticSink semantic_diag; + const auto symbol_result = build_symbol_table(*root, st, &semantic_diag); + + ASSERT_FALSE(symbol_result.ok()); + ASSERT_EQ(symbol_result.error_count, 1); + ASSERT_EQ(count_error_diagnostics(semantic_diag.diagnostics), 1); + + const auto *error = find_first_error(semantic_diag.diagnostics); + ASSERT_NE(error, nullptr); + EXPECT_EQ(error->span.begin.line, 10u); + EXPECT_NE(error->message.find("Class 'Foo' already declared"), + std::string::npos); +} + +TEST(SymbolTable, DuplicateMethodReportsDiagnostic) { + constexpr std::string_view source = R"(public class Main { + public static void main(String[] args) { + System.out.println(0); + } +} + +class Foo { + public int bar() { + return 0; + } + + public int bar() { + return 1; + } +} +)"; + + auto root = parse_program(source); + ASSERT_NE(root, nullptr); + + SymbolTable st; + CollectingDiagnosticSink semantic_diag; + const auto symbol_result = build_symbol_table(*root, st, &semantic_diag); + + ASSERT_FALSE(symbol_result.ok()); + ASSERT_EQ(symbol_result.error_count, 1); + ASSERT_EQ(count_error_diagnostics(semantic_diag.diagnostics), 1); + + const auto *error = find_first_error(semantic_diag.diagnostics); + ASSERT_NE(error, nullptr); + EXPECT_EQ(error->span.begin.line, 12u); + EXPECT_NE(error->message.find("Method 'bar' already declared"), + std::string::npos); +} + +TEST(SymbolTable, DuplicateParameterReportsDiagnostic) { + constexpr std::string_view source = R"(public class Main { + public static void main(String[] args) { + System.out.println(0); + } +} + +class Foo { + public int bar(int x, int x) { + return x; + } +} +)"; + + auto root = parse_program(source); + ASSERT_NE(root, nullptr); + + SymbolTable st; + CollectingDiagnosticSink semantic_diag; + const auto symbol_result = build_symbol_table(*root, st, &semantic_diag); + + ASSERT_FALSE(symbol_result.ok()); + ASSERT_EQ(symbol_result.error_count, 1); + ASSERT_EQ(count_error_diagnostics(semantic_diag.diagnostics), 1); + + const auto *error = find_first_error(semantic_diag.diagnostics); + ASSERT_NE(error, nullptr); + EXPECT_EQ(error->span.begin.line, 8u); + EXPECT_NE(error->message.find("Parameter 'x' already declared"), + std::string::npos); +} + +TEST(SymbolTable, DuplicateLocalVariableReportsDiagnostic) { + constexpr std::string_view source = R"(public class Main { + public static void main(String[] args) { + System.out.println(0); + } +} + +class Foo { + public int bar() { + int x; + int x; + return x; + } +} +)"; + + auto root = parse_program(source); + ASSERT_NE(root, nullptr); + + SymbolTable st; + CollectingDiagnosticSink semantic_diag; + const auto symbol_result = build_symbol_table(*root, st, &semantic_diag); + + ASSERT_FALSE(symbol_result.ok()); + ASSERT_EQ(symbol_result.error_count, 1); + ASSERT_EQ(count_error_diagnostics(semantic_diag.diagnostics), 1); + + const auto *error = find_first_error(semantic_diag.diagnostics); + ASSERT_NE(error, nullptr); + EXPECT_EQ(error->span.begin.line, 10u); + EXPECT_NE(error->message.find("Variable 'x' already declared"), + std::string::npos); +} + +TEST(SymbolTable, DuplicateFieldReportsDiagnostic) { + constexpr std::string_view source = R"(public class Main { + public static void main(String[] args) { + System.out.println(0); + } +} + +class Foo { + int x; + int x; + + public int bar() { + return x; + } +} +)"; + + auto root = parse_program(source); + ASSERT_NE(root, nullptr); + + SymbolTable st; + CollectingDiagnosticSink semantic_diag; + const auto symbol_result = build_symbol_table(*root, st, &semantic_diag); + + ASSERT_FALSE(symbol_result.ok()); + ASSERT_EQ(symbol_result.error_count, 1); + ASSERT_EQ(count_error_diagnostics(semantic_diag.diagnostics), 1); + + const auto *error = find_first_error(semantic_diag.diagnostics); + ASSERT_NE(error, nullptr); + EXPECT_EQ(error->span.begin.line, 9u); + EXPECT_NE(error->message.find("Variable 'x' already declared"), + std::string::npos); +}