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

1"""上下文管理与压缩""" 

2 

3import json 

4import logging 

5import re 

6 

7_logger = logging.getLogger("merco.context") 

8 

9 

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) 

17 

18 

19class ContextManager: 

20 """管理对话上下文,包括窗口控制与压缩""" 

21 

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 

28 

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) 

37 

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 

42 

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 

49 

50 def needs_compression(self) -> bool: 

51 """判断是否需要压缩""" 

52 return self.total_tokens > self.max_tokens * 0.8 

53 

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:] 

59 

60 

61 

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