from __future__ import annotations

import math
import random
import shutil
import subprocess
import wave
from pathlib import Path


class VideoDirector:
    """Premium short-form reel builder — cinematic beats, not slideshows."""

    MOTION = {
        "zoom_in": ("min(zoom+0.0018,1.4)", "iw/2-(iw/zoom/2)", "ih/2-(ih/zoom/2)"),
        "zoom_out": ("if(eq(on,1),1.4,max(zoom-0.0018,1.0))", "iw/2-(iw/zoom/2)", "ih/2-(ih/zoom/2)"),
        "pan_left": ("1.28", "if(eq(on,1),iw*0.18,max(x-2.2,0))", "ih/2-(ih/zoom/2)"),
        "pan_right": ("1.28", "if(eq(on,1),0,min(x+2.2,iw-iw/zoom))", "ih/2-(ih/zoom/2)"),
        "punch_in": ("if(lte(on,10),1+0.045*on,1.35)", "iw/2-(iw/zoom/2)", "ih/2-(ih/zoom/2)"),
        "parallax": ("1.22+0.0008*on", "iw/2-(iw/zoom/2)+8*sin(on/12)", "ih/2-(ih/zoom/2)"),
        "glitch": ("1.3", "iw/2-(iw/zoom/2)+if(lt(mod(on,12),2),15,0)", "ih/2-(ih/zoom/2)"),
    }

    TRANSITIONS = ("fade", "wipeleft", "wiperight", "slideup", "slidedown", "circlecrop", "pixelize", "hblur")

    def __init__(self, max_seconds: int = 30, enable_audio: bool = False):
        self.max_seconds = max(18, max_seconds)
        self.enable_audio = enable_audio

    def build(
        self,
        frames: list[Path],
        captions: list[str],
        out_path: Path,
        motion_style: str = "cinematic",
        motions: list[str] | None = None,
        durations: list[float] | None = None,
    ) -> Path:
        if not frames:
            raise RuntimeError("No frames for video")
        if not shutil.which("ffmpeg"):
            still = out_path.with_suffix(".png")
            shutil.copy(frames[0], still)
            return still

        out_path.parent.mkdir(parents=True, exist_ok=True)
        work = out_path.parent / "_video_work"
        if work.exists():
            shutil.rmtree(work, ignore_errors=True)
        work.mkdir(exist_ok=True)

        n = len(frames)
        if durations and len(durations) == n:
            pers = durations
        else:
            base = self.max_seconds / max(1, n)
            # Hook shorter, middle longer
            pers = []
            for i in range(n):
                if i == 0:
                    pers.append(max(2.0, base * 0.7))
                elif i == n - 1:
                    pers.append(max(2.5, base * 0.85))
                else:
                    pers.append(max(3.5, base * 1.15))
            # normalize to max_seconds
            total = sum(pers)
            pers = [p * (self.max_seconds / total) for p in pers]

        fps = 30
        clip_paths: list[Path] = []
        motion_cycle = list(self.MOTION.keys())
        for i, frame in enumerate(frames):
            motion = (motions[i] if motions and i < len(motions) else None) or motion_cycle[i % len(motion_cycle)]
            if motion_style == "glitch":
                motion = "glitch" if i % 2 == 0 else "punch_in"
            elif motion_style == "kinetic":
                motion = "punch_in" if i % 2 == 0 else "zoom_in"
            elif motion_style == "parallax":
                motion = "parallax"
            elif motion_style == "broadcast":
                motion = "pan_right" if i % 2 else "pan_left"
            caption = captions[i] if i < len(captions) else ""
            clip = work / f"clip_{i:02d}.mp4"
            self._animate_clip(frame, clip, duration=pers[i], fps=fps, motion=motion, caption=caption, beat_index=i)
            clip_paths.append(clip)

        merged = work / "merged.mp4"
        self._crossfade_concat(clip_paths, merged, pers)

        music = work / "score.wav"
        self._generate_score(music, duration=self.max_seconds + 2, style=motion_style)

        # Audio mux disabled at product level — copy silent/merged video only.
        # Keep score generation for future AudioModule wiring.
        if getattr(self, "enable_audio", False):
            self._mux_audio(merged, music, out_path)
        else:
            import shutil as _sh
            _sh.copy(merged, out_path)
        return out_path

    def _animate_clip(self, image, out, *, duration, fps, motion, caption, beat_index):
        frames_n = max(1, int(duration * fps))
        z, x, y = self.MOTION.get(motion, self.MOTION["zoom_in"])
        safe = caption.replace(":", "\\:").replace("'", "").replace('"', "")[:70]
        draw = ""
        if safe:
            # Kinetic caption bar — stronger on hook beat
            alpha = "0.65" if beat_index == 0 else "0.5"
            fontsize = 42 if beat_index == 0 else 34
            draw = (
                f",drawbox=x=32:y=ih-200:w=iw-64:h=120:color=black@{alpha}:t=fill"
                f",drawtext=text='{safe}':fontcolor=white:fontsize={fontsize}:"
                f"x=(w-text_w)/2:y=h-165"
            )
        vf = (
            f"scale=1500:1500:force_original_aspect_ratio=increase,"
            f"crop=1500:1500,"
            f"zoompan=z='{z}':x='{x}':y='{y}':d={frames_n}:s=1080x1080:fps={fps},"
            f"format=yuv420p{draw}"
        )
        cmd = [
            "ffmpeg", "-y", "-loop", "1", "-i", str(image),
            "-vf", vf, "-t", str(duration), "-r", str(fps),
            "-pix_fmt", "yuv420p", str(out),
        ]
        subprocess.run(cmd, check=True, capture_output=True)

    def _crossfade_concat(self, clips: list[Path], out: Path, pers: list[float]) -> None:
        if len(clips) == 1:
            shutil.copy(clips[0], out)
            return
        fade = 0.4
        inputs = []
        for c in clips:
            inputs.extend(["-i", str(c)])
        filter_parts = []
        prev = "[0:v]"
        offset = pers[0] - fade
        rng = random.Random(len(clips) * 17)
        for i in range(1, len(clips)):
            out_label = f"[v{i}]" if i < len(clips) - 1 else "[vout]"
            transition = rng.choice(self.TRANSITIONS)
            filter_parts.append(
                f"{prev}[{i}:v]xfade=transition={transition}:duration={fade}:offset={max(0.1, offset):.3f}{out_label}"
            )
            prev = out_label
            if i < len(pers):
                offset += pers[i] - fade
        filt = ";".join(filter_parts)
        cmd = ["ffmpeg", "-y", *inputs, "-filter_complex", filt, "-map", "[vout]", "-pix_fmt", "yuv420p", str(out)]
        try:
            subprocess.run(cmd, check=True, capture_output=True)
        except subprocess.CalledProcessError:
            lst = out.parent / "fallback.txt"
            lst.write_text("\n".join(f"file '{c.resolve()}'" for c in clips), encoding="utf-8")
            subprocess.run(
                ["ffmpeg", "-y", "-f", "concat", "-safe", "0", "-i", str(lst), "-c", "copy", str(out)],
                check=True,
                capture_output=True,
            )

    def _generate_score(self, path: Path, duration: float, style: str = "cinematic") -> None:
        """Layered ambient bed + soft rhythmic pulses (royalty-free synth)."""
        rate = 44100
        n = int(rate * duration)
        if style == "glitch":
            freqs = [98, 147, 220, 311]
            pulse = 7.0
            noise = 0.04
        elif style in {"kinetic", "product_reveal"}:
            freqs = [130.81, 196.0, 261.63, 329.63]
            pulse = 4.5
            noise = 0.02
        elif style == "broadcast":
            freqs = [110, 165, 220]
            pulse = 2.5
            noise = 0.015
        elif style == "data_motion":
            freqs = [146.83, 220, 293.66]
            pulse = 5.0
            noise = 0.025
        else:
            freqs = [98.0, 146.83, 196.0, 246.94]
            pulse = 2.2
            noise = 0.018

        samples = bytearray()
        for i in range(n):
            t = i / rate
            env = min(1.0, t * 3) * min(1.0, (duration - t) * 2.5)
            pulse_env = 0.5 + 0.5 * abs(math.sin(2 * math.pi * pulse * t))
            # soft click every pulse for UI feel
            click = 0.0
            if (t * pulse) % 1.0 < 0.02:
                click = 0.12 * math.sin(2 * math.pi * 1800 * t)
            val = 0.0
            for j, f in enumerate(freqs):
                val += math.sin(2 * math.pi * f * t) * (0.9 / (j + 1))
            val = val / len(freqs)
            val += noise * math.sin(2 * math.pi * 40 * t) * random.random()
            val = (val * 0.16 * env * pulse_env) + click * env
            s = max(-32767, min(32767, int(val * 32767)))
            samples += int(s).to_bytes(2, "little", signed=True)

        with wave.open(str(path), "w") as wf:
            wf.setnchannels(1)
            wf.setsampwidth(2)
            wf.setframerate(rate)
            wf.writeframes(bytes(samples))

    def _mux_audio(self, video: Path, audio: Path, out: Path) -> None:
        cmd = [
            "ffmpeg", "-y",
            "-i", str(video),
            "-i", str(audio),
            "-c:v", "copy",
            "-c:a", "aac",
            "-shortest",
            "-movflags", "+faststart",
            str(out),
        ]
        subprocess.run(cmd, check=True, capture_output=True)
