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

#include "Eigen/Core"

#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;

using Eigen::Map;
using Eigen::Array;

template<typename Type>
using ArrayView = Map<Array<Type, Eigen::Dynamic, 1>>;

template<typename Type>
using ConstArrayView = const Array<Type, Eigen::Dynamic, 1>;

template <typename Type>
void BinaryOpForward(
  MaybeLocal<Object>& lhs,
  MaybeLocal<Object>& rhs,
  MaybeLocal<Object>& dst,
  std::function<ConstArrayView<Type>(ArrayView<Type>, ArrayView<Type>)> op
) {
  TypedArrayContents<Type> lhsData(lhs.ToLocalChecked());
  TypedArrayContents<Type> rhsData(rhs.ToLocalChecked());
  TypedArrayContents<Type> dstData(dst.ToLocalChecked());

  ArrayView<Type> a(*lhsData, lhsData.length());
  ArrayView<Type> b(*rhsData, rhsData.length());
  ArrayView<Type> c(*dstData, dstData.length());

  c = op(a, b);
}

template <typename Type>
void BinaryOpBackward(
  MaybeLocal<Object>& lhs,
  MaybeLocal<Object>& rhs,
  MaybeLocal<Object>& output,
  MaybeLocal<Object>& grad,
  MaybeLocal<Object>& dst,
  std::function<ConstArrayView<Type>(ArrayView<Type>, ArrayView<Type>, ArrayView<Type>, ArrayView<Type>)> op
) {
  TypedArrayContents<Type> lhsData(lhs.ToLocalChecked());
  TypedArrayContents<Type> rhsData(rhs.ToLocalChecked());
  TypedArrayContents<Type> gradData(grad.ToLocalChecked());
  TypedArrayContents<Type> outputData(output.ToLocalChecked());
  TypedArrayContents<Type> dstData(dst.ToLocalChecked());

  ArrayView<Type> a(*lhsData, lhsData.length());
  ArrayView<Type> b(*rhsData, rhsData.length());
  ArrayView<Type> c(*gradData, gradData.length());
  ArrayView<Type> d(*outputData, outputData.length());
  ArrayView<Type> e(*dstData, dstData.length());

  e = op(a, b, c, d);
}

#define BINARY_OP_FORWARD(SYMBOL, TYPE, FUNCTOR)                        \
  static NAN_METHOD(SYMBOL) {                                           \
    MaybeLocal<Object> lhs = To<Object>(Arg(info, 0));                  \
    MaybeLocal<Object> rhs = To<Object>(Arg(info, 1));                  \
    MaybeLocal<Object> dst = To<Object>(Arg(info, 2));                  \
    BinaryOpForward<TYPE>(lhs, rhs, dst,                                \
      [&](ArrayView<TYPE> a, ArrayView<TYPE> b) -> ConstArrayView<TYPE> \
      {                                                                 \
        return (FUNCTOR).eval();                                        \
      }                                                                 \
    );                                                                  \
  }

#define BINARY_OP_BACKWARD(SYMBOL, TYPE, FUNCTOR)         \
  static NAN_METHOD(SYMBOL) {                             \
    MaybeLocal<Object> lhs = To<Object>(Arg(info, 0));    \
    MaybeLocal<Object> rhs = To<Object>(Arg(info, 1));    \
    MaybeLocal<Object> output = To<Object>(Arg(info, 2)); \
    MaybeLocal<Object> grad = To<Object>(Arg(info, 3));   \
    MaybeLocal<Object> dst = To<Object>(Arg(info, 4));    \
    BinaryOpBackward<TYPE>(lhs, rhs, output, grad, dst,   \
      [&](ArrayView<TYPE> a,                              \
          ArrayView<TYPE> b,                              \
          ArrayView<TYPE> c,                              \
          ArrayView<TYPE> d)                              \
          -> ConstArrayView<TYPE>                         \
      {                                                   \
        return (FUNCTOR).eval();                          \
      }                                                   \
    );                                                    \
  }

#define BINARY_OP_FORWARD3(SYMBOL,                                   \
                           ARRAYTYPE1, TYPE1,                        \
                           ARRAYTYPE2, TYPE2,                        \
                           ARRAYTYPE3, TYPE3, FUNCTOR)               \
  BINARY_OP_FORWARD(SYMBOL ## Forward ## ARRAYTYPE1, TYPE1, FUNCTOR) \
  BINARY_OP_FORWARD(SYMBOL ## Forward ## ARRAYTYPE2, TYPE2, FUNCTOR) \
  BINARY_OP_FORWARD(SYMBOL ## Forward ## ARRAYTYPE3, TYPE3, FUNCTOR)

  #define BINARY_OP_BACKWARD3(SYMBOL,                                        \
                              ARRAYTYPE1, TYPE1,                             \
                              ARRAYTYPE2, TYPE2,                             \
                              ARRAYTYPE3, TYPE3, FUNCTOR1, FUNCTOR2)         \
  BINARY_OP_BACKWARD(SYMBOL ## Backward1 ## ARRAYTYPE1, TYPE1, FUNCTOR1)     \
  BINARY_OP_BACKWARD(SYMBOL ## Backward2 ## ARRAYTYPE1, TYPE1, FUNCTOR2)     \
  BINARY_OP_BACKWARD(SYMBOL ## Backward1 ## ARRAYTYPE2, TYPE2, FUNCTOR1)     \
  BINARY_OP_BACKWARD(SYMBOL ## Backward2 ## ARRAYTYPE2, TYPE2, FUNCTOR2)     \
  BINARY_OP_BACKWARD(SYMBOL ## Backward1 ## ARRAYTYPE3, TYPE3, FUNCTOR1)     \
  BINARY_OP_BACKWARD(SYMBOL ## Backward2 ## ARRAYTYPE3, TYPE3, FUNCTOR2)

#define BINARY_OP3(SYMBOL,                                   \
                   ARRAYTYPE1, TYPE1,                        \
                   ARRAYTYPE2, TYPE2,                        \
                   ARRAYTYPE3, TYPE3,                        \
                   FUNCTOR1, FUNCTOR2, FUNCTOR3)             \
  BINARY_OP_FORWARD3(SYMBOL,                                 \
                     ARRAYTYPE1, TYPE1,                      \
                     ARRAYTYPE2, TYPE2,                      \
                     ARRAYTYPE3, TYPE3, FUNCTOR1)            \
  BINARY_OP_BACKWARD3(SYMBOL,                                \
                      ARRAYTYPE1, TYPE1,                     \
                      ARRAYTYPE2, TYPE2,                     \
                      ARRAYTYPE3, TYPE3, FUNCTOR2, FUNCTOR3)

#define EXPORT_BINARY_OP3(SYMBOL, ARRAYTYPE1, ARRAYTYPE2, ARRAYTYPE3)                               \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Forward" STRINGIFY(ARRAYTYPE1), SYMBOL ## Forward ## ARRAYTYPE1) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Forward" STRINGIFY(ARRAYTYPE2), SYMBOL ## Forward ## ARRAYTYPE2)     \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Forward" STRINGIFY(ARRAYTYPE3), SYMBOL ## Forward ## ARRAYTYPE3)     \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward1" STRINGIFY(ARRAYTYPE1), SYMBOL ## Backward1 ## ARRAYTYPE1) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward2" STRINGIFY(ARRAYTYPE1), SYMBOL ## Backward2 ## ARRAYTYPE1) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward1" STRINGIFY(ARRAYTYPE2), SYMBOL ## Backward1 ## ARRAYTYPE2) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward2" STRINGIFY(ARRAYTYPE2), SYMBOL ## Backward2 ## ARRAYTYPE2) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward1" STRINGIFY(ARRAYTYPE3), SYMBOL ## Backward1 ## ARRAYTYPE3) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward2" STRINGIFY(ARRAYTYPE3), SYMBOL ## Backward2 ## ARRAYTYPE3)

#define REGISTER_BINARY_OP(SYMBOL, FUNCTOR1, FUNCTOR2, FUNCTOR3) \
  BINARY_OP3(SYMBOL,                                             \
             int32, int32_t,                                     \
             float32, float,                                     \
             float64, double,                                    \
             FUNCTOR1, FUNCTOR2, FUNCTOR3)

#define EXPORT_BINARY_OP(SYMBOL)                     \
  EXPORT_BINARY_OP3(SYMBOL, int32, float32, float64)

REGISTER_BINARY_OP(AddOp, a + b, d * 1, d * 1)
REGISTER_BINARY_OP(SubOp, a - b, d * 1, d * -1)
REGISTER_BINARY_OP(MulOp, a * b, d * b, d * a)
REGISTER_BINARY_OP(DivOp, a / b, d / b, d * (-a / (b * b)))
REGISTER_BINARY_OP(PowOp, a.pow(b), d * (b * a.pow(b - 1)), d * c * a.log())

static void BinaryOpsInit(Handle<Object> exports) {
  EXPORT_BINARY_OP(AddOp)
  EXPORT_BINARY_OP(SubOp)
  EXPORT_BINARY_OP(MulOp)
  EXPORT_BINARY_OP(DivOp)
  EXPORT_BINARY_OP(PowOp)
}
