#include "MeshToVtkWriter.hpp"
#include "misc/miscFunctions.hpp"
#include <vtkIntArray.h>
#include <vtkHexahedron.h>
#include <vtkWedge.h>
#include <vtkPyramid.h>
#include <vtkTetra.h>
#include <vtkQuad.h>
#include <vtkTriangle.h>
#include <vtkLine.h>
#include <vtkQuadraticTetra.h>
#include <vtkQuadraticTriangle.h>
#include <vtkQuadraticEdge.h>
#include <vtkCellData.h>
#include "misc/Timer.hpp"

MeshToVtkWriter::MeshToVtkWriter(std::shared_ptr<Mesh> m): mesh(m)
{
}

MeshToVtkWriter::~MeshToVtkWriter()
{}

vtkSmartPointer<vtkUnstructuredGrid> MeshToVtkWriter::createVtkUnstructuredGrid(std::vector<MeshEntity_> entities, int grid_id)
{

	std::vector<Node_> nodes;
	std::set<int> node_ids;
	std::vector<Element_> elements;

	for (auto ent : entities)
	{
		for (auto nd : ent->getNodes())
		{
			if (node_ids.find(nd->getId()) == node_ids.end())
			{
				node_ids.insert(nd->getId());
				nodes.push_back(nd);
			}
		}

		for (auto el : ent->getElements())
			elements.push_back(el);
	}

	// sort nodes by id
	// this is important as result writing in dat files is done in ascending node order
	// this sorting might scramble the node order over different time steps
	// which can cause temporal issues in ParaView CP-1303
	std::sort(nodes.begin(), nodes.end(),
		[](const Node_ a, const Node_ b) -> bool
		{ return a->getId() < b->getId(); });


	vtkSmartPointer<vtkPoints> points =
		vtkSmartPointer<vtkPoints>::New();
	vtkSmartPointer<vtkUnstructuredGrid> ugrid =
		vtkSmartPointer<vtkUnstructuredGrid>::New();

	// map to match node ids
	// newNodeIds[oldNodeIds]
	std::map<int, int> newNodeIds;
	int id = 0;
	for (auto& point : nodes)
	{
		// we use only linear elements in post processing!
		if (point->isMidsideNode())
		{
			continue;
		}

		points->InsertNextPoint(point->x() * mesh->getMeshScale(), point->y() * mesh->getMeshScale(), point->z() * mesh->getMeshScale());		//SetPoint will be faster!!!
		newNodeIds[point->getId()] = id;
		id++;
	}

	ugrid->SetPoints(points);

	vtkSmartPointer<vtkIntArray> bID =
		vtkSmartPointer<vtkIntArray>::New();
	bID->SetNumberOfComponents(1);
	bID->SetName("BlockId");


	for (auto& element : elements )
	{
		std::vector<Node_> elementNodes = element->getNodes();
		switch (element->getElementVTKType())
		{
			case (12):
			{
				//HEXA
				vtkSmartPointer<vtkHexahedron> hexa =
					vtkSmartPointer<vtkHexahedron>::New();

				for (int i = 0; i < 8; i++)
				{
					hexa->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(hexa->GetCellType(), hexa->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (13):
			{
				//PRISM
				vtkSmartPointer<vtkWedge> prism =
					vtkSmartPointer<vtkWedge>::New();
				for (int i = 0; i < 6; i++)
				{
					prism->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(prism->GetCellType(), prism->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (14) :
			{
				//PYRAMID
				vtkSmartPointer<vtkPyramid> pyra =
					vtkSmartPointer<vtkPyramid>::New();
				for (int i = 0; i < 5; i++)
				{
					pyra->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(pyra->GetCellType(), pyra->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (10) :
			{
				//TETRA
				vtkSmartPointer<vtkTetra> tetra =
					vtkSmartPointer<vtkTetra>::New();
				for (int i = 0; i < 4; i++)
				{
					tetra->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(tetra->GetCellType(), tetra->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (9):
			{
				//QUAD
				vtkSmartPointer<vtkQuad> quad =
					vtkSmartPointer<vtkQuad>::New();
				for (int i = 0; i < 4; i++)
				{
					quad->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);

				}
				ugrid->InsertNextCell(quad->GetCellType(), quad->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (5):
			{
				//TRIANGLE
				vtkSmartPointer<vtkTriangle> triangle =
					vtkSmartPointer<vtkTriangle>::New();
				for (int i = 0; i < 3; i++)
				{
					triangle->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);

				}
				ugrid->InsertNextCell(triangle->GetCellType(), triangle->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (3):
			{
				//LINE
				vtkSmartPointer<vtkLine> line =
					vtkSmartPointer<vtkLine>::New();
				for (int i = 0; i < 2; i++)
				{
					line->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(line->GetCellType(), line->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (24):
			{
				//Quadratic TETRA
				vtkSmartPointer<vtkTetra> tetra =
					vtkSmartPointer<vtkTetra>::New();
				for (int i = 0; i < 4; i++)
				{
					tetra->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(tetra->GetCellType(), tetra->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case ( 22):
			{
				//Quadratic TRIANGLE
				vtkSmartPointer<vtkTriangle> triangle =
					vtkSmartPointer<vtkTriangle>::New();
				for (int i = 0; i < 3; i++)
				{
					triangle->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(triangle->GetCellType(), triangle->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}
			case (21):
			{
				//Quad LINE
				vtkSmartPointer<vtkLine> line =
					vtkSmartPointer<vtkLine>::New();
				for (int i = 0; i < 2; i++)
				{
					line->GetPointIds()->SetId(i, newNodeIds[elementNodes[i]->getId()]);
				}
				ugrid->InsertNextCell(line->GetCellType(), line->GetPointIds());
				bID->InsertNextValue(grid_id);
				break;
			}

		}
	}

	ugrid->GetCellData()->AddArray(bID);
	return ugrid;
}

void MeshToVtkWriter::createMultiblockMesh(vtkForCenos& vtuG, std::vector<int> bcsWithPost, GeometryData_ geom_data)
{
	std::vector<std::string> blockNames;
	std::vector<int> ids;

	for (auto& dom : geom_data->getDomains())
	{
		if (dom->getRole() != "_dummy_")
		{
			ids.push_back(dom->getId());
			blockNames.push_back(dom->getLabel());
		}
	}

	for (auto& bnd : geom_data->getBoundaries())
	{
		if (std::find(bcsWithPost.begin(), bcsWithPost.end(), bnd->getId()) != bcsWithPost.end() and (bnd->getRole() != "_dummy_"))
		{
			ids.push_back(bnd->getId());
			blockNames.push_back(bnd->getLabel());
		}
	}

	vtuG.setBlockNames(ids, blockNames);

	for (auto& dom : geom_data->getDomains())
	{
		if (dom->getRole() != "_dummy_")
		{
			std::vector<MeshEntity_> mesh_entities;
			for (auto ent : dom->getEntities())
			{
				auto mesh_ent = mesh->getEntityByName(ent->getName());
				mesh_entities.push_back(mesh_ent);
			}
			vtuG.addUGrid(createVtkUnstructuredGrid(mesh_entities, dom->getId()), dom->getId());
		}
	}

	for (auto& bnd : geom_data->getBoundaries())
	{
		if (std::find(bcsWithPost.begin(), bcsWithPost.end(), bnd->getId()) != bcsWithPost.end() and (bnd->getRole() != "_dummy_"))
		{
			std::vector<MeshEntity_> mesh_entities;
			for (auto ent : bnd->getEntities())
			{
				auto mesh_ent = mesh->getEntityByName(ent->getName());
				mesh_entities.push_back(mesh_ent);
			}
			vtuG.addUGrid(createVtkUnstructuredGrid(mesh_entities, bnd->getId()), bnd->getId());
		}
	}
}
