#include "instr-containers.h"
#include "llvmlayer.h"
#include "ast.h"
#include "query/containers.h"

using namespace std;
using namespace llvm;
using namespace xreate::containers;

xreate::containers::Instructions::Instructions(CodeScope *current, LLVMLayer* layer)
    : scope(current), llvm(layer),
      tyNum (static_cast<llvm::IntegerType*> (TypeAnnotation(TypePrimitive::Num).toLLVMType()))
{}

Value*
Instructions::compileIndex(const Symbol& dataSymbol, std::vector<llvm::Value*> indexes, std::string ident)
{
    #define NAME(x) (ident.empty()? x : ident)

    const ImplementationData& info = containers::Query::queryImplementation(dataSymbol);
    llvm::Value* data = dataSymbol.scope->__rawVars.at(dataSymbol.identifier);

    switch (info.impl)
    {
        case containers::LLVM_CONST_ARRAY: {
            assert(indexes.size() == 1);
            return llvm->builder.CreateExtractElement(data, indexes[0], NAME("el"));
        }

        case containers::LLVM_ARRAY: {
            indexes.insert(indexes.begin(), llvm::ConstantInt::get(tyNum, 0));
            Value *pEl = llvm->builder.CreateGEP(data, llvm::ArrayRef<llvm::Value *>(indexes));
            return llvm->builder.CreateLoad(pEl, NAME("el"));
        }

        default:
            // non supported container implementation
            assert(false);
    }
}

Value*
Instructions::compileGetElement(std::string varIn, Value* stateLoop, std::string ident)
{
    #define NAME(x) (ident.empty()? x : ident)
    llvm::IRBuilder<> &builder = llvm->builder;

    const Symbol& symbolIn = scope->findSymbol(varIn, *llvm);
    const Expression& exprIn = CodeScope::findDeclaration(symbolIn);

    if (exprIn.op == Operator::LIST_RANGE)
        {return stateLoop; }

    if (exprIn.op == Operator::LIST)
    {
        llvm::Value* dataIn = CodeScope::compileExpression(symbolIn, *llvm);
        return compileIndex(symbolIn, vector<Value*>{stateLoop}, NAME(string("el_") + varIn));
    }

    if(exprIn.op == Operator::MAP)
    {
        assert(exprIn.getOperands().size()==1);
        assert(exprIn.bindings.size());
        assert(exprIn.blocks.size());

        const ManagedScpPtr& scopeLoop = exprIn.blocks.front();
        const std::string& varIn2 = exprIn.getOperands()[0].getValueString();
        std::string varEl = exprIn.bindings[0];

        Value* elIn = compileGetElement(varIn2, stateLoop, varEl);
        scopeLoop->bindArg(elIn, move(varEl));
        return scopeLoop->compileExpression(scopeLoop->__body, *llvm, NAME(string("el_") + varIn));
    }

    assert(false);



    /*
     switch(implIn.impl)
    {
        case containers::LLVM_ARRAY:
        case containers::LLVM_CONST_ARRAY:
            stateLoopNext = builder.CreateAdd(stateLoop, ConstantInt::get(tyNum, 1));


        case containers::ON_THE_FLY:
            stateLoopNext = compileIndexLoad(stateLoad);
    }
    */

}

llvm::Value*
Instructions::compileMapArray(const Expression &expr, const std::string ident) {
        #define NAME(x) (ident.empty()? x : ident)

                //initialization
        std::string varIn = expr.getOperands()[0].getValueString();
        Symbol symbolIn = scope->findSymbol(varIn, *llvm);

        containers::ImplementationData implIn = containers::Query::queryImplementation(symbolIn); // impl of input list
        unsigned int size = implIn.size;
        ManagedScpPtr scopeLoop =  expr.blocks.front();
        std::string varEl = scopeLoop->__args[0];

        llvm::Value *rangeFrom = ConstantInt::get((tyNum), 0);
        llvm::Value *rangeTo = ConstantInt::get((tyNum), size-1);

                //definitions
    ArrayType* tyNumArray = (ArrayType*) (TypeAnnotation(tag_array, TypePrimitive::Num, size).toLLVMType());
    llvm::IRBuilder<> &builder = llvm->builder;

    llvm::BasicBlock *blockLoop = llvm::BasicBlock::Create(llvm::getGlobalContext(), "loop", llvm->context.function);
    llvm::BasicBlock *blockBeforeLoop = builder.GetInsertBlock();
    llvm::BasicBlock *blockAfterLoop = llvm::BasicBlock::Create(llvm::getGlobalContext(), "postloop", llvm->context.function);
    Value* dataOut = llvm->builder.CreateAlloca(tyNumArray, ConstantInt::get(tyNum, size), NAME("map"));

                // * initial check
    Value* condBefore = builder.CreateICmpSLE(rangeFrom, rangeTo);
    builder.CreateCondBr(condBefore, blockLoop, blockAfterLoop);

                // create PHI:
    builder.SetInsertPoint(blockLoop);
    llvm::PHINode *stateLoop = builder.CreatePHI(tyNum, 2, "mapIt");
    stateLoop->addIncoming(rangeFrom, blockBeforeLoop);

                // loop body:
    Value* elIn = compileGetElement(varIn, stateLoop, varEl);
    scopeLoop->bindArg(elIn, move(varEl));
    Value* elOut = scopeLoop->compileExpression(scopeLoop->__body, *llvm);
    Value *pElOut = builder.CreateGEP(dataOut, ArrayRef<Value *>(std::vector<Value*>{ConstantInt::get(tyNum, 0), stateLoop}));
    builder.CreateStore(elOut, pElOut);

                //next iteration preparing
    Value *stateLoopNext = builder.CreateAdd(stateLoop,llvm::ConstantInt::get(tyNum, 1));
    stateLoop->addIncoming(stateLoopNext, blockLoop);

                //next iteration checks:
    Value* condAfter = builder.CreateICmpSLE(stateLoopNext, rangeTo);
    builder.CreateCondBr(condAfter, blockLoop, blockAfterLoop);

                //finalization:
    builder.SetInsertPoint(blockAfterLoop);

    return dataOut;
}

llvm::Value*
Instructions::compileFold(const Expression& fold, const std::string& ident)
{
    #define NAME(x) (ident.empty()? x : ident)
    assert(fold.op == Operator::FOLD);

            //initialization:
    Symbol varInSymbol  = scope->findSymbol(fold.getOperands()[0].getValueString(), *llvm);
    ImplementationData info = Query::queryImplementation(varInSymbol);
    const pair<Expression, Expression>& range = info.getRange();
    llvm::Value* rangeFrom = scope->compileExpression(range.first, *llvm);
    llvm::Value* rangeTo = scope->compileExpression(range.second, *llvm);
    llvm::Value* accumInit = scope->compileExpression(fold.getOperands()[1], *llvm);
    std::string varIn = fold.getOperands()[0].getValueString();
    std::string varAccum = fold.bindings[1];
    std::string varEl = fold.bindings[0];
    llvm::Value* valSat;
    bool flagHasSaturation = false; //false; // TODO add `saturation` ann.

    llvm::BasicBlock *blockBeforeLoop = llvm->builder.GetInsertBlock();
    llvm::BasicBlock *blockLoop = llvm::BasicBlock::Create(llvm::getGlobalContext(), "fold", llvm->context.function);
    llvm::BasicBlock *blockAfterLoop = llvm::BasicBlock::Create(llvm::getGlobalContext(), "postfold", llvm->context.function);

            // * initial check
    Value* condBefore = llvm->builder.CreateICmpSLE(rangeFrom, rangeTo);
    llvm->builder.CreateCondBr(condBefore, blockLoop, blockAfterLoop);

    if (flagHasSaturation)
    {
        Value* condSat = llvm->builder.CreateICmpNE(accumInit, valSat);
        llvm->builder.CreateCondBr(condSat, blockLoop, blockAfterLoop);
    }

            // * create phi
    llvm->builder.SetInsertPoint(blockLoop);
    llvm::PHINode *accum = llvm->builder.CreatePHI(tyNum, 2, NAME("accum"));
    accum->addIncoming(accumInit, blockBeforeLoop);

    llvm::PHINode *stateLoop = llvm->builder.CreatePHI(tyNum, 2, "foldIt");
    stateLoop->addIncoming(rangeFrom, blockBeforeLoop);

            // * loop body
    ManagedScpPtr scopeLoop = fold.blocks.front();
    Value* elIn = compileGetElement(varIn, stateLoop);
    scopeLoop->bindArg(accum, move(varAccum));
    scopeLoop->bindArg(elIn, move(varEl));
    Value* accumNext = scopeLoop->compileExpression(scopeLoop->__body, *llvm);

            // * break checks, continue checks
    if (flagHasSaturation)
    {
        llvm::BasicBlock *blockChecks = llvm::BasicBlock::Create(llvm::getGlobalContext(), "checks", llvm->context.function);
        Value* condSat = llvm->builder.CreateICmpNE(accumNext, valSat);
        llvm->builder.CreateCondBr(condSat, blockChecks, blockAfterLoop);
        llvm->builder.SetInsertPoint(blockChecks);
    }
            // * computing next iteration state
    Value *stateLoopNext = llvm->builder.CreateAdd(stateLoop, llvm::ConstantInt::get(tyNum, 1));
    accum->addIncoming(accumNext,  llvm->builder.GetInsertBlock());
    stateLoop->addIncoming(stateLoopNext, llvm->builder.GetInsertBlock());

             // * next iteration checks
    Value* condAfter = llvm->builder.CreateICmpSLE(stateLoopNext, rangeTo);
    llvm->builder.CreateCondBr(condAfter, blockLoop, blockAfterLoop);

            // finalization:
    llvm->builder.SetInsertPoint(blockAfterLoop);

    return accum;
}

llvm::Value*
Instructions::compileIf(const Expression& exprIf, const std::string& ident)
{
    #define NAME(x) (ident.empty()? x : ident)

           //initialization:
    const Expression& condExpr = exprIf.getOperands()[0];
    llvm::IRBuilder<>& builder = llvm->builder;

    llvm::BasicBlock *blockAfter = llvm::BasicBlock::Create(llvm::getGlobalContext(), "ifAfter", llvm->context.function);
    llvm::BasicBlock *blockTrue = llvm::BasicBlock::Create(llvm::getGlobalContext(), "ifTrue", llvm->context.function);
    llvm::BasicBlock *blockFalse = llvm::BasicBlock::Create(llvm::getGlobalContext(), "ifFalse", llvm->context.function);

    llvm::Value* cond = scope->compileExpression(condExpr, *llvm);
    llvm->builder.CreateCondBr(cond, blockTrue, blockFalse);

    builder.SetInsertPoint(blockTrue);
    ManagedScpPtr scopeTrue = exprIf.blocks.front();
    llvm::Value* resultTrue = scopeTrue->compileExpression(scopeTrue->__body, *llvm);
    builder.CreateBr(blockAfter);

    builder.SetInsertPoint(blockFalse);
    ManagedScpPtr scopeFalse = exprIf.blocks.back();
    llvm::Value* resultFalse = scopeFalse->compileExpression(scopeFalse->__body, *llvm);
    builder.CreateBr(blockAfter);

    builder.SetInsertPoint(blockAfter);
    llvm::PHINode *ret =  builder.CreatePHI(tyNum, 2, NAME("if"));
    ret->addIncoming(resultTrue, blockTrue);
    ret->addIncoming(resultFalse, blockFalse);

    return ret;
}

llvm::Value*
Instructions::compileConstantArray(const Expression &expr, const std::string& hintRetVar) {
    const int& __size = expr.getOperands().size();
    const Expression& __data = expr;

    ArrayType* typList = (ArrayType*) (TypeAnnotation(tag_array, TypePrimitive::i32, __size).toLLVMType());
    Type*typI32 = TypeAnnotation(TypePrimitive::i32).toLLVMType();

    std::vector<Constant *> list;
    list.reserve(__size);

    const std::vector<Expression> operands = __data.getOperands();
    std::transform(operands.begin(), operands.end(), std::inserter(list, list.begin()),
        [typI32](const Expression& e){return ConstantInt::get(typI32, e.getValueDouble());});

    Value* listSource = ConstantArray::get(typList, ArrayRef<Constant*>(list));
        /*
    Value* listDest = l.builder.CreateAlloca(typList, ConstantInt::get(typI32, __size), *hintRetVar);
    l.buil1der.CreateMemCpy(listDest, listSource, __size, 16);
        */

    return listSource;
}