#pragma once
#define EIGEN_USE_THREADS

#include <node.h>
#include <nan.h>

#include "Eigen/Core"
#include "unsupported/Eigen/CXX11/Tensor"

#include "utils.h"

using v8::FunctionTemplate;
using v8::Handle;
using v8::Object;

using Nan::MaybeLocal;
using Nan::ThrowTypeError;
using Nan::TypedArrayContents;
using Nan::To;
using Nan::New;

template <typename Type, unsigned int Dims>
void MatMul(
  const Eigen::ThreadPoolDevice& device,
  MaybeLocal<Object>& lhs,
  MaybeLocal<Object>& rhs,
  MaybeLocal<Object>& dst,
  MaybeLocal<Object>& lhsShape,
  MaybeLocal<Object>& rhsShape,
  MaybeLocal<Object>& dstShape
) {
  using TMap = Eigen::TensorMap<Eigen::Tensor<Type, Dims, Eigen::RowMajor>, Eigen::Aligned>;
  using Shape = Eigen::array<Eigen::Index, Dims>;
  typedef typename Eigen::Tensor<Type, 2>::DimensionPair DimPair;

  TypedArrayContents<Type> lhsData(lhs.ToLocalChecked());
  TypedArrayContents<Type> rhsData(rhs.ToLocalChecked());
  TypedArrayContents<Type> dstData(dst.ToLocalChecked());

  Local<Object> lhsShapeArgs = lhsShape.ToLocalChecked();
  Local<Object> rhsShapeArgs = rhsShape.ToLocalChecked();
  Local<Object> dstShapeArgs = dstShape.ToLocalChecked();

  Shape lhsShape_;
  Shape rhsShape_;
  Shape dstShape_;

  for (unsigned int i = 0; i < Dims; i++) {
    lhsShape_[i] = lhsShapeArgs->Get(i)->Int32Value();
    rhsShape_[i] = rhsShapeArgs->Get(i)->Int32Value();
    dstShape_[i] = dstShapeArgs->Get(i)->Int32Value();
  }

  const Eigen::array<DimPair, 1> contractionPair({{DimPair(1, 0)}});

  TMap Ta(*lhsData, lhsShape_);
  TMap Tb(*rhsData, rhsShape_);
  TMap Tc(*dstData, dstShape_);

  if (Dims > 2) {
    for (int i = 0; i < dstShape_[0]; i++) {
      auto a = Ta.template chip<0>(i);
      auto b = Tb.template chip<0>(i);
      auto c = Tc.template chip<0>(i);
      c.device(device) = a.contract(b, contractionPair);
    }
  } else {
    Tc.device(device) = Ta.contract(Tb, contractionPair);
  }
}

#define MATMUL_OP(SYMBOL, TYPE, DIMS)                              \
  static NAN_METHOD(SYMBOL) {                                      \
    const Eigen::ThreadPoolDevice& device = GetDefaultDevice();    \
    MaybeLocal<Object> lhs = To<Object>(Arg(info, 0));             \
    MaybeLocal<Object> rhs = To<Object>(Arg(info, 1));             \
    MaybeLocal<Object> dst = To<Object>(Arg(info, 2));             \
    MaybeLocal<Object> lhsShape = To<Object>(Arg(info, 3));        \
    MaybeLocal<Object> rhsShape = To<Object>(Arg(info, 4));        \
    MaybeLocal<Object> dstShape = To<Object>(Arg(info, 5));        \
    MatMul<TYPE, DIMS>(                                            \
      device, lhs, rhs, dst,                                       \
      lhsShape, rhsShape, dstShape);                               \
  }

#define EXPORT_MATMUL_OP(SYMBOL, DIM)                                                            \
  EXPORT_SYMBOL(TOSTRING(SYMBOL) TOSTRING(DIM) "D" TOSTRING(int32), SYMBOL ## DIM ## Dint32)     \
  EXPORT_SYMBOL(TOSTRING(SYMBOL) TOSTRING(DIM) "D" TOSTRING(float32), SYMBOL ## DIM ## Dfloat32) \
  EXPORT_SYMBOL(TOSTRING(SYMBOL) TOSTRING(DIM) "D" TOSTRING(float64), SYMBOL ## DIM ## Dfloat64)

#define REGISTER_MATMUL_OP(SYMBOL, DIM)             \
  MATMUL_OP(SYMBOL ## DIM ## Dint32, int32_t, DIM)  \
  MATMUL_OP(SYMBOL ## DIM ## Dfloat32, float, DIM)  \
  MATMUL_OP(SYMBOL ## DIM ## Dfloat64, double, DIM)

REGISTER_MATMUL_OP(MatMulOp, 2)
REGISTER_MATMUL_OP(MatMulOp, 3)

static void LinAlgOpsInit(Handle<Object> exports) {
  EXPORT_MATMUL_OP(MatMulOp, 2)
  EXPORT_MATMUL_OP(MatMulOp, 3)
}
