#pragma once

#include "ipu_image.hpp"

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

namespace ipuimg {

struct WriteStats {
  std::uint64_t changedBytes = 0;
  std::uint64_t patchRuns = 0;
  std::uint64_t outputBytes = 0;
};

inline WriteStats write(const std::string &path,
                        const std::vector<std::vector<std::uint8_t>> &images,
                        std::uint32_t baseAddress, std::uint32_t entryPoint,
                        unsigned templateTile) {
  if (images.empty() || templateTile >= images.size() || images[0].empty())
    throw std::runtime_error("invalid tile images");
  const auto &base = images[templateTile];
  for (const auto &image : images)
    if (image.size() != base.size())
      throw std::runtime_error("tile image sizes do not match");

  Header header;
  header.tileCount = images.size();
  header.baseAddress = baseAddress;
  header.imageSize = base.size();
  header.entryPoint = entryPoint;
  header.templateCrc = crc32(base.data(), base.size());
  header.templateTile = templateTile;
  std::vector<TileRecord> records(images.size());

  std::ofstream output(path, std::ios::binary | std::ios::trunc);
  if (!output) throw std::runtime_error("cannot create " + path);
  output.write(magic.data(), magic.size());
  writeU32(output, version);
  writeU32(output, header.tileCount);
  writeU32(output, header.baseAddress);
  writeU32(output, header.imageSize);
  writeU32(output, header.entryPoint);
  writeU32(output, header.templateCrc);
  writeU32(output, header.templateTile);
  writeU32(output, 0);
  for (unsigned tile = 0; tile < images.size(); ++tile) {
    writeU64(output, 0);
    writeU32(output, 0);
    writeU32(output, 0);
  }
  output.write(reinterpret_cast<const char *>(base.data()), base.size());

  WriteStats stats;
  for (unsigned tile = 0; tile < images.size(); ++tile) {
    auto &record = records[tile];
    record.patchOffset = static_cast<std::uint64_t>(output.tellp());
    record.imageCrc = crc32(images[tile].data(), images[tile].size());
    std::size_t at = 0;
    while (at < base.size()) {
      while (at < base.size() && images[tile][at] == base[at]) ++at;
      if (at == base.size()) break;
      const std::size_t begin = at++;
      std::size_t lastDifference = at;
      while (at < base.size()) {
        if (images[tile][at] != base[at]) {
          lastDifference = at + 1;
        } else if (at - lastDifference >= 8) {
          break;
        }
        ++at;
      }
      const std::size_t size = lastDifference - begin;
      writeU32(output, begin);
      writeU32(output, size);
      output.write(reinterpret_cast<const char *>(images[tile].data() + begin),
                   size);
      stats.changedBytes += size;
      ++stats.patchRuns;
      at = lastDifference;
    }
    record.patchSize = static_cast<std::uint64_t>(output.tellp()) -
                       record.patchOffset;
  }
  if (!output) throw std::runtime_error("failed while writing IPU image");
  output.seekp(headerSize);
  for (const auto &record : records) {
    writeU64(output, record.patchOffset);
    writeU32(output, record.patchSize);
    writeU32(output, record.imageCrc);
  }
  output.seekp(0, std::ios::end);
  stats.outputBytes = static_cast<std::uint64_t>(output.tellp());
  output.close();
  if (!output) throw std::runtime_error("failed to finalize IPU image");
  return stats;
}

} // namespace ipuimg
