#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