#pragma once

#include <cstdint>
#include <stdexcept>

namespace hostxcode {

struct PacketHeader {
  std::uint32_t word0;
  std::uint32_t word1;

  bool operator==(const PacketHeader &other) const {
    return word0 == other.word0 && word1 == other.word1;
  }
};

enum class PacketSize { Short, Long };

inline PacketSize packetSize(std::uint32_t hostOffset, std::uint32_t bytes) {
  if (bytes != 0 && bytes <= 60 && (hostOffset & 3u) == 0 &&
      (bytes & 3u) == 0)
    return PacketSize::Short;
  if (bytes != 0 && bytes <= 1024 && (hostOffset & 63u) == 0 &&
      (bytes & 63u) == 0)
    return PacketSize::Long;
  throw std::invalid_argument(
      "host packet must be 4..60 bytes in 4-byte units or 64..1024 bytes "
      "in 64-byte units, with matching host alignment");
}

inline std::uint32_t routeWord0(unsigned physicalTile) {
  if (physicalTile > 0xfff)
    throw std::out_of_range("physical tile exceeds packet route field");
  return ((physicalTile >> 1) << 16) | ((physicalTile & 1u) << 15);
}

inline std::uint32_t routeWord1(unsigned physicalTile) {
  if (physicalTile > 0xfff)
    throw std::out_of_range("physical tile exceeds packet route field");
  return (physicalTile & 1u) << 31;
}

inline std::uint32_t hostAddressLength(std::uint32_t hostOffset,
                                       std::uint32_t bytes,
                                       PacketSize size) {
  const unsigned shift = size == PacketSize::Short ? 2 : 6;
  const std::uint32_t maximum = size == PacketSize::Short ? 60 : 1024;
  if (bytes == 0 || bytes > maximum ||
      (hostOffset & ((1u << shift) - 1)) != 0 ||
      (bytes & ((1u << shift) - 1)) != 0)
    throw std::invalid_argument("host packet address or length is unencodable");
  const std::uint32_t units = bytes >> shift;
  const std::uint32_t length =
      size == PacketSize::Long && bytes == 1024 ? 0 : units;
  const std::uint64_t encoded =
      (static_cast<std::uint64_t>(hostOffset >> shift) << 4) | length;
  if (encoded > 0x7fffffffu)
    throw std::invalid_argument("host packet address exceeds packet field");
  return static_cast<std::uint32_t>(encoded);
}

// Command sent before tile payload travelling from the IPU to the host.
inline PacketHeader tileToHost(unsigned physicalTile,
                               std::uint32_t hostOffset,
                               std::uint32_t bytes) {
  const auto size = packetSize(hostOffset, bytes);
  const std::uint32_t opcode =
      size == PacketSize::Short ? 0x80000000u : 0xa0000000u;
  return {opcode | routeWord0(physicalTile),
          routeWord1(physicalTile) |
              hostAddressLength(hostOffset, bytes, size)};
}

// Read request asking the host endpoint to write into tile SRAM. The local
// destination field is a 32-byte unit offset from 0x50000.
inline PacketHeader hostToTile(unsigned physicalTile,
                               std::uint32_t tileAddress,
                               std::uint32_t hostOffset,
                               std::uint32_t bytes) {
  constexpr std::uint32_t exchangeAddressBase = 0x50000;
  if (tileAddress < exchangeAddressBase || (tileAddress & 31u) != 0)
    throw std::invalid_argument(
        "host-to-tile destination must be 32-byte aligned at or above 0x50000");
  const std::uint32_t xAddress = (tileAddress - exchangeAddressBase) >> 5;
  if (xAddress > 0x1ff)
    throw std::invalid_argument("host-to-tile destination exceeds X address");
  const auto size = packetSize(hostOffset, bytes);
  const std::uint32_t opcode =
      size == PacketSize::Short ? 0xcc000200u : 0xec000200u;
  return {opcode | routeWord0(physicalTile) | xAddress,
          routeWord1(physicalTile) |
              hostAddressLength(hostOffset, bytes, size)};
}

// A zero-length read closes a tile-to-host packet sequence. It targets a
// caller-provided 32-byte dummy receive region and carries no host payload.
inline PacketHeader zeroByteRead(unsigned physicalTile,
                                 std::uint32_t dummyTileAddress) {
  constexpr std::uint32_t exchangeAddressBase = 0x50000;
  if (dummyTileAddress < exchangeAddressBase ||
      (dummyTileAddress & 31u) != 0)
    throw std::invalid_argument("zero-byte-read dummy address alignment");
  const std::uint32_t xAddress =
      (dummyTileAddress - exchangeAddressBase) >> 5;
  if (xAddress > 0x1ff)
    throw std::invalid_argument("zero-byte-read dummy address exceeds X address");
  return {0xcc000200u | routeWord0(physicalTile) | xAddress,
          routeWord1(physicalTile)};
}

} // namespace hostxcode
