mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
051ad73f5c
commit
5c064748c2
2 changed files with 54 additions and 3 deletions
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue