mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 20:13:04 +00:00
About 85 call sites read a setting as getattr(settings, "NAME", fallback), each carrying its own copy of the default. Every one of those names is a field with a default on the model, so the fallback could never apply to the real settings object; it only masked drift. Two had drifted: - OPENAI_PROMPT_CACHE_KEY defaults to True on the model but the reader fell back to False, and two test stubs relied on that. - SharePoint's MICROSOFT_AUTHORITY fallback to https://login.microsoftonline.com/<tenant> never fired, because the attribute always exists (as None), so MSAL got authority=None. The connector now derives the tenant authority when the setting is unset, as its test always assumed. Four places read EMBEDDINGS_KEY straight from os.environ, skipping the "None"/"" normalisation the model applies; they read the setting now. Test stubs that replaced a module's settings with a SimpleNamespace list every setting the code under test reads.
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 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(
|
|
settings.COMPRESSION_RECENT_FIELD_MAX_TOKENS 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 <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)"
|
|
)
|