"""
core/context_manager.py - 上下文管理器 (Token优化)
====================================================
业内最佳实践: Anthropic Context Engineering + LangChain SummaryBufferMemory

核心策略:
  1. Token 预算管理 — 系统提示10%、工具定义20%、记忆15%、历史40%、输出15%
  2. 滚动摘要缓冲 — 近期N条消息全量 + 历史增量摘要
  3. 工具输出压缩 — 仅保留关键结果，截断冗余输出
  4. 渐进式裁剪 — 超预算时按优先级逐层裁剪

目标: 将上下文利用率控制在模型窗口的60%以内，避免"context rot"
"""
from __future__ import annotations

import copy
import json
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple

from core.logger import get_logger
from core.utils import truncate_str

logger = get_logger("myagent.context")


# ==============================================================================
# 配置
# ==============================================================================

@dataclass
class ContextBudget:
    """Token 预算分配 (百分比)

    [FIX-失忆] 调整预算分配：
    - 增大 recent_messages (0.35→0.40): 对话历史是防止失忆的关键
    - 增大 history_summary (0.10→0.12): 滚动摘要提供更完整的上下文
    - 减小 tools_schema (0.15→0.10): 工具定义不需要那么多空间
    - 减小 safety_margin (0.10→0.08): 释放一些余量
    """
    system_prompt: float = 0.10    # 系统提示
    tools_schema: float = 0.10     # 工具定义 (JSON schema)
    memory_context: float = 0.15   # 记忆上下文 (偏好/经验/错误)
    history_summary: float = 0.12  # 历史摘要 (滚动摘要)
    recent_messages: float = 0.40  # 近期消息 (原始) - 防失忆核心
    user_message: float = 0.05     # 当前用户消息
    safety_margin: float = 0.08    # 安全余量

    @property
    def total_history(self) -> float:
        """历史部分总计 (摘要 + 近期)"""
        return self.history_summary + self.recent_messages


@dataclass
class ContextConfig:
    """上下文管理配置

    [FIX-失忆] 增加近期消息保留量和字符预算：
    - recent_message_count: 10 → 20 (更多轮次保留)
    - recent_message_max_chars: 8000 → 20000 (更大字符空间)
    - summary_max_chars: 3000 → 6000 (更完整的摘要)
    - max_message_chars: 每条消息最大字符数，超出截断 (默认 10000)
    """
    # 模型上下文窗口大小 (token数), 默认200K
    model_context_window: int = 200000
    # 预算分配
    budget: ContextBudget = field(default_factory=ContextBudget)
    # 近期消息保留条数 (原始消息, 不压缩)
    recent_message_count: int = 20
    # 近期消息最大字符数 (约等于token数, 中文1字符≈1token)
    recent_message_max_chars: int = 20000
    # 历史摘要最大字符数
    summary_max_chars: int = 6000
    # 单条工具输出最大字符数
    tool_output_max_chars: int = 2000
    # 记忆上下文最大字符数
    memory_context_max_chars: int = 3000
    # 每条消息最大字符数（超出截断）
    max_message_chars: int = 10000
    # 启用Token预算管理
    enabled: bool = True

    def get_budget_chars(self, category: str) -> int:
        """获取某类别的字符预算"""
        ratio = getattr(self.budget, category, 0)
        return int(self.model_context_window * ratio)


# ==============================================================================
# 文本压缩工具
# ==============================================================================

def compress_tool_output(content: str, max_chars: int = 2000) -> str:
    """
    压缩工具输出，仅保留关键结果。

    策略:
      - JSON: 保留顶层结构和关键字段
      - 列表: 保留前N项
      - 长文本: 截取首尾 + 省略中间
    """
    if not content or len(content) <= max_chars:
        return content

    # 尝试解析为 JSON
    try:
        data = json.loads(content)
        return _compress_json(data, max_chars)
    except (json.JSONDecodeError, TypeError):
        pass

    # 纯文本: 保留首尾
    half = max_chars // 2 - 50
    return content[:half] + f"\n... [省略 {len(content) - max_chars} 字符] ...\n" + content[-half:]


def _compress_json(data: Any, max_chars: int, depth: int = 0) -> str:
    """递归压缩JSON结构"""
    if depth > 3:
        return truncate_str(json.dumps(data, ensure_ascii=False), max_chars)

    if isinstance(data, dict):
        # 优先保留关键字段
        priority_keys = ["success", "error", "message", "result", "output", "summary",
                         "key_points", "status", "content", "url", "title"]
        result = {}
        for k in priority_keys:
            if k in data:
                val = data[k]
                if isinstance(val, str) and len(val) > 500:
                    val = val[:500] + "..."
                result[k] = val
        # 补充其他字段 (截断)
        remaining_chars = max_chars - len(json.dumps(result, ensure_ascii=False))
        if remaining_chars > 200:
            for k, v in data.items():
                if k not in result:
                    v_str = json.dumps(v, ensure_ascii=False) if not isinstance(v, str) else v
                    if len(v_str) > 300:
                        v_str = v_str[:300] + "..."
                    result[k] = json.loads(v_str) if not isinstance(v, str) else v_str
                    remaining_chars -= len(k) + len(v_str) + 10
                    if remaining_chars < 100:
                        break
        return json.dumps(result, ensure_ascii=False, indent=None)

    elif isinstance(data, list):
        # 列表: 保留前5项 + 总数
        if len(data) <= 5:
            return json.dumps(data, ensure_ascii=False)
        truncated = data[:5]
        result = json.dumps(truncated, ensure_ascii=False)
        if len(result) > max_chars:
            result = truncate_str(result, max_chars)
        result = result.rstrip("]") + f", ... (共 {len(data)} 项)]"
        return result

    return truncate_str(json.dumps(data, ensure_ascii=False), max_chars)


def estimate_tokens(text: str) -> int:
    """
    粗略估算token数量。
    中文: 1字符 ≈ 1~1.5 token
    英文: 1单词 ≈ 1.3 token (平均)
    JSON/代码: 1字符 ≈ 0.3 token
    """
    if not text:
        return 0

    chinese_count = sum(1 for c in text if '\u4e00' <= c <= '\u9fff')
    other_chars = len(text) - chinese_count

    # 中文字符每个约1.3token，其他字符约0.35token
    return int(chinese_count * 1.3 + other_chars * 0.35)


# ==============================================================================
# 滚动摘要缓冲
# ==============================================================================

@dataclass
class RollingSummary:
    """滚动摘要状态 (per session)"""
    session_id: str = ""
    summary: str = ""                   # 当前滚动摘要
    last_summarized_index: int = 0      # 已摘要到的消息索引
    total_summarized: int = 0           # 已摘要的总消息数
    version: int = 0                    # 摘要版本号

    def to_dict(self) -> dict:
        return {
            "session_id": self.session_id,
            "summary": self.summary,
            "last_summarized_index": self.last_summarized_index,
            "total_summarized": self.total_summarized,
            "version": self.version,
        }

    @classmethod
    def from_dict(cls, data: dict) -> "RollingSummary":
        return cls(
            session_id=data.get("session_id", ""),
            summary=data.get("summary", ""),
            last_summarized_index=data.get("last_summarized_index", 0),
            total_summarized=data.get("total_summarized", 0),
            version=data.get("version", 0),
        )


class RollingSummaryBuffer:
    """
    滚动摘要缓冲区。

    维护每个session的滚动摘要:
      - 近期N条消息保持原始
      - 更早的消息通过摘要保留
      - 每次新消息到来时，增量更新摘要

    与传统全量注入的区别:
      传统: 发送全部20条历史 → 约20K+ tokens
      摘要缓冲: 发送摘要3K + 近期10条8K → 约11K tokens
      节省: ~45%
    """

    def __init__(self, config: Optional[ContextConfig] = None):
        self.config = config or ContextConfig()
        self._summaries: Dict[str, RollingSummary] = {}

    def get_summary(self, session_id: str) -> RollingSummary:
        """获取session的滚动摘要"""
        if session_id not in self._summaries:
            self._summaries[session_id] = RollingSummary(session_id=session_id)
        return self._summaries[session_id]

    def update_summary(self, session_id: str, new_summary: str,
                       summarized_count: int) -> RollingSummary:
        """
        更新滚动摘要 (增量)。

        Args:
            session_id: 会话ID
            new_summary: 新的摘要文本 (可以是增量更新或全量替换)
            summarized_count: 本次摘要涵盖的消息数
        """
        state = self.get_summary(session_id)
        state.last_summarized_index += summarized_count
        state.total_summarized += summarized_count
        state.version += 1

        # 智能合并: 如果新摘要比旧的短很多，追加；否则替换
        if state.summary and len(new_summary) < len(state.summary) * 0.3:
            state.summary = state.summary + "\n" + new_summary
        else:
            state.summary = new_summary

        # 摘要本身也不宜太长
        if len(state.summary) > self.config.summary_max_chars * 2:
            state.summary = state.summary[-self.config.summary_max_chars:]

        return state

    def reset(self, session_id: str):
        """重置session的摘要"""
        if session_id in self._summaries:
            self._summaries[session_id] = RollingSummary(session_id=session_id)

    def get_summary_text(self, session_id: str) -> str:
        """获取格式化的摘要文本"""
        state = self.get_summary(session_id)
        if not state.summary:
            return ""
        return (
            f"[历史对话摘要 v{state.version}, 已总结 {state.total_summarized} 条消息]\n"
            f"{state.summary}"
        )

    def clear(self):
        """清空所有摘要"""
        self._summaries.clear()


# ==============================================================================
# 上下文管理器 (核心)
# ==============================================================================

class ContextManager:
    """
    上下文管理器 — 每次API调用前策划最优的消息列表。

    职责:
      1. 计算Token预算
      2. 构建系统提示 (带预算裁剪)
      3. 注入记忆上下文 (带预算裁剪)
      4. 管理滚动摘要缓冲
      5. 选择近期消息 (带工具输出压缩)
      6. 渐进式裁剪 (超预算时逐层削减)

    使用方式:
        cm = ContextManager(config)
        messages = cm.build_context(
            system_prompt="...",
            memory_context="...",
            conversation_history=[...],
            user_message="...",
            tools_schema=[...],
        )
    """

    def __init__(self, config: Optional[ContextConfig] = None):
        self.config = config or ContextConfig()
        self.summary_buffer = RollingSummaryBuffer(self.config)
        self._last_build_stats: Dict[str, Any] = {}

    def build_context(
        self,
        system_prompt: str = "",
        memory_context: str = "",
        conversation_history: Optional[list] = None,
        user_message: str = "",
        session_id: str = "",
        tools_schema: Optional[list] = None,
    ) -> list:
        """
        构建最优上下文消息列表。

        Returns:
            list of Message/Dict — 可直接发给LLM的消息列表
        """
        history = conversation_history or []

        if not self.config.enabled:
            # 未启用预算管理，使用原始逻辑 (兼容旧行为)
            messages = []
            if system_prompt:
                messages.append({"role": "system", "content": system_prompt})
            if memory_context:
                messages.append({"role": "system", "content": f"[记忆上下文]\n{memory_context}"})
            messages.extend(
                m.to_dict() if hasattr(m, 'to_dict') else m
                for m in history[-20:]
            )
            messages.append({"role": "user", "content": user_message})
            return messages

        # ── Step 1: 计算各部分预算 ──
        budget = self.config
        total_chars = budget.model_context_window * 0.6  # 使用60%窗口

        system_budget = int(total_chars * budget.budget.system_prompt)
        tools_budget = int(total_chars * budget.budget.tools_schema)
        memory_budget = int(total_chars * budget.budget.memory_context)
        summary_budget = int(total_chars * budget.budget.history_summary)
        recent_budget = int(total_chars * budget.budget.recent_messages)

        stats = {
            "total_budget": total_chars,
            "system_chars": 0,
            "tools_chars": 0,
            "memory_chars": 0,
            "summary_chars": 0,
            "recent_chars": 0,
            "recent_count": 0,
            "compressed_count": 0,
        }

        messages = []

        # ── Step 2: 系统提示 (可裁剪) ──
        if system_prompt:
            trimmed_system = system_prompt
            if len(trimmed_system) > system_budget:
                trimmed_system = self._trim_system_prompt(system_prompt, system_budget)
            messages.append({"role": "system", "content": trimmed_system})
            stats["system_chars"] = len(trimmed_system)

        # ── Step 3: 记忆上下文 (可裁剪) ──
        if memory_context:
            trimmed_memory = memory_context
            if len(trimmed_memory) > memory_budget:
                trimmed_memory = memory_context[:memory_budget] + "\n... [记忆上下文已裁剪]"
            messages.append({"role": "system", "content": f"[记忆上下文]\n{trimmed_memory}"})
            stats["memory_chars"] = len(trimmed_memory)

        # ── Step 4: 历史摘要 ──
        summary_text = self.summary_buffer.get_summary_text(session_id)
        if summary_text and session_id:
            if len(summary_text) > summary_budget:
                summary_text = summary_text[:summary_budget] + "\n... [摘要已裁剪]"
            messages.append({"role": "system", "content": summary_text})
            stats["summary_chars"] = len(summary_text)

        # ── Step 5: 近期消息 (带压缩) ──
        remaining_budget = (
            total_chars
            - stats["system_chars"]
            - stats["memory_chars"]
            - stats["summary_chars"]
            - len(user_message)
            - 500  # 安全余量
        )
        remaining_budget = max(remaining_budget, recent_budget // 2)

        recent_msgs = self._select_recent_messages(
            history,
            max_count=budget.recent_message_count,
            max_chars=remaining_budget,
        )
        for msg in recent_msgs:
            msg_dict = msg.to_dict() if hasattr(msg, 'to_dict') else msg
            messages.append(msg_dict)

        stats["recent_chars"] = sum(
            len(m.content if hasattr(m, 'content') else m.get('content', ''))
            for m in recent_msgs
        )
        stats["recent_count"] = len(recent_msgs)

        # ── Step 6: 当前用户消息 ──
        messages.append({"role": "user", "content": user_message})

        self._last_build_stats = stats
        logger.debug(
            f"上下文构建完成: system={stats['system_chars']}, "
            f"memory={stats['memory_chars']}, summary={stats['summary_chars']}, "
            f"recent={stats['recent_count']}条/{stats['recent_chars']}字, "
            f"budget={total_chars}"
        )

        return messages

    def _select_recent_messages(
        self,
        messages: list,
        max_count: int = 10,
        max_chars: int = 8000,
    ) -> list:
        """
        选择并压缩近期消息。

        策略:
          1. 取最后 max_count 条
          2. 工具输出压缩 (超过阈值截断)
          3. 总字符超预算时，从最早的消息开始裁剪
        """
        if not messages:
            return []

        recent = list(messages[-max_count:])
        compressed_count = 0

        # 压缩工具输出（使用副本，避免修改调用方的原始消息对象）
        for i, msg in enumerate(recent):
            content = msg.content if hasattr(msg, 'content') else msg.get('content', '')
            role = msg.role if hasattr(msg, 'role') else msg.get('role', '')

            new_content = None

            # 工具消息和assistant消息中的执行结果需要压缩
            if role in ("tool", "user") and len(content) > self.config.tool_output_max_chars:
                # 检测是否是执行结果
                if any(marker in content for marker in
                       ["[执行结果]", "tool_result", "```output", "[工具调用结果]"]):
                    compressed = compress_tool_output(content, self.config.tool_output_max_chars)
                    if len(compressed) < len(content):
                        compressed_count += 1
                        new_content = compressed

            # assistant消息也压缩超长内容
            if role == "assistant" and len(content) > self.config.max_message_chars:
                new_content = content[:self.config.max_message_chars] + "\n... [输出已压缩]"
                compressed_count += 1

            # 写入副本而非修改原始对象
            if new_content is not None:
                if isinstance(msg, dict):
                    recent[i] = dict(msg)
                    recent[i]['content'] = new_content
                elif hasattr(msg, '__dataclass_fields__'):
                    import dataclasses
                    recent[i] = dataclasses.replace(msg, content=new_content)
                elif hasattr(msg, 'content'):
                    # 通用对象: 创建浅拷贝
                    try:
                        recent[i] = copy.copy(msg)
                        recent[i].content = new_content
                    except Exception:
                        msg.content = new_content

        # 预算裁剪: 从最早的消息开始移除
        total = sum(
            len(m.content if hasattr(m, 'content') else m.get('content', ''))
            for m in recent
        )
        while total > max_chars and len(recent) > 3:
            removed = recent.pop(0)
            removed_len = len(removed.content if hasattr(removed, 'content') else removed.get('content', ''))
            total -= removed_len
            compressed_count += 1

        return recent

    def _trim_system_prompt(self, prompt: str, budget: int) -> str:
        """
        裁剪系统提示。

        策略: 保留前缀 (核心指令) + 精简格式要求
        """
        if len(prompt) <= budget:
            return prompt

        # 按行分割，保留前半部分 (核心指令通常在前)
        lines = prompt.split('\n')
        total = 0
        kept = []
        for line in lines:
            if total + len(line) + 1 > budget * 0.8:
                break
            kept.append(line)
            total += len(line) + 1

        result = '\n'.join(kept)
        if len(result) < len(prompt):
            result += f"\n... [系统提示已精简, 完整版约 {len(prompt)} 字符]"
        return result

    def get_build_stats(self) -> Dict[str, Any]:
        """获取上次构建的统计信息"""
        return dict(self._last_build_stats)

    def get_summary_buffer(self) -> RollingSummaryBuffer:
        """获取滚动摘要缓冲区 (供外部更新)"""
        return self.summary_buffer
