#include "backward.h"
#include "../tensor.h"
#include <torch/torch.h>

namespace TensorOps {

Napi::Value Backward(Tensor* self, const Napi::CallbackInfo& info) {
  Napi::Env env = info.Env();

  // Check if tensor requires grad
  if (!self->tensor.requires_grad()) {
    Napi::Error::New(env, "Tensor does not require grad and does not have a grad_fn")
      .ThrowAsJavaScriptException();
    return env.Undefined();
  }

  try {
    if (info.Length() == 0) {
      // Scalar backward (no gradient argument)
      self->tensor.backward();
    } else if (info.Length() >= 1) {
      // Non-scalar backward (with gradient argument)
      if (!info[0].IsObject()) {
        Napi::TypeError::New(env, "Gradient argument must be a Tensor")
          .ThrowAsJavaScriptException();
        return env.Undefined();
      }

      Tensor* gradient = Napi::ObjectWrap<Tensor>::Unwrap(info[0].As<Napi::Object>());
      self->tensor.backward(gradient->tensor);
    }
  } catch (const std::exception& e) {
    Napi::Error::New(env, e.what()).ThrowAsJavaScriptException();
    return env.Undefined();
  }

  return env.Undefined();
}

}  // namespace TensorOps
