"""Core compression service with simplified responsibilities.""" import logging import re from datetime import datetime, timezone from typing import Any, Dict, List, Optional from docsgpt.api.answer.services.compression.prompt_builder import ( CompressionPromptBuilder, ) from docsgpt.api.answer.services.compression.token_counter import TokenCounter from docsgpt.api.answer.services.compression.types import ( CompressionMetadata, is_compression_summary_row, latest_usable_compression_point, ) from docsgpt.core.settings import settings logger = logging.getLogger(__name__) class CompressionService: """ Service for compressing conversation history. Handles DB updates. """ def __init__( self, llm, model_id: str, conversation_service=None, prompt_builder: Optional[CompressionPromptBuilder] = None, ): """ Initialize compression service. Args: llm: LLM instance to use for compression model_id: Model ID for compression conversation_service: Service for DB operations (optional, for DB updates) prompt_builder: Custom prompt builder (optional) """ self.llm = llm self.model_id = model_id self.conversation_service = conversation_service self.prompt_builder = prompt_builder or CompressionPromptBuilder( version=settings.COMPRESSION_PROMPT_VERSION ) def compress_conversation( self, conversation: Dict[str, Any], compress_up_to_index: int, start_index: int = 0, persist_query_index: Optional[int] = None, ) -> CompressionMetadata: """ Compress conversation history up to specified index. Args: conversation: Full conversation document compress_up_to_index: Last query index to include in compression start_index: First query index that is new since the previous compression point (earlier queries are already summarised). persist_query_index: Index to record on the point when the conversation passed in is a shortened, synthetic one (the mid-execution path) and its positions do not match the saved conversation's. Defaults to ``compress_up_to_index``. Returns: CompressionMetadata with compression details Raises: ValueError: If compress_up_to_index is invalid """ try: queries = conversation.get("queries", []) if compress_up_to_index < 0 or compress_up_to_index >= len(queries): raise ValueError( f"Invalid compress_up_to_index: {compress_up_to_index} " f"(conversation has {len(queries)} queries)" ) if start_index < 0 or start_index > compress_up_to_index: raise ValueError( f"Nothing to compress: start_index {start_index} is past " f"compress_up_to_index {compress_up_to_index}" ) # Only the queries after the previous compression point are new; # the earlier ones are already inside that point's summary, and so # is the visible summary row that follows the point. queries_to_compress = [ q for q in queries[start_index : compress_up_to_index + 1] if not is_compression_summary_row(q) ] if not queries_to_compress: raise ValueError( "Nothing to compress: no new queries since the last " "compression point" ) # Check if there are existing compressions. ``compression_metadata`` # is a nullable JSONB column, so a never-compressed conversation # reads back as None; ``get(key, {})`` would return that None (the # default only applies to absent keys), so coalesce with ``or {}``. existing_compressions = (conversation.get("compression_metadata") or {}).get( "compression_points", [] ) previous_summary_tokens = 0 usable_point = latest_usable_compression_point(existing_compressions) existing_compressions = [usable_point] if usable_point else [] if existing_compressions: # Each point already folds in the ones before it, so only # the latest usable one matters — and it is part of what # the new summary replaces. previous_summary_tokens = TokenCounter.count_message_tokens( [{"content": existing_compressions[0].get("compressed_summary", "")}] ) logger.info( "Found a previous compression point (query %s) - " "the new summary builds on it", existing_compressions[0].get("query_index"), ) # Calculate original token count: everything the new summary replaces original_tokens = ( TokenCounter.count_query_tokens(queries_to_compress) + previous_summary_tokens ) # Log tool call stats self._log_tool_call_stats(queries_to_compress) # Build compression prompt messages = self.prompt_builder.build_prompt( queries_to_compress, existing_compressions ) # Call LLM to generate compression logger.info( f"Starting compression: {len(queries_to_compress)} queries " f"(messages {start_index}-{compress_up_to_index}, {original_tokens} tokens) " f"using model {self.model_id}" ) # See note in conversation_service.py: ``self.model_id`` is # the registry id (UUID for BYOM); the LLM's own model_id is # what the provider's API actually expects. response = self.llm.gen( model=getattr(self.llm, "model_id", None) or self.model_id, messages=messages, max_tokens=4000, ) # Extract summary from response compressed_summary = self._extract_summary(response) # Calculate compressed token count compressed_tokens = TokenCounter.count_message_tokens( [{"content": compressed_summary}] ) # An empty summary is not a compression: it replaced a 494k-token # conversation with nothing in prod (2026-08-27) while reporting # success. if not compressed_summary.strip() or compressed_tokens <= 0: raise ValueError( "Compression produced an empty summary; keeping original history" ) # Calculate compression ratio compression_ratio = ( original_tokens / compressed_tokens if compressed_tokens > 0 else 0 ) # Port of the in-memory path's guard: a "successful" summary # that isn't smaller than what it replaces must never become a # compression point — it would make every later rebuild WORSE # while reporting success (observed in prod: "successful" # compressions with negative savings). if compressed_tokens >= original_tokens: raise ValueError( f"Compression did not reduce token count " f"({original_tokens} → {compressed_tokens}); " f"keeping original history" ) logger.info( f"Compression complete: {original_tokens} → {compressed_tokens} tokens " f"({compression_ratio:.1f}x compression)" ) # Build compression metadata compression_metadata = CompressionMetadata( timestamp=datetime.now(timezone.utc), query_index=( persist_query_index if persist_query_index is not None else compress_up_to_index ), compressed_summary=compressed_summary, original_token_count=original_tokens, compressed_token_count=compressed_tokens, compression_ratio=compression_ratio, model_used=self.model_id, compression_prompt_version=self.prompt_builder.version, ) return compression_metadata except Exception as e: logger.error(f"Error compressing conversation: {str(e)}", exc_info=True) raise def compress_and_save( self, conversation_id: str, conversation: Dict[str, Any], compress_up_to_index: int, start_index: int = 0, persist_query_index: Optional[int] = None, ) -> CompressionMetadata: """ Compress conversation and save to database. Args: conversation_id: Conversation ID conversation: Full conversation document compress_up_to_index: Last query index to include Returns: CompressionMetadata Raises: ValueError: If conversation_service not provided or invalid index """ if not self.conversation_service: raise ValueError( "conversation_service required for compress_and_save operation" ) # Perform compression metadata = self.compress_conversation( conversation, compress_up_to_index, start_index=start_index, persist_query_index=persist_query_index, ) # Save to database self.conversation_service.update_compression_metadata( conversation_id, metadata.to_dict() ) logger.info(f"Compression metadata saved to database for {conversation_id}") return metadata def get_compressed_context( self, conversation: Dict[str, Any] ) -> tuple[Optional[str], List[Dict[str, Any]]]: """ Get compressed summary + recent uncompressed messages. Args: conversation: Full conversation document Returns: (compressed_summary, recent_messages) """ try: # ``or {}`` guards against a NULL ``compression_metadata`` column # (reads back as None), which would crash the ``.get`` calls below. compression_metadata = conversation.get("compression_metadata") or {} if not compression_metadata.get("is_compressed"): logger.debug("No compression metadata found - using full history") queries = conversation.get("queries", []) if queries is None: logger.error("Conversation queries is None - returning empty list") return None, [] return None, queries compression_points = compression_metadata.get("compression_points", []) if not compression_points: logger.debug("No compression points found - using full history") queries = conversation.get("queries", []) if queries is None: logger.error("Conversation queries is None - returning empty list") return None, [] return None, queries # The most recent point that can actually stand in for the # history it covers. An empty saved summary must not slice the # raw history away and replace it with nothing. latest_compression = latest_usable_compression_point(compression_points) if latest_compression is None: logger.warning( "No usable compression point (saved summaries are empty) - " "using full history" ) return None, conversation.get("queries", []) or [] compressed_summary = latest_compression.get("compressed_summary") last_compressed_index = latest_compression.get("query_index") compressed_tokens = latest_compression.get("compressed_token_count", 0) original_tokens = latest_compression.get("original_token_count", 0) # Get only messages after compression point queries = conversation.get("queries", []) total_queries = len(queries) # The visible summary rows appended after a compression are not # history: their content already rides in the system prompt. recent_queries = [ q for q in queries[last_compressed_index + 1 :] if not is_compression_summary_row(q) ] logger.info( f"Using compressed context: summary ({compressed_tokens} tokens, " f"compressed from {original_tokens}) + {len(recent_queries)} recent messages " f"(messages {last_compressed_index + 1}-{total_queries - 1})" ) return compressed_summary, self._bound_recent_queries(recent_queries) except Exception as e: logger.error( f"Error getting compressed context: {str(e)}", exc_info=True ) queries = conversation.get("queries", []) if queries is None: return None, [] return None, queries def _truncate_middle_tokens(self, text: str, max_tokens: int) -> str: """Middle-truncate ``text`` to roughly ``max_tokens`` tokens.""" from docsgpt.utils import num_tokens_from_string current = num_tokens_from_string(text) if current <= max_tokens: return text chars_per_token = len(text) / current if current > 0 else 4 target_chars = int(max_tokens * chars_per_token * 0.95) keep = int(target_chars * 0.4) marker = "\n\n[... trimmed to fit context after compression ...]\n\n" if keep <= 0: # ``text[-0:]`` would return the WHOLE string, not nothing. return marker.strip() return text[:keep] + marker + text[-keep:] def _bound_recent_queries( self, queries: List[Dict[str, Any]] ) -> List[Dict[str, Any]]: """Cap oversized verbatim fields in the post-compression-point tail. Compression summarizes everything up to the compression point, but the tail rides along verbatim — a single giant prompt / response / tool result there can defeat the whole compression (measured in prod: half of compressions produced no reduction in the next call's prompt). Returns copies; the caller's conversation dict is never mutated. """ max_tokens = int( getattr(settings, "COMPRESSION_RECENT_FIELD_MAX_TOKENS", 8000) or 0 ) if max_tokens <= 0: return queries from docsgpt.utils import num_tokens_from_string bounded: List[Dict[str, Any]] = [] trimmed = 0 for query in queries: if not isinstance(query, dict): bounded.append(query) continue out = query for field in ("prompt", "response"): value = query.get(field) if ( isinstance(value, str) and num_tokens_from_string(value) > max_tokens ): if out is query: out = dict(query) out[field] = self._truncate_middle_tokens(value, max_tokens) trimmed += 1 tool_calls = query.get("tool_calls") if isinstance(tool_calls, list): new_calls = None for idx, tc in enumerate(tool_calls): if not isinstance(tc, dict): continue result = tc.get("result") if ( isinstance(result, str) and num_tokens_from_string(result) > max_tokens ): if new_calls is None: new_calls = [ dict(c) if isinstance(c, dict) else c for c in tool_calls ] new_calls[idx]["result"] = self._truncate_middle_tokens( result, max_tokens ) trimmed += 1 if new_calls is not None: if out is query: out = dict(query) out["tool_calls"] = new_calls bounded.append(out) if trimmed: logger.info( f"Bounded {trimmed} oversized field(s) in recent " f"uncompressed queries (cap: {max_tokens} tokens each)" ) return bounded def _extract_summary(self, llm_response: str) -> str: """ Extract clean summary from LLM response. Args: llm_response: Raw LLM response Returns: Cleaned summary text """ try: # Try to extract content within tags summary_match = re.search( r"(.*?)", llm_response, re.DOTALL ) if summary_match: summary = summary_match.group(1).strip() else: # If no summary tags, remove analysis tags and use the rest summary = re.sub( r".*?", "", llm_response, flags=re.DOTALL ).strip() return summary except Exception as e: logger.warning(f"Error extracting summary: {str(e)}, using full response") return llm_response def _log_tool_call_stats(self, queries: List[Dict[str, Any]]) -> None: """Log statistics about tool calls in queries.""" total_tool_calls = 0 total_tool_result_chars = 0 tool_call_breakdown = {} for q in queries: for tc in q.get("tool_calls", []): total_tool_calls += 1 tool_name = tc.get("tool_name", "unknown") action_name = tc.get("action_name", "unknown") key = f"{tool_name}.{action_name}" tool_call_breakdown[key] = tool_call_breakdown.get(key, 0) + 1 # Track total tool result size result = tc.get("result", "") if result: total_tool_result_chars += len(str(result)) if total_tool_calls > 0: tool_breakdown_str = ", ".join( f"{tool}({count})" for tool, count in sorted(tool_call_breakdown.items()) ) tool_result_kb = total_tool_result_chars / 1024 logger.info( f"Tool call breakdown: {tool_breakdown_str} " f"(total result size: {tool_result_kb:.1f} KB, {total_tool_result_chars:,} chars)" )