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

namespace TensorOps {

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

  try {
    // MSE loss requires a target tensor
    if (info.Length() < 1 || !info[0].IsObject()) {
      Napi::TypeError::New(env, "mse_loss requires a target tensor").ThrowAsJavaScriptException();
      return env.Undefined();
    }

    // Get target tensor
    Tensor* target = Napi::ObjectWrap<Tensor>::Unwrap(info[0].As<Napi::Object>());

    // Parse optional reduction parameter (default: "mean")
    // Options: "none", "mean", "sum"
    std::string reduction = "mean";
    if (info.Length() > 1 && info[1].IsString()) {
      reduction = info[1].As<Napi::String>().Utf8Value();
    }

    // Compute MSE: (input - target)^2
    torch::Tensor diff = self->tensor - target->tensor;
    torch::Tensor squared = diff * diff;

    // Apply reduction
    torch::Tensor result;
    if (reduction == "none") {
      result = squared;
    } else if (reduction == "sum") {
      result = squared.sum();
    } else {  // "mean" (default)
      result = squared.mean();
    }

    return Tensor::NewInstance(env, result);
  } catch (const std::exception& e) {
    Napi::Error::New(env, e.what()).ThrowAsJavaScriptException();
    return env.Undefined();
  }
}

}  // namespace TensorOps
