#!/usr/bin/env python3
"""Measure planner budget sensitivity; retain failures, placement and device profiles."""
import argparse
import concurrent.futures
import hashlib
import json
import os
from pathlib import Path
import re
import signal
import subprocess
import time


def run_case(args, batch, budget, blocks, precision):
    name = f"{precision}-b{batch}-n{blocks}-" + (f"{budget}k" if budget else "default")
    path = args.output / name
    path.mkdir(parents=True, exist_ok=True)
    command = [str(args.binary), "c600-init.ipucfg", "--sdk", str(args.sdk),
               "--workload", "siglip-mlp-benchmark", "--mlp-batch", str(batch),
               "--mlp-blocks", str(blocks), "--device-lock", "/tmp/ipu-stack-device.lock",
               "--package", str(path / "model.ipuexe"),
               "--profile-output", str(path / "profile.capnp"),
               "--memory-profile-directory", str(path / "memory")]
    if precision == "fp8":
        command.append("--fp8-scale=-2")
    if budget:
        command += ["--tile-memory-budget-kib", str(budget)]
    (path / "command.json").write_text(json.dumps(command, indent=2) + "\n")
    start = time.monotonic()
    timed_out = False
    with (path / "run.log").open("w") as log:
        process = subprocess.Popen(command, stdout=log, stderr=subprocess.STDOUT,
                                   env=dict(os.environ, RAYON_NUM_THREADS=str(args.threads)),
                                   start_new_session=True)
        try:
            process.wait(timeout=args.timeout)
        except subprocess.TimeoutExpired:
            timed_out = True
            os.killpg(process.pid, signal.SIGKILL)
            process.wait()
    log = (path / "run.log").read_text()
    status = "timeout" if timed_out else "failed"
    if process.returncode == 0 and "hardwareTest=PASS" in log:
        status = "pass"
    elif "construction is not implemented" in log:
        status = "unsupported"
    elif "no fitting path through high operation boundary" in log:
        status = "planner-rejected"
    result = dict(case=name, batch=batch, blocks=blocks, precision=precision,
                  budget_kib=budget, status=status, returncode=process.returncode,
                  wall_seconds=round(time.monotonic() - start, 3))
    result["stages"] = dict(re.findall(r':([a-z_]+): close time.busy=([^ ]+)', log))
    for field in ["cycles", "maximumAbsoluteError", "weightBytes", "inputBytes"]:
        match = re.search(r'\b' + field + r'=([\d.]+)', log)
        if match:
            result[field] = float(match[1]) if "." in match[1] else int(match[1])
    result["diagnostic"] = log.splitlines()[-15:] if status != "pass" else []
    if (path / "profile.capnp").exists():
        query = subprocess.run([str(args.cli), "profile-query", str(path / "profile.capnp"), "--json"],
                               capture_output=True, text=True, check=True)
        (path / "summary.json").write_text(query.stdout)
        result["renderer_cycles"] = json.loads(query.stdout)["profileSpanCycles"]
    (path / "result.json").write_text(json.dumps(result, indent=2) + "\n")
    return result


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--sdk", type=Path, required=True)
    parser.add_argument("--binary", type=Path, default=Path("target/release/ipu-e2e-test"))
    parser.add_argument("--cli", type=Path, default=Path("target/release/ipu-stack"))
    parser.add_argument("--jobs", type=int, default=4)
    parser.add_argument("--threads", type=int, default=12)
    parser.add_argument("--timeout", type=int, default=600)
    parser.add_argument("--batches", type=int, nargs="+", default=[1, 2, 4])
    parser.add_argument("--budgets", type=int, nargs="+", default=[0, 384, 256, 128, 64])
    parser.add_argument("--precision", choices=["fp8", "fp16"], default="fp8")
    parser.add_argument("--blocks", type=int, default=1)
    args = parser.parse_args()
    args.output.mkdir(parents=True, exist_ok=False)
    (args.output / "build.json").write_text(json.dumps({
        "commit": subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(),
        "binary_sha256": hashlib.sha256(args.binary.read_bytes()).hexdigest(),
        "jobs": args.jobs, "threads_per_job": args.threads,
    }, indent=2) + "\n")
    results = []
    with concurrent.futures.ThreadPoolExecutor(max_workers=args.jobs) as pool:
        futures = [pool.submit(run_case, args, batch, budget, args.blocks, args.precision)
                   for budget in args.budgets for batch in args.batches]
        for future in concurrent.futures.as_completed(futures):
            result = future.result()
            results.append(result)
            print(json.dumps(result), flush=True)
            (args.output / "results.json").write_text(json.dumps(results, indent=2) + "\n")


if __name__ == "__main__":
    main()
