import os
import re
import subprocess

import folder_paths


class MergeSegmentVideos:
    """将 MiniMaxH3_segments/<运行名>/ 下的分段 MP4 按顺序合成一个完整视频。

    依赖容器内的 ffmpeg(VideoHelperSuite 自带的 imageio-ffmpeg 即可)。
    """

    @classmethod
    def INPUT_TYPES(cls):
        return {
            "required": {
                "run_name": (
                    "STRING",
                    {
                        "default": "robot_laser_003",
                        "multiline": False,
                        "tooltip": "与分段工作流的「运行名称」保持一致",
                    },
                ),
                "mode": (
                    ["copy_then_reencode", "copy", "reencode"],
                    {
                        "default": "copy_then_reencode",
                        "tooltip": "copy 无损最快;失败自动转 reencode 重新编码保证能合成",
                    },
                ),
                "only_latest_per_segment": (
                    "BOOLEAN",
                    {
                        "default": True,
                        "tooltip": "每段只取最新一次生成(重试产生的旧 MP4 自动忽略)",
                    },
                ),
            },
        }

    RETURN_TYPES = ("STRING",)
    RETURN_NAMES = ("merged_path",)
    FUNCTION = "merge"
    CATEGORY = "video/merge"
    OUTPUT_NODE = True

    def merge(self, run_name, mode, only_latest_per_segment):
        run_name = run_name.strip().strip("/")
        if not re.fullmatch(r"[A-Za-z0-9_\-\u4e00-\u9fff]+", run_name):
            raise ValueError(f"运行名称含不支持的字符: {run_name}")

        run_dir = os.path.join(
            folder_paths.get_output_directory(),
            "MiniMaxH3_segments",
            run_name,
        )
        if not os.path.isdir(run_dir):
            raise FileNotFoundError(f"找不到分段目录: {run_dir}")

        files = [
            f
            for f in os.listdir(run_dir)
            if f.startswith("segment_") and f.endswith(".mp4")
        ]
        if not files:
            raise FileNotFoundError(f"{run_dir} 下没有 segment_*.mp4 文件")

        def seg_key(name):
            m = re.match(r"segment_(\d+)_(\d+)_", name)
            if not m:
                return (10**9, 10**9, name)
            return (int(m.group(1)), int(m.group(2)), name)

        files.sort(key=seg_key)

        if only_latest_per_segment:
            latest = {}
            for f in files:
                m = re.match(r"segment_(\d+)_(\d+)_", f)
                if not m:
                    continue
                seg = int(m.group(1))
                cnt = int(m.group(2))
                if seg not in latest or cnt > latest[seg][0]:
                    latest[seg] = (cnt, f)
            files = [v[1] for _, v in sorted(latest.items())]

        ffmpeg = self._find_ffmpeg()
        list_file = os.path.join(run_dir, "_concat_list.txt")
        with open(list_file, "w", encoding="utf-8") as fh:
            for name in files:
                path = os.path.join(run_dir, name)
                escaped = path.replace("'", "'\\''")
                fh.write(f"file '{escaped}'\n")

        out_path = os.path.join(run_dir, f"{run_name}_merged.mp4")
        attempts = []
        if mode in ("copy", "copy_then_reencode"):
            attempts.append((["-c", "copy"], "copy"))
        if mode in ("reencode", "copy_then_reencode"):
            attempts.append(
                (
                    [
                        "-c:v",
                        "libx264",
                        "-preset",
                        "medium",
                        "-crf",
                        "18",
                        "-c:a",
                        "aac",
                        "-b:a",
                        "192k",
                    ],
                    "reencode",
                ),
            )

        last_error = ""
        for extra, label in attempts:
            cmd = [
                ffmpeg,
                "-y",
                "-f",
                "concat",
                "-safe",
                "0",
                "-i",
                list_file,
            ] + extra + [out_path]
            proc = subprocess.run(cmd, capture_output=True, text=True)
            if proc.returncode == 0 and os.path.exists(out_path) and os.path.getsize(out_path) > 0:
                print(f"[MergeSegmentVideos] 合成成功({label}): {out_path}")
                print(f"[MergeSegmentVideos] 片段数: {len(files)}")
                return (out_path,)
            last_error = (proc.stderr or "")[-2000:]

        raise RuntimeError(f"ffmpeg 合成失败:\n{last_error}")

    @staticmethod
    def _find_ffmpeg():
        try:
            import imageio_ffmpeg

            return imageio_ffmpeg.get_ffmpeg_exe()
        except Exception:
            pass
        import shutil

        exe = shutil.which("ffmpeg")
        if exe:
            return exe
        raise RuntimeError("找不到 ffmpeg,请确认 VideoHelperSuite 已安装")


NODE_CLASS_MAPPINGS = {
    "MergeSegmentVideos": MergeSegmentVideos,
}

NODE_DISPLAY_NAME_MAPPINGS = {
    "MergeSegmentVideos": "Merge Segment Videos 🎞 (ffmpeg 合成)",
}
