context_tracker.py 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191
  1. # -*- coding: utf-8 -*-
  2. """
  3. 上下文容量追踪模块
  4. 功能:
  5. - 基于 tiktoken 估算各组件 token 用量
  6. - 从 LLM 响应中提取实际 token 使用
  7. - 将会话上下文用量写入数据库
  8. 分类维度:
  9. messages / mcp / skills / system_prompt / other
  10. """
  11. import tiktoken
  12. from db.chat_store import save_context_usage, get_messages as db_get_messages
  13. # ── 编码器(cl100k_base 兼容 OpenAI / DeepSeek)──
  14. ENCODING = tiktoken.get_encoding("cl100k_base")
  15. # ── 模型上下文窗口上限 ──
  16. # 优先级:环境变量 CONTEXT_LIMIT → 模型名推断 → 兜底 131072
  17. _MODEL_CONTEXT_MAP = {
  18. "deepseek-v4-pro": 1000000, # 1M
  19. "deepseek-v4-flash": 1000000, # 1M
  20. "deepseek-v3": 131072, # 128K
  21. "gpt-4o": 131072, # 128K
  22. "gpt-4o-mini": 131072, # 128K
  23. "qwen3.7-plus": 131072, # 128K
  24. }
  25. def _resolve_context_limit() -> int:
  26. """解析上下文上限(每次调用时重新解析,确保读到最新环境变量)。
  27. 优先级:环境变量 CONTEXT_LIMIT → 模型名映射 → 兜底 1M。
  28. """
  29. # 直接从 .env 文件读取(不受 CWD / load_dotenv 时序影响)
  30. from pathlib import Path as _Path
  31. from dotenv import dotenv_values
  32. _env = dotenv_values(str(_Path(__file__).parent.parent / ".env"))
  33. env_val = _env.get("CONTEXT_LIMIT", "").strip()
  34. if env_val:
  35. try:
  36. return int(env_val)
  37. except ValueError:
  38. pass
  39. model = _env.get("DEEPAGENT_MODEL", "").strip().lower()
  40. # 去掉可能的 "openai:" 前缀
  41. if ":" in model:
  42. model = model.split(":", 1)[1]
  43. return _MODEL_CONTEXT_MAP.get(model, 1000000)
  44. def get_context_limit() -> int:
  45. """获取当前模型上下文窗口上限(懒加载)。"""
  46. return _resolve_context_limit()
  47. # ============================================================
  48. # 工具 schema 提取辅助
  49. # ============================================================
  50. def _tool_to_text(tool) -> str:
  51. """将工具函数转为可计数的文本表示(含名称 + 描述 + 参数 schema)。"""
  52. parts = []
  53. name = getattr(tool, "name", None) or getattr(tool, "__name__", str(tool))
  54. desc = getattr(tool, "description", "") or ""
  55. parts.append(f"Tool: {name}")
  56. if desc:
  57. parts.append(f"Description: {desc}")
  58. # 尝试提取参数 schema
  59. args_schema = getattr(tool, "args_schema", None)
  60. if args_schema and hasattr(args_schema, "schema"):
  61. try:
  62. import json
  63. parts.append(f"Args: {json.dumps(args_schema.schema(), ensure_ascii=False)}")
  64. except Exception:
  65. pass
  66. return "\n".join(parts)
  67. # ============================================================
  68. # 全局缓存 —— 系统组件 token 预估值(agent 启动时填充)
  69. # ============================================================
  70. _system_prompt_tokens = 0
  71. _skills_tokens = 0
  72. _mcp_tokens = 0
  73. def init_system_components(
  74. system_prompt: str = "",
  75. tool_defs: list[str] | None = None,
  76. skill_contents: list[str] | None = None,
  77. mcp_defs: list[str] | None = None,
  78. ):
  79. """初始化系统组件的 token 预估值。应在 agent 创建后调用一次。
  80. mcp 类别 = 本地工具 + MCP 远程工具(合并统计)。
  81. """
  82. global _system_prompt_tokens, _skills_tokens, _mcp_tokens
  83. _system_prompt_tokens = count_tokens(system_prompt)
  84. _skills_tokens = sum(count_tokens(s) for s in (skill_contents or []))
  85. # tools + mcp 合并为 mcp 类别
  86. _mcp_tokens = (
  87. sum(count_tokens(t) for t in (tool_defs or []))
  88. + sum(count_tokens(m) for m in (mcp_defs or []))
  89. )
  90. # ============================================================
  91. # Token 计数工具
  92. # ============================================================
  93. def count_tokens(text: str) -> int:
  94. """使用 tiktoken 精确计数 token 数。"""
  95. if not text:
  96. return 0
  97. try:
  98. return len(ENCODING.encode(text))
  99. except Exception:
  100. # 兜底估算:中文 ~1.5 字/token,英文 ~4 字/token
  101. return max(1, len(text) // 2)
  102. def _count_messages_tokens(session_id: str, current_message: str = "") -> int:
  103. """计算会话消息历史的 token 数。"""
  104. total = 0
  105. try:
  106. for msg in db_get_messages(session_id, limit=1000):
  107. total += count_tokens(msg.get("content", ""))
  108. except Exception:
  109. pass
  110. total += count_tokens(current_message)
  111. return total
  112. # ============================================================
  113. # 主入口:捕获并存储上下文用量
  114. # ============================================================
  115. async def capture_context_usage(
  116. session_id: str,
  117. prompt_tokens: int,
  118. completion_tokens: int,
  119. current_message: str = "",
  120. ):
  121. """在每次 agent 任务完成后调用,保存上下文用量快照。
  122. Args:
  123. session_id: 会话 ID
  124. prompt_tokens: 从 LLM 响应中提取的实际 prompt_tokens
  125. completion_tokens: 从 LLM 响应中提取的实际 completion_tokens
  126. current_message: 本次用户消息(用于计入 messages 估算)
  127. """
  128. if not session_id or prompt_tokens <= 0:
  129. return
  130. messages_tokens = _count_messages_tokens(session_id, current_message)
  131. # 系统组件 token(使用全局缓存估值)
  132. # mcp = 本地工具 + MCP 远程工具(已合并)
  133. skills_tokens = _skills_tokens
  134. mcp_tokens = _mcp_tokens
  135. system_prompt_tokens = _system_prompt_tokens
  136. # other = 提示词总量 - 已知各组件估算值(兜底 ≥0)
  137. known = messages_tokens + mcp_tokens + skills_tokens + system_prompt_tokens
  138. other_tokens = max(0, prompt_tokens - known)
  139. current_usage = prompt_tokens + completion_tokens
  140. try:
  141. save_context_usage(
  142. session_id=session_id,
  143. total_limit=get_context_limit(),
  144. current_usage=current_usage,
  145. messages_tokens=messages_tokens,
  146. mcp_tokens=mcp_tokens,
  147. skills_tokens=skills_tokens,
  148. system_prompt_tokens=system_prompt_tokens,
  149. other_tokens=other_tokens,
  150. )
  151. except Exception:
  152. pass # 上下文记录失败不应阻断对话