"""
core/task_persistence.py - 任务持久化与恢复
=====================================
将任务状态保存到磁盘，支持进程重启后恢复。
存储位置: ~/.myagent/data/tasks/
每个任务一个JSON文件: task_{id}.json
"""
from __future__ import annotations

import json
import os
from datetime import datetime, timezone, timedelta
from pathlib import Path
from typing import Any, Dict, List, Optional

from core.logger import get_logger
from core.utils import generate_id, timestamp

logger = get_logger("myagent.task_persistence")

# 允许的任务状态
VALID_STATUSES = {"pending", "running", "completed", "failed"}


class TaskPersistence:
    """
    任务持久化管理器。

    将群聊任务的状态保存为 JSON 文件到磁盘，
    以便进程重启后能检测到中断的任务。

    使用示例:
        tp = TaskPersistence(data_dir="~/.myagent/data")
        tp.initialize()

        tp.save_task(
            task_id="task_abc123",
            description="用户请求内容",
            session_id="group_xxx_agent_path",
            group_id="xxx",
            agent_path="程序员",
            status="running",
            metadata={"source": "group_chat"},
        )

        pending = tp.get_pending_tasks()
        for task in pending:
            print(f"未完成任务: {task['task_id']} - {task['description']}")
    """

    def __init__(self, data_dir: str | Path = ""):
        self.data_dir = Path(data_dir) if data_dir else Path()
        self._tasks_dir: Optional[Path] = None

    def initialize(self):
        """初始化任务存储目录"""
        self._tasks_dir = self.data_dir / "tasks"
        self._tasks_dir.mkdir(parents=True, exist_ok=True)
        logger.info(f"任务持久化已初始化 (dir={self._tasks_dir})")

    def _task_file(self, task_id: str) -> Path:
        """根据 task_id 返回 JSON 文件路径"""
        # 安全检查：防止路径遍历
        safe_id = os.path.basename(task_id)
        if safe_id != task_id or ".." in task_id or "/" in task_id:
            logger.warning(f"不安全的 task_id: {task_id}")
            safe_id = task_id.replace("/", "_").replace("..", "_")
        return self._tasks_dir / f"task_{safe_id}.json"

    def save_task(
        self,
        task_id: str,
        description: str,
        session_id: str,
        group_id: str = "",
        agent_path: str = "",
        status: str = "pending",
        metadata: Optional[Dict[str, Any]] = None,
        last_message: str = "",
    ) -> Dict[str, Any]:
        """
        保存或更新任务到磁盘。

        Args:
            task_id: 唯一任务标识
            description: 任务描述（通常是用户消息内容）
            session_id: 会话 ID
            group_id: 群聊 ID（如果是群聊任务）
            agent_path: Agent 的数字 aid
            status: 任务状态 (pending|running|completed|failed)
            metadata: 附加元数据
            last_message: 最后一条消息内容

        Returns:
            保存的任务数据字典
        """
        if not self._tasks_dir:
            self.initialize()

        if status not in VALID_STATUSES:
            logger.warning(f"无效的任务状态: {status}，默认设为 pending")
            status = "pending"

        # 如果任务文件已存在，合并已有数据
        existing = self._load_task_file(task_id)
        now = timestamp()

        task_data = {
            "task_id": task_id,
            "description": description,
            "session_id": session_id,
            "group_id": group_id,
            "agent_path": agent_path,
            "status": status,
            "created_at": existing.get("created_at", now) if existing else now,
            "updated_at": now,
            "last_message": last_message,
            "metadata": {
                **(existing.get("metadata", {}) if existing else {}),
                **(metadata or {}),
            },
        }

        file_path = self._task_file(task_id)
        try:
            file_path.write_text(
                json.dumps(task_data, ensure_ascii=False, indent=2),
                encoding="utf-8",
            )
            logger.debug(f"任务已保存: {task_id} (status={status})")
        except Exception as e:
            logger.error(f"保存任务失败 ({task_id}): {e}")

        return task_data

    def update_task_status(
        self,
        task_id: str,
        status: str,
        metadata: Optional[Dict[str, Any]] = None,
        last_message: str = "",
    ) -> bool:
        """
        更新任务状态。

        Args:
            task_id: 任务 ID
            status: 新状态
            metadata: 要合并的元数据
            last_message: 最后一条消息

        Returns:
            是否更新成功
        """
        if not self._tasks_dir:
            self.initialize()

        existing = self._load_task_file(task_id)
        if not existing:
            logger.warning(f"任务不存在，无法更新状态: {task_id}")
            return False

        existing["status"] = status
        existing["updated_at"] = timestamp()
        if metadata:
            existing["metadata"] = {**existing.get("metadata", {}), **metadata}
        if last_message:
            existing["last_message"] = last_message

        file_path = self._task_file(task_id)
        try:
            file_path.write_text(
                json.dumps(existing, ensure_ascii=False, indent=2),
                encoding="utf-8",
            )
            logger.debug(f"任务状态已更新: {task_id} -> {status}")
            return True
        except Exception as e:
            logger.error(f"更新任务状态失败 ({task_id}): {e}")
            return False

    def get_task(self, task_id: str) -> Optional[Dict[str, Any]]:
        """
        获取单个任务数据。

        Args:
            task_id: 任务 ID

        Returns:
            任务数据字典，不存在返回 None
        """
        if not self._tasks_dir:
            self.initialize()
        return self._load_task_file(task_id)

    def get_pending_tasks(self) -> List[Dict[str, Any]]:
        """
        获取所有未完成的任务 (pending + running)。

        Returns:
            未完成任务列表，按更新时间倒序
        """
        return self.get_all_tasks(status_filter=("pending", "running"))

    def get_all_tasks(
        self,
        status_filter: Optional[tuple] = None,
    ) -> List[Dict[str, Any]]:
        """
        获取所有任务。

        Args:
            status_filter: 可选状态过滤元组，如 ("pending", "running")

        Returns:
            任务列表，按更新时间倒序
        """
        if not self._tasks_dir:
            self.initialize()

        tasks = []
        if not self._tasks_dir:
            return tasks

        for f in self._tasks_dir.iterdir():
            if not f.is_file() or not f.name.startswith("task_") or not f.name.endswith(".json"):
                continue
            try:
                task = json.loads(f.read_text(encoding="utf-8"))
                if status_filter and task.get("status") not in status_filter:
                    continue
                tasks.append(task)
            except (json.JSONDecodeError, IOError) as e:
                logger.warning(f"读取任务文件失败 ({f.name}): {e}")

        # 按更新时间倒序
        tasks.sort(key=lambda t: t.get("updated_at", ""), reverse=True)
        return tasks

    def delete_task(self, task_id: str) -> bool:
        """
        删除任务记录。

        Args:
            task_id: 任务 ID

        Returns:
            是否删除成功
        """
        if not self._tasks_dir:
            self.initialize()

        file_path = self._task_file(task_id)
        if not file_path.exists():
            logger.warning(f"任务文件不存在: {task_id}")
            return False

        try:
            file_path.unlink()
            logger.debug(f"任务已删除: {task_id}")
            return True
        except Exception as e:
            logger.error(f"删除任务失败 ({task_id}): {e}")
            return False

    def cleanup_old_tasks(self, days: int = 7) -> int:
        """
        清理旧的已完成/失败任务。

        Args:
            days: 保留天数（超过此天数的 completed/failed 任务将被删除）

        Returns:
            清理的任务数量
        """
        if not self._tasks_dir:
            self.initialize()

        cutoff = datetime.now(timezone.utc) - timedelta(days=days)
        deleted_count = 0

        for f in self._tasks_dir.iterdir():
            if not f.is_file() or not f.name.startswith("task_") or not f.name.endswith(".json"):
                continue
            try:
                task = json.loads(f.read_text(encoding="utf-8"))
                status = task.get("status", "")
                # 只清理已完成或失败的任务
                if status not in ("completed", "failed"):
                    continue
                updated_at_str = task.get("updated_at", "")
                if not updated_at_str:
                    continue
                try:
                    updated_at = datetime.fromisoformat(updated_at_str)
                    # 如果没有时区信息，假定为 UTC
                    if updated_at.tzinfo is None:
                        updated_at = updated_at.replace(tzinfo=timezone.utc)
                    if updated_at < cutoff:
                        f.unlink()
                        deleted_count += 1
                except (ValueError, TypeError):
                    continue
            except (json.JSONDecodeError, IOError) as e:
                logger.warning(f"读取任务文件失败 ({f.name}): {e}")

        if deleted_count > 0:
            logger.info(f"已清理 {deleted_count} 个旧任务 (保留天数={days})")

        return deleted_count

    def recover_interrupted_tasks(self) -> List[Dict[str, Any]]:
        """
        恢复被中断的任务：将所有 "running" 状态的任务标记为 "failed"。

        在服务启动时调用，检测上次进程非正常退出时遗留的 running 任务。

        Returns:
            被恢复的任务列表
        """
        if not self._tasks_dir:
            self.initialize()

        running_tasks = self.get_all_tasks(status_filter=("running",))
        recovered = []

        for task in running_tasks:
            task_id = task.get("task_id", "")
            self.update_task_status(
                task_id,
                "failed",
                metadata={"interrupted": True, "interrupt_reason": "进程中断"},
                last_message=task.get("last_message", ""),
            )
            recovered.append(task)
            logger.warning(
                f"检测到中断任务: {task_id} "
                f"(description={task.get('description', '')[:50]}...)"
            )

        if recovered:
            logger.warning(f"共恢复 {len(recovered)} 个中断任务，已标记为 failed")

        return recovered

    def _load_task_file(self, task_id: str) -> Optional[Dict[str, Any]]:
        """加载单个任务 JSON 文件"""
        if not self._tasks_dir:
            return None

        file_path = self._task_file(task_id)
        if not file_path.exists():
            return None

        try:
            return json.loads(file_path.read_text(encoding="utf-8"))
        except (json.JSONDecodeError, IOError) as e:
            logger.warning(f"读取任务文件失败 ({task_id}): {e}")
            return None
