"""
Video Ingestion Module for VideoDetection.
Sử dụng FFmpeg trực tiếp qua subprocess để stream raw frames mà không gây nghẽn I/O đĩa cứng.
"""
from __future__ import annotations

import json
import subprocess
from dataclasses import dataclass
from pathlib import Path
from typing import Generator, List, Optional, Tuple

import numpy as np

from ..core.paths import TEMP_DIR, get_ffmpeg_path, get_ffprobe_path


@dataclass
class VideoMetadata:
    path: Path
    duration_sec: float
    width: int
    height: int
    fps: float
    total_frames: int


def get_video_metadata(video_path: Path) -> VideoMetadata:
    """Lấy thông tin kỹ thuật của video bằng ffprobe portable."""
    ffprobe_cmd = [
        get_ffprobe_path(),
        "-v", "error",
        "-select_streams", "v:0",
        "-show_entries", "stream=width,height,r_frame_rate,duration,nb_frames",
        "-show_entries", "format=duration",
        "-of", "json",
        str(video_path),
    ]

    result = subprocess.run(
        ffprobe_cmd,
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        text=True,
        check=False,
    )

    if result.returncode != 0:
        raise RuntimeError(f"Không thể đọc metadata video '{video_path}': {result.stderr}")

    info = json.loads(result.stdout)
    stream = info.get("streams", [{}])[0]
    fmt = info.get("format", {})

    width = int(stream.get("width", 0))
    height = int(stream.get("height", 0))

    # Xử lý frame rate dạng phân số "30/1" hoặc "30000/1001"
    fps_str = stream.get("r_frame_rate", "25/1")
    if "/" in fps_str:
        num, den = fps_str.split("/")
        fps = float(num) / float(den) if float(den) != 0 else 25.0
    else:
        fps = float(fps_str)

    # Duration
    duration_str = stream.get("duration") or fmt.get("duration") or "0"
    duration_sec = float(duration_str)

    total_frames = int(stream.get("nb_frames") or int(duration_sec * fps))

    return VideoMetadata(
        path=video_path,
        duration_sec=duration_sec,
        width=width,
        height=height,
        fps=fps,
        total_frames=total_frames,
    )


class VideoClipExtractor:
    """
    Trích xuất spatio-temporal video clips cho mô hình V-JEPA 2.
    Mỗi clip gồm T frames (mặc định 16) tại độ phân giải HxW (mặc định 224x224).
    """

    def __init__(
        self,
        clip_frames: int = 16,
        target_size: int = 224,
        target_fps: float = 16.0,
        stride_sec: float = 1.0,
    ) -> None:
        self.clip_frames = clip_frames
        self.target_size = target_size
        self.target_fps = target_fps
        self.stride_sec = stride_sec
        self.ffmpeg_bin = get_ffmpeg_path()

    def extract_clips(
        self,
        video_path: Path,
    ) -> Generator[Tuple[float, np.ndarray], None, None]:
        """
        Sinh các video clips theo dạng generator.
        Yield: (timestamp_start_sec, clip_array[16, 224, 224, 3] uint8)
        """
        meta = get_video_metadata(video_path)
        if meta.duration_sec <= 0:
            return

        cmd = [
            self.ffmpeg_bin,
            "-v", "error",
            "-i", str(video_path),
            "-vf", f"fps={self.target_fps},scale={self.target_size}:{self.target_size}",
            "-f", "rawvideo",
            "-pix_fmt", "rgb24",
            "-",
        ]

        frame_bytes = self.target_size * self.target_size * 3
        clip_buffer: List[np.ndarray] = []
        current_time_sec = 0.0
        frame_interval_sec = 1.0 / self.target_fps

        with subprocess.Popen(
            cmd,
            stdout=subprocess.PIPE,
            stderr=subprocess.PIPE,
            bufsize=10**7,
        ) as proc:
            assert proc.stdout is not None
            frame_idx = 0
            while True:
                raw_frame = proc.stdout.read(frame_bytes)
                if len(raw_frame) < frame_bytes:
                    break

                frame = np.frombuffer(raw_frame, dtype=np.uint8).reshape(
                    (self.target_size, self.target_size, 3)
                )
                clip_buffer.append(frame)

                if len(clip_buffer) == self.clip_frames:
                    clip_np = np.stack(clip_buffer, axis=0)  # [16, 224, 224, 3]
                    timestamp_start = max(0.0, current_time_sec - (self.clip_frames * frame_interval_sec))
                    yield (timestamp_start, clip_np)

                    # Áp dụng stride: số frame cần nhảy tương ứng với stride_sec
                    skip_frames = max(1, int(self.stride_sec * self.target_fps))
                    if skip_frames >= self.clip_frames:
                        clip_buffer = []
                    else:
                        clip_buffer = clip_buffer[skip_frames:]

                current_time_sec += frame_interval_sec
                frame_idx += 1
