| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191 |
- # -*- 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 # 上下文记录失败不应阻断对话
|