Coverage for agentos/conversation/conversation.py: 39%

205 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 09:19 +0800

1"""AgentOS v1.3.10 - Conversation Manager 模块。 

2 

3多轮对话上下文管理:滑动窗口、自动摘要、对话分支、token 感知裁剪。 

4适用于长会话场景,防止上下文溢出,同时保持关键信息不丢失。 

5""" 

6 

7from __future__ import annotations 

8 

9import hashlib 

10import time 

11from collections.abc import Callable 

12from dataclasses import dataclass, field 

13from enum import Enum, StrEnum, auto 

14 

15 

16class MessageRole(StrEnum): 

17 """消息角色。""" 

18 

19 SYSTEM = "system" 

20 USER = "user" 

21 ASSISTANT = "assistant" 

22 TOOL = "tool" 

23 

24 

25class TrimStrategy(Enum): 

26 """裁剪策略。""" 

27 

28 FIFO = auto() 

29 SUMMARIZE = auto() 

30 IMPORTANCE_WEIGHTED = auto() 

31 TOKEN_BUDGET = auto() 

32 

33 

34@dataclass 

35class Message: 

36 """单条对话消息。""" 

37 

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 = "" 

45 

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] 

50 

51 

52@dataclass 

53class ConversationConfig: 

54 """对话管理配置。""" 

55 

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 

64 

65 

66@dataclass 

67class ConversationStats: 

68 """对话统计。""" 

69 

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 

77 

78 

79@dataclass 

80class ConversationSnapshot: 

81 """对话快照(用于分支/恢复)。""" 

82 

83 messages: list[Message] 

84 stats: ConversationStats 

85 snapshot_id: str 

86 created_at: float = field(default_factory=time.time) 

87 label: str = "" 

88 

89 

90class ConversationManager: 

91 """多轮对话上下文管理器。 

92 

93 核心功能: 

94 - 滑动窗口:超出 max_messages/max_tokens 时自动裁剪 

95 - 自动摘要:超出阈值时压缩历史消息为摘要 

96 - 对话分支:支持 fork 创建分支,切换/合并分支 

97 - Token 感知:按 token 预算精确裁剪 

98 """ 

99 

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 

108 

109 # ── 消息管理 ────────────────────────────────────────────── 

110 

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 

131 

132 def add_many(self, messages: list[tuple[str, str]]) -> list[Message]: 

133 """批量添加消息。""" 

134 return [self.add(role, content) for role, content in messages] 

135 

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 

145 

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 "" 

152 

153 # ── 裁剪与压缩 ──────────────────────────────────────────── 

154 

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() 

170 

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() 

181 

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 

194 

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 

208 

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 

232 

233 def _update_summary(self): 

234 """更新对话摘要(调用方需通过 summarize_callback 注入 LLM 实现)。""" 

235 self._summary = f"[共 {len(self._messages)} 条消息, {self.stats.total_tokens} tokens]" 

236 

237 def set_summarizer(self, callback: Callable[[list[Message]], str]): 

238 """注入摘要回调。""" 

239 self._summarizer = callback 

240 

241 # ── 对话分支 ────────────────────────────────────────────── 

242 

243 def fork(self, label: str = "") -> ConversationSnapshot: 

244 """创建对话分支快照。""" 

245 import uuid 

246 

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 

265 

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 

282 

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() 

298 

299 def list_branches(self) -> dict[str, ConversationSnapshot]: 

300 """列出所有分支。""" 

301 return dict(self._branches) 

302 

303 # ── 工具方法 ────────────────────────────────────────────── 

304 

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 

310 

311 def _next_id(self) -> str: 

312 self._message_counter += 1 

313 return f"msg_{self._message_counter:06d}" 

314 

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() 

323 

324 @property 

325 def message_count(self) -> int: 

326 return len(self._messages) 

327 

328 @property 

329 def token_count(self) -> int: 

330 return self.stats.total_tokens 

331 

332 def __len__(self) -> int: 

333 return len(self._messages) 

334 

335 def __repr__(self) -> str: 

336 return f"<Conversation messages={len(self)} tokens={self.token_count} branches={len(self._branches)}>" 

337 

338 

339__all__ = [ 

340 "ConversationManager", 

341 "ConversationConfig", 

342 "ConversationStats", 

343 "ConversationSnapshot", 

344 "Message", 

345 "MessageRole", 

346 "TrimStrategy", 

347]