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
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
1"""CompressProcessor — 替代 ContextCompressor"""
2from __future__ import annotations
4import logging
5from merco.context.pipeline import ContextProcessor
6from merco.core.context import msg_tokens
8logger = logging.getLogger("merco.context.compress")
11class CompressProcessor(ContextProcessor):
12 """压缩:超过阈值时摘要旧消息"""
13 name = "compress"
15 def __init__(self, max_tokens: int = 64000, threshold: float = 0.75):
16 self.max_tokens = max_tokens
17 self.threshold = threshold
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
25 strategy = kwargs.get("compress_strategy", "sliding")
26 summary_fn = kwargs.get("summary_fn")
28 if strategy == "sliding":
29 return await self._sliding(messages, summary_fn)
30 elif strategy == "truncate":
31 return self._truncate(messages)
32 return messages
34 async def _sliding(self, messages: list[dict], summary_fn=None) -> list[dict]:
35 """滑动窗口压缩 — 保留最后 2 轮原文 + 摘要旧消息"""
36 TAIL_TURNS = 2
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
47 body = messages[:tail_start]
48 tail = messages[tail_start:]
50 if not body:
51 return messages
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)
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")
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
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)
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
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}