#include "unsqueeze.h"
#include "../tensor.h"
#include "../utils.h"
#include <torch/torch.h>

namespace TensorOps {

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

  if (info.Length() < 1) {
    Napi::TypeError::New(env, "Expected dimension argument").ThrowAsJavaScriptException();
    return env.Null();
  }

  if (!info[0].IsNumber()) {
    Napi::TypeError::New(env, "Dimension must be a number").ThrowAsJavaScriptException();
    return env.Null();
  }

  int64_t dim = info[0].As<Napi::Number>().Int64Value();

  // Use PyTorch's unsqueeze method
  torch::Tensor result = self->tensor.unsqueeze(dim);
  return Tensor::NewInstance(env, result);
}

}  // namespace TensorOps
