// Copyright (c) 2023 Graphcore Ltd. All rights reserved.

#include <cstdlib>
#include <cstdint>
#include <fstream>
#include <iostream>
#include <regex>
#include <string>
#include <vector>

#include <poplar/DeviceManager.hpp>
#include <poplar/Engine.hpp>
#include <poplar/Graph.hpp>
#include <poplar/Program.hpp>

#include "JDL.hpp"

using namespace poplar;

namespace {
constexpr unsigned numDataTiles = 1;
constexpr unsigned numElementsPerDataTile = 16;
constexpr unsigned lookupSize = 3;
constexpr unsigned receiverTileId = 1286;
constexpr int tileSelectorValue = 0;
constexpr int elementSelectorValue = 5;

std::vector<std::uint32_t> parseHexWords(const std::string &path) {
  std::ifstream in(path);
  if (!in) {
    throw std::runtime_error("failed to open plan file: " + path);
  }
  std::string text((std::istreambuf_iterator<char>(in)), std::istreambuf_iterator<char>());
  std::regex hexRegex("0x([0-9a-fA-F]{1,8})");
  std::vector<std::uint32_t> words;
  for (std::sregex_iterator it(text.begin(), text.end(), hexRegex), end; it != end; ++it) {
    words.push_back(static_cast<std::uint32_t>(std::stoul((*it)[1].str(), nullptr, 16)));
  }
  return words;
}
} // namespace

int main(int argc, char **argv) {
  if (argc != 2) {
    std::cerr << "usage: " << argv[0] << " <plan-file-from-jdl_plan_codegen.py>\n";
    return 2;
  }
  const std::string planPath = argv[1];

  // Setup device and graph.
  auto devManager = DeviceManager();
  auto devs = devManager.getDevices(TargetType::IPU, 1);
  if (devs.empty()) {
    std::cerr << "no IPU device available\n";
    return 3;
  }
  Device &device = devs[0];
  if (!device.attach()) {
    std::cerr << "failed to attach device\n";
    return 4;
  }
  Target target = device.getTarget();
  Graph graph(target);

  // Data setup.
  std::srand(0);
  std::vector<int> data_h(numDataTiles * numElementsPerDataTile);
  for (unsigned i = 0; i < data_h.size(); ++i) {
    data_h[i] = std::rand() % 100;
  }
  Tensor data = graph.addVariable(INT, {numDataTiles, numElementsPerDataTile}, "data");
  graph.setInitialValue<int>(data, data_h);
  for (unsigned tile = 0; tile < numDataTiles; ++tile) {
    graph.setTileMapping(data[tile], tile);
  }

  Tensor tileSelector = graph.addVariable(INT, {}, "tileSelector");
  Tensor elementSelector = graph.addVariable(INT, {}, "elementSelector");
  Tensor result = graph.addVariable(INT, {lookupSize}, "result");
  graph.setTileMapping(tileSelector, receiverTileId);
  graph.setTileMapping(elementSelector, receiverTileId);
  graph.setTileMapping(result, receiverTileId);

  JDL::Programs jdlPrograms = JDL::createPrograms(graph, data, tileSelector, elementSelector, result);
  const unsigned expectedPlanWords = jdlPrograms.numActiveTiles * 9;

  auto parsedWords = parseHexWords(planPath);
  if (parsedWords.size() < expectedPlanWords) {
    std::cerr << "plan file has " << parsedWords.size() << " hex words, need at least "
              << expectedPlanWords << "\n";
    return 5;
  }
  if (parsedWords.size() > expectedPlanWords) {
    // JSON output can contain extra metadata words before row arrays.
    parsedWords = std::vector<std::uint32_t>(
        parsedWords.end() - expectedPlanWords, parsedWords.end());
  }

  graph.createHostWrite("planWrite", jdlPrograms.planBuf.flatten());
  graph.createHostWrite("tileSelectorWrite", tileSelector);
  graph.createHostWrite("elementSelectorWrite", elementSelector);
  graph.createHostRead("resultRead", result);

  program::Sequence mainProgram({jdlPrograms.exchange});
  Engine engine(graph, mainProgram);
  engine.load(device);

  engine.writeTensor("planWrite", parsedWords.data(), parsedWords.data() + parsedWords.size());
  int tileSel = tileSelectorValue;
  int elemSel = elementSelectorValue;
  engine.writeTensor("tileSelectorWrite", &tileSel, &tileSel + 1);
  engine.writeTensor("elementSelectorWrite", &elemSel, &elemSel + 1);
  engine.run(0);

  std::vector<int> result_h(lookupSize, -1);
  engine.readTensor("resultRead", result_h.data(), result_h.data() + result_h.size());

  std::cout << "result:";
  for (auto x : result_h) std::cout << " " << x;
  std::cout << "\nexpected:";
  for (unsigned i = 0; i < lookupSize; ++i) {
    std::cout << " " << data_h[elementSelectorValue + i];
  }
  std::cout << "\n";

  bool ok = true;
  for (unsigned i = 0; i < lookupSize; ++i) {
    if (result_h[i] != data_h[elementSelectorValue + i]) {
      ok = false;
      break;
    }
  }
  std::cout << (ok ? "PASS\n" : "FAIL\n");
  return ok ? 0 : 1;
}
