#!/usr/bin/env python3
"""Build Boogu LoRA (boonude) workflow JSONs in ComfyUI UI format.

Graph (Base variant):
  UNETLoader -> LoraLoader(model,clip) -> ModelSamplingAuraFlow -> KSampler
  CLIPLoader -> LoraLoader.clip -> CLIPTextEncode(+/-) -> KSampler
  VAELoader -> VAEDecode ; EmptyLatentImage -> KSampler ; SaveImage
"""
import json

BASE_UNET = "boogu_image_base_fp8_scaled.safetensors"
TURBO_UNET = "boogu_image_turbo_fp8_scaled.safetensors"
LORA = "boonude.safetensors"
CLIP = "qwen3vl_8b_fp8_scaled.safetensors"
VAE = "ae.safetensors"
PROMPT = "A beautiful portrait of a woman, soft studio lighting, detailed skin texture, professional photography"
# Official Base negative prompt from the Boogu project's own workflow (boogu_base_t2i.json)
OFFICIAL_NEGATIVE = ("(((deformed))), blurry, over saturation, bad anatomy, disfigured, poorly drawn face, "
                     "mutation, mutated, (extra_limb), (ugly), (poorly drawn hands), fused fingers, "
                     "messy drawing, broken legs censor, censored, censor_bar")


def make_workflow(unet, prefix, sampler_widgets, auraflow, negative=OFFICIAL_NEGATIVE):
    """Return (workflow_dict). Nodes are defined with explicit edges; links are auto-numbered."""
    # node specs: id -> (type, widgets, inputs[(name,type)], outputs[(name,type)])
    specs = []
    edges = []  # (from_node, from_slot, to_node, to_slot, type)

    def add(ntype, widgets, inputs, outputs, pos, size=(300, 100)):
        nid = len(specs) + 1
        specs.append((nid, ntype, widgets, inputs, outputs, pos, size))
        return nid

    unet = add("UNETLoader", [unet, "default"],
               [("unet_name", "COMBO")], [("MODEL", "MODEL")], (0, 0))
    clip = add("CLIPLoader", [CLIP, "boogu", "default"],
               [("clip_name", "COMBO")], [("CLIP", "CLIP")], (0, 300))
    lora = add("LoraLoader", [LORA, 1.0, 1.0],
               [("model", "MODEL"), ("clip", "CLIP")],
               [("MODEL", "MODEL"), ("CLIP", "CLIP")], (360, 0))
    edges.append((unet, 0, lora, 0, "MODEL"))
    edges.append((clip, 0, lora, 1, "CLIP"))

    cur_model = lora
    cur_clip = lora
    if auraflow:
        af = add("ModelSamplingAuraFlow", [3.16],
                 [("model", "MODEL")], [("MODEL", "MODEL")], (720, 0))
        edges.append((lora, 0, af, 0, "MODEL"))
        cur_model = af

    vae = add("VAELoader", [VAE],
              [("vae_name", "COMBO")], [("VAE", "VAE")], (0, 480))
    pos = add("CLIPTextEncode", [PROMPT],
              [("clip", "CLIP")], [("CONDITIONING", "CONDITIONING")], (400, 300))
    edges.append((cur_clip, 1, pos, 0, "CLIP"))
    neg = add("CLIPTextEncode", [negative],
              [("clip", "CLIP")], [("CONDITIONING", "CONDITIONING")], (400, 520))
    edges.append((cur_clip, 1, neg, 0, "CLIP"))
    latent = add("EmptyLatentImage", [1024, 1024, 1],
                 [], [("LATENT", "LATENT")], (760, 480))
    ksampler = add("KSampler", sampler_widgets,
                   [("model", "MODEL"), ("positive", "CONDITIONING"),
                    ("negative", "CONDITIONING"), ("latent_image", "LATENT")],
                   [("LATENT", "LATENT")], (1100, 0))
    edges.append((cur_model, 0, ksampler, 0, "MODEL"))
    edges.append((pos, 0, ksampler, 1, "CONDITIONING"))
    edges.append((neg, 0, ksampler, 2, "CONDITIONING"))
    edges.append((latent, 0, ksampler, 3, "LATENT"))
    decode = add("VAEDecode", [],
                 [("samples", "LATENT"), ("vae", "VAE")],
                 [("IMAGE", "IMAGE")], (1460, 100))
    edges.append((ksampler, 0, decode, 0, "LATENT"))
    edges.append((vae, 0, decode, 1, "VAE"))
    save = add("SaveImage", [prefix],
               [("images", "IMAGE")], [], (1820, 100))
    edges.append((decode, 0, save, 0, "IMAGE"))

    # assign link ids
    links = []
    for i, (fn, fs, tn, ts, t) in enumerate(edges, start=1):
        links.append([i, fn, fs, tn, ts, t])
    # build nodes with consistent per-slot link ids
    in_link = {}  # (node, slot) -> link id
    out_link = {}  # (node, slot) -> list of link ids
    for i, (lid, fn, fs, tn, ts, t) in enumerate(links, start=1):
        in_link[(tn, ts)] = lid
        out_link.setdefault((fn, fs), []).append(lid)

    nodes = []
    for nid, ntype, widgets, inputs, outputs, pos, size in specs:
        node = {
            "id": nid,
            "type": ntype,
            "pos": list(pos),
            "size": list(size),
            "flags": {},
            "order": nid - 1,
            "mode": 0,
            "inputs": [
                {
                    "name": name,
                    "type": t,
                    "link": in_link.get((nid, i)),
                    **({"widget": {"name": name}} if t == "COMBO" else {}),
                }
                for i, (name, t) in enumerate(inputs)
            ],
            "outputs": [
                {"name": name, "type": t, "links": out_link.get((nid, i), []), "shape": 3}
                for i, (name, t) in enumerate(outputs)
            ],
            "properties": {"Node name for S&R": ntype},
            "widgets_values": widgets,
        }
        nodes.append(node)

    return {
        "last_node_id": len(specs),
        "last_link_id": len(links),
        "nodes": nodes,
        "links": links,
        "groups": [],
        "config": {},
        "extra": {"ds": {"scale": 0.5, "offset": [0, 0]}},
        "version": 0.4,
    }


base_wf = make_workflow(
    BASE_UNET, "Boogu_lora_base",
    # official Base recipe: 50 steps / cfg 4.0 (from boogu_base_t2i.json)
    [123456789012345, "randomize", 50, 4.0, "dpmpp_2m", "simple", 1.0],
    auraflow=True,
)
turbo_wf = make_workflow(
    TURBO_UNET, "Boogu_lora_turbo",
    [987654321098765, "randomize", 4, 1.0, "lcm", "sgm_uniform", 1.0],
    auraflow=False,
)

with open("boogu_lora_t2i_base.json", "w", encoding="utf-8") as f:
    json.dump(base_wf, f, ensure_ascii=False, indent=1)
with open("boogu_lora_t2i_turbo.json", "w", encoding="utf-8") as f:
    json.dump(turbo_wf, f, ensure_ascii=False, indent=1)

# sanity checks
for name, wf in (("base", base_wf), ("turbo", turbo_wf)):
    ids = {n["id"] for n in wf["nodes"]}
    assert all(l[1] in ids and l[3] in ids for l in wf["links"]), name
    # every node input link must reference an existing link id
    all_lids = {l[0] for l in wf["links"]}
    for n in wf["nodes"]:
        for inp in n["inputs"]:
            if inp["link"] is not None:
                assert inp["link"] in all_lids, f"{name} node{n['id']} input {inp['name']}"
    print(f"{name}: nodes={len(wf['nodes'])} links={len(wf['links'])} OK")
