"""用 ComfyUI HTTP API 运行前端格式工作流的最小工具。

用法（客户端部分完成后）:
    python comfy_run.py workflows/z_image_tur.json \
        --prompt "..." --width 1280 --height 720 --seed 42 \
        --prefix cover_42 --out images/cover_42.png
"""
import argparse
import json
import os
import time
import urllib.request
import uuid

SERVER = "http://127.0.0.1:8188"


def frontend_to_api(wf: dict) -> dict:
    """前端格式工作流 JSON -> API 格式图 {node_id: {class_type, inputs}}。

    参数取值用 widgets_values_named（按名称精确映射，避免
    control_after_generate 之类纯控件值混入位置序导致错位）。
    """
    links = {l[0]: (l[1], l[2]) for l in wf["links"]}  # link_id -> (from_node, from_slot)
    graph = {}
    for n in wf["nodes"]:
        inputs = {}
        named = n.get("widgets_values_named", {})
        positional = iter(n.get("widgets_values", []))
        for inp in n.get("inputs", []):
            name = inp["name"]
            if "widget" in inp:
                # 优先按名称映射；无 widgets_values_named 时退回位置序（可能错位）
                inputs[name] = named[name] if named else next(positional)
            elif inp.get("link") is not None:
                from_node, from_slot = links[inp["link"]]
                inputs[name] = [str(from_node), from_slot]
        graph[str(n["id"])] = {"class_type": n["type"], "inputs": inputs}
    return graph


def apply_overrides(graph: dict, overrides: dict) -> dict:
    """按 {node_id: {input_name: value}} 覆盖图中参数，原地修改并返回。"""
    for nid, kv in overrides.items():
        for k, v in kv.items():
            graph[str(nid)]["inputs"][k] = v
    return graph


def queue_prompt(server: str, graph: dict) -> str:
    body = json.dumps({"prompt": graph, "client_id": str(uuid.uuid4())}).encode()
    req = urllib.request.Request(server + "/prompt", data=body,
                                 headers={"Content-Type": "application/json"})
    with urllib.request.urlopen(req, timeout=30) as resp:
        return json.load(resp)["prompt_id"]


def wait_for_prompt(server: str, prompt_id: str, timeout: float = 600) -> list:
    deadline = time.monotonic() + timeout
    while time.monotonic() < deadline:
        with urllib.request.urlopen(server + f"/history/{prompt_id}", timeout=30) as resp:
            h = json.load(resp).get(prompt_id)
        if h:
            status = h.get("status", {})
            if status.get("status_str") == "success":
                return parse_image_records(h.get("outputs", {}))
            if status.get("status_str") == "error":
                raise RuntimeError("ComfyUI 执行失败: " + json.dumps(h)[:500])
        time.sleep(1.0)
    raise TimeoutError(f"等待生成超时 {timeout}s")


def parse_image_records(outputs: dict) -> list:
    for node_out in outputs.values():
        if node_out.get("images"):
            return node_out["images"]
    return []


def build_view_url(server: str, record: dict) -> str:
    return (server + "/view?filename=" + record["filename"]
            + "&subfolder=" + record.get("subfolder", "")
            + "&type=" + record.get("type", "output"))


def download_image(server: str, record: dict, dest: str) -> None:
    os.makedirs(os.path.dirname(os.path.abspath(dest)) or ".", exist_ok=True)
    urllib.request.urlretrieve(build_view_url(server, record), dest)


def main(argv=None):
    p = argparse.ArgumentParser(description="用 ComfyUI HTTP API 跑前端格式工作流")
    p.add_argument("workflow", help="前端格式工作流 JSON 路径")
    p.add_argument("--prompt", required=True, help="正面提示词")
    p.add_argument("--negative", default="low quality, bad anatomy, extra digits, "
                    "missing digits, extra limbs, missing limbs")
    p.add_argument("--width", type=int, default=1024)
    p.add_argument("--height", type=int, default=1024)
    p.add_argument("--seed", type=int, default=42)
    p.add_argument("--prefix", default="ComfyUI", help="SaveImage 文件名前缀")
    p.add_argument("--out", required=True, help="输出 PNG 路径")
    p.add_argument("--server", default=SERVER)
    args = p.parse_args(argv)

    wf = json.load(open(args.workflow, encoding="utf-8"))
    graph = apply_overrides(frontend_to_api(wf), {
        "67": {"text": args.prompt},
        "71": {"text": args.negative},
        "68": {"width": args.width, "height": args.height},
        "70": {"seed": args.seed},
        "9": {"filename_prefix": args.prefix},
    })
    prompt_id = queue_prompt(args.server, graph)
    print("已提交 prompt_id:", prompt_id)
    recs = wait_for_prompt(args.server, prompt_id)
    if not recs:
        raise RuntimeError("没有输出图片")
    download_image(args.server, recs[0], args.out)
    print("已保存:", args.out)


if __name__ == "__main__":
    main()
