Coverage for agentos/conversation/conversation.py: 39%
205 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-10 01:20 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-10 01:20 +0800
1"""AgentOS v1.3.10 - Conversation Manager 模块。
3多轮对话上下文管理:滑动窗口、自动摘要、对话分支、token 感知裁剪。
4适用于长会话场景,防止上下文溢出,同时保持关键信息不丢失。
5"""
7from __future__ import annotations
9import hashlib
10import time
11from collections.abc import Callable
12from dataclasses import dataclass, field
13from enum import Enum, StrEnum, auto
16class MessageRole(StrEnum):
17 """消息角色。"""
19 SYSTEM = "system"
20 USER = "user"
21 ASSISTANT = "assistant"
22 TOOL = "tool"
25class TrimStrategy(Enum):
26 """裁剪策略。"""
28 FIFO = auto()
29 SUMMARIZE = auto()
30 IMPORTANCE_WEIGHTED = auto()
31 TOKEN_BUDGET = auto()
34@dataclass
35class Message:
36 """单条对话消息。"""
38 role: MessageRole
39 content: str
40 timestamp: float = field(default_factory=time.time)
41 token_count: int = 0
42 importance: float = 1.0
43 metadata: dict = field(default_factory=dict)
44 message_id: str = ""
46 def __post_init__(self):
47 if not self.message_id:
48 raw = f"{self.role.value}:{self.content[:50]}:{self.timestamp}"
49 self.message_id = hashlib.md5(raw.encode()).hexdigest()[:12]
52@dataclass
53class ConversationConfig:
54 """对话管理配置。"""
56 max_messages: int = 50
57 max_tokens: int = 8000
58 trim_strategy: TrimStrategy = TrimStrategy.FIFO
59 preserve_system: bool = True
60 preserve_last_n: int = 4
61 summary_prompt: str = ""
62 auto_summarize_threshold: float = 0.75
63 token_counter: Callable[[str], int] | None = None
66@dataclass
67class ConversationStats:
68 """对话统计。"""
70 total_messages: int = 0
71 total_tokens: int = 0
72 trim_count: int = 0
73 summarize_count: int = 0
74 branch_count: int = 0
75 oldest_timestamp: float = 0.0
76 newest_timestamp: float = 0.0
79@dataclass
80class ConversationSnapshot:
81 """对话快照(用于分支/恢复)。"""
83 messages: list[Message]
84 stats: ConversationStats
85 snapshot_id: str
86 created_at: float = field(default_factory=time.time)
87 label: str = ""
90class ConversationManager:
91 """多轮对话上下文管理器。
93 核心功能:
94 - 滑动窗口:超出 max_messages/max_tokens 时自动裁剪
95 - 自动摘要:超出阈值时压缩历史消息为摘要
96 - 对话分支:支持 fork 创建分支,切换/合并分支
97 - Token 感知:按 token 预算精确裁剪
98 """
100 def __init__(self, config: ConversationConfig | None = None):
101 self.config = config or ConversationConfig()
102 self._messages: list[Message] = []
103 self._summary: str = ""
104 self.stats = ConversationStats()
105 self._branches: dict[str, ConversationSnapshot] = {}
106 self._current_branch: str = "main"
107 self._message_counter: int = 0
109 # ── 消息管理 ──────────────────────────────────────────────
111 def add(self, role: MessageRole | str, content: str, **meta) -> Message:
112 """添加一条消息,自动触发裁剪检查。"""
113 if isinstance(role, str):
114 role = MessageRole(role)
115 token_count = self._count_tokens(content)
116 msg = Message(
117 role=role,
118 content=content,
119 token_count=token_count,
120 message_id=self._next_id(),
121 metadata=meta,
122 )
123 self._messages.append(msg)
124 self.stats.total_messages += 1
125 self.stats.total_tokens += token_count
126 if not self.stats.oldest_timestamp:
127 self.stats.oldest_timestamp = msg.timestamp
128 self.stats.newest_timestamp = msg.timestamp
129 self._enforce_limits()
130 return msg
132 def add_many(self, messages: list[tuple[str, str]]) -> list[Message]:
133 """批量添加消息。"""
134 return [self.add(role, content) for role, content in messages]
136 def get_context(self, include_summary: bool = True, limit: int | None = None) -> list[dict]:
137 """获取当前对话上下文,返回 OpenAI 兼容格式。"""
138 result: list[dict] = []
139 if include_summary and self._summary:
140 result.append({"role": "system", "content": f"[对话摘要] {self._summary}"})
141 msgs = self._messages[-limit:] if limit else self._messages
142 for msg in msgs:
143 result.append({"role": msg.role.value, "content": msg.content})
144 return result
146 def get_system_prompt(self) -> str:
147 """提取 system 消息。"""
148 for msg in self._messages:
149 if msg.role == MessageRole.SYSTEM:
150 return msg.content
151 return ""
153 # ── 裁剪与压缩 ────────────────────────────────────────────
155 def _enforce_limits(self):
156 """检查并执行裁剪。"""
157 changed = False
158 while len(self._messages) > self.config.max_messages:
159 self._trim_one()
160 changed = True
161 while self.stats.total_tokens > self.config.max_tokens:
162 self._trim_one()
163 changed = True
164 if (
165 changed
166 and self.config.trim_strategy == TrimStrategy.SUMMARIZE
167 and self.config.summary_prompt
168 ):
169 self._update_summary()
171 def _trim_one(self):
172 """按裁剪策略移除一条消息。"""
173 if self.config.trim_strategy == TrimStrategy.FIFO:
174 self._trim_fifo()
175 elif self.config.trim_strategy == TrimStrategy.IMPORTANCE_WEIGHTED:
176 self._trim_lowest_importance()
177 elif self.config.trim_strategy == TrimStrategy.TOKEN_BUDGET:
178 self._trim_token_budget()
179 else:
180 self._trim_fifo()
182 def _trim_fifo(self):
183 """先进先出裁剪:移除最旧非保留消息。"""
184 preserve = self.config.preserve_last_n
185 for i, msg in enumerate(self._messages):
186 if self.config.preserve_system and msg.role == MessageRole.SYSTEM:
187 continue
188 if len(self._messages) - i <= preserve:
189 break
190 self.stats.total_tokens -= msg.token_count
191 self.stats.trim_count += 1
192 self._messages.pop(i)
193 return
195 def _trim_lowest_importance(self):
196 """移除重要性最低的消息。"""
197 preserve = self.config.preserve_last_n
198 candidates = list(enumerate(self._messages))
199 if self.config.preserve_system:
200 candidates = [(i, m) for i, m in candidates if m.role != MessageRole.SYSTEM]
201 if len(candidates) <= preserve:
202 return
203 candidates = candidates[:-preserve]
204 idx, _ = min(candidates, key=lambda x: x[1].importance)
205 msg = self._messages.pop(idx)
206 self.stats.total_tokens -= msg.token_count
207 self.stats.trim_count += 1
209 def _trim_token_budget(self):
210 """按 token 预算裁剪。"""
211 budget = int(self.config.max_tokens * self.config.auto_summarize_threshold)
212 preserve_last = self.config.preserve_last_n
213 system_count = sum(
214 1
215 for m in self._messages
216 if m.role == MessageRole.SYSTEM and self.config.preserve_system
217 )
218 while (
219 self.stats.total_tokens > budget and len(self._messages) > preserve_last + system_count
220 ):
221 for i, msg in enumerate(self._messages):
222 if self.config.preserve_system and msg.role == MessageRole.SYSTEM:
223 continue
224 if len(self._messages) - i <= preserve_last:
225 break
226 self.stats.total_tokens -= msg.token_count
227 self.stats.trim_count += 1
228 self._messages.pop(i)
229 break
230 else:
231 break
233 def _update_summary(self):
234 """更新对话摘要(调用方需通过 summarize_callback 注入 LLM 实现)。"""
235 self._summary = f"[共 {len(self._messages)} 条消息, {self.stats.total_tokens} tokens]"
237 def set_summarizer(self, callback: Callable[[list[Message]], str]):
238 """注入摘要回调。"""
239 self._summarizer = callback
241 # ── 对话分支 ──────────────────────────────────────────────
243 def fork(self, label: str = "") -> ConversationSnapshot:
244 """创建对话分支快照。"""
245 import uuid
247 sid = uuid.uuid4().hex[:8]
248 snapshot = ConversationSnapshot(
249 messages=list(self._messages),
250 stats=ConversationStats(
251 total_messages=self.stats.total_messages,
252 total_tokens=self.stats.total_tokens,
253 trim_count=self.stats.trim_count,
254 summarize_count=self.stats.summarize_count,
255 branch_count=self.stats.branch_count,
256 oldest_timestamp=self.stats.oldest_timestamp,
257 newest_timestamp=self.stats.newest_timestamp,
258 ),
259 snapshot_id=sid,
260 label=label or f"branch-{sid}",
261 )
262 self._branches[sid] = snapshot
263 self.stats.branch_count += 1
264 return snapshot
266 def switch_branch(self, snapshot_id: str):
267 """切换到指定分支。"""
268 snapshot = self._branches.get(snapshot_id)
269 if not snapshot:
270 raise KeyError(f"Branch '{snapshot_id}' not found")
271 self._messages = list(snapshot.messages)
272 self.stats = ConversationStats(
273 total_messages=snapshot.stats.total_messages,
274 total_tokens=snapshot.stats.total_tokens,
275 trim_count=snapshot.stats.trim_count,
276 summarize_count=snapshot.stats.summarize_count,
277 branch_count=snapshot.stats.branch_count,
278 oldest_timestamp=snapshot.stats.oldest_timestamp,
279 newest_timestamp=snapshot.stats.newest_timestamp,
280 )
281 self._current_branch = snapshot_id
283 def merge_branch(self, snapshot_id: str, strategy: str = "append"):
284 """合并分支消息到当前对话。"""
285 snapshot = self._branches.get(snapshot_id)
286 if not snapshot:
287 raise KeyError(f"Branch '{snapshot_id}' not found")
288 if strategy == "append":
289 existing_ids = {m.message_id for m in self._messages}
290 for msg in snapshot.messages:
291 if msg.message_id not in existing_ids:
292 self._messages.append(msg)
293 self.stats.total_messages += 1
294 self.stats.total_tokens += msg.token_count
295 elif strategy == "replace":
296 self._messages = list(snapshot.messages)
297 self._enforce_limits()
299 def list_branches(self) -> dict[str, ConversationSnapshot]:
300 """列出所有分支。"""
301 return dict(self._branches)
303 # ── 工具方法 ──────────────────────────────────────────────
305 def _count_tokens(self, text: str) -> int:
306 """估算 token 数。"""
307 if self.config.token_counter:
308 return self.config.token_counter(text)
309 return len(text) // 3
311 def _next_id(self) -> str:
312 self._message_counter += 1
313 return f"msg_{self._message_counter:06d}"
315 def clear(self, keep_system: bool = True):
316 """清空对话历史。"""
317 system_msgs = (
318 [m for m in self._messages if m.role == MessageRole.SYSTEM] if keep_system else []
319 )
320 self._messages = system_msgs
321 self._summary = ""
322 self.stats = ConversationStats()
324 @property
325 def message_count(self) -> int:
326 return len(self._messages)
328 @property
329 def token_count(self) -> int:
330 return self.stats.total_tokens
332 def __len__(self) -> int:
333 return len(self._messages)
335 def __repr__(self) -> str:
336 return f"<Conversation messages={len(self)} tokens={self.token_count} branches={len(self._branches)}>"
339__all__ = [
340 "ConversationManager",
341 "ConversationConfig",
342 "ConversationStats",
343 "ConversationSnapshot",
344 "Message",
345 "MessageRole",
346 "TrimStrategy",
347]