#include <cstdint>
#include <fstream>
#include <iostream>
#include <stdexcept>
#include <string>
#include <vector>

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

int main(int argc, char **argv) try {
  if (argc != 6)
    return 2;
  const std::size_t items = std::stoul(argv[2]);
  const unsigned sourceTile = std::stoul(argv[3]);
  const unsigned firstDestinationTile = std::stoul(argv[4]);
  const unsigned receiverPairs = std::stoul(argv[5]);
  if (receiverPairs == 0)
    return 2;

  poplar::DeviceManager manager;
  auto devices = manager.getDevices(poplar::TargetType::IPU, 1);
  if (devices.empty() || !devices.front().attach())
    return 3;
  poplar::Graph graph(devices.front().getTarget());
  graph.addCodelets("codelets64.cpp", poplar::CodeletFileType::Auto, "-O2",
                    "ipu2");

  auto source =
      graph.addVariable(poplar::UNSIGNED_LONGLONG, {items}, "source64");
  graph.setTileMapping(source, sourceTile);
  graph.createHostWrite("source-write", source);

  std::vector<poplar::Tensor> destinations;
  std::vector<poplar::Tensor> sourceCopies;
  std::vector<unsigned> destinationTiles;
  destinations.reserve(receiverPairs * 2);
  sourceCopies.reserve(receiverPairs * 2);
  for (unsigned tile = firstDestinationTile;
       destinationTiles.size() != receiverPairs * 2; tile += 2) {
    if (tile == (sourceTile & ~1u))
      continue;
    destinationTiles.push_back(tile);
    destinationTiles.push_back(tile + 1);
  }
  for (unsigned index = 0; index != destinationTiles.size(); ++index) {
    const auto name = "destination-" + std::to_string(index);
    auto destination =
        graph.addVariable(poplar::UNSIGNED_LONGLONG, {items}, name);
    graph.setTileMapping(destination, destinationTiles[index]);
    graph.createHostRead(name + "-read", destination);
    destinations.push_back(destination);
    sourceCopies.push_back(source);
  }

  auto keep = graph.addComputeSet("keep-destination-pairs-separated");
  for (unsigned pair = 0; pair != receiverPairs; ++pair) {
    auto vertex = graph.addVertex(keep, "KeepPair64Separated");
    graph.connect(vertex["first"], destinations[pair * 2]);
    graph.connect(vertex["second"], destinations[pair * 2 + 1]);
    graph.setTileMapping(vertex, destinationTiles[pair * 2]);
  }

  poplar::program::Sequence program;
  program.add(poplar::program::Copy(poplar::concat(sourceCopies),
                                    poplar::concat(destinations)));
  program.add(poplar::program::Execute(keep));
  poplar::Engine engine(graph, program);
  std::ofstream executable(argv[1], std::ios::binary);
  engine.serializeExecutable(executable);
  executable.close();
  engine.load(devices.front());

  std::vector<std::uint64_t> input(items);
  for (std::size_t index = 0; index != items; ++index)
    input[index] = 0x6400000000000000ull ^
                   (index * 0x9e3779b97f4a7c15ull);
  engine.writeTensor("source-write", input.data(), input.data() + input.size());
  engine.run(0);
  for (unsigned index = 0; index != receiverPairs * 2; ++index) {
    std::vector<std::uint64_t> result(items);
    const auto name = "destination-" + std::to_string(index) + "-read";
    engine.readTensor(name, result.data(), result.data() + result.size());
    if (result != input)
      throw std::runtime_error("paired 64-bit multicast result mismatch");
  }
  std::cout << "paired64-multicast hardware=PASS items=" << items
            << " receiverPairs=" << receiverPairs << '\n';
  return 0;
} catch (const std::exception &error) {
  std::cerr << error.what() << '\n';
  return 1;
}
