mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 04:12:36 +00:00
The backend import package is now docsgpt, the name it will carry on PyPI; application was far too generic to install into anyone's site-packages. git mv plus a mechanical rewrite of every import, dotted string and path reference: 734 Python files, the compose files, Dockerfile, workflows, docs, setup scripts, devcontainer, k8s manifests, vscode config, pytest and coverage config, .gitignore. Behaviour is unchanged. Kept for one release: - A top-level application package whose meta-path finder resolves application.x.y to the already-imported docsgpt.x.y object, so old imports and entry points (celery -A application.app.celery, uvicorn application.asgi:asgi_app) keep working with a FutureWarning. - Celery registers every application.* task name as an alias of its docsgpt.* task on start-up, so messages queued by the previous release still run. The redbeat key prefix moves to redbeat:docsgpt:v2: so schedule entries the previous release wrote are left unread instead of firing twice. The backend image builds from the repository root (docker build -f docsgpt/Dockerfile .) so it can ship the alias package; a root .dockerignore allow-lists docsgpt/ and application/ and keeps caches, local data, .env files, the sample index files and the Dockerfile out. Compose and the image workflows point at the new context.
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(
|
|
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 <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)"
|
|
)
|