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

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

extern "C" void ipuIoctlTraceSnapshot(const char *) __attribute__((weak));

namespace {
void snapshot(const char *phase) {
  if (ipuIoctlTraceSnapshot)
    ipuIoctlTraceSnapshot(phase);
}
} // namespace

int main(int argc, char **argv) {
  const std::size_t numWords =
      argc >= 2 ? std::strtoul(argv[1], nullptr, 0) : 16;
  const unsigned tile = argc >= 3 ? std::strtoul(argv[2], nullptr, 0) : 0;
  if (numWords == 0) {
    std::cerr << "word count must be nonzero\n";
    return 2;
  }
  std::vector<std::uint32_t> input(numWords);
  std::vector<std::uint32_t> output(numWords);
  for (std::size_t i = 0; i < numWords; ++i)
    input[i] = 0x71000000u + static_cast<std::uint32_t>(i * 0x10203u);

  poplar::DeviceManager manager;
  auto devices = manager.getDevices(poplar::TargetType::IPU, 1);
  if (devices.empty() || !devices.front().attach()) {
    std::cerr << "failed to attach one IPU\n";
    return 1;
  }

  poplar::Graph graph(devices.front().getTarget());
  if (tile >= graph.getTarget().getNumTiles()) {
    std::cerr << "tile is out of range\n";
    return 2;
  }
  auto tensor = graph.addVariable(poplar::UNSIGNED_INT, {numWords}, "tensor");
  graph.setTileMapping(tensor, tile);
  graph.createHostWrite("tensor-write", tensor);
  graph.createHostRead("tensor-read", tensor);

  poplar::OptionFlags options;
  if (std::getenv("IPU_HOST_EXCHANGE_DEBUG")) {
    const char *tiles = std::getenv("IPU_HOST_EXCHANGE_TILES");
    options.set("debug.printHostExchangeSchedule", tiles ? tiles : "0,1");
    options.set("debug.dumpHostExchangePackets", "true");
  }
  poplar::Engine engine(graph, poplar::program::Sequence{}, options);
  std::ofstream executable("sdk_local_host_tensor.popef", std::ios::binary);
  engine.serializeExecutable(executable);
  executable.close();

  engine.load(devices.front());
  snapshot("after-load");
  engine.writeTensor("tensor-write", input.data(), input.data() + input.size());
  snapshot("after-write");
  engine.readTensor("tensor-read", output.data(), output.data() + output.size());
  snapshot("after-read");
  if (output != input) {
    std::cerr << "local host tensor round trip corrupted data\n";
    return 3;
  }
  std::cout << "host tensor round trip passed on tile " << tile << '\n';
  return 0;
}
