#ifndef TENSOR_H
#define TENSOR_H

#include <napi.h>
#include <torch/torch.h>

// Include operation headers
#include "ops/add.h"
#include "ops/add_.h"
#include "ops/sub.h"
#include "ops/sub_.h"
#include "ops/mul.h"
#include "ops/mul_.h"
#include "ops/div.h"
#include "ops/div_.h"
#include "ops/matmul.h"
#include "ops/sum.h"
#include "ops/mean.h"
#include "ops/cpu.h"
#include "ops/cuda.h"
#include "ops/mps.h"
#include "ops/float.h"
#include "ops/double.h"
#include "ops/int.h"
#include "ops/long.h"
#include "ops/to.h"
#include "ops/shape.h"
#include "ops/dtype.h"
#include "ops/device.h"
#include "ops/to_string.h"
#include "ops/to_array.h"
#include "ops/reshape.h"
#include "ops/flatten.h"
#include "ops/unsqueeze.h"
#include "ops/squeeze.h"
#include "ops/transpose.h"
#include "ops/permute.h"
#include "ops/requires_grad.h"
#include "ops/backward.h"
#include "ops/grad.h"
#include "ops/zero_grad.h"
#include "ops/detach.h"
#include "ops/relu.h"
#include "ops/sigmoid.h"
#include "ops/tanh.h"
#include "ops/softmax.h"
#include "ops/log_softmax.h"
#include "ops/mse_loss.h"
#include "ops/cross_entropy.h"
#include "ops/nll_loss.h"
#include "ops/binary_cross_entropy.h"

class Tensor : public Napi::ObjectWrap<Tensor> {
public:
  static Napi::Object Init(Napi::Env env, Napi::Object exports);
  static Napi::Value NewInstance(Napi::Env env, torch::Tensor tensor);

  Tensor(const Napi::CallbackInfo& info);

  torch::Tensor tensor;

private:
  static Napi::FunctionReference constructor;

  // Arithmetic operations
  Napi::Value Add(const Napi::CallbackInfo& info);
  Napi::Value Sub(const Napi::CallbackInfo& info);
  Napi::Value Mul(const Napi::CallbackInfo& info);
  Napi::Value Div(const Napi::CallbackInfo& info);

  // In-place arithmetic operations
  Napi::Value AddInplace(const Napi::CallbackInfo& info);
  Napi::Value SubInplace(const Napi::CallbackInfo& info);
  Napi::Value MulInplace(const Napi::CallbackInfo& info);
  Napi::Value DivInplace(const Napi::CallbackInfo& info);

  // Matrix operations
  Napi::Value Matmul(const Napi::CallbackInfo& info);

  // Reduction operations
  Napi::Value Sum(const Napi::CallbackInfo& info);
  Napi::Value Mean(const Napi::CallbackInfo& info);

  // Conversion methods
  Napi::Value ToString(const Napi::CallbackInfo& info);
  Napi::Value ToArray(const Napi::CallbackInfo& info);
  Napi::Value To(const Napi::CallbackInfo& info);

  // Property accessors
  Napi::Value Shape(const Napi::CallbackInfo& info);
  Napi::Value Dtype(const Napi::CallbackInfo& info);
  Napi::Value Device(const Napi::CallbackInfo& info);

  // Device management
  Napi::Value Cpu(const Napi::CallbackInfo& info);
  Napi::Value Cuda(const Napi::CallbackInfo& info);
  Napi::Value Mps(const Napi::CallbackInfo& info);

  // Dtype shortcuts
  Napi::Value Float(const Napi::CallbackInfo& info);
  Napi::Value Double(const Napi::CallbackInfo& info);
  Napi::Value Int(const Napi::CallbackInfo& info);
  Napi::Value Long(const Napi::CallbackInfo& info);

  // Shape operations
  Napi::Value Reshape(const Napi::CallbackInfo& info);
  Napi::Value Flatten(const Napi::CallbackInfo& info);
  Napi::Value Unsqueeze(const Napi::CallbackInfo& info);
  Napi::Value Squeeze(const Napi::CallbackInfo& info);
  Napi::Value Transpose(const Napi::CallbackInfo& info);
  Napi::Value Permute(const Napi::CallbackInfo& info);

  // Autograd operations
  Napi::Value GetRequiresGrad(const Napi::CallbackInfo& info);
  Napi::Value SetRequiresGrad(const Napi::CallbackInfo& info);
  Napi::Value Backward(const Napi::CallbackInfo& info);
  Napi::Value GetGrad(const Napi::CallbackInfo& info);
  Napi::Value ZeroGrad(const Napi::CallbackInfo& info);
  Napi::Value Detach(const Napi::CallbackInfo& info);

  // Activation functions
  Napi::Value Relu(const Napi::CallbackInfo& info);
  Napi::Value Sigmoid(const Napi::CallbackInfo& info);
  Napi::Value Tanh(const Napi::CallbackInfo& info);
  Napi::Value Softmax(const Napi::CallbackInfo& info);
  Napi::Value LogSoftmax(const Napi::CallbackInfo& info);

  // Loss functions
  Napi::Value MseLoss(const Napi::CallbackInfo& info);
  Napi::Value CrossEntropy(const Napi::CallbackInfo& info);
  Napi::Value NllLoss(const Napi::CallbackInfo& info);
  Napi::Value BinaryCrossEntropy(const Napi::CallbackInfo& info);
};

#endif
