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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-25 07:19 +0800
1"""Conversation management for multi-turn LLM interactions.
3This module provides conversation management with context window handling,
4automatic message trimming, token counting, and system prompt management.
5"""
7from __future__ import annotations
9from datetime import UTC, datetime
10from typing import TYPE_CHECKING, Any, cast
12from lexigram.ai.llm.types import ChatMessage, Completion, Role
13from lexigram.contracts.ai.llm import TokenCounterProtocol
14from lexigram.contracts.core import Metadata
16if TYPE_CHECKING:
17 from lexigram.contracts import AbstractLLMClient
20__all__ = [
21 "ConversationConfig",
22 "ConversationManager",
23 "ConversationStats",
24]
27from lexigram.ai.llm.conversation.types import ConversationConfig, ConversationStats
30class ConversationManager:
31 """Manage multi-turn conversations with automatic context window management.
33 This class handles:
34 - Message history management
35 - Automatic token counting
36 - Context window trimming
37 - System prompt handling
38 - Conversation statistics
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 """
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.
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 {}
92 # Initialize configuration
93 self._config = ConversationConfig(
94 max_tokens=max_tokens,
95 reserve_tokens=reserve_tokens,
96 trim_strategy=trim_strategy,
97 )
99 # Initialize token counter
100 if token_counter is None:
101 from lexigram.ai.llm.pricing.tokens import CharEstimateCounter
103 token_counter = CharEstimateCounter() # type: ignore[assignment]
104 self._token_counter: TokenCounterProtocol = token_counter # type: ignore[assignment]
106 # Initialize stats
107 self._stats = ConversationStats()
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()
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.
124 Args:
125 message: Message content
126 role: Message role (default: USER)
127 **completion_kwargs: Additional kwargs for completion
129 Returns:
130 Completion response from LLM
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)
140 if role == Role.USER:
141 self._stats.user_messages += 1
142 elif role == Role.ASSISTANT:
143 self._stats.assistant_messages += 1
145 # Update token count and trim if needed
146 await self._update_token_count()
147 await self._trim_if_needed()
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()
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)
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
186 # Update stats
187 await self._update_token_count()
188 self._stats.last_updated = datetime.now(UTC)
190 return cast("Completion", completion)
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.
200 Args:
201 role: Message role
202 content: Message content
203 update_stats: Whether to update statistics
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)
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
220 await self._update_token_count()
221 self._stats.last_updated = datetime.now(UTC)
223 def get_history(
224 self,
225 include_system: bool = True,
226 limit: int | None = None,
227 ) -> list[ChatMessage]:
228 """Get conversation history.
230 Args:
231 include_system: Include system message in history
232 limit: Maximum number of messages to return (most recent)
234 Returns:
235 List of chat messages
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()
244 if not include_system:
245 messages = list(filter(lambda msg: msg.role != Role.SYSTEM, messages))
247 if limit is not None:
248 messages = messages[-limit:]
250 return messages
252 def get_stats(self) -> ConversationStats:
253 """Get conversation statistics.
255 Returns:
256 Conversation statistics
258 Example:
259 >>> stats = manager.get_stats()
260 >>> print(f"Total tokens: {stats.total_tokens}")
261 """
262 return cast("ConversationStats", self._stats.model_copy())
264 def clear_history(self, keep_system: bool = True) -> None:
265 """Clear conversation history.
267 Args:
268 keep_system: Keep system message when clearing
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()
281 self._update_token_count_sync()
283 def update_system_prompt(self, system_prompt: str) -> None:
284 """Update the system prompt.
286 Args:
287 system_prompt: New system prompt
289 Example:
290 >>> manager.update_system_prompt("You are a Python expert.")
291 """
292 self._system_prompt = system_prompt
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)
301 self._update_token_count_sync()
303 def get_token_count(self) -> int:
304 """Get current total token count.
306 Returns:
307 Total tokens in conversation
309 Example:
310 >>> tokens = manager.get_token_count()
311 >>> print(f"Current tokens: {tokens}")
312 """
313 return self._stats.total_tokens
315 def get_available_tokens(self) -> int:
316 """Get available tokens for completion.
318 Returns:
319 Available tokens (max_tokens - current_tokens - reserve_tokens)
320 Can be negative if context window is exceeded
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
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)
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)
341 async def _trim_if_needed(self) -> None:
342 """Trim messages if context window is exceeded."""
343 available = self.get_available_tokens()
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()
355 self._stats.trimmed_count += 1
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 = []
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)
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)
376 self._messages = system_messages + other_messages
377 await self._update_token_count()
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 = []
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)
391 if len(other_messages) <= self._config.min_messages:
392 return
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)
403 self._messages = system_messages + other_messages
404 await self._update_token_count()
406 def export_history(self) -> dict[str, Any]:
407 """Export conversation history to dictionary.
409 Returns:
410 Dictionary with conversation data (JSON-serializable)
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 }
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.
433 Args:
434 client: LLM client
435 history_data: Exported history data
437 Returns:
438 ConversationManager instance
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", {})
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 )
458 # Load messages
459 messages_data = history_data.get("messages", [])
460 manager._messages = [ChatMessage(**msg) for msg in messages_data]
462 # Load stats
463 stats_data = history_data.get("stats", {})
464 if stats_data:
465 manager._stats = ConversationStats(**stats_data)
467 manager._update_token_count_sync()
469 return manager
471 def __len__(self) -> int:
472 """Get number of messages in conversation."""
473 return len(self._messages)
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 )