Coverage for merco/context/processors/compress.py: 60%

91 statements  

« prev     ^ index     » next       coverage.py v7.15.0, created at 2026-07-07 14:04 +0800

1"""CompressProcessor — 替代 ContextCompressor""" 

2from __future__ import annotations 

3 

4import logging 

5from merco.context.pipeline import ContextProcessor 

6from merco.core.context import msg_tokens 

7 

8logger = logging.getLogger("merco.context.compress") 

9 

10 

11class CompressProcessor(ContextProcessor): 

12 """压缩:超过阈值时摘要旧消息""" 

13 name = "compress" 

14 

15 def __init__(self, max_tokens: int = 64000, threshold: float = 0.75): 

16 self.max_tokens = max_tokens 

17 self.threshold = threshold 

18 

19 async def process(self, messages: list[dict], **kwargs) -> list[dict]: 

20 total = sum(msg_tokens(m) for m in messages) 

21 trigger = int(self.max_tokens * self.threshold) 

22 if total <= trigger or len(messages) <= 4: 

23 return messages 

24 

25 strategy = kwargs.get("compress_strategy", "sliding") 

26 summary_fn = kwargs.get("summary_fn") 

27 

28 if strategy == "sliding": 

29 return await self._sliding(messages, summary_fn) 

30 elif strategy == "truncate": 

31 return self._truncate(messages) 

32 return messages 

33 

34 async def _sliding(self, messages: list[dict], summary_fn=None) -> list[dict]: 

35 """滑动窗口压缩 — 保留最后 2 轮原文 + 摘要旧消息""" 

36 TAIL_TURNS = 2 

37 

38 tail_start = 0 

39 user_count = 0 

40 for i in range(len(messages) - 1, -1, -1): 

41 if messages[i].get("role") == "user": 

42 user_count += 1 

43 if user_count >= TAIL_TURNS: 

44 tail_start = i 

45 break 

46 

47 body = messages[:tail_start] 

48 tail = messages[tail_start:] 

49 

50 if not body: 

51 return messages 

52 

53 if summary_fn: 

54 try: 

55 summary_text = await summary_fn(body) 

56 summary = {"role": "system", "content": summary_text} 

57 except Exception as e: 

58 logger.warning("LLM 摘要失败: %s, fallback", e) 

59 summary = self._build_summary(body) 

60 else: 

61 summary = self._build_summary(body) 

62 

63 result = [m for m in tail if m.get("role") == "system"][:1] 

64 result.append(summary) 

65 result.extend(m for m in tail if m.get("role") != "system") 

66 

67 before = sum(msg_tokens(m) for m in messages) 

68 after = sum(msg_tokens(m) for m in result) 

69 logger.debug("压缩: %d条(%dtok) → %d条(%dtok)", len(messages), before, len(result), after) 

70 return result 

71 

72 def _truncate(self, messages: list[dict]) -> list[dict]: 

73 """简单截断 fallback""" 

74 if len(messages) <= 6: 

75 return messages 

76 kept = messages[-6:] 

77 return self._extend_to_chain(messages, messages[:-6], kept) 

78 

79 def _extend_to_chain(self, all_messages, before, kept): 

80 """补全孤立 tool 消息的前导 assistant""" 

81 while True: 

82 orphan_at = None 

83 for i, msg in enumerate(kept): 

84 if msg.get("role") != "tool": 

85 continue 

86 prev = kept[i - 1] if i > 0 else None 

87 if not (prev and prev.get("role") == "assistant" and prev.get("tool_calls")): 

88 orphan_at = i 

89 break 

90 if orphan_at is None: 

91 break 

92 try: 

93 orig_idx = all_messages.index(kept[orphan_at]) 

94 except ValueError: 

95 break 

96 found = None 

97 for j in range(orig_idx - 1, -1, -1): 

98 msg = all_messages[j] 

99 if msg.get("role") == "assistant" and msg.get("tool_calls"): 

100 found = msg 

101 break 

102 if found is None or found in kept: 

103 break 

104 kept.insert(0, found) 

105 return kept 

106 

107 def _build_summary(self, messages: list[dict]) -> dict: 

108 """Fallback 摘要""" 

109 user_msgs = [m for m in messages if m.get("role") == "user"] 

110 preview = [] 

111 for um in user_msgs[-5:]: 

112 c = um.get("content", "")[:60] 

113 if c: 

114 preview.append(f"{c}") 

115 intro = ( 

116 f"[压缩了 {len(messages)} 条历史消息。" 

117 f"最近讨论: {'; '.join(preview) if preview else '无'}" 

118 f"详细历史请用 /search 查询。]" 

119 ) 

120 return {"role": "system", "content": intro}