#include <napi.h>
#include <torch/torch.h>
#include "tensor.h"
#include "ops/no_grad.h"

// Initialize the addon
Napi::Object Init(Napi::Env env, Napi::Object exports) {
  // Export Tensor class
  Tensor::Init(env, exports);

  // Export utility functions
  exports.Set("add", Napi::Function::New(env, [](const Napi::CallbackInfo& info) -> Napi::Value {
    Napi::Env env = info.Env();

    if (info.Length() < 2) {
      Napi::TypeError::New(env, "Expected two tensor arguments").ThrowAsJavaScriptException();
      return env.Null();
    }

    if (!info[0].IsObject() || !info[1].IsObject()) {
      Napi::TypeError::New(env, "Arguments must be Tensor objects").ThrowAsJavaScriptException();
      return env.Null();
    }

    Tensor* a = Napi::ObjectWrap<Tensor>::Unwrap(info[0].As<Napi::Object>());
    Tensor* b = Napi::ObjectWrap<Tensor>::Unwrap(info[1].As<Napi::Object>());

    torch::Tensor result = a->tensor + b->tensor;

    return Tensor::NewInstance(env, result);
  }));

  exports.Set("zeros", Napi::Function::New(env, [](const Napi::CallbackInfo& info) -> Napi::Value {
    Napi::Env env = info.Env();

    if (info.Length() < 1 || !info[0].IsArray()) {
      Napi::TypeError::New(env, "Expected array of dimensions").ThrowAsJavaScriptException();
      return env.Null();
    }

    Napi::Array dims = info[0].As<Napi::Array>();
    std::vector<int64_t> sizes;

    for (uint32_t i = 0; i < dims.Length(); i++) {
      Napi::Value val = dims[i];
      if (val.IsNumber()) {
        sizes.push_back(val.As<Napi::Number>().Int64Value());
      }
    }

    torch::Tensor tensor = torch::zeros(sizes);
    return Tensor::NewInstance(env, tensor);
  }));

  exports.Set("ones", Napi::Function::New(env, [](const Napi::CallbackInfo& info) -> Napi::Value {
    Napi::Env env = info.Env();

    if (info.Length() < 1 || !info[0].IsArray()) {
      Napi::TypeError::New(env, "Expected array of dimensions").ThrowAsJavaScriptException();
      return env.Null();
    }

    Napi::Array dims = info[0].As<Napi::Array>();
    std::vector<int64_t> sizes;

    for (uint32_t i = 0; i < dims.Length(); i++) {
      Napi::Value val = dims[i];
      if (val.IsNumber()) {
        sizes.push_back(val.As<Napi::Number>().Int64Value());
      }
    }

    torch::Tensor tensor = torch::ones(sizes);
    return Tensor::NewInstance(env, tensor);
  }));

  exports.Set("randn", Napi::Function::New(env, [](const Napi::CallbackInfo& info) -> Napi::Value {
    Napi::Env env = info.Env();

    if (info.Length() < 1 || !info[0].IsArray()) {
      Napi::TypeError::New(env, "Expected array of dimensions").ThrowAsJavaScriptException();
      return env.Null();
    }

    Napi::Array dims = info[0].As<Napi::Array>();
    std::vector<int64_t> sizes;

    for (uint32_t i = 0; i < dims.Length(); i++) {
      Napi::Value val = dims[i];
      if (val.IsNumber()) {
        sizes.push_back(val.As<Napi::Number>().Int64Value());
      }
    }

    torch::Tensor tensor = torch::randn(sizes);
    return Tensor::NewInstance(env, tensor);
  }));

  exports.Set("noGrad", Napi::Function::New(env, TensorOps::NoGrad));

  return exports;
}

NODE_API_MODULE(tytorch_native, Init)
