#!/usr/bin/env python3
# Copyright Advanced Micro Devices, Inc., or its affiliates.
# SPDX-License-Identifier: MIT

"""Run a command (typically hipblaslt-bench) under CU contention.

Launches a persistent HIP "busy" kernel pinned to a fixed number of CUs
(one workgroup per CU), waits until all its workgroups are resident, then runs
the given command while that contention is in effect. The cotenant is always
killed when the command finishes or this script is interrupted.

    hipblaslt-cotenant --cus 64 -- hipblaslt-bench -m 4096 -n 4096 -k 4096

A precompiled hipblaslt-cotenant-kernel beside this script is used if present;
otherwise it is built on first use with hipcc (override the compiler with
HIPCC=...) into a per-user cache. The cotenant log defaults to ./cotenant.<pid>.log.
"""

from __future__ import annotations

import argparse
import os
import signal
import subprocess
import sys
import tempfile
import time
from pathlib import Path
from typing import TextIO

HERE = Path(__file__).resolve().parent
SRC = HERE / "hipblaslt-cotenant-kernel.hip"
# A precompiled binary installed next to this script (the shipped case), if present.
SHIPPED_BIN = HERE / "hipblaslt-cotenant-kernel"


def cache_dir() -> Path:
    """Writable per-user dir for a self-built binary; the script's own dir may be read-only when installed."""
    base = os.environ.get("XDG_CACHE_HOME") or os.path.join(
        os.path.expanduser("~"), ".cache"
    )
    return Path(base) / "hipblaslt" / "cotenant"


def probe_device(env: dict[str, str]) -> tuple[str | None, int | None]:
    """Return (gfx_arch, total_cus) for the first GPU via rocminfo, or (None, None)."""
    try:
        out = subprocess.check_output(
            ["rocminfo"], text=True, env=env, stderr=subprocess.DEVNULL
        )
    except (FileNotFoundError, subprocess.CalledProcessError):
        return None, None
    arch, cus, in_gpu = None, None, False
    for line in out.splitlines():
        fields = line.split()
        if len(fields) >= 2 and fields[0] == "Name:" and fields[1].startswith("gfx"):
            arch, in_gpu = fields[1], True
        elif in_gpu and line.strip().startswith("Compute Unit:"):
            cus = int(fields[-1])
            break
    return arch, cus


def resolve_binary(arch: str | None, env: dict[str, str], explicit: str | None) -> Path:
    """Locate the cotenant binary: --binary, else a precompiled one beside the script, else build it."""
    if explicit:
        path = Path(explicit)
        if not (path.is_file() and os.access(path, os.X_OK)):
            sys.exit(f"ERROR: --binary {path} is not an executable file.")
        return path
    if os.access(SHIPPED_BIN, os.X_OK):
        return SHIPPED_BIN
    return build_cotenant(arch, env)


def build_cotenant(arch: str | None, env: dict[str, str]) -> Path:
    """Build the kernel into the per-user cache and return its path (rebuilds when stale)."""
    if not SRC.is_file():
        sys.exit(
            f"ERROR: no prebuilt kernel and source not found ({SRC}); pass --binary PATH."
        )
    out = cache_dir() / "hipblaslt-cotenant-kernel"
    stamp = cache_dir() / "arch"
    try:
        fresh = (
            out.exists()
            and out.stat().st_mtime >= SRC.stat().st_mtime
            and stamp.exists()
            and stamp.read_text().strip() == arch
        )
    except OSError:
        fresh = False
    if fresh:
        return out
    hipcc = os.environ.get("HIPCC", "hipcc")
    print(f"building hipblaslt-cotenant-kernel for {arch} (cache: {out.parent}) ...")
    try:
        out.parent.mkdir(parents=True, exist_ok=True)
    except OSError as e:
        sys.exit(f"ERROR: could not create cache dir {out.parent} ({e.strerror}).")
    # Build to a temp file and rename into place, so concurrent builders (a
    # parallel sweep with a cold cache) can't corrupt or ETXTBSY each other.
    fd, tmp_name = tempfile.mkstemp(dir=str(out.parent), prefix=".build-")
    os.close(fd)
    tmp = Path(tmp_name)
    try:
        subprocess.run(
            [
                hipcc,
                "-O2",
                "-std=c++17",
                f"--offload-arch={arch}",
                str(SRC),
                "-o",
                str(tmp),
            ],
            check=True,
            env=env,
        )
        os.chmod(tmp, 0o755)
        os.replace(tmp, out)
    except FileNotFoundError:
        tmp.unlink(missing_ok=True)
        sys.exit(f"ERROR: '{hipcc}' not found; install ROCm or set HIPCC to its path.")
    except subprocess.CalledProcessError as e:
        tmp.unlink(missing_ok=True)
        sys.exit(
            f"ERROR: building hipblaslt-cotenant-kernel failed (hipcc exited {e.returncode})."
        )
    except OSError as e:
        tmp.unlink(missing_ok=True)
        sys.exit(f"ERROR: could not write {out} ({e.strerror}).")
    try:
        stamp.write_text(arch)
    except OSError as e:
        print(
            f"WARNING: could not write build cache stamp {stamp} ({e.strerror}); will rebuild next run."
        )
    return out


def resolve_log_path(explicit: str | None) -> Path:
    """Pick a writable path for the cotenant log: --log, else the cwd, else a temp dir.

    The default name includes the PID so concurrent runs in the same directory
    don't stomp on each other's log.
    """
    if explicit:
        return Path(explicit)
    name = f"cotenant.{os.getpid()}.log"
    if os.access(Path.cwd(), os.W_OK):
        return Path.cwd() / name
    return Path(tempfile.gettempdir()) / name


def run_command(command: list[str], env: dict[str, str]) -> int:
    """Run the user's command, returning its exit code; fail cleanly if not found."""
    try:
        return subprocess.run(command, env=env).returncode
    except FileNotFoundError:
        sys.exit(f"ERROR: command not found: {command[0]}")


# Printed by the kernel once every workgroup is confirmed resident on the CUs.
READY_MARKER = "READY"


def wait_ready(
    proc: subprocess.Popen, log_handle: TextIO, log_path: Path, wait_s: float
) -> None:
    """Block until the cotenant logs READY (all workgroups resident); abort if it exits or times out."""
    deadline = time.monotonic() + wait_s
    while True:
        # Re-read the whole log each poll: the cotenant writes concurrently, so a
        # line-at-a-time read could split the marker across two reads and miss it.
        log_handle.seek(0)
        if READY_MARKER in log_handle.read():
            return
        if proc.poll() is not None:
            sys.exit(
                f"ERROR: cotenant exited (rc={proc.returncode}) before becoming resident; see {log_path}."
            )
        if time.monotonic() >= deadline:
            sys.exit(
                f"ERROR: cotenant not resident after {wait_s}s; aborting. See {log_path} "
                "(raise --wait if the device is slow to schedule the grid)."
            )
        time.sleep(0.1)


def main() -> int:
    # Route SIGTERM through the same cleanup path Ctrl-C already uses, so an
    # external `kill` of this process still tears the cotenant down.
    def _sigterm(_signum: int, _frame: object) -> None:
        raise KeyboardInterrupt

    signal.signal(signal.SIGTERM, _sigterm)

    ap = argparse.ArgumentParser(
        description="Run a command under CU contention from a busy cotenant kernel.",
        epilog="All options must precede `--`; everything after it is the command, "
        "e.g. hipblaslt-cotenant --cus 64 -- hipblaslt-bench -m 1024 -n 1024 -k 1024",
    )
    ap.add_argument(
        "--cus",
        type=int,
        required=True,
        metavar="N",
        help="number of CUs the cotenant occupies (0 = uncontended baseline, no cotenant)",
    )
    ap.add_argument(
        "--device", metavar="N", help="set HIP_VISIBLE_DEVICES for cotenant and command"
    )
    ap.add_argument(
        "--arch",
        metavar="GFX",
        help="build target arch (default: rocminfo auto-detect)",
    )
    ap.add_argument(
        "--binary",
        metavar="PATH",
        help="prebuilt hipblaslt-cotenant-kernel to use (default: one beside this script, else build into the cache)",
    )
    ap.add_argument(
        "--log",
        metavar="PATH",
        help="cotenant log path (default: ./cotenant.<pid>.log, or a temp dir if cwd is read-only)",
    )
    ap.add_argument(
        "--wait",
        type=float,
        default=30.0,
        metavar="S",
        help="max seconds to wait for residency (default: 30)",
    )
    ap.add_argument(
        "--grace",
        type=float,
        default=0.0,
        metavar="S",
        help="extra seconds to wait after residency is confirmed (default: 0)",
    )
    ap.add_argument(
        "command", nargs=argparse.REMAINDER, help="command to run after `--`"
    )
    args = ap.parse_args()

    if args.cus < 0:
        sys.exit("ERROR: --cus must be >= 0.")
    if not args.wait >= 0:
        sys.exit("ERROR: --wait must be >= 0.")
    if not args.grace >= 0:
        sys.exit("ERROR: --grace must be >= 0.")
    command = (
        args.command[1:] if args.command and args.command[0] == "--" else args.command
    )
    if not command:
        sys.exit("ERROR: no command given; pass it after `--`.")

    env = os.environ.copy()
    if args.device is not None:
        env["HIP_VISIBLE_DEVICES"] = args.device

    # --cus 0 is the uncontended baseline: run the command with no cotenant.
    if args.cus == 0:
        print(f"running uncontended (no cotenant): {' '.join(command)}")
        return run_command(command, env)

    # rocminfo selects via ROCR_VISIBLE_DEVICES; mirror the device the kernel
    # will actually use (HIP_VISIBLE_DEVICES, from --device or the environment)
    # so the arch/CU probe targets that GPU on heterogeneous multi-GPU hosts.
    probe_env = env.copy()
    hip_visible = env.get("HIP_VISIBLE_DEVICES")
    if hip_visible is not None:
        probe_env["ROCR_VISIBLE_DEVICES"] = hip_visible
    arch, total_cus = probe_device(probe_env)
    arch = args.arch or arch
    if arch is None:
        sys.exit("ERROR: could not detect GPU arch via rocminfo; pass --arch gfxNNN.")

    if total_cus is None:
        print(
            "WARNING: could not read CU count from rocminfo; skipping --cus bounds check."
        )
    elif args.cus > total_cus:
        sys.exit(f"ERROR: --cus {args.cus} exceeds the device CU count ({total_cus}).")
    elif args.cus == total_cus:
        sys.exit(
            f"ERROR: --cus {args.cus} would occupy all {total_cus} CUs and starve the benchmark; use fewer."
        )

    binary = resolve_binary(arch, env, args.binary)

    log_path = resolve_log_path(args.log)
    print(f"launching cotenant on {args.cus} CUs (log: {log_path})")
    try:
        log = open(log_path, "w")
    except OSError as e:
        sys.exit(
            f"ERROR: cannot write log to {log_path} ({e.strerror}); pass --log PATH."
        )
    with log, open(log_path, "r") as log_reader:
        # No start_new_session: keep the cotenant in this process's group so
        # terminal signals (SIGINT/SIGHUP) reach it directly even if our cleanup
        # never runs. It is not a group leader, so cleanup uses os.kill, not killpg.
        cotenant = subprocess.Popen(
            [str(binary), str(args.cus)],
            stdout=log,
            stderr=subprocess.STDOUT,
            env=env,
        )
        try:
            wait_ready(cotenant, log_reader, log_path, args.wait)
            if args.grace:
                time.sleep(args.grace)
            print(f"running: {' '.join(command)}")
            return run_command(command, env)
        finally:
            if cotenant.poll() is None:
                try:
                    os.kill(cotenant.pid, signal.SIGTERM)
                    try:
                        cotenant.wait(timeout=5)
                    except subprocess.TimeoutExpired:
                        os.kill(cotenant.pid, signal.SIGKILL)
                except ProcessLookupError:
                    pass  # cotenant exited between the poll and the signal


if __name__ == "__main__":
    try:
        sys.exit(main())
    except KeyboardInterrupt:
        # main()'s finally has already torn down the cotenant during unwinding.
        sys.exit(130)
