"""
agents/base.py - Agent 基类
============================
定义所有 Agent 的基础接口和通用能力。
"""
from __future__ import annotations

import asyncio
import json
import time
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Callable

from core.logger import get_logger
from core.llm import LLMClient, LLMResponse, Message
from core.utils import generate_id, timestamp

logger = get_logger("myagent.agent")


# 全局权限管理器实例（由 main.py 初始化）
_global_permission_manager = None


def set_permission_manager(pm):
    """设置全局权限管理器（应用启动时调用一次）"""
    global _global_permission_manager
    _global_permission_manager = pm


def get_permission_manager():
    """获取全局权限管理器"""
    return _global_permission_manager


@dataclass
class AgentContext:
    """Agent 上下文，在 Agent 之间传递"""
    task_id: str = ""
    session_id: str = "default"
    user_message: str = ""
    conversation_history: List[Message] = field(default_factory=list)
    working_memory: Dict[str, Any] = field(default_factory=dict)
    metadata: Dict[str, Any] = field(default_factory=dict)
    callbacks: Dict[str, Callable] = field(default_factory=dict)
    pending_injected_messages: List[str] = field(default_factory=list)


class BaseAgent(ABC):
    """
    Agent 基类。

    所有 Agent 继承此类，实现 process 方法。
    提供 LLM 调用、消息构建、错误处理等通用能力。
    """

    name: str = "base_agent"
    description: str = "基础 Agent"
    max_iterations: int = 30

    def __init__(
        self,
        llm: Optional[LLMClient] = None,
        memory_manager=None,
        executor=None,
        skill_registry=None,
        task_queue=None,
        config=None,
        config_broadcaster=None,
    ):
        self.llm = llm
        self.memory = memory_manager
        self.executor = executor
        self.skills = skill_registry
        self.task_queue = task_queue
        self.config = config or {}
        self.config_broadcaster = config_broadcaster
        self._permission_manager = None  # 延迟绑定，在 process 时获取
        self._stats = {"total_tasks": 0, "success": 0, "failed": 0}

    @abstractmethod
    async def process(self, context: AgentContext) -> AgentContext:
        """
        处理任务(子类必须实现)。

        Args:
            context: Agent 上下文

        Returns:
            更新后的 AgentContext
        """
        pass

    def _build_system_prompt(self) -> str:
        """构建系统提示词(子类可重写)"""
        return f"你是 {self.name}，{self.description}。"

    async def _call_llm(
        self,
        messages: List[Message],
        tools: Optional[List[Dict]] = None,
        **kwargs,
    ) -> LLMResponse:
        """调用 LLM"""
        if not self.llm:
            return LLMResponse(success=False, error="LLM 客户端未初始化")

        response = await self.llm.chat(messages, tools=tools, **kwargs)
        if not response.success:
            logger.error(f"{self.name} LLM 调用失败: {response.error}")
        return response

    async def _call_llm_stream(self, messages, tools=None, stream_response=None, text_delta_callback=None, reasoning_delta_callback=None, **kwargs):
        """调用LLM并流式输出token到SSE response

        当 stream_response 提供时，逐 token 将内容写入 SSE 流。
        同时积累 tool_call 增量，在流结束时返回完整的 LLMResponse。

        Args:
            text_delta_callback: 可选的回调函数 async (full_text_so_far, delta_text) -> None
                当提供时，不再自动发送 text_delta SSE 事件，而是调用此回调。
                回调可以自行决定如何处理文本增量（如过滤 JSON、提取 thought 等）。
            reasoning_delta_callback: 可选的推理增量回调 async (full_reasoning_so_far, delta_text) -> None
                专门用于处理模型的推理内容（thinking/reasoning tokens）。
        """
        if not self.llm:
            return LLMResponse(success=False, error="LLM 未初始化")

        # If no stream_response, fall back to non-streaming
        if not stream_response:
            return await self._call_llm(messages, tools=tools, **kwargs)

        import asyncio as _asyncio

        self.llm._ensure_client()
        msg_dicts = [m.to_dict() if hasattr(m, 'to_dict') else m for m in messages]
        request_kwargs = {
            "model": self.llm.model,
            "messages": msg_dicts,
            "temperature": self.llm.temperature,
            "max_tokens": self.llm.max_tokens,
            "stream": True,
        }
        if tools:
            request_kwargs["tools"] = tools
            request_kwargs["tool_choice"] = "auto"
        # ── 推理模式 (Reasoning) ──
        if self.llm.reasoning:
            if self.llm.provider in self.llm._OPENAI_COMPATIBLE_PROVIDERS or self.llm.provider == "zhipu":
                request_kwargs["reasoning_effort"] = self.llm.reasoning_effort
                if self.llm.max_tokens < 8192:
                    request_kwargs["max_tokens"] = 8192
        request_kwargs.update(kwargs)

        full_text = ""
        full_reasoning = ""  # Track reasoning tokens (for o1/o3/DeepSeek-R1 etc.)
        tool_calls_acc: Dict[int, Dict] = {}  # index -> {id, name, arguments_str}
        finish_reason = ""

        async def _write_sse(data: dict):
            """将一个事件写入 SSE 流，忽略客户端断开错误"""
            try:
                await stream_response.write(
                    ("data: " + json.dumps(data, ensure_ascii=False) + "\n\n").encode()
                )
            except Exception:
                pass  # Client disconnected

        async def _emit_text_delta(delta_text: str):
            """处理一个 text delta：如果提供了回调则调用回调，否则直接发送 SSE"""
            nonlocal full_text
            full_text += delta_text
            if text_delta_callback:
                await text_delta_callback(full_text, delta_text)
            else:
                await _write_sse({"type": "text_delta", "content": delta_text})

        async def _emit_reasoning_delta(delta_text: str):
            """处理推理 token：发送 reasoning_delta SSE 事件"""
            nonlocal full_reasoning
            full_reasoning += delta_text
            await _write_sse({"type": "reasoning_delta", "content": delta_text})
            # 同时调用 reasoning_delta_callback（如果提供）
            if reasoning_delta_callback:
                await reasoning_delta_callback(full_reasoning, delta_text)

        try:
            if self.llm.provider in self.llm._OPENAI_COMPATIBLE_PROVIDERS or self.llm.provider == "zhipu":
                # 使用异步客户端流式
                stream = await self.llm._client.chat.completions.create(**request_kwargs)
                
                async for chunk in stream:
                    if not chunk.choices:
                        if hasattr(chunk, 'usage') and chunk.usage:
                            self.llm._record_usage(
                                {"prompt_tokens": chunk.usage.prompt_tokens,
                                 "completion_tokens": chunk.usage.completion_tokens,
                                 "total_tokens": chunk.usage.total_tokens},
                                request_kwargs["model"],
                            )
                        continue

                    delta = chunk.choices[0].delta
                    if chunk.choices[0].finish_reason:
                        finish_reason = chunk.choices[0].finish_reason

                    # Handle content delta (stream to client)
                    if delta.content:
                        await _emit_text_delta(delta.content)

                    # Handle reasoning_content (OpenAI o1/o3, DeepSeek-R1, Qwen-QwQ etc.)
                    reasoning_content = getattr(delta, 'reasoning_content', None) or getattr(delta, 'reasoning', None)
                    if reasoning_content:
                        await _emit_reasoning_delta(reasoning_content)

                    # Handle tool_call deltas (accumulate)
                    if hasattr(delta, 'tool_calls') and delta.tool_calls:
                        for tc_delta in delta.tool_calls:
                            idx = tc_delta.index if hasattr(tc_delta, 'index') else 0
                            if idx not in tool_calls_acc:
                                tool_calls_acc[idx] = {"id": "", "name": "", "arguments": ""}
                            if tc_delta.id:
                                tool_calls_acc[idx]["id"] = tc_delta.id
                            if hasattr(tc_delta, 'function') and tc_delta.function:
                                if tc_delta.function.name:
                                    tool_calls_acc[idx]["name"] = tc_delta.function.name
                                if tc_delta.function.arguments:
                                    tool_calls_acc[idx]["arguments"] += tc_delta.function.arguments

                    # Handle usage in final chunk
                    if hasattr(chunk, 'usage') and chunk.usage:
                        self.llm._record_usage(
                            {"prompt_tokens": chunk.usage.prompt_tokens,
                             "completion_tokens": chunk.usage.completion_tokens,
                             "total_tokens": chunk.usage.total_tokens},
                            request_kwargs["model"],
                        )

            elif self.llm.provider == "anthropic":
                loop = _asyncio.get_running_loop()

                system_msg = ""
                anth_messages = []
                for m in messages:
                    role = m.role if hasattr(m, 'role') else m.get("role", "user")
                    content = m.content if hasattr(m, 'content') else m.get("content", "")
                    if role == "system":
                        system_msg = content
                        continue
                    # 转换 OpenAI Vision 格式为 Anthropic 格式
                    anth_content = self.llm._convert_to_anthropic_content(content)
                    anth_messages.append({"role": role, "content": anth_content})

                create_kwargs = {
                    "model": self.llm.model,
                    "messages": anth_messages,
                    "max_tokens": self.llm.max_tokens,
                    "stream": True,
                }
                if system_msg:
                    create_kwargs["system"] = system_msg

                def _create_stream():
                    return self.llm._client.messages.create(**create_kwargs)

                stream = await loop.run_in_executor(None, _create_stream)

                def _next_event(it):
                    try:
                        return next(it)
                    except StopIteration:
                        return None

                iterator = iter(stream)
                while True:
                    event = await loop.run_in_executor(None, _next_event, iterator)
                    if event is None:
                        break
                    if event.type == "content_block_delta":
                        # Handle text content
                        if hasattr(event.delta, "text") and event.delta.text:
                            await _emit_text_delta(event.delta.text)
                        # Handle extended thinking (Anthropic reasoning)
                        if hasattr(event.delta, "thinking") and event.delta.thinking:
                            await _emit_reasoning_delta(event.delta.thinking)
                        # Handle thinking delta via type
                        if hasattr(event.delta, "type") and event.delta.type == "thinking_delta":
                            thinking_text = getattr(event.delta, "thinking", "")
                            if thinking_text:
                                await _emit_reasoning_delta(thinking_text)
                    elif event.type == "message_stop":
                        finish_reason = "stop"

            elif self.llm.provider == "ollama":
                loop = _asyncio.get_running_loop()
                import requests as req_lib

                url = f"{self.llm.base_url}/api/chat"
                payload = {
                    "model": self.llm.model,
                    "messages": msg_dicts,
                    "stream": True,
                    "options": {
                        "temperature": self.llm.temperature,
                        "num_predict": self.llm.max_tokens,
                    },
                }

                def _request():
                    r = req_lib.post(url, json=payload, stream=True, timeout=self.llm.timeout)
                    r.raise_for_status()
                    return r.iter_lines()

                lines_iter = await loop.run_in_executor(None, _request)

                def _next_line(it):
                    try:
                        return next(it)
                    except StopIteration:
                        return None

                iterator = iter(lines_iter)
                while True:
                    line = await loop.run_in_executor(None, _next_line, iterator)
                    if line is None:
                        break
                    try:
                        data = json.loads(line.decode('utf-8') if isinstance(line, bytes) else line)
                        message = data.get("message", {})
                        content = message.get("content", "")
                        if content:
                            await _emit_text_delta(content)
                        # Handle Ollama thinking/reasoning (e.g., DeepSeek-R1 via Ollama)
                        thinking = message.get("thinking", "")
                        if thinking:
                            await _emit_reasoning_delta(thinking)
                        if data.get("done"):
                            finish_reason = "stop"
                            # Record usage from Ollama
                            usage = data.get("prompt_eval_count") or data.get("eval_count")
                            if data.get("prompt_eval_count"):
                                self.llm._record_usage(
                                    {"prompt_tokens": data.get("prompt_eval_count", 0),
                                     "completion_tokens": data.get("eval_count", 0),
                                     "total_tokens": data.get("prompt_eval_count", 0) + data.get("eval_count", 0)},
                                    self.llm.model,
                                )
                    except Exception:
                        continue
            else:
                return LLMResponse(success=False, error="未知提供商，不支持流式")

            # Build tool_calls list from accumulated deltas
            final_tool_calls = []
            for idx in sorted(tool_calls_acc.keys()):
                tc = tool_calls_acc[idx]
                _raw_args = tc["arguments"] if tc["arguments"] else "{}"
                try:
                    args = json.loads(_raw_args) if _raw_args else {}
                except (json.JSONDecodeError, TypeError):
                    # [v1.39] JSON解析失败时保留原始字符串，不静默丢弃
                    logger.warning(f"streaming tool_call arguments JSON解析失败: {_raw_args[:200]}")
                    args = {"raw_input": _raw_args}
                final_tool_calls.append({
                    "id": tc["id"],
                    "name": tc["name"],
                    "arguments": args,
                })

            # 对于推理模型（如 o1/DeepSeek-R1），如果 content 为空但有 reasoning 内容，
            # 使用 reasoning 内容作为最终回复
            final_content = full_text if full_text.strip() else (full_reasoning if full_reasoning.strip() else full_text)

            return LLMResponse(
                success=True,
                content=final_content,
                tool_calls=final_tool_calls,
                finish_reason=finish_reason,
                model=request_kwargs.get("model", self.llm.model),
                reasoning=full_reasoning,
            )
        except asyncio.CancelledError:
            # [v1.23.38] 传播取消信号，不吞掉
            raise
        except Exception as e:
            logger.error(f"LLM 流式调用失败: {e}")
            return LLMResponse(success=False, error=str(e))

    async def _call_llm_json(self, messages: List[Message], **kwargs) -> Dict[str, Any]:
        """调用 LLM 并获取 JSON 响应"""
        if not self.llm:
            return {"error": "LLM 客户端未初始化"}
        return await self.llm.chat_json(messages, **kwargs)

    def _message(self, role: str, content: str, **kwargs) -> Message:
        """快捷创建消息"""
        return Message(role=role, content=content, **kwargs)

    def update_stats(self, success: bool):
        """更新统计"""
        self._stats["total_tasks"] += 1
        if success:
            self._stats["success"] += 1
        else:
            self._stats["failed"] += 1

    def get_stats(self) -> Dict:
        return dict(self._stats)

    # ── 权限检查 ────────────────────────────────────────────

    @property
    def permission_manager(self):
        """获取权限管理器（延迟绑定）"""
        if self._permission_manager is None:
            from agents.base import get_permission_manager
            self._permission_manager = get_permission_manager()
        return self._permission_manager

    def check_permission(self, permission: str) -> bool:
        """
        检查当前 agent 是否拥有某项权限。

        Args:
            permission: 权限项名称 (execution/file_read/file_write/network/local_comm/remote_comm)

        Returns:
            True 如果权限开启，或权限管理器未初始化（向后兼容）
        """
        pm = self.permission_manager
        if pm is None:
            # 权限管理器未初始化，允许所有操作（向后兼容）
            return True
        return pm.check_permission(self.name, permission)

    def require_permission(self, permission: str) -> bool:
        """
        要求指定权限，权限不足时记录警告并返回 False。

        Args:
            permission: 权限项名称

        Returns:
            True 如果权限通过
        """
        pm = self.permission_manager
        if pm is None:
            return True
        return pm.require_permission(self.name, permission)
