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

namespace TensorOps {

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

  // Default values: flatten all dimensions to 1D (start_dim=0, end_dim=-1)
  int64_t start_dim = 0;
  int64_t end_dim = -1;

  // Parse optional start_dim argument
  if (info.Length() >= 1) {
    if (!info[0].IsNumber()) {
      Napi::TypeError::New(env, "start_dim must be a number").ThrowAsJavaScriptException();
      return env.Null();
    }
    start_dim = info[0].As<Napi::Number>().Int64Value();
  }

  // Parse optional end_dim argument
  if (info.Length() >= 2) {
    if (!info[1].IsNumber()) {
      Napi::TypeError::New(env, "end_dim must be a number").ThrowAsJavaScriptException();
      return env.Null();
    }
    end_dim = info[1].As<Napi::Number>().Int64Value();
  }

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

}  // namespace TensorOps
