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

"""
内容比对模块 (Content Comparator)

比对文本和图片的相似度，标注差异
"""

import Levenshtein
import difflib
import cv2
import numpy as np
from PIL import Image
import imagehash
from typing import Dict, List, Tuple, Any
from pathlib import Path


class ContentComparator:
    """内容比对器"""

    def __init__(self):
        self.text_weight = 0.5
        self.image_weight = 0.5

    def compare_text(
        self,
        original_text: str,
        candidate_texts: List[str]
    ) -> Dict[str, Any]:
        """
        比对文本相似度

        Args:
            original_text: 原始文本
            candidate_texts: 候选文本列表

        Returns:
            Dict: 比对结果
        """
        if not candidate_texts:
            return {"max": 0, "min": 0, "avg": 0, "details": []}

        similarities = []
        details = []

        for idx, candidate_text in enumerate(candidate_texts):
            # 计算Levenshtein距离
            distance = Levenshtein.distance(original_text, candidate_text)
            max_len = max(len(original_text), len(candidate_text))

            # 计算相似度（1 - 归一化距离）
            if max_len == 0:
                similarity = 1.0
            else:
                similarity = 1.0 - (distance / max_len)

            similarity = round(similarity, 4)
            similarities.append(similarity)

            # 生成差异标注
            diff = self._generate_text_diff(original_text, candidate_text)

            details.append({
                "index": idx,
                "similarity": similarity,
                "distance": distance,
                "diff": diff
            })

        result = {
            "max": round(max(similarities), 4),
            "min": round(min(similarities), 4),
            "avg": round(sum(similarities) / len(similarities), 4),
            "details": details
        }

        return result

    def compare_images(
        self,
        original_images: List[str],
        candidate_images: List[str]
    ) -> Dict[str, Any]:
        """
        比对图片相似度

        Args:
            original_images: 原始图片路径列表
            candidate_images: 候选图片路径列表

        Returns:
            Dict: 比对结果
        """
        if not original_images or not candidate_images:
            return {"max": 0, "min": 0, "avg": 0, "details": []}

        similarities = []
        details = []

        for orig_idx, orig_img in enumerate(original_images):
            for cand_idx, cand_img in enumerate(candidate_images):
                # 计算多种相似度指标
                ssim_score = self._calculate_ssim(orig_img, cand_img)
                hash_score = self._calculate_image_hash_similarity(orig_img, cand_img)

                # 综合相似度
                similarity = round((ssim_score + hash_score) / 2, 4)
                similarities.append(similarity)

                details.append({
                    "original_index": orig_idx,
                    "candidate_index": cand_idx,
                    "similarity": similarity,
                    "ssim": round(ssim_score, 4),
                    "hash": round(hash_score, 4)
                })

        if not similarities:
            return {"max": 0, "min": 0, "avg": 0, "details": []}

        result = {
            "max": round(max(similarities), 4),
            "min": round(min(similarities), 4),
            "avg": round(sum(similarities) / len(similarities), 4),
            "details": details
        }

        return result

    def _generate_text_diff(
        self,
        text1: str,
        text2: str
    ) -> str:
        """
        生成文本差异标注

        Args:
            text1: 文本1
            text2: 文本2

        Returns:
            str: 差异标注字符串
        """
        diff = difflib.unified_diff(
            text1.splitlines(keepends=True),
            text2.splitlines(keepends=True),
            fromfile='original',
            tofile='candidate',
            lineterm=''
        )

        return ''.join(diff)

    def _calculate_ssim(
        self,
        img1_path: str,
        img2_path: str
    ) -> float:
        """
        计算结构相似性（SSIM）

        Args:
            img1_path: 图片1路径
            img2_path: 图片2路径

        Returns:
            float: SSIM值（0-1）
        """
        try:
            # 读取图片
            img1 = cv2.imread(img1_path)
            img2 = cv2.imread(img2_path)

            if img1 is None or img2 is None:
                return 0.0

            # 调整大小
            h, w = img1.shape[:2]
            img2 = cv2.resize(img2, (w, h))

            # 转换为灰度
            gray1 = cv2.cvtColor(img1, cv2.COLOR_BGR2GRAY)
            gray2 = cv2.cvtColor(img2, cv2.COLOR_BGR2GRAY)

            # 计算SSIM
            C1 = (0.01 * 255) ** 2
            C2 = (0.03 * 255) ** 2

            mu1 = cv2.GaussianBlur(gray1, (11, 11), 1.5)
            mu2 = cv2.GaussianBlur(gray2, (11, 11), 1.5)

            mu1_sq = mu1 ** 2
            mu2_sq = mu2 ** 2
            mu1_mu2 = mu1 * mu2

            sigma1_sq = cv2.GaussianBlur(gray1 ** 2, (11, 11), 1.5) - mu1_sq
            sigma2_sq = cv2.GaussianBlur(gray2 ** 2, (11, 11), 1.5) - mu2_sq
            sigma12 = cv2.GaussianBlur(gray1 * gray2, (11, 11), 1.5) - mu1_mu2

            numerator = (2 * mu1_mu2 + C1) * (2 * sigma12 + C2)
            denominator = (mu1_sq + mu2_sq + C1) * (sigma1_sq + sigma2_sq + C2)

            ssim_map = numerator / denominator
            ssim = np.mean(ssim_map)

            return float(ssim)

        except Exception as e:
            print(f"计算SSIM失败: {e}")
            return 0.0

    def _calculate_image_hash_similarity(
        self,
        img1_path: str,
        img2_path: str
    ) -> float:
        """
        计算感知哈希相似度

        Args:
            img1_path: 图片1路径
            img2_path: 图片2路径

        Returns:
            float: 相似度（0-1）
        """
        try:
            # 读取图片
            img1 = Image.open(img1_path)
            img2 = Image.open(img2_path)

            # 计算感知哈希
            hash1 = imagehash.phash(img1)
            hash2 = imagehash.phash(img2)

            # 计算汉明距离
            distance = hash1 - hash2

            # 转换为相似度（最大距离为64）
            max_distance = 64
            similarity = 1.0 - (distance / max_distance)

            return float(similarity)

        except Exception as e:
            print(f"计算哈希相似度失败: {e}")
            return 0.0
