#!/usr/bin/env python3
"""调用 GPU 机 ComfyUI(8189) 的 Qwen-Image-2.1 批量文生图，并把结果拉回本机。

用法: python3 gen_qwen21.py prompts.json [输出目录] [--steps N] [--limit N]
"""
import json, os, sys, time, urllib.request, urllib.parse

HOST = "http://192.168.31.31:8189"
UNET = "qwen_image_2.1_int8_convrot.safetensors"
CLIP = "qwen3vl_8b_int8_convrot.safetensors"
VAE = "qwen_image_2.1_vae_bf16.safetensors"

# 局域网直连，绕过本机 ALL_PROXY
OPENER = urllib.request.build_opener(urllib.request.ProxyHandler({}))


def post(path, obj, timeout=120):
    req = urllib.request.Request(HOST + path, data=json.dumps(obj).encode(),
                                 headers={"Content-Type": "application/json"})
    return json.load(OPENER.open(req, timeout=timeout))


def get_json(path, timeout=120):
    return json.load(OPENER.open(HOST + path, timeout=timeout))


def build_wf(p, steps):
    return {
        "1": {"class_type": "UNETLoader", "inputs": {"unet_name": UNET, "weight_dtype": "default"}},
        "2": {"class_type": "CLIPLoader", "inputs": {"clip_name": CLIP, "type": "qwen_image", "device": "default"}},
        "3": {"class_type": "VAELoader", "inputs": {"vae_name": VAE}},
        "4": {"class_type": "TextEncodeQwenImage21",
              "inputs": {"clip": ["2", 0], "prompt": p["prompt"], "negative_prompt": p.get("negative", ""),
                         "resolution": 1024, "vae": ["3", 0]}},
        "5": {"class_type": "EmptyLatentImage",
              "inputs": {"width": p.get("width", 1536), "height": p.get("height", 864), "batch_size": 1}},
        "6": {"class_type": "KSampler",
              "inputs": {"model": ["1", 0], "positive": ["4", 0], "negative": ["4", 1],
                         "latent_image": ["5", 0], "seed": p.get("seed", 1), "steps": steps,
                         "cfg": 1.0, "sampler_name": "euler", "scheduler": "simple", "denoise": 1.0}},
        "7": {"class_type": "VAEDecode", "inputs": {"samples": ["6", 0], "vae": ["3", 0]}},
        "8": {"class_type": "SaveImage", "inputs": {"filename_prefix": p["name"], "images": ["7", 0]}},
    }


def main():
    cfg_path = sys.argv[1]
    outdir = sys.argv[2] if len(sys.argv) > 2 and not sys.argv[2].startswith("--") else "."
    steps_override = None
    limit = None
    for i, a in enumerate(sys.argv):
        if a == "--steps":
            steps_override = int(sys.argv[i + 1])
        if a == "--limit":
            limit = int(sys.argv[i + 1])

    items = json.load(open(cfg_path, encoding="utf-8"))
    if limit:
        items = items[:limit]
    os.makedirs(outdir, exist_ok=True)

    # 健康检查
    st = get_json("/system_stats", timeout=20)
    dev = st["devices"][0]
    print(f"ComfyUI OK: {dev['name']} vram_total={dev['vram_total']/2**20:.0f}MiB "
          f"vram_free={dev['vram_free']/2**20:.0f}MiB", flush=True)

    results = []
    for i, p in enumerate(items, 1):
        steps = steps_override or p.get("steps", 25)
        wf = build_wf(p, steps)
        t0 = time.time()
        r = post("/prompt", {"prompt": wf})
        pid = r["prompt_id"]
        print(f"[{i}/{len(items)}] {p['name']} {p.get('width')}x{p.get('height')} steps={steps} "
              f"seed={p.get('seed')} -> {pid}", flush=True)
        hist, err = None, None
        while time.time() - t0 < 2400:
            h = get_json("/history/" + pid)
            if pid in h:
                hist = h[pid]
                break
            time.sleep(3)
        el = time.time() - t0
        if hist is None:
            print(f"    TIMEOUT after {el:.1f}s", flush=True)
            continue
        status = hist.get("status", {})
        if status.get("status_str") == "error":
            for m in status.get("messages", []):
                print("    ERR:", json.dumps(m, ensure_ascii=False)[:600], flush=True)
            continue
        files = []
        for _nid, out in hist.get("outputs", {}).items():
            for im in out.get("images", []):
                q = urllib.parse.urlencode({"filename": im["filename"],
                                            "subfolder": im.get("subfolder", ""), "type": "output"})
                data = OPENER.open(HOST + "/view?" + q, timeout=180).read()
                dst = os.path.join(outdir, im["filename"])
                open(dst, "wb").write(data)
                files.append(dst)
                print(f"    saved {dst} ({len(data)} bytes)  {el:.1f}s ({el/steps:.2f}s/step)", flush=True)
        results.append({"name": p["name"], "elapsed_s": round(el, 1), "seed": p.get("seed"),
                        "size": f"{p.get('width')}x{p.get('height')}", "steps": steps, "files": files})

    json.dump(results, open(os.path.join(outdir, "_results.json"), "w", encoding="utf-8"),
              ensure_ascii=False, indent=2)
    print(f"done: {len(results)}/{len(items)} ok", flush=True)


if __name__ == "__main__":
    main()
