#include <cctype>
#include "node.hpp"
#include "parser.hpp"

static const char *name[] = {
            "NUM", "VART", "LFT", "RGT",
            "ADD", "SUB", "MUL", "DIV", "POW",
            "SIN", "COS", "ATAN", "LOG", "EXP", "SQRT", "DER",
            "LPAR", "RPAR", "SEMI", "BAD"
};

//-----------------------
//   Parser
//-----------------------
Parser::Parser(const std::string& str, Node* l, Node* r) 
											: std::istringstream(str)
{
	ch= ' ';
	lft= l;
	rgt= r;
}

Parser::~Parser(void)
{
}

void Parser::error(const char* msg)
{
	std::cerr << "ERROR: " << msg << std::endl;
	exit(EXIT_FAILURE);
}

void Parser::Lex(void)
{
	*this >> ws;
	int c= get();
	if (isdigit(c)) {
		putback(c);
		*this >> num;
		tok= NUM;
	}
	else
	if (isalpha(c)) {
		std::string fun;
		while (isalpha(c)) {
			fun+= char(c);
			c= get();
		}
		this->putback(c);
		if (fun=="t")
			tok= VART;
		else
		if (fun=="lft")
			tok= LFT;
		else
		if (fun=="rgt")
			tok= RGT;
		else
		if (fun=="sin")
			tok= SIN;
		else
		if (fun=="cos")
			tok= COS;
		else
		if (fun=="atan")
			tok= ATAN;
		else
		if (fun=="exp")
			tok= EXP;
		else
		if (fun=="log")
			tok= LOG;
		else
		if (fun=="sqrt")
			tok= SQRT;
		else
		if (fun=="der")
			tok= DER;
		else
			tok= BAD;
	}
	else
		switch (c) {
			case '+':  tok= ADD;
					   break;
			case '-':  tok= SUB;
					   break;
			case '*':  tok= MUL;
					   break;
			case '/':  tok= DIV;
					   break;
			case '^':  tok= POW;
					   break;
			case '(':  tok= LPAR;
					   break;
			case ')':  tok= RPAR;
					   break;
			case EOF:
			case ';':  tok= SEMI;
					   break;
			default:   tok= BAD;
					   break;
		}
//	std::cerr << "tok: " << name[tok] << std::endl;
}

Node* Parser::param(void)
{
	Node* p;
	if (tok==LPAR) {
		Lex();
		p=  expr();
		if (tok==RPAR)
			Lex();
		else {
			error("se espera ')'");
			p= NULL;
		}
	}
	else {
		error("se espera ´('");
		p= NULL;
	}
	return p;
}

Node* Parser::prim(void)
{
	Node *t;
	Node *p;
	switch (tok) {
		case NUM:	p= new Const(num);
					Lex();
					break;
		case VART:  p= new Var();
					Lex();
					break;
		case LFT:	p= lft->copy();
					Lex();
					break;
		case RGT:	p= rgt->copy();
					Lex();
					break;
		case LPAR:	p= param();
					break;
		case SIN:	Lex();
					p= new Sin(param());
					break;
		case COS:	Lex();
					p= new Cos(param());
					break;
		case ATAN:	Lex();
					p= new Atan(param());
					break;
		case EXP:	Lex();
					p= new Exp(param());
					break;
		case LOG:	Lex();
					p= new Log(param());
					break;
		case SQRT:	Lex();
					p= new Sqrt(param());
					break;
		case DER:	Lex();
					t= param();
					p= t->deriv();
					delete t;
					break;
		default:	error("missing <prim>");
					p= NULL;
	}
	return p;
}

Node* Parser::factor(void)
{
	Node* f= prim();
	while (tok==POW) {
		Lex();
		f= new Pow(f, factor());
	}
	return f;
}

Node* Parser::term(void)
{
	Node* t= factor();
	for (;;) 
		switch (tok) {
			case MUL:
					Lex();
					t= new Mul(t, factor());
					break;
			case DIV:
					Lex();
					t= new Div(t, factor());
					break;
			default:
					return t;
		}
}

Node* Parser::expr(void)
{
	Node* e;
	if (tok==ADD) {
		Lex();
		e= term();
	}
	else
	if (tok==SUB) {
		Lex();
		e= new Chs(term());
	}
	else
		e= term();

	for (;;) 
		switch (tok) {
			case ADD:
					Lex();
					e= new Add(e, term());
					break;
			case SUB:
					Lex();
					e= new Sub(e, term());
					break;
			default:
					return e;
		}
}

Node* Parser::Tree(void) 
{
	Lex();
	Node* e= expr();
	if (tok!=SEMI) {
		error("falta ';'");
		e= NULL;
	}
	return e;
}


