
#include <cmath>
#include <iostream>
#include <cstdlib>
#include <string>
#include <sstream>

#include "node.hpp"
#include "parser.hpp"

//-----------------------
//   Const
//-----------------------

Const::Const(double v)
{
	val= v;
}

double Const::eval(double t) const
{
	return val;
}

Node* Const::copy(void) const
{
	return new Const(val);
}

Node* Const::deriv(void) const
{
	Parser p("0", NULL, NULL);
	return p.Tree();
}

string Const::print(bool par) const
{
	ostringstream out;
	out << val;
	if (par)
		return "(" + out.str() + ")";
	else
		return out.str();
}

//-----------------------
//   Var
//-----------------------

Var::Var(void)
{
}

double Var::eval(double t) const
{
	return t;
}

Node* Var::copy(void) const
{
	return new Var();
}

Node* Var::deriv(void) const
{
	Parser p("1", NULL, NULL);
	return p.Tree();
}

string Var::print(bool par) const
{
	if (par)
		return "(t)";
	else
		return "t";
}

//-----------------------
//   BinOp   
//-----------------------

BinOp::BinOp(Node *l, Node *r)
{
	lft= l;
	rgt= r;
}

BinOp::~BinOp(void)
{
	delete rgt;
	delete lft;
}

//-----------------------
//   Add
//-----------------------

Add::Add(Node* l, Node* r) : BinOp( l, r )
{
}

double Add::eval(double t) const
{
	return lft->eval(t) + rgt->eval(t);
}

Node* Add::copy(void) const
{
	return new Add(lft->copy(), rgt->copy());
}

Node* Add::deriv(void) const
{
	Parser p("der(lft) + der(rgt)", lft, rgt);
	return p.Tree();
}

string Add::print(bool par) const
{
	string wrt= lft->print(false) + "+" + rgt->print(rgt->type()==CHS);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Sub
//-----------------------

Sub::Sub(Node* l, Node* r) : BinOp( l, r )
{
}

double Sub::eval(double t) const
{
	return lft->eval(t) - rgt->eval(t);
}

Node* Sub::copy(void) const
{
	return new Sub(lft->copy(), rgt->copy());
}

Node* Sub::deriv(void) const
{
	Parser p("der(lft) - der(rgt)", lft, rgt);
	return p.Tree();
}

string Sub::print(bool par) const
{
	int rt= rgt->type();
	string wrt= lft->print(false)
				+ "-" 
				+ rgt->print(rt==ADD || rt==SUB || rt==CHS);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Mul
//-----------------------

Mul::Mul(Node* l, Node* r) : BinOp( l, r )
{
}

double Mul::eval(double t) const
{
	return lft->eval(t) * rgt->eval(t);
}

Node* Mul::copy(void) const
{
	return new Mul(lft->copy(), rgt->copy());
}

Node* Mul::deriv(void) const
{
	Parser p("der(lft) * rgt + lft * der(rgt) ", lft, rgt);
	return p.Tree();
}

string Mul::print(bool par) const
{
	int lt= lft->type();
	int rt= rgt->type();
	string wrt= lft->print(lt==ADD || lt==SUB) 
			+ "*" 
			+ rgt->print(rt==ADD || rt==SUB || rt==CHS);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Div
//-----------------------

Div::Div(Node* l, Node* r) : BinOp( l, r )
{
}

double Div::eval(double t) const
{
	return lft->eval(t) / rgt->eval(t);
}

Node* Div::copy(void) const
{
	return new Div(lft->copy(), rgt->copy());
}

Node* Div::deriv(void) const
{
	Parser p("(rgt*der(lft) - lft*der(rgt)) / rgt^2", lft, rgt);
	return p.Tree();
}

string Div::print(bool par) const
{
	int lt= lft->type();
	int rt= rgt->type();
	string wrt= lft->print(lt==ADD || lt==SUB ) 
			+ "/" 
			+ rgt->print(rt==ADD || rt==SUB || rt==CHS || rt==MUL || rt==DIV);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Pow
//-----------------------

Pow::Pow(Node* l, Node* r) : BinOp( l, r )
{
}

double Pow::eval(double t) const
{
	return pow(lft->eval(t) , rgt->eval(t));
}

Node* Pow::copy(void) const
{
	return new Pow(lft->copy(), rgt->copy());
}

Node* Pow::deriv(void) const
{
	Parser p("rgt*lft^(rgt-1)*der(lft) + lft^rgt*log(lft)*der(rgt)", 
															lft, rgt );
	return p.Tree();
}

string Pow::print(bool par) const
{
	int lt= lft->type();
	int rt= rgt->type();
	string wrt= lft->print(lt==ADD || lt==SUB || lt==CHS || lt==MUL || lt==DIV || lt==POW) 
			+ "^" 
			+ rgt->print(rt==ADD || rt==SUB || rt==CHS || rt==MUL || rt==DIV);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}


//-----------------------
//   Fun
//-----------------------

Fun::Fun(Node* r)
{
	rgt= r;
}

Fun::~Fun(void)
{
	delete rgt;
}

//-----------------------
//   Sin
//-----------------------

Sin::Sin(Node* r) : Fun( r )
{
}

double Sin::eval(double t) const
{
	return sin(rgt->eval(t));
}

Node* Sin::copy(void) const
{
	return new Sin(rgt->copy());
}

Node* Sin::deriv(void) const
{
	Parser p("cos(rgt)*der(rgt)", NULL, rgt );
	return p.Tree();
}

string Sin::print(bool par) const
{
	string wrt= "sin" + rgt->print(true);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Cos
//-----------------------

Cos::Cos(Node* r) : Fun( r )
{
}

double Cos::eval(double t) const
{
	return cos(rgt->eval(t));
}

Node* Cos::copy(void) const
{
	return new Cos(rgt->copy());
}

Node* Cos::deriv(void) const
{
	Parser p("-sin(rgt)*der(rgt)", NULL, rgt );
	return p.Tree();
}

string Cos::print(bool par) const
{
	string wrt= "cos" + rgt->print(true);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Exp
//-----------------------

Exp::Exp(Node* r) : Fun( r )
{
}

double Exp::eval(double t) const
{
	return exp(rgt->eval(t));
}

Node* Exp::copy(void) const
{
	return new Exp(rgt->copy());
}

Node* Exp::deriv(void) const
{
	Parser p("exp(rgt)*der(rgt)", NULL, rgt );
	return p.Tree();
}

string Exp::print(bool par) const
{
	string wrt= "exp" + rgt->print(true);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Log
//-----------------------

Log::Log(Node* r) : Fun( r )
{
}

double Log::eval(double t) const
{
	return log(rgt->eval(t));
}

Node* Log::copy(void) const
{
	return new Log(rgt->copy());
}

Node* Log::deriv(void) const
{
	Parser p("der(rgt)/rgt", NULL, rgt );
	return p.Tree();
}

string Log::print(bool par) const
{
	string wrt= "log" + rgt->print(true);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Chs
//-----------------------

Chs::Chs(Node* r) : Fun( r )
{
}

double Chs::eval(double t) const
{
	return - rgt->eval(t);
}

Node* Chs::copy(void) const
{
	return new Chs(rgt->copy());
}

Node* Chs::deriv(void) const
{
	Parser p("-der(rgt)", NULL, rgt );
	return p.Tree();
}

string Chs::print(bool par) const
{
	int rt= rgt->type();
	string wrt= "-" + rgt->print(rt==ADD || rt==SUB || rt==CHS);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Atan
//-----------------------

Atan::Atan(Node* r) : Fun( r )
{
}

double Atan::eval(double t) const
{
	return  atan(rgt->eval(t));
}

Node* Atan::copy(void) const
{
	return new Atan(rgt->copy());
}

Node* Atan::deriv(void) const
{
	Parser p("der(rgt)/(1+rgt^2)", NULL, rgt );
	return p.Tree();
}

string Atan::print(bool par) const
{
	string wrt= "atan" + rgt->print(true);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

//-----------------------
//   Sqrt
//-----------------------

Sqrt::Sqrt(Node* r) : Fun( r )
{
}

double Sqrt::eval(double t) const
{
	return  sqrt(rgt->eval(t));
}

Node* Sqrt::copy(void) const
{
	return new Sqrt(rgt->copy());
}

Node* Sqrt::deriv(void) const
{
	Parser p("der(rgt)/(2*sqrt(rgt))", NULL, rgt );
	return p.Tree();
}

string Sqrt::print(bool par) const
{
	string wrt= "sqrt" + rgt->print(true);
	if (par)
		return "(" + wrt + ")";
	else
		return wrt;
}

