// viewer-routes.cpp: all API route implementations
//
// Models are loaded lazily on first access (gguf_init_from_file + mmap)
// and cached in server_state. KV metadata is read via gguf_get_* API,
// tensor data is dequantized via ggml_get_type_traits->to_float.

#include "viewer-http.h"

#include <ggml.h>
#include <gguf.h>

#include <algorithm>
#include <cctype>
#include <cstdlib>
#include <cstring>

// scan root directory for .gguf files
static std::vector<std::string> scan_gguf_files(const fs::path &root) {
  std::vector<std::string> paths;
  std::error_code ec;
  for (auto &entry : fs::recursive_directory_iterator(
           root, fs::directory_options::skip_permission_denied, ec)) {
    if (ec)
      break;
    if (!entry.is_regular_file(ec))
      continue;
    auto ext = entry.path().extension().string();
    for (auto &c : ext)
      c = static_cast<char>(std::tolower(static_cast<unsigned char>(c)));
    if (ext != ".gguf")
      continue;
    auto rel = fs::relative(entry.path(), root, ec);
    if (!ec)
      paths.push_back(rel.generic_string());
  }
  std::sort(paths.begin(), paths.end());
  return paths;
}

// get ?model= parameter
static std::string get_model_param(const httplib::Request &req) {
  return req.get_param_value("model");
}

// parse unsigned integer from string, return fallback on empty/invalid
static size_t parse_uint(const std::string &s, size_t fallback = 0) {
  if (s.empty())
    return fallback;
  int v = std::atoi(s.c_str());
  return v >= 0 ? static_cast<size_t>(v) : fallback;
}

// resolve a model: parse + mmap on first access, cache in state
static std::shared_ptr<model_state>
resolve_model(const std::shared_ptr<server_state> &state,
              const std::string &model_param, httplib::Response &res) {
  if (model_param.empty()) {
    set_error_response(res, "missing model parameter", 400);
    return nullptr;
  }

  // check cache
  {
    std::lock_guard<std::mutex> lock(state->mutex);
    auto it = state->models.find(model_param);
    if (it != state->models.end())
      return it->second;
  }

  // resolve path
  std::error_code ec;
  fs::path abs = fs::weakly_canonical(state->root / model_param, ec);
  if (ec || !fs::exists(abs, ec)) {
    set_error_response(res, "model not found", 404);
    return nullptr;
  }

  auto ms = std::make_shared<model_state>();
  ms->model_path = abs.string();
  ms->relative_path = model_param;

  // open GGUF: parse metadata + create tensor structs (no data allocated)
  struct ggml_context *tensor_ctx_raw = nullptr;
  struct gguf_init_params params = {};
  params.no_alloc = true;
  params.ctx = &tensor_ctx_raw;

  struct gguf_context *gguf_raw =
      gguf_init_from_file(ms->model_path.c_str(), params);
  if (!gguf_raw) {
    set_error_response(res, "failed to parse GGUF file", 500);
    return nullptr;
  }
  ms->gguf_ctx.reset(gguf_raw);
  ms->tensor_ctx.reset(tensor_ctx_raw);

  // mmap the file for tensor data access
  std::string mmap_error;
  if (!mmap_open(ms->model_path.c_str(), ms->mmap, mmap_error)) {
    set_error_response(res, mmap_error.c_str(), 500);
    return nullptr;
  }

  // populate tensor entries from ggml_context
  size_t data_offset = gguf_get_data_offset(gguf_raw);
  int64_t n_tensors = gguf_get_n_tensors(gguf_raw);
  ms->tensors.resize(static_cast<size_t>(n_tensors));
  ms->stats_cache.resize(static_cast<size_t>(n_tensors));

  for (int64_t i = 0; i < n_tensors; ++i) {
    size_t idx = static_cast<size_t>(i);
    tensor_entry &te = ms->tensors[idx];

    te.name = gguf_get_tensor_name(gguf_raw, i);
    te.type = gguf_get_tensor_type(gguf_raw, i);
    te.offset = gguf_get_tensor_offset(gguf_raw, i);
    te.n_bytes = gguf_get_tensor_size(gguf_raw, i);
    te.file_offset = data_offset + te.offset;
    te.block_size = static_cast<size_t>(ggml_blck_size(te.type));
    te.type_size = ggml_type_size(te.type);

    // get shape from ggml_context tensor struct
    struct ggml_tensor *gt =
        ggml_get_tensor(tensor_ctx_raw, te.name);
    if (gt) {
      te.n_dims = ggml_n_dims(gt);
      te.n_elements = ggml_nelements(gt);
      for (int d = 0; d < GGML_MAX_DIMS; ++d) {
        te.ne[d] = gt->ne[d];
      }
    } else {
      // fallback: compute from byte size
      te.n_dims = 1;
      te.n_elements =
          (te.block_size > 0 && te.type_size > 0)
              ? static_cast<int64_t>((te.n_bytes / te.type_size) *
                                     te.block_size)
              : 0;
      te.ne[0] = te.n_elements;
    }

    te.layout = compute_layout(te);
    ms->tensor_index[te.name] = idx;
    ms->stats_cache[idx].resize(te.layout.depth);
  }

  // store in cache
  {
    std::lock_guard<std::mutex> lock(state->mutex);
    auto &slot = state->models[model_param];
    if (!slot)
      slot = ms;
    ms = slot;
  }

  fprintf(stderr, "[Viewer] Loaded %s: %zu tensors, %zu KV pairs\n",
          ms->model_path.c_str(), ms->tensors.size(),
          static_cast<size_t>(gguf_get_n_kv(ms->gguf_ctx.get())));

  return ms;
}

// find a tensor by name, returns its index or SIZE_MAX if not found
static size_t find_tensor(const std::shared_ptr<model_state> &ms,
                          const std::string &name) {
  auto it = ms->tensor_index.find(name);
  if (it != ms->tensor_index.end())
    return it->second;
  return SIZE_MAX;
}

// convert a KV scalar value to JSON
static json kv_scalar_to_json(const struct gguf_context *ctx, int64_t key_id,
                              enum gguf_type type) {
  switch (type) {
  case GGUF_TYPE_UINT8:
    return gguf_get_val_u8(ctx, key_id);
  case GGUF_TYPE_INT8:
    return gguf_get_val_i8(ctx, key_id);
  case GGUF_TYPE_UINT16:
    return gguf_get_val_u16(ctx, key_id);
  case GGUF_TYPE_INT16:
    return gguf_get_val_i16(ctx, key_id);
  case GGUF_TYPE_UINT32:
    return gguf_get_val_u32(ctx, key_id);
  case GGUF_TYPE_INT32:
    return gguf_get_val_i32(ctx, key_id);
  case GGUF_TYPE_FLOAT32:
    return gguf_get_val_f32(ctx, key_id);
  case GGUF_TYPE_UINT64:
    return gguf_get_val_u64(ctx, key_id);
  case GGUF_TYPE_INT64:
    return gguf_get_val_i64(ctx, key_id);
  case GGUF_TYPE_FLOAT64:
    return gguf_get_val_f64(ctx, key_id);
  case GGUF_TYPE_BOOL:
    return gguf_get_val_bool(ctx, key_id);
  case GGUF_TYPE_STRING:
    return gguf_get_val_str(ctx, key_id);
  default:
    return nullptr;
  }
}

// convert array preview to JSON (up to limit elements)
static json array_preview_to_json(const struct gguf_context *ctx,
                                  int64_t key_id, size_t limit) {
  json out = json::array();
  enum gguf_type arr_type = gguf_get_arr_type(ctx, key_id);
  size_t n = gguf_get_arr_n(ctx, key_id);
  if (n > limit)
    n = limit;

  if (arr_type == GGUF_TYPE_STRING) {
    for (size_t i = 0; i < n; ++i) {
      out.push_back(gguf_get_arr_str(ctx, key_id, i));
    }
    return out;
  }

  // numeric and bool arrays: read from raw data pointer
  const uint8_t *data =
      static_cast<const uint8_t *>(gguf_get_arr_data(ctx, key_id));
  if (!data)
    return out;

  for (size_t i = 0; i < n; ++i) {
    switch (arr_type) {
    case GGUF_TYPE_UINT8:
      out.push_back(data[i]);
      break;
    case GGUF_TYPE_INT8:
      out.push_back(static_cast<int8_t>(data[i]));
      break;
    case GGUF_TYPE_UINT16: {
      uint16_t v;
      memcpy(&v, data + i * 2, 2);
      out.push_back(v);
    } break;
    case GGUF_TYPE_INT16: {
      int16_t v;
      memcpy(&v, data + i * 2, 2);
      out.push_back(v);
    } break;
    case GGUF_TYPE_UINT32: {
      uint32_t v;
      memcpy(&v, data + i * 4, 4);
      out.push_back(v);
    } break;
    case GGUF_TYPE_INT32: {
      int32_t v;
      memcpy(&v, data + i * 4, 4);
      out.push_back(v);
    } break;
    case GGUF_TYPE_FLOAT32: {
      float v;
      memcpy(&v, data + i * 4, 4);
      out.push_back(v);
    } break;
    case GGUF_TYPE_UINT64: {
      uint64_t v;
      memcpy(&v, data + i * 8, 8);
      out.push_back(v);
    } break;
    case GGUF_TYPE_INT64: {
      int64_t v;
      memcpy(&v, data + i * 8, 8);
      out.push_back(v);
    } break;
    case GGUF_TYPE_FLOAT64: {
      double v;
      memcpy(&v, data + i * 8, 8);
      out.push_back(v);
    } break;
    case GGUF_TYPE_BOOL:
      out.push_back(data[i] != 0);
      break;
    default:
      break;
    }
  }
  return out;
}

void setup_routes(httplib::Server &server,
                  std::shared_ptr<server_state> state) {

  // GET /api/models
  server.Get(
      "/api/models", [state](const httplib::Request &, httplib::Response &res) {
        auto files = scan_gguf_files(state->root);
        json items = json::array();
        for (const auto &path : files) {
          json item;
          item["path"] = path;
          auto slash = path.rfind('/');
          item["name"] =
              (slash != std::string::npos) ? path.substr(slash + 1) : path;
          std::error_code ec;
          item["size"] = ec ? 0 : fs::file_size(state->root / path, ec);
          items.push_back(std::move(item));
        }
        json body;
        body["root"] = state->root.generic_string();
        body["items"] = std::move(items);
        set_json_response(res, body);
      });

  // GET /api/info
  server.Get("/api/info",
             [state](const httplib::Request &req, httplib::Response &res) {
               auto ms = resolve_model(state, get_model_param(req), res);
               if (!ms)
                 return;
               const auto *ctx = ms->gguf_ctx.get();

               bool has_tokens = false;
               int64_t total_tokens = 0;
               int64_t tok_key =
                   gguf_find_key(ctx, "tokenizer.ggml.tokens");
               if (tok_key >= 0 &&
                   gguf_get_kv_type(ctx, tok_key) == GGUF_TYPE_ARRAY) {
                 has_tokens = true;
                 total_tokens =
                     static_cast<int64_t>(gguf_get_arr_n(ctx, tok_key));
               }

               json body;
               body["modelPath"] = ms->model_path;
               body["relativePath"] = ms->relative_path;
               body["fileSize"] = ms->mmap.size;
               body["nKv"] = gguf_get_n_kv(ctx);
               body["nTensors"] = ms->tensors.size();
               body["ggufVersion"] = gguf_get_version(ctx);
               body["alignment"] = gguf_get_alignment(ctx);
               body["dataOffset"] = gguf_get_data_offset(ctx);
               body["tokenizer"] = {{"hasTokens", has_tokens},
                                    {"totalTokens", total_tokens}};
               set_json_response(res, body);
             });

  // GET /api/kv
  server.Get(
      "/api/kv", [state](const httplib::Request &req, httplib::Response &res) {
        auto ms = resolve_model(state, get_model_param(req), res);
        if (!ms)
          return;
        const auto *ctx = ms->gguf_ctx.get();
        size_t preview_limit = parse_uint(req.get_param_value("preview"), 8);
        int64_t n_kv = gguf_get_n_kv(ctx);

        json kvs = json::array();
        for (int64_t i = 0; i < n_kv; ++i) {
          json node;
          node["key"] = gguf_get_key(ctx, i);
          enum gguf_type type = gguf_get_kv_type(ctx, i);
          node["type"] = gguf_type_name(type);

          if (type == GGUF_TYPE_ARRAY) {
            size_t arr_n = gguf_get_arr_n(ctx, i);
            node["length"] = arr_n;
            node["arrayType"] =
                gguf_type_name(gguf_get_arr_type(ctx, i));
            node["preview"] =
                array_preview_to_json(ctx, i, preview_limit);
            node["previewTruncated"] = arr_n > preview_limit;
          } else {
            node["value"] = kv_scalar_to_json(ctx, i, type);
          }
          kvs.push_back(std::move(node));
        }
        set_json_response(res, kvs);
      });

  // GET /api/tensors
  server.Get(R"(/api/tensors$)", [state](const httplib::Request &req,
                                         httplib::Response &res) {
    auto ms = resolve_model(state, get_model_param(req), res);
    if (!ms)
      return;

    json tensors = json::array();
    for (size_t i = 0; i < ms->tensors.size(); ++i) {
      const auto &t = ms->tensors[i];
      json node;
      node["name"] = t.name;
      node["type"] = ggml_type_name(t.type);
      node["nElements"] = t.n_elements;
      node["nBytes"] = t.n_bytes;
      node["offset"] = t.offset;
      node["fileOffset"] = t.file_offset;

      // shape as array (only non-trivial dims)
      json shape = json::array();
      for (int d = 0; d < t.n_dims; ++d) {
        shape.push_back(t.ne[d]);
      }
      node["shape"] = std::move(shape);
      node["ndim"] = t.n_dims;
      node["layout"] = {{"width", t.layout.width},
                         {"height", t.layout.height},
                         {"depth", t.layout.depth}};
      node["blockSize"] = t.block_size;
      tensors.push_back(std::move(node));
    }
    set_json_response(res, tensors);
  });

  // GET /api/tensors/:name/raw (tile of dequantized float values)
  server.Get(R"(/api/tensors/(.+)/raw)", [state](const httplib::Request &req,
                                                 httplib::Response &res) {
    auto ms = resolve_model(state, get_model_param(req), res);
    if (!ms)
      return;
    size_t ti = find_tensor(ms, req.matches[1].str());
    if (ti == SIZE_MAX) {
      set_error_response(res, "tensor not found", 404);
      return;
    }
    const auto &td = ms->tensors[ti];

    size_t slice = parse_uint(req.get_param_value("slice"));
    size_t x = parse_uint(req.get_param_value("x"));
    size_t y = parse_uint(req.get_param_value("y"));
    size_t w = parse_uint(req.get_param_value("width"), 1024);
    size_t h = parse_uint(req.get_param_value("height"), 1024);

    tensor_tile tile;
    std::string error;
    if (!tensor_read_tile(ms->mmap, td, slice, x, y, w, h, tile, error)) {
      set_error_response(res, error.c_str(), 500);
      return;
    }

    const auto &layout = td.layout;
    json body;
    body["layout"] = {{"width", layout.width},
                      {"height", layout.height},
                      {"depth", layout.depth}};
    body["origin"] = {{"x", tile.x}, {"y", tile.y}, {"slice", tile.slice}};
    body["viewport"] = {{"width", tile.width}, {"height", tile.height}};
    body["min"] = tile.valid > 0 ? json(tile.min) : json(nullptr);
    body["max"] = tile.valid > 0 ? json(tile.max) : json(nullptr);

    json values = json::array();
    for (size_t i = 0; i < tile.values.size(); ++i) {
      if (i < tile.mask.size() && tile.mask[i]) {
        values.push_back(tile.values[i]);
      } else {
        values.push_back(nullptr);
      }
    }
    body["values"] = std::move(values);
    set_json_response(res, body);
  });

  // GET /api/tensors/:name/slice/properties
  server.Get(
      R"(/api/tensors/(.+)/slice/properties)",
      [state](const httplib::Request &req, httplib::Response &res) {
        auto ms = resolve_model(state, get_model_param(req), res);
        if (!ms)
          return;
        size_t ti = find_tensor(ms, req.matches[1].str());
        if (ti == SIZE_MAX) {
          set_error_response(res, "tensor not found", 404);
          return;
        }
        const auto &td = ms->tensors[ti];

        size_t slice = parse_uint(req.get_param_value("slice"));
        slice_stats stats;
        std::string error;
        if (!tensor_slice_stats(ms->mmap, td, slice, stats, error)) {
          set_error_response(res, error.c_str(), 500);
          return;
        }

        json body;
        body["slice"] = slice;
        body["valid"] = stats.valid;
        body["min"] = (stats.computed && stats.valid > 0) ? json(stats.min)
                                                          : json(nullptr);
        body["max"] = (stats.computed && stats.valid > 0) ? json(stats.max)
                                                          : json(nullptr);
        body["percentiles"] = {
            {"lowerPercent", stats.p_lower},
            {"upperPercent", stats.p_upper},
            {"lower", (stats.computed && stats.valid > 0) ? json(stats.min)
                                                          : json(nullptr)},
            {"upper", (stats.computed && stats.valid > 0) ? json(stats.max)
                                                          : json(nullptr)}};
        set_json_response(res, body);
      });

  // GET /api/tensors/:name/value
  server.Get(R"(/api/tensors/(.+)/value)", [state](const httplib::Request &req,
                                                   httplib::Response &res) {
    auto ms = resolve_model(state, get_model_param(req), res);
    if (!ms)
      return;
    size_t ti = find_tensor(ms, req.matches[1].str());
    if (ti == SIZE_MAX) {
      set_error_response(res, "tensor not found", 404);
      return;
    }
    const auto &td = ms->tensors[ti];

    size_t slice = parse_uint(req.get_param_value("slice"));
    size_t x = parse_uint(req.get_param_value("x"));
    size_t y = parse_uint(req.get_param_value("y"));

    element_details details;
    std::string error;
    if (!tensor_element_at(ms->mmap, td, slice, x, y, details, error)) {
      set_error_response(res, error.c_str(), 400);
      return;
    }

    json body;
    body["tensor"] = td.name;
    body["type"] = ggml_type_name(td.type);
    body["coordinate"] = {{"x", x}, {"y", y}, {"slice", slice}};
    body["index"] = details.element_index;
    body["count"] = details.element_count;
    body["block"] = {{"index", details.block_index},
                     {"offset", details.index_in_block},
                     {"size", td.block_size}};
    body["tensorOffset"] = details.tensor_byte_offset;
    body["fileOffset"] = details.file_byte_offset;
    body["value"] = details.valid ? json(details.value) : json(nullptr);
    set_json_response(res, body);
  });

  // GET /api/tensors/:name/histogram
  server.Get(R"(/api/tensors/(.+)/histogram)",
             [state](const httplib::Request &req, httplib::Response &res) {
               auto ms = resolve_model(state, get_model_param(req), res);
               if (!ms)
                 return;
               size_t ti = find_tensor(ms, req.matches[1].str());
               if (ti == SIZE_MAX) {
                 set_error_response(res, "tensor not found", 404);
                 return;
               }
               const auto &td = ms->tensors[ti];

               size_t slice = parse_uint(req.get_param_value("slice"));
               size_t width = parse_uint(req.get_param_value("width"), 256);

               tensor_histogram hist;
               std::string error;
               if (!tensor_slice_histogram(ms->mmap, td, slice, width, hist,
                                           error)) {
                 set_error_response(res, error.c_str(), 500);
                 return;
               }

               json body;
               body["width"] = width;
               body["slice"] = hist.slice;
               body["range"] = {{"min", hist.range_min},
                                {"max", hist.range_max}};
               body["maxCount"] = hist.max_bin;
               body["total"] = hist.total;
               body["clippedLow"] = hist.clipped_lo;
               body["clippedHigh"] = hist.clipped_hi;
               body["zeroCount"] = hist.zero_count;
               body["bins"] = hist.bins;
               set_json_response(res, body);
             });

  // GET /api/tokenizer
  server.Get(R"(/api/tokenizer$)", [state](const httplib::Request &req,
                                           httplib::Response &res) {
    auto ms = resolve_model(state, get_model_param(req), res);
    if (!ms)
      return;
    const auto *ctx = ms->gguf_ctx.get();

    int64_t tok_key = gguf_find_key(ctx, "tokenizer.ggml.tokens");
    if (tok_key < 0 || gguf_get_kv_type(ctx, tok_key) != GGUF_TYPE_ARRAY) {
      json body;
      body["hasTokenizer"] = false;
      body["total"] = 0;
      body["offset"] = 0;
      body["limit"] = 0;
      body["items"] = json::array();
      set_json_response(res, body);
      return;
    }

    size_t total = gguf_get_arr_n(ctx, tok_key);
    size_t offset = parse_uint(req.get_param_value("offset"));
    size_t limit = parse_uint(req.get_param_value("limit"), 256);
    if (offset > total)
      offset = total;
    if (offset + limit > total)
      limit = total - offset;

    int64_t scores_key = gguf_find_key(ctx, "tokenizer.ggml.scores");
    int64_t types_key = gguf_find_key(ctx, "tokenizer.ggml.token_type");

    // raw data pointers for scores and types arrays
    const float *scores_data = nullptr;
    size_t scores_n = 0;
    if (scores_key >= 0 &&
        gguf_get_kv_type(ctx, scores_key) == GGUF_TYPE_ARRAY &&
        gguf_get_arr_type(ctx, scores_key) == GGUF_TYPE_FLOAT32) {
      scores_data = static_cast<const float *>(
          gguf_get_arr_data(ctx, scores_key));
      scores_n = gguf_get_arr_n(ctx, scores_key);
    }

    const int32_t *types_data = nullptr;
    size_t types_n = 0;
    if (types_key >= 0 &&
        gguf_get_kv_type(ctx, types_key) == GGUF_TYPE_ARRAY &&
        gguf_get_arr_type(ctx, types_key) == GGUF_TYPE_INT32) {
      types_data = static_cast<const int32_t *>(
          gguf_get_arr_data(ctx, types_key));
      types_n = gguf_get_arr_n(ctx, types_key);
    }

    json items = json::array();
    for (size_t i = offset; i < offset + limit; ++i) {
      json item;
      item["index"] = i;
      item["token"] = gguf_get_arr_str(ctx, tok_key, i);
      if (scores_data && i < scores_n)
        item["score"] = scores_data[i];
      if (types_data && i < types_n)
        item["tokenType"] = types_data[i];
      items.push_back(std::move(item));
    }

    json body;
    body["hasTokenizer"] = true;
    body["total"] = total;
    body["offset"] = offset;
    body["limit"] = limit;
    body["items"] = std::move(items);
    set_json_response(res, body);
  });

  // GET /health
  server.Get("/health", [](const httplib::Request &, httplib::Response &res) {
    res.set_content("{\"status\":\"ok\"}", "application/json");
  });
}
