#include "AnalyzeMesh.hpp"
#include "mesh/Element.hpp"
#include <iostream>
#include <fstream>
#include "misc/Timer.hpp"

using json = nlohmann::json;


AnalyzeMesh::AnalyzeMesh(std::shared_ptr<Mesh> m, json ms, std::string wdir) {
	mesh = m;
	meshSetup = ms;
	wDir = wdir;
}

void AnalyzeMesh::printElementCounts() {
	std::cout << "### Number of mesh elements ###" << std::endl;
	
	for (json ref : meshSetup["refinments"]) {
		std::string label = ref["label"];
		std::vector<std::string> entityNames;
		if (ref.contains("params3D")) {
			entityNames = ref["params3D"]["entities"].get<std::vector<std::string>>();
		}
		else {
			entityNames = ref["params2D"]["entities"].get<std::vector<std::string>>();
		}
		size_t elementCount = 0;

		for (std::string entName : entityNames) {
			MeshEntity_ ent = mesh->getEntityByName(entName);
			elementCount += ent->getElements().size();
		}
		std::cout << std::left << std::setw(20) << std::setfill(' ') << label;
		std::cout << std::to_string(elementCount) + "\n";
	}
	size_t total_count = mesh->getSizeElements();
	std::cout << std::left << std::setw(20) << std::setfill(' ') << "TOTAL:";
	std::cout << total_count << std::endl;

	std::cout << "### Number of mesh elements ###" << std::endl;

}

void AnalyzeMesh::analyze() {
	printElementCounts();
}


std::pair<Node_, Node_> orderNodePair(Node_ node1, Node_ node2) {
	// We sort the nodes, so that (2,3) and (3,2) wouldnt be recognized as different segments
	if (node1 > node2) {
		return std::pair(node2, node1);
	}
	else {
		return std::pair(node1, node2);
	}
}

std::vector<double> getSegments(std::vector<Element_> elements) {
	std::set<std::pair<Node_, Node_>> segments;

	for (Element_ elem : elements) {
		int elementType = elem->getElementType();
		std::vector nodes = elem->getNodes();
			
		// https://gmsh.info/doc/texinfo/gmsh.html#Node-ordering

		switch (elementType)
		{
		case ElementTypes::LINE:
		case ElementTypes::LINE_2:
			segments.insert(orderNodePair(nodes[0], nodes[1]));
			break;

		case ElementTypes::TRI:
		case ElementTypes::TRI_2:
			segments.insert(orderNodePair(nodes[0], nodes[1]));
			segments.insert(orderNodePair(nodes[1], nodes[2]));
			segments.insert(orderNodePair(nodes[0], nodes[2]));
			break;

		case ElementTypes::QUAD:
			segments.insert(orderNodePair(nodes[0], nodes[1]));
			segments.insert(orderNodePair(nodes[1], nodes[2]));
			segments.insert(orderNodePair(nodes[2], nodes[3]));
			segments.insert(orderNodePair(nodes[0], nodes[3]));
			break;

		case ElementTypes::TETRA:
		case ElementTypes::TETRA_2:
			segments.insert(orderNodePair(nodes[0], nodes[1]));
			segments.insert(orderNodePair(nodes[1], nodes[2]));
			segments.insert(orderNodePair(nodes[2], nodes[3]));
			segments.insert(orderNodePair(nodes[0], nodes[3]));
			break;

		case ElementTypes::HEXA:
			segments.insert(orderNodePair(nodes[0], nodes[1]));
			segments.insert(orderNodePair(nodes[1], nodes[2]));
			segments.insert(orderNodePair(nodes[2], nodes[3]));
			segments.insert(orderNodePair(nodes[3], nodes[0]));

			segments.insert(orderNodePair(nodes[4], nodes[5]));
			segments.insert(orderNodePair(nodes[5], nodes[6]));
			segments.insert(orderNodePair(nodes[6], nodes[7]));
			segments.insert(orderNodePair(nodes[7], nodes[4]));

			segments.insert(orderNodePair(nodes[0], nodes[4]));
			segments.insert(orderNodePair(nodes[1], nodes[5]));
			segments.insert(orderNodePair(nodes[2], nodes[6]));
			segments.insert(orderNodePair(nodes[3], nodes[7]));
			break;

		case ElementTypes::PRISM:
			segments.insert(orderNodePair(nodes[0], nodes[1]));
			segments.insert(orderNodePair(nodes[1], nodes[2]));
			segments.insert(orderNodePair(nodes[2], nodes[0]));

			segments.insert(orderNodePair(nodes[3], nodes[4]));
			segments.insert(orderNodePair(nodes[4], nodes[5]));
			segments.insert(orderNodePair(nodes[5], nodes[3]));

			segments.insert(orderNodePair(nodes[0], nodes[3]));
			segments.insert(orderNodePair(nodes[1], nodes[4]));
			segments.insert(orderNodePair(nodes[2], nodes[5]));
			break;

		case ElementTypes::PYRA:
			segments.insert(orderNodePair(nodes[0], nodes[1]));
			segments.insert(orderNodePair(nodes[1], nodes[2]));
			segments.insert(orderNodePair(nodes[2], nodes[3]));
			segments.insert(orderNodePair(nodes[3], nodes[0]));

			segments.insert(orderNodePair(nodes[4], nodes[0]));
			segments.insert(orderNodePair(nodes[4], nodes[1]));
			segments.insert(orderNodePair(nodes[4], nodes[2]));
			segments.insert(orderNodePair(nodes[4], nodes[3]));
			break;

		default:
			throw std::domain_error("Unexpected element type (" + std::to_string(elementType) + ") in getSegments");
			break;
		}
	}

	std::vector<double> edge_lenghts;
	for (std::pair<Node_, Node_> nodes : segments) {
		Node_ node1 = nodes.first;
		Node_ node2 = nodes.second;
		double x1 = node1->x();
		double y1 = node1->y();
		double z1 = node1->z();
		double x2 = node2->x();
		double y2 = node2->y();
		double z2 = node2->z();

		double d = 0, dd;
		dd = x1; dd -= x2; dd *= dd; d += dd;
		dd = y1; dd -= y2; dd *= dd; d += dd;
		dd = z1; dd -= z2; dd *= dd; d += dd;
		double len = sqrt(d);

		edge_lenghts.push_back(len);

	}
	sort(edge_lenghts.begin(), edge_lenghts.end());
	return edge_lenghts;
}




//void AnalyzeMesh::hackyHistograms() {
//	std::string resultFile = wDir + R"(\geometry\mesh_analysis_results.py)";
//	std::ofstream MyFile(resultFile);
//
//	std::string cenos_config(getenv("CENOS_CONFIG"));
//
//	MyFile << "data= {";
//
//	Timer t;
//	t.setMessage("MeshAnalysis");
//	// TEST for meshing edges
//	// geometry consists of several connected edges.
//	t.start("MeshAnalysis");
//
//	std::vector<MeshEntity_> meshEntities = mesh->getEntities();
//	for (MeshEntity_ meshEnt : meshEntities) {
//		// Warning: This currently will also return all the segments, including those belonging to a lower dimesions
//		// Because of this, segments from one dimensions lower will need to be subtracted
//		std::vector<Element_> elements = meshEnt->getElements();
//		std::vector<double> edge_lenghts = getSegments(elements);
//
//		MyFile << "\n\n\"" << meshEnt->getName() << "\": [\n";
//
//		for (double el : edge_lenghts) {
//			std::string str = std::to_string(el);
//			MyFile << str + ",\n";
//		}
//		MyFile << "],";
//	}
//
//	t.stop();
//
//	// Close the file
//	MyFile << "}\n";
//	MyFile << "meshSetup=";
//	MyFile << meshSetup.dump(4);
//	MyFile << "\npng_path = r'" + wDir + "/geometry/mesh_histogram.png'\n";
//	MyFile << R"(
//
//import numpy as np
//import matplotlib.pyplot as plt
//import itertools
//
//
//colormap = {"params1D":"#ff120a", "params2D":"#00a318", "params3D":"#59bfff"}
//legendmap = {"params1D":"1D: On edges", "params2D":"2D: On faces", "params3D":"3D: In volumes"}
//
//fig, axs = plt.subplots(len(meshSetup["refinments"]), sharex=True)
//
//for i, refinment in enumerate(reversed(meshSetup["refinments"])):
//
//    for param in ["params3D", "params2D", "params1D"]:
//        d = []
//        d = [data[ent] for ent in refinment[param]["entities"]]
//        d_list = list(itertools.chain.from_iterable(d))
//        binwidth = 0.2
//        bins=np.arange(min(d_list), max(d_list) + binwidth, binwidth)
//        # bins = 150
//        axs[i].hist(d_list, density=False, bins=bins, alpha=0.5, color=colormap[param], label=legendmap[param])
//        axs[i].set_title(refinment["label"])
//        axs[i].legend(loc="best")
//        # ax1.yscale('log')
//        if refinment["label"] == "Air":
//            axs[i].set_xlim([0,refinment["params3D"]["maxh"]*0.6])
//
//    
//    axs[i].axvspan(refinment[param]["minh"], refinment["params3D"]["maxh"], ymin=0.95, ymax=1.0, alpha=0.5, color=colormap["params3D"], linewidth=0)
//    axs[i].axvspan(refinment[param]["minh"], refinment["params2D"]["maxh"], ymin=0.9, ymax=0.95, alpha=0.5, color=colormap["params2D"], linewidth=0)
//    axs[i].axvspan(refinment[param]["minh"], refinment["params1D"]["maxh"], ymin=0.85, ymax=0.9, alpha=0.5, color=colormap["params1D"], linewidth=0)
//    
//fig.text(0.5, 0.04, "Mesh segment size, mm", ha='center', va='center')
//fig.text(0.04, 0.5, "Number of mesh segments", ha='center', va='center', rotation='vertical')
//
//fig.set_size_inches(9, 8)
//# plt.show()
//plt.savefig(png_path)
//
//
//)";
//	MyFile.close();
//	std::shared_ptr<FILE> pipe(_popen("where pvpython.exe ", "r"), _pclose);
//
//	if (!pipe) {
//		throw std::domain_error("paraview exe not found");	// could not start pipe
//	}
//
//	char buffer[512];
//	std::string result = "";
//
//	while (!feof(pipe.get())) {
//		if (fgets(buffer, 128, pipe.get()) != NULL) {
//			result += buffer;
//		}
//	}
//
//	std::size_t found = result.find("pvpython.exe");
//	if (found == std::string::npos)
//		throw std::domain_error("paraview exe not found2"); // could not find run_salome.bat
//
//	std::string paraview_path = result.substr(0, found);
//	std::string cmd = paraview_path + "pvpython.exe " + resultFile;
//	std::cout << cmd << std::endl;
//
//	std::shared_ptr<FILE> pipe2(_popen(cmd.c_str(), "r"), _pclose);
//}