#!/usr/bin/env python3
# -*- coding: utf-8 -*-

"""
视频关键帧提取器 (Video Keyframe Extractor)

从视频中提取关键帧，用于反向搜索
"""

import cv2
import numpy as np
from pathlib import Path
from typing import List
import hashlib


class VideoKeyframeExtractor:
    """视频关键帧提取器"""

    def __init__(self, output_dir: str = "./keyframes", max_keyframes: int = 5):
        """
        初始化关键帧提取器

        Args:
            output_dir: 输出目录
            max_keyframes: 最大提取关键帧数
        """
        self.output_dir = Path(output_dir)
        self.output_dir.mkdir(parents=True, exist_ok=True)
        self.max_keyframes = max_keyframes

        # 关键帧提取策略
        self.extraction_strategy = "uniform"  # uniform, scene_change, first_n

    async def extract_keyframes(self, video_path: str) -> List[Path]:
        """
        提取视频关键帧

        Args:
            video_path: 视频文件路径（URL或本地路径）

        Returns:
            List[Path]: 关键帧文件路径列表
        """
        # 如果是URL，先下载（这里简化实现，假设是本地路径）
        local_video_path = Path(video_path)

        if not local_video_path.exists():
            print(f"视频文件不存在: {video_path}")
            return []

        # 打开视频
        cap = cv2.VideoCapture(str(local_video_path))

        if not cap.isOpened():
            print(f"无法打开视频: {video_path}")
            return []

        # 获取视频信息
        total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
        fps = cap.get(cv2.CAP_PROP_FPS)
        duration = total_frames / fps

        print(f"视频信息: {total_frames}帧, {fps:.2f}fps, {duration:.2f}秒")

        # 提取关键帧
        keyframes = []

        if self.extraction_strategy == "uniform":
            keyframes = self._extract_uniform_keyframes(cap, total_frames)
        elif self.extraction_strategy == "scene_change":
            keyframes = self._extract_scene_change_keyframes(cap, total_frames)
        elif self.extraction_strategy == "first_n":
            keyframes = self._extract_first_n_keyframes(cap, total_frames)

        cap.release()

        print(f"提取了 {len(keyframes)} 个关键帧")
        return keyframes

    def _extract_uniform_keyframes(
        self,
        cap: cv2.VideoCapture,
        total_frames: int
    ) -> List[Path]:
        """
        均匀提取关键帧

        Args:
            cap: 视频捕获对象
            total_frames: 总帧数

        Returns:
            List[Path]: 关键帧路径列表
        """
        keyframes = []
        frame_indices = []

        # 计算帧间隔
        if total_frames <= self.max_keyframes:
            # 帧数少，提取所有帧
            frame_indices = list(range(total_frames))
        else:
            # 均匀采样
            interval = total_frames // self.max_keyframes
            frame_indices = [i * interval for i in range(self.max_keyframes)]

        # 提取帧
        for idx, frame_idx in enumerate(frame_indices):
            cap.set(cv2.CAP_PROP_POS_FRAMES, frame_idx)
            ret, frame = cap.read()

            if ret:
                # 保存关键帧
                keyframe_path = self._save_keyframe(frame, idx)
                keyframes.append(keyframe_path)

        return keyframes

    def _extract_scene_change_keyframes(
        self,
        cap: cv2.VideoCapture,
        total_frames: int
    ) -> List[Path]:
        """
        基于场景变化提取关键帧

        Args:
            cap: 视频捕获对象
            total_frames: 总帧数

        Returns:
            List[Path]: 关键帧路径列表
        """
        keyframes = []

        # 读取第一帧
        prev_frame = None
        frame_idx = 0
        keyframe_idx = 0

        # 场景变化阈值
        change_threshold = 0.3

        # 跳过前几帧（可能包含片头）
        start_frame = min(10, total_frames // 10)
        cap.set(cv2.CAP_PROP_POS_FRAMES, start_frame)
        frame_idx = start_frame

        while cap.isOpened() and len(keyframes) < self.max_keyframes:
            ret, frame = cap.read()

            if not ret:
                break

            if prev_frame is not None:
                # 计算帧差异
                diff = self._calculate_frame_difference(prev_frame, frame)

                # 如果差异超过阈值，保存关键帧
                if diff > change_threshold:
                    keyframe_path = self._save_keyframe(frame, keyframe_idx)
                    keyframes.append(keyframe_path)
                    keyframe_idx += 1

            prev_frame = frame
            frame_idx += 1

        # 如果没有检测到场景变化，提取第一帧
        if not keyframes:
            cap.set(cv2.CAP_PROP_POS_FRAMES, 0)
            ret, frame = cap.read()
            if ret:
                keyframe_path = self._save_keyframe(frame, 0)
                keyframes.append(keyframe_path)

        return keyframes

    def _extract_first_n_keyframes(
        self,
        cap: cv2.VideoCapture,
        total_frames: int
    ) -> List[Path]:
        """
        提取前N个关键帧

        Args:
            cap: 视频捕获对象
            total_frames: 总帧数

        Returns:
            List[Path]: 关键帧路径列表
        """
        keyframes = []
        num_frames = min(self.max_keyframes, total_frames)

        for idx in range(num_frames):
            ret, frame = cap.read()

            if ret:
                keyframe_path = self._save_keyframe(frame, idx)
                keyframes.append(keyframe_path)
            else:
                break

        return keyframes

    def _save_keyframe(self, frame: np.ndarray, index: int) -> Path:
        """
        保存关键帧

        Args:
            frame: 帧数据
            index: 帧索引

        Returns:
            Path: 保存的文件路径
        """
        # 生成唯一文件名
        frame_hash = hashlib.md5(frame.tobytes()).hexdigest()[:8]
        filename = f"keyframe_{index}_{frame_hash}.jpg"
        filepath = self.output_dir / filename

        # 保存帧
        cv2.imwrite(str(filepath), frame, [cv2.IMWRITE_JPEG_QUALITY, 90])

        return filepath

    def _calculate_frame_difference(
        self,
        frame1: np.ndarray,
        frame2: np.ndarray
    ) -> float:
        """
        计算两帧之间的差异

        Args:
            frame1: 帧1
            frame2: 帧2

        Returns:
            float: 差异值（0-1）
        """
        # 转换为灰度
        gray1 = cv2.cvtColor(frame1, cv2.COLOR_BGR2GRAY)
        gray2 = cv2.cvtColor(frame2, cv2.COLOR_BGR2GRAY)

        # 计算绝对差异
        diff = cv2.absdiff(gray1, gray2)

        # 计算差异均值
        mean_diff = np.mean(diff)

        # 归一化到0-1
        return mean_diff / 255.0

    def set_strategy(self, strategy: str):
        """
        设置关键帧提取策略

        Args:
            strategy: 策略类型（uniform, scene_change, first_n）
        """
        if strategy in ["uniform", "scene_change", "first_n"]:
            self.extraction_strategy = strategy
        else:
            raise ValueError(f"未知的策略: {strategy}")

    def clear_output_dir(self):
        """清空输出目录"""
        for file in self.output_dir.glob("*.jpg"):
            file.unlink()
