#pragma once

#include <array>
#include <cstdint>
#include <cstring>
#include <fstream>
#include <stdexcept>
#include <string>
#include <vector>

namespace ipuimg {

constexpr std::array<char, 8> magic = {'I', 'P', 'U', 'I', 'M', 'G', '1', 0};
constexpr std::uint32_t version = 1;
constexpr std::uint32_t headerSize = 40;
constexpr std::uint32_t recordSize = 16;

inline std::uint32_t crc32(const void *data, std::size_t size) {
  std::uint32_t crc = 0xffffffffu;
  const auto *bytes = static_cast<const std::uint8_t *>(data);
  for (std::size_t i = 0; i < size; ++i) {
    crc ^= bytes[i];
    for (unsigned bit = 0; bit < 8; ++bit)
      crc = (crc >> 1) ^ (0xedb88320u & (0u - (crc & 1u)));
  }
  return ~crc;
}

inline std::uint32_t readU32(std::istream &input) {
  std::array<unsigned char, 4> bytes{};
  input.read(reinterpret_cast<char *>(bytes.data()), bytes.size());
  if (!input) throw std::runtime_error("truncated IPU image");
  return static_cast<std::uint32_t>(bytes[0]) |
         (static_cast<std::uint32_t>(bytes[1]) << 8) |
         (static_cast<std::uint32_t>(bytes[2]) << 16) |
         (static_cast<std::uint32_t>(bytes[3]) << 24);
}

inline std::uint64_t readU64(std::istream &input) {
  const std::uint64_t low = readU32(input);
  return low | (static_cast<std::uint64_t>(readU32(input)) << 32);
}

inline void writeU32(std::ostream &output, std::uint32_t value) {
  const std::array<unsigned char, 4> bytes = {
      static_cast<unsigned char>(value),
      static_cast<unsigned char>(value >> 8),
      static_cast<unsigned char>(value >> 16),
      static_cast<unsigned char>(value >> 24)};
  output.write(reinterpret_cast<const char *>(bytes.data()), bytes.size());
}

inline void writeU64(std::ostream &output, std::uint64_t value) {
  writeU32(output, static_cast<std::uint32_t>(value));
  writeU32(output, static_cast<std::uint32_t>(value >> 32));
}

struct Header {
  std::uint32_t tileCount = 0;
  std::uint32_t baseAddress = 0;
  std::uint32_t imageSize = 0;
  std::uint32_t entryPoint = 0;
  std::uint32_t templateCrc = 0;
  std::uint32_t templateTile = 0;
};

struct TileRecord {
  std::uint64_t patchOffset = 0;
  std::uint32_t patchSize = 0;
  std::uint32_t imageCrc = 0;
};

class Image {
public:
  explicit Image(const std::string &path) : input_(path, std::ios::binary) {
    if (!input_) throw std::runtime_error("cannot open IPU image: " + path);
    std::array<char, 8> fileMagic{};
    input_.read(fileMagic.data(), fileMagic.size());
    if (fileMagic != magic) throw std::runtime_error("bad IPU image magic");
    if (readU32(input_) != version)
      throw std::runtime_error("unsupported IPU image version");
    header_.tileCount = readU32(input_);
    header_.baseAddress = readU32(input_);
    header_.imageSize = readU32(input_);
    header_.entryPoint = readU32(input_);
    header_.templateCrc = readU32(input_);
    header_.templateTile = readU32(input_);
    if (readU32(input_) != 0) throw std::runtime_error("invalid header flags");
    if (header_.tileCount == 0 || header_.imageSize == 0 ||
        header_.templateTile >= header_.tileCount)
      throw std::runtime_error("invalid IPU image dimensions");

    records_.resize(header_.tileCount);
    for (auto &record : records_) {
      record.patchOffset = readU64(input_);
      record.patchSize = readU32(input_);
      record.imageCrc = readU32(input_);
    }
    template_.resize(header_.imageSize);
    input_.read(reinterpret_cast<char *>(template_.data()), template_.size());
    if (!input_ || crc32(template_.data(), template_.size()) !=
                       header_.templateCrc)
      throw std::runtime_error("corrupt template image");
    const std::uint64_t patchBegin = headerSize +
        static_cast<std::uint64_t>(recordSize) * header_.tileCount +
        header_.imageSize;
    input_.seekg(0, std::ios::end);
    const auto fileEnd = input_.tellg();
    if (fileEnd < 0) throw std::runtime_error("cannot size IPU image");
    const std::uint64_t fileSize = static_cast<std::uint64_t>(fileEnd);
    for (const auto &record : records_)
      if (record.patchOffset < patchBegin || record.patchOffset > fileSize ||
          record.patchSize > fileSize - record.patchOffset)
        throw std::runtime_error("invalid tile patch extent");
  }

  const Header &header() const { return header_; }

  std::vector<std::uint8_t> tile(unsigned tile) {
    if (tile >= records_.size()) throw std::out_of_range("tile index");
    std::vector<std::uint8_t> result = template_;
    const TileRecord &record = records_[tile];
    input_.clear();
    input_.seekg(static_cast<std::streamoff>(record.patchOffset));
    const std::uint64_t end = record.patchOffset + record.patchSize;
    while (static_cast<std::uint64_t>(input_.tellg()) < end) {
      const std::uint32_t offset = readU32(input_);
      const std::uint32_t size = readU32(input_);
      if (size == 0 || offset > result.size() || size > result.size() - offset)
        throw std::runtime_error("invalid tile patch range");
      if (static_cast<std::uint64_t>(input_.tellg()) + size > end)
        throw std::runtime_error("truncated tile patch");
      input_.read(reinterpret_cast<char *>(result.data() + offset), size);
      if (!input_) throw std::runtime_error("truncated tile patch data");
    }
    if (static_cast<std::uint64_t>(input_.tellg()) != end)
      throw std::runtime_error("misaligned tile patch stream");
    if (crc32(result.data(), result.size()) != record.imageCrc)
      throw std::runtime_error("tile image checksum mismatch");
    return result;
  }

private:
  std::ifstream input_;
  Header header_;
  std::vector<TileRecord> records_;
  std::vector<std::uint8_t> template_;
};

} // namespace ipuimg
