#include "GetdpResolutionFns.hpp"
#include <sstream>
#include "easylogging++/easylogging++.h"
#include "cenos_exception.h"


std::string resolutionHeader(std::vector<std::string> solverKeys, std::vector<Physics::timeType> timeTypeIds, std::vector<std::string> initResolution,
	std::vector<std::string> preResolKeys, std::vector<Physics::timeType> preResolTypeIds )
{
	std::string multiPort = "";
	std::string outString = "Resolution \n"
		"{\n";

	for (auto ir : initResolution)
	{
		std::stringstream initStringStream;
		initStringStream << "{ Name Initialization ;\n";
		initStringStream << "  System {\n";
		initStringStream << "    { Name "<< ir << "_init ; NameOfFormulation "<< ir <<"Initialization ; DestinationSystem resol_" << ir <<"; }\n";
		initStringStream << "  }\n";
		initStringStream << "  Operation {\n";
		if (ir == "getdpThermal")
			initStringStream << "    GmshRead[\"../getdpresults/temp.pos\"];\n";
		initStringStream << "   Generate[" << ir << "_init]; Solve[" << ir << "_init]; TransferSolution[" << ir << "_init]; \n";
		initStringStream << "  }\n";
		initStringStream << "}\n";
		outString.append(initStringStream.str());
	}


	std::string headerString = "	{\n"
		"		Name analysis; \n"
		"		System \n"
		"		{\n";
	outString.append(headerString);

	for (unsigned i = 0; i < timeTypeIds.size(); i++)
	{
		if (solverKeys[i] == "getdpMicrowaves") {
			multiPort = "~{i}";
		}
		if (solverKeys[i] == "getdpMicrowaves") {

			outString.append("  		For i In{ 1:#Sur_UNIFPORT_N_N_List() }\n");
			outString.append("			{ Name resol_" + solverKeys[i] + multiPort + "; NameOfFormulation " + solverKeys[i] + "_formulation~{i}" + "; ");

		}
		else {
			outString.append("			{ Name resol_" + solverKeys[i] + "; NameOfFormulation " + solverKeys[i] + "_formulation" + "; ");
		}



		if ((timeTypeIds[i] == Physics::HARMONIC or timeTypeIds[i] == Physics::HARMONIC_WITH_TIME or timeTypeIds[i] == Physics::MULTIHARMONIC))
		{
			outString.append("Type ComplexValue; Frequency Freq;} \n");
			if (solverKeys[i] == "getdpMicrowaves") {
				outString.append("  		EndFor\n");
			}
			
		}
		else {
			outString.append("} \n");
		}

	}

	for (unsigned i = 0; i < preResolTypeIds.size(); i++)
	{
		outString.append("			{ Name resol_" + preResolKeys[i] + "; NameOfFormulation " + preResolKeys[i] + "_formulation" + "; ");

		if ((preResolTypeIds[i] == Physics::HARMONIC or timeTypeIds[i] == Physics::HARMONIC_WITH_TIME or preResolTypeIds[i] == Physics::MULTIHARMONIC))
		{
			outString.append("Type ComplexValue; Frequency Freq;");
		}
		outString.append("} \n");
	}


	outString.append("		} \n");
	return outString;
}




std::string initilizationOperations(std::vector<std::map<std::string, InputValue>> timePars, std::vector<Physics::timeType> timeTypeIds, std::vector<std::string> solverKeys, std::vector<std::string> preResolKeys)
{
	std::string multiPort = "";
	std::string outString = "		Operation \n";
	outString.append("		{ \n");

	outString.append("			DeleteFile[\"../getdpresults/times.dat\"]; \n");
	//initialization
	bool hasTransient = false;
	int transientIndex = 0;
	for (unsigned i = 0; i < timeTypeIds.size(); i++)
	{
		if (timeTypeIds[i] == Physics::TRANSIENT)
		{
			hasTransient = true;
			transientIndex = i;
		}
	}
	for (auto prk : preResolKeys)
	{
		outString.append("			InitSolution[resol_" + prk + "] ; \n");
		outString.append("			Generate[resol_" + prk + "] ; \n");
		outString.append("			Solve[resol_" + prk + "] ; \n");
		outString.append("			SaveSolution[resol_" + prk + "] ; \n");
	}



	for (unsigned i = 0; i < timeTypeIds.size(); i++)
	{
		if (solverKeys[i] == "getdpMicrowaves") {
			multiPort = "~{i}";
			outString.append("			For i In {1:#Sur_UNIFPORT_N_N_List()}\n");
		}
		if ( (timeTypeIds[i] == Physics::HARMONIC or timeTypeIds[i] == Physics::HARMONIC_WITH_TIME or timeTypeIds[i] == Physics::MULTIHARMONIC) )
		{
			
			outString.append("			InitSolution[resol_" + solverKeys[i] + multiPort + "] ; \n");
			if (hasTransient)
			{
				if (timePars[transientIndex]["tstart"].getValue() == timePars[transientIndex]["tend"].getValue())
					outString.append("			SaveSolution[resol_" + solverKeys[i] + "] ; \n");
			}
		}
		else
		{
			outString.append("			InitSolution[resol_" + solverKeys[i] + "] ; \n");
			if (timeTypeIds[i] != Physics::STEADY)
			{
				outString.append("			SaveSolution[resol_" + solverKeys[i] + "] ; \n");

				if (timePars[i].count("continueFromLast"))
				{
					if (std::stof(timePars[i]["continueFromLast"].getValue()) == 1)
					{
						outString.append("			SetTime[ " + timePars[i]["trestart"].getValue() + "]; \n");
					}
					else
					{
						outString.append("			SetTime[ " + timePars[i]["tstart"].getValue() + "]; \n");
					}
				}
				else
				{
					outString.append("			SetTime[ " + timePars[i]["tstart"].getValue() + "]; \n");
				}
				outString.append("			Print[{$Time}, File \"../getdpresults/times.dat\"]; \n");
			}
		}

	}
	return outString;
}



std::string timeLoopOpen(std::vector<std::map<std::string, InputValue>> timePars, std::vector<Physics::timeType> timeTypeIds, std::vector<std::string> solverKeys)
{
	std::string outString = "";
	std::string multiPort = "";
	//outer time loop
	//will go wrong, if several physics are transient
	for (unsigned i = 0; i < timePars.size(); i++)
	{
		if (solverKeys[i] == "getdpMicrowaves") {
			multiPort = "~{i}";
		}
		if (timeTypeIds[i] == Physics::TRANSIENT) 
		{
			// fixed step
			if (timePars[i]["isAdaptive"].getValue() == "0") 
			{
				std::string tstart;
				if (timePars[i].count("continueFromLast"))
				{
					if (std::stof(timePars[i]["continueFromLast"].getValue()) == 1)
					{
						tstart = timePars[i]["trestart"].getValue();
					}
					else
					{
						tstart = timePars[i]["tstart"].getValue();
					}
				}
				else
				{
					tstart = timePars[i]["tstart"].getValue();
				}

				std::string tend = timePars[i]["tend"].getValue();

				std::string tstep;
				if (timePars[i]["tstep"].getTypeId() == InputValue::SCALAR)
					tstep = timePars[i]["tstep"].getValue();
				else if (timePars[i]["tstep"].getTypeId() == InputValue::TABLE)
					tstep = "tsteps[$Time]";

				outString.append("        	TimeLoopTheta [");
				outString.append(tstart + ", " + tend + ", " + tstep);


				outString.append(", 1.0] \n");
				outString.append("        	{ \n");

			}

			//adaptive step
			if (timePars[i]["isAdaptive"].getValue() == "1") {
				outString.append("        	TimeLoopAdaptive[ ");

				if (timePars[i].count("continueFromLast"))
				{
					if (std::stof(timePars[i]["continueFromLast"].getValue()) == 1)
					{
						outString.append(timePars[i]["trestart"].getValue() + ", "
							+ timePars[i]["tend"].getValue() + ", "
							+ timePars[i]["initdt"].getValue() + ", "
							+ timePars[i]["mindt"].getValue() + ", "
							+ timePars[i]["maxdt"].getValue() + ", "
							"\"Euler\", "
							+ "List[Breakpoints],");
					}
					else
					{
						outString.append(timePars[i]["tstart"].getValue() + ", "
							+ timePars[i]["tend"].getValue() + ", "
							+ timePars[i]["initdt"].getValue() + ", "
							+ timePars[i]["mindt"].getValue() + ", "
							+ timePars[i]["maxdt"].getValue() + ", "
							"\"Euler\", "
							+ "List[Breakpoints],");
					}
				}
				else
				{
					outString.append(timePars[i]["tstart"].getValue() + ", "
						+ timePars[i]["tend"].getValue() + ", "
						+ timePars[i]["initdt"].getValue() + ", "
						+ timePars[i]["mindt"].getValue() + ", "
						+ timePars[i]["maxdt"].getValue() + ", "
						"\"Euler\", "
						+ "List[Breakpoints],");
				}
				outString.append(" PostOperation {");
				for (unsigned j = 0; j < timePars.size(); j++)
				{
					outString.append(" { adaptivePostOp_" + solverKeys[j] + ", " + timePars[i]["toler1"].getValue()
						+ ", " + timePars[i]["toler2"].getValue() + ", MeanL1Norm  } ");
				}
				outString.append("} ] \n");

				outString.append(" 		   	{ \n");
			}
		}

		if (timeTypeIds[i] == Physics::MULTIHARMONIC)
		{
			std::string fstart = timePars[i]["frequency"].getValue();
			std::string fend = timePars[i]["fend"].getValue();
			std::string fstep = timePars[i]["fstep"].getValue();

			if (timePars[i]["frequency"].getDoubleValue() >= timePars[i]["fend"].getDoubleValue())
			{
				std::string message = "End frequency should be smaller than starting frequency";
				throw cenos_exception(message);
			}

			outString.append("        	For f In {");
			outString.append( fstart + ": " + fend + " : " + fstep + "}\n");

			outString.append("        		SetFrequency[resol_" + solverKeys[i] + multiPort + ", f*" + std::to_string(timePars[i]["frequency"].getFactor()) + "];\n");
			outString.append("        		Print[{f*" + std::to_string(timePars[i]["frequency"].getFactor()) +"},Format \"Setting Frequency f = %g ...\" ];\n");
			outString.append("        		SetTime[f];\n");
		}

		else if (timeTypeIds[i] == Physics::HARMONIC)
		{
			if (timePars[i]["frequency"].getTypeId() == InputValue::valueType::SCALAR)
			{
				outString.append("        	SetTime[" + timePars[i]["frequency"].getValue() + "];\n");
				outString.append("        	DefineConstant[ f = " + timePars[i]["frequency"].getValue() + "];\n");

			}
			else
			{
				outString.append("        	SetTime[Freq / " + std::to_string(timePars[i]["frequency"].getFactor()) + "];\n");
				outString.append("        	DefineConstant[ f = Freq / " + std::to_string(timePars[i]["frequency"].getFactor()) + "];\n");
			}
		}
	}
	return outString;
}


std::string timeAdaptiveSplit(std::vector<std::map<std::string, InputValue>> timePars, std::vector<Physics::timeType> timeTypeIds)
{
	std::string outString = "";
	for (unsigned i = 0; i < timePars.size(); i++)
	{
		if (timeTypeIds[i] == Physics::TRANSIENT)
		{
			if (timePars[i]["isAdaptive"].getValue() == "1") {
				outString.append(" 		   	} \n");
				outString.append(" 		   	{ \n");
			}
		}
	}
	return outString;
}

std::string timeLoopClose(std::vector<Physics::timeType> timeTypeIds, std::vector<std::string> solverKeys)
{
	std::string outString = "";
	for (unsigned i = 0; i < timeTypeIds.size(); i++)
	{
		if  (timeTypeIds[i] == Physics::TRANSIENT )
		{
			outString.append("        	} \n");
		}
		else if (timeTypeIds[i] == Physics::MULTIHARMONIC)
		{
			if (solverKeys[i] == "getdpMicrowaves") {
				outString.append("        	EndFor \n");
			}
			
		}
	}
	return outString;
}





std::string outerLoopAccurateOpen(std::vector<bool> nonlinearities, Numerics::ALGO coupling, int maxIterations, float relaxFactor, float tolerance)
{
	std::string outString = "";
	for (size_t i = nonlinearities.size(); i > 0; i--)
	{

		if ((nonlinearities[i - 1]) and (coupling != Numerics::ALGO::Fast))
		{
			outString.append("			IterativeLoopN [ ");
			outString.append(std::to_string(maxIterations) + ", " + std::to_string(relaxFactor) + " , System { { resol_getdpThermal, " + std::to_string(tolerance) + ", 1e-5, Residual MeanL1Norm  } } ]\n");
			outString.append("				{ \n");
			break;
		}
	}
	return outString;
}


std::string outerLoopAccurateClose(std::vector<bool> nonlinearities, Numerics::ALGO coupling)
{
	std::string outString = "";
	for (size_t i = nonlinearities.size(); i > 0; i--)
	{

		if ((nonlinearities[i - 1]) and (coupling != Numerics::ALGO::Fast))
		{
			outString.append("				} \n");
			break;
		}
	}
	return outString;
}


std::string finalizeTimeStep(std::vector<std::map<std::string, InputValue>> timePars, std::vector<std::string> solverKeys, std::vector<Physics::timeType> timeTypeIds)
{
	std::string outString = "";
	bool stepWritten = false;
	//saving
	for (size_t i = 0; i < timePars.size(); i++)
	{
		if (timeTypeIds[i] == Physics::MULTIHARMONIC or timeTypeIds[i] == Physics::HARMONIC)
		{
			outString.append("        		SetTime[f];\n");
		}

		std::string multiPort = "";
		if (solverKeys[i] == "getdpMicrowaves") {
			multiPort = "~{i}";
		}


		outString.append("				SaveSolution[resol_" + solverKeys[i] + multiPort + "]; \n");
		if ( (timeTypeIds[i] == Physics::TRANSIENT) and (!stepWritten) )
		{
			outString.append("				Print[{$Time}, File \"../getdpresults/times.dat\"]; \n");
			stepWritten = true;
		}

		if ((timeTypeIds[i] == Physics::MULTIHARMONIC or timeTypeIds[i] == Physics::HARMONIC) and (!stepWritten))
		{
			outString.append("				Print[{f}, File \"../getdpresults/times.dat\"]; \n");
			//write frequencies only once

			stepWritten = true;
		}

	}

	if (!stepWritten)
		outString.append("				Print[{$Time}, File \"../getdpresults/times.dat\"]; \n");

	//update constraints
	for (size_t i = 0; i < timePars.size(); i++)
	{
		std::string multiPort = "";
		if (solverKeys[i] == "getdpMicrowaves") {
			multiPort = "~{i}";
		}
		outString.append("				PostOperation[postOp_" + solverKeys[i] + multiPort + "]; \n");
		if (solverKeys[i] == "getdpMicrowaves") {
			outString.append("			EndFor");
		}
	}
	outString.append("\n");
	return outString;
}




std::string solutionLoops(std::vector<bool> nonlinearities, Numerics::ALGO coupling, std::vector<std::map<std::string, InputValue>> timePars,
	int maxIterations, float relaxFactor, float tolerance, std::vector<std::shared_ptr<PostValues>> postValueList, std::vector<std::string> solverKeys,
	std::vector<std::map<std::string, InputValue>> solverPars, bool transientConstraints)
{
	std::string outString = "";
	for (unsigned i = nonlinearities.size(); i > 0; i--)
	{
		std::vector<int> registerList;
		std::shared_ptr<PostValues> scaling_post = nullptr;
		std::vector<std::string> registerNames;
		for (auto pv : postValueList)
		{
			if (!pv->isEnabled())
				continue;

			if ((pv->getType() == PostValues::MAX_VAL) and (pv->getSolverKey() == solverKeys[i - 1]))
			{
				registerList.push_back(pv->getValueRegister());
				registerNames.push_back(pv->getName());
			}

			if ((pv->getType() == PostValues::SCALING_VALUE) and (pv->getSolverKey() == solverKeys[i - 1]))
			{
				scaling_post = pv;  ;
			}
		}

		// this part only if no transient constraints!
		// otherwise, steps might get skipped!
		if ((registerList.size() != 0) and (!transientConstraints))
		{
			outString.append("				Test[ ($TimeStep<=1)] \n");
			outString.append("				{ \n");
			if ((nonlinearities[i - 1]) and (coupling == Numerics::ALGO::Fast))
			{
				outString.append("					Print [\"Solving formulation [" + solverKeys[i - 1] + "]\"]; \n");
				outString.append("					IterativeLoop [ ");
				outString.append(std::to_string(maxIterations) + ", " + std::to_string(tolerance) + ", " + std::to_string(relaxFactor) + "]\n");
				outString.append("						{ \n");
			}

			outString.append(subPhysicsLoop(i, nonlinearities[i - 1], timePars, 5, registerList, solverKeys, scaling_post));

			if ((nonlinearities[i - 1]) and (coupling == Numerics::ALGO::Fast))
			{
				outString.append("						} \n");
			}
			outString.append("				} \n");
			outString.append("				{ \n");
			for (auto rl : registerList)
			{
				outString.append("					Evaluate[$value" + std::to_string(rl) + " = $value" + std::to_string(rl) +
					" + #" + std::to_string(rl) + "]; \n");

			}

			for (int ir = 0; ir < registerList.size(); ir++)
			{
				outString.append("					Print[{#" + std::to_string(registerList[ir]) + "*100}, Format \'Maximal change of " + registerNames[ir] + " is %f ...\']; \n");
				outString.append("					Print[{$value" + std::to_string(registerList[ir]) + "*100}, Format \"Accumulated change of " + registerNames[ir] + " is %f ...\"]; \n");

			}

			outString.append("					Test[");
			int g = 0;
			for (auto rv : registerList)
			{
				if (g) outString.append(" || ");
				outString.append("($value" + std::to_string(rv) + " > " + solverPars[i - 1]["toleranceFactor"].getValue() + ") ");
				g++;
			}
			outString.append("] \n");

			outString.append("					{ \n");
			if ((nonlinearities[i - 1]) and (coupling == Numerics::ALGO::Fast))
			{
				outString.append("						Print [\"Solving formulation [" + solverKeys[i-1] + "]\"]; \n ");
				outString.append("						IterativeLoop [ ");
				outString.append(std::to_string(maxIterations) + ", " + std::to_string(tolerance) + ", " + std::to_string(relaxFactor) + "]\n");
				outString.append("							{ \n");
			}

			outString.append(subPhysicsLoop(i, nonlinearities[i - 1], timePars, 6, registerList, solverKeys, scaling_post));

			if ((nonlinearities[i - 1]) and (coupling == Numerics::ALGO::Fast))
			{
				outString.append("							} \n");
			}
			outString.append("					} \n");
			outString.append("					{ \n");
			outString.append("						Generate[resol_" + solverKeys[i - 1] + "] ;\n");
			outString.append("						Print[\"Skipping " + solverKeys[i - 1] + "...\"];\n");
			outString.append("					} \n");
			outString.append("				} \n");
			outString.append("\n");
		}
		else
		{
			if ((nonlinearities[i - 1]) and (coupling == Numerics::ALGO::Fast))
			{
				outString.append("				Print [\"Solving formulation [" + solverKeys[i - 1] + "]\"]; \n");
				outString.append("				IterativeLoop [ ");
				outString.append(std::to_string(maxIterations) + ", " + std::to_string(tolerance) + ", " + std::to_string(relaxFactor) + "]\n");
				outString.append("				{ \n");
			}

			outString.append(subPhysicsLoop(i, nonlinearities[i - 1], timePars, 5, registerList, solverKeys, scaling_post));

			if ((nonlinearities[i - 1]) and (coupling == Numerics::ALGO::Fast))
			{
				outString.append("				} \n");
			}
			outString.append("\n");
		}

	}
	return outString;
}



std::string subPhysicsLoop(int i, bool nonlin, std::vector<std::map<std::string, InputValue>> timePars, int nrOfTabs, std::vector<int> registerList, std::vector<std::string> solverKeys, std::shared_ptr<PostValues> scaling_post)
{
	std::string outString = "";
	std::string ts = "	";
	std::string tabString = "";

	for (int tb = 0; tb < nrOfTabs; tb++)
		tabString.append(ts);

	if (nonlin)
	{
		outString.append(tabString);
		outString.append("UpdateConstraint[resol_" + solverKeys[i - 1] + ", Region[All], Assign ]; \n");

		outString.append(tabString);
		outString.append("Print [\"Solving formulation [" + solverKeys[i - 1] + "]\"]; \n");

		outString.append(tabString);
		outString.append("GenerateJac[resol_" + solverKeys[i - 1] + "];\n");

		if (timePars[i - 1].count("frequency"))
		{
			if (timePars[i - 1]["frequency"].getTypeId() == InputValue::TABLE)
			{
				outString.append(tabString);
				outString.append("SetFrequency[resol_" + solverKeys[i - 1] + ", FreqFunc[$Time] ];\n");
				outString.append(tabString);
				outString.append("Print[{FreqFunc[$Time]},Format \"Setting Frequency f = %g ...\" ];\n");
			}
		}
		outString.append(tabString);
		outString.append("SolveJac[resol_" + solverKeys[i - 1] + "];\n");

		if (scaling_post != nullptr)
		{
			outString.append(tabString);
			outString.append("PostOperation[scalingPostOp_" + solverKeys[i - 1] + "] ;\n");
			outString.append(tabString);
			outString.append("Evaluate[$value" + std::to_string(scaling_post->getValueRegister()) + " = " + scaling_post->getEquationForScaling() + "];\n");
			outString.append(tabString);
			outString.append("MultiplySolution[resol_" + solverKeys[i - 1] + ", $value" + std::to_string(scaling_post->getValueRegister()) + "]; \n");
			for (auto cstr : scaling_post->getScalingConstraints())
			{
				outString.append(tabString);
				outString.append("MultiplyConstraint[resol_" + solverKeys[i - 1] + ", $value" + std::to_string(scaling_post->getValueRegister()) + ", " + cstr + "]; \n");
			}
		}

		for (auto rl : registerList)
		{
			outString.append(tabString);
			outString.append("Evaluate[$value" + std::to_string(rl) + " = 0];\n");
		}
	}
	else
	{
		std::string multiPort = "";
		if (solverKeys[i - 1] == "getdpMicrowaves") {
			multiPort = "~{i}";
		}
		outString.append(tabString);
		outString.append("UpdateConstraint[resol_" + solverKeys[i - 1] + multiPort + ", Region[All], Assign ]; \n");
		outString.append(tabString);
		outString.append("Print [\"Solving formulation [" + solverKeys[i - 1] + multiPort + "]\"]; \n");
		outString.append(tabString);
		outString.append("Generate[resol_" + solverKeys[i - 1] + multiPort + "] ;\n");
		outString.append(tabString);
		outString.append("Solve[resol_" + solverKeys[i - 1] + multiPort + "];\n");

		if (scaling_post != nullptr)
		{
			outString.append(tabString);
			outString.append("PostOperation[scalingPostOp_" + solverKeys[i - 1] + multiPort + "] ;\n");
			outString.append(tabString);
			outString.append("Evaluate[$value" + std::to_string(scaling_post->getValueRegister()) + " = " + scaling_post->getEquationForScaling() + "];\n");
			outString.append(tabString);
			outString.append("MultiplySolution[resol_" + solverKeys[i - 1] + multiPort + ", $value" + std::to_string(scaling_post->getValueRegister()) + "]; \n");
			for (auto cstr : scaling_post->getScalingConstraints())
			{
				outString.append(tabString);
				outString.append("MultiplyConstraint[resol_" + solverKeys[i - 1] + multiPort + ", $value" + std::to_string(scaling_post->getValueRegister()) + ", " + cstr + "]; \n");
			}
		}

		for (auto rl : registerList)
		{
			outString.append(tabString);
			outString.append("Evaluate[$value" + std::to_string(rl) + " = 0];\n");
		}
	}
	return outString;
}

std::string closeOperations()
{
	std::string outString = "";
	outString.append("		} \n");
	return outString;
}

std::string closeResolution()
{
	std::string outString = "";
	outString.append("	} \n");

	outString.append("} \n");
	return outString;
}
