#include <algorithm>
#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>

using namespace poplar;

namespace {
unsigned argOr(char **argv, int argc, int index, unsigned fallback) {
  return index < argc ? std::strtoul(argv[index], nullptr, 0) : fallback;
}
} // namespace

int main(int argc, char **argv) {
  const unsigned sender = argOr(argv, argc, 1, 0);
  const unsigned receiverA = argOr(argv, argc, 2, 274);
  const unsigned receiverB = argOr(argv, argc, 3, 1286);
  const unsigned count = argOr(argv, argc, 4, 3);
  const bool inspectOnly = argc > 5;
  const unsigned sourcePrefix = argOr(argv, argc, 6, 0);
  const unsigned receiverAPrefix = argOr(argv, argc, 7, 0);
  const unsigned receiverBPrefix = argOr(argv, argc, 8, 0);

  DeviceManager manager;
  auto devices = manager.getDevices(TargetType::IPU, 1);
  if (devices.empty() || !devices[0].attach()) {
    std::cerr << "failed to attach an IPU\n";
    return 2;
  }
  Device &device = devices[0];
  Graph graph(device.getTarget());
  std::vector<int> sourceHost(count);
  for (unsigned i = 0; i < count; ++i) {
    sourceHost[i] = static_cast<int>(0x12340000u + i * 17u);
  }
  Tensor sourceStorage =
      graph.addVariable(INT, {sourcePrefix + count}, "sourceStorage");
  Tensor outputAStorage =
      graph.addVariable(INT, {receiverAPrefix + count}, "outputAStorage");
  Tensor outputBStorage =
      graph.addVariable(INT, {receiverBPrefix + count}, "outputBStorage");
  Tensor source = sourceStorage.slice(sourcePrefix, sourcePrefix + count);
  Tensor outputA =
      outputAStorage.slice(receiverAPrefix, receiverAPrefix + count);
  Tensor outputB =
      outputBStorage.slice(receiverBPrefix, receiverBPrefix + count);
  Tensor outputs = concat({outputA, outputB}).reshape({2, count});
  std::vector<int> sourceStorageHost(sourcePrefix + count, 0x55667788);
  std::copy(sourceHost.begin(), sourceHost.end(),
            sourceStorageHost.begin() + sourcePrefix);
  graph.setInitialValue<int>(sourceStorage, sourceStorageHost);
  graph.setTileMapping(sourceStorage, sender);
  graph.setTileMapping(outputAStorage, receiverA);
  graph.setTileMapping(outputBStorage, receiverB);
  Tensor broadcast = source.expand({0}).broadcast(2, 0);
  if (!inspectOnly) {
    graph.createHostRead("outputsRead", outputs.flatten());
  }

  OptionFlags options;
  options.set("exchange.multicastPolicy", "balanced");
  options.set("debug.dumpGlobalExchangePackets", "true");
  options.set("debug.dumpDirectory", "sdk_dumps");
  options.set("debug.retainDebugInformation", "true");
  Engine engine(graph, program::Copy(broadcast, outputs), options);
  std::ofstream executable("sdk_multicast.popef", std::ios::binary);
  engine.serializeExecutable(executable);
  executable.close();
  engine.load(device);
  engine.run(0);

  if (inspectOnly) {
    std::cout << "serialized inspection executable\n";
    return 0;
  }

  std::vector<int> outputHost(outputs.numElements());
  engine.readTensor("outputsRead", outputHost.data(),
                    outputHost.data() + outputHost.size());
  bool ok = true;
  for (unsigned receiver = 0; receiver < 2; ++receiver) {
    for (unsigned i = 0; i < count; ++i) {
      ok &= outputHost[receiver * count + i] == sourceHost[i];
    }
  }
  std::cout << "sender=" << sender << " receivers=" << receiverA << ','
            << receiverB << " count=" << count << ' '
            << (ok ? "PASS\n" : "FAIL\n");
  return ok ? 0 : 1;
}
