#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""跑 SheetSage2：把参考音频转成 ABC/MIDI/分段结构（顺带验证 MERT2 完整性校验是否拦得住）。"""
from __future__ import annotations
import argparse, json, os, shutil, sys, time, urllib.request, urllib.parse

HOST = "http://192.168.31.31:8189"
OPENER = urllib.request.build_opener(urllib.request.ProxyHandler({}))


def api(path, data=None, timeout=120):
    body = json.dumps(data).encode() if data is not None else None
    req = urllib.request.Request(HOST + path, data=body,
                                headers={"Content-Type": "application/json"} if body else {})
    with OPENER.open(req, timeout=timeout) as r:
        raw = r.read()
    try:
        return json.loads(raw)
    except Exception:
        return raw


def build(audio_name, melody_only=False, preset="default", max_seconds=0.0, abc_error_mode="snap_invalid_notes"):
    return {
        "1": {"class_type": "LoadAudio", "inputs": {"audio": audio_name}},
        "2": {"class_type": "SheetSage2Loader", "inputs": {
            "model": "SheetSage2", "device": "cuda", "dtype": "bfloat16",
            "mert2_model": "MERT-v2-FullSong"}},
        "3": {"class_type": "SheetSage2Transcribe", "inputs": {
            "model": ["2", 0], "audio": ["1", 0],
            "melody_only": melody_only, "preset": preset,
            "max_seconds": max_seconds, "save_outputs": True,
            "abc_error_mode": abc_error_mode}},
        "4": {"class_type": "YuE2SaveMidiFile", "inputs": {
            "midi": ["3", 3], "filename_prefix": "YuE2/ref-midi"}},
    }


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--audio", required=True, help="ComfyUI input 目录里的文件名")
    p.add_argument("--melody-only", action="store_true")
    p.add_argument("--abc-error-mode", default="snap_invalid_notes",
                   choices=["strict", "snap_invalid_notes", "skip_invalid_notes", "fallback_full", "return_midi_only"])
    p.add_argument("--outdir", default="/home/zyw/Downloads/dl-hub/18-YuE2音乐生成")
    p.add_argument("--tag", default="sheetsage")
    a = p.parse_args()

    res = api("/prompt", {"prompt": build(a.audio, a.melody_only, abc_error_mode=a.abc_error_mode), "client_id": "dsh-ss2"})
    if "prompt_id" not in res:
        print("!! 提交失败:", json.dumps(res, ensure_ascii=False)[:3000]); sys.exit(1)
    pid = res["prompt_id"]
    print(f"[排队] prompt_id={pid} audio={a.audio}", flush=True)

    t0 = time.time()
    while True:
        time.sleep(10)
        el = time.time() - t0
        h = api(f"/history/{pid}")
        if pid in h:
            rec = h[pid]; st = rec.get("status", {})
            print(f"\n[完成] 用时 {el/60:.1f} 分钟 status={st.get('status_str')} completed={st.get('completed')}", flush=True)
            for m in st.get("messages") or []:
                if m[0] in ("execution_error", "execution_interrupted"):
                    print("!! 错误:", json.dumps(m[1], ensure_ascii=False)[:2000])
            outs = rec.get("outputs", {})
            print("[产物]", json.dumps(outs, ensure_ascii=False)[:1500], flush=True)
            os.makedirs(a.outdir, exist_ok=True)
            for nid, o in outs.items():
                for key in ("audio", "images", "gifs", "files"):
                    for item in o.get(key, []) or []:
                        fn = item.get("filename")
                        if not fn:
                            continue
                        url = (f"/view?filename={urllib.parse.quote(fn)}"
                               f"&subfolder={urllib.parse.quote(item.get('subfolder',''))}"
                               f"&type={item.get('type','output')}")
                        dst = os.path.join(a.outdir, f"{a.tag}_{fn}")
                        with OPENER.open(HOST + url, timeout=300) as r, open(dst, "wb") as w:
                            shutil.copyfileobj(r, w)
                        print(f"[下载] {dst} {os.path.getsize(dst)/1e6:.2f} MB", flush=True)
            return
        print(f"  等待中… {el/60:.1f} 分钟", flush=True)


if __name__ == "__main__":
    main()
