#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""额外生成一个「参考音频转谱（SheetSage2）」UI 工作流，并做连线一致性自检。"""
from __future__ import annotations
import json, pathlib

DST = pathlib.Path("/home/zyw/ComfyUI/user/default/workflows/YuE2")
DST.mkdir(parents=True, exist_ok=True)

NOTE = """【参考音频 → 乐谱】SheetSage2 单独跑，不需要 YuE2

步骤：
  1) 把参考音频拷进 ComfyUI/input/（mp3/wav/flac 都行）
  2) LoadAudio 的下拉里选它
  3) Queue 就出：ABC 乐谱 + MIDI + 分段结构 + 和弦
     产物落在 ComfyUI/output/SheetSage2/<时间戳>/

实测（本机）：63 秒的歌 → 2.9 秒出谱（F 调 70BPM，19 个和弦，Vocal 35 音符）

选项：
  • melody_only=true 只出旋律声部（做改编/翻唱用，配 YuE2 的 cot=melody）
  • max_seconds>0 可只转前 N 秒（长歌省时间）
  • abc_error_mode：ABC 导出失败时的兜底策略
      default=strict，可改 snap_invalid_notes / return_midi_only

MERT2 父模型用的是本地 models/MERT2/MERT-v2-FullSong
（完整性校验通过，无需改任何哈希）。
"""

TRANS = "SheetSage2Transcribe"


def node(nid, type_, pos, size, inputs, outputs, widgets, order, title=None):
    n = {"id": nid, "type": type_, "pos": list(pos), "size": list(size), "flags": {},
         "order": order, "mode": 0, "inputs": inputs, "outputs": outputs,
         "widgets_values": widgets,
         "properties": {"Node name for S&R": type_}}
    if title:
        n["title"] = title
    return n


def wi(name, type_, link=None):   # widget 型 input
    return {"name": name, "type": type_, "widget": {"name": name}, "link": link}


def si(name, type_, link):        # socket 型 input
    return {"name": name, "type": type_, "link": link}


nodes = [
    node(1, "LoadAudio", (40, 120), (300, 100),
         [wi("audio", "COMBO", None)],
         [{"name": "AUDIO", "type": "AUDIO", "links": [1], "slot_index": 0}],
         ["ref-song.flac"], 0, title="参考音频（选 input 目录里的文件）"),

    node(2, "SheetSage2Loader", (40, 300), (330, 130),
         [wi("model", "COMBO"), wi("device", "COMBO"), wi("dtype", "COMBO"),
          wi("mert2_model", "COMBO")],
         [{"name": "model", "type": "SS2_MODEL", "links": [2], "slot_index": 0}],
         ["SheetSage2", "cuda", "bfloat16", "MERT-v2-FullSong"], 1,
         title="SheetSage2 模型加载"),

    node(3, TRANS, (440, 120), (420, 320),
         [si("model", "SS2_MODEL", 2), si("audio", "AUDIO", 1),
          wi("melody_only", "BOOLEAN"), wi("preset", "COMBO"), wi("max_seconds", "FLOAT"),
          wi("save_outputs", "BOOLEAN"), wi("abc_error_mode", "COMBO"),
          wi("overlap_seconds", "FLOAT"), wi("lookahead_seconds", "FLOAT"),
          wi("auto_unload", "BOOLEAN")],
         [{"name": "abc", "type": "STRING", "links": [], "slot_index": 0},
          {"name": "structure", "type": "STRING", "links": [], "slot_index": 1},
          {"name": "info", "type": "STRING", "links": [], "slot_index": 2},
          {"name": "midi", "type": "YUE2_MIDI", "links": [3], "slot_index": 3}],
         [False, "default", 0.0, True, "strict", -1.0, -1.0, False], 2,
         title="转谱（出 ABC / MIDI / 分段）"),

    node(4, "YuE2SaveMidiFile", (920, 140), (300, 80),
         [si("midi", "YUE2_MIDI", 3), wi("filename_prefix", "STRING")],
         [{"name": "path", "type": "STRING", "links": [], "slot_index": 0}],
         ["YuE2/ref-midi"], 3, title="保存 MIDI"),

    node(5, "Note", (40, -300), (640, 340), [], [], [NOTE], 4),
]

links = [
    [1, 1, 0, 3, 1, "AUDIO"],
    [2, 2, 0, 3, 0, "SS2_MODEL"],
    [3, 3, 3, 4, 0, "YUE2_MIDI"],
]

doc = {
    "id": "yue2-sheetsage2-transcribe",
    "revision": 0,
    "last_node_id": 5,
    "last_link_id": 3,
    "nodes": nodes,
    "links": links,
    "groups": [],
    "config": {},
    "extra": {"ds": {"scale": 1.0, "offset": [0, 0]}},
    "version": 0.4,
}

# ---- 自检：连线引用的两端必须存在、类型一致、slot 合法 ----
errs = []
by_id = {n["id"]: n for n in doc["nodes"]}
for lid, oid, oslot, tid, tslot, ltype in links:
    o, t = by_id.get(oid), by_id.get(tid)
    if not o or not t:
        errs.append(f"link {lid}: 端节点不存在"); continue
    if oslot >= len(o["outputs"]):
        errs.append(f"link {lid}: 源 slot {oslot} 越界"); continue
    if tslot >= len(t["inputs"]):
        errs.append(f"link {lid}: 目标 slot {tslot} 越界"); continue
    if ltype not in (o["outputs"][oslot]["type"], t["inputs"][tslot]["type"]):
        errs.append(f"link {lid}: 类型不符 {ltype} vs {o['outputs'][oslot]['type']}/{t['inputs'][tslot]['type']}")
    if lid not in o["outputs"][oslot]["links"]:
        errs.append(f"link {lid}: 源 output.links 未登记")
    if t["inputs"][tslot].get("link") != lid:
        errs.append(f"link {lid}: 目标 input.link 不匹配")

for n in doc["nodes"]:
    for i, inp in enumerate(n["inputs"]):
        lid = inp.get("link")
        if lid is None:
            continue
        hit = [l for l in links if l[0] == lid]
        if not hit:
            errs.append(f"node {n['id']} input {i} 引用了不存在的 link {lid}")
        elif hit[0][3] != n["id"] or hit[0][4] != i:
            errs.append(f"node {n['id']} input {i} 与 link {lid} 的落点不一致")

out = DST / "YuE2-5-参考音频转谱-16GB.json"
out.write_text(json.dumps(doc, ensure_ascii=False, indent=2), encoding="utf-8")
print("自检:", "全部通过 ✅" if not errs else "发现问题 ❌")
for e in errs:
    print("  -", e)
print("写入", out, out.stat().st_size, "字节")
print("目录:")
for p in sorted(DST.iterdir()):
    print("  ", p.name, p.stat().st_size)
