#!/usr/bin/env python3

import argparse
import os
from pathlib import Path
import subprocess
import sys


PRESETS = {
    "small": (2, 8, 16 * 1024),
    "medium": (32, 64, 128 * 1024),
    "severe": (128, 256, 512 * 1024),
}


def run_git(
    repository: Path, *arguments: str, expect_failure: bool = False
) -> subprocess.CompletedProcess:
    environment = os.environ.copy()
    environment.update(
        {
            "GIT_CONFIG_GLOBAL": os.devnull,
            "GIT_CONFIG_SYSTEM": os.devnull,
            "GIT_AUTHOR_NAME": "Zed Benchmark",
            "GIT_AUTHOR_EMAIL": "benchmark@zed.dev",
            "GIT_COMMITTER_NAME": "Zed Benchmark",
            "GIT_COMMITTER_EMAIL": "benchmark@zed.dev",
        }
    )
    result = subprocess.run(
        ["git", "-C", repository, *arguments],
        env=environment,
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        check=False,
    )
    if expect_failure:
        if result.returncode == 0:
            raise RuntimeError(f"git {' '.join(arguments)} unexpectedly succeeded")
    elif result.returncode != 0:
        stderr = result.stderr.decode(errors="replace").strip()
        raise RuntimeError(f"git {' '.join(arguments)} failed: {stderr}")
    return result


def unchanged_separator(region_index: int) -> str:
    return "".join(
        f"// unchanged separator {region_index:04d} {line_index:02d}\n"
        for line_index in range(10)
    )


def conflict_file(
    file_index: int, region_count: int, variant: str, padding: int
) -> str:
    chunks = [
        f"// Synthetic merge-conflict workload: file {file_index:04d}, variant {variant}\n"
    ]
    for region_index in range(region_count):
        value = f"{variant}-file-{file_index:04d}-region-{region_index:04d}-"
        chunks.append(
            f'pub const VALUE_{region_index:04d}: &str = "{value}{"x" * padding}";\n'
        )
        chunks.append(unchanged_separator(region_index))
    return "".join(chunks)


def padding_for_target_size(file_index: int, region_count: int, target_size: int) -> int:
    unpadded_ours = conflict_file(file_index, region_count, "ours", 0)
    unpadded_theirs = conflict_file(file_index, region_count, "theirs", 0)
    marker_overhead = region_count * len(
        "<<<<<<< HEAD\n=======\n>>>>>>> conflict-side\n"
    )
    shared_separator_size = sum(
        len(unchanged_separator(region_index))
        for region_index in range(region_count)
    )
    minimum_size = (
        len(unpadded_ours)
        + len(unpadded_theirs)
        - shared_separator_size
        + marker_overhead
    )
    if target_size < minimum_size:
        raise ValueError(
            f"target size {target_size} is too small for {region_count} regions; "
            f"minimum is approximately {minimum_size} bytes"
        )
    return (target_size - minimum_size) // (2 * region_count)


def write_variant(
    repository: Path,
    file_count: int,
    region_count: int,
    target_size: int,
    variant: str,
) -> None:
    for file_index in range(file_count):
        path = (
            repository
            / "conflicts"
            / f"module_{file_index:04d}"
            / f"generated_conflict_{file_index:04d}.rs"
        )
        path.parent.mkdir(parents=True, exist_ok=True)
        padding = padding_for_target_size(file_index, region_count, target_size)
        path.write_text(
            conflict_file(file_index, region_count, variant, padding),
            encoding="utf-8",
        )


def validate(repository: Path, file_count: int, region_count: int) -> tuple[int, int]:
    merge_head = run_git(repository, "rev-parse", "-q", "--verify", "MERGE_HEAD")
    if not merge_head.stdout.strip():
        raise RuntimeError("MERGE_HEAD is missing")

    unresolved_output = run_git(
        repository, "diff", "--name-only", "--diff-filter=U", "-z"
    ).stdout
    unresolved_paths = [
        Path(path.decode()) for path in unresolved_output.split(b"\0") if path
    ]
    if len(unresolved_paths) != file_count:
        raise RuntimeError(
            f"expected {file_count} unresolved paths, found {len(unresolved_paths)}"
        )

    total_bytes = 0
    for relative_path in unresolved_paths:
        contents = (repository / relative_path).read_bytes()
        total_bytes += len(contents)
        marker_counts = (
            contents.count(b"<<<<<<< "),
            contents.count(b"======="),
            contents.count(b">>>>>>> "),
        )
        if marker_counts != (region_count, region_count, region_count):
            raise RuntimeError(
                f"{relative_path} has marker counts {marker_counts}, "
                f"expected {(region_count, region_count, region_count)}"
            )

    status_entries = [
        entry
        for entry in run_git(
            repository, "status", "--porcelain=v1", "-z"
        ).stdout.split(b"\0")
        if entry
    ]
    if len(status_entries) != file_count or any(
        not entry.startswith(b"UU ") for entry in status_entries
    ):
        raise RuntimeError("the repository contains unexpected status entries")

    return len(unresolved_paths), total_bytes


def main() -> int:
    parser = argparse.ArgumentParser(
        description="Create a deterministic unresolved Git merge for Zed performance testing."
    )
    parser.add_argument("destination", type=Path)
    parser.add_argument(
        "--preset", choices=PRESETS, default="medium", help="workload size"
    )
    arguments = parser.parse_args()

    file_count, region_count, target_size = PRESETS[arguments.preset]
    destination = arguments.destination.expanduser().resolve()
    try:
        destination.mkdir(parents=True, exist_ok=False)
        run_git(destination, "init", "-b", "main")

        write_variant(destination, file_count, region_count, target_size, "base")
        run_git(destination, "add", "conflicts")
        run_git(destination, "commit", "-m", "base")

        run_git(destination, "switch", "-c", "conflict-side")
        write_variant(destination, file_count, region_count, target_size, "theirs")
        run_git(destination, "add", "conflicts")
        run_git(destination, "commit", "-m", "theirs")

        run_git(destination, "switch", "main")
        write_variant(destination, file_count, region_count, target_size, "ours")
        run_git(destination, "add", "conflicts")
        run_git(destination, "commit", "-m", "ours")

        run_git(destination, "merge", "conflict-side", expect_failure=True)
        unresolved_count, total_bytes = validate(
            destination, file_count, region_count
        )
    except (OSError, RuntimeError, ValueError) as error:
        print(f"error: {error}", file=sys.stderr)
        return 1

    print(f"repository: {destination}")
    print(f"preset: {arguments.preset}")
    print(f"unresolved files: {unresolved_count}")
    print(f"conflict regions per file: {region_count}")
    print(f"total conflict regions: {unresolved_count * region_count}")
    print(f"working tree bytes: {total_bytes}")
    return 0


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