"""
core/config_broadcast.py - 配置热加载广播机制
===============================================
当配置发生变更（如添加/删除模型、修改LLM参数等）时，通知所有正在执行任务的 Agent，
使它们能够安全地暂停、等待配置更新完成、然后恢复继续执行任务。

核心设计:
  - 基于 asyncio.Event 实现无轮询的高效等待
  - Agent 在每次迭代循环中检查是否需要暂停
  - 暂停时 Agent 保存当前上下文状态（checkpoint），等待配置更新
  - 配置更新完成后 Agent 从 checkpoint 恢复继续执行
"""
from __future__ import annotations

import asyncio
import time
from enum import Enum
from typing import Dict, Any, Optional, Set
from dataclasses import dataclass, field

from core.logger import get_logger

logger = get_logger("myagent.broadcast")


class ReloadType(str, Enum):
    """重载类型枚举，用于区分不同类型的更新操作"""
    CONFIG = "config"         # 仅配置热重载
    CODE = "code"             # 代码模块热重载 (importlib.reload)
    DEPENDENCY = "dependency"  # 依赖包更新 (pip install)
    FULL = "full"             # 全量更新（含进程重启）

    def __str__(self) -> str:
        return self.value


@dataclass
class TaskCheckpoint:
    """任务检查点 - 保存 Agent 暂停时的状态"""
    task_id: str = ""
    agent_name: str = ""
    iteration: int = 0
    paused_at: str = ""       # 暂停时间戳
    resumed_at: str = ""      # 恢复时间戳
    reload_version: int = 0   # 配置版本号（用于判断是否需要重新加载）
    metadata: Dict[str, Any] = field(default_factory=dict)


class ConfigBroadcaster:
    """
    配置变更广播器。

    使用方式:
        # 初始化
        broadcaster = ConfigBroadcaster()

        # Agent 注册（开始任务时）
        broadcaster.register(task_id="task_xxx", agent_name="main_agent")

        # Agent 在循环中检查（每次迭代前）
        if await broadcaster.check_and_wait(task_id="task_xxx"):
            logger.info("配置已更新，继续执行")

        # 配置变更时（热重载入口）
        await broadcaster.request_reload()
        # ... 更新配置 ...
        await broadcaster.complete_reload()

        # Agent 注销（任务完成时）
        broadcaster.unregister(task_id="task_xxx")
    """

    def __init__(self):
        # 事件: reload_requested -> reload_done
        self._reload_requested: asyncio.Event = asyncio.Event()
        self._reload_done: asyncio.Event = asyncio.Event()
        self._lock: asyncio.Lock = asyncio.Lock()

        # 活跃任务注册表: task_id -> TaskCheckpoint
        self._active_tasks: Dict[str, TaskCheckpoint] = {}
        # 已暂停的任务集合（用于等待所有任务暂停）
        self._paused_tasks: Set[str] = set()

        # 配置版本计数器（每次热重载递增）
        self._reload_version: int = 0
        # 是否正在重载中
        self._reloading: bool = False
        # 当前重载类型
        self._reload_type: ReloadType = ReloadType.CONFIG

        # 统计
        self._total_reloads: int = 0
        self._total_tasks_paused: int = 0

        # 更新回调（供 UpdateManager 注册进度通知）
        self._on_reload_callbacks: list = []

    @property
    def reload_version(self) -> int:
        return self._reload_version

    @property
    def is_reloading(self) -> bool:
        return self._reloading

    @property
    def active_task_count(self) -> int:
        return len(self._active_tasks)

    def register(self, task_id: str, agent_name: str = ""):
        """
        注册一个活跃任务。

        Agent 在开始执行任务时调用此方法，以便在配置变更时能收到通知。

        Args:
            task_id: 任务唯一标识符
            agent_name: Agent 名称（用于日志）
        """
        self._active_tasks[task_id] = TaskCheckpoint(
            task_id=task_id,
            agent_name=agent_name,
            reload_version=self._reload_version,
        )
        logger.debug(f"任务已注册: {task_id} (agent={agent_name}, "
                     f"活跃任务数={len(self._active_tasks)})")

    def unregister(self, task_id: str):
        """
        注销一个任务。

        Agent 在任务完成或失败时调用此方法。

        Args:
            task_id: 任务唯一标识符
        """
        removed = self._active_tasks.pop(task_id, None)
        self._paused_tasks.discard(task_id)
        if removed:
            logger.debug(f"任务已注销: {task_id} (活跃任务数={len(self._active_tasks)})")

    async def check_and_wait(self, task_id: str) -> tuple[bool, str]:
        """
        检查是否需要暂停并等待配置重载完成。

        Agent 应在其处理循环的每次迭代开始时调用此方法。
        如果配置正在重载，此方法会阻塞（async）直到重载完成。

        Returns:
            tuple[bool, str]: (是否发生了重载, 重载类型字符串)
            - (False, "") 表示无配置变更
            - (True, "config") 表示配置热重载
            - (True, "code") 表示代码模块热重载
        """
        if not self._reload_requested.is_set():
            return (False, "")

        checkpoint = self._active_tasks.get(task_id)
        if not checkpoint:
            return (False, "")

        reload_type = self._reload_type

        # 标记为暂停
        self._paused_tasks.add(task_id)
        checkpoint.paused_at = time.strftime("%Y-%m-%d %H:%M:%S")
        self._total_tasks_paused += 1

        logger.info(f"[{task_id}] ⏸️ 任务暂停，等待{reload_type.value}重载完成... "
                     f"(agent={checkpoint.agent_name}, iteration={checkpoint.iteration})")

        # 等待重载完成（全量更新给更长时间）
        wait_timeout = 120.0 if reload_type == ReloadType.FULL else 30.0
        try:
            await asyncio.wait_for(self._reload_done.wait(), timeout=wait_timeout)
        except asyncio.TimeoutError:
            logger.warning(f"[{task_id}] ⚠️ 等待重载超时({wait_timeout}s)，继续执行")
            return (False, "")

        checkpoint.resumed_at = time.strftime("%Y-%m-%d %H:%M:%S")
        checkpoint.reload_version = self._reload_version
        checkpoint.metadata["reload_type"] = reload_type.value
        self._paused_tasks.discard(task_id)

        logger.info(f"[{task_id}] ▶️ 任务恢复，{reload_type.value}已更新 "
                     f"(version={self._reload_version})")

        return (True, reload_type.value)

    async def request_reload(self, reload_type: ReloadType = ReloadType.CONFIG) -> int:
        """
        请求重载 - 通知所有活跃任务暂停。

        Args:
            reload_type: 重载类型（CONFIG/CODE/DEPENDENCY/FULL），
                       Agent 恢复后可根据类型决定是否需要额外处理。

        Returns:
            int: 当前活跃任务数量
        """
        async with self._lock:
            if self._reloading:
                logger.warning(f"重载已在进行中 (类型={self._reload_type.value})，跳过")
                return 0

            self._reloading = True
            self._reload_type = reload_type
            active_count = len(self._active_tasks)

            if active_count == 0:
                logger.info("📡 配置重载广播: 无活跃任务，直接执行")
                return 0

            # 设置重载请求信号
            self._reload_requested.set()
            self._reload_done.clear()

            logger.info(f"📡 配置重载广播: 已通知 {active_count} 个活跃任务暂停")

            # 等待所有活跃任务暂停（最多 10 秒）
            if active_count > 0:
                try:
                    await asyncio.wait_for(
                        self._wait_all_paused(),
                        timeout=10.0
                    )
                except asyncio.TimeoutError:
                    paused_count = len(self._paused_tasks)
                    logger.warning(
                        f"⚠️ 等待任务暂停超时: {paused_count}/{active_count} 已暂停"
                    )

            return active_count

    async def _wait_all_paused(self):
        """等待所有活跃任务都暂停"""
        while True:
            all_paused = all(
                tid in self._paused_tasks
                for tid in self._active_tasks
            )
            if all_paused:
                return
            await asyncio.sleep(0.1)

    async def complete_reload(self):
        """
        通知所有暂停的任务配置重载已完成。

        在配置实际更新完成后调用此方法，所有暂停的任务将恢复执行。
        """
        self._reload_version += 1
        self._total_reloads += 1

        # 清除重载请求信号（防止重复触发）
        self._reload_requested.clear()

        # 设置重载完成信号（唤醒所有等待的 Agent）
        self._reload_done.set()

        # 给 Agent 一点时间恢复
        await asyncio.sleep(0.3)

        self._reloading = False

        logger.info(f"✅ 配置重载广播完成 (version={self._reload_version}, "
                     f"累计重载={self._total_reloads})")

    def on_reload(self, callback):
        """注册重载事件回调（供 UpdateManager 等外部模块监听）"""
        self._on_reload_callbacks.append(callback)

    def get_stats(self) -> Dict[str, Any]:
        """获取广播器统计信息"""
        return {
            "reload_version": self._reload_version,
            "reload_type": self._reload_type.value,
            "total_reloads": self._total_reloads,
            "active_tasks": len(self._active_tasks),
            "paused_tasks": len(self._paused_tasks),
            "total_tasks_paused": self._total_tasks_paused,
            "is_reloading": self._reloading,
            "tasks": [
                {
                    "task_id": cp.task_id,
                    "agent_name": cp.agent_name,
                    "iteration": cp.iteration,
                    "paused": cp.task_id in self._paused_tasks,
                    "paused_at": cp.paused_at,
                    "resumed_at": cp.resumed_at,
                    "reload_type": cp.metadata.get("reload_type", ""),
                }
                for cp in self._active_tasks.values()
            ]
        }

    def force_reset(self):
        """
        强制重置广播器状态。

        仅用于测试或异常恢复场景。
        """
        self._reload_requested.clear()
        self._reload_done.clear()
        self._reloading = False
        self._paused_tasks.clear()
        logger.warning("广播器状态已强制重置")
