#!/usr/bin/env python3
"""Insert boonude LoRA into the OFFICIAL Boogu Turbo t2i template (minimal surgical edit).

- Source: venv-installed official template (untouched copy kept separately)
- Edit: add LoraLoader node inside the template's subgraph, between
  UNETLoader(2)->[LoraLoader]->KSampler(32) and CLIPLoader(7)->[LoraLoader]->CLIPTextEncode(11)
- Everything else stays byte-identical.
"""
import json

SRC = "/home/zyw/ComfyUI/venv/lib/python3.12/site-packages/comfyui_workflow_templates_json/templates/image_boogu_image_0_1_turbo_t2i.json"
OUT = "/home/zyw/ComfyUI/user/default/workflows/boogu_image_0_1_turbo_t2i_official_lora.json"

d = json.load(open(SRC, encoding="utf-8"))
sub = d["definitions"]["subgraphs"][0]
nodes = {n["id"]: n for n in sub["nodes"]}
links = sub["links"]

# sanity checks on expected structure
assert nodes[2]["type"] == "UNETLoader", nodes[2]["type"]
assert nodes[7]["type"] == "CLIPLoader", nodes[7]["type"]
assert nodes[32]["type"] == "KSampler", nodes[32]["type"]
assert nodes[11]["type"] == "CLIPTextEncode", nodes[11]["type"]

# collect ALL node/link ids across the whole document (top graph + subgraphs)
all_node_ids = set()
all_link_ids = set()


def collect(obj):
    if isinstance(obj, dict):
        if "nodes" in obj and isinstance(obj["nodes"], list):
            for n in obj["nodes"]:
                all_node_ids.add(n["id"])
        if "links" in obj and isinstance(obj["links"], list):
            for l in obj["links"]:
                if isinstance(l, dict):
                    all_link_ids.add(l["id"])
                elif isinstance(l, (list, tuple)):
                    all_link_ids.add(l[0])
                else:  # plain int link id (node output slot links)
                    all_link_ids.add(l)
        for v in obj.values():
            collect(v)
    elif isinstance(obj, list):
        for v in obj:
            collect(v)


collect(d)
new_node_id = max(all_node_ids) + 1
lid = max(all_link_ids) + 1


def find_link(pred):
    for l in links:
        if pred(l):
            return l
    return None


l_model = find_link(lambda l: l["origin_id"] == 2 and l["origin_slot"] == 0 and l["target_id"] == 32 and l["target_slot"] == 0)
l_clip = find_link(lambda l: l["origin_id"] == 7 and l["origin_slot"] == 0 and l["target_id"] == 11 and l["target_slot"] == 0)
assert l_model is not None and l_clip is not None, "expected links not found"

# remove the two links being rewired
links.remove(l_model)
links.remove(l_clip)

# new links: UNET->Lora, CLIP->Lora, Lora->KSampler, Lora->CLIPTextEncode
L1 = {"id": lid, "origin_id": 2, "origin_slot": 0, "target_id": new_node_id, "target_slot": 0, "type": "MODEL"}; lid += 1
L2 = {"id": lid, "origin_id": 7, "origin_slot": 0, "target_id": new_node_id, "target_slot": 1, "type": "CLIP"}; lid += 1
L3 = {"id": lid, "origin_id": new_node_id, "origin_slot": 0, "target_id": 32, "target_slot": 0, "type": "MODEL"}; lid += 1
L4 = {"id": lid, "origin_id": new_node_id, "origin_slot": 1, "target_id": 11, "target_slot": 0, "type": "CLIP"}; lid += 1
links.extend([L1, L2, L3, L4])

# update per-node slot link fields
nodes[2]["outputs"][0]["links"] = [L1["id"]]
nodes[7]["outputs"][0]["links"] = [L2["id"]]
for inp in nodes[32]["inputs"]:
    if inp["name"] == "model":
        inp["link"] = L3["id"]
for inp in nodes[11]["inputs"]:
    if inp["name"] == "clip":
        inp["link"] = L4["id"]

# new LoraLoader node (placed next to UNETLoader)
lora_node = {
    "id": new_node_id,
    "type": "LoraLoader",
    "pos": [nodes[2]["pos"][0] + 420, nodes[2]["pos"][1]],
    "size": [320, 130],
    "flags": {},
    "order": len(sub["nodes"]),
    "mode": 0,
    "inputs": [
        {"name": "model", "type": "MODEL", "link": L1["id"]},
        {"name": "clip", "type": "CLIP", "link": L2["id"]},
    ],
    "outputs": [
        {"name": "MODEL", "type": "MODEL", "links": [L3["id"]], "shape": 3},
        {"name": "CLIP", "type": "CLIP", "links": [L4["id"]], "shape": 3},
    ],
    "properties": {
        "Node name for S&R": "LoraLoader",
        "models": [{"name": "boonude.safetensors", "url": "https://civitai.red/models/2726783", "directory": "loras"}],
    },
    "widgets_values": ["boonude.safetensors", 1.0, 1.0],
}
sub["nodes"].append(lora_node)

# bump template-level counters to cover new ids
d["last_node_id"] = max(d["last_node_id"], new_node_id)
d["last_link_id"] = max(d["last_link_id"], lid - 1)

# --- validate ---
valid_ids = all_node_ids | {new_node_id, -10, -20}  # -10/-20 = subgraph input/output nodes
valid_lids = all_link_ids | {L1["id"], L2["id"], L3["id"], L4["id"]}
for l in links:
    assert l["origin_id"] in valid_ids and l["target_id"] in valid_ids, f"bad link {l}"
    assert l["id"] in valid_lids
for n in sub["nodes"]:
    for inp in n.get("inputs", []):
        if inp.get("link") is not None:
            assert inp["link"] in valid_lids, f"node {n['id']} input {inp['name']} link {inp['link']}"
for n in sub["nodes"]:
    for out in n.get("outputs", []):
        for lk in out.get("links", []):
            assert lk in valid_lids

with open(OUT, "w", encoding="utf-8") as f:
    json.dump(d, f, ensure_ascii=False, indent=1)
print("OK:", OUT)
print("inserted LoraLoader node", new_node_id, "| links", L1["id"], L2["id"], L3["id"], L4["id"])
