#pragma once
#include <node.h>
#include <nan.h>
#include <math.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 UnaryOpForward(
  MaybeLocal<Object>& src,
  MaybeLocal<Object>& dst,
  std::function<ConstArrayView<Type>(ArrayView<Type>)> op
) {
  TypedArrayContents<Type> srcData(src.ToLocalChecked());
  TypedArrayContents<Type> dstData(dst.ToLocalChecked());

  ArrayView<Type> a(*srcData, srcData.length());
  ArrayView<Type> b(*dstData, dstData.length());

  b = op(a);
}

template <typename Type>
void UnaryOpBackward(
  MaybeLocal<Object>& src,
  MaybeLocal<Object>& output,
  MaybeLocal<Object>& grad,
  MaybeLocal<Object>& dst,
  std::function<ConstArrayView<Type>(ArrayView<Type>, ArrayView<Type>, ArrayView<Type>)> op
) {
  TypedArrayContents<Type> srcData(src.ToLocalChecked());
  TypedArrayContents<Type> gradData(grad.ToLocalChecked());
  TypedArrayContents<Type> outputData(output.ToLocalChecked());
  TypedArrayContents<Type> dstData(dst.ToLocalChecked());

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

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

#define UNARY_OP_FORWARD(SYMBOL, TYPE, FUNCTOR)        \
  static NAN_METHOD(SYMBOL) {                          \
    MaybeLocal<Object> src = To<Object>(Arg(info, 0)); \
    MaybeLocal<Object> dst = To<Object>(Arg(info, 1)); \
    UnaryOpForward<TYPE>(src, dst,                     \
      [&](ArrayView<TYPE> a) -> ConstArrayView<TYPE>   \
      {                                                \
        return (FUNCTOR).eval();                       \
      }                                                \
    );                                                 \
  }

#define UNARY_OP_BACKWARD(SYMBOL, TYPE, FUNCTOR)                                           \
  static NAN_METHOD(SYMBOL) {                                                              \
    MaybeLocal<Object> src = To<Object>(Arg(info, 0));                                     \
    MaybeLocal<Object> output = To<Object>(Arg(info, 1));                                  \
    MaybeLocal<Object> grad = To<Object>(Arg(info, 2));                                    \
    MaybeLocal<Object> dst = To<Object>(Arg(info, 3));                                     \
    UnaryOpBackward<TYPE>(src, output, grad, dst,                                          \
      [&](ArrayView<TYPE> a, ArrayView<TYPE> b, ArrayView <TYPE>c) -> ConstArrayView<TYPE> \
      {                                                                                    \
        using Array = ConstArrayView<TYPE>; return (FUNCTOR).eval();                       \
      }                                                                                    \
    );                                                                                     \
  }

#define UNARY_OP_FORWARD3(SYMBOL,                                   \
                          ARRAYTYPE1, TYPE1,                        \
                          ARRAYTYPE2, TYPE2,                        \
                          ARRAYTYPE3, TYPE3, FUNCTOR)               \
  UNARY_OP_FORWARD(SYMBOL ## Forward ## ARRAYTYPE1, TYPE1, FUNCTOR) \
  UNARY_OP_FORWARD(SYMBOL ## Forward ## ARRAYTYPE2, TYPE2, FUNCTOR) \
  UNARY_OP_FORWARD(SYMBOL ## Forward ## ARRAYTYPE3, TYPE3, FUNCTOR)

  #define UNARY_OP_BACKWARD3(SYMBOL,                                   \
                             ARRAYTYPE1, TYPE1,                        \
                             ARRAYTYPE2, TYPE2,                        \
                             ARRAYTYPE3, TYPE3, FUNCTOR1)              \
  UNARY_OP_BACKWARD(SYMBOL ## Backward ## ARRAYTYPE1, TYPE1, FUNCTOR1) \
  UNARY_OP_BACKWARD(SYMBOL ## Backward ## ARRAYTYPE2, TYPE2, FUNCTOR1) \
  UNARY_OP_BACKWARD(SYMBOL ## Backward ## ARRAYTYPE3, TYPE3, FUNCTOR1)

#define UNARY_OP3(SYMBOL,                          \
                  ARRAYTYPE1, TYPE1,               \
                  ARRAYTYPE2, TYPE2,               \
                  ARRAYTYPE3, TYPE3,               \
                  FUNCTOR1, FUNCTOR2)              \
  UNARY_OP_FORWARD3(SYMBOL,                        \
                    ARRAYTYPE1, TYPE1,             \
                    ARRAYTYPE2, TYPE2,             \
                    ARRAYTYPE3, TYPE3, FUNCTOR1)   \
  UNARY_OP_BACKWARD3(SYMBOL,                       \
                     ARRAYTYPE1, TYPE1,            \
                     ARRAYTYPE2, TYPE2,            \
                     ARRAYTYPE3, TYPE3, FUNCTOR2)

#define EXPORT_UNARY_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) "Backward" STRINGIFY(ARRAYTYPE1), SYMBOL ## Backward ## ARRAYTYPE1) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward" STRINGIFY(ARRAYTYPE2), SYMBOL ## Backward ## ARRAYTYPE2) \
  EXPORT_SYMBOL(STRINGIFY(SYMBOL) "Backward" STRINGIFY(ARRAYTYPE3), SYMBOL ## Backward ## ARRAYTYPE3)

#define REGISTER_UNARY_OP(SYMBOL, FUNCTOR1, FUNCTOR2) \
  UNARY_OP3(SYMBOL,                                   \
            int32, int32_t,                           \
            float32, float,                           \
            float64, double,                          \
            FUNCTOR1, FUNCTOR2)

#define EXPORT_UNARY_OP(SYMBOL)                                   \
  EXPORT_UNARY_OP3(SYMBOL, int32, float32, float64)

REGISTER_UNARY_OP(NegOp, -a, -c)
REGISTER_UNARY_OP(AbsOp, a.abs(), c * (a >= 0).select(Array::Ones(a.size()), -Array::Ones(a.size())))
REGISTER_UNARY_OP(SqrtOp, a.sqrt(), c / (2 * b))
REGISTER_UNARY_OP(ExpOp, a.exp(), c * b)
REGISTER_UNARY_OP(LogOp, a.log(), c / a)
REGISTER_UNARY_OP(SinOp, a.sin(), c * a.cos())
REGISTER_UNARY_OP(CosOp, a.cos(), c * -a.sin());
REGISTER_UNARY_OP(TanOp, a.tan(), c * (1 + b * b))
REGISTER_UNARY_OP(SinhOp, a.sinh(), c * a.cosh())
REGISTER_UNARY_OP(CoshOp, a.cosh(), c * a.sinh())
REGISTER_UNARY_OP(TanhOp, a.tanh(), c * (1 - a * a));
REGISTER_UNARY_OP(AsinOp, a.asin(), c / (1 - a * a).sqrt())
REGISTER_UNARY_OP(AcosOp, a.acos(), -c / (1 - a * a).sqrt())
REGISTER_UNARY_OP(AtanOp, a.atan(), c / (1 + a * a))
REGISTER_UNARY_OP(AsinhOp, (a + (a * a + 1).sqrt()).log(), c / (a * a + 1).sqrt())
REGISTER_UNARY_OP(AcoshOp, (a + (a * a - 1).sqrt()).log(), c / (a * a - 1).sqrt())
REGISTER_UNARY_OP(AtanhOp, ((a + 1).log() - (1 - a).log()) / 2, c / (1 - a * a))

static void UnaryOpsInit(Handle<Object> exports) {
  EXPORT_UNARY_OP(NegOp)
  EXPORT_UNARY_OP(AbsOp)
  EXPORT_UNARY_OP(SqrtOp)
  EXPORT_UNARY_OP(ExpOp)
  EXPORT_UNARY_OP(LogOp)
  EXPORT_UNARY_OP(SinOp)
  EXPORT_UNARY_OP(CosOp)
  EXPORT_UNARY_OP(TanOp)
  EXPORT_UNARY_OP(SinhOp)
  EXPORT_UNARY_OP(CoshOp)
  EXPORT_UNARY_OP(TanhOp)
  EXPORT_UNARY_OP(AsinOp)
  EXPORT_UNARY_OP(AcosOp)
  EXPORT_UNARY_OP(AtanOp)
  EXPORT_UNARY_OP(AsinhOp)
  EXPORT_UNARY_OP(AcoshOp)
  EXPORT_UNARY_OP(AtanhOp)
}
