fix(compress): truncate high-scoring messages instead of fully stubbing them

When a relevant message was too large to fit in the token budget it was
replaced with a stub, leaving the LLM with no real content to work with.
Now the highest-scoring overflow message is truncated (first 70% + last 30%
of words) to fill the remaining budget, so the LLM always receives actual
content rather than just a retrieval pointer.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-04-13 08:44:14 -07:00
parent 051ad73f5c
commit 5c064748c2
2 changed files with 54 additions and 3 deletions

View file

@ -6,7 +6,7 @@ and retrieval tool injection.
from typing import Dict, List, Optional, Set
from litellm.caching.dual_cache import DualCache
from litellm.compression.message_stubbing import extract_key, stub_message
from litellm.compression.message_stubbing import extract_key, stub_message, truncate_message
from litellm.compression.retrieval_tool import build_retrieval_tool
from litellm.compression.scoring.bm25 import bm25_score_messages
from litellm.litellm_core_utils.token_counter import token_counter
@ -166,15 +166,34 @@ def compress(
model=model, text=messages[i].get("content", "") or ""
)
# Fill token budget from highest-scoring messages
# Fill token budget from highest-scoring messages.
# For each candidate (ranked by relevance):
# - If it fits entirely → keep it as-is.
# - If it doesn't fit but there's meaningful remaining budget → truncate it
# to fill that budget (so the LLM always has real content to work with).
# - Otherwise → stub it (pointer only, content goes to cache).
# We only truncate one message (the highest-scoring one that overflows) so
# the budget is consumed and the rest are stubbed cleanly.
truncated_overrides: Dict[int, dict] = {} # idx -> truncated message dict
for idx in ranked_indices:
if idx in kept_indices:
continue
msg_content = messages[idx].get("content", "") or ""
msg_tokens = token_counter(model=model, text=msg_content)
remaining = compression_target - current_tokens
if current_tokens + msg_tokens <= compression_target:
# Fits entirely
kept_indices.add(idx)
current_tokens += msg_tokens
elif remaining >= 100 and idx not in truncated_overrides:
# Too large to fit whole, but we have budget — truncate it.
# Only do this once (the highest-scoring overflow message).
truncated = truncate_message(messages[idx], remaining)
truncated_overrides[idx] = truncated
kept_indices.add(idx)
current_tokens = compression_target # budget consumed
# Build compressed messages and cache
compressed_messages: List[dict] = []
@ -183,7 +202,8 @@ def compress(
for i, msg in enumerate(messages):
if i in kept_indices:
compressed_messages.append(msg)
# Use the truncated version if we made one, otherwise the original
compressed_messages.append(truncated_overrides.get(i, msg))
else:
key = extract_key(msg, fallback_index=i, used_keys=used_keys)
content = msg.get("content", "")

View file

@ -75,3 +75,34 @@ def stub_message(message: dict, key: str) -> dict:
)
return {**message, "content": stub_content}
def truncate_message(message: dict, max_tokens: int) -> dict:
"""
Truncate a message's content to approximately max_tokens by keeping
the first 70% and last 30% of words with a separator in between.
Used when a message is too large to fit entirely in the budget but
too relevant to fully stub out.
"""
content = message.get("content", "")
if isinstance(content, list):
content = " ".join(
p.get("text", "") if isinstance(p, dict) else str(p) for p in content
)
# Rough conversion: 1 token ≈ 0.75 words
target_words = max(1, int(max_tokens * 0.75))
words = content.split()
if len(words) <= target_words:
return {**message, "content": content}
first_count = (target_words * 2) // 3
last_count = target_words - first_count
truncated = (
" ".join(words[:first_count])
+ "\n...[truncated for context window]...\n"
+ " ".join(words[-last_count:])
)
return {**message, "content": truncated}