"""
Temporal Alignment and Multimodal Matching Module.
Căn chỉnh trục thời gian bằng Diagonal Hough Transform / Offset Clustering để phát hiện
chính xác các đoạn video bị sao chép [target_start, target_end] vs [ref_start, ref_end].
"""
from __future__ import annotations

from collections import defaultdict
from dataclasses import dataclass
from typing import Dict, List, Optional, Tuple

import numpy as np

from ..indexing.database import OwnerDatabase


@dataclass
class MatchSegment:
    target_start_sec: float
    target_end_sec: float
    reference_video_id: str
    ref_start_sec: float
    ref_end_sec: float
    confidence_score: float
    visual_score: float
    audio_score: float
    match_type: str  # "multimodal", "visual_only", "audio_only"

    @property
    def duration_sec(self) -> float:
        return max(0.0, self.target_end_sec - self.target_start_sec)

    def to_dict(self) -> Dict[str, object]:
        return {
            "target_range": f"{self._format_time(self.target_start_sec)} - {self._format_time(self.target_end_sec)}",
            "ref_range": f"{self._format_time(self.ref_start_sec)} - {self._format_time(self.ref_end_sec)}",
            "reference_video_id": self.reference_video_id,
            "target_start_sec": round(self.target_start_sec, 2),
            "target_end_sec": round(self.target_end_sec, 2),
            "ref_start_sec": round(self.ref_start_sec, 2),
            "ref_end_sec": round(self.ref_end_sec, 2),
            "duration_sec": round(self.duration_sec, 2),
            "confidence_score": round(self.confidence_score, 4),
            "visual_score": round(self.visual_score, 4),
            "audio_score": round(self.audio_score, 4),
            "match_type": self.match_type,
        }

    @staticmethod
    def _format_time(seconds: float) -> str:
        m, s = divmod(int(seconds), 60)
        h, m = divmod(m, 60)
        if h > 0:
            return f"{h:02d}:{m:02d}:{s:02d}"
        return f"{m:02d}:{s:02d}"


@dataclass
class RawPointMatch:
    t_target: float
    t_ref: float
    ref_video_id: str
    score: float
    modality: str  # "visual" hoặc "audio"


class TemporalMatcher:
    def __init__(
        self,
        db: OwnerDatabase,
        visual_threshold: float = 0.78,
        audio_threshold: float = 0.82,
        min_match_duration_sec: float = 3.0,
        temporal_bin_size_sec: float = 2.0,
        fusion_weight_visual: float = 0.65,
        fusion_weight_audio: float = 0.35,
    ) -> None:
        self.db = db
        self.visual_threshold = visual_threshold
        self.audio_threshold = audio_threshold
        self.min_match_duration_sec = min_match_duration_sec
        self.temporal_bin_size_sec = temporal_bin_size_sec
        self.w_visual = fusion_weight_visual
        self.w_audio = fusion_weight_audio

    def match(
        self,
        query_visual_features: List[Tuple[float, np.ndarray]],  # [(t_sec, emb_1024)]
        query_audio_features: List[Tuple[float, np.ndarray]],   # [(t_sec, fp_128)]
    ) -> List[MatchSegment]:
        """
        Đối soát toàn diện giữa các đặc trưng của target video và cơ sở dữ liệu chủ sở hữu.
        """
        raw_matches: List[RawPointMatch] = []

        # 1. Tìm kiếm thị giác (Visual search)
        if query_visual_features and self.db.visual_index.total > 0:
            query_times = [f[0] for f in query_visual_features]
            query_vecs = np.stack([f[1] for f in query_visual_features], axis=0)

            distances, indices = self.db.visual_index.search(query_vecs, top_k=3)

            for q_idx in range(len(query_times)):
                t_t = query_times[q_idx]
                for k_idx in range(distances.shape[1]):
                    sim = float(distances[q_idx, k_idx])
                    idx = int(indices[q_idx, k_idx])
                    if idx < 0 or idx >= len(self.db.visual_meta):
                        continue
                    if sim >= self.visual_threshold:
                        meta = self.db.visual_meta[idx]
                        raw_matches.append(RawPointMatch(
                            t_target=t_t,
                            t_ref=float(meta["timestamp"]),
                            ref_video_id=meta["video_id"],
                            score=sim,
                            modality="visual",
                        ))

        # 2. Tìm kiếm âm thanh (Audio search)
        if query_audio_features and self.db.audio_index.total > 0:
            query_times_a = [f[0] for f in query_audio_features]
            query_vecs_a = np.stack([f[1] for f in query_audio_features], axis=0)

            distances_a, indices_a = self.db.audio_index.search(query_vecs_a, top_k=3)

            for q_idx in range(len(query_times_a)):
                t_t = query_times_a[q_idx]
                for k_idx in range(distances_a.shape[1]):
                    sim = float(distances_a[q_idx, k_idx])
                    idx = int(indices_a[q_idx, k_idx])
                    if idx < 0 or idx >= len(self.db.audio_meta):
                        continue
                    if sim >= self.audio_threshold:
                        meta = self.db.audio_meta[idx]
                        raw_matches.append(RawPointMatch(
                            t_target=t_t,
                            t_ref=float(meta["timestamp"]),
                            ref_video_id=meta["video_id"],
                            score=sim,
                            modality="audio",
                        ))

        if not raw_matches:
            return []

        # 3. Gom nhóm theo từng video_id trong kho gốc
        by_video: Dict[str, List[RawPointMatch]] = defaultdict(list)
        for m in raw_matches:
            by_video[m.ref_video_id].append(m)

        detected_segments: List[MatchSegment] = []

        # 4. Căn chỉnh trục thời gian (Diagonal Hough Clustering)
        for vid, matches in by_video.items():
            segments = self._cluster_and_align(vid, matches)
            detected_segments.extend(segments)

        # Áp dụng Temporal NMS để loại bỏ các đoạn trùng lặp chồng lấn
        return self._temporal_nms(detected_segments)

    def _cluster_and_align(self, vid: str, matches: List[RawPointMatch]) -> List[MatchSegment]:
        """
        Tính offset Delta = t_ref - t_target.
        Các điểm khớp cùng đoạn sao chép sẽ có Delta xấp xỉ nhau.
        """
        # Phân rổ theo Delta (binned offset)
        offset_bins: Dict[int, List[RawPointMatch]] = defaultdict(list)
        for m in matches:
            delta = m.t_ref - m.t_target
            bin_idx = int(np.round(delta / self.temporal_bin_size_sec))
            offset_bins[bin_idx].append(m)

        segments: List[MatchSegment] = []

        for bin_idx, cluster in offset_bins.items():
            if len(cluster) < 2:
                continue

            # Sắp xếp theo t_target
            cluster.sort(key=lambda x: x.t_target)

            # Chia nhỏ nếu khoảng cách giữa 2 điểm liên tiếp quá xa (> 4 giây)
            sub_clusters: List[List[RawPointMatch]] = []
            current_sub: List[RawPointMatch] = [cluster[0]]

            for i in range(1, len(cluster)):
                if cluster[i].t_target - cluster[i - 1].t_target <= 4.0:
                    current_sub.append(cluster[i])
                else:
                    if len(current_sub) >= 2:
                        sub_clusters.append(current_sub)
                    current_sub = [cluster[i]]
            if len(current_sub) >= 2:
                sub_clusters.append(current_sub)

            for sub in sub_clusters:
                t_target_start = sub[0].t_target
                t_target_end = sub[-1].t_target + 1.0  # Cộng thêm 1 clip duration
                duration = t_target_end - t_target_start

                if duration < self.min_match_duration_sec:
                    continue

                t_ref_start = sub[0].t_ref
                t_ref_end = sub[-1].t_ref + 1.0

                # Tính điểm visual & audio riêng rẽ
                vis_scores = [pt.score for pt in sub if pt.modality == "visual"]
                aud_scores = [pt.score for pt in sub if pt.modality == "audio"]

                vis_avg = float(np.mean(vis_scores)) if vis_scores else 0.0
                aud_avg = float(np.mean(aud_scores)) if aud_scores else 0.0

                if vis_scores and aud_scores:
                    conf = (vis_avg * self.w_visual) + (aud_avg * self.w_audio)
                    mtype = "multimodal"
                elif vis_scores:
                    conf = vis_avg
                    mtype = "visual_only"
                else:
                    conf = aud_avg
                    mtype = "audio_only"

                segments.append(MatchSegment(
                    target_start_sec=t_target_start,
                    target_end_sec=t_target_end,
                    reference_video_id=vid,
                    ref_start_sec=t_ref_start,
                    ref_end_sec=t_ref_end,
                    confidence_score=conf,
                    visual_score=vis_avg,
                    audio_score=aud_avg,
                    match_type=mtype,
                ))

        return segments

    def match_cached(
        self,
        target_features: Any,
        reference_features_list: List[Any],
    ) -> List[MatchSegment]:
        """
        Đối soát tức thì (<0.1s) giữa 1 video target và danh sách video tham chiếu (1, vài video hoặc cả kho)
        bằng ma trận đặc trưng đã tính toán và lưu cache từ trước.
        """
        all_segments: List[MatchSegment] = []

        t_vis_times = target_features.visual_times
        t_vis_embs = target_features.visual_embs
        t_aud_times = target_features.audio_times
        t_aud_fps = target_features.audio_fps

        for ref in reference_features_list:
            raw_matches: List[RawPointMatch] = []

            # 1. Đối soát Visual
            if len(t_vis_embs) > 0 and len(ref.visual_embs) > 0:
                # Cosine similarity matrix: [N_target, N_ref]
                sim_vis = np.dot(t_vis_embs, ref.visual_embs.T)
                q_idxs, r_idxs = np.where(sim_vis >= self.visual_threshold)
                for q_i, r_i in zip(q_idxs, r_idxs):
                    raw_matches.append(RawPointMatch(
                        t_target=float(t_vis_times[q_i]),
                        t_ref=float(ref.visual_times[r_i]),
                        ref_video_id=ref.video_id,
                        score=float(sim_vis[q_i, r_i]),
                        modality="visual",
                    ))

            # 2. Đối soát Audio (Ưu tiên Hamming Bit-Vector 128-bit POPCOUNT siêu tốc)
            if (
                target_features.audio_bits is not None
                and ref.audio_bits is not None
                and len(target_features.audio_bits) > 0
                and len(ref.audio_bits) > 0
            ):
                from ..indexing.quantizer import compute_hamming_similarity
                sim_aud = compute_hamming_similarity(target_features.audio_bits, ref.audio_bits)
                q_idxs, r_idxs = np.where(sim_aud >= self.audio_threshold)
                for q_i, r_i in zip(q_idxs, r_idxs):
                    raw_matches.append(RawPointMatch(
                        t_target=float(t_aud_times[q_i]),
                        t_ref=float(ref.audio_times[r_i]),
                        ref_video_id=ref.video_id,
                        score=float(sim_aud[q_i, r_i]),
                        modality="audio",
                    ))
            elif len(t_aud_fps) > 0 and len(ref.audio_fps) > 0:
                sim_aud = np.dot(t_aud_fps, ref.audio_fps.T)
                q_idxs, r_idxs = np.where(sim_aud >= self.audio_threshold)
                for q_i, r_i in zip(q_idxs, r_idxs):
                    raw_matches.append(RawPointMatch(
                        t_target=float(t_aud_times[q_i]),
                        t_ref=float(ref.audio_times[r_i]),
                        ref_video_id=ref.video_id,
                        score=float(sim_aud[q_i, r_i]),
                        modality="audio",
                    ))

            if raw_matches:
                segments = self._cluster_and_align(ref.video_id, raw_matches)
                all_segments.extend(segments)

        return self._temporal_nms(all_segments)

    def _temporal_nms(self, segments: List[MatchSegment], iou_threshold: float = 0.3) -> List[MatchSegment]:
        """
        Non-Maximum Suppression (NMS) theo trục thời gian:
        Loại bỏ các đoạn trùng lặp chồng lấn của cùng một video tham chiếu,
        chỉ giữ lại phân đoạn có độ dài và confidence cao nhất.
        """
        if not segments:
            return []

        # Sắp xếp theo confidence giảm dần, ưu tiên độ dài lớn hơn khi điểm bằng nhau
        sorted_segs = sorted(segments, key=lambda s: (s.confidence_score, s.duration_sec), reverse=True)
        kept: List[MatchSegment] = []

        for seg in sorted_segs:
            overlap = False
            for k in kept:
                # Kiểm tra chồng lấn thời gian trên video target
                inter_start = max(seg.target_start_sec, k.target_start_sec)
                inter_end = min(seg.target_end_sec, k.target_end_sec)
                inter = max(0.0, inter_end - inter_start)
                if inter > 0:
                    union = (seg.duration_sec + k.duration_sec) - inter
                    iou = inter / union if union > 0 else 0
                    if iou > iou_threshold or inter >= min(seg.duration_sec, k.duration_sec) * 0.7:
                        overlap = True
                        break
            if not overlap:
                kept.append(seg)

        # Sắp xếp lại theo thứ tự thời gian xuất hiện trong target video
        kept.sort(key=lambda s: s.target_start_sec)
        return kept

