import os
import sys
import cv2
import numpy as np
from pathlib import Path

# Đảm bảo UTF-8
if sys.platform == "win32":
    try:
        sys.stdout.reconfigure(encoding="utf-8")
        sys.stderr.reconfigure(encoding="utf-8")
    except Exception:
        pass

def extract_head_tail(video_path: str):
    cap = cv2.VideoCapture(video_path)
    total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
    fps = cap.get(cv2.CAP_PROP_FPS) or 24.0
    
    # Head frame (Frame 0)
    cap.set(cv2.CAP_PROP_POS_FRAMES, 0)
    ret_head, head_frame = cap.read()
    
    # Tail frame (Last frame)
    cap.set(cv2.CAP_PROP_POS_FRAMES, max(0, total_frames - 1))
    ret_tail, tail_frame = cap.read()
    
    cap.release()
    return head_frame, tail_frame, total_frames, fps

def compare_frames(frame_a, frame_b):
    """
    So sánh độ tương đồng thị giác giữa 2 khung hình:
    - Histogram Correlation (đo lường tương đồng màu sắc & ánh sáng)
    - Structural Template Match / Normalized Absolute Difference
    """
    if frame_a is None or frame_b is None:
        return 0.0
    
    # Resize về cùng kích thước chuẩn để so sánh
    h, w = 360, 640
    img_a = cv2.resize(frame_a, (w, h))
    img_b = cv2.resize(frame_b, (w, h))
    
    # 1. So sánh Histogram màu sắc (HSV)
    hsv_a = cv2.cvtColor(img_a, cv2.COLOR_BGR2HSV)
    hsv_b = cv2.cvtColor(img_b, cv2.COLOR_BGR2HSV)
    hist_a = cv2.calcHist([hsv_a], [0, 1], None, [50, 60], [0, 180, 0, 256])
    hist_b = cv2.calcHist([hsv_b], [0, 1], None, [50, 60], [0, 180, 0, 256])
    cv2.normalize(hist_a, hist_a, 0, 1, cv2.NORM_MINMAX)
    cv2.normalize(hist_b, hist_b, 0, 1, cv2.NORM_MINMAX)
    color_corr = cv2.compareHist(hist_a, hist_b, cv2.HISTCMP_CORREL)
    
    # 2. So sánh Gray Difference
    gray_a = cv2.cvtColor(img_a, cv2.COLOR_BGR2GRAY)
    gray_b = cv2.cvtColor(img_b, cv2.COLOR_BGR2GRAY)
    diff = cv2.absdiff(gray_a, gray_b)
    diff_score = 1.0 - (np.mean(diff) / 255.0)
    
    # Tổng hợp độ khớp (0.0 đến 1.0)
    match_score = (max(0, color_corr) * 0.6) + (diff_score * 0.4)
    return float(match_score)

def check_shot_continuity(video_prev_path: str, video_next_path: str, threshold: float = 0.65):
    print(f"\n[*] Đang kiểm tra nối tiếp:")
    print(f"    - Shot trước: {Path(video_prev_path).name}")
    print(f"    - Shot sau:   {Path(video_next_path).name}")
    
    _, tail_prev, _, _ = extract_head_tail(video_prev_path)
    head_next, _, _, _ = extract_head_tail(video_next_path)
    
    if tail_prev is None or head_next is None:
        print("    [!] Lỗi: Không thể trích xuất khung hình từ một trong hai video.")
        return False, 0.0

    score = compare_frames(tail_prev, head_next)
    print(f"    👉 Điểm tương đồng chuyển cảnh: {score:.2%}")
    
    if score >= threshold:
        print(f"    [✓] KHỚP MẠCH PHIM: Hai shot liền mạch, không cần chuyển cảnh phụ.")
        return True, score
    else:
        print(f"    [!] CẢNH BÁO KHỰNG HÌNH: Độ lệch thị giác vượt ngưỡng an toàn.")
        print(f"    👉 Cần bổ sung đoạn nối (Bridging Transition) hoặc Cross-Dissolve 1-2 giây.")
        return False, score

if __name__ == "__main__":
    if len(sys.argv) >= 3:
        check_shot_continuity(sys.argv[1], sys.argv[2])
    else:
        print("Sử dụng: python continuity_checker.py <video_truoc.mp4> <video_sau.mp4>")
