#include "compilepass.h"
#include "clasplayer.h"
#include "llvmlayer.h"
#include <ast.h>

#include "query/containers.h"
#include "query/context.h"
#include "compilation/instr-containers.h"
#include "ExternLayer.h"
#include "pass/adhocpass.h"

#include <boost/optional.hpp>
#include <memory>
#include <iostream>

using namespace std;
using namespace xreate;
using namespace xreate::compilation;
using namespace llvm;

//SECTIONTAG types/convert implementation
llvm::Value*
doAutomaticTypeConversion(llvm::Value* source, llvm::Type* tyTarget, LLVMLayer* llvm){
	if (tyTarget->isIntegerTy() && source->getType()->isIntegerTy())
	{
		llvm::IntegerType* tyTargetInt = llvm::dyn_cast<IntegerType>(tyTarget);
		llvm::IntegerType* tySourceInt = llvm::dyn_cast<IntegerType>(source->getType());

		if (tyTargetInt->getBitWidth() < tySourceInt->getBitWidth()){
			return llvm->builder.CreateCast(llvm::Instruction::Trunc, source, tyTarget);
		}

		if (tyTargetInt->getBitWidth() > tySourceInt->getBitWidth()){
			return llvm->builder.CreateCast(llvm::Instruction::SExt, source, tyTarget);
		}
	}

	return source;
}


CodeScopeUnit::CodeScopeUnit(CodeScope* codeScope, FunctionUnit* f, CompilePass* compilePass)
    : scope(codeScope), pass(compilePass), function(f)
{}

namespace xreate {

class CallStatement {
public:
	virtual llvm::Value* operator() (std::vector<llvm::Value *>&& args, const std::string& hintDecl, llvm::IRBuilder<>&) = 0;
};

class CallStatementRaw: public CallStatement{
public:
	CallStatementRaw(llvm::Value* callee): __callee(callee) {}

	llvm::Value* operator() (std::vector<llvm::Value *>&& args, const std::string& hintDecl, llvm::IRBuilder<>& builder) {
		return builder.CreateCall(__callee, args, hintDecl);
	}

private:
	llvm::Value* __callee;
};

class CallStatementInline: public CallStatement{
public:
	CallStatementInline(FunctionUnit* caller, FunctionUnit* callee)
		: __caller(caller), __callee(callee) {}

	llvm::Value* operator() (std::vector<llvm::Value *>&& args, const std::string& hintDecl, llvm::IRBuilder<>& builder) {
		//TEST inlining
		return __callee->compileInline(move(args), __caller);
	}

private:
	FunctionUnit* __caller;
	FunctionUnit* __callee;
};

}

//SECTIONTAG late-context find callee function
//TEST static late context decisions
//TEST dynamic late context decisions
CallStatement*
CodeScopeUnit::findFunction(const std::string& calleeName){
		LLVMLayer* llvm = pass->man->llvm;
		ClaspLayer* clasp = pass->man->clasp;

		ContextQuery* queryContext = pass->queryContext;

		const std::list<ManagedFnPtr>& specializations = pass->man->root->getFunctionVariants(calleeName);

		//check external function
		if (!specializations.size()){
        	llvm::Function* external = pass->man->llvm->layerExtern->lookupFunction(calleeName);

        	return new CallStatementRaw(external);
		}

		//no decisions required
		if (specializations.size()==1){
			if (!specializations.front()->guardContext.isValid()) {
				return new CallStatementRaw( pass->getFunctionUnit(specializations.front())->compile());
			}
		}

		//prepare specializations dictionary
		typedef ExpressionSerialization<>::Serializer Serializer;
		typedef ExpressionSerialization<>::Code Code;

		auto adapter = [](const ManagedFnPtr& p){ return p->guardContext; };
		Serializer serializer(specializations, adapter);

		std::map<Code, ManagedFnPtr> dictSpecializations;
		boost::optional<ManagedFnPtr> variantDefault;
		boost::optional<ManagedFnPtr> variant;

		for(const ManagedFnPtr& f: specializations){
			const Expression& guard = f->guardContext;

			//default case:
			if (!guard.isValid()){
				variantDefault = f;
				continue;
			}

			assert(dictSpecializations.emplace(serializer.getId(guard), f).second && "Found several appropriate specializations");
		}

		//check static context
		const ScopeContextDecisions& decisions = queryContext->getStaticDecisions(clasp->pack(this->scope));
		if (decisions.count(calleeName)){
			variant =  dictSpecializations.at(serializer.getId(decisions.at(calleeName)));
		}

		size_t sizeDemand = this->function->contextCompiler.getFunctionDemandSize();

		//decision made if static context found or no late context exists(and there is default variant)
		bool flagHasStaticDecision =  variant || (variantDefault && !sizeDemand);

		//if no late context exists
		if (flagHasStaticDecision) {
			FunctionUnit* calleeUnit = pass->getFunctionUnit(variant? *variant: *variantDefault);

			//inlining possible based on static decision only
	        if (calleeUnit->isInline()) {
	            return new CallStatementInline(function, calleeUnit);
	        }

			return new CallStatementRaw(calleeUnit->compile());
		}

		//require default variant if no static decision made
		assert(variantDefault);

		llvm::Function* functionVariantDefault = this->pass->getFunctionUnit(*variantDefault)->compile();
		return new CallStatementRaw(this->function->contextCompiler.findFunction(calleeName, functionVariantDefault));
}

void
CodeScopeUnit::bindArg(llvm::Value* var, std::string&& name)
{
    assert(scope->__vartable.count(name));
    VID id = scope->__vartable.at(name);
    __rawVars[id] = var;
}

llvm::Value*
CodeScopeUnit::process(const Expression& expr, const std::string& hintVarDecl){
#define DEFAULT(x) (hintVarDecl.empty()? x: hintVarDecl)
    llvm::Value *left; llvm::Value *right;
    LLVMLayer& l = *pass->man->llvm;
    containers::Instructions instructions = containers::Instructions({function, this, pass});

    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 = process(expr.operands[0]);
            right = process(expr.operands[1]);

            //SECTIONTAG types/convert binary operation
           	right =	doAutomaticTypeConversion(right, left->getType(), &l);
            break;

        default:;
    }

    switch (expr.op) {
        case Operator::ADD:
            return l.builder.CreateAdd(left, right, DEFAULT("tmp_add"));
            break;

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

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

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

        case Operator::EQU:
            left->dump();
            right->dump();
            return l.builder.CreateICmpEQ(left, right, DEFAULT("tmp_equ"));
            break;

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

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

        case Operator::NEG:
            left = process(expr.operands[0]);
            return l.builder.CreateNeg(left, DEFAULT("tmp_neg"));
            break;

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

            std::string nameCallee = expr.getValueString();
            unique_ptr<CallStatement> callee(findFunction(nameCallee));

            //prepare arguments
            std::vector<llvm::Value *> args;
            args.reserve(expr.operands.size());

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

            ScopePacked outerScopeId = pass->man->clasp->pack(this->scope);

            //SECTIONTAG late-context propagation arg
            size_t calleeDemandSize = pass->queryContext->getFunctionDemand(nameCallee).size();
            if (calleeDemandSize){
            	llvm::Value* argLateContext = function->contextCompiler.compileArgument(nameCallee, outerScopeId);
            	args.push_back(argLateContext);
            }

            return (*callee)(move(args), DEFAULT("res_"+nameCallee), l.builder);
        }

        case Operator::IF:
        {
        	return instructions.compileIf(expr, DEFAULT("tmp_if"));
        }

        case Operator::SWITCH:
        {
        	return instructions.compileSwitch(expr, DEFAULT("tmp_switch"));
        }

        case Operator::LOOP_CONTEXT:
        {
        	return instructions.compileLoopContext(expr, DEFAULT("tmp_loop"));
        }

        case Operator::LOGIC_AND: {
        	assert(expr.operands.size() == 1);
        	return process (expr.operands[0]);
        }

        case Operator::LIST:
        {
           return instructions.compileConstantArray(expr, DEFAULT("tmp_list"));
        };

        case Operator::LIST_RANGE:
        {
            assert(false); //no compilation phase for a range list
          //  return InstructionList(this).compileConstantArray(expr, l, hintRetVar);
        };

        case Operator::LIST_NAMED:
        {
            typedef Expanded<TypeAnnotation> ExpandedType;

            ExpandedType tyRaw = l.ast->expandType(expr.type);

            const std::vector<string> fields = (tyRaw.get().__operator == TypeOperator::CUSTOM)?
                l.layerExtern->getStructFields(l.layerExtern->lookupType(tyRaw.get().__valueCustom))
                : tyRaw.get().fields;

            std::map<std::string, size_t> indexFields;
            for(size_t i=0, size = fields.size(); i<size; ++i){
                indexFields.emplace(fields[i], i);
            }

            llvm::StructType* tyRecord = llvm::cast<llvm::StructType>(l.toLLVMType(tyRaw));
            llvm::Value* record = llvm::UndefValue::get(tyRecord);

            for (size_t i=0; i<expr.operands.size(); ++i){
                const Expression& operand = expr.operands.at(i);
                unsigned int fieldId = indexFields.at(expr.bindings.at(i));

                llvm::Value* result = 0;

//TODO Null ad hoc Llvm implementation
//                if (operand.isNone()){
//                    llvm::Type* tyNullField = tyRecord->getElementType(fieldId);
//                    result = llvm::UndefValue::get(tyNullField);
//
//                } else {
                    result = process(operand);
//                }

                assert (result);
                record = l.builder.CreateInsertValue(record, result, llvm::ArrayRef<unsigned>({fieldId}));
            }

            return record;
        };



        case Operator::MAP:
        {
            assert(expr.blocks.size());
            return instructions.compileMapSolid(expr, DEFAULT("map"));
        };

        case Operator::FOLD:
        {
            return instructions.compileFold(expr, DEFAULT("fold"));
        };

        case Operator::INDEX:
        {
                //TODO allow multiindex
            assert(expr.operands.size()==1);
            const std::string &ident = expr.getValueString();
            Symbol s = scope->findSymbol(ident);
            const TypeAnnotation& t = s.scope->findDefinition(s);
            const ExpandedType& t2 = pass->man->root->expandType(t);

            switch (t2.get().__operator)
            {
                case TypeOperator::STRUCT: case TypeOperator::CUSTOM:
                {
                    Expression idx = expr.operands.at(0);
                    assert(idx.__state == Expression::STRING);
                    std::string idxField = idx.getValueString();

                    llvm::Value* aggr  = compileSymbol(s, ident);
                    return instructions.compileStructIndex(aggr, t2, idxField);
                };

                case TypeOperator::ARRAY: {
                    std::vector<llvm::Value*> indexes;
                    std::transform(++expr.operands.begin(), expr.operands.end(), std::inserter(indexes, indexes.end()),
                                   [this] (const Expression& op){return process(op);}
                    );

                    return instructions.compileArrayIndex(s, indexes, DEFAULT(string("el_") + ident));
                };

                default:
                    assert(false);
            }
        };

        	//SECTIONTAG adhoc actual compilation
        case Operator::ADHOC: {
        	assert(function->adhocImplementation && "Adhoc implementation not found");
        	string comm = expr.operands[0].getValueString();

        	CodeScope* scope = function->adhocImplementation->getImplementationForCommand(comm);
        	CodeScopeUnit* unitScope = function->getScopeUnit(scope);
        	return unitScope->compile();
        };

        case Operator::SEQUENCE: {
        	assert (expr.getOperands().size());

        	llvm::Value* result;
        	for(const Expression &op: expr.getOperands()){
        		result = process(op, "");
        	}

			return result;
        }

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

            switch (expr.__state) {
                case Expression::IDENT: {
                    const std::string &ident = expr.getValueString();
                    Symbol s = scope->findSymbol(ident);
                    return compileSymbol(s, ident);
                }

                case Expression::NUMBER: {
                    int literal = expr.getValueDouble();
                    return llvm::ConstantInt::get(llvm::Type::getInt32Ty(llvm::getGlobalContext()), literal);
                }

                case Expression::STRING: {
                    return instructions.compileConstantStringAsPChar(expr.getValueString(), DEFAULT("tmp_str"));
                };

                case Expression::VARIANT: {
                	const ExpandedType& typVariant = pass->man->root->expandType(expr.type);
                	llvm::Type* typRaw = l.toLLVMType(typVariant);
                	int value = expr.getValueDouble();
                    return llvm::ConstantInt::get(typRaw, value);
                }

                default: {
                    break;
                }
            };

            break;

        default: break;

    }

    assert(false);
    return 0;
}

llvm::Value*
CodeScopeUnit::compile(const std::string& hintBlockDecl){
    if (raw != nullptr) return raw;


    if (!hintBlockDecl.empty()) {
        llvm::BasicBlock *block = llvm::BasicBlock::Create(llvm::getGlobalContext(), hintBlockDecl, function->raw);
        pass->man->llvm->builder.SetInsertPoint(block);
    }

    raw = process(scope->__body);
    return raw;
}

llvm::Value*
CodeScopeUnit::compileSymbol(const Symbol& s, std::string hintRetVar)
{
    CodeScope* scope = s.scope;
    CodeScopeUnit* self = function->getScopeUnit(scope);

    if (self->__rawVars.count(s.identifier))     {
        return self->__rawVars[s.identifier];
    }

    return self->__rawVars[s.identifier] = self->process(scope->findDeclaration(s), hintRetVar);
}

bool
FunctionUnit::isInline(){
    Symbol ret = Symbol{0, function->__entry};
    bool flagOnTheFly = SymbolAttachments::get<IsImplementationOnTheFly>(ret, false);

    return flagOnTheFly;
}

llvm::Function*
FunctionUnit::compile(){
    if (raw != nullptr) return raw;

    std::vector<llvm::Type *> types;
    LLVMLayer* llvm = pass->man->llvm;
    llvm::IRBuilder<>& builder = llvm->builder;
    CodeScope* entry = function->__entry;
    AST* ast = pass->man->root;

    const string& functionName = ast->getFunctionVariants(function->__name).size() > 1? function->__name  + std::to_string(function.id()) : function->__name;

    std::transform(entry->__args.begin(), entry->__args.end(), std::inserter(types, types.end()),
            [this, llvm, ast, entry](const std::string &arg)->llvm::Type* {
                assert(entry->__vartable.count(arg));
                VID argid = entry->__vartable.at(arg);
                assert(entry->__definitions.count(argid));
                return llvm->toLLVMType(ast->expandType(entry->__definitions.at(argid)));
            });

    	//SECTIONTAG late-context signature type
    size_t sizeLateContextDemand = contextCompiler.getFunctionDemandSize();
    if (sizeLateContextDemand) {
    	llvm::Type* ty32 = llvm::Type::getInt32Ty(llvm::getGlobalContext());
    	llvm::Type* tyDemand = llvm::ArrayType::get(ty32, sizeLateContextDemand);
    	types.push_back(tyDemand);
    }

    	//SECTIONTAG adhoc func signature determination
    llvm::Type* expectedResultType;
    if (function->isPrefunction){
    	AdhocPass* adhocpass = reinterpret_cast<AdhocPass*>(pass->man->getPassById(PassId::AdhocPass));

    	adhocImplementation = adhocpass->determineForScope(entry);
    	expectedResultType = llvm->toLLVMType(ast->expandType(adhocImplementation->getResultType()));
    } else {
    	expectedResultType = llvm->toLLVMType(ast->expandType(entry->__definitions[0]));
    }

    llvm::FunctionType *ft = llvm::FunctionType::get(expectedResultType, types, false);

    raw = llvm::cast<llvm::Function>(llvm->module->getOrInsertFunction(functionName, ft));

    CodeScopeUnit* entryCompilation = getScopeUnit(entry);
    llvm::Function::arg_iterator fargsI = raw->arg_begin();
    for (std::string &arg : entry->__args) {
        VID argid = entry->__vartable[arg];

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

    if (sizeLateContextDemand){
    	fargsI->setName("latecontext");
    	contextCompiler.raw = fargsI;
    }

    const std::string&blockName =  "entry";
    llvm::BasicBlock* blockCurrent = builder.GetInsertBlock();

    llvm::Value* result = entryCompilation->compile(blockName);
    assert(result);

    //SECTIONTAG types/convert function ret value
    builder.CreateRet(doAutomaticTypeConversion(result, expectedResultType, llvm));

    if (blockCurrent){
    	builder.SetInsertPoint(blockCurrent);
    }

    llvm->moveToGarbage(ft);
    return raw;
}

llvm::Value*
FunctionUnit::compileInline(std::vector<llvm::Value *> &&args, FunctionUnit* outer){
	CodeScopeUnit* entryCompilation = outer->getScopeUnit(function->__entry);
    for(int i=0, size = args.size(); i<size; ++i) {
        entryCompilation->bindArg(args.at(i), string(entryCompilation->scope->__args.at(i)));
    }


    return entryCompilation->compile();
}

CodeScopeUnit*
FunctionUnit::getScopeUnit(CodeScope* scope){
    if (!scopes.count(scope)){
    	CodeScopeUnit* unit = new CodeScopeUnit(scope, this, pass);
        scopes.emplace(scope, std::unique_ptr<CodeScopeUnit>(unit));
    }

    return scopes.at(scope).get();
}

CodeScopeUnit*
FunctionUnit::getEntry(){
	return getScopeUnit(function->getEntryScope());
}

CodeScopeUnit*
FunctionUnit::getScopeUnit(ManagedScpPtr scope){
    return getScopeUnit(&*scope);
}

FunctionUnit*
CompilePass::getFunctionUnit(const ManagedFnPtr& function){
	unsigned int id = function.id();

    if (!functions.count(id)){
    	FunctionUnit* unit = new FunctionUnit(function, this);
        functions.emplace(id, std::unique_ptr<FunctionUnit>(unit));
        return unit;
    }

    return functions.at(id).get();
}

void
CompilePass::run(){
	queryContext = reinterpret_cast<ContextQuery*> (man->clasp->getQuery(QueryId::ContextQuery));

    //Find out main function;
    ClaspLayer::ModelFragment model = man->clasp->query(Config::get("function-entry"));
    assert(model && "Error: No entry function found");
    assert(model->first != model->second && "Error: Ambiguous entry function");

    string nameMain = std::get<0>(ClaspLayer::parse<std::string>(model->first->second));
    FunctionUnit* unitMain = getFunctionUnit(man->root->findFunction(nameMain));
    entry = unitMain->compile();
}

llvm::Function*
CompilePass::getEntryFunction(){
	assert(entry);
	return entry;
}

void
CompilePass::prepareQueries(ClaspLayer* clasp){
	clasp->registerQuery(new containers::Query(), QueryId::ContainersQuery);
	clasp->registerQuery(new ContextQuery(), QueryId::ContextQuery);
}

//CODESCOPE COMPILATION PHASE


//FIND SYMBOL(compilation phase):
    //if (!forceCompile)
    //{
    //    return result;
    //}

    //    //search in already compiled vars
    //if (__rawVars.count(vId))
    //{
    //    return result;
    //}

    //if (!__declarations.count(vId)) {
    //    //error: symbol is uncompiled scope arg
    //    assert(false);
    //}

    //const Expression& e = __declarations.at(vId);

    //__rawVars[vId] = process(e, l, name);


//FIND FUNCTION
    //llvm::Function*
    //CompilePass::findFunction(const std::string& name){
    //    ManagedFnPtr calleeFunc = man->root->findFunction(name);
    //    assert(calleeFunc.isValid());

    //    return  nullptr;
    //}
