"""
Unit & Integration Tests for VideoDetection Pipeline.
Kiểm thử toàn bộ hệ thống: Paths, Extractor, Database, Temporal Alignment.
"""
from __future__ import annotations

import unittest
from pathlib import Path
import numpy as np

from src.core.paths import APP_ROOT, DATA_DIR, TARGET_DIR, TEMP_DIR, clean_temp_dir, get_ffmpeg_path
from src.core.config import CONFIG
from src.visual.vjepa_extractor import VJEPAExtractor
from src.audio.neuralfp_extractor import NeuralFPExtractor
from src.indexing.database import OwnerDatabase, VideoRecord
from src.alignment.matcher import TemporalMatcher


class TestVideoDetectionPipeline(unittest.TestCase):
    def test_paths_and_bins(self) -> None:
        self.assertTrue(TARGET_DIR.exists())
        self.assertTrue(DATA_DIR.exists())
        self.assertTrue(TEMP_DIR.exists())
        ffmpeg_p = get_ffmpeg_path()
        self.assertTrue(Path(ffmpeg_p).exists(), f"FFmpeg path không tồn tại: {ffmpeg_p}")

    def test_clean_temp(self) -> None:
        # Tạo file rác trong Temp/
        dummy_file = TEMP_DIR / "dummy_frame.tmp"
        dummy_file.write_text("test_content", encoding="utf-8")
        self.assertTrue(dummy_file.exists())

        removed = clean_temp_dir()
        self.assertGreaterEqual(removed, 1)
        self.assertFalse(dummy_file.exists())
        # .gitkeep phải còn nguyên
        gitkeep = TEMP_DIR / ".gitkeep"
        if gitkeep.exists():
            self.assertTrue(gitkeep.exists())

    def test_config_loading(self) -> None:
        self.assertIn("similarity_threshold", CONFIG.visual)
        self.assertIn("similarity_threshold", CONFIG.audio)
        self.assertIn("min_match_duration_sec", CONFIG.alignment)

    def test_vjepa_extractor(self) -> None:
        extractor = VJEPAExtractor(embedding_dim=1024)
        # Dummy clip: [16 frames, 224, 224, 3] uint8
        dummy_clip = np.zeros((16, 224, 224, 3), dtype=np.uint8)
        emb = extractor.extract(dummy_clip)

        self.assertEqual(emb.shape, (1024,))
        # Kiểm tra L2-norm xấp xỉ 1.0
        norm = np.linalg.norm(emb)
        self.assertAlmostEqual(norm, 1.0, places=3)

    def test_neuralfp_extractor(self) -> None:
        extractor = NeuralFPExtractor(fingerprint_dim=128)
        # Dummy audio chunk: 1.0s @ 16000Hz = 16000 samples float32
        dummy_chunk = np.random.randn(16000).astype(np.float32)
        fp = extractor.extract(dummy_chunk)

        self.assertEqual(fp.shape, (128,))
        norm = np.linalg.norm(fp)
        self.assertAlmostEqual(norm, 1.0, places=3)

    def test_database_and_temporal_alignment(self) -> None:
        """
        Kiểm thử end-to-end: Nạp video gốc (10s), tạo target video copy 5s từ đoạn [2s - 7s],
        và kiểm chứng TemporalMatcher phát hiện chính xác đoạn trùng lặp.
        """
        # Tạo database cô lập tạm trong Temp/test_db
        test_db_dir = TEMP_DIR / "test_db"
        test_db_dir.mkdir(parents=True, exist_ok=True)
        db = OwnerDatabase(data_dir=test_db_dir)

        # 1. Tạo 10 clips visual và audio cho video gốc (mỗi clip cách nhau 1.0s)
        ref_visual_feats = []
        ref_audio_feats = []
        np.random.seed(42)

        for sec in range(10):
            # Tạo vector ngẫu nhiên L2-normalized
            v_vec = np.random.randn(1024).astype(np.float32)
            v_vec /= np.linalg.norm(v_vec)
            ref_visual_feats.append((float(sec), v_vec))

            a_vec = np.random.randn(128).astype(np.float32)
            a_vec /= np.linalg.norm(a_vec)
            ref_audio_feats.append((float(sec), a_vec))

        rec = VideoRecord(
            video_id="goc_video_01",
            filename="goc_video_01.mp4",
            rel_path="goc_video_01.mp4",
            duration_sec=10.0,
            visual_clips_count=10,
            audio_chunks_count=10,
            indexed_at="2026-10-11 12:00:00",
        )
        db.add_reference(rec, ref_visual_feats, ref_audio_feats)

        # 2. Tạo video target dài 8 giây:
        # Trong đó từ giây 1.0s đến 6.0s (5 giây) là copy y nguyên từ đoạn [2.0s -> 7.0s] của video gốc!
        # Tức là: Delta = t_ref - t_target = 2.0 - 1.0 = +1.0s
        query_vis = []
        query_aud = []

        # Giây 0 của target: không liên quan
        noise_v = np.random.randn(1024).astype(np.float32)
        noise_v /= np.linalg.norm(noise_v)
        query_vis.append((0.0, noise_v))

        noise_a = np.random.randn(128).astype(np.float32)
        noise_a /= np.linalg.norm(noise_a)
        query_aud.append((0.0, noise_a))

        # Giây 1 -> 5 của target: copy từ giây 2 -> 6 của ref (5 clips = 5 giây trùng lặp)
        for t_target in range(1, 6):
            t_ref = t_target + 1  # 2, 3, 4, 5, 6
            query_vis.append((float(t_target), ref_visual_feats[t_ref][1]))
            query_aud.append((float(t_target), ref_audio_feats[t_ref][1]))

        # 3. Chạy matcher
        matcher = TemporalMatcher(
            db=db,
            visual_threshold=0.75,
            audio_threshold=0.75,
            min_match_duration_sec=3.0,
        )

        matches = matcher.match(query_vis, query_aud)

        # Kiểm chứng
        self.assertGreater(len(matches), 0, "Matcher không tìm thấy đoạn trùng lặp!")
        top_match = matches[0]

        self.assertEqual(top_match.reference_video_id, "goc_video_01")
        self.assertAlmostEqual(top_match.target_start_sec, 1.0, delta=0.5)
        self.assertAlmostEqual(top_match.ref_start_sec, 2.0, delta=0.5)
        self.assertGreaterEqual(top_match.duration_sec, 3.0)
        self.assertGreaterEqual(top_match.confidence_score, 0.9)

    def test_feature_store_and_match_cached(self) -> None:
        """
        Kiểm thử các kịch bản so sánh linh hoạt:
        - 1 Target vs Cả kho [Ref A, Ref B]
        - 1 Target vs 1 Video chọn lọc [Ref B]
        - 1 Target vs Video không liên quan [Ref A]
        """
        from src.indexing.feature_store import CachedVideoFeatures

        np.random.seed(100)

        # Ref A (10 giây)
        ref_a_vis = np.random.randn(10, 1024).astype(np.float32)
        ref_a_vis /= np.linalg.norm(ref_a_vis, axis=1, keepdims=True)
        ref_a_aud = np.random.randn(10, 128).astype(np.float32)
        ref_a_aud /= np.linalg.norm(ref_a_aud, axis=1, keepdims=True)
        times_10 = np.arange(10, dtype=np.float32)

        ref_a = CachedVideoFeatures(
            video_id="ref_movie_A",
            filename="ref_movie_A.mp4",
            video_path="Data/ref_movie_A.mp4",
            duration_sec=10.0,
            file_size_bytes=1000,
            file_mtime=1.0,
            visual_times=times_10,
            visual_embs=ref_a_vis,
            audio_times=times_10,
            audio_fps=ref_a_aud,
        )

        # Ref B (15 giây)
        ref_b_vis = np.random.randn(15, 1024).astype(np.float32)
        ref_b_vis /= np.linalg.norm(ref_b_vis, axis=1, keepdims=True)
        ref_b_aud = np.random.randn(15, 128).astype(np.float32)
        ref_b_aud /= np.linalg.norm(ref_b_aud, axis=1, keepdims=True)
        times_15 = np.arange(15, dtype=np.float32)

        ref_b = CachedVideoFeatures(
            video_id="ref_movie_B",
            filename="ref_movie_B.mp4",
            video_path="Data/ref_movie_B.mp4",
            duration_sec=15.0,
            file_size_bytes=1500,
            file_mtime=1.0,
            visual_times=times_15,
            visual_embs=ref_b_vis,
            audio_times=times_15,
            audio_fps=ref_b_aud,
        )

        # Target (8 giây), trong đó giây [2 -> 6] (5s) sao chép từ giây [5 -> 9] của Ref B!
        target_vis = np.random.randn(8, 1024).astype(np.float32)
        target_vis /= np.linalg.norm(target_vis, axis=1, keepdims=True)
        target_aud = np.random.randn(8, 128).astype(np.float32)
        target_aud /= np.linalg.norm(target_aud, axis=1, keepdims=True)

        for i in range(2, 7):
            ref_idx = i + 3  # 5, 6, 7, 8, 9
            target_vis[i] = ref_b_vis[ref_idx]
            target_aud[i] = ref_b_aud[ref_idx]

        target_feat = CachedVideoFeatures(
            video_id="target_clip_X",
            filename="target_clip_X.mp4",
            video_path="Target/target_clip_X.mp4",
            duration_sec=8.0,
            file_size_bytes=800,
            file_mtime=1.0,
            visual_times=np.arange(8, dtype=np.float32),
            visual_embs=target_vis,
            audio_times=np.arange(8, dtype=np.float32),
            audio_fps=target_aud,
        )

        # Khởi tạo matcher
        db_mock = OwnerDatabase(data_dir=TEMP_DIR / "dummy_db")
        matcher = TemporalMatcher(
            db=db_mock,
            visual_threshold=0.8,
            audio_threshold=0.8,
            min_match_duration_sec=3.0,
        )

        # KỊCH BẢN 1: So sánh Target với CẢ KHO [Ref A, Ref B]
        matches_all = matcher.match_cached(target_feat, [ref_a, ref_b])
        self.assertEqual(len(matches_all), 1)
        self.assertEqual(matches_all[0].reference_video_id, "ref_movie_B")
        self.assertAlmostEqual(matches_all[0].target_start_sec, 2.0, delta=0.5)
        self.assertAlmostEqual(matches_all[0].ref_start_sec, 5.0, delta=0.5)

        # KỊCH BẢN 2: Người dùng chỉ chọn đối đầu với [Ref B]
        matches_b = matcher.match_cached(target_feat, [ref_b])
        self.assertEqual(len(matches_b), 1)
        self.assertEqual(matches_b[0].reference_video_id, "ref_movie_B")

        # KỊCH BẢN 3: Người dùng chọn so sánh với [Ref A] (Không liên quan)
        matches_a = matcher.match_cached(target_feat, [ref_a])
        self.assertEqual(len(matches_a), 0, "Không được có trùng lặp với Ref A!")

    def test_sqlite_database(self) -> None:
        """Kiểm thử CSDL SQLite nhúng: cấu trúc bảng, lưu trữ blob vector và truy vấn theo tên video."""
        from src.indexing.sqlite_db import VideoDatabase
        from src.indexing.feature_store import CachedVideoFeatures
        from src.alignment.matcher import MatchSegment

        test_db_file = TEMP_DIR / "test_videodetection.db"
        if test_db_file.exists():
            test_db_file.unlink()

        db = VideoDatabase(test_db_file)

        # Tạo sample features
        dummy_p = TEMP_DIR / "sample_video.mp4"
        dummy_p.write_bytes(b"0" * 1024)

        vis_vecs = np.ones((5, 1024), dtype=np.float32)
        aud_vecs = np.ones((5, 128), dtype=np.float32)
        times = np.arange(5, dtype=np.float32)

        feat = CachedVideoFeatures(
            video_id="sample_video",
            filename="sample_video.mp4",
            video_path=str(dummy_p),
            duration_sec=5.0,
            file_size_bytes=1024,
            file_mtime=dummy_p.stat().st_mtime,
            visual_times=times,
            visual_embs=vis_vecs,
            audio_times=times,
            audio_fps=aud_vecs,
        )

        # 1. Upsert video
        db.upsert_video(dummy_p, category="reference", features=feat)
        self.assertTrue(db.is_video_up_to_date(dummy_p))

        # 2. Truy vấn theo tên video
        loaded = db.get_features_by_name("sample_video.mp4")
        self.assertIsNotNone(loaded)
        self.assertEqual(loaded.filename, "sample_video.mp4")
        self.assertEqual(loaded.visual_embs.shape, (5, 1024))
        self.assertEqual(loaded.audio_fps.shape, (5, 128))
        np.testing.assert_array_almost_equal(loaded.visual_embs, vis_vecs)

        # 3. Ghi nhận lịch sử đối soát
        seg = MatchSegment(
            target_start_sec=1.0,
            target_end_sec=4.0,
            reference_video_id="sample_video",
            ref_start_sec=0.0,
            ref_end_sec=3.0,
            confidence_score=0.95,
            visual_score=0.96,
            audio_score=0.94,
            match_type="multimodal",
        )
        db.log_match_history("target_test.mp4", [seg])

        history = db.get_comparison_history(limit=10)
        self.assertEqual(len(history), 1)
        self.assertEqual(history[0]["target_video_name"], "target_test.mp4")
        self.assertEqual(history[0]["reference_video_name"], "sample_video")

        # 4. Thống kê
        stats = db.get_database_stats()
        self.assertEqual(stats["total_reference_videos"], 1)
        self.assertEqual(stats["total_match_history"], 1)

    def test_quantizer_and_hamming(self) -> None:
        """Kiểm thử PQ64 Product Quantization và 128-bit Hamming Bit-Vector."""
        from src.indexing.quantizer import (
            ProductQuantizer64,
            binarize_audio_fingerprint,
            compute_hamming_similarity,
        )

        # 1. Kiểm tra Binarization Audio (128 floats -> 16 bytes)
        float_fp = np.random.randn(128).astype(np.float32)
        bit_fp = binarize_audio_fingerprint(float_fp)
        self.assertEqual(bit_fp.shape, (16,))
        self.assertEqual(bit_fp.dtype, np.uint8)

        # Bản sao giống hệt phải có Hamming similarity = 1.0 (0 bit khác nhau)
        sim_self = compute_hamming_similarity(bit_fp[None, :], bit_fp[None, :])
        self.assertAlmostEqual(float(sim_self[0, 0]), 1.0, places=4)

        # Bản sao bị nhiễu nhẹ (chỉ lật 2 bit trong số 128 bit)
        noisy_bit_fp = bit_fp.copy()
        noisy_bit_fp[0] ^= 0b00000011
        sim_noisy = compute_hamming_similarity(bit_fp[None, :], noisy_bit_fp[None, :])
        # 126 / 128 = 0.984375
        self.assertGreaterEqual(float(sim_noisy[0, 0]), 0.98)

        # 2. Kiểm tra Product Quantization PQ64 (1024 floats -> 64 bytes uint8)
        pq = ProductQuantizer64(d=1024, m=64, nbits=8)
        vecs = np.random.randn(5, 1024).astype(np.float32)
        vecs /= np.linalg.norm(vecs, axis=1, keepdims=True)

        codes = pq.encode(vecs)
        self.assertEqual(codes.shape, (5, 64))
        self.assertEqual(codes.dtype, np.uint8)

        # Kiểm tra Asymmetric Distance Computation (ADC)
        adc_sim = pq.compute_asymmetric_distances(vecs, codes)
        self.assertEqual(adc_sim.shape, (5, 5))
        # Đường chéo chính (vector so với chính mã nén của nó) phải có điểm cao nhất trong hàng
        for i in range(5):
            self.assertEqual(int(np.argmax(adc_sim[i])), i)


if __name__ == "__main__":
    unittest.main()

