#include "config.h"

#include <iomanip>
#include <iostream>
#include <sstream>
#include <string>
#include <vector>

namespace {

bool to_number(const std::string &s, int &out) {
  try {
    size_t used = 0;
    const int v = std::stoi(s, &used);
    return used == s.size() && (out = v, true);
  } catch (const std::exception &) {
    return false;
  }
}

bool to_number(const std::string &s, float &out) {
  try {
    size_t used = 0;
    const float v = std::stof(s, &used);
    return used == s.size() && (out = v, true);
  } catch (const std::exception &) {
    return false;
  }
}

template <class T>
struct Named {
  const char *name;
  T value;
};

constexpr Named<SampleStrategy> kSamplers[] = {{"proposal", SampleStrategy::PROPOSAL},
                                               {"stratified", SampleStrategy::STRATIFIED},
                                               {"uniform", SampleStrategy::UNIFORM}};
constexpr Named<DevicePreference> kDevices[] = {{"auto", DevicePreference::AUTO},
                                                {"cpu", DevicePreference::CPU},
                                                {"cuda", DevicePreference::CUDA}};
constexpr Named<LossType> kLosses[] = {{"huber", LossType::PSEUDO_HUBER},
                                       {"mse", LossType::MSE}};

template <class T, size_t N>
bool from_name(const Named<T> (&table)[N], const std::string &s, T &out) {
  for (const auto &e : table)
    if (s == e.name) return (out = e.value, true);
  return false;
}

template <class T, size_t N>
const char *to_name(const Named<T> (&table)[N], T value) {
  for (const auto &e : table)
    if (value == e.value) return e.name;
  return "?";
}

// One CLI flag. A row with no setter is a section header in --help.
struct Opt {
  const char *flag;
  const char *meta;
  const char *help;
  bool (*set)(Config &, const std::string &) = nullptr;
  void (*show)(std::ostream &, const Config &) = nullptr;
};

template <auto Field>
constexpr Opt number(const char *flag, const char *meta, const char *help) {
  return {flag, meta, help,
          [](Config &c, const std::string &v) { return to_number(v, c.*Field); },
          [](std::ostream &o, const Config &c) { o << c.*Field; }};
}

template <auto Field>
constexpr Opt path(const char *flag, const char *meta, const char *help) {
  return {flag, meta, help,
          [](Config &c, const std::string &v) { return (c.*Field = v, true); },
          [](std::ostream &o, const Config &c) {
            o << ((c.*Field).empty() ? "none" : (c.*Field).string());
          }};
}

const Opt kOpts[] = {
    {nullptr, nullptr, "Training"},
    path<&Config::render_checkpoint>("--render", "FILE",
                                     "skip training: load this checkpoint and render the orbit"),
    number<&Config::n_iters>("--iters", "N", "training iterations"),
    number<&Config::image_size>("--size", "N", "images are resized to N x N"),
    number<&Config::seed>("--seed", "N", "RNG seed"),
    {"--device", "auto|cpu|cuda", "auto falls back to CPU, cuda exits without a GPU",
     [](Config &c, const std::string &v) { return from_name(kDevices, v, c.device_pref); },
     [](std::ostream &o, const Config &c) { o << to_name(kDevices, c.device_pref); }},

    {nullptr, nullptr, "Sampling"},
    {"--sampler", "proposal|stratified|uniform", "proposal is hierarchical",
     [](Config &c, const std::string &v) { return from_name(kSamplers, v, c.sampler); },
     [](std::ostream &o, const Config &c) { o << to_name(kSamplers, c.sampler); }},
    number<&Config::n_samples>("--samples", "N", "coarse samples per ray"),
    number<&Config::n_importance>("--importance", "N", "fine importance samples"),
    number<&Config::ray_batch>("--ray-batch", "N", "rays per step, pooled across images"),
    number<&Config::z_near>("--near", "F", "near plane"),
    number<&Config::z_far>("--far", "F", "far plane"),
    number<&Config::warmup_iters>("--warmup", "N", "iters the full model drives the coarse pass"),
    number<&Config::interlevel_weight>("--interlevel-weight", "F", "weight on the interlevel loss"),

    {nullptr, nullptr, "Model and optimiser"},
    number<&Config::width>("--width", "N", "trunk width"),
    number<&Config::depth>("--depth", "N", "hidden SIREN layers in the trunk"),
    number<&Config::learning_rate>("--lr", "F", "AdamW learning rate"),
    number<&Config::weight_decay>("--weight-decay", "F", "AdamW weight decay"),
    {"--loss", "huber|mse", "photometric loss",
     [](Config &c, const std::string &v) { return from_name(kLosses, v, c.loss); },
     [](std::ostream &o, const Config &c) { o << to_name(kLosses, c.loss); }},
    number<&Config::huber_c>("--huber-c", "F", "pseudo-Huber transition point"),

    {nullptr, nullptr, "Throughput and output"},
    number<&Config::batch_size>("--batch-size", "N", "sample points per forward chunk"),
    number<&Config::log_freq>("--log-every", "N", "iterations between loss lines"),
    number<&Config::eval_freq>("--eval-every", "N", "iterations between test-view evaluations"),
    number<&Config::plot_freq>("--preview-every", "N", "iterations between previews and checkpoints"),
    number<&Config::n_preview_frames>("--preview-frames", "N", "views per preview render"),
    number<&Config::n_final_frames>("--final-frames", "N", "views in the final orbit"),
};

bool fail(const std::string &message) {
  std::cerr << "Error: " << message << std::endl;
  return false;
}

void print_usage(const char *program) {
  const Config d;
  std::cout << "Usage: " << program << " <data_path> <output_path> [options]\n\n"
            << "  data_path    directory holding transforms.json and the images it references\n"
            << "  output_path  directory for checkpoints, preview renders and metrics\n";
  for (const auto &o : kOpts) {
    if (!o.set) {
      std::cout << "\n" << o.help << "\n";
      continue;
    }
    std::ostringstream lhs;
    lhs << "  " << o.flag << " " << o.meta;
    std::cout << std::left << std::setw(38) << lhs.str() << ' ' << o.help << " (default ";
    o.show(std::cout, d);
    std::cout << ")\n";
  }
  std::cout << "  -h, --help\n";
}

}  // namespace

const char *sampler_name(SampleStrategy s) { return to_name(kSamplers, s); }

bool parse_arguments(int argc, char *argv[], Config &cfg, bool &help) {
  help = false;
  const char *program = argc > 0 ? argv[0] : "NeRF.cpp";
  std::vector<std::string> positional;

  for (int i = 1; i < argc; i++) {
    const std::string arg = argv[i];
    if (arg == "-h" || arg == "--help") {
      print_usage(program);
      help = true;
      return false;
    }
    if (!arg.starts_with("--")) {
      positional.push_back(arg);
      continue;
    }

    const Opt *opt = nullptr;
    for (const auto &o : kOpts)
      if (o.flag && arg == o.flag) opt = &o;
    if (!opt)
      return fail("unknown option '" + arg + "'\nRun " + program + " --help for the list.");
    if (i + 1 >= argc) return fail(arg + " needs a value");

    const std::string value = argv[++i];
    if (!opt->set(cfg, value))
      return fail(arg + " expects " + opt->meta + ", got '" + value + "'");
  }

  if (positional.size() != 2)
    return fail("expected <data_path> and <output_path>, got " +
                std::to_string(positional.size()) + "\nUsage: " + program +
                " <data_path> <output_path> [options]");
  cfg.data_path = positional[0];
  cfg.output_path = positional[1];

  // Settings that would otherwise fail confusingly much later on.
  const struct {
    bool bad;
    const char *message;
  } checks[] = {
      {cfg.n_iters < 0, "--iters must be >= 0"},
      {cfg.image_size <= 0, "--size must be > 0"},
      {cfg.n_samples <= 0, "--samples must be > 0"},
      {cfg.n_importance < 0, "--importance must be >= 0"},
      {cfg.ray_batch <= 0, "--ray-batch must be > 0"},
      {cfg.batch_size <= 0, "--batch-size must be > 0"},
      {cfg.width <= 0 || cfg.depth <= 0, "--width and --depth must be > 0"},
      {cfg.log_freq <= 0 || cfg.plot_freq <= 0 || cfg.eval_freq <= 0,
       "--log-every, --eval-every and --preview-every must be > 0"},
      {cfg.z_near <= 0.0f || cfg.z_far <= cfg.z_near, "need 0 < --near < --far"},
      {cfg.hierarchical() && cfg.n_importance == 0,
       "--sampler proposal needs --importance > 0; use --sampler stratified for one pass"},
  };
  for (const auto &c : checks)
    if (c.bad) return fail(c.message);

  return true;
}
