fix(compress.py): working prompt compression for claude code

ensures claude code messages can run through proxy easily
This commit is contained in:
Krrish Dholakia 2026-04-14 17:32:07 -07:00
parent d1b9036dbf
commit 91436a666d
5 changed files with 387 additions and 54 deletions

View file

@ -105,6 +105,79 @@ def _extract_last_user_message(messages: List[dict]) -> str:
return ""
def _extract_tool_use_ids(content: Any) -> List[str]:
if not isinstance(content, list):
return []
tool_use_ids: List[str] = []
for part in content:
if not isinstance(part, dict):
continue
if part.get("type") != "tool_use":
continue
tool_use_id = part.get("id")
if isinstance(tool_use_id, str) and tool_use_id:
tool_use_ids.append(tool_use_id)
return tool_use_ids
def _extract_tool_result_ids(content: Any) -> Set[str]:
if not isinstance(content, list):
return set()
tool_result_ids: Set[str] = set()
for part in content:
if not isinstance(part, dict):
continue
if part.get("type") != "tool_result":
continue
tool_use_id = part.get("tool_use_id")
if isinstance(tool_use_id, str) and tool_use_id:
tool_result_ids.add(tool_use_id)
return tool_result_ids
def _extract_anthropic_tool_exchange_spans(
messages: List[dict],
) -> Tuple[List[Set[int]], Optional[str]]:
"""
Return atomic 2-message spans for Anthropic tool exchanges.
Each assistant message containing `tool_use` must be immediately followed by a
user message containing matching `tool_result` blocks for all tool_use ids.
"""
spans: List[Set[int]] = []
i = 0
while i < len(messages):
current = messages[i]
if current.get("role") != "assistant":
i += 1
continue
tool_use_ids = _extract_tool_use_ids(current.get("content"))
if not tool_use_ids:
i += 1
continue
if i + 1 >= len(messages):
return [], "invalid_anthropic_tool_sequence"
next_msg = messages[i + 1]
if next_msg.get("role") != "user":
return [], "invalid_anthropic_tool_sequence"
tool_result_ids = _extract_tool_result_ids(next_msg.get("content"))
if not tool_result_ids:
return [], "invalid_anthropic_tool_sequence"
for tool_use_id in tool_use_ids:
if tool_use_id not in tool_result_ids:
return [], "invalid_anthropic_tool_sequence"
spans.append({i, i + 1})
i += 2
return spans, None
def _get_protected_indices(messages: List[dict]) -> List[int]:
"""
Return indices of messages that must never be compressed:
@ -156,6 +229,94 @@ def _combine_scores(
return [bm25_weight * b + emb_weight * e for b, e in zip(norm_bm25, norm_emb)]
def _select_kept_indices_for_budget(
normalized_messages: List[dict],
original_messages: List[dict],
combined_scores: List[float],
compression_target: int,
model: str,
initial_kept_indices: Set[int],
tool_exchange_spans: List[Set[int]],
) -> Tuple[Set[int], Dict[int, dict]]:
kept_indices = set(initial_kept_indices)
current_tokens = 0
for i in kept_indices:
current_tokens += token_counter(
model=model,
text=cast(str, normalized_messages[i].get("content", "") or ""),
)
# Fill token budget from highest-scoring units.
# A unit is either:
# 1) a single message index, or
# 2) an Anthropic tool-exchange span that must be kept/dropped atomically.
truncated_overrides: Dict[int, dict] = {} # idx -> truncated message dict
span_id_by_index: Dict[int, int] = {}
for span_id, span in enumerate(tool_exchange_spans):
for idx in span:
span_id_by_index[idx] = span_id
# Build single-message candidate units (non-span messages).
candidate_units: List[Tuple[float, Tuple[int, ...], bool]] = []
for idx in range(len(normalized_messages)):
if idx in span_id_by_index or idx in kept_indices:
continue
candidate_units.append((combined_scores[idx], (idx,), True))
# Build span candidate units (atomic keep/drop for tool exchanges).
for span in tool_exchange_spans:
span_indices = tuple(sorted(span))
if any(idx in kept_indices for idx in span_indices):
continue
span_score = max(combined_scores[idx] for idx in span_indices)
candidate_units.append((span_score, span_indices, False))
# Sort by descending relevance score.
candidate_units.sort(key=lambda item: item[0], reverse=True)
for _score, indices, can_truncate in candidate_units:
if any(idx in kept_indices for idx in indices):
continue
msg_tokens = 0
for idx in indices:
msg_tokens += token_counter(
model=model,
text=cast(str, normalized_messages[idx].get("content", "") or ""),
)
remaining = compression_target - current_tokens
if remaining <= 0:
break # budget exhausted
if current_tokens + msg_tokens <= compression_target:
# Fits entirely
kept_indices.update(indices)
current_tokens += msg_tokens
elif can_truncate and len(indices) == 1 and remaining >= 100:
# Too large to fit whole single message, but we have budget — truncate it.
idx = indices[0]
truncated = truncate_message(original_messages[idx], remaining)
truncated_tokens = token_counter(
model=model,
text=truncated.get("content", "") or "",
)
truncated_overrides[idx] = truncated
kept_indices.add(idx)
current_tokens += truncated_tokens
return kept_indices, truncated_overrides
def _get_dropped_tool_span_indices(
kept_indices: Set[int], tool_exchange_spans: List[Set[int]]
) -> Set[int]:
dropped_tool_span_indices: Set[int] = set()
for span in tool_exchange_spans:
if not any(idx in kept_indices for idx in span):
dropped_tool_span_indices.update(span)
return dropped_tool_span_indices
def compress(
messages: List[dict],
model: str,
@ -218,6 +379,7 @@ def compress(
compression_ratio=0.0,
cache={},
tools=[],
compression_skipped_reason="below_trigger",
)
# Extract query for relevance scoring
@ -242,68 +404,52 @@ def compress(
else:
combined_scores = bm25_scores
# Sort message indices by score descending
ranked_indices = sorted(
range(len(normalized_messages)),
key=lambda i: combined_scores[i],
reverse=True,
)
# Protected messages are never compressed
protected_indices = _get_protected_indices(normalized_messages)
kept_indices: Set[int] = set(protected_indices)
# Count tokens for protected messages
current_tokens = 0
for i in kept_indices:
current_tokens += token_counter(
model=model,
text=cast(str, normalized_messages[i].get("content", "") or ""),
tool_exchange_spans: List[Set[int]] = []
if input_type == "anthropic_messages":
tool_exchange_spans, tool_sequence_error = _extract_anthropic_tool_exchange_spans(
original_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 as much of the budget as possible.
# - Otherwise → stub it (pointer only, content goes to cache).
# Multiple messages may be truncated so we preserve partial content from
# several high-scoring messages rather than fully stubbing all but one.
truncated_overrides: Dict[int, dict] = {} # idx -> truncated message dict
for idx in ranked_indices:
if idx in kept_indices:
continue
msg_tokens = token_counter(
model=model,
text=cast(str, normalized_messages[idx].get("content", "") or ""),
)
remaining = compression_target - current_tokens
if remaining <= 0:
break # budget exhausted
if current_tokens + msg_tokens <= compression_target:
# Fits entirely
kept_indices.add(idx)
current_tokens += msg_tokens
elif remaining >= 100:
# Too large to fit whole, but we have budget — truncate it.
truncated = truncate_message(original_messages[idx], remaining)
truncated_tokens = token_counter(
model=model,
text=truncated.get("content", "") or "",
if tool_sequence_error is not None:
return CompressedResult(
messages=original_messages,
original_tokens=original_tokens,
compressed_tokens=original_tokens,
compression_ratio=0.0,
cache={},
tools=[],
compression_skipped_reason=tool_sequence_error,
)
truncated_overrides[idx] = truncated
kept_indices.add(idx)
current_tokens += truncated_tokens
for span in tool_exchange_spans:
# If any message in the span is protected, keep the whole span.
if any(idx in kept_indices for idx in span):
kept_indices.update(span)
kept_indices, truncated_overrides = _select_kept_indices_for_budget(
normalized_messages=normalized_messages,
original_messages=original_messages,
combined_scores=combined_scores,
compression_target=compression_target,
model=model,
initial_kept_indices=kept_indices,
tool_exchange_spans=tool_exchange_spans,
)
# Build compressed messages and cache
compressed_messages: List[dict] = []
cache: Dict[str, str] = {}
used_keys: Set[str] = set()
dropped_tool_span_indices = _get_dropped_tool_span_indices(
kept_indices=kept_indices, tool_exchange_spans=tool_exchange_spans
)
for i, msg in enumerate(original_messages):
if i in dropped_tool_span_indices:
continue
if i in kept_indices:
# Use the truncated version if we made one, otherwise the original
compressed_messages.append(truncated_overrides.get(i, msg))

View file

@ -113,6 +113,7 @@ class CompressionInterceptionLogger(CustomLogger):
)
cache = cast(Dict[str, str], compressed.get("cache", {}))
skip_reason = cast(Optional[str], compressed.get("compression_skipped_reason"))
if cache:
call_id = cast(Optional[str], kwargs.get("litellm_call_id"))
if not call_id:
@ -126,6 +127,13 @@ class CompressionInterceptionLogger(CustomLogger):
compressed.get("compressed_tokens"),
len(cache),
)
elif skip_reason is not None:
verbose_logger.debug(
"CompressionInterception: compression skipped [reason=%s original=%d compressed=%d]",
skip_reason,
compressed.get("original_tokens"),
compressed.get("compressed_tokens"),
)
return kwargs

View file

@ -36,7 +36,7 @@ litellm_settings:
compression_interception_params:
enabled: true
compression_trigger: 1000
# optional:
# embedding_model: "text-embedding-3-small"
# embedding_model_params:
# dimensions: 512
# # optional:
# # embedding_model: "text-embedding-3-small"
# # embedding_model_params:
# # dimensions: 512

View file

@ -2,7 +2,7 @@
Type definitions for litellm.compress().
"""
from typing import Dict, List, Literal, TypedDict
from typing import Dict, List, Literal, NotRequired, TypedDict
CompressionInputType = Literal["anthropic_messages", "openai_chat_completions"]
@ -14,3 +14,4 @@ class CompressedResult(TypedDict):
compression_ratio: float # fraction reduced, e.g. 0.6 means 60% reduction
cache: Dict[str, str] # key -> original content (for retrieval tool responses)
tools: List[dict] # [litellm_content_retrieve tool definition]
compression_skipped_reason: NotRequired[str]

View file

@ -3,6 +3,7 @@ Unit tests for litellm.compress().
"""
import os
import importlib
import pytest
@ -482,3 +483,180 @@ def test_simple_compression(final_user_message, expected_content):
assert "Unrelated cooking recipes " not in result["messages"][1]["content"]
else:
raise ValueError(f"Unexpected expected_content: {expected_content}")
def test_compress_anthropic_drops_irrelevant_tool_exchange_span(monkeypatch):
compress_module = importlib.import_module("litellm.compression.compress")
def fake_bm25_score_messages(query, messages):
assert "final query" in query
assert len(messages) == 5
# Prefer idx=0 and de-prioritize the tool exchange span (idx=1,2)
return [0.95, 0.01, 0.02, 0.8, 1.0]
def fake_token_counter(model, messages=None, text=None):
if messages is not None:
return 1000
if text is None:
return 0
if "final query" in text:
return 50
if "assistant_tail" in text:
return 20
if "other_blob" in text:
return 220
if "tool_payload_relevant" in text:
return 200
if text == "":
return 1
return 10
monkeypatch.setattr(compress_module, "bm25_score_messages", fake_bm25_score_messages)
monkeypatch.setattr(compress_module, "token_counter", fake_token_counter)
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_drop",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_drop",
"content": [{"type": "text", "text": "tool_payload_relevant"}],
}
],
},
{"role": "assistant", "content": "assistant_tail"},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
input_type=ANTHROPIC_INPUT_TYPE,
compression_trigger=100,
compression_target=280,
)
# idx=1,2 should be dropped atomically (no orphan tool blocks left behind)
assert len(result["messages"]) == 3
assert result["messages"][0]["role"] == "user"
assert "other_blob" in result["messages"][0]["content"]
assert result["messages"][1]["content"] == "assistant_tail"
assert result["messages"][2]["content"] == "final query"
assert result["cache"] == {}
def test_compress_anthropic_keeps_relevant_tool_exchange_span(monkeypatch):
compress_module = importlib.import_module("litellm.compression.compress")
def fake_bm25_score_messages(query, messages):
assert "final query" in query
assert len(messages) == 5
# Prefer the tool exchange span over idx=0
return [0.05, 0.01, 0.92, 0.8, 1.0]
def fake_token_counter(model, messages=None, text=None):
if messages is not None:
return 1000
if text is None:
return 0
if "final query" in text:
return 50
if "assistant_tail" in text:
return 20
if "other_blob" in text:
return 220
if "tool_payload_relevant" in text:
return 200
if text == "":
return 1
return 10
monkeypatch.setattr(compress_module, "bm25_score_messages", fake_bm25_score_messages)
monkeypatch.setattr(compress_module, "token_counter", fake_token_counter)
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_keep",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_keep",
"content": [{"type": "text", "text": "tool_payload_relevant"}],
}
],
},
{"role": "assistant", "content": "assistant_tail"},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
input_type=ANTHROPIC_INPUT_TYPE,
compression_trigger=100,
compression_target=280,
)
assert len(result["messages"]) == 5
assert result["messages"][1]["role"] == "assistant"
assert result["messages"][2]["role"] == "user"
# idx=0 should be compressed instead
assert "litellm_content_retrieve" in result["messages"][0]["content"]
assert len(result["cache"]) == 1
def test_compress_anthropic_malformed_tool_sequence_passes_through():
messages = [
{"role": "user", "content": "other_blob " * 300},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_broken",
"name": "litellm_content_retrieve",
"input": {"key": "message_1"},
}
],
},
{"role": "user", "content": [{"type": "text", "text": "missing tool_result"}]},
{"role": "user", "content": "final query"},
]
result = litellm.compress(
messages=messages,
model="claude-sonnet-4-20250514",
input_type=ANTHROPIC_INPUT_TYPE,
compression_trigger=100,
compression_target=280,
)
assert result["messages"] == messages
assert result["cache"] == {}
assert result["tools"] == []
assert result["compression_skipped_reason"] == "invalid_anthropic_tool_sequence"