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

"""
传播路径分析模块 (Propagation Path Analyzer)

分析信息在社交媒体上的传播路径，识别关键节点
"""

from typing import Dict, List, Any
from datetime import datetime
from collections import defaultdict


class PathAnalyzer:
    """传播路径分析器"""

    def __init__(self):
        self.amplifier_threshold = 10000  # 粉丝数超过1万视为放大器
        self.time_gap_threshold = 3600  # 1小时内的连续传播

    async def analyze(
        self,
        search_results: List[Dict[str, Any]]
    ) -> Dict[str, Any]:
        """
        分析传播路径

        Args:
            search_results: 反向搜索结果

        Returns:
            Dict: 传播路径分析结果
        """
        if not search_results:
            return {
                "path_type": "无数据",
                "key_nodes": [],
                "visual_representation": None,
                "total_hops": 0,
                "time_span": None
            }

        # 按时间排序
        sorted_results = sorted(
            search_results,
            key=lambda x: self._parse_timestamp(x.get('timestamp', '')),
            reverse=False
        )

        # 识别关键节点
        key_nodes = self._identify_key_nodes(sorted_results)

        # 构建传播链路
        propagation_chain = self._build_propagation_chain(sorted_results)

        # 计算传播跳数
        total_hops = len(propagation_chain) - 1

        # 计算时间跨度
        time_span = self._calculate_time_span(sorted_results)

        # 生成可视化数据
        visual_representation = self._generate_visual_representation(
            propagation_chain,
            key_nodes
        )

        result = {
            "path_type": "时间线",
            "key_nodes": key_nodes,
            "propagation_chain": propagation_chain,
            "total_hops": total_hops,
            "time_span": time_span,
            "visual_representation": visual_representation
        }

        return result

    def _identify_key_nodes(
        self,
        sorted_results: List[Dict[str, Any]]
    ) -> List[Dict[str, Any]]:
        """
        识别关键节点

        关键节点类型：
        - source: 原始来源（最早）
        - amplifier: 放大器（大V、媒体账号）
        - bridge: 桥接节点（跨平台传播）
        - hub: 传播枢纽（转发数量多）
        """
        if not sorted_results:
            return []

        key_nodes = []

        # 原始来源
        if sorted_results:
            source = {
                "node_id": "node_0",
                "account": sorted_results[0].get('account_name', 'Unknown'),
                "platform": sorted_results[0].get('platform', 'Unknown'),
                "time": sorted_results[0].get('timestamp', ''),
                "role": "源头",
                "followers": sorted_results[0].get('followers', 0)
            }
            key_nodes.append(source)

        # 识别放大器
        amplifiers = []
        platforms_seen = {source['platform']} if key_nodes else set()

        for idx, result in enumerate(sorted_results[1:], start=1):
            followers = result.get('followers', 0)
            platform = result.get('platform', 'Unknown')

            # 大V或媒体账号
            if followers >= self.amplifier_threshold:
                amplifier = {
                    "node_id": f"node_{idx}",
                    "account": result.get('account_name', 'Unknown'),
                    "platform": platform,
                    "time": result.get('timestamp', ''),
                    "role": "放大器",
                    "followers": followers
                }
                amplifiers.append(amplifier)

            # 跨平台桥接
            if platform not in platforms_seen:
                bridge = {
                    "node_id": f"node_{idx}_bridge",
                    "account": result.get('account_name', 'Unknown'),
                    "platform": platform,
                    "time": result.get('timestamp', ''),
                    "role": "跨平台传播节点",
"                    "followers": followers
                }
                key_nodes.append(bridge)
                platforms_seen.add(platform)

        # 添加放大器
        key_nodes.extend(amplifiers[:5])  # 最多添加5个放大器

        return key_nodes

    def _build_propagation_chain(
        self,
        sorted_results: List[Dict[str, Any]]
    ) -> List[Dict[str, Any]]:
        """
        构建传播链路

        Args:
            sorted_results: 按时间排序的搜索结果

        Returns:
            List: 传播链路
        """
        chain = []

        for idx, result in enumerate(sorted_results):
            node = {
                "hop": idx,
                "account": result.get('account_name', 'Unknown'),
                "platform": result.get('platform', 'Unknown'),
                "time": result.get('timestamp', ''),
                "url": result.get('source_url', ''),
                "similarity": result.get('similarity', 0)
            }
            chain.append(node)

        return chain

    def _calculate_time_span(
        self,
        sorted_results: List[Dict[str, Any]]
    ) -> Dict[str, str]:
        """
        计算时间跨度

        Args:
            sorted_results: 按时间排序的搜索结果

        Returns:
            Dict: 时间跨度信息
        """
        if len(sorted_results) < 2:
            return None

        first_time = self._parse_timestamp(sorted_results[0].get('timestamp', ''))
        last_time = self._parse_timestamp(sorted_results[-1].get('timestamp', ''))

        if first_time is None or last_time is None:
            return None

        time_diff = last_time - first_time

        return {
            "start": sorted_results[0].get('timestamp', ''),
            "end": sorted_results[-1].get('timestamp', ''),
            "duration_hours": time_diff.total_seconds() / 3600,
            "duration_days": time_diff.total_seconds() / (24 * 3600)
        }

    def _generate_visual_representation(
        self,
        propagation_chain: List[Dict[str, Any]],
        key_nodes: List[Dict[str, Any]]
    ) -> Dict[str, Any]:
        """
        生成可视化数据（用于前端渲染）

        Args:
            propagation_chain: 传播链路
            key_nodes: 关键节点

        Returns:
            Dict: 可视化数据
        """
        # 生成节点数据
        nodes = []
        edges = []

        for node in propagation_chain:
            nodes.append({
                "id": node['hop'],
                "label": f"{node['account']}@{node['platform']}",
                "time": node['time'],
                "type": "normal"
            })

            # 生成边
            if node['hop'] > 0:
                edges.append({
                    "source": node['hop'] - 1,
                    "target": node['hop'],
                    "similarity": node['similarity']
                })

        # 标记关键节点
        key_node_ids = {kn['account'] for kn in key_nodes}
        for node in nodes:
            if node['label'].split('@')[0] in key_node_ids:
                node['type'] = 'key'

        return {
            "nodes": nodes,
            "edges": edges,
            "type": "directed_graph"
        }

    def _parse_timestamp(self, timestamp_str: str) -> datetime:
        """
        解析时间戳

        Args:
            timestamp_str: 时间戳字符串

        Returns:
            datetime: 解析后的时间对象
        """
        formats = [
            "%Y-%m-%dT%H:%M:%SZ",
            "%Y-%m-%dT%H:%M:%S.%fZ",
            "%Y-%m-%d %H:%M:%S",
            "%Y-%m-%d",
            "%Y/%m/%d %H:%M:%S"
        ]

        for fmt in formats:
            try:
                return datetime.strptime(timestamp_str, fmt)
            except ValueError:
                continue

        return None
