mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-05 14:14:39 +00:00
Three review findings on the bounded-chain change. Mid-execution compression rebuilt the conversation from the in-flight messages, which after a turn-start reuse hold only the recent turns: the summary living in the system prompt never reached the compressor, so the new summary replaced the old one, and the persisted point's query_index was relative to that shortened list. The summary the agent is running under now rides into the synthetic conversation as its latest point (query_index -1, so every in-flight query is new), for both the database and the in-memory path, and the database path persists the index of the saved conversation's last row. Saved points with an empty summary, which earlier versions wrote, were treated as reusable: get_compressed_context sliced the raw history away and the effective token count made the conversation look small. Point selection everywhere now takes the latest usable point (non-blank summary, positive token count) and falls back to the raw history when there is none. The chained system-head hash was committed while building the request, so a transport failure followed by the same-primary retry omitted a changed system message. The hash is now staged per request and committed only when the provider records the response.
484 lines
19 KiB
Python
484 lines
19 KiB
Python
"""Core compression service with simplified responsibilities."""
|
|
|
|
import logging
|
|
import re
|
|
from datetime import datetime, timezone
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
from application.api.answer.services.compression.prompt_builder import (
|
|
CompressionPromptBuilder,
|
|
)
|
|
from application.api.answer.services.compression.token_counter import TokenCounter
|
|
from application.api.answer.services.compression.types import (
|
|
CompressionMetadata,
|
|
is_compression_summary_row,
|
|
latest_usable_compression_point,
|
|
)
|
|
from application.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 application.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 application.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 <summary> tags
|
|
summary_match = re.search(
|
|
r"<summary>(.*?)</summary>", 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"<analysis>.*?</analysis>", "", 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)"
|
|
)
|