#include "exchange_plan.hpp"

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

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

using namespace poplar;

namespace {
constexpr unsigned guardWords = 16;
constexpr int guard = static_cast<int>(0x6badf00dU);

unsigned arg(char **argv, int index) {
  return static_cast<unsigned>(std::strtoul(argv[index], nullptr, 0));
}

int value(unsigned index) {
  return static_cast<int>(0x12340000u ^ (index * 0x9e3779b9u));
}
} // namespace

int main(int argc, char **argv) {
  if (argc < 4) {
    std::cerr << "usage: proper_multicast SENDER COUNT RECEIVER...\n";
    return 2;
  }
  const unsigned sender = arg(argv, 1);
  const unsigned count = arg(argv, 2);
  std::vector<unsigned> receivers;
  for (int i = 3; i < argc; ++i) receivers.push_back(arg(argv, i));
  constexpr unsigned iterations = 3;

  xcom::FanOutPlan multicast;
  try {
    multicast = xcom::assembleMulticast(sender, receivers, count);
  } catch (const std::exception &error) {
    std::cerr << "cannot assemble multicast: " << error.what() << '\n';
    return 2;
  }

  DeviceManager manager;
  auto devices = manager.getDevices(TargetType::IPU, 1);
  if (devices.empty() || !devices[0].attach()) {
    std::cerr << "failed to attach an IPU\n";
    return 3;
  }

  Device &device = devices[0];
  Graph graph(device.getTarget());
  graph.addCodelets(std::vector<std::string>{"no_blob_codelets.cpp"}, "-O2",
                    "ipu2");

  const unsigned sourceWords = count + 2 * guardWords + iterations;
  std::vector<int> sourceHost(sourceWords);
  for (unsigned i = 0; i < sourceWords; ++i) sourceHost[i] = value(i);
  Tensor source = graph.addVariable(INT, {sourceWords}, "source");
  graph.setInitialValue<int>(source, sourceHost);
  graph.setTileMapping(source, sender);

  const unsigned resultStride = count + 2 * guardWords;
  Tensor results =
      graph.addVariable(INT, {receivers.size(), resultStride}, "results");
  graph.setInitialValue<int>(results,
                             std::vector<int>(results.numElements(), guard));
  for (unsigned i = 0; i < receivers.size(); ++i)
    graph.setTileMapping(results[i], receivers[i]);

  Tensor plans =
      graph.addVariable(UNSIGNED_INT, {1 + receivers.size(), 9}, "plans");
  Tensor dummy =
      graph.addVariable(UNSIGNED_INT, {1 + receivers.size(), 1}, "dummy");
  graph.setTileMapping(plans[0], sender);
  graph.setTileMapping(dummy[0], sender);
  for (unsigned i = 0; i < receivers.size(); ++i) {
    graph.setTileMapping(plans[1 + i], receivers[i]);
    graph.setTileMapping(dummy[1 + i], receivers[i]);
  }

  Tensor sourceOffset = graph.addVariable(INT, {}, "sourceOffset");
  graph.setTileMapping(sourceOffset, sender);
  ComputeSet exchange = graph.addComputeSet("properMulticast");
  auto send = graph.addVertex(exchange, "NoBlobMulticastSend",
                              {{"planBuf", plans[0]},
                               {"nonexecutableDummy", dummy[0]},
                               {"elementSelector", sourceOffset},
                               {"data", source}});
  graph.setTileMapping(send, sender);
  for (unsigned i = 0; i < receivers.size(); ++i) {
    auto destination = results[i].slice(guardWords, guardWords + count);
    auto receive = graph.addVertex(exchange, "NoBlobMulticastRecv",
                                   {{"planBuf", plans[1 + i]},
                                    {"nonexecutableDummy", dummy[1 + i]},
                                    {"result", destination}});
    graph.setTileMapping(receive, receivers[i]);
  }
  const unsigned numTiles = graph.getTarget().getNumTiles();
  for (unsigned tile = 0; tile < numTiles; ++tile) {
    bool active = tile == sender;
    for (unsigned receiver : receivers) active |= tile == receiver;
    if (active) continue;
    auto inactive = graph.addVertex(exchange, "NoBlobNonParticipation");
    graph.setTileMapping(inactive, tile);
  }

  graph.createHostWrite("plansWrite", plans.flatten());
  graph.createHostWrite("sourceOffsetWrite", sourceOffset);
  graph.createHostRead("resultsRead", results.flatten());
  graph.createHostRead("sourceRead", source);
  Engine engine(graph, program::Sequence({program::Sync(SyncType::INTERNAL),
                                          program::Execute(exchange),
                                          program::Sync(SyncType::INTERNAL)}));
  engine.load(device);

  std::vector<std::uint32_t> words;
  words.insert(words.end(), multicast.sender.begin(), multicast.sender.end());
  for (const auto &row : multicast.receivers)
    words.insert(words.end(), row.begin(), row.end());
  engine.writeTensor("plansWrite", words.data(), words.data() + words.size());

  bool valuesCorrect = true;
  std::vector<bool> receiverCorrect(receivers.size(), true);
  bool guardsPreserved = true;
  bool mismatchPrinted = false;
  for (unsigned iteration = 0; iteration < iterations; ++iteration) {
    const int offset = static_cast<int>(guardWords + iteration);
    engine.writeTensor("sourceOffsetWrite", &offset, &offset + 1);
    engine.run(0);
    std::vector<int> output(results.numElements());
    engine.readTensor("resultsRead", output.data(), output.data() + output.size());
    for (unsigned receiver = 0; receiver < receivers.size(); ++receiver) {
      for (unsigned i = 0; i < guardWords; ++i) {
        guardsPreserved &= output[receiver * resultStride + i] == guard;
        guardsPreserved &=
            output[receiver * resultStride + guardWords + count + i] == guard;
      }
      for (unsigned i = 0; i < count; ++i) {
        const int actual = output[receiver * resultStride + guardWords + i];
        const int expected = sourceHost[static_cast<unsigned>(offset) + i];
        if (actual != expected && !mismatchPrinted) {
          std::cout << "firstMismatch iteration=" << iteration
                    << " receiver=" << receivers[receiver] << " word=" << i
                    << " actual=0x" << std::hex
                    << static_cast<unsigned>(actual) << " expected=0x"
                    << static_cast<unsigned>(expected) << std::dec << '\n';
          mismatchPrinted = true;
        }
        receiverCorrect[receiver] = receiverCorrect[receiver] && actual == expected;
      }
    }
  }
  std::vector<int> sourceAfter(sourceWords);
  engine.readTensor("sourceRead", sourceAfter.data(),
                    sourceAfter.data() + sourceAfter.size());
  const bool sourcePreserved = sourceAfter == sourceHost;
  for (bool correct : receiverCorrect) valuesCorrect &= correct;
  const bool ok = valuesCorrect && guardsPreserved && sourcePreserved;
  std::cout << "singlePacket=yes source=" << sender << " receivers=";
  for (unsigned i = 0; i < receivers.size(); ++i)
    std::cout << (i == 0 ? "" : ",") << receivers[i];
  std::cout << " count=" << count << " iterations=" << iterations
            << " valuesCorrect=" << (valuesCorrect ? "yes" : "NO")
            << " sourcePreserved=" << (sourcePreserved ? "yes" : "NO")
            << " guardsPreserved=" << (guardsPreserved ? "yes" : "NO") << '\n'
            << (ok ? "PASS\n" : "FAIL\n");
  if (!valuesCorrect) {
    std::cout << "failedReceivers=";
    bool first = true;
    for (unsigned i = 0; i < receivers.size(); ++i) {
      if (!receiverCorrect[i]) {
        std::cout << (first ? "" : ",") << receivers[i] << "(p"
                  << xcom::logicalToPhysical(receivers[i]) << ')';
        first = false;
      }
    }
    std::cout << '\n';
  }
  return ok ? 0 : 1;
}
