"""
Feature Store and Cache Manager for VideoDetection.
Tự động đồng bộ, trích xuất và lưu trữ đặc trưng (V-JEPA 2 + NeuralFP) theo từng video
vào thư mục con .features/ của Data/ và Target/.
Đảm bảo khi mở phần mềm thì toàn bộ số liệu đã sẵn sàng, so sánh tức thì (<0.1s).
"""
from __future__ import annotations

import json
import logging
import os
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any, Callable, Dict, List, Optional, Tuple

import numpy as np

from ..audio.neuralfp_extractor import NeuralFPExtractor
from ..core.paths import DATA_DIR, TARGET_DIR, clean_temp_dir
from ..ingestion.audio_loader import AudioChunkExtractor
from ..ingestion.video_loader import VideoClipExtractor, get_video_metadata
from ..visual.vjepa_extractor import VJEPAExtractor
from .quantizer import ProductQuantizer64, binarize_audio_fingerprint

logger = logging.getLogger(__name__)

SUPPORTED_EXTENSIONS = {".mp4", ".mkv", ".avi", ".mov", ".flv", ".webm", ".ts", ".m4v", ".wmv"}


@dataclass
class CachedVideoFeatures:
    video_id: str
    filename: str
    video_path: str
    duration_sec: float
    file_size_bytes: int
    file_mtime: float
    visual_times: np.ndarray  # [N_vis] float32
    visual_embs: np.ndarray   # [N_vis, 1024] float32
    audio_times: np.ndarray   # [N_aud] float32
    audio_fps: np.ndarray     # [N_aud, 128] float32
    visual_codes: Optional[np.ndarray] = None  # [N_vis, 64] uint8 (Product Quantization PQ64)
    audio_bits: Optional[np.ndarray] = None    # [N_aud, 16] uint8 (128-bit Hamming Bit-Vector)

    @property
    def visual_count(self) -> int:
        return len(self.visual_times)

    @property
    def audio_count(self) -> int:
        return len(self.audio_times)

    def to_metadata_dict(self) -> Dict[str, Any]:
        return {
            "video_id": self.video_id,
            "filename": self.filename,
            "video_path": self.video_path,
            "duration_sec": round(self.duration_sec, 2),
            "file_size_bytes": self.file_size_bytes,
            "file_mtime": self.file_mtime,
            "visual_count": self.visual_count,
            "audio_count": self.audio_count,
        }


class FeatureStore:
    """
    Quản lý kho cache đặc trưng độc lập cho cả Data/ và Target/.
    """

    def __init__(
        self,
        clip_extractor: VideoClipExtractor,
        chunk_extractor: AudioChunkExtractor,
        visual_extractor: VJEPAExtractor,
        audio_extractor: NeuralFPExtractor,
    ) -> None:
        self.clip_extractor = clip_extractor
        self.chunk_extractor = chunk_extractor
        self.visual_extractor = visual_extractor
        self.audio_extractor = audio_extractor
        # Bộ lượng tử hóa vector Product Quantization PQ64
        self.quantizer = ProductQuantizer64(d=1024, m=64, nbits=8)

    @staticmethod
    def get_features_dir(base_dir: Path) -> Path:
        fdir = base_dir / ".features"
        fdir.mkdir(parents=True, exist_ok=True)
        return fdir

    @staticmethod
    def get_cache_path(video_path: Path) -> Path:
        fdir = FeatureStore.get_features_dir(video_path.parent)
        # Tên file cache: ten_video.npz
        safe_name = f"{video_path.stem}_{video_path.suffix.lstrip('.')}.npz"
        return fdir / safe_name

    def is_cached_and_valid(self, video_path: Path) -> bool:
        cache_p = self.get_cache_path(video_path)
        if not cache_p.exists():
            return False

        try:
            stat = video_path.stat()
            # Đọc nhanh header metadata
            with np.load(cache_p, allow_pickle=True) as data:
                cached_size = int(data.get("file_size_bytes", -1))
                cached_mtime = float(data.get("file_mtime", -1.0))
                # So sánh kích thước và thời gian sửa đổi file
                if cached_size == stat.st_size and abs(cached_mtime - stat.st_mtime) < 1.0:
                    return True
        except Exception:
            pass
        return False

    def load_features(self, video_path: Path) -> Optional[CachedVideoFeatures]:
        cache_p = self.get_cache_path(video_path)
        if not cache_p.exists():
            return None

        try:
            with np.load(cache_p, allow_pickle=True) as data:
                v_embs = data["visual_embs"]
                a_fps = data["audio_fps"]

                # Lấy mã lượng tử PQ và Hamming bits, nếu bản cache cũ chưa có thì tự động sinh ngay
                v_codes = data.get("visual_codes", None)
                if v_codes is None and len(v_embs) > 0:
                    v_codes = self.quantizer.encode(v_embs)

                a_bits = data.get("audio_bits", None)
                if a_bits is None and len(a_fps) > 0:
                    a_bits = binarize_audio_fingerprint(a_fps)

                return CachedVideoFeatures(
                    video_id=str(data["video_id"]),
                    filename=str(data["filename"]),
                    video_path=str(data["video_path"]),
                    duration_sec=float(data["duration_sec"]),
                    file_size_bytes=int(data["file_size_bytes"]),
                    file_mtime=float(data["file_mtime"]),
                    visual_times=data["visual_times"],
                    visual_embs=v_embs,
                    audio_times=data["audio_times"],
                    audio_fps=a_fps,
                    visual_codes=v_codes,
                    audio_bits=a_bits,
                )
        except Exception as e:
            logger.error(f"Lỗi nạp cache từ {cache_p}: {e}")
            return None

    def process_and_cache_video(
        self,
        video_path: Path,
        progress_callback: Optional[Callable[[int, str], None]] = None,
    ) -> CachedVideoFeatures:
        """
        Trích xuất đặc trưng V-JEPA 2 và NeuralFP rồi lưu file .npz vào .features/.
        """
        stat = video_path.stat()
        meta = get_video_metadata(video_path)

        if progress_callback:
            progress_callback(10, f"Đang bóc tách video clips: {video_path.name}")

        # 1. Visual Clips
        vis_times: List[float] = []
        vis_embs_list: List[np.ndarray] = []
        for t_sec, clip in self.clip_extractor.extract_clips(video_path):
            emb = self.visual_extractor.extract(clip)
            vis_times.append(t_sec)
            vis_embs_list.append(emb)

        if progress_callback:
            progress_callback(50, f"Đang bóc tách audio fingerprints: {video_path.name}")

        # 2. Audio Chunks
        aud_times: List[float] = []
        aud_fps_list: List[np.ndarray] = []
        for t_sec, chunk in self.chunk_extractor.extract_chunks(video_path):
            fp = self.audio_extractor.extract(chunk)
            aud_times.append(t_sec)
            aud_fps_list.append(fp)

        # Chuyển đổi numpy arrays
        visual_times_np = np.array(vis_times, dtype=np.float32)
        visual_embs_np = (
            np.stack(vis_embs_list, axis=0).astype(np.float32)
            if vis_embs_list
            else np.empty((0, self.visual_extractor.embedding_dim), dtype=np.float32)
        )

        audio_times_np = np.array(aud_times, dtype=np.float32)
        audio_fps_np = (
            np.stack(aud_fps_list, axis=0).astype(np.float32)
            if aud_fps_list
            else np.empty((0, self.audio_extractor.fingerprint_dim), dtype=np.float32)
        )

        # Mã hóa lượng tử hóa PQ64 cho Visual (1024 float -> 64 uint8)
        visual_codes = (
            self.quantizer.encode(visual_embs_np)
            if len(visual_embs_np) > 0
            else np.empty((0, self.quantizer.m), dtype=np.uint8)
        )

        # Nhị phân hóa Hamming 128-bit cho Audio (128 float -> 16 bytes uint8)
        audio_bits = (
            binarize_audio_fingerprint(audio_fps_np)
            if len(audio_fps_np) > 0
            else np.empty((0, 16), dtype=np.uint8)
        )

        features = CachedVideoFeatures(
            video_id=video_path.stem,
            filename=video_path.name,
            video_path=str(video_path),
            duration_sec=meta.duration_sec,
            file_size_bytes=stat.st_size,
            file_mtime=stat.st_mtime,
            visual_times=visual_times_np,
            visual_embs=visual_embs_np,
            audio_times=audio_times_np,
            audio_fps=audio_fps_np,
            visual_codes=visual_codes,
            audio_bits=audio_bits,
        )

        # Lưu cache file nén chứa cả vector gốc và mã lượng tử hóa siêu nhẹ
        cache_p = self.get_cache_path(video_path)
        np.savez_compressed(
            cache_p,
            video_id=features.video_id,
            filename=features.filename,
            video_path=features.video_path,
            duration_sec=features.duration_sec,
            file_size_bytes=features.file_size_bytes,
            file_mtime=features.file_mtime,
            visual_times=features.visual_times,
            visual_embs=features.visual_embs,
            audio_times=features.audio_times,
            audio_fps=features.audio_fps,
            visual_codes=features.visual_codes,
            audio_bits=features.audio_bits,
        )

        clean_temp_dir()

        if progress_callback:
            progress_callback(100, f"Đã lưu cache đặc trưng: {video_path.name}")

        return features

    def sync_directory(
        self,
        directory: Path,
        progress_callback: Optional[Callable[[int, str], None]] = None,
    ) -> Tuple[int, int, List[CachedVideoFeatures]]:
        """
        Quét thư mục (Data/ hoặc Target/).
        Nếu có video mới chưa trích xuất: tự động trích xuất và cache.
        Nếu đã cache: nạp nhanh từ đĩa.
        Trả về: (số lượng mới xử lý, số lượng dùng cache cũ, danh sách tất cả features)
        """
        video_files = [
            f for f in directory.iterdir()
            if f.is_file() and f.suffix.lower() in SUPPORTED_EXTENSIONS
        ]

        newly_processed = 0
        from_cache = 0
        all_features: List[CachedVideoFeatures] = []
        total = len(video_files)

        if total == 0:
            if progress_callback:
                progress_callback(100, f"Thư mục {directory.name}/ hiện không có video.")
            return 0, 0, []

        for idx, v_path in enumerate(video_files):
            def file_cb(pct: int, msg: str):
                if progress_callback:
                    overall = int(((idx * 100) + pct) / total)
                    progress_callback(overall, f"[{idx+1}/{total}] {msg}")

            if self.is_cached_and_valid(v_path):
                feats = self.load_features(v_path)
                if feats:
                    all_features.append(feats)
                    from_cache += 1
                    if progress_callback:
                        file_cb(100, f"Đã nạp sẵn từ cache: {v_path.name}")
                    continue

            # Chưa có cache -> xử lý mới
            feats = self.process_and_cache_video(v_path, file_cb)
            all_features.append(feats)
            newly_processed += 1

        if progress_callback:
            progress_callback(100, f"Đồng bộ xong {directory.name}/: {len(all_features)} video.")

        return newly_processed, from_cache, all_features
