#include "type_checking.hh"
#include "symbol_checking.hh"
#include "functions.hh"
#include "ast.hh"
#include <map>
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*> *rexpr_list = where_expr.rexpr_list();
list<RExpr*>::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<WhenExpr*> *expr_list = values.when_expr_list();
list<WhenExpr*>::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*> *expr_list = matrix_expr.expr_list();
int arity;
list<Expr*>::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*> *expr_list = fun_call.expr_list();
int arity;
list<Expr*>::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<const char *> *indices = update.indices();
list<const char *>::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*> *rexpr_list = dprog.rexpr_list();
list<RExpr*>::const_iterator re_itr;
for (re_itr = rexpr_list->begin(); re_itr != rexpr_list->end(); ++re_itr)
(*re_itr)->accept(tcheck);
const list<Update*> *update_list = dprog.update_list();
list<Update*>::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;
}
syntax highlighted by Code2HTML, v. 0.9.1