From 5c064748c2f5259d4bf453b2211f7e20232fd272 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 13 Apr 2026 08:44:14 -0700 Subject: [PATCH] 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 --- litellm/compression/compress.py | 26 ++++++++++++++++++--- litellm/compression/message_stubbing.py | 31 +++++++++++++++++++++++++ 2 files changed, 54 insertions(+), 3 deletions(-) diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index b738744e73c..8a79f7274ca 100644 --- a/litellm/compression/compress.py +++ b/litellm/compression/compress.py @@ -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", "") diff --git a/litellm/compression/message_stubbing.py b/litellm/compression/message_stubbing.py index e7d54097a61..6ca0a8399e0 100644 --- a/litellm/compression/message_stubbing.py +++ b/litellm/compression/message_stubbing.py @@ -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}