#include "type_checking.hh" #include "symbol_checking.hh" #include "functions.hh" #include "ast.hh" #include using namespace ast; using namespace std; using namespace type_checking; using namespace symbol_checking; void MatrixType::print (ostream &os) const { os << "matrix " << name() << '[' << arity() << "] : " << cell_type(); } void MatrixType::accept(TypeVisitor &v) { v.visit(*this); } void FunctionType::print (ostream &os) const { os << "fun " << name() << '(' << arity() << ") : " << return_type(); } void FunctionType::accept(TypeVisitor &v) { v.visit(*this); } void ValType::print (ostream &os) const { os << "val " << name() << " : " << type(); } void ValType::accept(TypeVisitor &v) { v.visit(*this); } void IndexType::print (ostream &os) const { os << "index " << name(); } void IndexType::accept(TypeVisitor &v) { v.visit(*this); } namespace type_checking { class TypeChecking : public Visitor{ protected: SymbolTable *_scope; public: TypeChecking(SymbolTable *scope = 0) : _scope(scope) {} const TypeInfo* check_defined (Ast &ast, const char *symbol) const; void check_index (Ast &ast, const char *var) const; void check_index_or_val (Ast &ast, const char *var) const; void check_matrix (Ast &ast, const char *name, int arity) const; void check_fun_call (Ast &ast, const char *name, int arity) const; void check_builtin_values_fun (Ast &ast, const char *name) const; void check_builtin_range_fun (Ast &ast, const char *name) const; virtual void visit(RExpr &rexpr); virtual void visit(WhenExpr &when_expr); virtual void visit(WhereExpr &where_expr); virtual void visit(Values &values); virtual void visit(Range &range); virtual void visit(SimpleFun &simple_fun); virtual void visit(MatrixExpr &matrix_expr); virtual void visit(FunCallExpr &fun_call_expr); virtual void visit(IDExpr &id_expr); virtual void visit(BinOpExpr &binop_expr); virtual void visit(NEGExpr &neg_expr); virtual void visit(RelBExpr &bexpr); virtual void visit(ANDBExpr &bexpr); virtual void visit(ORBExpr &bexpr); virtual void visit(NOTBExpr &bexpr); virtual void visit(Update &update); virtual void visit(DProg &dprog); }; }; const TypeInfo* TypeChecking::check_defined(Ast &ast, const char *symbol) const { assert(symbol != 0); const TypeInfo *tinfo = _scope->lookup(symbol); if (!tinfo) throw new UndefinedSymbol(ast, symbol); return tinfo; } void TypeChecking::check_index (Ast &ast, const char *symbol) const { assert(symbol != 0); const TypeInfo *tinfo = check_defined(ast, symbol); IndexType expected(symbol); if (expected != (*tinfo)) throw new WrongType(ast, new IndexType(symbol), tinfo); } void TypeChecking::check_index_or_val (Ast &ast, const char *symbol) const { assert(symbol != 0); const TypeInfo *tinfo = check_defined(ast, symbol); IndexType expected_index(symbol); if (expected_index == (*tinfo)) return; ValType expected_val(symbol, TYPE_INT); // FIXME: expected basic type? if (expected_val == (*tinfo)) return; // no go, throw exception throw new WrongType(ast, new IndexType(symbol), tinfo); } void TypeChecking::check_matrix (Ast &ast, const char *name, int arity) const { assert(name != 0); const TypeInfo *tinfo = check_defined(ast, name); MatrixType expected(name, arity); if (expected != (*tinfo)) throw new WrongType(ast, new MatrixType(name, arity), tinfo); } void TypeChecking::check_fun_call (Ast &ast, const char *name, int arity) const { assert(name != 0); const TypeInfo *tinfo = check_defined(ast, name); FunctionType expected(name, arity); if (expected != (*tinfo)) throw new WrongType(ast, new FunctionType(name, arity), tinfo); } void TypeChecking::check_builtin_values_fun(Ast &ast, const char *name) const { assert(name != 0); if (strcmp(name, "select") == 0) return; // special builtin if (! functions::defined(name)) throw new BuiltinNotDefined(ast, name); } void TypeChecking::check_builtin_range_fun(Ast &ast, const char *name) const { assert(name != 0); if (! functions::defined(name)) throw new BuiltinNotDefined(ast, name); } void TypeChecking::visit (RExpr &rexpr) { rexpr.begin()->accept(*this); rexpr.end()->accept(*this); } void TypeChecking::visit (WhenExpr &when_expr) { // check bexpr in this scope when_expr.bexpr()->accept(*this); // check fun in nested scope SymbolTable scope(*when_expr.fun(), _scope); TypeChecking tcheck(&scope); when_expr.fun()->accept(tcheck); } void TypeChecking::visit (WhereExpr &where_expr) { // check ranges in this scope list *rexpr_list = where_expr.rexpr_list(); list::iterator i; for (i = rexpr_list->begin(); i != rexpr_list->end(); ++i) (*i)->accept(*this); // check fun in nested scope SymbolTable scope(*where_expr.fun(), _scope); TypeChecking tcheck(&scope); where_expr.fun()->accept(tcheck); } void TypeChecking::visit (Values &values) { check_builtin_values_fun(values, values.id()); list *expr_list = values.when_expr_list(); list::iterator i; for (i = expr_list->begin(); i != expr_list->end(); ++i) (*i)->accept(*this); } void TypeChecking::visit(Range &range) { check_builtin_range_fun(range, range.id()); range.where_expr()->accept(*this); } void TypeChecking::visit (SimpleFun &simple_fun) { simple_fun.expr()->accept(*this); } void TypeChecking::visit (IDExpr &id_expr) { check_index_or_val(id_expr, id_expr.id()); } void TypeChecking::visit (MatrixExpr &matrix_expr) { const list *expr_list = matrix_expr.expr_list(); int arity; list::const_iterator i; for (i = expr_list->begin(), arity = 0; i != expr_list->end(); ++i, ++arity) (*i)->accept(*this); check_matrix(matrix_expr, matrix_expr.id(), arity); } void TypeChecking::visit (FunCallExpr &fun_call) { const list *expr_list = fun_call.expr_list(); int arity; list::const_iterator i; for (i = expr_list->begin(), arity = 0; i != expr_list->end(); ++i, ++arity) (*i)->accept(*this); check_fun_call(fun_call, fun_call.id(), arity); } void TypeChecking::visit (BinOpExpr &binop_expr) { binop_expr.left()->accept(*this); binop_expr.right()->accept(*this); } void TypeChecking::visit (NEGExpr &neg_expr) { neg_expr.expr()->accept(*this); } void TypeChecking::visit(RelBExpr &bexpr) { bexpr.left()->accept(*this); bexpr.right()->accept(*this); } void TypeChecking::visit(ANDBExpr &bexpr) { bexpr.left()->accept(*this); bexpr.right()->accept(*this); } void TypeChecking::visit(ORBExpr &bexpr) { bexpr.left()->accept(*this); bexpr.right()->accept(*this); } void TypeChecking::visit(NOTBExpr &bexpr) { bexpr.bexpr()->accept(*this); } void TypeChecking::visit (Update &update) { // check matrix check_matrix(update, update.name(), update.indices()->size()); // check indices in this scope const list *indices = update.indices(); list::const_iterator i; for (i = indices->begin(); i != indices->end(); ++i) check_index(update, *i); // check fun in nested scope SymbolTable scope(*update.fun(), _scope); TypeChecking tcheck(&scope); update.fun()->accept(tcheck); } void TypeChecking::visit (DProg &dprog) { SymbolTable scope(dprog, _scope); TypeChecking tcheck(&scope); const list *rexpr_list = dprog.rexpr_list(); list::const_iterator re_itr; for (re_itr = rexpr_list->begin(); re_itr != rexpr_list->end(); ++re_itr) (*re_itr)->accept(tcheck); const list *update_list = dprog.update_list(); list::const_iterator u_itr; for (u_itr = update_list->begin(); u_itr != update_list->end(); ++u_itr) (*u_itr)->accept(tcheck); } void type_checking::check (DProg* dprog) { TypeChecking type_checker; dprog->accept(type_checker); } void UndefinedSymbol::print_error_msg (std::ostream &os) { os << "Symbol " << _symbol << " undefined." << std::endl; } void BuiltinNotDefined::print_error_msg (std::ostream &os) { os << "Builtin function " << _symbol << " not defined." << std::endl; } void WrongType::print_error_msg (std::ostream &os) { os << "Type error, expected type: " << *_expected << " actual type: " << *_actual << std::endl; }