#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""RTX VSR 2x 放大（分块执行）—— 避免整片一次放大把内存打满。

用法:
  python3 rtx_upscale.py <input.mp4> <outdir> [--chunk-frames 130] [--scale 2.0] [--prefix rtx]

流程:
  1. 本地 ffmpeg 按帧切块（-c copy，保持原始编码）
  2. 逐块上传到 GPU 机 ComfyUI input 目录
  3. VHS_LoadVideoPath -> DaSiWa_RTX_UpscalerRefiner(VSR) -> VHS_VideoCombine
  4. 下载每块结果，最后 ffmpeg concat 成整片
"""
import argparse
import json
import subprocess
import sys
import time
from pathlib import Path

import requests

sys.path.insert(0, "/home/zyw/Downloads/dasiwa-h3-research")
import v18_pipeline as p  # noqa: E402

COMFY_URL = "http://192.168.31.31:8189"


def build_api(chunk_abs_path: str, w: int, h: int, fps: float, prefix: str,
              scale: float, quality: str) -> dict:
    return {
        "l1": {"class_type": "VHS_LoadVideoPath",
               "inputs": {"video": chunk_abs_path, "force_rate": 0,
                          "custom_width": 0, "custom_height": 0, "frame_load_cap": 0,
                          "skip_first_frames": 0, "select_every_nth": 1,
                          "start_frame": 0, "end_frame": 0, "maximum_frames": 0,
                          "force_size": "Disabled", "custom_audio": 0,
                          "audio_start_index": 0, "pingpong": False}},
        "r1": {"class_type": "DaSiWa_RTX_UpscalerRefiner",
               "inputs": {"images": ["l1", 0],
                          "denoise": False, "denoise_quality": "Low",
                          "deblur": False, "deblur_quality": "Low",
                          "upscale": "VSR", "upscale_quality": quality,
                          "resize_type": "Scale", "scale": scale,
                          "megapixels": 2.0, "width": w, "height": h,
                          "divisible_by": "8", "ratio_preset": "16:9",
                          "resize_method": "Center Crop (Fill)", "device_id": 0,
                          "empty_cache": False, "use_mmap": False,
                          "auto_unload_models": True}},
        "c1": {"class_type": "VHS_VideoCombine",
               "inputs": {"images": ["r1", 0], "audio": ["l1", 2],
                          "frame_rate": fps, "loop_count": 0,
                          "filename_prefix": f"video/{prefix}",
                          "format": "video/h264-mp4", "pix_fmt": "yuv420p",
                          "crf": 16, "save_output": True, "pingpong": False}},
    }


def probe(path: Path):
    out = subprocess.run(
        ["ffprobe", "-v", "error", "-select_streams", "v:0",
         "-count_frames", "-show_entries", "stream=nb_read_frames,width,height,r_frame_rate",
         "-show_entries", "format=duration", "-of", "json", str(path)],
        capture_output=True, text=True).stdout
    d = json.loads(out)
    st = d["streams"][0]
    frames = int(st.get("nb_read_frames") or 0)
    fps = eval(st["r_frame_rate"]) if st.get("r_frame_rate") else 24.0
    return {"frames": frames, "w": int(st["width"]), "h": int(st["height"]),
            "fps": float(fps), "duration": float(d["format"]["duration"])}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("input")
    ap.add_argument("outdir")
    ap.add_argument("--chunk-frames", type=int, default=130)
    ap.add_argument("--scale", type=float, default=2.0)
    ap.add_argument("--quality", default="High")
    ap.add_argument("--prefix", default="xianjia_rtx")
    ap.add_argument("--batch", default="")
    args = ap.parse_args()

    src = Path(args.input).resolve()
    outdir = Path(args.outdir).resolve()
    chunkdir = outdir / "chunks"
    updir = outdir / "upscaled"
    chunkdir.mkdir(parents=True, exist_ok=True)
    updir.mkdir(parents=True, exist_ok=True)

    info = probe(src)
    target_w = int(info["w"] * args.scale)
    target_h = int(info["h"] * args.scale)
    print(f"[rtx] 源：{info['frames']} 帧 {info['w']}x{info['h']} @{info['fps']}fps "
          f"{info['duration']:.2f}s → 目标 {target_w}x{target_h}")

    # 1) 切块（画面按帧精确 select，音轨同步 atrim；VHS 需要音轨，故视频内必须带音频）
    n = (info["frames"] + args.chunk_frames - 1) // args.chunk_frames
    chunks = []
    for i in range(n):
        start = i * args.chunk_frames
        end = min(start + args.chunk_frames - 1, info["frames"] - 1)
        cp = chunkdir / f"{src.stem}_chunk{i:02d}.mp4"
        if not cp.exists():
            fc = (f"[0:v]select='between(n\\,{start}\\,{end})',setpts=PTS-STARTPTS[v];"
                  f"[0:a]atrim=start={start / info['fps']:.6f}:end={(end + 2) / info['fps']:.6f},"
                  f"asetpts=PTS-STARTPTS[a]")
            subprocess.run(["ffmpeg", "-y", "-v", "error", "-i", str(src),
                            "-filter_complex", fc, "-map", "[v]", "-map", "[a]",
                            "-vsync", "0", "-c:v", "libx264", "-crf", "14",
                            "-pix_fmt", "yuv420p", "-c:a", "aac", "-shortest",
                            str(cp)], check=True)
        chunks.append((cp, start, end))
    print(f"[rtx] 切成 {len(chunks)} 块（每块 ≤{args.chunk_frames} 帧，带音轨）")

    tag = args.batch or time.strftime("%H%M%S")
    results = []
    for i, (cp, start, end) in enumerate(chunks):
        remote = f"rtx_{tag}_{i:02d}.mp4"
        p._upload_file_to_comfy(str(cp), remote)
        api = build_api(f"/home/zyw/ComfyUI/input/{remote}", target_w, target_h,
                        info["fps"], f"{args.prefix}_{tag}_{i:02d}", args.scale, args.quality)
        (outdir / f"api_rtx_{i:02d}.json").write_text(
            json.dumps(api, ensure_ascii=False, indent=1), encoding="utf-8")
        print(f"[rtx] 提交块 {i+1}/{len(chunks)}（起始帧 {start}）...")
        t0 = time.time()
        entry = p.submit_and_wait(api, timeout_min=90)
        vi = p.extract_video_output(entry)
        lp = p.download_video(vi, updir)
        merged = updir / f"{src.stem}_rtx_{i:02d}.mp4"
        if lp != merged:
            if merged.exists():
                merged.unlink()
            lp.replace(merged)
        results.append(merged)
        print(f"[rtx] 块 {i+1} 完成：{merged.name}（{(time.time()-t0)/60:.1f} 分钟）")

    # 2) 拼接
    listfile = outdir / "concat_rtx.txt"
    listfile.write_text("".join(f"file '{r.resolve()}'\n" for r in results), encoding="utf-8")
    final = outdir / f"{src.stem}_rtx{args.scale:g}x.mp4"
    subprocess.run(["ffmpeg", "-y", "-v", "error", "-f", "concat", "-safe", "0",
                    "-i", str(listfile), "-c", "copy", str(final)], check=True)
    print(f"[rtx] 成片: {final}")
    print(json.dumps(probe(final), ensure_ascii=False))


if __name__ == "__main__":
    main()
