#!/usr/bin/env python3
"""h3-run.py — build and run one H3 Turbo generation against a ComfyUI endpoint.

Bakes in every session-verified gotcha:
  - model paths carry their subfolder prefixes (diffusion_models/, vae/,
    text_encoders/) while the LoRA name stays bare — the HTTP-400 trap
  - frames snap UP to the 17k+5 grid at 24fps
  - talks to http://127.0.0.1:8188 by default (the SSH tunnel endpoint —
    RunPod's public 8188 usually refuses; run
    `ssh -i ~/.runpod/ssh/runpodctl-ssh-key -f -N -L 8188:127.0.0.1:8188 root@<ip> -p <port>` first)
  - records t0 (submit) / t1 (first progress) / t2 (sampling done) /
    t3 (outputs in /history) to timings.csv — feed these to
    lvrged_factory_job action=finish as load_s/sample_s/decode_s

Usage:
  h3-run.py --prompt "..." [--seconds 15] [--width 864] [--height 480]
            [--steps 8] [--seed 42] [--prefix h3/out] [--first-frame img.png]
            [--endpoint http://127.0.0.1:8188]
"""
import argparse, csv, json, os, sys, time, urllib.request, urllib.error, uuid

UNET = "diffusion_models/minimax_h3_fl2va_pruned_int8_convrot.safetensors"
CLIP = "text_encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
VAE_V = "vae/minimax_h3_video_vae_fp16.safetensors"
VAE_A = "vae/minimax_h3_audio_vae_fp32.safetensors"
LORA = "minimax_h3_turbo_v4_step600_ema.safetensors"  # bare — loras/ prefix would 400


def frames_for(seconds: float) -> int:
    n = max(5, round(seconds * 24))
    return n + (5 - n % 17) % 17  # snap UP to the 17k+5 grid


def http(url, payload=None, timeout=60):
    data = json.dumps(payload).encode() if payload is not None else None
    req = urllib.request.Request(url, data=data, headers={"Content-Type": "application/json"} if data else {})
    with urllib.request.urlopen(req, timeout=timeout) as r:
        return json.loads(r.read())


def build_graph(a, frames, first_frame_name=None):
    g = {
        "1": {"class_type": "UNETLoader", "inputs": {"unet_name": UNET, "weight_dtype": "default"}},
        "2": {"class_type": "MiniMaxH3TurboLoRA", "inputs": {"model": ["1", 0], "lora_name": LORA, "strength": 1.0, "low_vram": False}},
        "3": {"class_type": "CLIPLoader", "inputs": {"clip_name": CLIP, "type": "minimax", "device": "default"}},
        "4": {"class_type": "VAELoader", "inputs": {"vae_name": VAE_V}},
        "5": {"class_type": "VAELoader", "inputs": {"vae_name": VAE_A}},
        "6": {"class_type": "MiniMaxH3ImageToVideo", "inputs": {"clip": ["3", 0], "vae": ["4", 0], "prompt": a.prompt, "width": a.width, "height": a.height, "length": frames}},
        "7": {"class_type": "MiniMaxH3TurboSampler", "inputs": {}},
        "8": {"class_type": "BasicGuider", "inputs": {"model": ["2", 0], "conditioning": ["6", 0]}},  # 0.30.0: 'conditioning', not 'positive'
        "9": {"class_type": "RandomNoise", "inputs": {"noise_seed": a.seed}},
        "10": {"class_type": "BasicScheduler", "inputs": {"model": ["2", 0], "scheduler": "beta", "steps": a.steps, "denoise": 1.0}},
        "11": {"class_type": "SamplerCustomAdvanced", "inputs": {"noise": ["9", 0], "guider": ["8", 0], "sampler": ["7", 0], "sigmas": ["10", 0], "latent_image": ["6", 1]}},
        "12": {"class_type": "VAEDecode", "inputs": {"samples": ["11", 0], "vae": ["4", 0]}},
        "13": {"class_type": "VAEDecodeAudio", "inputs": {"samples": ["11", 0], "vae": ["5", 0]}},
        "14": {"class_type": "CreateVideo", "inputs": {"images": ["12", 0], "fps": 24, "audio": ["13", 0]}},
        "15": {"class_type": "SaveVideo", "inputs": {"video": ["14", 0], "filename_prefix": a.prefix, "format": "auto", "codec": "auto"}},
    }
    if first_frame_name:  # chain clips: last frame of clip N = first_frame of N+1
        g["20"] = {"class_type": "LoadImage", "inputs": {"image": first_frame_name}}
        g["6"]["inputs"]["first_frame"] = ["20", 0]
    return g


def upload_image(endpoint, path):
    import mimetypes
    boundary = uuid.uuid4().hex
    name = os.path.basename(path)
    ctype = mimetypes.guess_type(name)[0] or "application/octet-stream"
    with open(path, "rb") as f:
        filedata = f.read()
    body = (f"--{boundary}\r\nContent-Disposition: form-data; name=\"image\"; filename=\"{name}\"\r\n"
            f"Content-Type: {ctype}\r\n\r\n").encode() + filedata + f"\r\n--{boundary}--\r\n".encode()
    req = urllib.request.Request(f"{endpoint}/upload/image", data=body,
                                 headers={"Content-Type": f"multipart/form-data; boundary={boundary}"})
    with urllib.request.urlopen(req, timeout=120) as r:
        return json.loads(r.read())["name"]


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--prompt", required=True, help="include the audio in the prompt — H3 generates it natively")
    p.add_argument("--seconds", type=float, default=15)
    p.add_argument("--width", type=int, default=864)
    p.add_argument("--height", type=int, default=480)
    p.add_argument("--steps", type=int, default=8)
    p.add_argument("--seed", type=int, default=42)
    p.add_argument("--prefix", default="h3/out")
    p.add_argument("--first-frame", help="local image path; uploaded and wired as the i2v keyframe")
    p.add_argument("--endpoint", default=os.environ.get("COMFY_ENDPOINT", "http://127.0.0.1:8188"))
    p.add_argument("--label", default="run")
    a = p.parse_args()

    if a.width % 32 or a.height % 32:
        sys.exit(f"width/height must be multiples of 32 (got {a.width}x{a.height})")
    if a.width * a.height < 1_000_000 and (a.width, a.height) != (864, 480):
        print(f"note: {a.width}x{a.height} is below the ~1MP quality floor (864x480 is the benchmark exception)", file=sys.stderr)

    frames = frames_for(a.seconds)
    first_name = upload_image(a.endpoint, a.first_frame) if a.first_frame else None
    graph = build_graph(a, frames, first_name)

    t0 = time.time()
    try:
        resp = http(f"{a.endpoint}/prompt", {"prompt": graph, "client_id": "lvrged-factory"})
    except urllib.error.HTTPError as e:
        body = e.read().decode()[:1200]
        print(f"HTTP {e.code} — validation failed. The body names the node and valid values:\n{body}", file=sys.stderr)
        print("(most common cause: model path prefixes — see the comfyui skill)", file=sys.stderr)
        sys.exit(1)
    except urllib.error.URLError as e:
        print(f"cannot reach {a.endpoint}: {e.reason}\nis the SSH tunnel up? ssh -f -N -L 8188:127.0.0.1:8188 root@<ip> -p <port>", file=sys.stderr)
        sys.exit(1)
    pid = resp["prompt_id"]
    print(f"[{a.label}] queued {pid} · {a.width}x{a.height} · {frames}f (~{frames/24:.0f}s) · {a.steps} steps · seed {a.seed}")

    t1 = t2 = None
    last_frac = 0.0
    while True:
        time.sleep(5)
        try:
            hist = http(f"{a.endpoint}/history/{pid}", timeout=30)
        except Exception:
            hist = {}
        entry = hist.get(pid) or {}
        if entry.get("outputs"):
            t3 = time.time()
            break
        try:
            pr = http(f"{a.endpoint}/progress", timeout=15)
            frac = float(pr.get("progress") or 0)
        except Exception:
            frac = last_frac
        if frac > 0 and t1 is None:
            t1 = time.time()
        if frac >= 1 and t2 is None:
            t2 = time.time()
        if frac != last_frac:
            print(f"  sampling {frac*100:3.0f}%", file=sys.stderr)
            last_frac = frac
        if time.time() - t0 > 3600:
            sys.exit("gave up after 1h — run the gpu-ops failure protocol")

    files = [f.get("filename") for o in entry["outputs"].values() for k in ("videos", "gifs", "images") for f in o.get(k, [])]
    load_s = round((t1 or t3) - t0, 1)
    sample_s = round(((t2 or t3) - (t1 or t3)), 1)
    decode_s = round(t3 - (t2 or t3), 1)
    total = round(t3 - t0, 1)
    print(f"[{a.label}] DONE {total}s (load {load_s} · sample {sample_s} · decode {decode_s}) → {files}")
    print(f"  fetch: scp -P <ssh_port> root@<ip>:<ComfyUI>/output/{a.prefix}* .")

    log = os.path.join(os.path.dirname(os.path.abspath(__file__)), "timings.csv")
    new = not os.path.exists(log)
    with open(log, "a", newline="") as f:
        w = csv.writer(f)
        if new:
            w.writerow(["ts", "label", "prompt_id", "wxh", "frames", "steps", "seed", "load_s", "sample_s", "decode_s", "total_s", "files"])
        w.writerow([int(t0), a.label, pid, f"{a.width}x{a.height}", frames, a.steps, a.seed, load_s, sample_s, decode_s, total, ";".join(files)])


if __name__ == "__main__":
    main()
