#include <ast.h>

llvm::Value*
Function::compileExpression(const Expression& expr, LLVMLayer& l, std::string* hintRetVar=0)
{
    std::string var;
    if (hintRetVar && this->__vartable.count(*hintRetVar))
    {
        var = *hintRetVar;
    }

    #define VARNAME(x) (var.size()? var : x)
    llvm::Value* left; llvm::Value* right;


    switch (expr.__op)
    {
        case Operator::ADD: case Operator::SUB: case Operator::MUL: case Operator::DIV: case Operator::EQU: case Operator::LSS: case Operator::GTR:
            assert(expr.__state == Expression::COMPOUND);
            assert(expr.operands.size() == 2);

            left = compileExpression(expr.operands[0], l);
            right = compileExpression(expr.operands[1], l);
            break;

        default: ;
    }

    switch (expr.__op)
    {
        case Operator::ADD:
        return l.builder.CreateFAdd(left, right, VARNAME("tmp_add"));
        break;

        case Operator::SUB:
        return  l.builder.CreateFSub(left, right, VARNAME("tmp_sub"));
        break;

        case Operator::MUL:
        return  l.builder.CreateFMul(left, right, VARNAME("tmp_mul"));
        break;

        case Operator::DIV:
        return  l.builder.CreateFDiv(left, right, VARNAME("tmp_div"));
        break;

        case Operator::EQU:
        return  l.builder.CreateFCmpOEQ(left, right, VARNAME("tmp_equ"));
        break;

        case Operator::LSS:
        return  l.builder.CreateFCmpOLT(left, right, VARNAME("tmp_lss"));
        break;

        case Operator::GTR:
        return l.builder.CreateFCmpOGT(left, right, VARNAME("tmp_gtr"));
        break;

        case Operator::NEG:
        left = compileExpression(expr.operands[0], l);
        return  l.builder.CreateFNeg(left, VARNAME("tmp_neg"));
        break;

        case Operator::CALL:
    {
            assert(expr.__state == Expression::COMPOUND);

            const std::string& fname = expr.__valueS;
            assert(root->__indexFunctions.count(fname));

            const Function& calleeFunc = root->getFunctionById(root->__indexFunctions[fname]);
            llvm::Function* callee = calleeFunc.__raw;

            std::vector<llvm::Value*> args;
            args.reserve(expr.operands.size()-1);

            std::transform(expr.operands.begin(), expr.operands.end(), std::inserter(args, args.end()),
               [&l, this](const Expression &operand){return compileExpression(operand, l); }
            );

            return  l.builder.CreateCall(callee, args, VARNAME("tmp_call"));
    }
        break;


        case Operator::NONE:
            assert(expr.__state != Expression::COMPOUND);

            switch (expr.__state)
            {
                case Expression::IDENT:
            {
                std::string vname = expr.__valueS;
                assert(__vartable.count(vname));

                VID vId  =__vartable.at(vname);
                if (__rawVars.count(vId))
                {
                    return this->__rawVars.at(vId);
                }

                Expression& e = __declarations[vId];
                llvm::Value* result = compileExpression(e, l, &vname);
                __rawVars[vId] = result;
                return result;
            break;
            }

            case Expression::NUMBER:
                double literal = expr.__valueD;
                return llvm::ConstantFP::get(llvm::getGlobalContext(), llvm::APFloat(literal));
            };

        break;

    }

    assert(false);
    return 0;
}

llvm::Function*
Function::compile(LLVMLayer &l)
{

    std::vector<llvm::Type*> types;
    std::transform(__args.begin(), __args.end(), std::inserter(types, types.end()),
         [this](std::string& arg)     {
           assert(__vartable.count(arg));
           VID argid = __vartable[arg];

           assert(__definitions.count(argid));
           return __definitions[argid].toLLVMType();
    } );

    llvm::FunctionType* ft  = llvm::FunctionType::get(__retType.toLLVMType(), types,false);
    __raw = llvm::cast<llvm::Function>(l.module->getOrInsertFunction(__name,ft));

    llvm::Function::arg_iterator fargsI = __raw->arg_begin();
    for(std::string& arg : __args)
    {
     VID argid = __vartable[arg];

     __rawVars[argid] = fargsI;
     fargsI->setName(arg);
     ++fargsI;
    }

    llvm::BasicBlock *block = llvm::BasicBlock::Create(llvm::getGlobalContext(), "entry", __raw);
    l.builder.SetInsertPoint(block);

    l.builder.CreateRet(compileExpression(__body, l, 0));
    l.moveToGarbage(ft);

    return __raw;
};

void
AST::compile(LLVMLayer &layer)
{
    layer.module = new llvm::Module(getModuleName(),llvm::getGlobalContext());

    for(Function& f: __functions)
    {
        llvm::Function* rawf = f.compile(layer);

    }
}
