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

namespace TensorOps {

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

  // If no argument provided, squeeze all dimensions of size 1
  if (info.Length() == 0) {
    torch::Tensor result = self->tensor.squeeze();
    return Tensor::NewInstance(env, result);
  }

  // If dimension argument provided, squeeze only that dimension
  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 squeeze method with dimension
  torch::Tensor result = self->tensor.squeeze(dim);
  return Tensor::NewInstance(env, result);
}

}  // namespace TensorOps
