#!/usr/bin/env python3
"""Run current kernel timing/correctness fixtures, serializing device access."""
import argparse
from concurrent.futures import ThreadPoolExecutor
import json
from pathlib import Path
import subprocess


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--sdk", required=True)
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--cases", nargs="+", help="Run only these named fixtures")
    args = parser.parse_args()
    args.output.mkdir(parents=True, exist_ok=False)
    cases = [
        ("elementwise", "elementwise_check", []),
        ("residual", "elementwise_check", ["--residual"]),
        ("elementwise-fp8", "elementwise_check", ["--fp8"]),
        ("bias-gelu-fp8", "elementwise_check", ["--fp8", "--bias-gelu"]),
        ("cast", "cast_check", []),
        ("cast-inplace", "cast_check", ["--in-place-only"]),
        ("unpack", "unpack_check", []),
        ("pack", "pack_check", ["--specialized"]),
        ("pack-fp16", "pack_check", ["--specialized", "--fp16"]),
        ("softmax", "softmax_check", []),
        ("softmax-interleaved", "softmax_check", ["--interleaved-output"]),
        ("softmax-split", "softmax_check", ["--split-rows"]),
        ("softmax-fp8", "softmax_check", ["--fp8-output"]),
        ("gelu-reduce", "ipu-kernel-equivalence", ["--reference", "device", "--exact"]),
        ("bias-gelu", "ipu-kernel-equivalence", ["--reference", "device", "--bias-gelu", "--bias-rows", "3", "--exact"]),
    ]
    for word in [2, 4, 8]:
        for contiguous in [True, False]:
            if word == 2 and not contiguous:
                continue
            name = f"copy-{word}-{'dense' if contiguous else 'strided'}"
            tasks = []
            for rows in ([1] if contiguous else [1, 2, 3, 4, 5, 6, 7, 12, 32, 73]):
                for words in [1, 2, 6, 7, 32, 64, 192, 448]:
                    stride = (words + 2) * word
                    size = (rows * stride + 3) // 4 * 4
                    if size > 65536:
                        continue
                    tasks.append(dict(bytes=size, tasks=[dict(
                        source=0, destination=0, row_bytes=words * word,
                        rows=rows, source_stride=stride, destination_stride=stride)]))
            path = args.output / f"{name}.json"
            path.write_text(json.dumps(tasks, indent=2) + "\n")
            cases.append((name, "copy_check", [str(path), "--word-bytes", str(word)]
                          + (["--contiguous"] if contiguous else [])))
            if word == 8:
                cases.append(("fill" if contiguous else "fill-strided", "copy_check",
                              [str(path), "--word-bytes", "8", "--fill"]
                              + (["--contiguous"] if contiguous else [])))
                if not contiguous:
                    cases.append(("fill-shared", "copy_check",
                                  [str(path), "--word-bytes", "8", "--fill", "--fill-shared"]))
    if args.cases:
        unknown = set(args.cases) - {name for name, _, _ in cases}
        if unknown:
            parser.error(f"Unknown fixtures: {sorted(unknown)}")
        cases = [case for case in cases if case[0] in args.cases]

    def run(case):
        name, binary, options = case
        directory = args.output / name
        directory.mkdir()
        command = [f"target/release/{binary}", "--sdk", args.sdk,
                   "--device-lock", "/tmp/ipu-stack-device.lock",
                   "--output", str(directory), *options]
        (directory / "command.json").write_text(json.dumps(command, indent=2) + "\n")
        with (directory / "run.log").open("w") as log:
            result = subprocess.run(command, stdout=log, stderr=subprocess.STDOUT)
        print(f"{name}: {result.returncode}", flush=True)
        return {"case": name, "returncode": result.returncode}

    with ThreadPoolExecutor(max_workers=2) as pool:
        results = list(pool.map(run, cases))
    (args.output / "results.json").write_text(json.dumps(results, indent=2) + "\n")
    raise SystemExit(any(result["returncode"] for result in results))


if __name__ == "__main__":
    main()
