"""Generates Rust literals containing SPIR-V bytecode from shader files.

This script identifies shader files (*.vert or *.frag) from a specified
directory, compiles them into SPIR-V binary format using the `glslc` compiler,
and then embeds the resulting SPIR-V bytecode into Rust `&[u32]` literals.

Usage:
  compile_spv_shaders.py <shaders_dir> <output_rust_file>

  <shaders_dir>: The path to the directory containing shader source files.

Requirements:
  - `glslc` can be obtained from the Shaderc project:
  https://github.com/google/shaderc
"""

import os
import struct
import subprocess
import sys
from typing import List


def compile_spv_code(filepath):
  """Compiles a shader to SPIR-V, returning the bytecode."""
  result = subprocess.run(
      [
          "glslc",
          "--target-env=vulkan1.0",
          filepath,
          "-o",
          "-",
      ],
      check=True,
      capture_output=True,
  )
  return result.stdout


def parse_spv_binary(spv_binary, source_filepath):
  """Parses SPIR-V binary into a list of 32-bit word strings."""
  if len(spv_binary) % 4 != 0:
    raise ValueError(
        f"SPIR-V binary size for {source_filepath} is not a multiple of 4"
    )
  words = struct.unpack(f"<{len(spv_binary)//4}I", spv_binary)
  return [f"0x{word:08x}" for word in words]


def generate_rust_literal(filepath, bytecode_words):
  """Generates a Rust literal from the given SPIR-V bytecode words."""
  column_width = 100
  indent = "    "
  if not bytecode_words:
    return ""
  # Calculate how many '0x, ' entries fit on one line.
  num_words_per_line = int(
      (column_width - len(indent)) / (len(bytecode_words[0]) + 2)
  )

  filename = os.path.basename(filepath)
  name, ext = os.path.splitext(filename)

  # Add original shader source as a doc comment.
  doc_comment = ""
  with open(filepath, "r") as shader_source:
    doc_comment += "\n"
    doc_comment += f"/// {filename} source:\n"
    doc_comment += "///\n"
    for line in shader_source:
      stripped_line = line.lstrip()
      if stripped_line.startswith("//"):
        # If the line (after stripping leading whitespace) starts with '//',
        # prepend '///  ' and then the content of the comment while preserving
        # original leading whitespace.
        leading_spaces_count = len(line) - len(stripped_line)
        doc_comment += (
            "///  "
            + (" " * leading_spaces_count)
            + "// "
            + stripped_line[2:]
        )
      else:
        # Otherwise, simply prepend '///' to the line.
        if line.strip():
          doc_comment += f"/// {line}"
        else:
          # If the line is empty, add a newline.
          doc_comment += "///\n"
    doc_comment += "\n"  # Add an extra newline after the comments block

  # For e.g., MM21SHADER_VERT_SPV.
  const_name = f"{name.upper()}{ext[1:].upper()}_SPV"

  literal = doc_comment
  literal += f"pub const {const_name}: &[u32] = &["
  for i, word in enumerate(bytecode_words):
    if i % num_words_per_line == 0:
      if i != 0:
        literal += ","
      literal += "\n" + indent + word
    else:
      literal += ", " + word
  literal += ",\n];\n"  # Add an extra newline after the const block
  return literal


def main(args: List[str]):
  if len(args) != 3:
    raise ValueError(
        "Usage: compile_spv_shaders.py <shaders_dir> <output_rust_file>"
    )

  shader_dir = args[1]
  output_rust_file = args[2]

  if not os.path.isdir(shader_dir):
    raise ValueError(f"Directory {shader_dir} does not exist")

  shader_files = []
  for name in os.listdir(shader_dir):
    filepath = os.path.join(shader_dir, name)
    if os.path.isfile(filepath):
      _, ext = os.path.splitext(name)
      if ext[1:] not in ["vert", "frag"]:
        continue
      shader_files.append(filepath)

  all_literals = ""
  for filepath in shader_files:
    spv_binary = compile_spv_code(filepath)
    bytecode_words = parse_spv_binary(spv_binary, filepath)
    all_literals += generate_rust_literal(filepath, bytecode_words)

    print(f"Successfully parsed {len(bytecode_words)} words for {filepath}.")

  with open(output_rust_file, "w") as f:
    f.write(
        "// This file is automatically generated by compile_spv_shaders.py. Do"
        " not edit.\n"
    )
    f.write(all_literals)
  print(f"Successfully wrote SPIR-V literals to {output_rust_file}")


if __name__ == "__main__":
  main(sys.argv)
