mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(compress.py): working prompt compression for claude code
ensures claude code messages can run through proxy easily
This commit is contained in:
parent
d1b9036dbf
commit
91436a666d
5 changed files with 387 additions and 54 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue