"""
Timeline Analyzer Module
时间线分析模块

功能：时间排序、传播路径构建、关键节点识别
"""

from typing import List, Dict, Optional
from datetime import datetime
from dataclasses import dataclass, field
import pytz


@dataclass
class Post:
    """帖子信息"""
    unique_id: str
    platform: str
    account: str
    content: str = ""
    image_urls: List[str] = field(default_factory=list)
    publish_time: Optional[datetime] = None
    repost_of: Optional[str] = None  # 转发自哪个帖子ID
    metadata: Dict = field(default_factory=dict)


@dataclass
class Node:
    """传播节点"""
    account: str
    platform: str
    post_id: str
    publish_time: datetime
    role: str = "unknown"  # "original", "repost", "amplifier", "modifier"
    connections: List[str] = field(default_factory=list)  # 连接的节点ID
    metadata: Dict = field(default_factory=dict)


@dataclass
class Timeline:
    """传播时间线"""
    nodes: List[Node] = field(default_factory=list)
    events: List[Dict] = field(default_factory=list)
    start_time: Optional[datetime] = None
    end_time: Optional[datetime] = None


class TimelineAnalyzer:
    """时间线分析器"""

    def __init__(self, timezone: str = "UTC"):
        """
        初始化时间线分析器

        Args:
            timezone: 时区，默认UTC
        """
        self.timezone = pytz.timezone(timezone)

    def parse_time(self, time_str: str) -> Optional[datetime]:
        """
        解析时间字符串为datetime对象

        Args:
            time_str: 时间字符串（ISO 8601格式或其他常见格式）

        Returns:
            datetime对象，解析失败返回None
        """
        try:
            # 尝试ISO 8601格式
            dt = datetime.fromisoformat(time_str.replace("Z", "+00:00"))

"))

            # 转换为指定时区
            if dt.tzinfo is None:
                dt = self.timezone.localize(dt)
            else:
                dt = dt.astimezone(self.timezone)

            return dt

        except Exception as e:
            print(f"Error parsing time '{time_str}': {e}")
            return None

    def build_timeline(
        self,
        posts: List[Post],
        target_post_id: Optional[str] = None
    ) -> Timeline:
        """
        构建传播时间线

        Args:
            posts: 帖子列表
            target_post_id: 目标帖子ID（用于定位时间线终点）

        Returns:
            Timeline对象
        """
        timeline = Timeline()

        if not posts:
            return timeline

        # 过滤掉没有时间的帖子
        valid_posts = [p for p in posts if p.publish_time is not None]

        if not valid_posts:
            return timeline

        # 按时间排序
        valid_posts.sort(key=lambda p: p.publish_time)

        # 转换为节点
        nodes = []
        for post in valid_posts:
            node = Node(
                account=post.account,
                platform=post.platform,
                post_id=post.unique_id,
                publish_time=post.publish_time,
                metadata={
                    "content": post.content,
                    "image_urls": post.image_urls,
                    "repost_of": post.repost_of
                }
            )
            nodes.append(node)

        # 构建事件
        events = []
        for i, node in enumerate(nodes):
            event = {
                "time": node.publish_time,
                "event": "original_post" if i == 0 else "repost",
                "account": node.account,
                "platform": node.platform,
                "post_id": node.post_id
            }

            if node.metadata.get("repost_of"):
                event["repost_of"] = node.metadata["repost_of"]
                event["event"] = "repost"

            events.append(event)

        # 设置时间范围
        timeline.start_time = valid_posts[0].publish_time
        timeline.end_time = valid_posts[-1].publish_time
        timeline.nodes = nodes
        timeline.events = events

        return timeline

    def identify_original_source(
        self,
        posts: List[Post]
    ) -> Optional[Post]:
        """
        识别可能的原始来源

        Args:
            posts: 帖子列表

        Returns:
            最可能的原始来源帖子
        """
        if not posts:
            return None

        # 过滤掉没有时间的帖子
        valid_posts = [p forre p in posts if p.publish_time is not None]

        if not valid_posts:
            return None

        # 按时间排序
        valid_posts.sort(key=lambda p: p.publish_time)

        # 最早的帖子可能是原始来源
        earliest_post = valid_posts[0]

        # 检查是否有repost_of字段指向更早的帖子
        post_id_map = {p.unique_id: p for p in valid_posts}

        for post in valid_posts:
            if post.repost_of and post.repost_of in post_id_map:
                referenced_post = post_id_map[post.repost_of]
                if referenced_post.publish_time < earliest_post.publish_time:
                    earliest_post = referenced_post

        return earliest_post

    def find_key_nodes(
        self,
        timeline: Timeline,
        min_connections: int = 2
    ) -> List[Node]:
        """
        识别关键传播节点

        Args:
            timeline: Timeline对象
            min_connections: 最小连接数阈值

        Returns:
            关键节点列表
        """
        if not timeline.nodes:
            return []

        # 计算每个节点的连接数
        connection_counts = {}
        for node in timeline.nodes:
            connection_counts[node.post_id] = len(node.connections)

        # 筛选关键节点
        key_nodes = [
            node for node in timeline.nodes
            if connection_counts.get(node.post_id, 0) >= min_connections
        ]

        # 如果没有满足条件的节点，返回前3个
        if not key_nodes:
            key_nodes = timeline.nodes[:3]

        return key_nodes

    def find_amplifiers(
        self,
        timeline: Timeline,
        followers_threshold: int = 1000
    ) -> List[Node]:
        """
        识别主要放大器（高影响力账号）

        Args:
            timeline: Timeline对象
            followers_threshold: 粉丝数阈值

        Returns:
            放大器节点列表
        """
        amplifiers = []

        for node in timeline.nodes:
            followers = node.metadata.get("followers", 0)

            if followers >= followers_threshold:
                node.role = "amplifier"
                amplifiers.append(node)

        return amplifiers

    def calculate_propagation_speed(
        self,
        timeline: Timeline
    ) -> Dict[str, float]:
        """
        计算传播速度

        Args:
            timeline: Timeline对象

        Returns:
            传播速度指标
        """
        if not timeline.start_time or not timeline.end_time:
            return {}

        total_duration = (timeline.end_time - timeline.start_time).total_seconds()

        if total_duration <= 0:
            return {}

        node_count = len(timeline.nodes)

        metrics = {
            "total_duration_seconds": total_duration,
            "total_duration_minutes": total_duration / 60,
            "total_duration_hours": total_duration / 3600,
            "node_count": node_count,
            "average_time_between_nodes": total_duration / max(node_count - 1, 1),
            "propagation_rate_nodes_per_hour": node_count / (total_duration / 3600)
        }

        return metrics

    def visualize_timeline(
        self,
        timeline: Timeline
    ) -> str:
        """
        生成时间线的文本可视化

        Args:
            timeline: Timeline对象

        Returns:
            时间线可视化文本
        """
        if not timeline.events:
            return "Empty timeline"

        lines = []
        lines.append("=" * 60)
        lines.append("Propagation Timeline")
        lines.append("=" * 60)

        for i, event in enumerate(timeline.events):
            time_str = event["time"].strftime("%Y-%m-%d %H:%M:%S")
            event_type = event["event"]
            account = event["account"]
            platform = event["platform"]
            post_id = event["post_id"]

            line = f"\n[{i+1}] {time_str} | {event_type} | {account}@{platform}"
            line += f"\n    Post ID: {post_id}"

            if "repost_of" in event:
                line += f"\n    Repost of: {event['repost_of']}"

            lines.append(line)

        # 添加传播速度
        metrics = self.calculate_propagation_speed(timeline)
        if metrics:
            lines.append("\n" + "=" * 60)
            lines.append("Propagation Metrics")
            lines.append("=" * 60)
            lines.append(f"Total Duration: {metrics.get('total_duration_hours', 0):.2f} hours")
            lines.append(f"Node Count: {metrics.get('node_count', 0)}")
            lines.append(f"Average Time Between Nodes: {metrics.get('average_time_between_nodes', 0):.2f} seconds")
            lines.append(f"Propagation Rate: {metrics.get('propagation_rate_nodes_per_hour', 0):.2f} nodes/hour")

        return "\n".join(lines)
