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

#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() {
  constexpr std::size_t numWords = 16;
  std::array<std::uint32_t, numWords> input{};
  std::array<std::uint32_t, numWords> output{};
  for (std::size_t i = 0; i < numWords; ++i)
    input[i] = 0x62000000u + static_cast<std::uint32_t>(i * 0x10101u);

  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;
  }
  auto &device = devices.front();

  poplar::Graph graph(device.getTarget());
  auto source = graph.addVariable(poplar::UNSIGNED_INT, {numWords}, "source");
  auto destination =
      graph.addVariable(poplar::UNSIGNED_INT, {numWords}, "destination");
  graph.setTileMapping(source, 0);
  graph.setTileMapping(destination, 1);
  graph.createHostWrite("source-write", source);
  graph.createHostRead("destination-read", destination);

  poplar::program::Sequence program;
  program.add(poplar::program::Copy(source, destination));
  poplar::OptionFlags options;
  if (std::getenv("IPU_HOST_EXCHANGE_DEBUG")) {
    options.set("debug.printHostExchangeSchedule", "0,1");
    options.set("debug.dumpHostExchangePackets", "true");
  }
  poplar::Engine engine(graph, program, options);
  std::ofstream executable("sdk_host_tensor.popef", std::ios::binary);
  engine.serializeExecutable(executable);
  executable.close();
  if (!executable) {
    std::cerr << "failed to serialize executable\n";
    return 3;
  }

  std::cerr << "ORACLE load begin\n";
  engine.load(device);
  snapshot("after-load");
  std::cerr << "ORACLE load end\nORACLE writeTensor begin\n";
  engine.writeTensor("source-write", input.data(), input.data() + input.size());
  snapshot("after-write");
  std::cerr << "ORACLE writeTensor end\nORACLE run begin\n";
  engine.run(0);
  snapshot("after-run");
  std::cerr << "ORACLE run end\nORACLE readTensor begin\n";
  engine.readTensor("destination-read", output.data(),
                    output.data() + output.size());
  snapshot("after-read");
  std::cerr << "ORACLE readTensor end\n";

  if (output != input) {
    std::cerr << "tensor round trip corrupted data\n";
    return 2;
  }
  std::cout << "tensor round trip passed: " << std::hex << output.front()
            << " .. " << output.back() << '\n';
  return 0;
}
