#!/usr/bin/env python3
"""Render a pause-safe talking-head timeline with semantic B-roll and ducked BGM."""
from __future__ import annotations

import argparse
import json
import subprocess
from pathlib import Path


def ff(value: float) -> str:
    return f"{value:.6f}".rstrip("0").rstrip(".")


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("timeline")
    parser.add_argument("plan")
    parser.add_argument("output")
    parser.add_argument("--filter-report")
    args = parser.parse_args()

    timeline_path = Path(args.timeline).resolve()
    plan_path = Path(args.plan).resolve()
    timeline = json.loads(timeline_path.read_text(encoding="utf-8"))
    plan = json.loads(plan_path.read_text(encoding="utf-8"))
    clips = timeline["clips"]
    if not clips:
        raise SystemExit("Timeline has no clips")

    source = Path(plan["source"]).resolve()
    music = Path(plan["music"]).resolve()
    width = int(plan.get("width", 1080))
    height = int(plan.get("height", 1920))
    fps = int(plan.get("fps", 30))
    expected = float(timeline["expected_duration"])
    concat_duration = sum(float(item["source_out"]) - float(item["source_in"]) for item in clips)
    tail = max(0.0, expected - concat_duration)

    command = ["ffmpeg", "-y", "-v", "warning", "-i", str(source)]
    for item in plan.get("broll", []):
        command += ["-i", str(Path(item["path"]).resolve())]
    music_input = 1 + len(plan.get("broll", []))
    command += ["-stream_loop", "-1", "-i", str(music)]

    graph: list[str] = []
    split_v = "".join(f"[sv{i}]" for i in range(len(clips)))
    split_a = "".join(f"[sa{i}]" for i in range(len(clips)))
    graph.append(f"[0:v]split={len(clips)}{split_v}")
    graph.append(f"[0:a]asplit={len(clips)}{split_a}")
    concat_inputs = []
    for index, item in enumerate(clips):
        source_in = float(item["source_in"])
        source_out = float(item["source_out"])
        graph.append(
            f"[sv{index}]trim=start={ff(source_in)}:end={ff(source_out)},"
            f"setpts=PTS-STARTPTS[cv{index}]"
        )
        graph.append(
            f"[sa{index}]atrim=start={ff(source_in)}:end={ff(source_out)},"
            f"asetpts=PTS-STARTPTS[ca{index}]"
        )
        concat_inputs.append(f"[cv{index}][ca{index}]")

    graph.append(
        "".join(concat_inputs)
        + f"concat=n={len(clips)}:v=1:a=1[vrec][arec]"
    )
    graph.append(
        f"[vrec]scale={width}:{height}:force_original_aspect_ratio=increase,"
        f"crop={width}:{height},"
        "eq=contrast=1.04:brightness=-0.01:saturation=1.02:gamma=0.98,"
        "colorbalance=rs=0.015:bs=-0.01,"
        f"tpad=stop_mode=clone:stop_duration={ff(tail)},"
        "setsar=1,format=yuv420p[vbase]"
    )

    previous = "vbase"
    for index, item in enumerate(plan.get("broll", []), start=1):
        start = float(item["start"])
        duration = float(item["duration"])
        source_in = float(item.get("source_in", 0))
        source_out = source_in + duration
        label = f"b{index}"
        graph.append(
            f"[{index}:v]trim=start={ff(source_in)}:end={ff(source_out)},"
            f"setpts=PTS-STARTPTS+{ff(start)}/TB,"
            f"scale={width}:{height}:force_original_aspect_ratio=increase,"
            f"crop={width}:{height},"
            "eq=contrast=1.06:brightness=-0.015:saturation=0.92:gamma=0.97,"
            "colorbalance=rs=0.02:bs=-0.015,setsar=1,format=yuv420p"
            f"[{label}]"
        )
        out = f"vo{index}"
        graph.append(
            f"[{previous}][{label}]overlay=0:0:"
            f"enable='between(t,{ff(start)},{ff(start + duration)})':eof_action=pass[{out}]"
        )
        previous = out

    graph.append(
        f"[arec]apad=pad_dur={ff(tail)},atrim=duration={ff(expected)},"
        "highpass=f=70,lowpass=f=15500,"
        "acompressor=threshold=-18dB:ratio=3:attack=10:release=180:makeup=3,"
        "aformat=sample_rates=48000:channel_layouts=stereo[voice]"
    )
    graph.append(
        f"[{music_input}:a]atrim=duration={ff(expected)},asetpts=PTS-STARTPTS,"
        f"afade=t=in:st=0:d=1.2,afade=t=out:st={ff(max(0, expected - 2.0))}:d=2,"
        "volume=0.14,aformat=sample_rates=48000:channel_layouts=stereo[bed]"
    )
    graph.append(
        "[bed][voice]sidechaincompress=threshold=0.025:ratio=12:"
        "attack=15:release=350:makeup=1[ducked]"
    )
    graph.append(
        "[voice][ducked]amix=inputs=2:weights='1 0.85':normalize=0,"
        "loudnorm=I=-14:TP=-1.5:LRA=9[aout]"
    )

    filter_graph = ";\n".join(graph)
    report = Path(args.filter_report).resolve() if args.filter_report else Path(args.output).with_suffix(".filter.txt")
    report.parent.mkdir(parents=True, exist_ok=True)
    report.write_text(filter_graph + "\n", encoding="utf-8")

    output = Path(args.output).resolve()
    output.parent.mkdir(parents=True, exist_ok=True)
    command += [
        "-filter_complex_script", str(report),
        "-map", f"[{previous}]",
        "-map", "[aout]",
        "-r", str(fps),
        "-c:v", "libx264",
        "-preset", "slow",
        "-crf", "17",
        "-profile:v", "high",
        "-pix_fmt", "yuv420p",
        "-c:a", "aac",
        "-b:a", "192k",
        "-ar", "48000",
        "-ac", "2",
        "-movflags", "+faststart",
        "-t", ff(expected),
        str(output),
    ]
    completed = subprocess.run(command)
    if completed.returncode:
        raise SystemExit(completed.returncode)
    print(json.dumps({
        "output": str(output),
        "duration": expected,
        "clips": len(clips),
        "broll": len(plan.get("broll", [])),
        "filter_report": str(report),
        "status": "ok",
    }, ensure_ascii=False))


if __name__ == "__main__":
    main()
