#!/usr/bin/env python3
"""LTX-2.5 distilled: first/last-frame guided video for seamless loops.

Rebuilt (flattened) from ComfyUI blueprint "First & Last Frame to Video (LTX-2.5)".

usage:
  ltx25_loop.py <promptfile> <image> <prefix> <W> <H> <FRAMES> <SEED> [mode=flf|i2v]
    mode flf : same image as first AND last frame guide -> seamless loop
    mode i2v : image as first frame only
"""
import json, sys, time, threading, subprocess, urllib.request, urllib.error

HOST = "http://127.0.0.1:8189"
UNET = "ltx-2.5-22b-distilled-transformer-comfy-int8-convrot.safetensors"
VVAE = "ltx-2.5-video-vae-bf16.safetensors"
AVAE = "ltx-2.5-audio-vae-bf16.safetensors"
CLIPNAME = "gemma4-12b-with-proj-ltx-2.5-comfy-int8-convrot.safetensors"
SIGMAS = "1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0"
NEG = ("blurry, out of focus, overexposed, underexposed, low contrast, washed out colors, "
       "excessive noise, grainy texture, poor lighting, flickering, motion blur, distorted proportions, "
       "deformed hands, extra limbs, watermark, text, logo, jpeg artifacts, still frame, static image")


def post(path, obj):
    req = urllib.request.Request(HOST + path, data=json.dumps(obj).encode(),
                                 headers={"Content-Type": "application/json"})
    try:
        return json.load(urllib.request.urlopen(req, timeout=120))
    except urllib.error.HTTPError as e:
        body = e.read().decode("utf-8", "replace")
        print("HTTP %s on %s\n%s" % (e.code, path, body[:4000]))
        try:
            det = json.loads(body)
            for k, v in (det.get("node_errors") or {}).items():
                print("  node", k, "->", json.dumps(v, ensure_ascii=False)[:600])
        except Exception:
            pass
        raise


def get(path):
    return json.load(urllib.request.urlopen(HOST + path, timeout=120))


def vram_watch(stop, out):
    while not stop.is_set():
        try:
            for line in subprocess.run(["nvidia-smi", "--query-gpu=memory.used", "--format=csv,noheader,nounits"],
                                       capture_output=True, text=True, timeout=5).stdout.split():
                out.append(int(line))
        except Exception:
            pass
        time.sleep(2)


def build(prompt, image, w, h, frames, seed, mode, prefix, strength=0.7):
    wf = {
        "285": {"class_type": "UNETLoader", "inputs": {"unet_name": UNET, "weight_dtype": "default"}},
        "286": {"class_type": "VAELoader", "inputs": {"vae_name": VVAE}},
        "287": {"class_type": "VAELoader", "inputs": {"vae_name": AVAE}},
        "288": {"class_type": "CLIPLoader", "inputs": {"clip_name": CLIPNAME, "type": "ltxv", "device": "default"}},
        "290": {"class_type": "LoadImage", "inputs": {"image": image, "upload": "image"}},
        "280": {"class_type": "LTXVPreprocess", "inputs": {"image": ["290", 0], "img_compression": 18}},
        "276": {"class_type": "CLIPTextEncode", "inputs": {"clip": ["288", 0], "text": prompt}},
        "275": {"class_type": "CLIPTextEncode", "inputs": {"clip": ["288", 0], "text": NEG}},
        "270": {"class_type": "LTXVConditioning", "inputs": {"positive": ["276", 0], "negative": ["275", 0],
                                                             "frame_rate": 24.0}},
        "268": {"class_type": "EmptyLTXVLatentVideo", "inputs": {"width": w, "height": h, "length": frames,
                                                                 "batch_size": 1}},
        "263": {"class_type": "LTXVAddGuide", "inputs": {"positive": ["270", 0], "negative": ["270", 1],
                                                         "vae": ["286", 0], "latent": ["268", 0], "image": ["280", 0],
                                                         "frame_idx": 0, "strength": strength}},
        "266": {"class_type": "LTXVEmptyLatentAudio", "inputs": {"frames_number": frames, "frame_rate": 24,
                                                                 "batch_size": 1, "audio_vae": ["287", 0]}},
        "264": {"class_type": "ManualSigmas", "inputs": {"sigmas": SIGMAS}},
        "271": {"class_type": "SamplerEulerAncestral", "inputs": {"eta": 0.0, "s_noise": 1.0}},
        "274": {"class_type": "RandomNoise", "inputs": {"noise_seed": seed}},
        "260": {"class_type": "VAEDecodeTiled", "inputs": {"samples": ["267", 2], "vae": ["286", 0],
                                                           "tile_size": 512, "overlap": 64, "temporal_size": 64,
                                                           "temporal_overlap": 16}},
        "261": {"class_type": "LTXVAudioVAEDecode", "inputs": {"samples": ["262", 1], "audio_vae": ["287", 0]}},
        "272": {"class_type": "CreateVideo", "inputs": {"images": ["260", 0], "audio": ["261", 0], "fps": 24.0}},
        "300": {"class_type": "SaveVideo", "inputs": {"video": ["272", 0], "filename_prefix": prefix,
                                                      "format": "auto"}},
    }
    if mode == "flf":
        wf["291"] = {"class_type": "LoadImage", "inputs": {"image": image, "upload": "image"}}
        wf["279"] = {"class_type": "LTXVPreprocess", "inputs": {"image": ["291", 0], "img_compression": 18}}
        wf["265"] = {"class_type": "LTXVAddGuide", "inputs": {"positive": ["263", 0], "negative": ["263", 1],
                                                              "vae": ["286", 0], "latent": ["263", 2],
                                                              "image": ["279", 0], "frame_idx": -1,
                                                              "strength": strength}}
        pos_neg = ["265", 0]
        neg_neg = ["265", 1]
        latent = ["265", 2]
    else:
        pos_neg, neg_neg, latent = ["263", 0], ["263", 1], ["263", 2]

    wf["259"] = {"class_type": "LTXVConcatAVLatent", "inputs": {"video_latent": latent, "audio_latent": ["266", 0]}}
    wf["273"] = {"class_type": "LTXVDualCFGGuider", "inputs": {"model": ["285", 0], "positive": pos_neg,
                                                               "negative": neg_neg, "video_cfg": 1.0,
                                                               "audio_cfg": 1.0}}
    wf["269"] = {"class_type": "SamplerCustomAdvanced", "inputs": {"noise": ["274", 0], "guider": ["273", 0],
                                                                   "sampler": ["271", 0], "sigmas": ["264", 0],
                                                                   "latent_image": ["259", 0]}}
    wf["262"] = {"class_type": "LTXVSeparateAVLatent", "inputs": {"av_latent": ["269", 1]}}
    wf["267"] = {"class_type": "LTXVCropGuides", "inputs": {"positive": pos_neg, "negative": neg_neg,
                                                            "latent": ["262", 0]}}
    return wf


def main():
    promptfile, image, prefix = sys.argv[1], sys.argv[2], sys.argv[3]
    w, h, frames, seed = (int(x) for x in sys.argv[4:8])
    mode = sys.argv[8] if len(sys.argv) > 8 else "flf"
    strength = float(sys.argv[9]) if len(sys.argv) > 9 else 0.7
    prompt = open(promptfile).read().strip()

    wf = build(prompt, image, w, h, frames, seed, mode, prefix, strength)
    stop, vram = threading.Event(), []
    threading.Thread(target=vram_watch, args=(stop, vram), daemon=True).start()

    t0 = time.time()
    r = post("/prompt", {"prompt": wf})
    pid = r["prompt_id"]
    print("prompt_id:", pid, "| mode=%s %dx%d frames=%d seed=%d" % (mode, w, h, frames, seed), flush=True)

    hist = None
    while time.time() - t0 < 5400:
        h_ = get("/history/" + pid)
        if pid in h_:
            hist = h_[pid]
            break
        el = time.time() - t0
        print("  ...%ds  vram=%s MiB" % (el, vram[-1] if vram else "?"), flush=True)
        time.sleep(30)
    elapsed = time.time() - t0
    stop.set()

    if hist is None:
        print("TIMEOUT after %.0fs" % elapsed)
        sys.exit(1)
    st = hist.get("status", {})
    print("elapsed: %.1fs  status=%s" % (elapsed, st.get("status_str")))
    print("vram_used MiB: start=%s peak=%s end=%s" % (vram[0] if vram else "?", max(vram) if vram else "?",
                                                      vram[-1] if vram else "?"))
    if st.get("status_str") == "error":
        for m in st.get("messages", []):
            print("ERR:", json.dumps(m)[:1200])
        sys.exit(2)
    for node_id, out in hist.get("outputs", {}).items():
        for key in ("images", "videos", "gifs"):
            for f in out.get(key, []):
                print("OUTPUT[%s]:" % key, json.dumps(f))


if __name__ == "__main__":
    main()
