Coverage for merco/core/context.py: 93%
44 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
1"""上下文管理与压缩"""
3import json
4import logging
5import re
7_logger = logging.getLogger("merco.context")
10def estimate_tokens(text: str) -> int:
11 """估算 token 数。CJK 约 1.5 token/字,英文约 4 字符/token。"""
12 if not text:
13 return 0
14 cjk = len(re.findall(r'[\u4e00-\u9fff\u3400-\u4dbf\uf900-\ufaff]', text))
15 other = len(text) - cjk
16 return int(cjk * 1.5 + other / 4)
19class ContextManager:
20 """管理对话上下文,包括窗口控制与压缩"""
22 def __init__(self, max_tokens: int = 128000):
23 self.max_tokens = max_tokens
24 self.current_tokens = 0
25 self.messages = []
26 self._overhead_tokens = 0 # system prompt + tool defs 等固定开销
27 self.last_actual_tokens = 0 # API 返回的实测 prompt_tokens
29 def add(self, message: dict):
30 """添加消息到上下文"""
31 r = message.get("reasoning", "")
32 if r:
33 _logger.warning("context.add: 消息包含 reasoning (%d chars, first 100: %s…)",
34 len(r), r[:100].replace("\n", "\\n"))
35 self.messages.append(message)
36 self.current_tokens += msg_tokens(message)
38 def set_overhead(self, system_prompt: str, tool_count: int):
39 """设置固定开销(system prompt + tool definitions)"""
40 # system prompt + tool 定义平均 ~200 token/个
41 self._overhead_tokens = estimate_tokens(system_prompt) + tool_count * 200
43 @property
44 def total_tokens(self) -> int:
45 """发送给 LLM 的 token 数(优先 API 实测值,回退估算)"""
46 if self.last_actual_tokens > 0:
47 return self.last_actual_tokens
48 return self.current_tokens + self._overhead_tokens
50 def needs_compression(self) -> bool:
51 """判断是否需要压缩"""
52 return self.total_tokens > self.max_tokens * 0.8
54 def get_window(self, n: int = None) -> list:
55 """获取最近 N 条消息"""
56 if n is None:
57 return self.messages
58 return self.messages[-n:]
62def msg_tokens(message: dict) -> int:
63 """单条消息 token 估算(content + tool_calls JSON)"""
64 total = 0
65 content = message.get("content", "")
66 if isinstance(content, str):
67 total += estimate_tokens(content)
68 for tc in message.get("tool_calls", []) or []:
69 total += estimate_tokens(json.dumps(tc, ensure_ascii=False))
70 return total