#!/usr/bin/env -S uv run python
"""Small public grader for the Phase 2 inference-systems Lab."""

from __future__ import annotations

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


LAB_ROOT = Path(__file__).resolve().parent


class Exercise(NamedTuple):
    name: str
    title: str
    points: int
    command: tuple[str, ...]


EXERCISES = (
    Exercise(
        "m2-c",
        "cached multi-head attention",
        100,
        (
            "-m",
            "unittest",
            "discover",
            "-s",
            "grader_tests",
            "-p",
            "test_m2_c_cached_mha.py",
            "-v",
        ),
    ),
)


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Run all public Lab gates or one named exercise."
    )
    parser.add_argument(
        "exercises",
        nargs="*",
        metavar="EXERCISE",
        help="one or more of: " + ", ".join(item.name for item in EXERCISES),
    )
    parser.add_argument(
        "--list",
        action="store_true",
        help="list available exercises without running them",
    )
    return parser.parse_args()


def selected_exercises(names: list[str]) -> tuple[Exercise, ...]:
    if not names:
        return EXERCISES

    by_name = {exercise.name: exercise for exercise in EXERCISES}
    unknown = [name for name in names if name not in by_name]
    if unknown:
        choices = ", ".join(by_name)
        raise SystemExit(f"unknown exercise: {', '.join(unknown)}; choose from {choices}")
    return tuple(by_name[name] for name in names)


def main() -> int:
    args = parse_args()
    if args.list:
        for exercise in EXERCISES:
            print(f"{exercise.name:6}  {exercise.points:>3} points  {exercise.title}")
        return 0

    exercises = selected_exercises(args.exercises)
    available_points = sum(exercise.points for exercise in exercises)
    earned_points = 0
    failed = []
    environment = os.environ.copy()
    source_directory = environment.get("LAB_SOURCE_DIR", "src")
    source_path = (LAB_ROOT / source_directory).resolve()
    try:
        source_path.relative_to(LAB_ROOT)
    except ValueError as error:
        raise SystemExit("LAB_SOURCE_DIR must stay inside the Lab directory") from error
    if not source_path.is_dir():
        raise SystemExit(f"Lab source directory does not exist: {source_path}")
    environment["PYTHONPATH"] = os.pathsep.join(
        filter(None, (str(source_path), environment.get("PYTHONPATH")))
    )
    environment.setdefault("OPENBLAS_NUM_THREADS", "1")
    environment.setdefault("OMP_NUM_THREADS", "1")

    for exercise in exercises:
        print(f"\n== {exercise.name}: {exercise.title} ==", flush=True)
        result = subprocess.run(
            (sys.executable, *exercise.command),
            cwd=LAB_ROOT,
            env=environment,
            check=False,
        )
        if result.returncode == 0:
            earned_points += exercise.points
            print(f"== {exercise.name}: PASS ({exercise.points}/{exercise.points}) ==")
        else:
            failed.append(exercise.name)
            print(f"== {exercise.name}: FAIL (0/{exercise.points}) ==")

    print(f"\nScore: {earned_points}/{available_points}")
    if failed:
        print("Still working: " + ", ".join(failed))
        return 1
    print("All selected exercises passed.")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
