# -*- coding: utf-8 -*- """ 上下文容量追踪模块 功能: - 基于 tiktoken 估算各组件 token 用量 - 从 LLM 响应中提取实际 token 使用 - 将会话上下文用量写入数据库 分类维度: messages / mcp / skills / system_prompt / other """ import tiktoken from db.chat_store import save_context_usage, get_messages as db_get_messages # ── 编码器(cl100k_base 兼容 OpenAI / DeepSeek)── ENCODING = tiktoken.get_encoding("cl100k_base") # ── 模型上下文窗口上限 ── # 优先级:环境变量 CONTEXT_LIMIT → 模型名推断 → 兜底 131072 _MODEL_CONTEXT_MAP = { "deepseek-v4-pro": 1000000, # 1M "deepseek-v4-flash": 1000000, # 1M "deepseek-v3": 131072, # 128K "gpt-4o": 131072, # 128K "gpt-4o-mini": 131072, # 128K "qwen3.7-plus": 131072, # 128K } def _resolve_context_limit() -> int: """解析上下文上限(每次调用时重新解析,确保读到最新环境变量)。 优先级:环境变量 CONTEXT_LIMIT → 模型名映射 → 兜底 1M。 """ # 直接从 .env 文件读取(不受 CWD / load_dotenv 时序影响) from pathlib import Path as _Path from dotenv import dotenv_values _env = dotenv_values(str(_Path(__file__).parent.parent / ".env")) env_val = _env.get("CONTEXT_LIMIT", "").strip() if env_val: try: return int(env_val) except ValueError: pass model = _env.get("DEEPAGENT_MODEL", "").strip().lower() # 去掉可能的 "openai:" 前缀 if ":" in model: model = model.split(":", 1)[1] return _MODEL_CONTEXT_MAP.get(model, 1000000) def get_context_limit() -> int: """获取当前模型上下文窗口上限(懒加载)。""" return _resolve_context_limit() # ============================================================ # 工具 schema 提取辅助 # ============================================================ def _tool_to_text(tool) -> str: """将工具函数转为可计数的文本表示(含名称 + 描述 + 参数 schema)。""" parts = [] name = getattr(tool, "name", None) or getattr(tool, "__name__", str(tool)) desc = getattr(tool, "description", "") or "" parts.append(f"Tool: {name}") if desc: parts.append(f"Description: {desc}") # 尝试提取参数 schema args_schema = getattr(tool, "args_schema", None) if args_schema and hasattr(args_schema, "schema"): try: import json parts.append(f"Args: {json.dumps(args_schema.schema(), ensure_ascii=False)}") except Exception: pass return "\n".join(parts) # ============================================================ # 全局缓存 —— 系统组件 token 预估值(agent 启动时填充) # ============================================================ _system_prompt_tokens = 0 _skills_tokens = 0 _mcp_tokens = 0 def init_system_components( system_prompt: str = "", tool_defs: list[str] | None = None, skill_contents: list[str] | None = None, mcp_defs: list[str] | None = None, ): """初始化系统组件的 token 预估值。应在 agent 创建后调用一次。 mcp 类别 = 本地工具 + MCP 远程工具(合并统计)。 """ global _system_prompt_tokens, _skills_tokens, _mcp_tokens _system_prompt_tokens = count_tokens(system_prompt) _skills_tokens = sum(count_tokens(s) for s in (skill_contents or [])) # tools + mcp 合并为 mcp 类别 _mcp_tokens = ( sum(count_tokens(t) for t in (tool_defs or [])) + sum(count_tokens(m) for m in (mcp_defs or [])) ) # ============================================================ # Token 计数工具 # ============================================================ def count_tokens(text: str) -> int: """使用 tiktoken 精确计数 token 数。""" if not text: return 0 try: return len(ENCODING.encode(text)) except Exception: # 兜底估算:中文 ~1.5 字/token,英文 ~4 字/token return max(1, len(text) // 2) def _count_messages_tokens(session_id: str, current_message: str = "") -> int: """计算会话消息历史的 token 数。""" total = 0 try: for msg in db_get_messages(session_id, limit=1000): total += count_tokens(msg.get("content", "")) except Exception: pass total += count_tokens(current_message) return total # ============================================================ # 主入口:捕获并存储上下文用量 # ============================================================ async def capture_context_usage( session_id: str, prompt_tokens: int, completion_tokens: int, current_message: str = "", ): """在每次 agent 任务完成后调用,保存上下文用量快照。 Args: session_id: 会话 ID prompt_tokens: 从 LLM 响应中提取的实际 prompt_tokens completion_tokens: 从 LLM 响应中提取的实际 completion_tokens current_message: 本次用户消息(用于计入 messages 估算) """ if not session_id or prompt_tokens <= 0: return messages_tokens = _count_messages_tokens(session_id, current_message) # 系统组件 token(使用全局缓存估值) # mcp = 本地工具 + MCP 远程工具(已合并) skills_tokens = _skills_tokens mcp_tokens = _mcp_tokens system_prompt_tokens = _system_prompt_tokens # other = 提示词总量 - 已知各组件估算值(兜底 ≥0) known = messages_tokens + mcp_tokens + skills_tokens + system_prompt_tokens other_tokens = max(0, prompt_tokens - known) current_usage = prompt_tokens + completion_tokens try: save_context_usage( session_id=session_id, total_limit=get_context_limit(), current_usage=current_usage, messages_tokens=messages_tokens, mcp_tokens=mcp_tokens, skills_tokens=skills_tokens, system_prompt_tokens=system_prompt_tokens, other_tokens=other_tokens, ) except Exception: pass # 上下文记录失败不应阻断对话