"""
Video Processor Module
视频关键模块提取模块

功能：提取视频关键帧、缩略图、帧相似度比对
"""

import os
import subprocess
import tempfile
from typing import List, Tuple, Optional
from pathlib import Path

import cv2
import numpy as np
from PIL import Image


class VideoProcessor:
    """视频处理器"""

    def __init__(self, temp_dir: str = None):
        """
        初始化视频处理器

        Args:
            temp_dir: 临时文件目录，默认为系统临时目录
        """
        self.temp_dir = temp_dir or tempfile.gettempdir()
        os.makedirs(self.temp_dir, exist_ok=True)

    def _extract_frame_at_timestamp(
        self,
        video_path: str,
        timestamp: float,
        output_path: str
    ) -> bool:
        """
        提取指定时间戳的帧

        Args:
            video_path: 视频路径
            timestamp: 时间戳（秒）
            output_path: 输出图片路径

        Returns:
            是否成功
        """
        try:
            # 使用ffmpeg提取帧
            cmd = [
                "ffmpeg",
                "-i", video_path,
                "-ss", str(timestamp),
                "-vframes", "1",
                "-q:v", "2",
                "-y",  # 覆盖输出文件
                output_path
            ]

            result = subprocess.run(
                cmd,
                capture_output=True,
                text=True,
                timeout=30
            )

            return result.returncode == 0 and os.path.exists(output_path)

        except subprocess.TimeoutExpired:
            print(f"Timeout extracting frame at {timestamp}s")
            return False
        except Exception as e:
            print(f"Error extracting frame: {e}")
            return False

    def generate_thumbnail(
        self,
        video_path: str,
        timestamp: float = 0.0,
        width: int = 320,
        height: int = None
    ) -> Optional[str]:
        """
        生成视频缩略图

        Args:
            video_path: 视频路径
            timestamp: 时间戳（秒），默认为0
            width: 宽度，默认为320
            height: 高度，默认为None（保持宽高比）

        Returns:
            缩略图路径，失败返回None
        """
        try:
            # 生成临时文件名
            video_hash = self._get_file_hash(video_path)
            thumb_path = os.path.join(
                self.temp_dir,
                f"thumb_{video{hash}}_{int(timestamp)}.jpg"
            )

            # 提取帧
            if not self._extract_frame_at_timestamp(video_path, timestamp, thumb_path):
                return None

            # 调整大小
            img = Image.open(thumb_path)

            if height is None:
                # 保持宽高比
                aspect_ratio = img.height / img.width
                height = int(width * aspect_ratio)

            img_resized = img.resize((width, height), Image.Resampling.LANCZOS)
            img_resized.save(thumb_path, quality=85)

            return thumb_path

        except Exception as e:
            print(f"Error generating thumbnail: {e}")
            return None

    def extract_keyframes(
        self,
        video_path: str,
        method: str = "scene_detect",
        max_frames: int = 5,
        min_interval: float = 1.0
    ) -> List[str]:
        """
        提取关键帧

        Args:
            video_path: 视频路径
            method: 提取方法
                - "scene_detect": 场景变化检测（推荐）
                - "uniform": 均匀间隔
                - "faces": 检测人脸
                - "first_n": 前N帧
            max_frames: 最大帧数
            min_interval: 最小间隔（秒），避免提取太密集的帧

        Returns:
            关键帧路径列表
        """
        if method == "scene_detect":
            return self._extract_keyframes_scene_detect(
                video_path, max_frames, min_interval
            )
        elif method == "uniform":
            return self._extract_keyframes_uniform(
                video_path, max_frames
            )
        elif method == "faces":
            return self._extract_keyframes_faces(
                video_path, max_frames, min_interval
            )
        elif method == "first_n":
            return self._extract_keyframes_first_n(
                video_path, max_frames
            )
        else:
            print(f"Unknown method: {method}, using scene_detect")
            return self._extract_keyframes_scene_detect(
                video_path, max_frames, min_interval
            )

    def _extract_keyframes_scene_detect(
        self,
        video_path: str,
        max_frames: int,
        min_interval: float
    ) -> List[str]:
        """
        使用场景变化检测提取关键帧
        """
        keyframes = []

        try:
            cap = cap2.VideoCapture(video_path)
            if not cap.isOpened():
                print(f"Cannot open video: {video_path}")
                return []

            fps = cap.get(cv2.CAP_PROP_FPS)
            total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
            duration = total_frames / fps

            # 计算采样间隔
            if total_frames <= 0:
                cap.release()
                return []

            interval = max(int(fps * min_interval), 1)
            step = max(int(total_frames / max_frames), interval)

            prev_frame = None
            prev_gray = None

            frame_count = 0
            last_extract_time = -min_interval - 1

            video_hash = self._get_file_hash(video_path)

            while True:
                ret, frame = cap.read()
                if not ret:
                    break

                current_time = frame_count / fps

                # 检查是否达到最小间隔
                if current_time - last_extract_time < min_interval:
                    frame_count += 1
                    continue

                # 计算帧差异
                gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)

                if prev_gray is None:
                    # 第一帧总是关键帧
                    keyframe_path = self._save_keyframe(
                        frame, video_hash, frame_count, current_time
                    )
                    if keyframe_path:
                        keyframes.append((keyframe_path, current_time))
                        last_extract_time = current_time
                else:
                    # 计算差异
                    diff = cv2.absdiff(gray, prev_gray)
                    diff_score = np.mean(diff)

                    # 如果差异足够大，则认为是关键帧
                    if diff_score > 30:  # 差异阈值，可根据需要调整
                        keyframe_path = self._save_keyframe(
                            frame, video_hash, frame_count, current_time
                        )
                        if keyframe_path:
                            keyframes.append((keyframe_path, current_time))
                            last_extract_time = current_time

                prev_gray = gray.copy()
                frame_count += 1

                # 检查是否达到最大帧数
                if len(keyframes) >= max_frames:
                    break

                # 跳过一些帧以加速处理
                for _ in range(step - 1):
                    ret, _ = cap.read()
                    if not ret:
                        break
                    frame_count += 1

            cap.release()

            # 按时间排序并返回路径
            keyframes.sort(key=lambda x: x[1])
            return [kf[0] for kf in keyframes[:max_frames]]

        except Exception as e:
            print(f"Error in scene detection: {e}")
            return []

    def _extract_keyframes_uniform(
        self,
        video_path: str,
        max_frames: int
    ) -> List[str]:
        """
        均匀间隔提取关键帧
        """
        keyframes = []

        try:
            cap = cv2.VideoCapture(video_path)
            if not cap.isOpened():
                return []

            fps = cap.get(cv2.CAP_PROP_FPS)
            total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))

            if total_frames <= 0:
                cap.release()
                return []

            duration = total_frames / fps
            interval = duration / max_frames

            video_hash = self._get_file_hash(video_path)

            for i in range(max_frames):
                timestamp = i * interval
                frame_num = int(timestamp * fps)

                cap.set(cv2.CAP_PROP_POS_FRAMES, frame_num)
                ret, frame = cap.read()

                if ret:
                    keyframe_path = self._save_keyframe(
                        frame, video_hash, frame_num, timestamp
                    )
                    if keyframe_path:
                        keyframes.append(keyframe_path)

            cap.release()
            return keyframes

        except Exception as e:
            print(f"Error in uniform extraction: {e}")
            return []

    def _extract_keyframes_faces(
        self,
        video_path: str,
        max_frames: int,
        min_interval: float
    ) -> List[str]:
        """
        检测人脸提取关键帧
        """
        keyframes = []

        try:
            # 加载人脸检测器（使用Haar级联分类器）
            face_cascade = cv2.CascadeClassifier(
                cv2.data.haarcascades + "haarcascade_frontalface_default.xml"
            )

            if face_cascade.empty():
                print("Failed to load face cascade")
                return self._extract_keyframes_uniform(video_path, max_frames)

            cap = cv2.VideoCapture(video_path)
            if not cap.isOpened():
                return []

            fps = cap.get(cv2.CAP_PROP_FPS)
            total_frames = int(cv2.CAP_PROP_FRAME_COUNT)

            if total_frames <= 0:
                cap.release()
                return []

            step = max(int(total_frames / max_frames), int(fps * min_interval))

            frame_count = 0
            last_extract_time = -min_interval - 1

            video_hash = self._get_file_hash(video_path)

            while True:
                ret, frame = cap.read()
                if not ret:
                    break

                current_time = frame_count / fps

                # 检查间隔
                if current_time - last_extract_time < min_interval:
                    frame_count += 1
                    continue

                # 检测人脸
                gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
                faces = face_cascade.detectMultiScale(
                    gray,
                    scaleFactor=1.1,
                    minNeighbors=5,
                    minSize=(30, 30)
                )

                # 如果检测到人脸，保存关键帧
                if len(faces) > 0:
                    keyframe_path = self._save_keyframe(
                        frame, video_hash, frame_count, current_time
                    )
                    if keyframe_path:
                        keyframes.append((keyframe_path, current_time))
                        last_extract_time = current_time

                        if len(keyframes) >= max_frames:
                            break

                frame_count += 1

                # 跳帧
                for _ in range(step - 1):
                    ret, _ = cap.read()
                    if not ret:
                        break
                    frame_count += 1

            cap.release()

            keyframes.sort(key=lambda x: x[1])
            return [kf[0] for kf in keyframes[:max_frames]]

        except Exception as e:
            print(f"Error in face detection: {e}")
            return []

    def _extract_keyframes_first_n(
        self,
        video_path: str,
        max_frames: int
    ) -> List[str]:
        """
        提取前N帧
        """
        keyframes = []

        try:
            cap = cv2.VideoCapture(video_path)
            if not cap.isOpened():
                return []

            fps = cap.get(cv2.CAP_PROP_FPS)

            video_hash = self._get_file_hash(video_path)

            for i in range(max_frames):
                ret, frame = cap.read()
                if not ret:
                    break

                timestamp = i / fps
                keyframe_path = self._save_keyframe(
                    frame, video_hash, i, timestamp
                )
                if keyframe_path:
                    keyframes.append(keyframe_path)

            cap.release()
            return keyframes

        except Exception as e:
            print(f"Error extracting first N frames: {e}")
            return []

    def _save_keyframe(
        self,
        frame: np.ndarray,
        video_hash: str,
        frame_num: int,
        timestamp: float
    ) -> Optional[str]:
        """
        保存关键帧到临时文件
        """
        try:
            keyframe_path = os.path.join(
                self.temp_dir,
                f"keyframe_{video_hash}_{frame_num}_{int(timestamp * 1000)}.jpg"
            )

            cv2.imwrite(keyframe_path, frame, [cv2.IMWRITE_JPEG_QUALITY, 85])
            return keyframe_path

        except Exception as e:
            print(f"Error saving keyframe: {e}")
            return None

    def compare_frames(
        self,
        frame1_path: str,
        frame2_path: str,
        method: str = "ssim"
    ) -> float:
        """
        比较两帧的相似度

        Args:
            frame1_path: 第一帧路径
            frame2_path: 第二帧路径
            method: 比较方法
                - "ssim": 结构相似性（推荐）
                - "mse": 均方误差
                - "hash": 感知哈希

        Returns:
            相似度分数（0.0-1.0），1.0表示完全相同
        """
        try:
            frame1 = cv2.imread(frame1_path)
            frame2 = cv2.imread(frame2_path)

            if frame1 is None or frame2 is None:
                return 0.0

            # 调整大小到相同尺寸
            h1, w1 = frame1.shape[:2]
            h2, w2 = frame2.shape[:2]

            if h1 != h2 or w1 != w2:
                frame2 = cv2.resize(frame2, (w1, h1))

            if method == "ssim":
                return self._calculate_ssim(frame1, frame2)
            elif method == "mse":
                return self._calculate_mse(frame1, frame2)
            elif method == "hash":
                return self._calculate_hash_similarity(frame1_path, frame2_path)
            else:
                print(f"Unknown method: {method}, using ssim")
                return self._calculate_ssim(frame1, frame2)

        except Exception as e:
            print(f"Error comparing frames: {e}")
            return 0.0

    def _calculate_ssim(
        self,
        frame1: np.ndarray,
        frame2: np.ndarray
    ) -> float:
        """
        计算结构相似性（SSIM）
        """
        try:
            # 转换为灰度图
            gray1 = cv2.cvtColor(frame1, cv2.COLOR_BGR2GRAY)
            gray2 = cv2.cvtColor(frame2, cv2.COLOR_BGR2GRAY)

            # 计算均值
            mu1 = cv2.GaussianBlur(gray1, (11, 11), 1.5)
            mu2 = cv2.GaussianBlur(gray2, (11, 11), 1.5)

            mu1_sq = mu1 ** 2
            mu2_sq = mu2 ** 2
            mu1_mu2 = mu1 * mu2

            # 计算方差
            sigma1_sq = cv2.GaussianBlur(gray1 ** 2, (11, 11), 1.5) - mu1_sq
            sigma2_sq = cv2.GaussianBlur(gray2 ** 2, (11, 11), 1.5) - mu2_sq
            sigma12 = cv2.GaussianBlur(gray1 * gray2, (11, 11), 1.5) - mu1_mu2

            # SSIM常数
            C1 = (0.01 * 255) ** 2
            C2 = (0.03 * 255) ** 2

            # SSIM公式
            ssim_map = ((2 * mu1_mu2 + C1) * (2 * sigma12 + C2)) / \
                       ((mu1_sq + mu2_sq + C1) * (sigma1_sq + sigma2_sq + C2))

            return float(np.mean(ssim_map))

        except Exception as e:
            print(f"Error calculating SSIM: {e}")
            return 0.0

    def _calculate_mse(
        self,
        frame1: np.ndarray,
        frame2: np.ndarray
    ) -> float:
        """
        计算均方误差（MSE），转换为相似度
        """
        try:
            diff = frame1.astype(float) - frame2.astype(float)
            mse = np.mean(diff ** 2)

            # 转换为相似度（MSE越小，相似度越高）
            max_mse = (255 ** 2) * 3  # RGB三通道
            similarity = 1.0 - min(mse / max_mse, 1.0)

            return similarity

        except Exception as ever:
            print(f"Error calculating MSE: {e}")
            return 0.0

    def _calculate_hash_similarity(
        self,
        frame1_path: str,
        frame2_path: str
    ) -> float:
        """
        使用感知哈希计算相似度
        """
        try:
            from imagehash import phash, average_hash, dhash

            img1 = Image.open(frame1_path)
            img2 = Image.open(frame2_path)

            # 计算多种哈希
            hash1_p = phash(img1)
            hash2_p = phash(img2)

            hash1_a = average_hash(img1)
            hash2_a = average_hash(img2)

            hash1_d = dhash(img1)
            hash2_d = dhash(img2)

            # 计算哈希距离（越小越相似）
            dist_p = hash1_p - hash2_p
            dist_a = hash1_a - hash2_a
            dist_d = hash1_d - hash2_d

            # 归一化为相似度
            max_dist = 64  # 哈希的最大距离
            similarity = 1.0 - min((dist_p + dist_a + dist_d) / (3 * max_dist), 1.0)

            return similarity

        except Exception as e:
            print(f"Error calculating hash similarity: {e}")
            return 0.0

    def _get_file_hash(self, file_path: str) -> str:
        """
        获取文件哈希值（用于生成唯一文件名）
        """
        try:
            import hashlib

            md5 = hashlib.md5()
            with open(file_path, "rb") as f:
                for chunk in iter(lambda: f.read(4096), b""):
                    md5.update(chunk)

            return md5.hexdigest()[:12]

        except Exception as e:
            print(f"Error calculating file hash: {e}")
            return str(hash(file_path))[:12]

    def cleanup(self, max_age_hours: int = 24):
        """
        清理临时文件

        Args:
            max_age_hours: 最大文件年龄（小时），超过此年龄的文件将被删除
        """
        try:
            import time

            current_time = time.time()
            max_age_seconds = max_age_hours * 3600

            for filename in os.listdir(self.temp_dir):
                if filename.startswith(("thumb_", "keyframe_")):
                    filepath = os.path.join(self.temp_dir, filename)
                    file_age = current_time - os.path.getmtime(filepath)

                    if file_age > max_age_seconds:
                        os.remove(filepath)
                        print(f"Removed old temp file: {filename}")

        except Exception as e:
            print(f"Error cleaning up temp files: {e}")
