#include <iostream>
#include <tuple>
#include <unordered_map>

#include <Eigen/Cholesky>
#include <Eigen/SparseCore>

#include <OpenMesh/Core/Mesh/PolyMesh_ArrayKernelT.hh>
#include <OpenMesh/Core/IO/MeshIO.hh>

// Parameters
double eps_energy = 1;
bool dynamic_epsilon = true;   // * 0.9 in each iteration
double eps_regularizer = 1e-3;
size_t iterations = 10;

using Vector = OpenMesh::Vec3d;
struct Traits : public OpenMesh::DefaultTraits {
  using Point  = Vector;
  using Normal = Vector;
};
using Mesh = OpenMesh::PolyMesh_ArrayKernelT<Traits>;

using Triplet = Eigen::Triplet<double>;
using SpVec = Eigen::SparseVector<double>;
using SpMat = Eigen::SparseMatrix<double>;

int main(int argc, char **argv) {
  if (argc != 2) {
    std::cerr << "Usage: " << argv[0] << " <input-mesh>" << std::endl;
    return 1;
  }

  Mesh mesh;
  if (!OpenMesh::IO::read_mesh(mesh, argv[1])) {
    std::cerr << "Cannot read mesh from file: " << argv[1] << std::endl;
    return 2;
  }
  mesh.request_face_normals();
  mesh.update_face_normals();

  std::unordered_map<Mesh::FaceHandle, size_t> fidx;
  std::vector<Mesh::FaceHandle> fhdl;
  std::unordered_map<Mesh::VertexHandle, size_t> vidx;
  std::vector<Mesh::VertexHandle> vhdl;
  std::vector<std::tuple<Mesh::VertexHandle, Mesh::VertexHandle, Mesh::FaceHandle>> edges;

  // Count entities and populate maps
  size_t n = 0;
  for (auto v : mesh.vertices())
    if (!v.is_boundary()) {
      vidx[v] = n++;
      vhdl.push_back(v);
    }
  size_t m = 0, ne = 0;
  for (auto f : mesh.faces()) {
    fidx[f] = m++;
    fhdl.push_back(f);
    for (auto e : f.edges()) {
      edges.push_back(std::make_tuple(e.v0(), e.v1(), f));
      ne++;
    }
  }

  size_t nv = 3*n + 3*m + n;  // # of variables   : px py pz ... nx ny nz ... d ...
  size_t nc = m + n + ne;     // # of constraints : |n|=1 ... |p-p'|=d ... (pj-pi)n=0 ...

  // Initialize x
  Eigen::VectorXd x = Eigen::VectorXd::Zero(nv);

  for (size_t i = 0; i < n; ++i)
    x.segment(3 * i, 3) = Eigen::Vector3d(mesh.point(vhdl[i]).data());
  for (size_t i = 0; i < m; ++i)
    x.segment(3 * (n + i), 3) = Eigen::Vector3d(mesh.normal(fhdl[i]).data());

  // Setup constraints
  std::vector<SpMat> H;
  std::vector<SpVec> b;
  std::vector<double> c;
  for (size_t i = 0; i < m; ++i) { // |n|=1
    H.emplace_back(nv, nv);
    b.emplace_back(nv);
    std::vector<Triplet> data;
    for (size_t j = 0; j < 3; ++j)
      data.emplace_back(3 * (n + i) + j, 3 * (n + i) + j, 2);
    H.back().setFromTriplets(data.begin(), data.end());
    c.push_back(-1);
  }
  for (size_t i = 0; i < n; ++i) { // |p-p'|=d
    H.emplace_back(nv, nv);
    b.emplace_back(nv);
    std::vector<Triplet> data;
    for (size_t j = 0; j < 3; ++j)
      data.emplace_back(3 * i + j, 3 * i + j, 2);
    H.back().setFromTriplets(data.begin(), data.end());
    for (size_t j = 0; j < 3; ++j)
      b.back().coeffRef(3 * i + j) = -2 * mesh.point(vhdl[i])[j];
    b.back().coeffRef(3 * (n + m) + i) = -1;
    c.push_back(mesh.point(vhdl[i]).sqrnorm());
  }
  for (size_t i = 0; i < ne; ++i) { // (pj-pi)n=0
    H.emplace_back(nv, nv);
    b.emplace_back(nv);
    auto &[v1, v2, f] = edges[i];
    auto k = fidx.at(f);
    std::vector<Triplet> data;
    if (vidx.contains(v1)) {
      auto i1 = vidx.at(v1);
      for (size_t j = 0; j < 3; ++j) {
        data.emplace_back(3 * (n + k) + j, 3 * i1 + j, -1);
        data.emplace_back(3 * i1 + j, 3 * (n + k) + j, -1);
      }
    } else
      for (size_t j = 0; j < 3; ++j)
        b.back().coeffRef(3 * (n + k) + j) -= mesh.point(v1)[j];
    if (vidx.contains(v2)) {
      auto i2 = vidx.at(v2);
      for (size_t j = 0; j < 3; ++j) {
        data.emplace_back(3 * (n + k) + j, 3 * i2 + j, 1);
        data.emplace_back(3 * i2 + j, 3 * (n + k) + j, 1);
      }
    } else
      for (size_t j = 0; j < 3; ++j)
        b.back().coeffRef(3 * (n + k) + j) += mesh.point(v2)[j];
    H.back().setFromTriplets(data.begin(), data.end());
    c.push_back(0);
  }

  // Setup the energy
  SpMat K(1, nv);
  for (size_t i = 0; i < n; ++i)
    K.coeffRef(0, 3 * (n + m) + i) = eps_energy;
  Eigen::VectorXd s(1);
  s(0) = 0;

  // Start the iteration
  auto quad = [&](size_t i, const Eigen::VectorXd &x) {
    return 0.5 * x.dot(H[i] * x) + b[i].dot(x) + c[i];
  };
  for (size_t iter = 1; iter <= iterations; ++iter) {
    std::cout << "Iteration " << iter << std::endl;

    SpMat Ht(nv, nc);
    Eigen::VectorXd r(nc);
    for (size_t i = 0; i < nc; ++i) {
      Ht.col(i) = H[i] * x + b[i];
      r(i) = Ht.col(i).dot(x) - quad(i, x);
    }

    if (dynamic_epsilon)
      K *= 0.9;

    Eigen::MatrixXd A =
      Ht * Ht.transpose() + K.transpose() * K +
      std::pow(eps_regularizer, 2) * Eigen::MatrixXd::Identity(nv, nv);
    Eigen::VectorXd b = Ht * r + K.transpose() * s + std::pow(eps_regularizer, 2) * x;

    x = A.llt().solve(b);
    std::cout << "Error: " << (Ht.transpose() * x - r).squaredNorm() << std::endl;
    std::cout << "Energy: " << (K * x - s).squaredNorm() << std::endl;
  }

  // Modfiy the mesh and write the output
  for (size_t i = 0; i < n; ++i)
    mesh.set_point(vhdl[i], { x(3 * i), x(3 * i + 1), x(3 * i + 2) });
  OpenMesh::IO::write_mesh(mesh, "output.obj");
}
