Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-llm/src/lexigram/ai/llm/conversation/manager.py: 18%

157 statements  

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

1"""Conversation management for multi-turn LLM interactions. 

2 

3This module provides conversation management with context window handling, 

4automatic message trimming, token counting, and system prompt management. 

5""" 

6 

7from __future__ import annotations 

8 

9from datetime import UTC, datetime 

10from typing import TYPE_CHECKING, Any, cast 

11 

12from lexigram.ai.llm.types import ChatMessage, Completion, Role 

13from lexigram.contracts.ai.llm import TokenCounterProtocol 

14from lexigram.contracts.core import Metadata 

15 

16if TYPE_CHECKING: 

17 from lexigram.contracts import AbstractLLMClient 

18 

19 

20__all__ = [ 

21 "ConversationConfig", 

22 "ConversationManager", 

23 "ConversationStats", 

24] 

25 

26 

27from lexigram.ai.llm.conversation.types import ConversationConfig, ConversationStats 

28 

29 

30class ConversationManager: 

31 """Manage multi-turn conversations with automatic context window management. 

32 

33 This class handles: 

34 - Message history management 

35 - Automatic token counting 

36 - Context window trimming 

37 - System prompt handling 

38 - Conversation statistics 

39 

40 Example: 

41 >>> from lexigram.ai.llm import OpenAIClient, ConversationManager 

42 >>> 

43 >>> client = OpenAIClient(api_key="sk-...", model="gpt-4") 

44 >>> manager = ConversationManager( 

45 ... client=client, 

46 ... system_prompt="You are a helpful assistant.", 

47 ... max_tokens=4096 

48 ... ) 

49 >>> 

50 >>> # Add user message and get response 

51 >>> response = await manager.chat("What is Python?") 

52 >>> print(response.content) 

53 >>> 

54 >>> # Continue conversation 

55 >>> response = await manager.chat("Tell me more about it") 

56 >>> print(response.content) 

57 >>> 

58 >>> # Get conversation history 

59 >>> history = manager.get_history() 

60 >>> stats = manager.get_stats() 

61 >>> print(f"Total messages: {stats.total_messages}") 

62 >>> print(f"Total tokens: {stats.total_tokens}") 

63 """ 

64 

65 def __init__( 

66 self, 

67 client: AbstractLLMClient, 

68 system_prompt: str | None = None, 

69 max_tokens: int = 4096, 

70 reserve_tokens: int = 1000, 

71 trim_strategy: str = "oldest", 

72 metadata: Metadata | None = None, 

73 token_counter: TokenCounterProtocol | None = None, 

74 ) -> None: 

75 """Initialize conversation manager. 

76 

77 Args: 

78 client: LLM client for completions 

79 system_prompt: Optional system prompt (prepended to all conversations) 

80 max_tokens: Maximum context window size 

81 reserve_tokens: Tokens to reserve for completion 

82 trim_strategy: Message trimming strategy ('oldest', 'middle', 'summary') 

83 metadata: Additional metadata for the conversation 

84 token_counter: Optional TokenCounterProtocol implementation. 

85 If not provided, uses CharEstimateCounter. 

86 """ 

87 self._client = client 

88 self._system_prompt = system_prompt 

89 self._messages: list[ChatMessage] = [] 

90 self._metadata = metadata or {} 

91 

92 # Initialize configuration 

93 self._config = ConversationConfig( 

94 max_tokens=max_tokens, 

95 reserve_tokens=reserve_tokens, 

96 trim_strategy=trim_strategy, 

97 ) 

98 

99 # Initialize token counter 

100 if token_counter is None: 

101 from lexigram.ai.llm.pricing.tokens import CharEstimateCounter 

102 

103 token_counter = CharEstimateCounter() # type: ignore[assignment] 

104 self._token_counter: TokenCounterProtocol = token_counter # type: ignore[assignment] 

105 

106 # Initialize stats 

107 self._stats = ConversationStats() 

108 

109 # Add system message if provided 

110 if system_prompt: 

111 system_msg = ChatMessage(role=Role.SYSTEM, content=system_prompt) 

112 self._messages.append(system_msg) 

113 self._stats.system_messages += 1 

114 self._update_token_count_sync() 

115 

116 async def chat( 

117 self, 

118 message: str, 

119 role: Role = Role.USER, 

120 **completion_kwargs: Any, 

121 ) -> Completion: 

122 """Send a message and get a response. 

123 

124 Args: 

125 message: Message content 

126 role: Message role (default: USER) 

127 **completion_kwargs: Additional kwargs for completion 

128 

129 Returns: 

130 Completion response from LLM 

131 

132 Example: 

133 >>> response = await manager.chat("Hello!") 

134 >>> print(response.content) 

135 """ 

136 # Add user message 

137 user_msg = ChatMessage(role=role, content=message) 

138 self._messages.append(user_msg) 

139 

140 if role == Role.USER: 

141 self._stats.user_messages += 1 

142 elif role == Role.ASSISTANT: 

143 self._stats.assistant_messages += 1 

144 

145 # Update token count and trim if needed 

146 await self._update_token_count() 

147 await self._trim_if_needed() 

148 

149 # Get completion — pass ChatMessage objects directly to preserve thinking_blocks 

150 result = await self._client.complete( 

151 messages=self._messages, 

152 **completion_kwargs, 

153 ) 

154 if result.is_err(): 

155 raise result.unwrap_err() 

156 completion = result.unwrap() 

157 

158 # Normalize string completions into Completion model for downstream consumers 

159 if isinstance(completion, str): 

160 model_name = completion_kwargs.get("model") 

161 if not isinstance(model_name, str): 

162 model_name = getattr(self._client, "model", None) 

163 if not isinstance(model_name, str): 

164 model_name = "unknown" 

165 completion = Completion(content=completion, model=model_name) 

166 

167 # Add assistant response to history, preserving thinking blocks for 

168 # providers (e.g. Anthropic) that require re-injection in subsequent turns. 

169 thinking_blocks: list[dict[str, Any]] | None = None 

170 if completion.thinking and completion.thinking.signature: 

171 thinking_blocks = [ 

172 { 

173 "type": "thinking", 

174 "thinking": completion.thinking.content, 

175 "signature": completion.thinking.signature, 

176 } 

177 ] 

178 assistant_msg = ChatMessage( 

179 role=Role.ASSISTANT, 

180 content=completion.content, 

181 thinking_blocks=thinking_blocks, 

182 ) 

183 self._messages.append(assistant_msg) 

184 self._stats.assistant_messages += 1 

185 

186 # Update stats 

187 await self._update_token_count() 

188 self._stats.last_updated = datetime.now(UTC) 

189 

190 return cast("Completion", completion) 

191 

192 async def add_message( 

193 self, 

194 role: Role, 

195 content: str, 

196 update_stats: bool = True, 

197 ) -> None: 

198 """Add a message to conversation history without getting a response. 

199 

200 Args: 

201 role: Message role 

202 content: Message content 

203 update_stats: Whether to update statistics 

204 

205 Example: 

206 >>> await manager.add_message(Role.USER, "Hello") 

207 >>> await manager.add_message(Role.ASSISTANT, "Hi there!") 

208 """ 

209 msg = ChatMessage(role=role, content=content) 

210 self._messages.append(msg) 

211 

212 if update_stats: 

213 if role == Role.USER: 

214 self._stats.user_messages += 1 

215 elif role == Role.ASSISTANT: 

216 self._stats.assistant_messages += 1 

217 elif role == Role.SYSTEM: 

218 self._stats.system_messages += 1 

219 

220 await self._update_token_count() 

221 self._stats.last_updated = datetime.now(UTC) 

222 

223 def get_history( 

224 self, 

225 include_system: bool = True, 

226 limit: int | None = None, 

227 ) -> list[ChatMessage]: 

228 """Get conversation history. 

229 

230 Args: 

231 include_system: Include system message in history 

232 limit: Maximum number of messages to return (most recent) 

233 

234 Returns: 

235 List of chat messages 

236 

237 Example: 

238 >>> history = manager.get_history(limit=10) 

239 >>> for msg in history: 

240 ... print(f"{msg.role}: {msg.content}") 

241 """ 

242 messages = self._messages.copy() 

243 

244 if not include_system: 

245 messages = list(filter(lambda msg: msg.role != Role.SYSTEM, messages)) 

246 

247 if limit is not None: 

248 messages = messages[-limit:] 

249 

250 return messages 

251 

252 def get_stats(self) -> ConversationStats: 

253 """Get conversation statistics. 

254 

255 Returns: 

256 Conversation statistics 

257 

258 Example: 

259 >>> stats = manager.get_stats() 

260 >>> print(f"Total tokens: {stats.total_tokens}") 

261 """ 

262 return cast("ConversationStats", self._stats.model_copy()) 

263 

264 def clear_history(self, keep_system: bool = True) -> None: 

265 """Clear conversation history. 

266 

267 Args: 

268 keep_system: Keep system message when clearing 

269 

270 Example: 

271 >>> manager.clear_history() 

272 """ 

273 if keep_system and self._system_prompt: 

274 system_msg = ChatMessage(role=Role.SYSTEM, content=self._system_prompt) 

275 self._messages = [system_msg] 

276 self._stats = ConversationStats(system_messages=1) 

277 else: 

278 self._messages = [] 

279 self._stats = ConversationStats() 

280 

281 self._update_token_count_sync() 

282 

283 def update_system_prompt(self, system_prompt: str) -> None: 

284 """Update the system prompt. 

285 

286 Args: 

287 system_prompt: New system prompt 

288 

289 Example: 

290 >>> manager.update_system_prompt("You are a Python expert.") 

291 """ 

292 self._system_prompt = system_prompt 

293 

294 # Remove old system message and add new one 

295 self._messages = list( 

296 filter(lambda msg: msg.role != Role.SYSTEM, self._messages), 

297 ) 

298 system_msg = ChatMessage(role=Role.SYSTEM, content=system_prompt) 

299 self._messages.insert(0, system_msg) 

300 

301 self._update_token_count_sync() 

302 

303 def get_token_count(self) -> int: 

304 """Get current total token count. 

305 

306 Returns: 

307 Total tokens in conversation 

308 

309 Example: 

310 >>> tokens = manager.get_token_count() 

311 >>> print(f"Current tokens: {tokens}") 

312 """ 

313 return self._stats.total_tokens 

314 

315 def get_available_tokens(self) -> int: 

316 """Get available tokens for completion. 

317 

318 Returns: 

319 Available tokens (max_tokens - current_tokens - reserve_tokens) 

320 Can be negative if context window is exceeded 

321 

322 Example: 

323 >>> available = manager.get_available_tokens() 

324 >>> print(f"Available for completion: {available}") 

325 """ 

326 used = self._stats.total_tokens 

327 reserved = self._config.reserve_tokens 

328 max_tokens = self._config.max_tokens 

329 return max_tokens - used - reserved 

330 

331 def _update_token_count_sync(self) -> None: 

332 """Update total token count synchronously (for initialization).""" 

333 self._stats.total_tokens = self._token_counter.count_messages(self._messages) # type: ignore[arg-type] 

334 self._stats.total_messages = len(self._messages) 

335 

336 async def _update_token_count(self) -> None: 

337 """Update total token count.""" 

338 self._stats.total_tokens = self._token_counter.count_messages(self._messages) # type: ignore[arg-type] 

339 self._stats.total_messages = len(self._messages) 

340 

341 async def _trim_if_needed(self) -> None: 

342 """Trim messages if context window is exceeded.""" 

343 available = self.get_available_tokens() 

344 

345 if available < 0: 

346 # Need to trim messages 

347 if self._config.trim_strategy == "oldest": 

348 await self._trim_oldest() 

349 elif self._config.trim_strategy == "middle": 

350 await self._trim_middle() 

351 else: 

352 # Default to oldest if strategy not recognized 

353 await self._trim_oldest() 

354 

355 self._stats.trimmed_count += 1 

356 

357 async def _trim_oldest(self) -> None: 

358 """Trim oldest messages (keeping system message).""" 

359 # Keep system message if configured 

360 system_messages = [] 

361 other_messages = [] 

362 

363 for msg in self._messages: 

364 if msg.role == Role.SYSTEM and self._config.keep_system: 

365 system_messages.append(msg) 

366 else: 

367 other_messages.append(msg) 

368 

369 # Keep minimum number of messages 

370 min_keep = self._config.min_messages 

371 if len(other_messages) > min_keep: 

372 # Remove oldest messages until we're under token limit 

373 while self.get_available_tokens() < 0 and len(other_messages) > min_keep: 

374 other_messages.pop(0) 

375 

376 self._messages = system_messages + other_messages 

377 await self._update_token_count() 

378 

379 async def _trim_middle(self) -> None: 

380 """Trim middle messages (keeping system, first few, and last few).""" 

381 # Keep system message 

382 system_messages = [] 

383 other_messages = [] 

384 

385 for msg in self._messages: 

386 if msg.role == Role.SYSTEM and self._config.keep_system: 

387 system_messages.append(msg) 

388 else: 

389 other_messages.append(msg) 

390 

391 if len(other_messages) <= self._config.min_messages: 

392 return 

393 

394 # Keep first and last messages, remove middle 

395 keep_each_side = self._config.min_messages // 2 

396 while ( 

397 self.get_available_tokens() < 0 and len(other_messages) > keep_each_side * 2 

398 ): 

399 # Remove from middle 

400 middle_idx = len(other_messages) // 2 

401 other_messages.pop(middle_idx) 

402 

403 self._messages = system_messages + other_messages 

404 await self._update_token_count() 

405 

406 def export_history(self) -> dict[str, Any]: 

407 """Export conversation history to dictionary. 

408 

409 Returns: 

410 Dictionary with conversation data (JSON-serializable) 

411 

412 Example: 

413 >>> data = manager.export_history() 

414 >>> from lexigram import serialization as json 

415 >>> with open("conversation.json", "w") as f: 

416 ... json.dump(data, f) 

417 """ 

418 return { 

419 "messages": [msg.model_dump(mode="json") for msg in self._messages], 

420 "stats": self._stats.model_dump(mode="json"), 

421 "config": self._config.model_dump(mode="json"), 

422 "metadata": self._metadata, 

423 } 

424 

425 @classmethod 

426 def from_history( 

427 cls, 

428 client: AbstractLLMClient, 

429 history_data: dict[str, Any], 

430 ) -> ConversationManager: 

431 """Create conversation manager from exported history. 

432 

433 Args: 

434 client: LLM client 

435 history_data: Exported history data 

436 

437 Returns: 

438 ConversationManager instance 

439 

440 Example: 

441 >>> from lexigram import serialization as json 

442 >>> with open("conversation.json") as f: 

443 ... data = json.load(f) 

444 >>> manager = ConversationManager.from_history(client, data) 

445 """ 

446 config = history_data.get("config", {}) 

447 metadata = history_data.get("metadata", {}) 

448 

449 manager = cls( 

450 client=client, 

451 system_prompt=None, # Will be loaded from messages 

452 max_tokens=config.get("max_tokens", 4096), 

453 reserve_tokens=config.get("reserve_tokens", 1000), 

454 trim_strategy=config.get("trim_strategy", "oldest"), 

455 metadata=metadata, 

456 ) 

457 

458 # Load messages 

459 messages_data = history_data.get("messages", []) 

460 manager._messages = [ChatMessage(**msg) for msg in messages_data] 

461 

462 # Load stats 

463 stats_data = history_data.get("stats", {}) 

464 if stats_data: 

465 manager._stats = ConversationStats(**stats_data) 

466 

467 manager._update_token_count_sync() 

468 

469 return manager 

470 

471 def __len__(self) -> int: 

472 """Get number of messages in conversation.""" 

473 return len(self._messages) 

474 

475 def __repr__(self) -> str: 

476 """String representation.""" 

477 return ( 

478 f"ConversationManager(" 

479 f"messages={len(self._messages)}, " 

480 f"tokens={self._stats.total_tokens}, " 

481 f"available={self.get_available_tokens()})" 

482 )