#include <poplar/Vertex.hpp>

namespace {
// The public target headers omit the exchange CSRs. These indices are exposed
// by libipu_arch_info and are stable across IPU21 revisions.
constexpr unsigned IncomingMuxPair = 161;
constexpr unsigned IncomingMux = 160;
constexpr unsigned IncomingDelta = 162;
constexpr unsigned IncomingFormat = 163;
constexpr unsigned IncomingBase = 164;
constexpr unsigned IncomingSinit = 165;
constexpr unsigned IncomingDcount = 166;
constexpr unsigned OutgoingBase = 167;
constexpr unsigned OutgoingDelta = 168;
constexpr unsigned ExchangeCtl = 169;
constexpr unsigned ExchangeAdj = 171;
} // namespace

class [[poplar::constraint("elem(*first) != elem(*second)")]]
KeepPair64Separated : public poplar::Vertex {
public:
  poplar::InOut<poplar::Vector<unsigned long long>> first;
  poplar::InOut<poplar::Vector<unsigned long long>> second;
  bool compute() { return true; }
};

class KeepPadding64 : public poplar::Vertex {
public:
  poplar::InOut<poplar::Vector<unsigned long long>> values;
  bool compute() { return true; }
};

class ReadExchangeConfig : public poplar::SupervisorVertex {
public:
  poplar::Output<poplar::Vector<unsigned>> values;

  __attribute__((target("supervisor"))) bool compute() {
#if defined(__IPU__) && defined(__POPC__)
    values[0] = __builtin_ipu_get(IncomingMux);
    values[1] = __builtin_ipu_get(IncomingMuxPair);
    values[2] = __builtin_ipu_get(IncomingDelta);
    values[3] = __builtin_ipu_get(IncomingFormat);
    values[4] = __builtin_ipu_get(IncomingBase);
    values[5] = __builtin_ipu_get(IncomingSinit);
    values[6] = __builtin_ipu_get(IncomingDcount);
    values[7] = __builtin_ipu_get(OutgoingBase);
    values[8] = __builtin_ipu_get(OutgoingDelta);
    values[9] = __builtin_ipu_get(ExchangeCtl);
    values[10] = __builtin_ipu_get(ExchangeAdj);
#else
    for (unsigned index = 0; index != values.size(); ++index)
      values[index] = 0;
#endif
    return true;
  }
};
