diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index adcf0729aea..604e216a832 100644 --- a/litellm/compression/compress.py +++ b/litellm/compression/compress.py @@ -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)) diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py index f312ea68ae2..af55a084c9f 100644 --- a/litellm/integrations/compression_interception/handler.py +++ b/litellm/integrations/compression_interception/handler.py @@ -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 diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 08d383dc0e7..10d51c31244 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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 \ No newline at end of file +# # optional: +# # embedding_model: "text-embedding-3-small" +# # embedding_model_params: +# # dimensions: 512 \ No newline at end of file diff --git a/litellm/types/compression.py b/litellm/types/compression.py index 9cb6a2580c9..cc6e53819be 100644 --- a/litellm/types/compression.py +++ b/litellm/types/compression.py @@ -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] diff --git a/tests/test_litellm/test_compression.py b/tests/test_litellm/test_compression.py index 69b593a70df..16fd9bb69ec 100644 --- a/tests/test_litellm/test_compression.py +++ b/tests/test_litellm/test_compression.py @@ -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"